summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2024-04-28 15:06:43 -0700
committerRose Hogenson <rosehogenson@posteo.net>2024-04-28 15:06:43 -0700
commit5737e8430d43b3b5bd448f761cd2f3c35707a174 (patch)
treeb24f44ea5245a56b2f712c2eef6a3897f53bd098
parent9f5889c1263d63b953b3a324da07321deef5ca58 (diff)
downloadsml-5737e8430d43b3b5bd448f761cd2f3c35707a174.tar.zst
Add case over int.
-rw-r--r--bytecode/src/encoding.rs43
-rw-r--r--bytecode/src/heap.rs4
-rw-r--r--bytecode/src/main.rs28
-rw-r--r--codegen.sml15
-rw-r--r--cps.sml53
-rw-r--r--elab.sml22
-rw-r--r--linker.sml17
-rw-r--r--main.sml1
-rw-r--r--parser.sml26
-rw-r--r--run_tests.fish8
-rw-r--r--sort.sml30
-rw-r--r--syntax.sml16
-rw-r--r--tests/11-case.sml1
-rw-r--r--tests/12-case-int.sml5
14 files changed, 248 insertions, 21 deletions
diff --git a/bytecode/src/encoding.rs b/bytecode/src/encoding.rs
index 8bec753..80fcc58 100644
--- a/bytecode/src/encoding.rs
+++ b/bytecode/src/encoding.rs
@@ -68,6 +68,26 @@ pub struct Div {
}
#[derive(Debug, Clone, Copy)]
+pub struct Less {
+ pub out: u8,
+ pub val1: Arg,
+ pub val2: Arg,
+}
+
+#[derive(Debug, Clone, Copy)]
+pub struct Eq {
+ pub out: u8,
+ pub val1: Arg,
+ pub val2: Arg,
+}
+
+#[derive(Debug, Clone, Copy)]
+pub struct If {
+ pub test: Arg,
+ pub target: usize,
+}
+
+#[derive(Debug, Clone, Copy)]
pub enum Op {
Alloc(Alloc),
Call,
@@ -79,6 +99,9 @@ pub enum Op {
Sub(Sub),
Mul(Mul),
Div(Div),
+ Less(Less),
+ Eq(Eq),
+ If(If),
}
struct Reader<'a> {
@@ -198,6 +221,26 @@ impl Op {
let val2 = r.parse_arg(arg2_const)?;
Op::Div(Div { out, val1, val2 })
}
+ 11 => {
+ let out = r.parse_local()?;
+ let val1 = r.parse_arg(arg1_const)?;
+ let val2 = r.parse_arg(arg2_const)?;
+ Op::Less(Less { out, val1, val2 })
+ }
+ 12 => {
+ let out = r.parse_local()?;
+ let val1 = r.parse_arg(arg1_const)?;
+ let val2 = r.parse_arg(arg2_const)?;
+ Op::Eq(Eq { out, val1, val2 })
+ }
+ 13 => {
+ let test = r.parse_arg(arg1_const)?;
+ let Some(itarget) = r.parse_value()?.to_int() else {
+ return Err(Box::from("invalid target"));
+ };
+ let target = usize::try_from(itarget)?;
+ Op::If(If { test, target })
+ }
_ => {
return Err(Box::from(format!("invalid code {}", code)));
}
diff --git a/bytecode/src/heap.rs b/bytecode/src/heap.rs
index 45ce283..970db7d 100644
--- a/bytecode/src/heap.rs
+++ b/bytecode/src/heap.rs
@@ -1,7 +1,7 @@
use crate::value::Value;
use std::error::Error;
-pub const NUM_LOCALS: usize = 8;
+pub const NUM_LOCALS: usize = 256;
#[derive(Debug)]
pub struct Heap {
@@ -27,7 +27,7 @@ impl Heap {
let Some(q) = Value(self.buf[p]).to_pointer() else {
return false;
};
- return self.first_active == (q < self.buf.len() / 2);
+ self.first_active == (q < self.buf.len() / 2)
}
fn alloc_size(&self, p: usize) -> usize {
diff --git a/bytecode/src/main.rs b/bytecode/src/main.rs
index bbc0ddd..4cba4ef 100644
--- a/bytecode/src/main.rs
+++ b/bytecode/src/main.rs
@@ -96,6 +96,34 @@ impl State {
};
self.heap.locals[usize::from(op.out)] = Value::from_int(res);
}
+ Op::Less(op) => {
+ let Some(v1) = self.read_arg(op.val1).to_int() else {
+ return Err(Box::from("less needs an int"));
+ };
+ let Some(v2) = self.read_arg(op.val2).to_int() else {
+ return Err(Box::from("less needs an int"));
+ };
+ let res = if v1 < v2 { 1 } else { 0 };
+ self.heap.locals[usize::from(op.out)] = Value::from_int(res);
+ }
+ Op::Eq(op) => {
+ let Some(v1) = self.read_arg(op.val1).to_int() else {
+ return Err(Box::from("eq needs an int"));
+ };
+ let Some(v2) = self.read_arg(op.val2).to_int() else {
+ return Err(Box::from("eq needs an int"));
+ };
+ let res = if v1 == v2 { 1 } else { 0 };
+ self.heap.locals[usize::from(op.out)] = Value::from_int(res);
+ }
+ Op::If(op) => {
+ let Some(t) = self.read_arg(op.test).to_int() else {
+ return Err(Box::from("if needs an int"));
+ };
+ if t != 0 {
+ self.i = op.target;
+ }
+ }
}
Ok(())
}
diff --git a/codegen.sml b/codegen.sml
index 30928d4..a4fef6c 100644
--- a/codegen.sml
+++ b/codegen.sml
@@ -2,8 +2,8 @@ structure CodeGen =
struct
fun enumerate l = ListPair.zip (List.tabulate (length l, (fn x => x)), l)
- (* There are 8 registers *)
- val tempReg = 7
+ (* There are 256 registers *)
+ val tempReg = 255
structure VarMap = Map (type k = Syntax.var
val cmp = Int.compare)
@@ -149,6 +149,17 @@ struct
| Syntax.CPrimop (Syntax.PSub, [x, y], [res], [k]) => Syntax.OSub (translate res, translateVal x, translateVal y) :: go k
| Syntax.CPrimop (Syntax.PMul, [x, y], [res], [k]) => Syntax.OMul (translate res, translateVal x, translateVal y) :: go k
| Syntax.CPrimop (Syntax.PDiv, [x, y], [res], [k]) => Syntax.ODiv (translate res, translateVal x, translateVal y) :: go k
+ | Syntax.CPrimop (Syntax.PLess, [x, y], [res], [k]) => Syntax.OLess (translate res, translateVal x, translateVal y) :: go k
+ | Syntax.CPrimop (Syntax.PEq, [x, y], [res], [k]) => Syntax.OEq (translate res, translateVal x, translateVal y) :: go k
+ | Syntax.CPrimop (Syntax.PIf, [b], [], [k1, k2]) =>
+ let
+ val trueLabel = Gensym.new ()
+ in
+ Syntax.OIf (translateVal b, trueLabel)
+ :: go k2
+ @ Syntax.OLabel trueLabel
+ :: go k1
+ end
| _ => raise Fail ("malformed CPS:\n" ^ Syntax.cexpToString expr)
in go expr
end
diff --git a/cps.sml b/cps.sml
index adb29c0..ad3e0e8 100644
--- a/cps.sml
+++ b/cps.sml
@@ -57,6 +57,59 @@ struct
toCPS expr (fn v => go exprs (v :: vars))
in go exprs []
end
+ | Syntax.LSwitch (expr, arms, otherwise) =>
+ let
+ val sortedArms = Sort.sort (fn ((x, _), (y, _)) => Int.compare (x, y)) arms
+ val addr = Gensym.new ()
+ val arg = Gensym.new ()
+ fun go _ [] cont = toCPS otherwise cont
+ | go v [(x, arm)] cont =
+ let val b = Gensym.new ()
+ in
+ Syntax.CPrimop
+ ( Syntax.PEq
+ , [v, Syntax.VInt x]
+ , [b]
+ , [ Syntax.CPrimop
+ ( Syntax.PIf
+ , [Syntax.VVar b]
+ , []
+ , [toCPS arm cont, toCPS otherwise cont]
+ )
+ ]
+ )
+ end
+ | go v arms cont =
+ let
+ val b = Gensym.new ()
+ val h = length arms div 2
+ val half1 = List.take (arms, h)
+ val half2 = List.drop (arms, h)
+ val (x, _) = hd half2
+ in
+ Syntax.CPrimop
+ ( Syntax.PLess
+ , [v, Syntax.VInt x]
+ , [b]
+ , [ Syntax.CPrimop
+ ( Syntax.PIf
+ , [Syntax.VVar b]
+ , []
+ , [go v half1 cont, go v half2 cont]
+ )
+ ]
+ )
+ end
+ in
+ Syntax.CFix
+ ( [(addr, [arg], cont (Syntax.VVar arg))]
+ , toCPS expr
+ (fn v =>
+ go v sortedArms
+ (fn x =>
+ Syntax.CApp (Syntax.VVar addr, [x])))
+ )
+ end
| _ => raise Fail ("malformed expression " ^ Syntax.lexpToString e)
fun hoist (expr : Syntax.cexp) : Syntax.cexp =
diff --git a/elab.sml b/elab.sml
index cc73c98..f6f2e38 100644
--- a/elab.sml
+++ b/elab.sml
@@ -39,17 +39,23 @@ struct
val v = Gensym.new ()
val env' =
case pat of
- Syntax.PWild => env
- | Syntax.PVar name => StringMap.insert name v env
+ Syntax.PVar name => StringMap.insert name v env
+ | _ => env
in Syntax.LFn (v, elab env' body)
end
| Syntax.ECase (expr, []) => raise Fail "nonexhaustive match"
- | Syntax.ECase (expr, (pat, body) :: rest) =>
- if not (null rest) then raise Fail "redundant match" else
- let val v = Gensym.new () in
- case pat of
- Syntax.PWild => Syntax.LApp (Syntax.LFn (v, elab env body), elab env expr)
- | Syntax.PVar name => Syntax.LApp (Syntax.LFn (v, elab (StringMap.insert name v env) body), elab env expr)
+ | 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 []
end
fun elaborate (p : Syntax.expr) : Syntax.lexp = elab StringMap.empty p
diff --git a/linker.sml b/linker.sml
index 927bf48..1515079 100644
--- a/linker.sml
+++ b/linker.sml
@@ -102,6 +102,23 @@ struct
; writeValue w v1
; writeValue w v2
)
+ | Syntax.OLess (r, v1, v2) =>
+ ( makeOpcode w 11 (isConst v1) (isConst v2)
+ ; writeVar w r
+ ; writeValue w v1
+ ; writeValue w v2
+ )
+ | Syntax.OEq (r, v1, v2) =>
+ ( makeOpcode w 12 (isConst v1) (isConst v2)
+ ; writeVar w r
+ ; writeValue w v1
+ ; writeValue w v2
+ )
+ | Syntax.OIf (condition, label) =>
+ ( makeOpcode w 13 (isConst condition) false
+ ; writeValue w condition
+ ; writeValue w (Syntax.VLabel label)
+ )
| Syntax.OLabel _ => ()
end
diff --git a/main.sml b/main.sml
index 4b22f65..8a01ec4 100644
--- a/main.sml
+++ b/main.sml
@@ -1,3 +1,4 @@
+use "sort.sml";
use "format.sml";
use "result.sml";
use "buffer.sml";
diff --git a/parser.sml b/parser.sml
index 6efc4aa..3c044af 100644
--- a/parser.sml
+++ b/parser.sml
@@ -275,9 +275,10 @@ struct
List.exists (fn c => c = start) symbolicIdentifierLetters
in
lexeme
- (try (parseString s) >>
- notFollowedBy (if isSymbolic then symbolicIdentifierLetter else alphaNumIdentifierLetter) >>
- const s)
+ (try
+ (parseString s >>
+ notFollowedBy (if isSymbolic then symbolicIdentifierLetter else alphaNumIdentifierLetter) >>
+ const s))
<?> s
end
@@ -349,6 +350,7 @@ struct
val rec atpat : Syntax.pat parser = fn st =>
(Syntax.PWild <$ reserved "_"
+ <|> Syntax.PInt <$> integer
<|> between (symbol "(") (symbol ")") pat
<|> Syntax.PVar <$> identifier) st
and pat : Syntax.pat parser = fn st => atpat st
@@ -437,16 +439,24 @@ struct
bind orelseExpr (fn e2 =>
const (Syntax.EOrElse (e1, e2))))
<|> const e1) st
- and handleExpr : Syntax.expr parser = fn st => orelseExpr st
- and raiseExpr : Syntax.expr parser = fn st => handleExpr st
- and fnExpr : Syntax.expr parser = fn st =>
+ and expr : Syntax.expr parser = fn st =>
((reserved "fn" >>
bind pat (fn p =>
reserved "=>" >>
bind expr (fn e =>
const (Syntax.ELambda (p, e)))))
- <|> raiseExpr) st
- and expr : Syntax.expr parser = fn st => fnExpr st
+ <|> (reserved "case" >>
+ bind expr (fn e =>
+ reserved "of" >>
+ bind
+ (sepBy1
+ (bind pat (fn p =>
+ reserved "=>" >>
+ bind expr (fn e =>
+ const (p, e))))
+ (reserved "|")) (fn arms =>
+ const (Syntax.ECase (e, arms)))))
+ <|> orelseExpr) st
and dec : Syntax.dec parser = fn st =>
(reserved "val" >>
bind pat (fn p =>
diff --git a/run_tests.fish b/run_tests.fish
index f91b8ae..1005fac 100644
--- a/run_tests.fish
+++ b/run_tests.fish
@@ -7,7 +7,13 @@ cargo build
cd $d
-for f in $d/tests/*.sml
+if test (count $argv) -gt 0
+ set files $d/tests/$argv*
+else
+ set files $d/tests/*
+end
+
+for f in $files
$d/bytecode/target/debug/bytecode (sml $d/main.sml -o /dev/stderr $f >/dev/null 2>| psub)
set -l got $status
if test $got -eq 42
diff --git a/sort.sml b/sort.sml
new file mode 100644
index 0000000..97e4a8b
--- /dev/null
+++ b/sort.sml
@@ -0,0 +1,30 @@
+signature SORT =
+sig
+ val sort : ('a * 'a -> order) -> 'a list -> 'a list
+end
+
+structure Sort :> SORT =
+struct
+ fun split (l : 'a list, n : int) : ('a list * int) * ('a list * int) =
+ let val h = n div 2
+ in ((l, h), (List.drop (l, h), n - h))
+ end
+
+ fun merge (_ : 'a * 'a -> order) ([] : 'a list) (l2 : 'a list) : 'a list = l2
+ | merge _ l1 [] = l1
+ | merge cmp (xl as (x :: xs)) (yl as (y :: ys)) =
+ (case cmp (x, y) of
+ GREATER => y :: merge cmp xl ys
+ | _ => x :: merge cmp xs yl)
+
+ fun sort' (cmp : 'a * 'a -> order) (s as (l : 'a list, n : int)) : 'a list =
+ if n < 2 then List.take (l, n) else
+ let
+ val (first, second) = split s
+ val firstSorted = sort' cmp first
+ val secondSorted = sort' cmp second
+ in merge cmp firstSorted secondSorted
+ end
+
+ fun sort (cmp : 'a * 'a -> order) (l : 'a list) : 'a list = sort' cmp (l, length l)
+end
diff --git a/syntax.sml b/syntax.sml
index b91e0a2..f85e6f2 100644
--- a/syntax.sml
+++ b/syntax.sml
@@ -10,6 +10,7 @@ struct
datatype pat =
PWild
| PVar of string
+ | PInt of int
datatype expr =
EIdent of string list
@@ -37,6 +38,9 @@ struct
| PSub
| PMul
| PDiv
+ | PLess
+ | PEq
+ | PIf
datatype lexp =
LVar of var
@@ -47,6 +51,7 @@ struct
| LString of string
| LRecord of lexp list
| LPrim of primop
+ | LSwitch of lexp * (int * lexp) list * lexp
(* CPS *)
datatype value =
@@ -73,6 +78,9 @@ struct
| OSub of var * value * value
| OMul of var * value * value
| ODiv of var * value * value
+ | OLess of var * value * value
+ | OEq of var * value * value
+ | OIf of value * var
| OLabel of var
fun listToString (show : 'a -> string) (l : 'a list) =
@@ -100,6 +108,7 @@ struct
case p of
PWild => "PWild"
| PVar v => "PVar " ^ quote v
+ | PInt i => "PInt " ^ Int.toString i
fun exprToStringI (indent : string) (x : expr) : string =
let val self = exprToStringI indent
@@ -134,6 +143,9 @@ struct
| PSub => "PSub"
| PMul => "PMul"
| PDiv => "PDiv"
+ | PLess => "PLess"
+ | PEq => "PEq"
+ | PIf => "PIf"
fun lexpToString (x : lexp) : string =
case x of
@@ -145,6 +157,7 @@ struct
| LString s => "LString " ^ quote s
| LRecord l => "LRecord " ^ listToString lexpToString l
| LPrim p => "LPrim " ^ primopToString p
+ | LSwitch (e, arms, otherwise) => "LSwitch (" ^ lexpToString e ^ ", " ^ listToString (fn (x, e) => "(" ^ Int.toString x ^ ", " ^ lexpToString e ^ ")") arms ^ ", " ^ lexpToString otherwise ^ ")"
fun valueToString (x : value) : string =
case x of
@@ -179,5 +192,8 @@ struct
| OSub (r, v1, v2) => "Var " ^ Int.toString r ^ " = OSub (" ^ valueToString v1 ^ ", " ^ valueToString v2 ^ ")"
| OMul (r, v1, v2) => "Var " ^ Int.toString r ^ " = OMul (" ^ valueToString v1 ^ ", " ^ valueToString v2 ^ ")"
| ODiv (r, v1, v2) => "Var " ^ Int.toString r ^ " = ODiv (" ^ valueToString v1 ^ ", " ^ valueToString v2 ^ ")"
+ | OLess (r, v1, v2) => "Var " ^ Int.toString r ^ " = OLess (" ^ valueToString v1 ^ ", " ^ valueToString v2 ^ ")"
+ | OEq (r, v1, v2) => "Var " ^ Int.toString r ^ " = OEq (" ^ valueToString v1 ^ ", " ^ valueToString v2 ^ ")"
+ | OIf (condition, target) => "OIf (" ^ valueToString condition ^ ") goto " ^ Int.toString target
| OLabel l => "OLabel " ^ Int.toString l
end
diff --git a/tests/11-case.sml b/tests/11-case.sml
new file mode 100644
index 0000000..297ae5b
--- /dev/null
+++ b/tests/11-case.sml
@@ -0,0 +1 @@
+val _ = case 42 of x => __builtin "exit" x
diff --git a/tests/12-case-int.sml b/tests/12-case-int.sml
new file mode 100644
index 0000000..98967e4
--- /dev/null
+++ b/tests/12-case-int.sml
@@ -0,0 +1,5 @@
+val _ =
+ case 18 of
+ 17 => __builtin "exit" 41
+ | 18 => __builtin "exit" 42
+ | _ => __builtin "exit" 43