math-parser-rs 0.1.0

A simple handwritten dsl for interpreting math.
Documentation
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)
    }
}