use crate::eval::Num; use crate::expr::{BinOp, Expr, UnOp}; use crate::lexer::Lexer; use std::str::FromStr; enum Op { Paren, Un(UnOp), Bin(BinOp, Expr), } struct Parser<'a> { r: Lexer<'a>, first_token: Option<&'a str>, stack: Vec, } impl<'a> Parser<'a> { fn peek(&mut self) -> Option<&'a str> { if let Some(t) = self.first_token { return Some(t); } let t = self.r.next()?; self.first_token = Some(t); Some(t) } fn next(&mut self) -> Option<&'a str> { if let Some(t) = std::mem::take(&mut self.first_token) { return Some(t); } self.r.next() } fn symbol(&mut self, s: &str) -> bool { let Some(t) = self.peek() else { return false; }; if t != s { return false; } let _ = self.next(); true } fn parse_num(&mut self) -> Option { let t = self.next()?; if t.contains('.') || t.contains('e') || t.contains('E') { let Ok(f) = f64::from_str(t) else { return None; }; return Some(Num::Float(f)); } let Ok(i) = i128::from_str(t) else { return None; }; Some(Num::Int(i)) } fn parse_const(&mut self) -> Option { if self.symbol("e") { return Some(Num::Float(std::f64::consts::E)); } if self.symbol("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::Un(UnOp::Neg)); return true; } if self.symbol("sin") { self.stack.push(Op::Un(UnOp::Sin)); return true; } if self.symbol("cos") { self.stack.push(Op::Un(UnOp::Cos)); return true; } if self.symbol("tan") { self.stack.push(Op::Un(UnOp::Tan)); return true; } if self.symbol("asin") || self.symbol("arcsin") { self.stack.push(Op::Un(UnOp::Asin)); return true; } if self.symbol("acos") || self.symbol("arccos") { self.stack.push(Op::Un(UnOp::Acos)); return true; } if self.symbol("atan") || self.symbol("arctan") { self.stack.push(Op::Un(UnOp::Atan)); return true; } if self.symbol("sqrt") { self.stack.push(Op::Un(UnOp::Sqrt)); return true; } if self.symbol("log") || self.symbol("ln") { self.stack.push(Op::Un(UnOp::Log)); return true; } if self.symbol("log10") { self.stack.push(Op::Un(UnOp::Log10)); return true; } if self.symbol("log2") { self.stack.push(Op::Un(UnOp::Log2)); return true; } if self.symbol("floor") { self.stack.push(Op::Un(UnOp::Floor)); return true; } if self.symbol("ceil") || self.symbol("ceiling") { self.stack.push(Op::Un(UnOp::Ceil)); return true; } if self.symbol("round") { self.stack.push(Op::Un(UnOp::Round)); return true; } if self.symbol("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, } } 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!(), } } e } fn parse_op(&mut self, e: Expr) { if self.symbol("^") { let e = self.eval(e, 3); self.stack.push(Op::Bin(BinOp::Pow, e)); return; } if self.symbol("*") { let e = self.eval(e, 1); self.stack.push(Op::Bin(BinOp::Mul, e)); return; } if self.symbol("/") { let e = self.eval(e, 1); self.stack.push(Op::Bin(BinOp::Div, e)); return; } if self.symbol("//") { let e = self.eval(e, 1); self.stack.push(Op::Bin(BinOp::IntDiv, e)); return; } if self.symbol("%") { let e = self.eval(e, 1); self.stack.push(Op::Bin(BinOp::Mod, e)); return; } if self.symbol("+") { let e = self.eval(e, 0); self.stack.push(Op::Bin(BinOp::Add, e)); return; } if self.symbol("-") { let e = self.eval(e, 0); self.stack.push(Op::Bin(BinOp::Sub, e)); return; } let e = self.eval(e, 1); self.stack.push(Op::Bin(BinOp::Mul, e)); } fn parse(&mut self) -> Option { let mut e = loop { if self.symbol("(") { self.stack.push(Op::Paren); continue; } if self.parse_fun() { continue; } let mut e = Expr::Num(self.parse_const()?); while self.symbol(")") { loop { 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 e; } self.parse_op(e); }; 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), } } } } pub fn parse(expr: &str) -> Option { let mut p = Parser { r: Lexer { buf: expr }, first_token: None, stack: Vec::new(), }; 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,); } }