summaryrefslogtreecommitdiffstats
path: root/elab.sml
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2024-05-05 21:12:16 -0700
committerRose Hogenson <rosehogenson@posteo.net>2024-05-05 21:12:16 -0700
commit457c482397e2ca745067e5a3ef2c7d71d7daf187 (patch)
treeaef2b8d88421829105a3c6b8887d8ffde39d5e1e /elab.sml
parent048895eb0c85b591b2dd338dd06514bc1b18bacb (diff)
downloadsml-457c482397e2ca745067e5a3ef2c7d71d7daf187.tar.zst
Add tuples.
Diffstat (limited to 'elab.sml')
-rw-r--r--elab.sml204
1 files changed, 182 insertions, 22 deletions
diff --git a/elab.sml b/elab.sml
index af6bec2..784515d 100644
--- a/elab.sml
+++ b/elab.sml
@@ -1,6 +1,7 @@
structure Elab =
struct
- structure StringMap = Map(type k = string val cmp = String.compare);
+ structure StringMap = Map(type k = string val cmp = String.compare)
+ structure IntMap = Map(type k = int val cmp = Int.compare)
fun primop (s : string) : Syntax.primop =
case s of
@@ -11,7 +12,181 @@ struct
| "div" => Syntax.PDiv
| _ => raise Fail ("invalid op: " ^ s)
- fun elab (env : int StringMap.map) (p : Syntax.expr) : Syntax.lexp =
+ fun enumerate l = ListPair.zip (List.tabulate (length l, (fn x => x)), l)
+
+ fun hdstls ([] : 'a list list) : ('a list * 'a list list) option = SOME ([], [])
+ | hdstls ([] :: _) = NONE
+ | hdstls ((x :: xs) :: ls) =
+ case hdstls ls of
+ NONE => NONE
+ | SOME (heads, tails) => SOME (x :: heads, xs :: tails)
+
+ fun transpose ([] : 'a list list) : 'a list list = []
+ | transpose l =
+ case hdstls l of
+ NONE => []
+ | SOME (heads, tails) => heads :: transpose tails
+
+ fun patternMatrix (arms : Syntax.pat list) : Syntax.pat list list =
+ let val ts =
+ foldl
+ (fn (Syntax.PWild, tupleSize) => tupleSize
+ | (Syntax.PVar _, tupleSize) => tupleSize
+ | (Syntax.PInt _, _) => ~2
+ | (Syntax.PTuple t, ~1) => length t
+ | (Syntax.PTuple t, tupleSize) =>
+ if tupleSize = length t
+ then tupleSize
+ else raise Fail "a type checker would have caught this")
+ ~1
+ arms
+ in if ts < 0
+ then [arms]
+ else let
+ val expandedArms =
+ map
+ (fn Syntax.PTuple t => t
+ | _ => List.tabulate (ts, fn _ => Syntax.PWild))
+ arms
+ val cols = transpose expandedArms
+ in List.concat (map patternMatrix cols)
+ end
+ end
+
+ fun occurrenceVector (expr : Syntax.lexp) (arms : Syntax.pat list) : Syntax.lexp list =
+ let
+ val expandedArms =
+ List.mapPartial
+ (fn Syntax.PTuple t => SOME t
+ | _ => NONE)
+ arms
+ val cols = transpose expandedArms
+ in
+ case cols of
+ [] => [expr]
+ | _ => List.concat (map (fn (i, col) => occurrenceVector (Syntax.LSelect (i, expr)) col) (enumerate cols))
+ end
+
+ fun patternBindings (expr : Syntax.lexp) (Syntax.PVar v) : (string * Syntax.lexp) list = [(v, expr)]
+ | patternBindings expr (Syntax.PTuple t) =
+ List.concat (map (fn (i, p) => patternBindings (Syntax.LSelect (i, expr)) p) (enumerate t))
+ | patternBindings _ _ = []
+
+ fun swap1 0 (l : 'a list) : 'a list = l
+ | swap1 n (first :: rest) =
+ (case swap1 (n - 1) rest of
+ x :: rest => x :: first :: rest
+ | _ => raise Fail "swap1: index out of bounds")
+ | swap1 _ _ = raise Fail "swap1: index out of bounds"
+
+ fun swap (n : int) (patterns : Syntax.pat list list) (occurrences : Syntax.lexp list) : Syntax.pat list list * Syntax.lexp list =
+ (map (swap1 n) patterns, swap1 n occurrences)
+
+ fun specialize (n : int) (patterns : Syntax.pat list list, occurrences : Syntax.lexp list, actions : Syntax.lexp list) : Syntax.pat list list * Syntax.lexp list * Syntax.lexp list =
+ let
+ val (patterns, actions) =
+ ListPair.unzip
+ (List.mapPartial
+ (fn (Syntax.PInt i :: rest, a) => if i = n then SOME (rest, a) else NONE
+ | (Syntax.PWild :: rest, a) => SOME (rest, a)
+ | (Syntax.PVar _ :: rest, a) => SOME (rest, a)
+ | _ => raise Fail "unexpected pattern in the matrix")
+ (ListPair.zip (patterns, actions)))
+ in (patterns, tl occurrences, actions)
+ end
+
+ fun default (patterns : Syntax.pat list list, occurrences : Syntax.lexp list, actions : Syntax.lexp list) : Syntax.pat list list * Syntax.lexp list * Syntax.lexp list =
+ let
+ val (patterns, actions) =
+ ListPair.unzip
+ (List.mapPartial
+ (fn (Syntax.PWild :: rest, a) => SOME (rest, a)
+ | (Syntax.PVar _ :: rest, a) => SOME (rest, a)
+ | _ => NONE)
+ (ListPair.zip (patterns, actions)))
+ in (patterns, tl occurrences, actions)
+ end
+
+ fun compilePatternMatching ([] : Syntax.pat list list, _ : Syntax.lexp list, _ : Syntax.lexp list) : Syntax.lexp =
+ raise Fail "nonexhaustive match"
+ | compilePatternMatching (patterns as firstRow :: rows, occurrences, actions) =
+ let val refutablePattern =
+ List.find
+ (fn (_, Syntax.PInt _) => true
+ | _ => false)
+ (enumerate firstRow)
+ in
+ case refutablePattern of
+ NONE =>
+ if null rows
+ then hd actions
+ else raise Fail "redundant match"
+ | SOME (i, _) =>
+ let
+ val (patterns, occurrences) =
+ if i = 0
+ then (patterns, occurrences)
+ else swap i patterns occurrences
+ val signatures =
+ map (fn (x, _) => x)
+ (IntMap.toList
+ (foldl
+ (fn (Syntax.PInt i, acc) => IntMap.insert i true acc
+ | (_, acc) => acc)
+ IntMap.empty
+ (map hd patterns)))
+ in
+ Syntax.LSwitch
+ ( hd occurrences
+ , map
+ (fn i => (i, compilePatternMatching (specialize i (patterns, occurrences, actions))))
+ signatures
+ , compilePatternMatching (default (patterns, occurrences, actions))
+ )
+ end
+ end
+
+ fun actionVector (env : int StringMap.map) (expr : Syntax.lexp) (arms : (Syntax.pat * Syntax.expr) list) : Syntax.lexp list =
+ map
+ (fn (p, body) =>
+ let
+ val bindings = patternBindings expr p
+ val env =
+ foldl
+ (fn ((name, _), env) =>
+ StringMap.insert name (Gensym.new ()) env)
+ env
+ bindings
+ in
+ foldl
+ (fn ((name, binding), acc) =>
+ Syntax.LApp (Syntax.LFn (valOf (StringMap.lookup name env), acc), binding))
+ (elab env body)
+ bindings
+ end)
+ arms
+
+ and elabCase (env : int StringMap.map) (expr : Syntax.lexp) (arms : (Syntax.pat * Syntax.expr) list) =
+ let
+ val patterns = transpose (patternMatrix (map (fn (x, _) => x) arms))
+ val occurrences = occurrenceVector expr (map (fn (x, _) => x) arms)
+ val actions = actionVector env expr arms
+ val actionFns =
+ map
+ (fn a => (Gensym.new (), Gensym.new (), a))
+ actions
+ val smallActions =
+ map
+ (fn (f, _, _) => Syntax.LApp (Syntax.LVar f, Syntax.LInt 0))
+ actionFns
+ in
+ Syntax.LFix
+ ( actionFns
+ , compilePatternMatching (patterns, occurrences, smallActions)
+ )
+ end
+
+ and elab (env : int StringMap.map) (p : Syntax.expr) : Syntax.lexp =
case p of
Syntax.EIdent [i] => (case StringMap.lookup i env of
NONE => raise Fail ("unbound identifier " ^ i)
@@ -46,28 +221,13 @@ struct
Syntax.LFix ([(n, arg, fnBody)], elab env' (Syntax.ELet (decls, body)))
end
| Syntax.ELet (Syntax.DValRec _ :: _, _) => raise Fail "invalid val rec"
- | Syntax.ELambda (pat, body) =>
- let
- val v = Gensym.new ()
- val env' =
- case pat of
- Syntax.PVar name => StringMap.insert name v env
- | _ => env
- in Syntax.LFn (v, elab env' body)
+ | Syntax.ELambda body =>
+ let val v = Gensym.new ()
+ in Syntax.LFn (v, elabCase env (Syntax.LVar v) [body])
end
- | Syntax.ECase (expr, []) => raise Fail "nonexhaustive match"
| Syntax.ECase (expr, arms) =>
- let
- fun go [] _ = raise Fail "nonexhaustive match"
- | go ((Syntax.PWild, body) :: []) acc = Syntax.LSwitch (elab env expr, rev acc, elab env body)
- | go ((Syntax.PVar name, body) :: []) acc =
- let val v = Gensym.new () in
- Syntax.LApp (Syntax.LFn (v, Syntax.LSwitch (Syntax.LVar v, rev acc, elab (StringMap.insert name v env) body)), elab env expr)
- end
- | go ((Syntax.PInt i, body) :: rest) acc =
- go rest ((i, elab env body) :: acc)
- | go _ _ = raise Fail "redundant match"
- in go arms []
+ let val v = Gensym.new ()
+ in Syntax.LApp (Syntax.LFn (v, elabCase env (Syntax.LVar v) arms), elab env expr)
end
fun elaborate (p : Syntax.expr) : Syntax.lexp = elab StringMap.empty p