summaryrefslogtreecommitdiffstats
path: root/elab.sml
diff options
context:
space:
mode:
Diffstat (limited to 'elab.sml')
-rw-r--r--elab.sml114
1 files changed, 88 insertions, 26 deletions
diff --git a/elab.sml b/elab.sml
index b62adc9..3a361a6 100644
--- a/elab.sml
+++ b/elab.sml
@@ -12,7 +12,12 @@ struct
| "div" => Syntax.PDiv
| _ => raise Fail ("invalid op: " ^ s)
- fun enumerate l = ListPair.zip (List.tabulate (length l, (fn x => x)), l)
+ fun mapIdx (f : int * 'a -> 'b) (l : 'a list) : 'b list =
+ let
+ fun go _ [] = []
+ | go i (x :: xs) = f (i, x) :: go (i + 1) xs
+ in go 0 l
+ end
fun hdstls ([] : 'a list list) : ('a list * 'a list list) option = SOME ([], [])
| hdstls ([] :: _) = NONE
@@ -27,11 +32,34 @@ struct
NONE => []
| SOME (heads, tails) => heads :: transpose tails
+ type env = int StringMap.map * int StringMap.map
+
+ fun bindVar (name : string) (sym : int) ((vars, types) : env) : env = (StringMap.insert name sym vars, types)
+
+ fun lookupVar (name : string) ((vars, _) : env) : int =
+ case StringMap.lookup name vars of
+ SOME x => x
+ | NONE => raise Fail ("unbound identifier " ^ name)
+
+ fun bindDataCons (cons : (string * Syntax.etype option) list) ((vars, types) : env) : env =
+ let
+ val (_, vars, types) =
+ foldl
+ (fn ((name, _), (i, vars, types)) =>
+ (i + 1, StringMap.insert name (Gensym.new ()) vars, StringMap.insert name i types))
+ (0, vars, types)
+ cons
+ in (vars, types)
+ end
+
+ fun lookupCon (name : string) ((_, types) : env) : int option = StringMap.lookup name types
+
fun patternMatrix (arms : Syntax.pat list) : Syntax.pat list list =
let val ts =
foldl
(fn (Syntax.PWild, tupleSize) => tupleSize
| (Syntax.PVar _, tupleSize) => tupleSize
+ | (Syntax.PCon _, _) => ~2
| (Syntax.PInt _, _) => ~2
| (Syntax.PTuple t, ~1) => length t
| (Syntax.PTuple t, tupleSize) =>
@@ -64,12 +92,13 @@ struct
in
case cols of
[] => [expr]
- | _ => List.concat (map (fn (i, col) => occurrenceVector (Syntax.LSelect (i, expr)) col) (enumerate cols))
+ | _ => List.concat (mapIdx (fn (i, col) => occurrenceVector (Syntax.LSelect (i, expr)) col) 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))
+ List.concat (mapIdx (fn (i, p) => patternBindings (Syntax.LSelect (i, expr)) p) t)
+ | patternBindings expr (Syntax.PCon (_, arg)) = patternBindings (Syntax.LSelect (1, expr)) arg
| patternBindings _ _ = []
fun swap1 0 (l : 'a list) : 'a list = l
@@ -82,15 +111,26 @@ struct
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 =
+ fun specialize (env : env) (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
+ fun specializeRow (Syntax.PInt i :: rest) =
+ if i = n then SOME (Syntax.PWild :: rest) else NONE
+ | specializeRow (Syntax.PWild :: rest) = SOME (Syntax.PWild :: rest)
+ | specializeRow (Syntax.PVar _ :: rest) = SOME (Syntax.PWild :: rest)
+ | specializeRow (Syntax.PCon (con, Syntax.PTuple []) :: rest) =
+ specializeRow (Syntax.PCon (con, Syntax.PTuple [Syntax.PWild]) :: rest)
+ | specializeRow (Syntax.PCon (con, Syntax.PTuple args) :: rest) =
+ if valOf (lookupCon con env) = n
+ then SOME (args @ rest)
+ else NONE
+ | specializeRow (Syntax.PCon (con, obj) :: rest) =
+ specializeRow (Syntax.PCon (con, Syntax.PTuple [obj]) :: rest)
+ | specializeRow _ = raise Fail "unexpected pattern in the matrix"
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")
+ (fn (p, a) =>
+ Option.map (fn p => (p, a)) (specializeRow p))
(ListPair.zip (patterns, actions)))
in (patterns, tl occurrences, actions)
end
@@ -107,14 +147,26 @@ struct
in (patterns, tl occurrences, actions)
end
- fun compilePatternMatching ([] : Syntax.pat list list, _ : Syntax.lexp list, _ : Syntax.lexp list) : Syntax.lexp =
+ (* There's ambiguity between pattern variables and constructors that can only
+ * be resolved by looking up each pattern variable in the constructors map *)
+ fun fixConstructors (env : env) (expr : Syntax.pat) =
+ case expr of
+ Syntax.PTuple t => Syntax.PTuple (map (fixConstructors env) t)
+ | Syntax.PVar v =>
+ if isSome (lookupCon v env)
+ then Syntax.PCon (v, Syntax.PTuple [])
+ else Syntax.PVar v
+ | Syntax.PCon (s, p) => Syntax.PCon (s, fixConstructors env p)
+ | _ => expr
+
+ fun compilePatternMatching (env : env) ([] : Syntax.pat list list, _ : Syntax.lexp list, _ : Syntax.lexp list) : Syntax.lexp =
raise Fail "nonexhaustive match"
- | compilePatternMatching (patterns as firstRow :: rows, occurrences, actions) =
+ | compilePatternMatching env (patterns as firstRow :: rows, occurrences, actions) =
let val refutablePattern =
List.find
(fn (_, Syntax.PInt _) => true
| _ => false)
- (enumerate firstRow)
+ (mapIdx (fn x => x) firstRow)
in
case refutablePattern of
NONE => hd actions
@@ -136,14 +188,14 @@ struct
Syntax.LSwitch
( hd occurrences
, map
- (fn i => (i, compilePatternMatching (specialize i (patterns, occurrences, actions))))
+ (fn i => (i, compilePatternMatching env (specialize env i (patterns, occurrences, actions))))
signatures
- , compilePatternMatching (default (patterns, occurrences, actions))
+ , compilePatternMatching env (default (patterns, occurrences, actions))
)
end
end
- fun actionVector (env : int StringMap.map) (expr : Syntax.lexp) (arms : (Syntax.pat * Syntax.expr) list) : Syntax.lexp list =
+ fun actionVector (env : env) (expr : Syntax.lexp) (arms : (Syntax.pat * Syntax.expr) list) : Syntax.lexp list =
map
(fn (p, body) =>
let
@@ -151,20 +203,21 @@ struct
val env =
foldl
(fn ((name, _), env) =>
- StringMap.insert name (Gensym.new ()) env)
+ bindVar name (Gensym.new ()) env)
env
bindings
in
foldl
(fn ((name, binding), acc) =>
- Syntax.LApp (Syntax.LFn (valOf (StringMap.lookup name env), acc), binding))
+ Syntax.LApp (Syntax.LFn (lookupVar 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) =
+ and elabCase (env : env) (expr : Syntax.lexp) (arms : (Syntax.pat * Syntax.expr) list) =
let
+ val arms = map (fn (p, e) => (fixConstructors env p, e)) arms
val patterns = transpose (patternMatrix (map (fn (x, _) => x) arms))
val occurrences = occurrenceVector expr (map (fn (x, _) => x) arms)
val actions = actionVector env expr arms
@@ -179,15 +232,13 @@ struct
in
Syntax.LFix
( actionFns
- , compilePatternMatching (patterns, occurrences, smallActions)
+ , compilePatternMatching env (patterns, occurrences, smallActions)
)
end
- and elab (env : int StringMap.map) (p : Syntax.expr) : Syntax.lexp =
+ and elab (env : env) (p : Syntax.expr) : Syntax.lexp =
case p of
- Syntax.EIdent [i] => (case StringMap.lookup i env of
- NONE => raise Fail ("unbound identifier " ^ i)
- | SOME x => Syntax.LVar x)
+ Syntax.EIdent [i] => Syntax.LVar (lookupVar i env)
| Syntax.EIdent _ => raise Fail "long identifiers are not supported"
| Syntax.EBuiltin builtin => Syntax.LPrim (primop builtin)
| Syntax.EInt i => Syntax.LInt i
@@ -204,18 +255,29 @@ struct
| Syntax.EAndAlso (_, _) => raise Fail "unimplemented"
| Syntax.EOrElse (_, _) => raise Fail "unimplemented"
| Syntax.ELet ([], body) => elab env body
+ | Syntax.ELet (Syntax.DDatatype (name, cons) :: decls, body) =>
+ let
+ val env = bindDataCons cons env
+ fun go _ [] = []
+ | go i ((name, _) :: cons) =
+ let val v = Gensym.new ()
+ in (lookupVar name env, v, Syntax.LRecord [Syntax.LInt i, Syntax.LVar v]) :: go (i + 1) cons
+ end
+ in
+ Syntax.LFix (go 0 cons, elab env (Syntax.ELet (decls, body)))
+ end
| Syntax.ELet (Syntax.DVal (pat, v) :: decls, body) =>
elab env (Syntax.ECase (v, [(pat, Syntax.ELet (decls, body))]))
| Syntax.ELet (Syntax.DValRec (Syntax.PVar name, f as Syntax.ELambda _) :: decls, body) =>
let
val n = Gensym.new ()
- val env' = StringMap.insert name n env
+ val env = bindVar name n env
val (arg, fnBody) =
- case elab env' f of
+ case elab env f of
Syntax.LFn x => x
| _ => raise Fail "Syntax.ELambda should expand to Syntax.LFn"
in
- Syntax.LFix ([(n, arg, fnBody)], elab env' (Syntax.ELet (decls, body)))
+ Syntax.LFix ([(n, arg, fnBody)], elab env (Syntax.ELet (decls, body)))
end
| Syntax.ELet (Syntax.DValRec _ :: _, _) => raise Fail "invalid val rec"
| Syntax.ELambda body =>
@@ -227,5 +289,5 @@ struct
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
+ fun elaborate (p : Syntax.expr) : Syntax.lexp = elab (StringMap.empty, StringMap.empty) p
end