summaryrefslogtreecommitdiffstats
path: root/Elab.sml
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2025-11-04 18:58:28 -0800
committerRose Hogenson <rosehogenson@posteo.net>2025-11-04 21:05:18 -0800
commiteaa183b7dfc841df75d42ace7a90ada1e40ff282 (patch)
tree1d627b8088d5996a7bbf7c0a7439360e31de3dc3 /Elab.sml
parentc48a992a6ed6ebd79c344b37364a680ddd948dea (diff)
downloadsml-main.tar.zst
Add listsHEADmain
Diffstat (limited to 'Elab.sml')
-rw-r--r--Elab.sml36
1 files changed, 30 insertions, 6 deletions
diff --git a/Elab.sml b/Elab.sml
index ba436f5..13bc890 100644
--- a/Elab.sml
+++ b/Elab.sml
@@ -59,6 +59,8 @@ struct
fun patternBindings (expr : Syntax.lexp) (Syntax.TPVar v, _ : Syntax.ty) : (string * Syntax.lexp) list = [(v, expr)]
| patternBindings expr (Syntax.TPTuple t, _) =
List.concat (map (fn (i, p) => patternBindings (Syntax.LSelect (i, expr)) p) (enumerate t))
+ | patternBindings _ (Syntax.TPList [], _) = []
+ | patternBindings expr (Syntax.TPList (x :: xs), t) = patternBindings (Syntax.LSelect (0, Syntax.LSelect (1, expr))) x @ patternBindings (Syntax.LSelect (1, Syntax.LSelect (1, expr))) (Syntax.TPList xs, t)
| patternBindings expr (Syntax.TPCon (_, arg), _) = patternBindings (Syntax.LSelect (1, expr)) arg
| patternBindings _ _ = []
@@ -117,11 +119,15 @@ struct
val con1 =
List.find
(fn (Syntax.TPCon (name, _), Syntax.TDatatype cons) => lookupCon (List.last name) cons = n
+ | (Syntax.TPList [], _) => n = 0
+ | (_, Syntax.TList _) => n = 1
| _ => false)
(map hd patterns)
val newTupleSize =
case con1 of
SOME (Syntax.TPCon (_, (Syntax.TPTuple t, _)), _) => length t
+ | SOME (Syntax.TPList [], _) => 0
+ | SOME (_, Syntax.TList _) => 2
| _ => 0
val occHead = hd occurrences
val occRest = tl occurrences
@@ -130,10 +136,17 @@ struct
then Syntax.LSelect (1, occHead) :: occRest
else
List.tabulate (newTupleSize, fn i => Syntax.LSelect (i, Syntax.LSelect (1, occHead))) @ occRest
+ val newTupleSize = if newTupleSize = 0 then 1 else newTupleSize
fun specializeRow ((Syntax.TPInt i, ty):: rest) =
if i = n then SOME ((Syntax.TPWild, ty) :: rest) else NONE
- | specializeRow ((Syntax.TPWild, ty) :: rest) = SOME ((Syntax.TPWild, ty) :: rest)
- | specializeRow ((Syntax.TPVar _, ty) :: rest) = SOME ((Syntax.TPWild, ty) :: rest)
+ | specializeRow ((Syntax.TPWild, ty) :: rest) = SOME (List.tabulate (newTupleSize, fn _ => (Syntax.TPWild, ty)) @ rest)
+ | specializeRow ((Syntax.TPVar _, ty) :: rest) = SOME (List.tabulate (newTupleSize, fn _ => (Syntax.TPWild, ty)) @ rest)
+ | specializeRow ((Syntax.TPList [], Syntax.TList ty) :: rest) =
+ if n = 0 then SOME ((Syntax.TPWild, ty) :: rest) else NONE
+ | specializeRow ((Syntax.TPList (pat :: pats), ty) :: rest) =
+ if n = 1 then SOME ([pat, (Syntax.TPList pats, ty)] @ rest) else NONE
+ | specializeRow ((Syntax.TPCon (_, (Syntax.TPTuple pats, _)), Syntax.TList _) :: rest) =
+ if n = 1 then SOME (pats @ rest) else NONE
| specializeRow ((Syntax.TPCon (con, (Syntax.TPTuple [], tupleTy)), conTy) :: rest) =
specializeRow ((Syntax.TPCon (con, (Syntax.TPTuple [(Syntax.TPWild, Syntax.TTuple [])], tupleTy)), conTy) :: rest)
| specializeRow ((Syntax.TPCon (con, (Syntax.TPTuple args, _)), Syntax.TDatatype cons) :: rest) =
@@ -172,6 +185,7 @@ struct
List.find
(fn (_, (Syntax.TPInt _, _)) => true
| (_, (Syntax.TPCon _, _)) => true
+ | (_, (Syntax.TPList _, _)) => true
| _ => false)
(enumerate firstRow)
in
@@ -189,13 +203,16 @@ struct
(IntMap.toList
(foldl
(fn ((Syntax.TPInt i, _), acc) => IntMap.insert i true acc
+ | ((Syntax.TPList [], _), acc) => IntMap.insert 0 true acc
+ | ((_, Syntax.TList _), acc) => IntMap.insert 1 true acc
| ((Syntax.TPCon (c, _), Syntax.TDatatype cons), acc) => IntMap.insert (lookupCon (List.last c) cons) true acc
| (_, acc) => acc)
IntMap.empty
firstCol))
val nCons =
- case List.find (fn (Syntax.TPCon _, _) => true | _ => false) firstCol of
+ case List.find (fn (Syntax.TPCon _, _) => true | (Syntax.TPList _, _) => true | _ => false) firstCol of
SOME (Syntax.TPCon _, Syntax.TDatatype cons) => length cons
+ | SOME (_, Syntax.TList _) => 2
| _ => ~1
val defaultCase =
if length signatures = nCons
@@ -273,8 +290,8 @@ struct
| Syntax.TEList exprs =>
foldr
(fn (x, acc) =>
- Syntax.LRecord [elab env x, acc])
- (Syntax.LInt 0)
+ Syntax.LRecord [Syntax.LInt 1, Syntax.LRecord [elab env x, acc]])
+ (Syntax.LRecord [Syntax.LInt 0])
exprs
| Syntax.TEApp (f, x) => Syntax.LApp (elab env f, elab env x)
| Syntax.TEAndAlso (_, _) => raise Fail "unimplemented"
@@ -400,5 +417,12 @@ struct
in Syntax.LApp (func, tuple) end
| _ => raise Fail ("invalid expression " ^ ShowSyntax.typedExprToString p)
- fun elaborate (p : Syntax.typedExpr * Syntax.ty) : Syntax.lexp = elab emptyEnv p
+ fun elaborate (p : Syntax.typedExpr * Syntax.ty) : Syntax.lexp =
+ let
+ val cons = GenSym.new ()
+ val alpha = GenSym.new ()
+ val env = IdentMap.insert (Syntax.ITVar, "::") cons emptyEnv
+ in
+ Syntax.LApp (Syntax.LFn (cons, elab env p), Syntax.LFn (alpha, Syntax.LRecord [Syntax.LInt 1, Syntax.LVar alpha]))
+ end
end