diff options
| author | Rose Hogenson <rosehogenson@posteo.net> | 2024-04-28 15:06:43 -0700 |
|---|---|---|
| committer | Rose Hogenson <rosehogenson@posteo.net> | 2024-04-28 15:06:43 -0700 |
| commit | 5737e8430d43b3b5bd448f761cd2f3c35707a174 (patch) | |
| tree | b24f44ea5245a56b2f712c2eef6a3897f53bd098 | |
| parent | 9f5889c1263d63b953b3a324da07321deef5ca58 (diff) | |
| download | sml-5737e8430d43b3b5bd448f761cd2f3c35707a174.tar.zst | |
Add case over int.
| -rw-r--r-- | bytecode/src/encoding.rs | 43 | ||||
| -rw-r--r-- | bytecode/src/heap.rs | 4 | ||||
| -rw-r--r-- | bytecode/src/main.rs | 28 | ||||
| -rw-r--r-- | codegen.sml | 15 | ||||
| -rw-r--r-- | cps.sml | 53 | ||||
| -rw-r--r-- | elab.sml | 22 | ||||
| -rw-r--r-- | linker.sml | 17 | ||||
| -rw-r--r-- | main.sml | 1 | ||||
| -rw-r--r-- | parser.sml | 26 | ||||
| -rw-r--r-- | run_tests.fish | 8 | ||||
| -rw-r--r-- | sort.sml | 30 | ||||
| -rw-r--r-- | syntax.sml | 16 | ||||
| -rw-r--r-- | tests/11-case.sml | 1 | ||||
| -rw-r--r-- | tests/12-case-int.sml | 5 |
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 @@ -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 = @@ -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 @@ -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 @@ -1,3 +1,4 @@ +use "sort.sml"; use "format.sml"; use "result.sml"; use "buffer.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 @@ -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 |
