diff options
Diffstat (limited to 'Elab.sml')
| -rw-r--r-- | Elab.sml | 306 |
1 files changed, 126 insertions, 180 deletions
@@ -35,80 +35,42 @@ struct NONE => [] | SOME (heads, tails) => heads :: transpose tails - datatype env = Env of { vars: int StringMap.map, types: (int * int) StringMap.map, structTypes: env StringMap.map } + datatype env = Env of int StringMap.map - val emptyEnv = Env { vars = StringMap.empty, types = StringMap.empty, structTypes = StringMap.empty } + val emptyEnv = Env StringMap.empty - fun bindVar (name : string) (sym : int) (Env env) : env = Env { vars = StringMap.insert name sym (#vars env), types = #types env, structTypes = #structTypes env } + fun bindVar (name : string) (sym : int) (Env env) : env = Env (StringMap.insert name sym env) fun lookupVar (name : string) (Env env) : int = - case StringMap.lookup name (#vars env) of + case StringMap.lookup name env of SOME x => x | NONE => raise Fail ("unbound identifier " ^ name) - fun bindDataCons (cons : (string * Syntax.etype option) list) (Env env) : env = - let - val nCons = length cons - val (_, vars, types) = - foldl - (fn ((name, _), (i, vars, types)) => - (i + 1, StringMap.insert name (GenSym.new ()) vars, StringMap.insert name (i, nCons) types)) - (0, #vars env, #types env) - cons - in Env { vars = vars, types = types, structTypes = #structTypes env } - end - - fun lookupStructType (name : string) (Env env) : env = - case StringMap.lookup name (#structTypes env) of - SOME x => x - | NONE => raise Fail ("unbound structure " ^ name) - - fun lookupCon ([name] : string list) (Env env) : int option = - Option.map (fn (i, _) => i) (StringMap.lookup name (#types env)) - | lookupCon (structName :: names) env = - lookupCon names (lookupStructType structName env) - | lookupCon [] _ = raise Fail "lookupCon empty" - - fun nConstructors ([name] : string list) (Env env) : int = - let val (_, n) = valOf (StringMap.lookup name (#types env)) - in n - end - | nConstructors (structName :: names) env = - nConstructors names (lookupStructType structName env) - | nConstructors [] _ = raise Fail "nConstructors empty" + fun bindDataCons (cons : (string * Syntax.ty option) list) (Env env) : env = + Env + (foldl + (fn ((name, _), acc) => + StringMap.insert name (GenSym.new ()) acc) + env + cons) - 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) => - 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 patternMatrix (arms : (Syntax.typedPat * Syntax.ty) list) : (Syntax.typedPat * Syntax.ty) list list = + if not (isSome (List.find (fn (Syntax.TPTuple t, _) => true | _ => false) arms)) then [arms] else + let + val expandedArms = + map + (fn (Syntax.TPTuple t, _) => t + | (_, Syntax.TTuple tys) => map (fn ty => (Syntax.TPWild, ty)) tys + | _ => raise Fail "unreachable") + arms + val cols = transpose expandedArms + in List.concat (map patternMatrix cols) end - fun occurrenceVector (expr : Syntax.lexp) (arms : Syntax.pat list) : Syntax.lexp list = + fun occurrenceVector (expr : Syntax.lexp) (arms : (Syntax.typedPat * Syntax.ty) list) : Syntax.lexp list = let val expandedArms = List.mapPartial - (fn Syntax.PTuple t => SOME t + (fn (Syntax.TPTuple t, _) => SOME t | _ => NONE) arms val cols = transpose expandedArms @@ -118,10 +80,10 @@ struct | _ => 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) = + 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 expr (Syntax.PCon (_, arg)) = patternBindings (Syntax.LSelect (1, expr)) arg + | patternBindings expr (Syntax.TPCon (_, arg), _) = patternBindings (Syntax.LSelect (1, expr)) arg | patternBindings _ _ = [] fun swap1 0 (l : 'a list) : 'a list = l @@ -131,19 +93,24 @@ struct | _ => 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 = + fun swap (n : int) (patterns : (Syntax.typedPat * Syntax.ty) list list) (occurrences : Syntax.lexp list) : (Syntax.typedPat * Syntax.ty) list list * Syntax.lexp list = (map (swap1 n) patterns, swap1 n occurrences) - 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 = + fun lookupCon (name : string) (ty : string list) : int = + case List.find (fn (_, x) => x = name) (enumerate ty) of + SOME (i, _) => i + | NONE => raise Fail ("Unknown field " ^ name) + + fun specialize (n : int) (patterns : (Syntax.typedPat * Syntax.ty) list list, occurrences : Syntax.lexp list, actions : Syntax.lexp list) : (Syntax.typedPat * Syntax.ty) list list * Syntax.lexp list * Syntax.lexp list = let val con1 = List.find - (fn Syntax.PCon (name, _) => valOf (lookupCon name env) = n + (fn (Syntax.TPCon (name, _), Syntax.TDatatype cons) => lookupCon (List.last name) cons = n | _ => false) (map hd patterns) val newTupleSize = case con1 of - SOME (Syntax.PCon (_, Syntax.PTuple t)) => length t + SOME (Syntax.TPCon (_, (Syntax.TPTuple t, _)), _) => length t | _ => 0 val occHead = hd occurrences val occRest = tl occurrences @@ -152,19 +119,19 @@ struct then Syntax.LSelect (1, occHead) :: occRest else List.tabulate (newTupleSize, fn i => Syntax.LSelect (i, Syntax.LSelect (1, occHead))) @ occRest - 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 + 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.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) = + if lookupCon (List.last con) cons = 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" + | specializeRow ((Syntax.TPCon (con, (obj, objTy)), ty) :: rest) = + specializeRow ((Syntax.TPCon (con, (Syntax.TPTuple [(obj, objTy)], Syntax.TTuple [objTy])), ty) :: rest) + | specializeRow _ = raise Fail "you think that's air you're breathing now?" val (patterns, actions) = ListPair.unzip (List.mapPartial @@ -174,26 +141,26 @@ struct in (patterns, 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 = + fun default (patterns : (Syntax.typedPat * Syntax.ty) list list, occurrences : Syntax.lexp list, actions : Syntax.lexp list) : (Syntax.typedPat * Syntax.ty) list list * Syntax.lexp list * Syntax.lexp list = let val (patterns, actions) = ListPair.unzip (List.filter - (fn (Syntax.PWild :: _, _) => true - | (Syntax.PVar _ :: _, _) => true + (fn ((Syntax.TPWild, _) :: _, _) => true + | ((Syntax.TPVar _, _) :: _, _) => true | _ => false) (ListPair.zipEq (patterns, actions))) in (patterns, occurrences, actions) end (* https://compiler.club/compiling-pattern-matching/ *) - fun compilePatternMatching (env : env) ([] : Syntax.pat list list, _ : Syntax.lexp list, _ : Syntax.lexp list) : Syntax.lexp = + fun compilePatternMatching ([] : (Syntax.typedPat * Syntax.ty) list list, _ : Syntax.lexp list, _ : Syntax.lexp list) : Syntax.lexp = raise Fail "nonexhaustive match" - | compilePatternMatching env (patterns as firstRow :: rows, occurrences, actions) = + | compilePatternMatching (patterns as firstRow :: rows, occurrences, actions) = let val refutablePattern = List.find - (fn (_, Syntax.PInt _) => true - | (_, Syntax.PCon _) => true + (fn (_, (Syntax.TPInt _, _)) => true + | (_, (Syntax.TPCon _, _)) => true | _ => false) (enumerate firstRow) in @@ -210,19 +177,19 @@ struct map (fn (x, _) => x) (IntMap.toList (foldl - (fn (Syntax.PInt i, acc) => IntMap.insert i true acc - | (Syntax.PCon (c, _), acc) => IntMap.insert (valOf (lookupCon c env)) true acc + (fn ((Syntax.TPInt i, _), acc) => IntMap.insert i 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.PCon _) => true | _ => false) firstCol of - SOME (Syntax.PCon (name, _)) => nConstructors name env + case List.find (fn (Syntax.TPCon _, _) => true | _ => false) firstCol of + SOME (Syntax.TPCon _, Syntax.TDatatype cons) => length cons | _ => ~1 val defaultCase = if length signatures = nCons then NONE - else SOME (compilePatternMatching env (default (patterns, occurrences, actions))) + else SOME (compilePatternMatching (default (patterns, occurrences, actions))) val switchOperand = if nCons < 0 then hd occurrences @@ -231,49 +198,22 @@ struct Syntax.LSwitch ( switchOperand , map - (fn i => (i, compilePatternMatching env (specialize env i (patterns, occurrences, actions)))) + (fn i => (i, compilePatternMatching (specialize i (patterns, occurrences, actions)))) signatures , defaultCase ) end end - fun declBoundVars (Syntax.DVal (p, _)) : string list = map (fn (x, _) => x) (patternBindings (Syntax.LInt 0) p) - | declBoundVars (Syntax.DValRec (p, _)) = map (fn (x, _) => x) (patternBindings (Syntax.LInt 0) p) - | declBoundVars (Syntax.DFun (name, _)) = [name] - | declBoundVars (Syntax.DDatatype (_, cons)) = map (fn (x, _) => x) cons - | declBoundVars (Syntax.DStruct (name, _)) = [name] - - fun structBoundVars (decls : Syntax.dec list) : string list = List.concatMap declBoundVars decls + fun declBoundVars (Syntax.TDVal (p, _)) : string list = map (fn (x, _) => x) (patternBindings (Syntax.LInt 0) p) + | declBoundVars (Syntax.TDValRec (p, _)) = map (fn (x, _) => x) (patternBindings (Syntax.LInt 0) p) + | declBoundVars (Syntax.TDFun (name, _)) = [name] + | declBoundVars (Syntax.TDDatatype (_, cons)) = map (fn (x, _) => x) cons + | declBoundVars (Syntax.TDStruct (name, _)) = [name] - fun bindStruct (name : string) (structEnv : env) (Env { vars, types, structTypes }) : env = - Env { vars = vars, types = types, structTypes = StringMap.insert name structEnv structTypes } - - fun bindStructType (name : string) (decls : Syntax.dec list) (env : env) : env = - let - val structEnv = - foldl - (fn (Syntax.DDatatype (_, cons), env) => bindDataCons cons env - | (Syntax.DStruct (name, Syntax.SStruct decls), env) => bindStructType name decls env - | (Syntax.DStruct (name, Syntax.SIdent ident), env) => - let - fun lookup [name] env = lookupStructType name env - | lookup (name :: names) env = lookup names (lookupStructType name env) - | lookup _ _ = raise Fail "invalid struct identifier" - in lookup ident env end - | (_, env) => env) - emptyEnv - decls - val structEnv = - foldl - (fn ((i, n), env) => bindVar n i env) - structEnv - (enumerate (structBoundVars decls)) - in - bindStruct name structEnv env - end + fun structBoundVars (decls : Syntax.typedDec list) : string list = List.concatMap declBoundVars decls - fun actionVector (env : env) (expr : Syntax.lexp) (arms : (Syntax.pat * Syntax.expr) list) : Syntax.lexp list = + fun actionVector (env : env) (expr : Syntax.lexp) (arms : ((Syntax.typedPat * Syntax.ty) * (Syntax.typedExpr * Syntax.ty)) list) : Syntax.lexp list = map (fn (p, body) => let @@ -293,7 +233,7 @@ struct end) arms - and elabCase (env : env) (expr : Syntax.lexp) (arms : (Syntax.pat * Syntax.expr) list) = + and elabCase (env : env) (expr : Syntax.lexp) (arms : ((Syntax.typedPat * Syntax.ty) * (Syntax.typedExpr * Syntax.ty)) list) = let val patterns = transpose (patternMatrix (map (fn (x, _) => x) arms)) val occurrences = occurrenceVector expr (map (fn (x, _) => x) arms) @@ -309,41 +249,35 @@ struct in Syntax.LFix ( actionFns - , compilePatternMatching env (patterns, occurrences, smallActions) + , compilePatternMatching (patterns, occurrences, smallActions) ) end - and elab (env : env) (p : Syntax.expr) : Syntax.lexp = + and elab (env : env) (p : Syntax.typedExpr, ty : Syntax.ty) : Syntax.lexp = case p of - Syntax.EIdent [i] => Syntax.LVar (lookupVar i env) - | Syntax.EIdent (structName :: accessors) => + Syntax.TEIdent i => Syntax.LVar (lookupVar i env) + | Syntax.TEDot (structExpr as (_, Syntax.TStruct (_, fields)), field) => let - val s = Syntax.LVar (lookupVar structName env) - val env = lookupStructType structName env - fun go env [i] acc = Syntax.LSelect (lookupVar i env, acc) - | go env (accessor :: accessors) acc = - let val env = lookupStructType accessor env - in go env accessors (Syntax.LSelect (lookupVar accessor env, acc)) end - | go _ [] _ = raise Fail "go empty" - in go env accessors s - end - | Syntax.EIdent [] => raise Fail "invalid syntax" - | Syntax.EBuiltin builtin => Syntax.LPrim (primop builtin) - | Syntax.EInt i => Syntax.LInt i - | Syntax.EStr s => Syntax.LString s - | Syntax.ETuple exprs => Syntax.LRecord (map (elab env) exprs) - | Syntax.EList exprs => + val i = + case List.find (fn (_, (x, _)) => x = field) (enumerate fields) of + NONE => raise Fail ("Unknown field " ^ field) + | SOME (i, _) => i + in Syntax.LSelect (i, elabStructExpr env structExpr) end + | Syntax.TEBuiltin builtin => Syntax.LPrim (primop builtin) + | Syntax.TEInt i => Syntax.LInt i + | Syntax.TEStr s => Syntax.LString s + | Syntax.TETuple exprs => Syntax.LRecord (map (elab env) exprs) + | Syntax.TEList exprs => foldr (fn (x, acc) => Syntax.LRecord [elab env x, acc]) (Syntax.LInt 0) exprs - | Syntax.EApp (f, x) => Syntax.LApp (elab env f, elab env x) - | Syntax.ETyped (e, _) => elab env e - | Syntax.EAndAlso (_, _) => raise Fail "unimplemented" - | Syntax.EOrElse (_, _) => raise Fail "unimplemented" - | Syntax.ELet ([], body) => elab env body - | Syntax.ELet (Syntax.DDatatype (name, cons) :: decls, body) => + | Syntax.TEApp (f, x) => Syntax.LApp (elab env f, elab env x) + | Syntax.TEAndAlso (_, _) => raise Fail "unimplemented" + | Syntax.TEOrElse (_, _) => raise Fail "unimplemented" + | Syntax.TELet ([], body) => elab env body + | Syntax.TELet (Syntax.TDDatatype (name, cons) :: decls, body) => let val env = bindDataCons cons env val funs = @@ -363,12 +297,12 @@ struct foldl (fn ((v, x), acc) => Syntax.LApp (Syntax.LFn (v, acc), x)) - (Syntax.LFix (funs, elab env (Syntax.ELet (decls, body)))) + (Syntax.LFix (funs, elab env (Syntax.TELet (decls, body), ty))) vals 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) => + | Syntax.TELet (Syntax.TDVal (pat, v) :: decls, body) => + elab env (Syntax.TECase (v, [(pat, (Syntax.TELet (decls, body), ty))]), ty) + | Syntax.TELet (Syntax.TDValRec ((Syntax.TPVar name, _), f as (Syntax.TELambda _, _)) :: decls, body) => let val n = GenSym.new () val env = bindVar name n env @@ -377,10 +311,10 @@ struct 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.TELet (decls, body), ty)) end - | Syntax.ELet (Syntax.DValRec _ :: _, _) => raise Fail "invalid val rec" - | Syntax.ELet (Syntax.DFun (name, cases) :: decls, body) => + | Syntax.TELet (Syntax.TDValRec _ :: _, _) => raise Fail "invalid val rec" + | Syntax.TELet (Syntax.TDFun (name, cases) :: decls, body) => let val (ps1, _) = hd cases val nPats = length ps1 @@ -391,7 +325,12 @@ struct val temps = List.tabulate (nPats, fn _ => GenSym.new ()) val env = bindVar name n env val t = GenSym.new () - val innerCase = elabCase env (Syntax.LVar t) (map (fn (ps, b) => (Syntax.PTuple ps, b)) cases) + val innerCase = + elabCase env (Syntax.LVar t) + (map + (fn (ps, b) => + ((Syntax.TPTuple ps, Syntax.TTuple (map (fn (_, t) => t) ps)), b)) + cases) in Syntax.LFix ( [ ( n @@ -402,37 +341,44 @@ struct (tl temps) ) ] - , elab env (Syntax.ELet (decls, body)) + , elab env (Syntax.TELet (decls, body), ty) ) end end - | Syntax.ELet (Syntax.DStruct (name, Syntax.SStruct structDecls) :: decls, body) => + | Syntax.TELet (Syntax.TDStruct (name, (Syntax.TSStruct structDecls, _)) :: decls, body) => let val names = structBoundVars structDecls - val tuple = elab env (Syntax.ELet (structDecls, Syntax.ETuple (map (fn n => Syntax.EIdent [n]) names))) + val tuple = elab env (Syntax.TELet (structDecls, (Syntax.TETuple (map (fn n => (Syntax.TEIdent n, Syntax.TTuple [])) names), Syntax.TTuple [] (* TODO *))), Syntax.TTuple [] (* TODO *)) + val v = GenSym.new () + val env = bindVar name v env + in Syntax.LApp (Syntax.LFn (v, elab env (Syntax.TELet (decls, body), ty)), tuple) end + | Syntax.TELet (Syntax.TDStruct (name, (Syntax.TSIdent structName, _)) :: decls, body) => + let val v = GenSym.new () - val env = bindStructType name structDecls env val env = bindVar name v env - in Syntax.LApp (Syntax.LFn (v, elab env (Syntax.ELet (decls, body))), tuple) end - | Syntax.ELet (Syntax.DStruct (name, Syntax.SIdent (structName :: accessors)) :: decls, body) => + in Syntax.LApp (Syntax.LFn (v, elab env (Syntax.TELet (decls, body), ty)), Syntax.LVar (lookupVar structName env)) end + | Syntax.TELet (Syntax.TDStruct (name, (Syntax.TSDot (parent as (_, Syntax.TStruct (fields, _)), field), _)) :: decls, body) => let - val s = Syntax.LVar (lookupVar structName env) - fun go env [] acc = (env, acc) - | go env (accessor :: accessors) acc = - go - (lookupStructType accessor env) - accessors - (Syntax.LSelect (lookupVar accessor env, acc)) - val (structEnv, structExpr) = go (lookupStructType structName env) accessors s val v = GenSym.new () - val env = bindStruct name structEnv env val env = bindVar name v env - in Syntax.LApp (Syntax.LFn (v, elab env (Syntax.ELet (decls, body))), structExpr) end - | Syntax.ELambda body => + val i = + case List.find (fn (_, (x, _)) => x = field) (enumerate fields) of + NONE => raise Fail ("Unknown field " ^ field) + | SOME (i, _) => i + in Syntax.LApp (Syntax.LFn (v, elab env (Syntax.TELet (decls, body), ty)), Syntax.LSelect (i, elabStructExpr env parent)) end + | Syntax.TELambda body => let val v = GenSym.new () in Syntax.LFn (v, elabCase env (Syntax.LVar v) [body]) end - | Syntax.ECase (expr, arms) => + | Syntax.TECase (expr, 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 emptyEnv p + and elabStructExpr (env : env) (Syntax.TSIdent structName, _ : Syntax.structType) : Syntax.lexp = Syntax.LVar (lookupVar structName env) + | elabStructExpr env (Syntax.TSDot (expr as (_, Syntax.TStruct (_, fields)), field), _) = + let val i = + case List.find (fn (_, (x, _)) => x = field) (enumerate fields) of + SOME (i, _) => i + | NONE => raise Fail "unbound field" + in Syntax.LSelect (i, elabStructExpr env expr) end + + fun elaborate (p : Syntax.typedExpr * Syntax.ty) : Syntax.lexp = elab emptyEnv p end |
