use super::*;
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
enum State {
Value,
Operator,
}
struct FnVal {
pfn: Function,
pre: Order,
nargs: u8,
}
pub struct Expr<'a> {
env: &'a dyn Env,
fns: Vec<FnVal>,
vals: Vec<Value>,
next: State,
position: usize,
}
impl<'a> Expr<'a> {
#[inline]
pub const fn new(env: &'a dyn Env) -> Expr<'a> {
Expr {
env,
fns: Vec::new(),
vals: Vec::new(),
next: State::Value,
position: 0,
}
}
#[inline]
fn with_capacity(env: &'a dyn Env, capacity: usize) -> Expr<'a> {
Expr {
env,
fns: Vec::with_capacity(capacity),
vals: Vec::with_capacity(capacity),
next: State::Value,
position: 0,
}
}
pub fn parse(&mut self, tok: &Token) -> Result<(), Error> {
self.position = tok.position;
let result = match self.next {
State::Operator => self.parse_op(&tok.kind),
State::Value => self.parse_val(&tok.kind),
};
result.map_err(|kind| Error { kind, position: tok.position })
}
pub fn feed(&mut self, input: &str) -> Result<(), Error> {
for tok in tokenize(input) {
self.parse(&tok)?;
}
Ok(())
}
pub fn result(mut self) -> Result<Value, Error> {
let position = self.position;
let wrap_err = |kind| Error { kind, position };
if self.next == State::Value {
return Err(wrap_err(ErrorKind::UnfinishedExpression));
}
self.eval_gt(Order::FnBarrier).map_err(wrap_err)?;
if self.vals.len() != 1 || self.fns.len() != 0 {
return Err(wrap_err(ErrorKind::UnbalancedParens));
}
Ok(self.vals[0])
}
}
impl<'a> Expr<'a> {
fn parse_val(&mut self, tok: &TokenKind) -> Result<(), ErrorKind> {
match tok {
TokenKind::Unk(_) => {
Err(ErrorKind::InvalidToken)
},
TokenKind::Lit(val) => {
self.vals.push(*val);
self.next = State::Operator;
Ok(())
},
TokenKind::Op(op) => {
let desc = op.desc();
if desc.unary {
self.fns.push(FnVal {
pfn: desc.pfn,
pre: Order::Unary,
nargs: 1,
});
self.next = State::Value;
Ok(())
}
else {
Err(ErrorKind::DisallowedUnary)
}
},
TokenKind::Var(name) => {
let result = self.env.value(name)?;
self.vals.push(result);
self.next = State::Operator;
Ok(())
},
TokenKind::Open(name) => {
let pfn = self.env.function(name)?;
let pre = Order::FnBarrier; let nargs = 1;
self.fns.push(FnVal { pfn, pre, nargs });
self.next = State::Value;
Ok(())
},
TokenKind::Comma => {
Err(ErrorKind::NaExpression)
},
TokenKind::Close => {
if self.fns.last().map(|f| f.pre == Order::FnBarrier && f.nargs == 1).unwrap_or(false) {
Err(ErrorKind::BadArgument)
}
else {
Err(ErrorKind::NaExpression)
}
},
}
}
fn parse_op(&mut self, tok: &TokenKind) -> Result<(), ErrorKind> {
match tok {
TokenKind::Unk(_) => {
Err(ErrorKind::InvalidToken)
},
TokenKind::Lit(_) => {
Err(ErrorKind::ExpectOperator)
},
TokenKind::Op(op) => {
let desc = op.desc();
match desc.assoc {
Assoc::Left => self.eval_ge(desc.pre)?,
Assoc::Right => self.eval_gt(desc.pre)?,
};
self.fns.push(FnVal {
pfn: desc.pfn,
pre: desc.pre,
nargs: 2,
});
self.next = State::Value;
Ok(())
},
TokenKind::Var(_) => {
self.parse_op(&TokenKind::Op(Operator::IMul))?;
self.parse_val(tok)
},
TokenKind::Open(_) => {
self.parse_op(&TokenKind::Op(Operator::IMul))?;
self.parse_val(tok)
},
TokenKind::Comma => {
self.eval_gt(Order::FnBarrier)?;
self.fns.last_mut().ok_or(ErrorKind::MisplacedComma)?.nargs += 1;
self.next = State::Value;
Ok(())
},
TokenKind::Close => {
self.eval_gt(Order::FnBarrier)?;
self.eval_apply()?;
self.next = State::Operator;
Ok(())
},
}
}
fn eval_ge(&mut self, pre: Order) -> Result<(), ErrorKind> {
while self.fns.last().map(|f| f.pre >= pre).unwrap_or(false) {
self.eval_apply()?;
}
Ok(())
}
fn eval_gt(&mut self, pre: Order) -> Result<(), ErrorKind> {
while self.fns.last().map(|f| f.pre > pre).unwrap_or(false) {
self.eval_apply()?;
}
Ok(())
}
fn eval_apply(&mut self) -> Result<(), ErrorKind> {
if let Some(f) = self.fns.pop() {
if f.nargs as usize > self.vals.len() {
return Err(ErrorKind::InternalError);
}
let args = self.vals.len() - f.nargs as usize..;
let result = {
let vals = &mut self.vals[args.clone()];
(f.pfn)(self.env, vals)?
};
let _ = self.vals.drain(args.clone());
self.vals.push(result);
Ok(())
}
else {
Err(ErrorKind::UnbalancedParens)
}
}
}
pub fn eval(env: &dyn Env, input: &str) -> Result<Value, Error> {
let mut expr = Expr::new(env);
expr.feed(input)?;
expr.result()
}
pub fn eval_tokens(env: &dyn Env, tokens: &[Token]) -> Result<Value, Error> {
let mut expr = Expr::with_capacity(env, tokens.len() / 2 + 1);
for tok in tokens {
expr.parse(tok)?;
}
expr.result()
}
#[test]
fn basics() {
let env = BasicEnv::default();
assert_eq!(eval(&env, "2 + 3"), Ok(5.0));
assert_eq!(eval(&env, "2-3*4"), Ok(-10.0));
assert_eq!(eval(&env, "2*3+4"), Ok(10.0));
assert_eq!(eval(&env, "3^2-2"), Ok(7.0));
assert_eq!(eval(&env, "2+---2"), Ok(0.0));
assert_eq!(eval(&env, "-1"), Ok(-1.0));
assert_eq!(eval(&env, "-2^2 + 3*4 + sin(pi / 2)"), Ok(9.0));
}
#[test]
fn funcs() {
let env = BasicEnv::default();
assert_eq!(eval(&env, "2*(3+4)"), Ok(14.0));
assert_eq!(eval(&env, "mul(2,add(3,4))"), Ok(14.0));
}
#[test]
fn errors() {
let env = BasicEnv::default();
let err_kind = |input: &str| eval(&env, input).map_err(|e| e.kind);
assert_eq!(err_kind(""), Err(ErrorKind::UnfinishedExpression));
assert_eq!(err_kind("12 5"), Err(ErrorKind::ExpectOperator));
assert_eq!(err_kind(","), Err(ErrorKind::NaExpression));
assert_eq!(err_kind(")"), Err(ErrorKind::NaExpression));
assert_eq!(err_kind("*2"), Err(ErrorKind::DisallowedUnary));
assert_eq!(err_kind("2 +"), Err(ErrorKind::UnfinishedExpression));
assert_eq!(err_kind("(2"), Err(ErrorKind::UnbalancedParens));
assert_eq!(err_kind("(3))"), Err(ErrorKind::UnbalancedParens));
assert_eq!(err_kind("2,"), Err(ErrorKind::MisplacedComma));
assert_eq!(err_kind("pi()"), Err(ErrorKind::NameNotFound));
assert_eq!(err_kind("mean"), Err(ErrorKind::NameNotFound));
assert_eq!(err_kind("hello(5)"), Err(ErrorKind::NameNotFound));
assert_eq!(err_kind("hi"), Err(ErrorKind::NameNotFound));
}