From c39c43e144733973ebf9cd6842df75b5a5ceee0b Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Fri, 17 May 2024 12:36:56 -0700 Subject: Add datatypes. --- elab.sml | 114 ++++++++++++++++++++++++++++++++++++++++++++++++--------------- 1 file changed, 88 insertions(+), 26 deletions(-) (limited to 'elab.sml') 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 -- cgit v1.3.1