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 | _ => raise Fail ("invalid op: " ^ s) 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 => hd actions | 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) | SOME x => Syntax.LVar x) | Syntax.EIdent _ => raise Fail "long identifiers are not supported" | 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 => 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.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 (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.ELet (decls, body))) end | Syntax.ELet (Syntax.DValRec _ :: _, _) => raise Fail "invalid val rec" | Syntax.ELambda body => let val v = Gensym.new () in Syntax.LFn (v, elabCase env (Syntax.LVar v) [body]) end | Syntax.ECase (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 StringMap.empty p end