summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2024-05-17 12:36:56 -0700
committerRose Hogenson <rosehogenson@posteo.net>2024-05-17 12:36:56 -0700
commitc39c43e144733973ebf9cd6842df75b5a5ceee0b (patch)
tree91712d2b3dd9dfe30806e447830e1f8665286c05
parent8ec4edc977e0ed87f86c7e74c13d3633cb788069 (diff)
downloadsml-c39c43e144733973ebf9cd6842df75b5a5ceee0b.tar.zst
Add datatypes.
-rw-r--r--elab.sml114
-rw-r--r--parser.sml23
-rw-r--r--syntax.sml9
-rw-r--r--tests/16-datatype.sml6
4 files changed, 123 insertions, 29 deletions
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
diff --git a/parser.sml b/parser.sml
index 9325c98..8057873 100644
--- a/parser.sml
+++ b/parser.sml
@@ -335,7 +335,7 @@ struct
bind parseSingleType parseTycons st
and parseTupleType : Syntax.etype parser =
fn st =>
- bind (sepBy1 parseTycon (symbol "*")) (fn types =>
+ bind (sepBy1 parseTycon (reserved "*")) (fn types =>
const
(case types of
[ty] => ty
@@ -357,7 +357,12 @@ struct
[p] => p
| _ => Syntax.PTuple pats))
<|> Syntax.PVar <$> identifier) st
- and pat : Syntax.pat parser = fn st => atpat st
+ and pat : Syntax.pat parser = fn st =>
+ (bind identifier (fn ident =>
+ bind atpat (fn arg =>
+ const (Syntax.PCon (ident, arg)))
+ <|> const (Syntax.PVar ident))
+ <|> atpat) st
fun leftOp (i : int) : Syntax.expr parser =
bind getUserState (fn UserState (_, opTable) =>
@@ -462,7 +467,19 @@ struct
const (Syntax.ECase (e, arms)))))
<|> orelseExpr) st
and dec : Syntax.dec parser = fn st =>
- ((reserved "val" >>
+ ((reserved "datatype" >>
+ bind identifier (fn name =>
+ reserved "=" >>
+ bind
+ (sepBy1
+ (bind identifier (fn con =>
+ (reserved "of" >>
+ bind parseType (fn ty =>
+ const (con, SOME ty)))
+ <|> const (con, NONE)))
+ (reserved "|")) (fn cons =>
+ const (Syntax.DDatatype (name, cons)))))
+ <|> (reserved "val" >>
bind (true <$ reserved "rec" <|> const false) (fn isRec =>
bind pat (fn p =>
reserved "=" >>
diff --git a/syntax.sml b/syntax.sml
index fbb4dfd..767851d 100644
--- a/syntax.sml
+++ b/syntax.sml
@@ -12,6 +12,7 @@ struct
| PVar of string
| PInt of int
| PTuple of pat list
+ | PCon of string * pat
datatype expr =
EIdent of string list
@@ -31,6 +32,7 @@ struct
and dec =
DVal of pat * expr
| DValRec of pat * expr
+ | DDatatype of string * (string * etype option) list
(* Lambda language *)
type var = int
@@ -101,6 +103,11 @@ struct
fun quote (s : string) : string = "\"" ^ String.toString s ^ "\""
+ fun optionToString (show : 'a -> string) (x : 'a option) =
+ case x of
+ NONE => "NONE"
+ | SOME x => "SOME " ^ show x
+
fun etypeToString (x : etype) : string =
case x of
Tyvar s => "Tyvar " ^ quote s
@@ -114,6 +121,7 @@ struct
| PVar v => "PVar " ^ quote v
| PInt i => "PInt " ^ Int.toString i
| PTuple pats => "PTuple " ^ listToString patToString pats
+ | PCon (con, v) => "PCon " ^ "(" ^ quote con ^ ", " ^ patToString v ^ ")"
fun exprToStringI (indent : string) (x : expr) : string =
let val self = exprToStringI indent
@@ -137,6 +145,7 @@ struct
case x of
DVal (p, e) => "DVal (" ^ patToString p ^ ", " ^ exprToStringI indent e ^ ")"
| DValRec (p, e) => "DValRec (" ^ patToString p ^ ", " ^ exprToStringI indent e ^ ")"
+ | DDatatype (name, arms) => "DDatatype (" ^ quote name ^ ", " ^ listToString (fn (con, v) => "(" ^ quote con ^ ", " ^ optionToString etypeToString v ^ ")") arms ^ ")"
val exprToString : expr -> string = exprToStringI ""
diff --git a/tests/16-datatype.sml b/tests/16-datatype.sml
new file mode 100644
index 0000000..ee4dd16
--- /dev/null
+++ b/tests/16-datatype.sml
@@ -0,0 +1,6 @@
+datatype D = A | B of int
+
+val _ =
+ case B 42 of
+ B x => __builtin "exit" x
+ | _ => __builtin "exit" 0