diff options
| author | Rose Hogenson <rosehogenson@posteo.net> | 2024-05-30 20:00:08 -0700 |
|---|---|---|
| committer | Rose Hogenson <rosehogenson@posteo.net> | 2024-05-30 21:09:54 -0700 |
| commit | 5e88902a161674605c6835c7be96d412d82e7dd4 (patch) | |
| tree | b14073a4270e0051dfa6dcfb0e3d4112ef5cb452 /src/parser.rs | |
| parent | 5f47e7b0613b7477ef91568f4a4ed1e10dc5f707 (diff) | |
| download | qc-5e88902a161674605c6835c7be96d412d82e7dd4.tar.zst | |
Write tests for parser.
Diffstat (limited to 'src/parser.rs')
| -rw-r--r-- | src/parser.rs | 390 |
1 files changed, 253 insertions, 137 deletions
diff --git a/src/parser.rs b/src/parser.rs index 5e02037..6eeb9d8 100644 --- a/src/parser.rs +++ b/src/parser.rs @@ -1,14 +1,12 @@ -use crate::eval; use crate::eval::Num; +use crate::expr::{BinOp, Expr, UnOp}; use crate::lexer::Lexer; use std::str::FromStr; enum Op { Paren, - Op { - prec: i8, - eval: Box<dyn FnOnce(Num) -> Num>, - }, + Un(UnOp), + Bin(BinOp, Expr), } struct Parser<'a> { @@ -51,213 +49,161 @@ impl<'a> Parser<'a> { let Ok(f) = f64::from_str(t) else { return None; }; - return Some(eval::float(f)); + return Some(Num::Float(f)); } let Ok(i) = i128::from_str(t) else { return None; }; - Some(eval::int(i)) + Some(Num::Int(i)) } fn parse_const(&mut self) -> Option<Num> { if self.symbol("e") { - return Some(eval::float(std::f64::consts::E)); + return Some(Num::Float(std::f64::consts::E)); } if self.symbol("pi") { - return Some(eval::float(std::f64::consts::PI)); + return Some(Num::Float(std::f64::consts::PI)); } self.parse_num() } fn parse_fun(&mut self) -> bool { if self.symbol("-") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(|n| eval::int(0).sub(n)), - }); + self.stack.push(Op::Un(UnOp::Neg)); return true; } if self.symbol("sin") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::sin), - }); + self.stack.push(Op::Un(UnOp::Sin)); return true; } if self.symbol("cos") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::cos), - }); + self.stack.push(Op::Un(UnOp::Cos)); return true; } if self.symbol("tan") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::tan), - }); + self.stack.push(Op::Un(UnOp::Tan)); return true; } if self.symbol("asin") || self.symbol("arcsin") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::asin), - }); + self.stack.push(Op::Un(UnOp::Asin)); return true; } if self.symbol("acos") || self.symbol("arccos") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::acos), - }); + self.stack.push(Op::Un(UnOp::Acos)); return true; } if self.symbol("atan") || self.symbol("arctan") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::atan), - }); + self.stack.push(Op::Un(UnOp::Atan)); return true; } if self.symbol("sqrt") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::sqrt), - }); + self.stack.push(Op::Un(UnOp::Sqrt)); return true; } if self.symbol("log") || self.symbol("ln") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::log), - }); + self.stack.push(Op::Un(UnOp::Log)); return true; } if self.symbol("log10") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::log10), - }); + self.stack.push(Op::Un(UnOp::Log10)); return true; } if self.symbol("log2") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::log2), - }); + self.stack.push(Op::Un(UnOp::Log2)); return true; } if self.symbol("floor") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::floor), - }); + self.stack.push(Op::Un(UnOp::Floor)); return true; } if self.symbol("ceil") || self.symbol("ceiling") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::ceil), - }); + self.stack.push(Op::Un(UnOp::Ceil)); return true; } if self.symbol("round") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::round), - }); + self.stack.push(Op::Un(UnOp::Round)); return true; } if self.symbol("abs") { - self.stack.push(Op::Op { - prec: 3, - eval: Box::new(Num::abs), - }); + self.stack.push(Op::Un(UnOp::Abs)); return true; } false } +} + +fn prec(op: BinOp) -> i8 { + match op { + BinOp::Pow => 2, + BinOp::Mul => 1, + BinOp::Div => 1, + BinOp::IntDiv => 1, + BinOp::Mod => 1, + BinOp::Add => 0, + BinOp::Sub => 0, + } +} - fn eval(&mut self, n: Num, target_prec: i8) -> Num { - let mut n = n; - while let Some(Op::Op { prec, .. }) = self.stack.last() { - if *prec < target_prec { - break; +impl<'a> Parser<'a> { + fn eval(&mut self, e: Expr, target_prec: i8) -> Expr { + let mut e = e; + loop { + match self.stack.last() { + None | Some(Op::Paren) => break, + Some(Op::Bin(op, _)) if prec(*op) < target_prec => break, + _ => (), + } + match self.stack.pop().unwrap() { + Op::Un(op) => e = op.expr(e), + Op::Bin(op, e1) => e = op.expr(e1, e), + Op::Paren => unreachable!(), } - let Some(Op::Op { eval, .. }) = self.stack.pop() else { - unreachable!(); - }; - n = eval(n); } - n + e } - fn parse_op(&mut self, n: Num) { + fn parse_op(&mut self, e: Expr) { if self.symbol("^") { - let n1 = self.eval(n, 3); - self.stack.push(Op::Op { - prec: 2, - eval: Box::new(move |n2| n1.pow(n2)), - }); + let e = self.eval(e, 3); + self.stack.push(Op::Bin(BinOp::Pow, e)); return; } if self.symbol("*") { - let n1 = self.eval(n, 1); - self.stack.push(Op::Op { - prec: 1, - eval: Box::new(move |n2| n1.mul(n2)), - }); + let e = self.eval(e, 1); + self.stack.push(Op::Bin(BinOp::Mul, e)); return; } if self.symbol("/") { - let n1 = self.eval(n, 1); - self.stack.push(Op::Op { - prec: 1, - eval: Box::new(move |n2| n1.div(n2)), - }); + let e = self.eval(e, 1); + self.stack.push(Op::Bin(BinOp::Div, e)); return; } if self.symbol("//") { - let n1 = self.eval(n, 1); - self.stack.push(Op::Op { - prec: 1, - eval: Box::new(move |n2| n1.int_div(n2)), - }); + let e = self.eval(e, 1); + self.stack.push(Op::Bin(BinOp::IntDiv, e)); return; } if self.symbol("%") { - let n1 = self.eval(n, 1); - self.stack.push(Op::Op { - prec: 1, - eval: Box::new(move |n2| n1.modulo(n2)), - }); + let e = self.eval(e, 1); + self.stack.push(Op::Bin(BinOp::Mod, e)); return; } if self.symbol("+") { - let n1 = self.eval(n, 0); - self.stack.push(Op::Op { - prec: 0, - eval: Box::new(move |n2| n1.add(n2)), - }); + let e = self.eval(e, 0); + self.stack.push(Op::Bin(BinOp::Add, e)); return; } if self.symbol("-") { - let n1 = self.eval(n, 0); - self.stack.push(Op::Op { - prec: 0, - eval: Box::new(move |n2| n1.sub(n2)), - }); + let e = self.eval(e, 0); + self.stack.push(Op::Bin(BinOp::Sub, e)); return; } - let n1 = self.eval(n, 1); - self.stack.push(Op::Op { - prec: 1, - eval: Box::new(move |n2| n1.mul(n2)), - }); + let e = self.eval(e, 1); + self.stack.push(Op::Bin(BinOp::Mul, e)); } - fn parse(&mut self) -> Option<Num> { - let mut n = loop { + fn parse(&mut self) -> Option<Expr> { + let mut e = loop { if self.symbol("(") { self.stack.push(Op::Paren); continue; @@ -265,32 +211,33 @@ impl<'a> Parser<'a> { if self.parse_fun() { continue; } - let mut n = self.parse_const()?; + let mut e = Expr::Num(self.parse_const()?); while self.symbol(")") { loop { - let op = self.stack.pop()?; - let Op::Op { eval, .. } = op else { - break; - }; - n = eval(n); + match self.stack.pop()? { + Op::Un(op) => e = op.expr(e), + Op::Bin(op, e1) => e = op.expr(e1, e), + Op::Paren => break, + } } } if self.peek().is_none() { - break n; + break e; } - self.parse_op(n); + self.parse_op(e); }; - while let Some(op) = self.stack.pop() { - let Op::Op { eval, .. } = op else { - return None; - }; - n = eval(n); + loop { + match self.stack.pop() { + Some(Op::Paren) => return None, + Some(Op::Un(op)) => e = op.expr(e), + Some(Op::Bin(op, e1)) => e = op.expr(e1, e), + None => return Some(e), + } } - Some(n) } } -pub fn parse(expr: &str) -> Option<Num> { +pub fn parse(expr: &str) -> Option<Expr> { let mut p = Parser { r: Lexer { buf: expr }, first_token: None, @@ -298,3 +245,172 @@ pub fn parse(expr: &str) -> Option<Num> { }; p.parse() } + +#[cfg(test)] +mod tests { + use super::*; + + fn int(i: i128) -> Expr { + Expr::Num(Num::Int(i)) + } + + fn float(f: f64) -> Expr { + Expr::Num(Num::Float(f)) + } + + #[test] + fn parse_int() { + assert_eq!(parse("500"), Some(int(500))); + } + + #[test] + fn parse_float() { + assert_eq!(parse("1e2"), Some(float(100.))); + } + + #[test] + fn parse_fun() { + assert_eq!( + parse("sin pi"), + Some(UnOp::Sin.expr(float(std::f64::consts::PI))) + ); + } + + #[test] + fn parse_nested_fun() { + assert_eq!( + parse("log log 100"), + Some(UnOp::Log.expr(UnOp::Log.expr(int(100)))) + ); + } + + #[test] + fn parse_power() { + assert_eq!( + parse("2^1^2"), + Some(BinOp::Pow.expr(int(2), BinOp::Pow.expr(int(1), int(2)))) + ); + } + + #[test] + fn parse_fun_power() { + assert_eq!( + parse("log 2^2"), + Some(BinOp::Pow.expr(UnOp::Log.expr(int(2)), int(2))) + ); + } + + #[test] + fn parse_power_fun() { + assert_eq!( + parse("2^log 2"), + Some(BinOp::Pow.expr(int(2), UnOp::Log.expr(int(2)))) + ); + } + + #[test] + fn parse_pow_mul() { + assert_eq!( + parse("2^2*2"), + Some(BinOp::Mul.expr(BinOp::Pow.expr(int(2), int(2)), int(2))) + ); + } + + #[test] + fn parse_mul_pow() { + assert_eq!( + parse("2*2^2"), + Some(BinOp::Mul.expr(int(2), BinOp::Pow.expr(int(2), int(2)))) + ); + } + + #[test] + fn parse_mul_div() { + assert_eq!( + parse("2*2/2"), + Some(BinOp::Div.expr(BinOp::Mul.expr(int(2), int(2)), int(2))) + ); + } + + #[test] + fn parse_add_mul() { + assert_eq!( + parse("2+2*2"), + Some(BinOp::Add.expr(int(2), BinOp::Mul.expr(int(2), int(2)))) + ); + } + + #[test] + fn parse_mul_add() { + assert_eq!( + parse("2*2+2"), + Some(BinOp::Add.expr(BinOp::Mul.expr(int(2), int(2)), int(2))) + ); + } + + #[test] + fn parse_fun_add() { + assert_eq!( + parse("log 2+2"), + Some(BinOp::Add.expr(UnOp::Log.expr(int(2)), int(2))), + ); + } + + #[test] + fn parse_parens() { + assert_eq!( + parse("(1+2)*3"), + Some(BinOp::Mul.expr(BinOp::Add.expr(int(1), int(2)), int(3))) + ); + } + + #[test] + fn parse_implicit_multiplication() { + assert_eq!( + parse("2pi"), + Some(BinOp::Mul.expr(int(2), float(std::f64::consts::PI))), + ); + } + + #[test] + fn parse_implicit_multiplication_neg() { + assert_eq!( + parse("-2 2 -2"), + Some(BinOp::Sub.expr(BinOp::Mul.expr(UnOp::Neg.expr(int(2)), int(2)), int(2))) + ); + } + + #[test] + fn parse_fun_implicit_multiplication() { + assert_eq!( + parse("sin 2pi"), + Some(BinOp::Mul.expr(UnOp::Sin.expr(int(2)), float(std::f64::consts::PI))), + ); + } + + #[test] + fn parse_unary_negate() { + assert_eq!( + parse("1---2"), + Some(BinOp::Sub.expr(int(1), UnOp::Neg.expr(UnOp::Neg.expr(int(2))))) + ); + } + + #[test] + fn parse_negate_fun() { + assert_eq!( + parse("log-log 2"), + Some(UnOp::Log.expr(UnOp::Neg.expr(UnOp::Log.expr(int(2))))) + ); + } + + #[test] + fn parse_unmatched_parens() { + assert_eq!(parse("0))))"), None,); + } + + #[test] + fn parse_unmatched_op() { + assert_eq!(parse("2^2+"), None,); + } +} |
