Emit typed Kotlin class declarations

ober

05e6b0f863c58dbc7e363f9fda30985a3014930b

diff --git a/lib/jerboa/typed/checker.ss b/lib/jerboa/typed/checker.ss
index fc6c577..3b8354c 100644
--- a/lib/jerboa/typed/checker.ss
+++ b/lib/jerboa/typed/checker.ss
@@ -43,6 +43,7 @@
   (def *variant-env* (make-parameter '()))
   (def *current-effects* (make-parameter '()))
   (def *global-value-env* (make-parameter '()))
+  (def *elaboration-scope* (make-parameter #f))
 
   (def builtin-type-names
     '(Unit Bool Char Int Nat Fixnum Float String Bytes Symbol Keyword))
@@ -269,6 +270,7 @@
       [(typed-resource? decl) (typed-resource-name decl)]
       [(typed-variant? decl) (typed-variant-name decl)]
       [(typed-object-decl? decl) (typed-object-decl-name decl)]
+      [(typed-class-decl? decl) (typed-class-decl-name decl)]
       [else #f]))
 
   (def (declared-type-names declarations)
@@ -306,6 +308,7 @@
       [(typed-def? decl) (list (typed-def-name decl))]
       [(typed-property? decl) (list (typed-property-name decl))]
       [(typed-object-decl? decl) (list (typed-object-decl-name decl))]
+      [(typed-class-decl? decl) (list (typed-class-decl-name decl))]
       [(typed-extern? decl) (list (typed-extern-name decl))]
       [(typed-record? decl) (record-value-names decl)]
       [(typed-variant? decl) (variant-value-names decl)]
@@ -875,11 +878,15 @@
                    out)))]
         [else (loop (cdr rest) out)])))
 
+  (def (elaboration-key name)
+    (let ([scope (*elaboration-scope*)])
+      (if scope (list scope name) name)))
+
   (def (record-elaboration-by-name! name body-ir)
     (let ([acc (*elaboration-acc*)])
       (when (and acc body-ir)
         (set-box! acc
-          (cons (cons name body-ir) (unbox acc))))))
+          (cons (cons (elaboration-key name) body-ir) (unbox acc))))))
 
   (def (check-property prop type-names)
     (let-values ([(init-ir init-errors)
@@ -919,7 +926,8 @@
         ([*call-env* (append local-calls (*call-env*))]
          [*variant-env* (append local-variants (*variant-env*))]
          [*global-value-env*
-          (append local-global-values (*global-value-env*))])
+          (append local-global-values (*global-value-env*))]
+         [*elaboration-scope* (typed-object-decl-name object)])
         (append
           (duplicate-errors 'duplicate-type
             "duplicate object type declaration"
@@ -932,6 +940,69 @@
               (check-declaration decl visible-type-names))
             declarations)))))
 
