Refactoring exprs

This commit is contained in:
2023-07-08 22:31:53 -05:00
parent bc3a1ad88f
commit f682a234c7
4 changed files with 111 additions and 129 deletions
+3 -2
View File
@@ -7,5 +7,6 @@ edition = "2021"
[dependencies] [dependencies]
vcore = { path = "../vcore" } vcore = { path = "../vcore" }
pest = "2.1.3" pest = "2.7.0"
pest_derive = "2.1.0" pest_derive = "2.7.0"
lazy_static = "1.2.0"
+1 -1
View File
@@ -144,7 +144,7 @@ pub enum Node {
ArrayRef(ArrayRef), ArrayRef(ArrayRef),
Name(String), Name(String),
Address(String), Address(String),
Expr(BoxNode, Vec<(Operator, Node)>), Expr(BoxNode, Operator, BoxNode),
} }
impl From<i32> for Node { impl From<i32> for Node {
+8 -4
View File
@@ -20,10 +20,14 @@ assignment = { lvalue ~ "=" ~ rvalue }
lvalue = { arrayref | name } lvalue = { arrayref | name }
rvalue = { expr | string } rvalue = { expr | string }
sign = { "+" | "-" } add = { "+" }
term_op = { "*" | "/" | "%" } sub = { "-" }
expr = { term ~ (sign ~ term)* } mul = { "*" }
term = { val ~ (term_op ~ val)* } div = { "/" }
modulus = { "%" }
operator = _{ add | sub | mul | div | modulus }
expr = { val ~ (operator ~ val)* }
val = { // These are what get evaluated and left on the stack: val = { // These are what get evaluated and left on the stack:
number | // Literal numbers number | // Literal numbers
("(" ~ expr ~ ")") | // Nested parenthesized exprs ("(" ~ expr ~ ")") | // Nested parenthesized exprs
+99 -122
View File
@@ -1,13 +1,23 @@
use pest::pratt_parser::PrattParser;
use pest::Parser; use pest::Parser;
mod inner { #[derive(Parser)]
#[derive(Parser)] #[grammar = "forge.pest"]
#[grammar = "forge.pest"] struct ForgeParser;
pub struct ForgeParser;
lazy_static::lazy_static! {
static ref PRATT_PARSER: PrattParser<Rule> = {
use pest::pratt_parser::{Assoc::*, Op};
use Rule::*;
// Precedence is defined lowest to highest
PrattParser::new()
.op(Op::infix(add, Left) | Op::infix(sub, Left))
.op(Op::infix(mul, Left) | Op::infix(div, Left) | Op::infix(modulus, Left))
};
} }
use crate::ast::*; use crate::ast::*;
use inner::*;
use pest::error::{Error, LineColLocation}; use pest::error::{Error, LineColLocation};
use std::iter::Peekable; use std::iter::Peekable;
use std::str::FromStr; use std::str::FromStr;
@@ -110,11 +120,11 @@ trait AstNode: Sized {
} }
trait Parseable: Sized { trait Parseable: Sized {
fn parse(src: &str) -> Result<Self, ParseError>; fn from_str(src: &str) -> Result<Self, ParseError>;
} }
impl<T: AstNode> Parseable for T { impl<T: AstNode> Parseable for T {
fn parse(src: &str) -> Result<Self, ParseError> { fn from_str(src: &str) -> Result<Self, ParseError> {
let pair = ForgeParser::parse(Self::RULE, src) let pair = ForgeParser::parse(Self::RULE, src)
.map_err(ParseError::from)? .map_err(ParseError::from)?
.next() .next()
@@ -436,54 +446,28 @@ impl AstNode for RepeatLoop {
impl AstNode for Node { impl AstNode for Node {
const RULE: Rule = Rule::expr; // Also used for vals / terms / the whole tree const RULE: Rule = Rule::expr; // Also used for vals / terms / the whole tree
fn from_pair(pair: Pair) -> Self { fn from_pair(pair: Pair) -> Self {
let expr = match pair.as_rule() { PRATT_PARSER
Rule::val => Node::from_pair(pair.first()), .map_primary(|val| {
Rule::number => Node::Number(pair.into_number()), let primary = val.first();
Rule::name => Node::Name(String::from(pair.as_str())), match primary.as_rule() {
Rule::call => Node::Call(Call::from_pair(pair)), Rule::number => Node::Number(primary.into_number()),
Rule::arrayref => Node::ArrayRef(ArrayRef::from_pair(pair)), Rule::name => Node::Name(String::from(primary.as_str())),
Rule::address => Node::Address(pair.first_as_string()), Rule::call => Node::Call(Call::from_pair(primary)),
Rule::expr | Rule::term => { Rule::arrayref => Node::ArrayRef(ArrayRef::from_pair(primary)),
let mut children = pair.into_inner(); Rule::address => Node::Address(primary.first_as_string()),
let car = Node::from_pair(children.next().unwrap()); Rule::expr => Node::from_pair(primary),
Node::Expr(car.into(), parse_expr(children)) rule => unreachable!("Expr::parse expected atom, found {:?}", rule),
} }
_ => unreachable!(), })
}; .map_infix(|lhs, op, rhs| Node::Expr(lhs.into(), Operator::from_pair(op), rhs.into()))
shake_expr(expr) .parse(pair.into_inner())
}
}
fn parse_expr(mut cdr: Pairs) -> Vec<(Operator, Node)> {
let mut terms = Vec::new();
while let Some(operator) = cdr.next() {
let rhs = cdr.next().unwrap();
let rhs = Node::from_pair(rhs);
let op = Operator::from_pair(operator);
terms.push((op, rhs));
}
terms
}
fn shake_expr(node: Node) -> Node {
match node {
Node::Expr(car, cdr) => {
let shaken_car = shake_expr(car.into());
if cdr.is_empty() {
shaken_car
} else {
let shaken_cdr = cdr.into_iter().map(|(op, n)| (op, shake_expr(n))).collect();
Node::Expr(shaken_car.into(), shaken_cdr)
}
}
_ => node,
} }
} }
/////////////////////////////////////////////////////////////////////////////////////////// ///////////////////////////////////////////////////////////////////////////////////////////
impl AstNode for Operator { impl AstNode for Operator {
const RULE: Rule = Rule::sign; // also term_op const RULE: Rule = Rule::operator; // also term_op
fn from_pair(pair: Pair) -> Self { fn from_pair(pair: Pair) -> Self {
match pair.as_str() { match pair.as_str() {
"+" => Self::Add, "+" => Self::Add,
@@ -537,7 +521,7 @@ mod test {
#[test] #[test]
fn parse_globals() { fn parse_globals() {
assert_eq!( assert_eq!(
Global::parse("global foo;"), Global::from_str("global foo;"),
Ok(Global { Ok(Global {
name: "foo".into(), name: "foo".into(),
typename: None, typename: None,
@@ -545,7 +529,7 @@ mod test {
}) })
); );
assert_eq!( assert_eq!(
Global::parse("global foo:Thing;"), Global::from_str("global foo:Thing;"),
Ok(Global { Ok(Global {
name: "foo".into(), name: "foo".into(),
typename: Some("Thing".into()), typename: Some("Thing".into()),
@@ -553,7 +537,7 @@ mod test {
}) })
); );
assert_eq!( assert_eq!(
Global::parse("global foo:Thing[10];"), Global::from_str("global foo:Thing[10];"),
Ok(Global { Ok(Global {
name: "foo".into(), name: "foo".into(),
typename: Some("Thing".into()), typename: Some("Thing".into()),
@@ -561,7 +545,7 @@ mod test {
}) })
); );
assert_eq!( assert_eq!(
Global::parse("global foo[10];"), Global::from_str("global foo[10];"),
Ok(Global { Ok(Global {
name: "foo".into(), name: "foo".into(),
typename: None, typename: None,
@@ -573,7 +557,7 @@ mod test {
#[test] #[test]
fn parse_consts() { fn parse_consts() {
assert_eq!( assert_eq!(
Const::parse("const a = 123;"), Const::from_str("const a = 123;"),
Ok(Const { Ok(Const {
name: "a".into(), name: "a".into(),
value: Some(123), value: Some(123),
@@ -581,7 +565,7 @@ mod test {
}) })
); );
assert_eq!( assert_eq!(
Const::parse("const a = 0xaa;"), Const::from_str("const a = 0xaa;"),
Ok(Const { Ok(Const {
name: "a".into(), name: "a".into(),
value: Some(0xaa), value: Some(0xaa),
@@ -589,7 +573,7 @@ mod test {
}) })
); );
assert_eq!( assert_eq!(
Const::parse("const a = -7;"), Const::from_str("const a = -7;"),
Ok(Const { Ok(Const {
name: "a".into(), name: "a".into(),
value: Some(-7), value: Some(-7),
@@ -597,7 +581,7 @@ mod test {
}) })
); );
assert_eq!( assert_eq!(
Const::parse("const a = \"foo bar\";"), Const::from_str("const a = \"foo bar\";"),
Ok(Const { Ok(Const {
name: "a".into(), name: "a".into(),
value: None, value: None,
@@ -609,7 +593,7 @@ mod test {
#[test] #[test]
fn parse_structs() { fn parse_structs() {
assert_eq!( assert_eq!(
Struct::parse("struct Point { x, y }"), Struct::from_str("struct Point { x, y }"),
Ok(Struct { Ok(Struct {
name: "Point".into(), name: "Point".into(),
members: vec![ members: vec![
@@ -628,7 +612,7 @@ mod test {
); );
assert_eq!( assert_eq!(
Struct::parse("struct Foo { bar[100] }"), Struct::from_str("struct Foo { bar[100] }"),
Ok(Struct { Ok(Struct {
name: "Foo".into(), name: "Foo".into(),
members: vec![Member { members: vec![Member {
@@ -640,7 +624,7 @@ mod test {
); );
assert_eq!( assert_eq!(
Struct::parse("struct Foo { bar:Thing[100] }"), Struct::from_str("struct Foo { bar:Thing[100] }"),
Ok(Struct { Ok(Struct {
name: "Foo".into(), name: "Foo".into(),
members: vec![Member { members: vec![Member {
@@ -655,7 +639,7 @@ mod test {
#[test] #[test]
fn parse_function_headers() { fn parse_function_headers() {
assert_eq!( assert_eq!(
Function::parse("fn foo() {}"), Function::from_str("fn foo() {}"),
Ok(Function { Ok(Function {
name: "foo".into(), name: "foo".into(),
inline: false, inline: false,
@@ -666,7 +650,7 @@ mod test {
}) })
); );
assert_eq!( assert_eq!(
Function::parse("fn foo(a, b) {}"), Function::from_str("fn foo(a, b) {}"),
Ok(Function { Ok(Function {
name: "foo".into(), name: "foo".into(),
inline: false, inline: false,
@@ -686,7 +670,7 @@ mod test {
}) })
); );
assert_eq!( assert_eq!(
Function::parse("fn foo(a:Blah, b) {}"), Function::from_str("fn foo(a:Blah, b) {}"),
Ok(Function { Ok(Function {
name: "foo".into(), name: "foo".into(),
inline: false, inline: false,
@@ -719,10 +703,13 @@ mod test {
}; };
assert_eq!( assert_eq!(
Function::parse("fn foo<inline, org=0x400>(a) {}"), Function::from_str("fn foo<inline, org=0x400>(a) {}"),
Ok(func.clone()) Ok(func.clone())
); );
assert_eq!(Function::parse("fn foo<org=0x400, inline>(a) {}"), Ok(func)); assert_eq!(
Function::from_str("fn foo<org=0x400, inline>(a) {}"),
Ok(func)
);
} }
#[test] #[test]
@@ -730,47 +717,37 @@ mod test {
use Node::*; use Node::*;
use Operator::*; use Operator::*;
// A very, very basic expression // A very, very basic expression
assert_eq!(Node::parse("23"), Ok(Number(23))); assert_eq!(Node::from_str("23"), Ok(Number(23)));
// Two vals with an operator // Two vals with an operator
assert_eq!( assert_eq!(
Node::parse("23 + 5"), Node::from_str("23 + 5"),
Ok(Expr(23.into(), vec![(Add, Number(5))])) Ok(Expr(23.into(), Add, Number(5).into()))
); );
// Multiple terms at the same precedence level // Multiple terms at the same precedence level
assert_eq!( assert_eq!(
Node::parse("1 + 2 + 3"), Node::from_str("1 + 2 + 3"),
Ok(Expr(1.into(), vec![(Add, 2.into()), (Add, 3.into())])) Ok(Expr(Expr(1.into(), Add, 2.into()).into(), Add, 3.into()))
); );
// Higher precedence levels // Higher precedence levels
assert_eq!( assert_eq!(
Node::parse("1 + 2 * 3"), Node::from_str("1 + 2 * 3"),
Ok(Node::Expr( Ok(Expr(1.into(), Add, Expr(2.into(), Mul, 3.into()).into()))
1.into(),
vec![(Add, Expr(2.into(), vec![(Mul, 3.into())]))]
))
); );
assert_eq!(Node::from_str("2 * 3"), Ok(Expr(2.into(), Mul, 3.into())));
assert_eq!( assert_eq!(
Node::parse("2 * 3"), Node::from_str("2 * 3 + 4"),
Ok(Expr(Number(2).into(), vec![(Mul, 3.into())])) Ok(Expr(Expr(2.into(), Mul, 3.into()).into(), Add, 4.into()))
);
assert_eq!(
Node::parse("2 * 3 + 4"),
Ok(Expr(
Expr(2.into(), vec![(Mul, 3.into())]).into(),
vec![(Add, 4.into())]
))
); );
// Parens // Parens
assert_eq!( assert_eq!(
Node::parse("(1 + 2) * 3"), Node::from_str("(1 + 2) * 3"),
Ok(Node::Expr( Ok(Expr(Expr(1.into(), Add, 2.into()).into(), Mul, 3.into()))
Expr(1.into(), vec![(Add, 2.into())]).into(),
vec![(Mul, 3.into())]
))
); );
} }
@@ -782,7 +759,7 @@ mod test {
// Normal numbers // Normal numbers
assert_eq!( assert_eq!(
AR::parse("foo[7]"), AR::from_str("foo[7]"),
Ok(AR { Ok(AR {
name: "foo".into(), name: "foo".into(),
subscript: Number(7).into() subscript: Number(7).into()
@@ -791,16 +768,16 @@ mod test {
// Full exprs (this is the last one of these; the full expr test above covers it // Full exprs (this is the last one of these; the full expr test above covers it
assert_eq!( assert_eq!(
AR::parse("foo[7+x]"), AR::from_str("foo[7+x]"),
Ok(AR { Ok(AR {
name: "foo".into(), name: "foo".into(),
subscript: Expr(Number(7).into(), vec![(Add, Name("x".into()))]).into() subscript: Node::from_str("7+x").unwrap().into()
}) })
); );
// Exprs that are actually arrayrefs // Exprs that are actually arrayrefs
assert_eq!( assert_eq!(
Node::parse("foo[7]"), Node::from_str("foo[7]"),
Ok(ArrayRef(AR { Ok(ArrayRef(AR {
name: "foo".into(), name: "foo".into(),
subscript: Number(7).into() subscript: Number(7).into()
@@ -810,7 +787,7 @@ mod test {
#[test] #[test]
fn parse_addresses() { fn parse_addresses() {
assert_eq!(Node::parse("&foo"), Ok(Node::Address("foo".into()))); assert_eq!(Node::from_str("&foo"), Ok(Node::Address("foo".into())));
} }
#[test] #[test]
@@ -823,14 +800,14 @@ mod test {
}; };
// Can Node parse a call? // Can Node parse a call?
assert_eq!(Node::parse("blah()"), Ok(Node::Call(blah.clone()))); assert_eq!(Node::from_str("blah()"), Ok(Node::Call(blah.clone())));
// Can Statement parse a call? // Can Statement parse a call?
assert_eq!(Statement::parse("blah();"), Ok(Statement::Call(blah))); assert_eq!(Statement::from_str("blah();"), Ok(Statement::Call(blah)));
// Calls with args // Calls with args
assert_eq!( assert_eq!(
Call::parse("blah(1, 2)"), Call::from_str("blah(1, 2)"),
Ok(Call { Ok(Call {
name: "blah".into(), name: "blah".into(),
args: vec![Rvalue::Expr(Number(1)), Rvalue::Expr(Number(2))] args: vec![Rvalue::Expr(Number(1)), Rvalue::Expr(Number(2))]
@@ -839,7 +816,7 @@ mod test {
// Calls with strings // Calls with strings
assert_eq!( assert_eq!(
Call::parse("blah(\"foo\", 2)"), Call::from_str("blah(\"foo\", 2)"),
Ok(Call { Ok(Call {
name: "blah".into(), name: "blah".into(),
args: vec![Rvalue::String("foo".into()), Rvalue::Expr(Number(2))] args: vec![Rvalue::String("foo".into()), Rvalue::Expr(Number(2))]
@@ -850,12 +827,12 @@ mod test {
#[test] #[test]
fn parse_return() { fn parse_return() {
assert_eq!( assert_eq!(
Statement::parse("return;"), Statement::from_str("return;"),
Ok(Statement::Return(Return(None))) Ok(Statement::Return(Return(None)))
); );
assert_eq!( assert_eq!(
Statement::parse("return 17;"), Statement::from_str("return 17;"),
Ok(Statement::Return(Return(Some(Node::Number(17))))) Ok(Statement::Return(Return(Some(Node::Number(17)))))
); );
} }
@@ -863,7 +840,7 @@ mod test {
#[test] #[test]
fn parse_assignment() { fn parse_assignment() {
assert_eq!( assert_eq!(
Statement::parse("foo = 7;"), Statement::from_str("foo = 7;"),
Ok(Statement::Assignment(Assignment { Ok(Statement::Assignment(Assignment {
lvalue: Lvalue::Name("foo".into()), lvalue: Lvalue::Name("foo".into()),
rvalue: Rvalue::Expr(Node::Number(7)) rvalue: Rvalue::Expr(Node::Number(7))
@@ -871,7 +848,7 @@ mod test {
); );
assert_eq!( assert_eq!(
Assignment::parse("foo[45] = 7"), Assignment::from_str("foo[45] = 7"),
Ok(Assignment { Ok(Assignment {
lvalue: Lvalue::ArrayRef(ArrayRef { lvalue: Lvalue::ArrayRef(ArrayRef {
name: "foo".into(), name: "foo".into(),
@@ -885,7 +862,7 @@ mod test {
#[test] #[test]
fn parse_var_decl() { fn parse_var_decl() {
assert_eq!( assert_eq!(
Statement::parse("var blah;"), Statement::from_str("var blah;"),
Ok(Statement::VarDecl(VarDecl { Ok(Statement::VarDecl(VarDecl {
name: "blah".into(), name: "blah".into(),
typename: None, typename: None,
@@ -895,7 +872,7 @@ mod test {
); );
assert_eq!( assert_eq!(
VarDecl::parse("var blah:Foo[7] = 35"), VarDecl::from_str("var blah:Foo[7] = 35"),
Ok(VarDecl { Ok(VarDecl {
name: "blah".into(), name: "blah".into(),
typename: Some("Foo".into()), typename: Some("Foo".into()),
@@ -908,7 +885,7 @@ mod test {
#[test] #[test]
fn parse_block() { fn parse_block() {
assert_eq!( assert_eq!(
Block::parse("{ foo(); bar(); }"), Block::from_str("{ foo(); bar(); }"),
Ok(Block(vec![ Ok(Block(vec![
Statement::Call(Call { Statement::Call(Call {
name: "foo".into(), name: "foo".into(),
@@ -925,20 +902,20 @@ mod test {
#[test] #[test]
fn parse_conditional() { fn parse_conditional() {
assert_eq!( assert_eq!(
Statement::parse("if(cond) { foo(); }"), Statement::from_str("if(cond) { foo(); }"),
Ok(Statement::Conditional(Conditional { Ok(Statement::Conditional(Conditional {
condition: Node::parse("cond").unwrap(), condition: Node::from_str("cond").unwrap(),
body: Block::parse("{ foo(); }").unwrap(), body: Block::from_str("{ foo(); }").unwrap(),
alternative: None alternative: None
})) }))
); );
assert_eq!( assert_eq!(
Statement::parse("if(cond) { foo(); } else { bar(); }"), Statement::from_str("if(cond) { foo(); } else { bar(); }"),
Ok(Statement::Conditional(Conditional { Ok(Statement::Conditional(Conditional {
condition: Node::parse("cond").unwrap(), condition: Node::from_str("cond").unwrap(),
body: Block::parse("{ foo(); }").unwrap(), body: Block::from_str("{ foo(); }").unwrap(),
alternative: Some(Block::parse("{ bar(); }").unwrap()), alternative: Some(Block::from_str("{ bar(); }").unwrap()),
})) }))
); );
} }
@@ -946,10 +923,10 @@ mod test {
#[test] #[test]
fn parse_while_loops() { fn parse_while_loops() {
assert_eq!( assert_eq!(
Statement::parse("while(cond) { foo(); }"), Statement::from_str("while(cond) { foo(); }"),
Ok(Statement::WhileLoop(WhileLoop { Ok(Statement::WhileLoop(WhileLoop {
condition: Node::parse("cond").unwrap(), condition: Node::from_str("cond").unwrap(),
body: Block::parse("{ foo(); }").unwrap(), body: Block::from_str("{ foo(); }").unwrap(),
})) }))
); );
} }
@@ -957,20 +934,20 @@ mod test {
#[test] #[test]
fn parse_repeat_loops() { fn parse_repeat_loops() {
assert_eq!( assert_eq!(
Statement::parse("repeat(10) x { foo(x); }"), Statement::from_str("repeat(10) x { foo(x); }"),
Ok(Statement::RepeatLoop(RepeatLoop { Ok(Statement::RepeatLoop(RepeatLoop {
count: Node::Number(10), count: Node::Number(10),
name: Some("x".into()), name: Some("x".into()),
body: Block::parse("{ foo(x); }").unwrap(), body: Block::from_str("{ foo(x); }").unwrap(),
})) }))
); );
assert_eq!( assert_eq!(
Statement::parse("repeat(10) { foo(); }"), Statement::from_str("repeat(10) { foo(); }"),
Ok(Statement::RepeatLoop(RepeatLoop { Ok(Statement::RepeatLoop(RepeatLoop {
count: Node::Number(10), count: Node::Number(10),
name: None, name: None,
body: Block::parse("{ foo(); }").unwrap(), body: Block::from_str("{ foo(); }").unwrap(),
})) }))
); );
} }
@@ -978,10 +955,10 @@ mod test {
#[test] #[test]
fn parse_program() { fn parse_program() {
assert_eq!( assert_eq!(
Program::parse("global foo; struct Point { x, y }"), Program::from_str("global foo; struct Point { x, y }"),
Ok(Program(vec![ Ok(Program(vec![
Declaration::parse("global foo;").unwrap(), Declaration::from_str("global foo;").unwrap(),
Declaration::parse("struct Point { x, y }").unwrap(), Declaration::from_str("struct Point { x, y }").unwrap(),
])) ]))
) )
} }