diff options
| author | Rose Hogenson <rosehogenson@posteo.net> | 2025-05-18 08:17:20 -0700 |
|---|---|---|
| committer | Rose Hogenson <rosehogenson@posteo.net> | 2025-06-13 19:14:58 -0700 |
| commit | 8b3a9b8f0d80e7dd789f363deb6bb36189a31f01 (patch) | |
| tree | 414ac8c234ded42f1bf4d8da7a96a4c7078db95f /Types.sml | |
| parent | cc3f514265fc7c4166b8e30fcdbb26e69286f339 (diff) | |
| download | sml-8b3a9b8f0d80e7dd789f363deb6bb36189a31f01.tar.zst | |
Add a typechecker
Diffstat (limited to 'Types.sml')
| -rw-r--r-- | Types.sml | 464 |
1 files changed, 464 insertions, 0 deletions
diff --git a/Types.sml b/Types.sml new file mode 100644 index 0000000..2722063 --- /dev/null +++ b/Types.sml @@ -0,0 +1,464 @@ +structure Types = struct + structure StringMap = Map(type k = string val cmp = String.compare) + structure IntMap = Map(type k = int val cmp = Int.compare) + + type tyvar = int + datatype tycon = Bool | Int | Str | Fun | Tuple | List | Datatype of int + and ty = + TyVar of tyvar + | TyCon of tycon * ty list + | TyStruct of structTy + and structTy = Struct of { + structs : structTy StringMap.map, + vals : ty StringMap.map + } + + structure TyVarMap = IntMap + + fun listToString (show : 'a -> string) (l : 'a list) : string = + "[" ^ String.concatWith ", " (map show l) ^ "]" + + fun printTyCon Int : string = "Int" + | printTyCon Fun = "Fun" + | printTyCon Tuple = "Tuple" + | printTyCon (Datatype tag) = "Datatype " ^ Int.toString tag + + and printTy (TyVar a) : string = "TyVar " ^ Int.toString a + | printTy (TyCon (n, tys)) = "TyCon (" ^ printTyCon n ^ ", " ^ listToString printTy tys ^ ")" + + (* map of type variables to types *) + val substitution : ty TyVarMap.map ref = ref TyVarMap.empty + + (* map of datatype tags to type description *) + val datatypes : string list IntMap.map ref = ref IntMap.empty + + fun find (TyVar a) : ty = + let val representative = + case getOpt (TyVarMap.lookup a (!substitution), TyVar a) of + TyVar b => if a = b then TyVar a else find (TyVar b) + | t => find t + in + substitution := TyVarMap.insert a representative (!substitution) ; + representative + end + | find (TyCon (n, ts)) = TyCon (n, map find ts) + | find (TyStruct structType) = TyStruct structType + + fun unify (t1 : ty) (t2 : ty) : unit = + case (find t1, find t2) of + (TyVar a, t2) => substitution := TyVarMap.insert a t2 (!substitution) + | (t1, TyVar a) => substitution := TyVarMap.insert a t1 (!substitution) + | (TyCon (n1, tys1), TyCon (n2, tys2)) => + if n1 <> n2 orelse length tys1 <> length tys2 then raise Fail ("attempt to unify distinct types " ^ printTy (TyCon (n1, tys1)) ^ " and " ^ printTy (TyCon (n2, tys2))) else + ListPair.appEq (fn (t1, t2) => unify t1 t2) (tys1, tys2) + | _ => raise Fail "unify invalid types" + + datatype binding = Let of tyvar | Arg of tyvar + + type env = { + bindings : binding StringMap.map, + boundVars : unit TyVarMap.map, + typesByName : tycon StringMap.map, + structs : structTy StringMap.map + } + + fun bind (makeBinding : tyvar -> binding) (s : string) (v : tyvar) ({bindings, boundVars, typesByName, structs} : env) : env = { + bindings = StringMap.insert s (makeBinding v) bindings, + boundVars = TyVarMap.insert v () boundVars, + typesByName = typesByName, + structs = structs + } + + val userTypeVariables : tyvar StringMap.map ref = ref StringMap.empty + + fun etypeToTy (_ : env) (Syntax.Tyvar "int") = TyCon (Int, []) + | etypeToTy _ (Syntax.Tyvar "bool") = TyCon (Bool, []) + | etypeToTy _ (Syntax.Tyvar "string") = TyCon (Str, []) + | etypeToTy env (Syntax.Tyvar v) = + if String.isPrefix "'" v then + case StringMap.lookup v (!userTypeVariables) of + SOME v => TyVar v + | NONE => + let val t = GenSym.new () in + userTypeVariables := StringMap.insert v t (!userTypeVariables); + TyVar t + end + else + (case StringMap.lookup v (#typesByName env) of + SOME t => TyCon (t, []) + | NONE => raise Fail "unknown type") + | etypeToTy env (Syntax.Tycon (tys, ty)) = + (case StringMap.lookup ty (#typesByName env) of + SOME t => TyCon (t, map (etypeToTy env) tys) + | NONE => raise Fail ("unknown type " ^ ty)) + | etypeToTy env (Syntax.TyTuple tys) = TyCon (Tuple, map (etypeToTy env) tys) + | etypeToTy env (Syntax.Tyfun (arg, result)) = TyCon (Fun, [etypeToTy env arg, etypeToTy env result]) + + fun bindType ((vars, name, constructors) : string list * string * (string * Syntax.etype option) list) (env as {bindings, boundVars, typesByName, structs} : env) : env = + let + val tag = GenSym.new () + val env = { + bindings = bindings, + boundVars = boundVars, + typesByName = StringMap.insert name (Datatype tag) typesByName, + structs = structs + } + in + datatypes := IntMap.insert tag (map (fn (x, _) => x) constructors) (!datatypes) ; + foldl + (fn ((x, t), env) => + let + val g = GenSym.new () + val resType = TyCon (Datatype tag, map (etypeToTy env o Syntax.Tyvar) vars) + val conType = + case t of + NONE => resType + | SOME t => TyCon (Fun, [etypeToTy env t, resType]) + in + unify (TyVar g) conType; + bind Let x g env + end) + env + constructors + end + + fun bindStruct (name : string) (str : structTy) ({bindings, boundVars, typesByName, structs} : env) : env = { + bindings = bindings, + boundVars = boundVars, + typesByName = typesByName, + structs = StringMap.insert name str structs + } + + fun bindPat (makeBinding : tyvar -> binding) (Syntax.PVar v) (env : env) : env = + bind makeBinding v (GenSym.new ()) env + | bindPat makeBinding (Syntax.PTuple pats) env = + foldl (fn (pat, env) => bindPat makeBinding pat env) env pats + | bindPat makeBinding (Syntax.PCon (_, pat)) env = bindPat makeBinding pat env + | bindPat _ _ env = env + + fun genericVars (env : env) (TyVar a) : unit TyVarMap.map = + if isSome (TyVarMap.lookup a (#boundVars env)) + then TyVarMap.empty + else TyVarMap.fromList [(a, ())] + | genericVars env (TyCon (_, tys)) = + foldl (fn (x, acc) => TyVarMap.union acc (genericVars env x)) TyVarMap.empty tys + | genericVars _ (TyStruct _) = TyVarMap.empty + + fun instantiate (env : env) (t : ty) : ty = + let + val genericVars = genericVars env t + val newVars = TyVarMap.fromList (map (fn (a, _) => (a, GenSym.new ())) (TyVarMap.toList genericVars)) + fun replace (TyVar a) = + TyVar (getOpt (TyVarMap.lookup a newVars, a)) + | replace (TyCon (n, tys)) = TyCon (n, map replace tys) + | replace (TyStruct structType) = TyStruct structType + in replace t end + + fun lookupBinding (name : string) (env : env) : ty = + case StringMap.lookup name (#bindings env) of + SOME (Let t) => instantiate env (find (TyVar t)) + | SOME (Arg t) => TyVar t + | NONE => raise Fail ("unbound var " ^ name) + + fun patBindings (Syntax.PVar v) : string list = [v] + | patBindings (Syntax.PTuple pats) = List.concat (map patBindings pats) + | patBindings (Syntax.PCon (_, pat)) = patBindings pat + | patBindings _ = [] + + fun tyToSyntaxTy (TyVar v) = Syntax.TVar v + | tyToSyntaxTy (TyCon (Bool, [])) = Syntax.TBool + | tyToSyntaxTy (TyCon (Int, [])) = Syntax.TInt + | tyToSyntaxTy (TyCon (Str, [])) = Syntax.TString + | tyToSyntaxTy (TyCon (Fun, [f, x])) = Syntax.TFun (tyToSyntaxTy f, tyToSyntaxTy x) + | tyToSyntaxTy (TyCon (Tuple, ts)) = Syntax.TTuple (map tyToSyntaxTy ts) + | tyToSyntaxTy (TyCon (List, [t])) = Syntax.TList (tyToSyntaxTy t) + | tyToSyntaxTy (TyCon (Datatype tag, _)) = + (case (IntMap.lookup tag (!datatypes)) of + SOME cases => Syntax.TDatatype cases + | NONE => raise Fail ("unknown datatype tag " ^ Int.toString tag)) + | tyToSyntaxTy t = raise Fail ("invalid ty " ^ printTy t) + + fun structTypeToSyntaxStructType (Struct s) = + Syntax.TStruct + ( map (fn (name, ty) => (name, structTypeToSyntaxStructType ty)) (StringMap.toList (#structs s)) + , map (fn (name, ty) => (name, tyToSyntaxTy ty)) (StringMap.toList (#vals s)) + ) + (* | structTypeToSyntaxStructType _ = raise Fail "we really need to combine structs and expressions" *) + + fun tagPat (_ : env) Syntax.PWild : Syntax.typedPat * ty = (Syntax.TPWild, TyVar (GenSym.new ())) + | tagPat env (Syntax.PVar v) = (Syntax.TPVar v, TyVar (case (StringMap.lookup v (#bindings env)) of SOME (Let t) => t | SOME (Arg t) => t | NONE => raise Fail "unknown pattern variable")) + | tagPat _ (Syntax.PInt i) = (Syntax.TPInt i, TyCon (Int, [])) + | tagPat env (Syntax.PTuple pats) = + let + val taggedPats = map (tagPat env) pats + val types = map (fn (_, ty) => ty) taggedPats + val taggedPats = map (fn (p, ty) => (p, tyToSyntaxTy ty)) taggedPats + in (Syntax.TPTuple taggedPats, TyCon (Tuple, types)) end + | tagPat env (Syntax.PCon (con, pat)) = + let + val conBinding = + case con of + [ident] => lookupBinding ident env + | str :: fields => + case StringMap.lookup str (#structs env) of + SOME (Struct structTy) => + let + fun go structTy [field] = + (case StringMap.lookup field (#vals structTy) of + SOME t => instantiate env (find t) + | NONE => raise Fail "unbound something or other") + | go structTy (str :: fields) = + case StringMap.lookup str (#structs structTy) of + SOME (Struct structTy) => go structTy fields + | NONE => raise Fail ("unbound field " ^ str) + in go structTy fields end + | NONE => raise Fail ("unbound struct " ^ str) + in + case conBinding of + TyCon (Fun, [argType, resType]) => + let val (pat, patType) = tagPat env pat + in + unify patType argType ; + (Syntax.TPCon (con, (pat, tyToSyntaxTy patType)), resType) + end + | conType => + case pat of + Syntax.PTuple [] => (Syntax.TPCon (con, (Syntax.TPTuple [], Syntax.TTuple [])), conType) + | _ => raise Fail ("non-function " ^ (String.concatWith "." con) ^ " applied to argument in pattern") + end + + fun F (env : env) (Syntax.EIdent i : Syntax.expr) : Syntax.typedExpr * ty = + (Syntax.TEIdent i, lookupBinding i env) + | F env (Syntax.EDot (expr, field)) = + let + val (str, Struct structType) = tagStructExpr env expr + val fieldType = instantiate env (find (valOf (StringMap.lookup field (#vals structType)))) + in (Syntax.TEDot ((str, structTypeToSyntaxStructType (Struct structType)), field), fieldType) end + | F _ (Syntax.EBuiltin "exit") = (Syntax.TEBuiltin "exit", TyCon (Fun, [TyCon (Int, []), TyVar (GenSym.new ())])) + | F _ (Syntax.EBuiltin "add") = (Syntax.TEBuiltin "add", TyCon (Fun, [TyCon (Tuple, [TyCon (Int, []), TyCon (Int, [])]), TyCon (Int, [])])) + | F _ (Syntax.EBuiltin "sub") = (Syntax.TEBuiltin "sub", TyCon (Fun, [TyCon (Tuple, [TyCon (Int, []), TyCon (Int, [])]), TyCon (Int, [])])) + | F _ (Syntax.EBuiltin "mul") = (Syntax.TEBuiltin "mul", TyCon (Fun, [TyCon (Tuple, [TyCon (Int, []), TyCon (Int, [])]), TyCon (Int, [])])) + | F _ (Syntax.EBuiltin "div") = (Syntax.TEBuiltin "div", TyCon (Fun, [TyCon (Tuple, [TyCon (Int, []), TyCon (Int, [])]), TyCon (Int, [])])) + | F _ (Syntax.EBuiltin b) = raise Fail ("unknown builtin " ^ b) + | F _ (Syntax.EInt i) = (Syntax.TEInt i, TyCon (Int, [])) + | F _ (Syntax.EStr s) = (Syntax.TEStr s, TyCon (Str, [])) + | F env (Syntax.ETuple exprs) = + let + val taggedExprs = map (F env) exprs + val types = map (fn (_, x) => x) taggedExprs + in (Syntax.TETuple (map (fn (e, t) => (e, tyToSyntaxTy t)) taggedExprs), TyCon (Tuple, types)) end + | F env (Syntax.EList exprs) = + let + val taggedExprs = map (fn expr => F env expr) exprs + val alpha = TyVar (GenSym.new ()) + val _ = app (fn (_, t) => unify alpha t) taggedExprs + in (Syntax.TEList (map (fn (e, t) => (e, tyToSyntaxTy t)) taggedExprs), TyCon (List, [alpha])) end + | F env (Syntax.EApp (func, arg)) = + let + val (func, funcType) = F env func + val (arg, argType) = F env arg + val beta = TyVar (GenSym.new ()) + val _ = unify funcType (TyCon (Fun, [argType, beta])) + in (Syntax.TEApp ((func, tyToSyntaxTy funcType), (arg, tyToSyntaxTy argType)), beta) end + | F env (Syntax.ETyped (expr, _)) = F env expr + | F env (Syntax.EAndAlso (e1, e2)) = + let + val (e1, e1Type) = F env e1 + val (e2, e2Type) = F env e2 + in + unify e1Type (TyCon (Bool, [])) ; + unify e2Type (TyCon (Bool, [])) ; + (Syntax.TEAndAlso ((e1, tyToSyntaxTy e1Type), (e2, tyToSyntaxTy e2Type)), TyCon (Bool, [])) + end + | F env (Syntax.EOrElse (e1, e2)) = + let + val (e1, e1Type) = F env e1 + val (e2, e2Type) = F env e2 + in + unify e1Type (TyCon (Bool, [])) ; + unify e2Type (TyCon (Bool, [])) ; + (Syntax.TEOrElse ((e1, tyToSyntaxTy e1Type), (e2, tyToSyntaxTy e2Type)), TyCon (Bool, [])) + end + | F env (Syntax.ELet (decs, body)) = + let + val (env, decs) = tagDecs env decs + val (body, bodyType) = F env body + in (Syntax.TELet (decs, (body, tyToSyntaxTy bodyType)), bodyType) + end + | F env (Syntax.ELambda (pat, body)) = + let + val env = bindPat Arg pat env + val (pat, patType) = tagPat env pat + val (body, bodyType) = F env body + in (Syntax.TELambda ((pat, tyToSyntaxTy patType), (body, tyToSyntaxTy bodyType)), TyCon (Fun, [patType, bodyType])) end + | F env (Syntax.ECase (arg, arms)) = + let + val (arg, argType) = F env arg + val resultType = TyVar (GenSym.new ()) + val arms = + map + (fn (pat, body) => + let + val env = bindPat Arg pat env + val (pat, patType) = tagPat env pat + val (body, bodyType) = F env body + in + unify patType argType; + unify bodyType resultType; + ((pat, tyToSyntaxTy patType), (body, tyToSyntaxTy bodyType)) + end) + arms + in (Syntax.TECase ((arg, tyToSyntaxTy argType), arms), resultType) end + + and tagDec (env : env) (Syntax.DVal (pat, expr)) : env * Syntax.typedDec option = + let + val (expr, exprType) = F env expr + val env = bindPat Let pat env + val (pat, patType) = tagPat env pat + in + unify patType exprType ; + (env, SOME (Syntax.TDVal ((pat, tyToSyntaxTy patType), (expr, tyToSyntaxTy exprType)))) + end + | tagDec env (Syntax.DValRec (pat, expr)) = + let + val recEnv = bindPat Arg pat env + val (expr, exprType) = F recEnv expr + val (taggedPat, patType) = tagPat recEnv pat + in + unify patType exprType ; + (bindPat Let pat env, SOME (Syntax.TDValRec ((taggedPat, tyToSyntaxTy patType), (expr, tyToSyntaxTy exprType)))) + end + | tagDec env (Syntax.DFun (name, (args, body) :: cases)) = + let + val fnType = GenSym.new () + val args1Env = + foldl + (fn (arg, env) => bindPat Arg arg env) + (bind Arg name fnType env) + args + val args = map (tagPat args1Env) args + val argTypes = map (fn (_, ty) => ty) args + val (body, bodyType) = F args1Env body + val cases = + (map (fn (arg, ty) => (arg, tyToSyntaxTy ty)) args, (body, tyToSyntaxTy bodyType)) + :: map + (fn (args, body) => + let + val env = + foldl + (fn (arg, env) => bindPat Arg arg env) + (bind Arg name fnType env) + args + val args = map (tagPat env) args + val (body, bt) = F env body + in + ListPair.appEq + (fn ((_, myArgType), argType) => unify myArgType argType) + (args, argTypes) ; + unify bodyType bt ; + (map (fn (arg, ty) => (arg, tyToSyntaxTy ty)) args, (body, tyToSyntaxTy bt)) + end) + cases + in (bind Let name fnType env, SOME (Syntax.TDFun (name, cases))) + end + | tagDec env (Syntax.DDatatype (vars, name, data)) = + let val env = bindType (vars, name, data) env + in (env, SOME (Syntax.TDDatatype (name, map (fn (name, ty) => (name, Option.map (tyToSyntaxTy o (etypeToTy env)) ty)) data))) end + | tagDec _ (Syntax.DType _) = raise Fail "TODO" + | tagDec env (Syntax.DStruct (name, str)) = + let val (str, strType) = tagStructExpr env str + in (bindStruct name strType env, SOME (Syntax.TDStruct (name, (str, structTypeToSyntaxStructType strType)))) end + | tagDec _ _ = raise Fail "invalid expr" + + and tagDecs env [] = (env, []) + | tagDecs env (dec :: decs) = + let + val (env, dec) = tagDec env dec + val (env, decs) = tagDecs env decs + val decs = + case dec of + NONE => decs + | SOME dec => dec :: decs + in (env, decs) + end + + and tagStructExpr (env : env) (Syntax.SIdent s) : Syntax.typedStructExpr * structTy = + (case StringMap.lookup s (#structs env) of + SOME t => (Syntax.TSIdent s, t) + | NONE => raise Fail ("unknown struct type " ^ s)) + | tagStructExpr env (Syntax.SDot (sExpr, field)) = + let + val (parent, Struct parentType) = tagStructExpr env sExpr + val fieldType = + case StringMap.lookup field (#structs parentType) of + SOME x => x + | NONE => raise Fail ("unknown struct field " ^ field) + in (Syntax.TSDot ((parent, structTypeToSyntaxStructType (Struct parentType)), field), fieldType) end + | tagStructExpr env (Syntax.SStruct decls) = + let val (env, decls) = tagDecs env decls + in + (Syntax.TSStruct decls, Struct { + structs = #structs env, + vals = + StringMap.fromList + (map + (fn (name, Let v) => (name, TyVar v) + | (name, Arg v) => (name, TyVar v)) + (StringMap.toList (#bindings env))) + }) + end + + fun reexpandType (Syntax.TVar i) : Syntax.ty = tyToSyntaxTy (find (TyVar i)) + | reexpandType ty = ty + + fun reexpandPat (pat : Syntax.typedPat, ty : Syntax.ty) : Syntax.typedPat * Syntax.ty = + let + val pat = + case pat of + Syntax.TPTuple pats => Syntax.TPTuple (map reexpandPat pats) + | Syntax.TPCon (con, arg) => Syntax.TPCon (con, reexpandPat arg) + | _ => pat + in (pat, reexpandType ty) end + + fun reexpand (expr : Syntax.typedExpr, ty : Syntax.ty) : Syntax.typedExpr * Syntax.ty = + let + val expr = + case expr of + Syntax.TEDot (str, field) => Syntax.TEDot (reexpandStructExpr str, field) + | Syntax.TETuple exprs => Syntax.TETuple (map reexpand exprs) + | Syntax.TEList exprs => Syntax.TEList (map reexpand exprs) + | Syntax.TEApp (f, arg) => Syntax.TEApp (reexpand f, reexpand arg) + | Syntax.TEAndAlso (expr1, expr2) => Syntax.TEAndAlso (reexpand expr1, reexpand expr2) + | Syntax.TEOrElse (expr1, expr2) => Syntax.TEOrElse (reexpand expr1, reexpand expr2) + | Syntax.TELet (decls, body) => Syntax.TELet (map reexpandDecl decls, reexpand body) + | Syntax.TELambda (arg, body) => Syntax.TELambda (reexpandPat arg, reexpand body) + | Syntax.TECase (arg, arms) => Syntax.TECase (reexpand arg, map (fn (pat, body) => (reexpandPat pat, reexpand body)) arms) + | _ => expr + in (expr, reexpandType ty) end + + and reexpandDecl (Syntax.TDVal (pat, expr)) : Syntax.typedDec = Syntax.TDVal (reexpandPat pat, reexpand expr) + | reexpandDecl (Syntax.TDValRec (pat, expr)) = Syntax.TDValRec (reexpandPat pat, reexpand expr) + | reexpandDecl (Syntax.TDFun (name, arms)) = Syntax.TDFun (name, map (fn (pats, body) => (map reexpandPat pats, reexpand body)) arms) + | reexpandDecl (decl as Syntax.TDDatatype _) = decl + | reexpandDecl (Syntax.TDStruct (name, str)) = Syntax.TDStruct (name, reexpandStructExpr str) + + and reexpandStructExpr (str : Syntax.typedStructExpr, structType : Syntax.structType) : Syntax.typedStructExpr * Syntax.structType = + let + val str = + case str of + Syntax.TSIdent _ => str + | Syntax.TSDot (str, field) => Syntax.TSDot (reexpandStructExpr str, field) + | Syntax.TSStruct decls => Syntax.TSStruct (map reexpandDecl decls) + in (str, structType) end + + fun tag (expr : Syntax.expr) : Syntax.typedExpr * Syntax.ty = + let + val env = { + bindings = StringMap.empty, + boundVars = TyVarMap.empty, + typesByName = StringMap.empty, + structs = StringMap.empty + } + val (expr, ty) = F env expr + in reexpand (expr, tyToSyntaxTy ty) end +end |
