Add typed Kotlin lambda lowering

ober

2112b909dff4150a763859878d45cc81dcc05037

diff --git a/lib/jerboa/typed/checker.ss b/lib/jerboa/typed/checker.ss
index 0b7ac60..c8642d5 100644
--- a/lib/jerboa/typed/checker.ss
+++ b/lib/jerboa/typed/checker.ss
@@ -753,9 +753,9 @@
            [arity-entry (assq head compound-type-arities)])
       (cond
         [(eq? head '->)
-         (if (< (length args) 2)
+         (if (< (length args) 1)
            (list (make-check-error 'bad-type-arity
-                   "function type needs at least one argument and a result"
+                   "function type needs zero or more arguments and a result"
                    type))
            (append-map (lambda (arg) (check-type arg type-names)) args))]
         [(not arity-entry)
@@ -989,6 +989,79 @@
             (append cond-errors then-errors else-errors
                     condition-errors branch-errors))))))
 
+  (def (lambda-param-value param)
+    (strip-source-annotations param))
+
+  (def (lambda-param-name param)
+    (let ([v (lambda-param-value param)])
+      (and (pair? v) (car v))))
+
+  (def (parse-lambda-param param type-names)
+    (let ([source (expr-source param)]
+          [v (lambda-param-value param)])
+      (cond
+        [(and (list? v)
+              (= (length v) 3)
+              (symbol? (car v))
+              (eq? (cadr v) ':))
+         (let* ([type (parse-typed-type (caddr v))]
+                [type-errors (check-type type type-names)])
+           (values (and (null? type-errors)
+                        (make-typed-param (car v) type source))
+                   type-errors))]
+        [else
+         (values #f
+           (list (make-check-error 'bad-lambda-param
+                   "lambda parameter must be shaped (name : Type)"
+                   v
+                   source)))])))
+
+  (def (infer-lambda args env type-names expr)
+    (cond
+      [(< (length args) 2)
+       (values #f
+         (list (error-at expr 'bad-lambda
+                 "lambda expects a parameter list and body"
+                 expr)))]
+      [(not (expr-list? (car args)))
+       (values #f
+         (list (error-at expr 'bad-lambda
+                 "lambda parameters must be a list"
+                 (car args))))]
+      [else
+       (let* ([param-exprs (expr->list (car args))]
+              [duplicate-name-errors
+               (duplicate-errors 'duplicate-param
+                 "duplicate lambda parameter"
+                 (map lambda-param-name param-exprs))])
+         (let loop ([rest param-exprs]
+                    [params '()]
+                    [errors duplicate-name-errors])
+           (if (null? rest)
+             (let ([ok-params (reverse params)])
+               (if (not (null? errors))
+                 (values #f (reverse errors))
+                 (let* ([local-env
+                         (append (param-env ok-params) env)]
+                        [param-types (map typed-param-type ok-params)])
+                   (let-values ([(body-ir body-errors)
+                                 (infer-body (cdr args) local-env type-names
+                                             (expr-source expr))])
+                     (let ([body-type (ir-type body-ir)])
+                       (values
+                         (and body-type
+                              (make-typed-ir-lambda
+                                (cons '-> (append param-types (list body-type)))
+                                (expr-source expr)
+                                ok-params
+                                body-ir))
+                         body-errors))))))
+             (let-values ([(param param-errors)
+                           (parse-lambda-param (car rest) type-names)])
+               (loop (cdr rest)
+                     (if param (cons param params) params)
+                     (append (reverse param-errors) errors))))))]))
+
   (def (lookup-variant name)
     (lookup-name name (*variant-env*)))
 
@@ -2594,6 +2667,10 @@
                   (values after-cond
                           (append cond-errors then-errors else-errors)))
                 (values moved '()))]
+             [(lambda)
+              (if (< (length args) 2)
+                (values moved '())
+                (check-ownership-body (cdr args) env moved))]
              [(match)
               (if (< (length args) 1)
                 (values moved '())
@@ -2698,6 +2775,8 @@
                 (infer-let (car args) (cdr args) env type-names expr))]
              [(if)
               (infer-if args env type-names expr)]
+             [(lambda)
+              (infer-lambda args env type-names expr)]
              [(match)
               (infer-match args env type-names expr)]
              [(+ - * / mod)
diff --git a/lib/jerboa/typed/core.ss b/lib/jerboa/typed/core.ss
index c84c22f..70ae1b7 100644
--- a/lib/jerboa/typed/core.ss
+++ b/lib/jerboa/typed/core.ss
@@ -39,6 +39,11 @@
     typed-ir-if-type typed-ir-if-source
     typed-ir-if-test typed-ir-if-then typed-ir-if-else
 
+    typed-ir-lambda?
+    make-typed-ir-lambda
+    typed-ir-lambda-type typed-ir-lambda-source
+    typed-ir-lambda-params typed-ir-lambda-body
+
     typed-ir-for-fold?
     make-typed-ir-for-fold
     typed-ir-for-fold-type typed-ir-for-fold-source
@@ -88,6 +93,7 @@
   (defstruct typed-ir-let (type source bindings body))
   (defstruct typed-ir-binding (name expr))
   (defstruct typed-ir-if (type source test then else))
+  (defstruct typed-ir-lambda (type source params body))
   ;; single-accumulator fold over an in-range index: the loop variable runs
   ;; [range-start, range-end); body (of acc's type) becomes acc's next value.
   (defstruct typed-ir-for-fold
@@ -205,6 +211,7 @@
         (typed-ir-begin? x)
         (typed-ir-let? x)
         (typed-ir-if? x)
+        (typed-ir-lambda? x)
         (typed-ir-for-fold? x)
         (typed-ir-bytes-build? x)
         (typed-ir-match? x)
@@ -217,6 +224,7 @@
       [(typed-ir-begin? node) (typed-ir-begin-type node)]
       [(typed-ir-let? node) (typed-ir-let-type node)]
       [(typed-ir-if? node) (typed-ir-if-type node)]
+      [(typed-ir-lambda? node) (typed-ir-lambda-type node)]
       [(typed-ir-for-fold? node) (typed-ir-for-fold-type node)]
       [(typed-ir-bytes-build? node) (typed-ir-bytes-build-type node)]
       [(typed-ir-match? node) (typed-ir-match-type node)]
@@ -230,6 +238,7 @@
       [(typed-ir-begin? node) (typed-ir-begin-source node)]
       [(typed-ir-let? node) (typed-ir-let-source node)]
       [(typed-ir-if? node) (typed-ir-if-source node)]
+      [(typed-ir-lambda? node) (typed-ir-lambda-source node)]
       [(typed-ir-for-fold? node) (typed-ir-for-fold-source node)]
       [(typed-ir-bytes-build? node) (typed-ir-bytes-build-source node)]
       [(typed-ir-match? node) (typed-ir-match-source node)]
diff --git a/lib/jerboa/typed/kotlin/lower.ss b/lib/jerboa/typed/kotlin/lower.ss
index abc361b..eff9c82 100644
--- a/lib/jerboa/typed/kotlin/lower.ss
+++ b/lib/jerboa/typed/kotlin/lower.ss
@@ -123,6 +123,21 @@
        (make-kt-type 'JbResult #f
          (list (typed-type->kotlin-type (cadr type))
                (typed-type->kotlin-type (caddr type))))]
+      [(and (pair? type) (eq? (car type) '->))
+       (let* ([args (cdr type)]
+              [arity (- (length args) 1)]
+              [return-type (let loop ([xs args])
+                             (if (null? (cdr xs)) (car xs) (loop (cdr xs))))])
+         (make-kt-type
+           (string->symbol (string-append "Function" (number->string arity)))
+           #f
+           (append
+             (map typed-type->kotlin-type
+                  (let loop ([xs args])
+                    (if (or (null? xs) (null? (cdr xs)))
+                      '()
+                      (cons (car xs) (loop (cdr xs))))))
+             (list (typed-type->kotlin-type return-type)))))]
       [(and (pair? type) (memq (car type) '(Owned Borrow MutBorrow)))
        (typed-type->kotlin-type (cadr type))]
       [else (error 'typed-type->kotlin-type "unsupported typed Jerboa type" type)]))
@@ -735,6 +750,11 @@
             'toByte
             '())))))
 
+  (def (lower-lambda ir)
+    (make-kt-lambda
+      (map typed-param-name (typed-ir-lambda-params ir))
+      (lower-expr (typed-ir-lambda-body ir))))
+
   (def (lower-expr ir)
     (cond
       [(typed-ir-lit? ir)
@@ -750,6 +770,8 @@
          (lower-expr (typed-ir-if-test ir))
          (lower-expr (typed-ir-if-then ir))
          (lower-expr (typed-ir-if-else ir)))]
+      [(typed-ir-lambda? ir)
+       (lower-lambda ir)]
       [(typed-ir-for-fold? ir)
        (lower-for-fold ir)]
       [(typed-ir-bytes-build? ir)
diff --git a/lib/jerboa/typed/kotlin/print.ss b/lib/jerboa/typed/kotlin/print.ss
index ec983ab..e59e4d9 100644
--- a/lib/jerboa/typed/kotlin/print.ss
+++ b/lib/jerboa/typed/kotlin/print.ss
@@ -269,12 +269,17 @@
          (kotlin-expr->string (kt-index-get-index expr))
          "]")]
       [(kt-lambda? expr)
-       (string-append
-         "{ "
-         (join-strings (map kotlin-symbol-name (kt-lambda-params expr)) ", ")
-         " -> "
-         (kotlin-expr->string (kt-lambda-body expr))
-         " }")]
+       (if (null? (kt-lambda-params expr))
+         (string-append
+           "{ "
+           (kotlin-expr->string (kt-lambda-body expr))
+           " }")
+         (string-append
+           "{ "
+           (join-strings (map kotlin-symbol-name (kt-lambda-params expr)) ", ")
+           " -> "
+           (kotlin-expr->string (kt-lambda-body expr))
+           " }"))]
       [(kt-binary? expr)
        (parenthesize
          (string-append
diff --git a/tests/test-typed-checker.ss b/tests/test-typed-checker.ss
index 06d0503..fa12d96 100644
--- a/tests/test-typed-checker.ss
+++ b/tests/test-typed-checker.ss
@@ -469,6 +469,35 @@
          (g x))))
   '(return-type-mismatch))
 
+(test "typed lambda return value"
+  (error-kinds
+    '(typed-library (body lambda-ok)
+       (export make-inc)
+       (type Int32)
+       (def (make-inc) : (-> Int32 Int32)
+         (lambda ((x : Int32))
+           (+ x (int32 1))))))
+  '())
+
+(test "typed zero-argument lambda return value"
+  (error-kinds
+    '(typed-library (body lambda-zero-ok)
+       (export noop)
+       (def (noop) : (-> Unit)
+         (lambda ()
+           (begin)))))
+  '())
+
+(test "typed lambda return mismatch"
+  (error-kinds
+    '(typed-library (body lambda-return-mismatch)
+       (export bad)
+       (type Int32)
+       (def (bad) : (-> String String)
+         (lambda ((x : Int32))
+           x))))
+  '(return-type-mismatch))
+
 (test "extern kotlin call can be used from typed code"
   (error-kinds
     '(typed-library (body extern-ok)
diff --git a/tests/test-typed-kotlin.ss b/tests/test-typed-kotlin.ss
index f221a65..a087bf7 100644
--- a/tests/test-typed-kotlin.ss
+++ b/tests/test-typed-kotlin.ss
@@ -63,6 +63,13 @@
       (make-kt-member-call (make-kt-name '(ch)) 'isLetterOrDigit '())))
   "{ ch -> ch.isLetterOrDigit() }")
 
+(test "zero-argument lambda expression printer"
+  (kotlin-expr->string
+    (make-kt-lambda
+      '()
+      (make-kt-lit 'Unit '())))
+  "{ Unit }")
+
 (define ast-file
   (make-kt-file
     '(sample generated)
@@ -327,6 +334,35 @@
   member-extern-kotlin
   "view.selectedGroupId = id")
 
+(define lambda-form
+  '(typed-library (sample typed lambdas)
+     (export plusOne noop)
+     (type Int32)
+     (def (plusOne) : (-> Int32 Int32)
+       (lambda ((x : Int32))
+         (+ x (int32 1))))
+     (def (noop) : (-> Unit)
+       (lambda ()
+         (begin)))))
+
+(define lambda-kotlin (typed-library-form->kotlin-string lambda-form))
+
+(test-contains "typed lambda return lowers to Kotlin Function1"
+  lambda-kotlin
+  "fun plusOne(): Function1<Int, Int>")
+
+(test-contains "typed lambda body lowers to Kotlin lambda"
+  lambda-kotlin
+  "return { x -> (x + 1) }")
+
+(test-contains "typed zero-arg lambda return lowers to Kotlin Function0"
+  lambda-kotlin
+  "fun noop(): Function0<Unit>")
+
+(test-contains "typed zero-arg lambda body omits arrow"
+  lambda-kotlin
+  "return { Unit }")
+
 (define geometry-form
   '(typed-library (sample typed geometry)
      (export make-SsdCell SsdCell? SsdCell-x SsdCell-y SsdCell-w SsdCell-h