+  (def (class-param-env params)
+    (map (lambda (param)
+           (cons (typed-param-name param)
+                 (typed-param-type param)))
+         params))
+
+  (def (check-class-super super env type-names)
+    (if (not super)
+      '()
+      (let-values ([(arg-irs arg-errors)
+                    (infer-args
+                      (typed-class-super-args super)
+                      env
+                      type-names)])
+        (let ([errors
+               (append
+                 (check-type (typed-class-super-name super) type-names)
+                 arg-errors)])
+          (when (null? errors)
+            (record-elaboration-by-name! '*super-args* arg-irs))
+          errors))))
+
+  (def (check-class-decl class type-names)
+    (let* ([params (typed-class-decl-params class)]
+           [declarations (typed-class-decl-declarations class)]
+           [local-type-names (declared-type-names declarations)]
+           [local-value-names (declared-value-names declarations)]
+           [local-calls (call-env declarations)]
+           [local-variants (declared-variant-env declarations)]
+           [local-global-values (global-property-env declarations)]
+           [constructor-env (class-param-env params)]
+           [visible-type-names (append type-names local-type-names)]
+           [member-env (append constructor-env
+                               local-global-values
+                               (*global-value-env*))])
+      (parameterize
+        ([*call-env* (append local-calls (*call-env*))]
+         [*variant-env* (append local-variants (*variant-env*))]
+         [*global-value-env* member-env]
+         [*elaboration-scope* (typed-class-decl-name class)])
+        (append
+          (duplicate-errors 'duplicate-param
+            "duplicate class constructor parameter"
+            (map typed-param-name params))
+          (append-map
+            (lambda (param)
+              (check-type (typed-param-type param) visible-type-names))
+            params)
+          (check-class-super
+            (typed-class-decl-super class)
+            constructor-env
+            visible-type-names)
+          (duplicate-errors 'duplicate-type
+            "duplicate class type declaration"
+            local-type-names)
+          (duplicate-errors 'duplicate-value
+            "duplicate class value declaration"
+            local-value-names)
+          (append-map
+            (lambda (decl)
+              (check-declaration decl visible-type-names))
+            declarations)))))
+
   (def (infer-body exprs env type-names . src*)
     (let ([src (if (pair? src*) (car src*) #f)])
       (cond
@@ -3548,6 +3619,7 @@
       [(typed-variant? decl) (check-variant decl type-names)]
       [(typed-property? decl) (check-property decl type-names)]
       [(typed-object-decl? decl) (check-object-decl decl type-names)]
+      [(typed-class-decl? decl) (check-class-decl decl type-names)]
       [(typed-extern? decl) (check-extern decl type-names)]
       [(typed-def? decl) (check-def decl type-names)]
       [else '()]))
@@ -3584,6 +3656,8 @@
                       (memq (typed-property-name decl) exports)]
                      [(typed-object-decl? decl)
                       (memq (typed-object-decl-name decl) exports)]
+                     [(typed-class-decl? decl)
+                      (memq (typed-class-decl-name decl) exports)]
                      [(typed-record? decl)
                       (or (memq (typed-record-name decl) exports)
                           (let any-loop ([rs values])
@@ -3759,7 +3833,7 @@
                     [(typed-def? (car rest))
                      (let* ([def (car rest)]
                             [name (typed-def-name def)]
-                            [entry (assq name entries)])
+                            [entry (assoc name entries)])
                        (loop (cdr rest)
                              (if entry
                                (cons (make-elaborated-def
@@ -3769,7 +3843,7 @@
                     [(typed-property? (car rest))
                      (let* ([prop (car rest)]
                             [name (typed-property-name prop)]
-                            [entry (assq name entries)])
+                            [entry (assoc name entries)])
                        (loop (cdr rest)
                              (if entry
                                (cons (make-elaborated-def
@@ -3782,7 +3856,7 @@
                   (let loop ([rest entries] [out '()])
                     (cond
                       [(null? rest) (reverse out)]
-                      [(memq (caar rest) returned-names)
+                      [(member (caar rest) returned-names)
                        (loop (cdr rest) out)]
                       [else
                        (loop (cdr rest)
diff --git a/lib/jerboa/typed/kotlin/lower.ss b/lib/jerboa/typed/kotlin/lower.ss
index af2f8f6..6e6fba1 100644
--- a/lib/jerboa/typed/kotlin/lower.ss
+++ b/lib/jerboa/typed/kotlin/lower.ss
@@ -22,6 +22,7 @@
   (def *kotlin-variant-env* (make-parameter '()))
   (def *kotlin-variant-case-env* (make-parameter '()))
   (def *kotlin-ir-env* (make-parameter '()))
+  (def *kotlin-decl-scope* (make-parameter #f))
 
   (def (append-map f xs)
     (let loop ([rest xs] [out '()])
@@ -72,6 +73,14 @@
     (let ([entry (assq key alist)])
       (and entry (cdr entry))))
 
+  (def (kotlin-ir-key name)
+    (let ([scope (*kotlin-decl-scope*)])
+      (if scope (list scope name) name)))
+
+  (def (lookup-kotlin-ir name)
+    (let ([entry (assoc (kotlin-ir-key name) (*kotlin-ir-env*))])
+      (and entry (cdr entry))))
+
   (def (info-ref info key default)
     (let ([entry (assq key info)])
       (if entry (cdr entry) default)))
@@ -979,8 +988,8 @@
       #f))
 
   (def (lower-def def)
-    (let ([ir-entry (assq (typed-def-name def) (*kotlin-ir-env*))])
-      (unless ir-entry
+    (let ([ir (lookup-kotlin-ir (typed-def-name def))])
+      (unless ir
         (error 'lower-def "missing elaborated IR for typed definition" (typed-def-name def)))
       (make-kt-function
         #f
@@ -989,13 +998,13 @@
         (map lower-param (typed-def-params def))
         (typed-type->kotlin-type (typed-def-return-type def))
         (if (eq? (typed-def-return-type def) 'Unit)
-          (lower-unit-statements (cdr ir-entry))
-          (list (make-kt-return (lower-expr (cdr ir-entry)))))
+          (lower-unit-statements ir)
+          (list (make-kt-return (lower-expr ir))))
         '())))
 
   (def (lower-property prop)
-    (let ([ir-entry (assq (typed-property-name prop) (*kotlin-ir-env*))])
-      (unless ir-entry
+    (let ([ir (lookup-kotlin-ir (typed-property-name prop))])
+      (unless ir
         (error 'lower-property
           "missing elaborated IR for typed property"
           (typed-property-name prop)))
@@ -1004,7 +1013,7 @@
         (typed-property-mutable? prop)
         (typed-property-name prop)
         (typed-type->kotlin-type (typed-property-type prop))
-        (lower-expr (cdr ir-entry))
+        (lower-expr ir)
         '())))
 
   (def (lower-extern decl)
@@ -1022,20 +1031,49 @@
              '()))))
 
   (def (lower-object-declaration decl)
-    (make-kt-class
-      'object
-      #f
-      (typed-object-decl-name decl)
-      '()
-      '()
-      (let loop ([decls (typed-object-decl-declarations decl)] [out '()])
-        (cond
-          [(null? decls) (reverse out)]
-          [else
-           (let ([lowered (lower-declaration (car decls))])
-             (loop (cdr decls)
-                   (if lowered (cons lowered out) out)))]))
-      '()))
+    (parameterize ([*kotlin-decl-scope* (typed-object-decl-name decl)])
+      (make-kt-class
+        'object
+        #f
+        (typed-object-decl-name decl)
+        '()
+        '()
+        (lower-nested-declarations
+          (typed-object-decl-declarations decl))
+        '())))
+
+  (def (lower-class-super super)
+    (and super
+         (let ([arg-irs (lookup-kotlin-ir '*super-args*)])
+           (unless arg-irs
+             (error 'lower-class-super
+               "missing elaborated IR for class superclass args"
+               (typed-class-super-name super)))
+           (make-kt-new
+             (typed-type->kotlin-type (typed-class-super-name super))
+             (map lower-expr arg-irs)))))
+
+  (def (lower-nested-declarations decls)
+    (let loop ([decls decls] [out '()])
+      (cond
+        [(null? decls) (reverse out)]
+        [else
+         (let ([lowered (lower-declaration (car decls))])
+           (loop (cdr decls)
+                 (if lowered (cons lowered out) out)))])))
+
+  (def (lower-class-declaration decl)
+    (parameterize ([*kotlin-decl-scope* (typed-class-decl-name decl)])
+      (make-kt-class
+        'class
+        #f
+        (typed-class-decl-name decl)
+        (map lower-param (typed-class-decl-params decl))
+        (let ([super (lower-class-super (typed-class-decl-super decl))])
+          (if super (list super) '()))
+        (lower-nested-declarations
+          (typed-class-decl-declarations decl))
+        '())))
 
   (def (lower-declaration decl)
     (cond
@@ -1043,6 +1081,7 @@
       [(typed-variant? decl) (lower-variant decl)]
       [(typed-property? decl) (lower-property decl)]
       [(typed-object-decl? decl) (lower-object-declaration decl)]
+      [(typed-class-decl? decl) (lower-class-declaration decl)]
       [(typed-extern? decl) (lower-extern decl)]
       [(typed-def? decl) (lower-def decl)]
       [else #f]))
diff --git a/lib/jerboa/typed/kotlin/print.ss b/lib/jerboa/typed/kotlin/print.ss
index ad5a7ca..edaf1e9 100644
--- a/lib/jerboa/typed/kotlin/print.ss
+++ b/lib/jerboa/typed/kotlin/print.ss
@@ -489,7 +489,7 @@
       (if (null? supers)
         ""
         (string-append " : "
-                       (join-strings (map kotlin-type->string supers) ", ")))))
+                       (join-strings (map kotlin-super->string supers) ", ")))))
 
   (def (kotlin-modifier->string modifier)
     (cond
diff --git a/lib/jerboa/typed/parser.ss b/lib/jerboa/typed/parser.ss
index 710970f..0eb6080 100644
--- a/lib/jerboa/typed/parser.ss
+++ b/lib/jerboa/typed/parser.ss
@@ -42,6 +42,15 @@
     typed-object-decl-name typed-object-decl-declarations
     typed-object-decl-source
 
+    typed-class-decl? make-typed-class-decl
+    typed-class-decl-name typed-class-decl-params
+    typed-class-decl-super typed-class-decl-declarations
+    typed-class-decl-source
+
+    typed-class-super? make-typed-class-super
+    typed-class-super-name typed-class-super-args
+    typed-class-super-source
+
     typed-extern? make-typed-extern
     typed-extern-name typed-extern-params typed-extern-return-type
     typed-extern-kotlin-path typed-extern-source
@@ -71,6 +80,8 @@
   (defstruct typed-resource (name close source))
   (defstruct typed-property (mutable? name type init source))
   (defstruct typed-object-decl (name declarations source))
+  (defstruct typed-class-decl (name params super declarations source))
+  (defstruct typed-class-super (name args source))
   (defstruct typed-extern (name params return-type kotlin-path source))
   (defstruct typed-variant (name cases source))
   (defstruct typed-variant-case (name fields source))
@@ -539,6 +550,58 @@
         (map parse-typed-declaration (cddr raw-form))
         source)))
 
+  (def (extends-form? form)
+    (let ([form (strip-source-annotations form)])
+      (and (pair? form)
+           (eq? (car form) 'extends))))
+
+  (def (parse-class-super form)
+    (let* ([source (datum-source form)]
+           [raw-form (datum-value form)]
+           [form (strip-source-annotations form)])
+      (expect-proper-list 'parse-typed-class form)
+      (unless (and (>= (length form) 2)
+                   (eq? (car form) 'extends)
+                   (symbol? (cadr form)))
+        (error 'parse-typed-class
+          "expected (extends Super arg ...)"
+          form))
+      (make-typed-class-super
+        (cadr form)
+        (cddr raw-form)
+        source)))
+
+  (def (parse-class-decl form)
+    (let* ([source (datum-source form)]
+           [raw-form (datum-value form)]
+           [form (strip-source-annotations form)])
+      (expect-proper-list 'parse-typed-class form)
+      (unless (>= (length form) 3)
+        (error 'parse-typed-class
+          "expected (class Name ((param : Type) ...) declaration ...)"
+          form))
+      (let* ([name (expect-symbol 'parse-typed-class (cadr form) form)]
+             [raw-params (caddr raw-form)]
+             [params-form (caddr form)])
+        (unless (and (list? params-form)
+                     (proper-list? params-form))
+          (error 'parse-typed-class
+            "class constructor params must be a list"
+            params-form))
+        (let* ([tail-raw (cdddr raw-form)]
+               [tail (cdddr form)]
+               [has-super? (and (pair? tail)
+                                (extends-form? (car tail-raw)))]
+               [super (and has-super?
+                           (parse-class-super (car tail-raw)))]
+               [decls (if has-super? (cdr tail-raw) tail-raw)])
+          (make-typed-class-decl
+            name
+            (map parse-param raw-params)
+            super
+            (map parse-typed-declaration decls)
+            source)))))
+
   (def (parse-typed-declaration form)
     (let ([stripped-form (strip-source-annotations form)])
       (expect-proper-list 'parse-typed-declaration stripped-form)
@@ -550,6 +613,7 @@
         [(resource) (parse-resource form)]
         [(val var) (parse-property form)]
         [(object) (parse-object-decl form)]
+        [(class) (parse-class-decl form)]
         [(variant) (parse-variant form)]
         [(extern) (parse-extern form)]
         [(def) (parse-def form)]
diff --git a/tests/test-typed-kotlin.ss b/tests/test-typed-kotlin.ss
index a0f4af8..055576d 100644
--- a/tests/test-typed-kotlin.ss
+++ b/tests/test-typed-kotlin.ss
@@ -1467,5 +1467,35 @@
   object-declaration-kotlin
   "fun loadErrorMessage(): String {\n        return throwableMessageOrEmpty(loadError)")
 
+(define class-declaration-form
+  '(typed-library (sample typed classdecl)
+     (export ReviewView)
+     (type Context)
+     (type View)
+     (type Float32)
+     (extern (contextDensity (context : Context)) : Float32
+       (kotlin-call contextDensity))
+     (extern (densityReady (density : Float32)) : Bool
+       (kotlin-call densityReady))
+     (class ReviewView ((context : Context))
+       (extends View context)
+       (val density : Float32
+         (contextDensity context))
+       (def (ready) : Bool
+         (densityReady density)))))
+
+(define class-declaration-kotlin
+  (typed-library-form->kotlin-string class-declaration-form))
+
+(test-contains "typed class declaration lowers constructor and superclass"
+  class-declaration-kotlin
+  "class ReviewView(context: Context) : View(context) {")
+(test-contains "typed class declaration lowers member property"
+  class-declaration-kotlin
+  "val density: Float = contextDensity(context)")
+(test-contains "typed class declaration lowers member function"
+  class-declaration-kotlin
+  "fun ready(): Boolean {\n        return densityReady(density)")
+
 (printf "typed-kotlin tests: ~a passed, ~a failed~%" pass fail)
 (when (> fail 0) (exit 1))