Check Typed Jerboa variant matches

ober

997963fc3abf1a983f08afce6b575f3312bf418b

diff --git a/docs/typed-jerboa.md b/docs/typed-jerboa.md
index 30641d0..3be4b68 100644
--- a/docs/typed-jerboa.md
+++ b/docs/typed-jerboa.md
@@ -149,9 +149,10 @@ Current landing:
 - This is still a front-end milestone. Function-body checking currently covers
   literals, variables, `begin`, simple `let`, `if`, arithmetic primitives,
   numeric comparisons, boolean primitives, calls to typed functions defined in
-  the same module, and generated record/variant operations. Imported calls and
-  richer forms are reported as unsupported. It does not yet resolve imports,
-  lower to typed core IR, or emit Rust/LLVM.
+  the same module, generated record/variant operations, and exhaustive
+  `match` over same-module variants. Imported calls and richer forms are
+  reported as unsupported. It does not yet resolve imports, lower to typed core
+  IR, or emit Rust/LLVM.
 
 ## Surface Syntax
 
@@ -285,8 +286,10 @@ Non-exhaustive matches should be compile errors unless there is an explicit
 fallback branch.
 
 The current checker recognizes variant constructors and the variant predicate
-for variants declared in the same typed module. Exhaustive `match` checking is
-still a later milestone.
+for variants declared in the same typed module. It also checks `match` forms
+over same-module variants for case arity, field binding types, duplicate cases,
+consistent branch result types, and exhaustiveness. A `_` or `else` fallback is
+the explicit opt-out from exhaustiveness.
 
 ## Functions
 
@@ -755,7 +758,9 @@ Minimum excluded features:
 - Check records and variants. Same-module record constructors, predicates,
   accessors, mutable setters, variant constructors, and variant predicates now
   check arity, argument types, and return type flow.
-- Check match exhaustiveness.
+- Check match exhaustiveness. Same-module variant `match` now checks case
+  coverage, duplicate cases, pattern arity, field bindings, and branch result
+  types.
 - Produce useful errors.
 
 ### Milestone 3: Rust Backend
diff --git a/lib/jerboa/typed/checker.ss b/lib/jerboa/typed/checker.ss
index 164cad8..d782819 100644
--- a/lib/jerboa/typed/checker.ss
+++ b/lib/jerboa/typed/checker.ss
@@ -20,6 +20,7 @@
   (defstruct typed-call-sig (param-types return-type))
 
   (def *call-env* (make-parameter '()))
+  (def *variant-env* (make-parameter '()))
 
   (def builtin-type-names
     '(Unit Bool Char Int Nat Fixnum Float String Bytes Symbol Keyword))
@@ -206,6 +207,16 @@
                (append (reverse (declaration-call-signatures (car rest)))
                        out))])))
 
