summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2024-05-05 21:12:16 -0700
committerRose Hogenson <rosehogenson@posteo.net>2024-05-05 21:12:16 -0700
commit457c482397e2ca745067e5a3ef2c7d71d7daf187 (patch)
treeaef2b8d88421829105a3c6b8887d8ffde39d5e1e
parent048895eb0c85b591b2dd338dd06514bc1b18bacb (diff)
downloadsml-457c482397e2ca745067e5a3ef2c7d71d7daf187.tar.zst
Add tuples.
-rw-r--r--buffer.sml2
-rw-r--r--cps.sml4
-rw-r--r--elab.sml204
-rw-r--r--format.sml7
-rw-r--r--main.sml1
-rw-r--r--parser.sml6
-rw-r--r--syntax.sml3
-rw-r--r--tests/15-tuple.sml3
8 files changed, 198 insertions, 32 deletions
diff --git a/buffer.sml b/buffer.sml
index 19cd99a..054f29c 100644
--- a/buffer.sml
+++ b/buffer.sml
@@ -9,7 +9,7 @@ struct
if i + aLen + bLen <= baseLen
then
(Word8ArraySlice.copy { src = b, dst = base, di = i + aLen } ;
- Word8ArraySlice.slice (base, i, SOME (aLen + bLen)))
+ Word8ArraySlice.slice (base, i, SOME (aLen + bLen)))
else
let val newBuf = Word8Array.array (baseLen * 2 + bLen, Word8.fromInt 0)
in
diff --git a/cps.sml b/cps.sml
index cd61a7a..894c57e 100644
--- a/cps.sml
+++ b/cps.sml
@@ -47,6 +47,10 @@ struct
| Syntax.LInt i => cont (Syntax.VInt i)
| Syntax.LString s => cont (Syntax.VString s)
| Syntax.LRecord [] => cont (Syntax.VInt 0)
+ | Syntax.LSelect (i, expr) =>
+ let val temp = Gensym.new ()
+ in toCPS expr (fn x => Syntax.CSelect (i, x, temp, cont (Syntax.VVar temp)))
+ end
| Syntax.LRecord exprs =>
let
fun go [] vars =
diff --git a/elab.sml b/elab.sml
index af6bec2..784515d 100644
--- a/elab.sml
+++ b/elab.sml
@@ -1,6 +1,7 @@
structure Elab =
struct
- structure StringMap = Map(type k = string val cmp = String.compare);
+ 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
@@ -11,7 +12,181 @@ struct
| "div" => Syntax.PDiv
| _ => raise Fail ("invalid op: " ^ s)
- fun elab (env : int StringMap.map) (p : Syntax.expr) : Syntax.lexp =
+ 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 =>
+ if null rows
+ then hd actions
+ else raise Fail "redundant match"
+ | 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)
@@ -46,28 +221,13 @@ struct
Syntax.LFix ([(n, arg, fnBody)], elab env' (Syntax.ELet (decls, body)))
end
| Syntax.ELet (Syntax.DValRec _ :: _, _) => raise Fail "invalid val rec"
- | Syntax.ELambda (pat, body) =>
- let
- val v = Gensym.new ()
- val env' =
- case pat of
- Syntax.PVar name => StringMap.insert name v env
- | _ => env
- in Syntax.LFn (v, elab env' body)
+ | Syntax.ELambda body =>
+ let val v = Gensym.new ()
+ in Syntax.LFn (v, elabCase env (Syntax.LVar v) [body])
end
- | Syntax.ECase (expr, []) => raise Fail "nonexhaustive match"
| Syntax.ECase (expr, arms) =>
- let
- fun go [] _ = raise Fail "nonexhaustive match"
- | go ((Syntax.PWild, body) :: []) acc = Syntax.LSwitch (elab env expr, rev acc, elab env body)
- | go ((Syntax.PVar name, body) :: []) acc =
- let val v = Gensym.new () in
- Syntax.LApp (Syntax.LFn (v, Syntax.LSwitch (Syntax.LVar v, rev acc, elab (StringMap.insert name v env) body)), elab env expr)
- end
- | go ((Syntax.PInt i, body) :: rest) acc =
- go rest ((i, elab env body) :: acc)
- | go _ _ = raise Fail "redundant match"
- in go 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
diff --git a/format.sml b/format.sml
deleted file mode 100644
index 98e8230..0000000
--- a/format.sml
+++ /dev/null
@@ -1,7 +0,0 @@
-structure Format =
-struct
- fun listToString (show : 'a -> string) (l : 'a list) =
- "[" ^ String.concatWith ", " (map show l) ^ "]"
-
- fun pairToString showA showB (a, b) = "(" ^ showA a ^ ", " ^ showB b ^ ")"
-end
diff --git a/main.sml b/main.sml
index 8a01ec4..762fd76 100644
--- a/main.sml
+++ b/main.sml
@@ -1,5 +1,4 @@
use "sort.sml";
-use "format.sml";
use "result.sml";
use "buffer.sml";
use "gensym.sml";
diff --git a/parser.sml b/parser.sml
index c1432f2..9325c98 100644
--- a/parser.sml
+++ b/parser.sml
@@ -351,7 +351,11 @@ struct
val rec atpat : Syntax.pat parser = fn st =>
(Syntax.PWild <$ reserved "_"
<|> Syntax.PInt <$> integer
- <|> between (symbol "(") (symbol ")") pat
+ <|> bind (between (symbol "(") (symbol ")") (sepBy pat (symbol ","))) (fn pats =>
+ const
+ (case pats of
+ [p] => p
+ | _ => Syntax.PTuple pats))
<|> Syntax.PVar <$> identifier) st
and pat : Syntax.pat parser = fn st => atpat st
diff --git a/syntax.sml b/syntax.sml
index 5160e38..fbb4dfd 100644
--- a/syntax.sml
+++ b/syntax.sml
@@ -53,6 +53,7 @@ struct
| LInt of int
| LString of string
| LRecord of lexp list
+ | LSelect of int * lexp
| LPrim of primop
| LSwitch of lexp * (int * lexp) list * lexp
@@ -112,6 +113,7 @@ struct
PWild => "PWild"
| PVar v => "PVar " ^ quote v
| PInt i => "PInt " ^ Int.toString i
+ | PTuple pats => "PTuple " ^ listToString patToString pats
fun exprToStringI (indent : string) (x : expr) : string =
let val self = exprToStringI indent
@@ -160,6 +162,7 @@ struct
| 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 ^ ")"
| LPrim p => "LPrim " ^ primopToString p
| LSwitch (e, arms, otherwise) => "LSwitch (" ^ lexpToString e ^ ", " ^ listToString (fn (x, e) => "(" ^ Int.toString x ^ ", " ^ lexpToString e ^ ")") arms ^ ", " ^ lexpToString otherwise ^ ")"
diff --git a/tests/15-tuple.sml b/tests/15-tuple.sml
new file mode 100644
index 0000000..114fd25
--- /dev/null
+++ b/tests/15-tuple.sml
@@ -0,0 +1,3 @@
+fun f (x, y) = __builtin "exit" (__builtin "add" (x, y))
+
+val _ = f (40, 2)