use std::collections::HashMap;
pub type FnMap = HashMap<String, Box<dyn Fn(&[f64]) -> f64>>;
#[derive(Debug)]
pub enum EvalError {
UnknownFunction(String),
}
impl std::fmt::Display for EvalError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
EvalError::UnknownFunction(func) => write!(f, "Unknown function: {}", func),
}
}
}
#[derive(Debug, Clone, Copy)]
pub enum OperatorKind {
Plus,
Minus,
Multiply,
Divide,
}
impl std::fmt::Display for OperatorKind {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let op_str = match self {
OperatorKind::Plus => "+",
OperatorKind::Minus => "-",
OperatorKind::Multiply => "*",
OperatorKind::Divide => "/",
};
write!(f, "{op_str}")
}
}
#[derive(Clone)]
pub enum AstNode {
Number(f64),
Variable(String),
FunctionCall {
name: String,
args: Vec<AstNode>,
},
BinaryOp {
op: OperatorKind,
left: Box<AstNode>,
right: Box<AstNode>,
},
}
impl AstNode {
pub fn eval(&self, get_var: &dyn Fn(&str) -> f64, fn_map: &FnMap) -> Result<f64, EvalError> {
match self {
AstNode::Number(n) => Ok(*n),
AstNode::Variable(var) => Ok(get_var(var)),
AstNode::BinaryOp { op, left, right } => {
let l = left.eval(get_var, fn_map)?;
let r = right.eval(get_var, fn_map)?;
match op {
OperatorKind::Plus => Ok(l + r),
OperatorKind::Minus => Ok(l - r),
OperatorKind::Multiply => Ok(l * r),
OperatorKind::Divide => Ok(l / r),
}
}
AstNode::FunctionCall { name, args } => {
let evaluated_args: Result<Vec<f64>, _> =
args.iter().map(|arg| arg.eval(get_var, fn_map)).collect();
let evaluated_args = evaluated_args?;
if let Some(f) = fn_map.get(name) {
Ok(f(&evaluated_args))
} else {
Err(EvalError::UnknownFunction(name.clone()))
}
}
}
}
}
pub fn default_fn_map() -> FnMap {
let mut map: FnMap = HashMap::new();
map.insert("sin".to_string(), Box::new(|args: &[f64]| args[0].sin()));
map
}
impl std::fmt::Display for AstNode {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
AstNode::Number(n) => write!(f, "{n}"),
AstNode::Variable(var) => write!(f, "{var}"),
AstNode::BinaryOp { op, left, right } => {
write!(f, "({left} {op} {right})")
}
AstNode::FunctionCall { name, args } => {
let args_str = args
.iter()
.map(|arg| format!("{arg}"))
.collect::<Vec<String>>()
.join(", ");
write!(f, "{name}({args_str})")
}
}
}
}
#[derive(Debug)]
pub struct ParserError {
pub message: String,
}
impl std::fmt::Display for ParserError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "ParserError: {}", self.message)
}
}