+  (def (declared-variant-env declarations)
+    (let loop ([rest declarations] [out '()])
+      (cond
+        [(null? rest) (reverse out)]
+        [(typed-variant? (car rest))
+         (loop (cdr rest)
+               (cons (cons (typed-variant-name (car rest)) (car rest)) out))]
+        [else
+         (loop (cdr rest) out)])))
+
   (def (check-compound-type type type-names)
     (let* ([head (car type)]
            [args (cdr type)]
@@ -297,6 +308,32 @@
   (def (extend-env name type env)
     (cons (cons name type) env))
 
+  (def (symbol-list? xs)
+    (and (list? xs)
+         (let loop ([rest xs])
+           (cond
+             [(null? rest) #t]
+             [(symbol? (car rest)) (loop (cdr rest))]
+             [else #f]))))
+
+  (def (pattern-binding-names names)
+    (let loop ([rest names] [out '()])
+      (cond
+        [(null? rest) (reverse out)]
+        [(eq? (car rest) '_) (loop (cdr rest) out)]
+        [else (loop (cdr rest) (cons (car rest) out))])))
+
+  (def (extend-env* names types env)
+    (let loop ([rest-names names] [rest-types types] [out env])
+      (cond
+        [(or (null? rest-names) (null? rest-types)) out]
+        [(eq? (car rest-names) '_)
+         (loop (cdr rest-names) (cdr rest-types) out)]
+        [else
+         (loop (cdr rest-names)
+               (cdr rest-types)
+               (extend-env (car rest-names) (car rest-types) out))])))
+
   (def (binding-name binding)
     (and (pair? binding) (car binding)))
 
@@ -373,6 +410,211 @@
             (append cond-errors then-errors else-errors
                     condition-errors branch-errors))))))
 
+  (def (lookup-variant name)
+    (lookup-name name (*variant-env*)))
+
+  (def (lookup-variant-case variant case-name)
+    (let loop ([rest (typed-variant-cases variant)])
+      (cond
+        [(null? rest) #f]
+        [(eq? (typed-variant-case-name (car rest)) case-name) (car rest)]
+        [else (loop (cdr rest))])))
+
+  (def (variant-case-names variant)
+    (map typed-variant-case-name (typed-variant-cases variant)))
+
+  (def (missing-variant-cases variant seen)
+    (let loop ([rest (variant-case-names variant)] [out '()])
+      (cond
+        [(null? rest) (reverse out)]
+        [(memq (car rest) seen) (loop (cdr rest) out)]
+        [else (loop (cdr rest) (cons (car rest) out))])))
+
+  (def (fallback-match-pattern? pattern)
+    (and (symbol? pattern) (memq pattern '(else _)) #t))
+
+  (def (known-branch-types types)
+    (let loop ([rest types] [out '()])
+      (cond
+        [(null? rest) (reverse out)]
+        [(car rest) (loop (cdr rest) (cons (car rest) out))]
+        [else (loop (cdr rest) out)])))
+
+  (def (consistent-branch-type types)
+    (let ([known (known-branch-types types)])
+      (cond
+        [(null? known) #f]
+        [else
+         (let ([first (car known)])
+           (let loop ([rest (cdr known)])
+             (cond
+               [(null? rest) first]
+               [(equal? (car rest) first) (loop (cdr rest))]
+               [else #f])))])))
+
+  (def (branch-type-errors types detail)
+    (let ([known (known-branch-types types)])
+      (cond
+        [(or (null? known) (null? (cdr known))) '()]
+        [else
+         (let ([first (car known)])
+           (let loop ([rest (cdr known)])
+             (cond
+               [(null? rest) '()]
+               [(equal? (car rest) first) (loop (cdr rest))]
+               [else
+                (list (make-check-error 'branch-type-mismatch
+                        "match branches must have the same type"
+                        detail))])))])))
+
+  (def (infer-match-case-clause clause pattern body variant env type-names seen)
+    (let* ([case-name (car pattern)]
+           [vars (cdr pattern)]
+           [case-decl (lookup-variant-case variant case-name)])
+      (cond
+        [(not case-decl)
+         (values #f
+                 (list (make-check-error 'unknown-match-case
+                         "match pattern is not a case of the variant"
+                         case-name))
+                 #f
+                 #f)]
+        [(not (symbol-list? vars))
+         (values #f
+                 (list (make-check-error 'bad-match-pattern
+                         "match case fields must be symbols"
+                         clause))
+                 case-name
+                 #f)]
+        [(memq case-name seen)
+         (values #f
+                 (list (make-check-error 'duplicate-match-case
+                         "duplicate variant case in match"
+                         case-name))
+                 case-name
+                 #f)]
+        [else
+         (let* ([case-field-types
+                 (field-types (typed-variant-case-fields case-decl))]
+                [arity-errors
+                 (if (= (length vars) (length case-field-types))
+                   '()
+                   (list (make-check-error 'bad-match-arity
+                           "match pattern field count does not match variant case"
+                           (list case-name
+                                 (length case-field-types)
+                                 (length vars)))))]
+                [duplicate-binding-errors
+                 (duplicate-errors 'duplicate-pattern-binding
+                   "duplicate binding in match pattern"
+                   (pattern-binding-names vars))])
+           (if (or (not (null? arity-errors))
+                   (not (null? duplicate-binding-errors)))
+             (values #f
+                     (append duplicate-binding-errors arity-errors)
+                     case-name
+                     #f)
+             (let-values ([(body-type body-errors)
+                           (infer-body
+                             body
+                             (extend-env* vars case-field-types env)
+                             type-names)])
+               (values body-type body-errors case-name #f))))])))
+
+  (def (infer-match-clause clause variant env type-names seen)
+    (cond
+      [(or (not (pair? clause)) (not (list? clause)))
+       (values #f
+               (list (make-check-error 'bad-match-clause
+                       "match clause must be a proper list"
+                       clause))
+               #f
+               #f)]
+      [(null? (cdr clause))
+       (values #f
+               (list (make-check-error 'bad-match-clause
+                       "match clause needs a body"
+                       clause))
+               #f
+               #f)]
+      [else
+       (let ([pattern (car clause)]
+             [body (cdr clause)])
+         (cond
+           [(fallback-match-pattern? pattern)
+            (let-values ([(body-type body-errors)
+                          (infer-body body env type-names)])
+              (values body-type body-errors #f #t))]
+           [(and (pair? pattern) (list? pattern) (symbol? (car pattern)))
+            (infer-match-case-clause
+              clause pattern body variant env type-names seen)]
+           [else
+            (values #f
+                    (list (make-check-error 'bad-match-pattern
+                            "match pattern must be a variant case, _, or else"
+                            pattern))
+                    #f
+                    #f)]))]))
+
+  (def (infer-match-clauses clauses variant env type-names expr)
+    (let loop ([rest clauses]
+               [seen '()]
+               [fallback? #f]
+               [branch-types '()]
+               [errors '()])
+      (if (null? rest)
+        (let* ([ordered-types (reverse branch-types)]
+               [branch-errors
+                (if (null? errors)
+                  (branch-type-errors ordered-types expr)
+                  '())]
+               [missing
+                (if (or fallback? (not (null? errors)))
+                  '()
+                  (missing-variant-cases variant seen))]
+               [exhaustiveness-errors
+                (if (null? missing)
+                  '()
+                  (list (make-check-error 'non-exhaustive-match
+                          "match does not cover all variant cases"
+                          missing)))]
+               [result-type (consistent-branch-type ordered-types)])
+          (values
+            (if (null? branch-errors) result-type #f)
+            (append errors exhaustiveness-errors branch-errors)))
+        (let-values ([(branch-type clause-errors case-name fallback-clause?)
+                      (infer-match-clause
+                        (car rest) variant env type-names seen)])
+          (loop (cdr rest)
+                (if case-name (add-unique case-name seen) seen)
+                (or fallback? fallback-clause?)
+                (cons branch-type branch-types)
+                (append errors clause-errors))))))
+
+  (def (infer-match args env type-names expr)
+    (if (< (length args) 2)
+      (values #f
+        (list (make-check-error 'bad-match
+                "match expects a target expression and at least one clause"
+                expr)))
+      (let-values ([(target-type target-errors)
+                    (infer-expression (car args) env type-names)])
+        (let ([variant (and target-type (lookup-variant target-type))])
+          (if (not variant)
+            (values #f
+              (append
+                target-errors
+                (if target-type
+                  (list (make-check-error 'match-type-mismatch
+                          "match target must have a declared variant type"
+                          target-type))
+                  '())))
+            (let-values ([(result-type clause-errors)
+                          (infer-match-clauses
+                            (cdr args) variant env type-names expr)])
+              (values result-type
+                      (append target-errors clause-errors))))))))
+
   (def (numeric-type? type)
     (and type (memq type '(Nat Int Fixnum Float)) #t))
 
@@ -524,6 +766,8 @@
             (infer-let (cadr expr) (cddr expr) env type-names expr))]
          [(if)
           (infer-if (cdr expr) env type-names expr)]
+         [(match)
+          (infer-match (cdr expr) env type-names expr)]
          [(+ - * /)
           (infer-arithmetic (car expr) (cdr expr) env type-names expr)]
          [(= < <= > >=)
@@ -590,8 +834,10 @@
     (let* ([declarations (typed-module-declarations module)]
            [type-names (declared-type-names declarations)]
            [value-names (declared-value-names declarations)]
-           [calls (call-env declarations)])
-      (parameterize ([*call-env* calls])
+           [calls (call-env declarations)]
+           [variants (declared-variant-env declarations)])
+      (parameterize ([*call-env* calls]
+                     [*variant-env* variants])
         (append
           (duplicate-errors 'duplicate-type
             "duplicate type declaration"
diff --git a/tests/fixtures/typed/valid-split-tree.ss b/tests/fixtures/typed/valid-split-tree.ss
index 1338998..82b7d67 100644
--- a/tests/fixtures/typed/valid-split-tree.ss
+++ b/tests/fixtures/typed/valid-split-tree.ss
@@ -1,6 +1,7 @@
 (typed-library (sample typed split-tree)
   (export make-Pane Pane? Pane-id Pane-focused? Pane-focused?-set!
-          split-size make-insert make-noop edit-op? EditOp? Insert Noop)
+          split-size make-insert make-noop edit-op? edit-size
+          EditOp? Insert Noop)
 
   (record Pane
     ((id : Nat)
@@ -20,4 +21,9 @@
     (Noop))
 
   (def (edit-op? (op : EditOp)) : Bool
-    (EditOp? op)))
+    (EditOp? op))
+
+  (def (edit-size (op : EditOp)) : Nat
+    (match op
+      ((Insert at text) at)
+      ((Noop) 0))))
diff --git a/tests/test-typed-checker.ss b/tests/test-typed-checker.ss
index e2ace10..17afb0f 100644
--- a/tests/test-typed-checker.ss
+++ b/tests/test-typed-checker.ss
@@ -344,6 +344,129 @@
          (Insert at at))))
   '(argument-type-mismatch))
 
+(test "variant match exhaustive"
+  (error-kinds
+    '(typed-library (body match-ok)
+       (export edit-size)
+       (variant EditOp
+         (Insert (at : Nat) (text : String))
+         (Noop))
+       (def (edit-size (op : EditOp)) : Nat
+         (match op
+           ((Insert at text) at)
+           ((Noop) 0)))))
+  '())
+
+(test "variant match binds field types"
+  (error-kinds
+    '(typed-library (body match-field)
+       (export edit-text)
+       (variant EditOp
+         (Insert (at : Nat) (text : String))
+         (Noop))
+       (def (edit-text (op : EditOp)) : String
+         (match op
+           ((Insert _ text) text)
+           ((Noop) "")))))
+  '())
+
+(test "variant match fallback suppresses exhaustiveness"
+  (error-kinds
+    '(typed-library (body match-fallback)
+       (export edit-size)
+       (variant EditOp
+         (Insert (at : Nat) (text : String))
+         (Noop))
+       (def (edit-size (op : EditOp)) : Nat
+         (match op
+           ((Insert at text) at)
+           (else 0)))))
+  '())
+
+(test "variant match detects missing case"
+  (error-kinds
+    '(typed-library (body match-missing)
+       (export edit-size)
+       (variant EditOp
+         (Insert (at : Nat) (text : String))
+         (Noop))
+       (def (edit-size (op : EditOp)) : Nat
+         (match op
+           ((Insert at text) at)))))
+  '(non-exhaustive-match))
+
+(test "variant match rejects non-variant target"
+  (error-kinds
+    '(typed-library (body match-target)
+       (export f)
+       (variant EditOp
+         (Noop))
+       (def (f (x : Nat)) : Nat
+         (match x
+           ((Noop) 0)))))
+  '(match-type-mismatch))
+
+(test "variant match rejects unknown case"
+  (error-kinds
+    '(typed-library (body match-unknown-case)
+       (export f)
+       (variant EditOp
+         (Noop))
+       (def (f (op : EditOp)) : Nat
+         (match op
+           ((Insert at text) at)
+           ((Noop) 0)))))
+  '(unknown-match-case))
+
+(test "variant match rejects bad case arity"
+  (error-kinds
+    '(typed-library (body match-bad-arity)
+       (export f)
+       (variant EditOp
+         (Insert (at : Nat) (text : String))
+         (Noop))
+       (def (f (op : EditOp)) : Nat
+         (match op
+           ((Insert at) at)
+           ((Noop) 0)))))
+  '(bad-match-arity))
+
+(test "variant match rejects duplicate case"
+  (error-kinds
+    '(typed-library (body match-duplicate-case)
+       (export f)
+       (variant EditOp
+         (Noop))
+       (def (f (op : EditOp)) : Nat
+         (match op
+           ((Noop) 0)
+           ((Noop) 1)))))
+  '(duplicate-match-case))
+
+(test "variant match rejects duplicate binding"
+  (error-kinds
+    '(typed-library (body match-duplicate-binding)
+       (export f)
+       (variant EditOp
+         (Insert (at : Nat) (text : String)))
+       (def (f (op : EditOp)) : Nat
+         (match op
+           ((Insert x x) x)))))
+  '(duplicate-pattern-binding))
+
+(test "variant match rejects branch type mismatch"
+  (error-kinds
+    '(typed-library (body match-branch-mismatch)
+       (export f)
+       (variant EditOp
+         (Insert (at : Nat) (text : String))
+         (Noop))
+       (def (f (op : EditOp)) : Nat
+         (match op
+           ((Insert at text) text)
+           ((Noop) 0)))))
+  '(branch-type-mismatch))
+
 (test "arithmetic primitive returns numeric type"
   (error-kinds
     '(typed-library (body arithmetic)