summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2024-05-18 13:08:31 -0700
committerRose Hogenson <rosehogenson@posteo.net>2024-05-18 13:08:31 -0700
commit2581ac99323ead839964e9b5f2a799cf61fd5bdc (patch)
treea733c96894f4e8fb37dffb964f57fd39d9996d08
parent92baa4f6bf99efe229abf2ab895cec8fd7a8936d (diff)
downloadsml-2581ac99323ead839964e9b5f2a799cf61fd5bdc.tar.zst
Allow pattern matching in function definitions.
-rw-r--r--elab.sml26
-rw-r--r--parser.sml23
-rw-r--r--syntax.sml18
-rw-r--r--tests/18-fun-case.sml5
4 files changed, 56 insertions, 16 deletions
diff --git a/elab.sml b/elab.sml
index eb168c5..2895894 100644
--- a/elab.sml
+++ b/elab.sml
@@ -330,6 +330,32 @@ struct
Syntax.LFix ([(n, arg, fnBody)], elab env (Syntax.ELet (decls, body)))
end
| Syntax.ELet (Syntax.DValRec _ :: _, _) => raise Fail "invalid val rec"
+ | Syntax.ELet (Syntax.DFun (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.PTuple 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.ELet (decls, body))
+ )
+ end
+ end
| Syntax.ELambda body =>
let val v = Gensym.new ()
in Syntax.LFn (v, elabCase env (Syntax.LVar v) [body])
diff --git a/parser.sml b/parser.sml
index 8057873..08cb992 100644
--- a/parser.sml
+++ b/parser.sml
@@ -489,15 +489,20 @@ struct
then Syntax.DValRec (p, e)
else Syntax.DVal (p, e))))))
<|> (reserved "fun" >>
- bind identifier (fn name =>
- bind (many1 atpat) (fn args =>
- reserved "=" >>
- bind expr (fn body =>
- const
- (Syntax.DValRec
- ( Syntax.PVar name
- , foldr Syntax.ELambda body args
- ))))))) st
+ bind
+ (sepBy1
+ (bind identifier (fn name =>
+ bind (many1 atpat) (fn args =>
+ reserved "=" >>
+ bind expr (fn body =>
+ const (name, args, body)))))
+ (reserved "|")) (fn cases =>
+ let val (name, _, _) = hd cases
+ in
+ if not (List.all (fn (n, _, _) => n = name) cases)
+ then raise Fail "clauses do not all have same function name"
+ else const (Syntax.DFun (name, map (fn (_, x, y) => (x, y)) cases))
+ end))) st
val program : Syntax.expr parser = (fn decs => Syntax.ELet (decs, Syntax.EInt 0)) <$> many dec
fun parse (f : string) : (string, Syntax.expr) Result.either = runParser program f
diff --git a/syntax.sml b/syntax.sml
index aa94b09..5672a5b 100644
--- a/syntax.sml
+++ b/syntax.sml
@@ -32,6 +32,7 @@ struct
and dec =
DVal of pat * expr
| DValRec of pat * expr
+ | DFun of string * (pat list * expr) list
| DDatatype of string * (string * etype option) list
(* Lambda language *)
@@ -145,6 +146,7 @@ struct
case x of
DVal (p, e) => "DVal (" ^ patToString p ^ ", " ^ exprToStringI indent e ^ ")"
| DValRec (p, e) => "DValRec (" ^ patToString p ^ ", " ^ exprToStringI indent e ^ ")"
+ | DFun (name, cases) => "DFun (" ^ quote name ^ ", " ^ multilineListToString (fn indent => fn (ps, b) => "(" ^ listToString patToString ps ^ ", " ^ exprToStringI indent b ^ ")") indent cases ^ ")"
| DDatatype (name, arms) => "DDatatype (" ^ quote name ^ ", " ^ listToString (fn (con, v) => "(" ^ quote con ^ ", " ^ optionToString etypeToString v ^ ")") arms ^ ")"
val exprToString : expr -> string = exprToStringI ""
@@ -162,18 +164,20 @@ struct
| PEq => "PEq"
| PIf => "PIf"
- fun lexpToString (x : lexp) : string =
+ fun lexpToStringI (indent : string) (x : lexp) : string =
case x of
LVar v => "LVar " ^ Int.toString v
- | LFn (arg, expr) => "LFun (" ^ Int.toString arg ^ ", " ^ lexpToString expr ^ ")"
- | LFix (decls, body) => "LFix (" ^ listToString (fn (arg, var, expr) => "(" ^ Int.toString arg ^ ", " ^ Int.toString var ^ ", " ^ lexpToString expr ^ ")") decls ^ ", " ^ lexpToString body ^ ")"
- | LApp (a, b) => "LApp (" ^ lexpToString a ^ ", " ^ lexpToString b ^ ")"
+ | LFn (arg, expr) => "LFun (" ^ Int.toString arg ^ ",\n" ^ indent ^ "\t" ^ lexpToStringI (indent ^ "\t") expr ^ ")"
+ | LFix (decls, body) => "LFix (" ^ multilineListToString (fn indent => fn (arg, var, expr) => "(" ^ Int.toString arg ^ ", " ^ Int.toString var ^ ", " ^ lexpToStringI indent expr ^ ")") indent decls ^ ",\n" ^ indent ^ lexpToStringI indent body ^ ")"
+ | LApp (a, b) => "LApp (" ^ lexpToStringI indent a ^ ",\n" ^ indent ^ "\t" ^ lexpToStringI (indent ^ "\t") b ^ ")"
| LInt i => "LInt " ^ Int.toString i
| LString s => "LString " ^ quote s
- | LRecord l => "LRecord " ^ listToString lexpToString l
- | LSelect (i, r) => "LSelect (" ^ Int.toString i ^ ", " ^ lexpToString r ^ ")"
+ | LRecord l => "LRecord " ^ listToString (lexpToStringI indent) l
+ | LSelect (i, r) => "LSelect (" ^ Int.toString i ^ ", " ^ lexpToStringI indent r ^ ")"
| LPrim p => "LPrim " ^ primopToString p
- | LSwitch (e, arms, otherwise) => "LSwitch (" ^ lexpToString e ^ ", " ^ listToString (fn (x, e) => "(" ^ Int.toString x ^ ", " ^ lexpToString e ^ ")") arms ^ ", " ^ optionToString lexpToString otherwise ^ ")"
+ | LSwitch (e, arms, otherwise) => "LSwitch (" ^ lexpToStringI indent e ^ ",\n" ^ indent ^ "\t" ^ multilineListToString (fn indent => fn (x, e) => "(" ^ Int.toString x ^ ", " ^ lexpToStringI indent e ^ ")") (indent ^ "\t") arms ^ ",\n" ^ indent ^ "\t" ^ optionToString (lexpToStringI (indent ^ "\t")) otherwise ^ ")"
+
+ fun lexpToString (x : lexp) : string = lexpToStringI "" x
fun valueToString (x : value) : string =
case x of
diff --git a/tests/18-fun-case.sml b/tests/18-fun-case.sml
new file mode 100644
index 0000000..d64b7d7
--- /dev/null
+++ b/tests/18-fun-case.sml
@@ -0,0 +1,5 @@
+fun fib 1 = 1
+ | fib 2 = 2
+ | fib n = __builtin "add" (fib (__builtin "sub" (n, 1)), fib (__builtin "sub" (n, 2)))
+
+val _ = __builtin "exit" (__builtin "add" (fib 8, 8))