use crate::eval; use crate::eval::Num; use crate::lexer::Lexer; use num::BigInt; use std::str::FromStr; enum ParseError { Empty, Consumed, } fn seq(r: Result) -> Result { match r { Err(ParseError::Empty) => Err(ParseError::Consumed), x => x, } } fn or(r: Result) -> Result, ParseError> { match r { Ok(x) => Ok(Some(x)), Err(ParseError::Empty) => Ok(None), Err(ParseError::Consumed) => Err(ParseError::Consumed), } } struct Parser<'a> { r: Lexer<'a>, first_token: Option<&'a str>, } impl<'a> Parser<'a> { fn peek(&mut self) -> Result<&'a str, ParseError> { if let Some(t) = self.first_token { return Ok(t); } let Some(t) = self.r.next() else { return Err(ParseError::Empty); }; self.first_token = Some(t); Ok(t) } fn next(&mut self) -> Result<&'a str, ParseError> { if let Some(t) = std::mem::take(&mut self.first_token) { return Ok(t); } let Some(t) = self.r.next() else { return Err(ParseError::Empty); }; Ok(t) } fn symbol(&mut self, s: &str) -> bool { let Ok(t) = self.peek() else { return false; }; if t != s { return false; } let _ = self.next(); true } fn parse_num(&mut self) -> Result { let negative = self.symbol("-"); let t = seq(self.next())?; if t.contains('.') || t.contains('e') || t.contains('E') { let Ok(f) = f64::from_str(t) else { return Err(ParseError::Consumed); }; if negative { return Ok(eval::float(-f)); } return Ok(eval::float(f)); } let Ok(i) = BigInt::from_str(t) else { return Err(ParseError::Consumed); }; if negative { return Ok(eval::int(-i)); } Ok(eval::int(i)) } fn paren_expr(&mut self) -> Result { if !self.symbol("(") { return Err(ParseError::Empty); } let n = seq(self.expr())?; if !self.symbol(")") { return Err(ParseError::Consumed); } Ok(n) } fn parse_const(&mut self) -> Result { if self.symbol("e") { return Ok(eval::float(std::f64::consts::E)); } if self.symbol("pi") { return Ok(eval::float(std::f64::consts::PI)); } Err(ParseError::Empty) } fn atom(&mut self) -> Result { if let Some(n) = or(self.paren_expr())? { return Ok(n); } if let Some(n) = or(self.parse_const())? { return Ok(n); } self.parse_num() } fn parse_fun(&mut self) -> Result { if self.symbol("sin") { return Ok(seq(self.fun_expr())?.sin()); } if self.symbol("cos") { return Ok(seq(self.fun_expr())?.cos()); } if self.symbol("tan") { return Ok(seq(self.fun_expr())?.tan()); } if self.symbol("sqrt") { return Ok(seq(self.fun_expr())?.sqrt()); } if self.symbol("log") || self.symbol("ln") { return Ok(seq(self.fun_expr())?.log()); } if self.symbol("log10") { return Ok(seq(self.fun_expr())?.log10()); } if self.symbol("log2") { return Ok(seq(self.fun_expr())?.log2()); } if self.symbol("floor") { return Ok(seq(self.fun_expr())?.floor()); } if self.symbol("ceil") || self.symbol("ceiling") { return Ok(seq(self.fun_expr())?.ceil()); } if self.symbol("round") { return Ok(seq(self.fun_expr())?.round()); } if self.symbol("abs") { return Ok(seq(self.fun_expr())?.abs()); } Err(ParseError::Empty) } fn fun_expr(&mut self) -> Result { if let Some(n) = or(self.parse_fun())? { return Ok(n); } self.atom() } fn fun_expr_no_num(&mut self) -> Result { if let Some(n) = or(self.parse_fun())? { return Ok(n); } if let Some(n) = or(self.paren_expr())? { return Ok(n); } self.parse_const() } fn expt_expr(&mut self) -> Result { let e1 = self.fun_expr()?; if self.symbol("^") { let e2 = seq(self.expt_expr())?; return Ok(e1.pow(e2)); } Ok(e1) } fn expt_expr_no_num(&mut self) -> Result { let e1 = self.fun_expr_no_num()?; if self.symbol("^") { let e2 = seq(self.expt_expr())?; return Ok(e1.pow(e2)); } Ok(e1) } fn mul_expr_fold(&mut self, e1: Num) -> Result { if self.symbol("*") { let e2 = seq(self.expt_expr())?; return self.mul_expr_fold(e1.mul(e2)); } if self.symbol("/") { let e2 = seq(self.expt_expr())?; return self.mul_expr_fold(e1.div(e2)); } if self.symbol("//") { let e2 = seq(self.expt_expr())?; return self.mul_expr_fold(e1.int_div(e2)); } if self.symbol("%") { let e2 = seq(self.expt_expr())?; return self.mul_expr_fold(e1.modulo(e2)); } // Allow implicit multiplication if let Some(e2) = or(self.expt_expr_no_num())? { return self.mul_expr_fold(e1.mul(e2)); } Ok(e1) } fn mul_expr(&mut self) -> Result { let n = self.expt_expr()?; self.mul_expr_fold(n) } fn add_expr_fold(&mut self, e1: Num) -> Result { if self.symbol("+") { let e2 = seq(self.mul_expr())?; let e = e1.add(e2); return self.add_expr_fold(e); } if self.symbol("-") { let e2 = seq(self.mul_expr())?; let e = e1.sub(e2); return self.add_expr_fold(e); } Ok(e1) } fn expr(&mut self) -> Result { let n = self.mul_expr()?; self.add_expr_fold(n) } } pub fn parse(expr: &str) -> Option { let mut p = Parser { r: Lexer { buf: expr }, first_token: None, }; let n = p.expr().ok()?; if p.r.next().is_some() { // Trailing characters. return None; } Some(n) }