diff options
Diffstat (limited to 'src/simplify.rs')
| -rw-r--r-- | src/simplify.rs | 101 |
1 files changed, 101 insertions, 0 deletions
diff --git a/src/simplify.rs b/src/simplify.rs new file mode 100644 index 0000000..ee51fb9 --- /dev/null +++ b/src/simplify.rs @@ -0,0 +1,101 @@ +use crate::sexpr::{Atom, Expr, Number}; + +fn simplify_once(expr: Expr) -> Expr { + match expr { + Expr::Atom(_) => expr, + Expr::Application(items) => { + let simplified_items: Vec<Expr> = items.into_iter().map(simplify_once).collect(); + + if simplified_items.is_empty() { + return Expr::Application(simplified_items); + } + + match &simplified_items[0] { + Expr::Atom(Atom::Builtin(op)) if op == "+" => { + let mut sum = 0.0; + for arg in &simplified_items[1..] { + if let Expr::Atom(Atom::Number(i)) = arg { + sum += i.0; + } else { + return Expr::Application(simplified_items); + } + } + Expr::Atom(Atom::Number(Number(sum))) + } + Expr::Atom(Atom::Builtin(op)) if op == "*" => { + let mut product = 1.0; + for arg in &simplified_items[1..] { + if let Expr::Atom(Atom::Number(i)) = arg { + product *= i.0; + } else { + return Expr::Application(simplified_items); + } + } + Expr::Atom(Atom::Number(Number(product))) + } + _ => Expr::Application(simplified_items), + } + } + Expr::OrderedList(items) => { + Expr::OrderedList(items.into_iter().map(simplify_once).collect()) + } + } +} + +pub fn simplify(expr: Expr) -> Expr { + let mut current = expr; + + loop { + let next = simplify_once(current.clone()); + if next == current { + return next; + } + current = next; + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn num(i: f64) -> Expr { + Expr::Atom(Atom::Number(Number(i))) + } + + fn builtin(name: &str) -> Expr { + Expr::Atom(Atom::Builtin(name.to_string())) + } + + fn app(items: Vec<Expr>) -> Expr { + Expr::Application(items) + } + + #[test] + fn simplifies_addition_of_integers() { + let expr = app(vec![builtin("+"), num(2.0), num(3.0), num(4.0)]); + assert_eq!(simplify(expr), num(9.0)); + } + + #[test] + fn simplifies_multiplication_of_integers() { + let expr = app(vec![builtin("*"), num(2.0), num(3.0), num(4.0)]); + assert_eq!(simplify(expr), num(24.0)); + } + + #[test] + fn simplifies_nested_numeric_forms() { + let expr = app(vec![ + builtin("+"), + app(vec![builtin("*"), num(2.0), num(3.0)]), + app(vec![builtin("+"), num(4.0), num(5.0)]), + ]); + + assert_eq!(simplify(expr), num(15.0)); + } + + #[test] + fn does_not_simplify_when_non_integer_argument_present() { + let expr = app(vec![builtin("+"), num(2.0), builtin("x")]); + assert_eq!(simplify(expr.clone()), expr); + } +} |
