structure Elab = struct 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 "exit" => Syntax.PExit | "add" => Syntax.PAdd | "sub" => Syntax.PSub | "mul" => Syntax.PMul | "div" => Syntax.PDiv | "read" => Syntax.PRead | "write" => Syntax.PWrite | "writeErr" => Syntax.PWriteErr | _ => raise Fail ("invalid op: " ^ s) fun enumerate (l : 'a list) : (int * 'a) list = let fun go _ [] = [] | go i (x :: xs) = (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 | 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 datatype env = Env of int StringMap.map val emptyEnv = Env StringMap.empty 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 env of SOME x => x | NONE => raise Fail ("unbound identifier " ^ name) 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.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.typedPat * Syntax.ty) list) : Syntax.lexp list = let val expandedArms = List.mapPartial (fn (Syntax.TPTuple 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.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.TPCon (_, arg), _) = patternBindings (Syntax.LSelect (1, expr)) arg | 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.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 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.TPCon (name, _), Syntax.TDatatype cons) => lookupCon (List.last name) cons = n | _ => false) (map hd patterns) val newTupleSize = case con1 of SOME (Syntax.TPCon (_, (Syntax.TPTuple t, _)), _) => length t | _ => 0 val occHead = hd occurrences val occRest = tl occurrences val occurrences = if newTupleSize = 0 then Syntax.LSelect (1, occHead) :: occRest else List.tabulate (newTupleSize, fn i => Syntax.LSelect (i, Syntax.LSelect (1, occHead))) @ occRest 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.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 (fn (p, a) => Option.map (fn p => (p, a)) (specializeRow p)) (ListPair.zipEq (patterns, actions))) in (patterns, occurrences, actions) end 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.TPWild, _) :: _, _) => true | ((Syntax.TPVar _, _) :: _, _) => true | _ => false) (ListPair.zipEq (patterns, actions))) in (patterns, occurrences, actions) end (* https://compiler.club/compiling-pattern-matching/ *) fun compilePatternMatching ([] : (Syntax.typedPat * Syntax.ty) 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.TPInt _, _)) => true | (_, (Syntax.TPCon _, _)) => true | _ => false) (enumerate firstRow) in case refutablePattern of NONE => hd actions | SOME (i, _) => let val (patterns, occurrences) = if i = 0 then (patterns, occurrences) else swap i patterns occurrences val firstCol = map hd patterns val signatures = map (fn (x, _) => x) (IntMap.toList (foldl (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.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 (default (patterns, occurrences, actions))) val switchOperand = if nCons < 0 then hd occurrences else Syntax.LSelect (0, hd occurrences) in Syntax.LSwitch ( switchOperand , map (fn i => (i, compilePatternMatching (specialize i (patterns, occurrences, actions)))) signatures , defaultCase ) end end 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 structBoundVars (decls : Syntax.typedDec list) : string list = List.concatMap declBoundVars decls 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 val bindings = patternBindings expr p val env = foldl (fn ((name, _), env) => bindVar name (GenSym.new ()) env) env bindings in foldl (fn ((name, binding), acc) => Syntax.LApp (Syntax.LFn (lookupVar name env, acc), binding)) (elab env body) bindings end) arms 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) 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 : env) (p : Syntax.typedExpr, ty : Syntax.ty) : Syntax.lexp = case p of Syntax.TEIdent i => Syntax.LVar (lookupVar i env) | Syntax.TEDot (structExpr as (_, Syntax.TStruct (_, fields)), field) => let 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.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 = List.mapPartial (fn (_, (_, NONE)) => NONE | (i, (name, _)) => let val v = GenSym.new () in SOME (lookupVar name env, v, Syntax.LRecord [Syntax.LInt i, Syntax.LVar v]) end) (enumerate cons) val vals = List.mapPartial (fn (i, (name, NONE)) => SOME (lookupVar name env, Syntax.LRecord [Syntax.LInt i]) | _ => NONE) (enumerate cons) in foldl (fn ((v, x), acc) => Syntax.LApp (Syntax.LFn (v, acc), x)) (Syntax.LFix (funs, elab env (Syntax.TELet (decls, body), ty))) vals end | 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 val (arg, fnBody) = 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.TELet (decls, body), ty)) end | 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 in if not (List.all (fn (ps, _) => length ps = nPats) cases) then raise Fail "clauses do not all have same number of patterns" else let val n = GenSym.new () 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.TPTuple ps, Syntax.TTuple (map (fn (_, t) => t) ps)), b)) cases) in Syntax.LFix ( [ ( n , hd temps , foldr Syntax.LFn (Syntax.LApp (Syntax.LFn (t, innerCase), Syntax.LRecord (map Syntax.LVar temps))) (tl temps) ) ] , elab env (Syntax.TELet (decls, body), ty) ) end end | Syntax.TELet (Syntax.TDStruct (name, (Syntax.TSStruct structDecls, _)) :: decls, body) => let val names = structBoundVars structDecls 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 = bindVar name v env 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 v = GenSym.new () val env = bindVar name v env 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.TECase (expr, arms) => let val v = GenSym.new () in Syntax.LApp (Syntax.LFn (v, elabCase env (Syntax.LVar v) arms), elab env expr) end 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