From 3a03bdd7e7a2dbe1d83c33f5fbe8e94d6bac9f71 Mon Sep 17 00:00:00 2001 From: Collin Williams <96917990+bluedragon1221@users.noreply.github.com> Date: Sun, 14 Jun 2026 16:41:52 -0500 Subject: messy code --- src/format_latex.rs | 194 +++++++++++++++++++++++ src/lib.rs | 4 +- src/main.rs | 132 ++++++++++++++-- src/match_rules.rs | 432 +++++++++++++++++++++++++++++++++++++++++++++++----- src/sexpr.rs | 217 ++++++++++++++++++++++---- src/simplify.rs | 101 ++++++++++++ 6 files changed, 998 insertions(+), 82 deletions(-) create mode 100644 src/format_latex.rs create mode 100644 src/simplify.rs (limited to 'src') diff --git a/src/format_latex.rs b/src/format_latex.rs new file mode 100644 index 0000000..d754cb0 --- /dev/null +++ b/src/format_latex.rs @@ -0,0 +1,194 @@ +use crate::sexpr::{Atom, Expr}; + +const PREC_EQUAL: u8 = 10; +const PREC_ADD: u8 = 20; +const PREC_MUL: u8 = 30; +const PREC_POW: u8 = 40; +const PREC_ATOM: u8 = 50; + +pub fn to_latex(expr: &Expr) -> String { + render_expr(expr, 0) +} + +fn render_expr(expr: &Expr, parent_prec: u8) -> String { + let (rendered, prec) = render_expr_inner(expr); + if prec < parent_prec { + format!("\\left({}\\right)", rendered) + } else { + rendered + } +} + +fn render_expr_inner(expr: &Expr) -> (String, u8) { + match expr { + Expr::Atom(atom) => (render_atom(atom), PREC_ATOM), + Expr::OrderedList(items) => { + let pieces: Vec = items.iter().map(|item| render_expr(item, 0)).collect(); + (format!("\\left[{}\\right]", pieces.join(", ")), PREC_ATOM) + } + Expr::Application(items) => render_application(items), + } +} + +fn render_application(items: &[Expr]) -> (String, u8) { + let Some(Expr::Atom(Atom::Builtin(head))) = items.first() else { + let pieces: Vec = items.iter().map(|item| render_expr(item, 0)).collect(); + return (format!("\\left({}\\right)", pieces.join(" ")), PREC_ATOM); + }; + + let args = &items[1..]; + + match head.as_str() { + "+" => { + let pieces: Vec = args.iter().map(|arg| render_expr(arg, PREC_ADD)).collect(); + (pieces.join(" + "), PREC_ADD) + } + "*" => { + let pieces: Vec = args.iter().map(|arg| render_expr(arg, PREC_MUL)).collect(); + (pieces.join(" \\cdot "), PREC_MUL) + } + "=" => { + let pieces: Vec = args.iter().map(|arg| render_expr(arg, PREC_EQUAL)).collect(); + (pieces.join(" = "), PREC_EQUAL) + } + "sin" => unary_func("\\sin", args), + "cos" => unary_func("\\cos", args), + "d/dx" => { + if args.len() == 1 { + ( + format!( + "\\frac{{d}}{{dx}}\\left({}\\right)", + render_expr(&args[0], 0) + ), + PREC_ATOM, + ) + } else { + fallback_call(head, args) + } + } + _ => { + if let Some(power) = power_from_head(head) { + if args.len() == 1 { + let base = render_expr(&args[0], PREC_POW); + (format!("{}^{{{}}}", base, power), PREC_POW) + } else { + fallback_call(head, args) + } + } else { + fallback_call(head, args) + } + } + } +} + +fn unary_func(name: &str, args: &[Expr]) -> (String, u8) { + if args.len() == 1 { + (format!("{}\\left({}\\right)", name, render_expr(&args[0], 0)), PREC_ATOM) + } else { + fallback_call(name, args) + } +} + +fn fallback_call(head: &str, args: &[Expr]) -> (String, u8) { + let pieces: Vec = args.iter().map(|arg| render_expr(arg, 0)).collect(); + ( + format!( + "{}\\left({}\\right)", + latex_symbol(head), + pieces.join(", ") + ), + PREC_ATOM, + ) +} + +fn power_from_head(head: &str) -> Option<&str> { + head.strip_prefix('^') + .filter(|suffix| !suffix.is_empty() && suffix.chars().all(|c| c.is_ascii_digit())) +} + +fn render_atom(atom: &Atom) -> String { + match atom { + Atom::Number(n) => n.to_string(), + Atom::Builtin(name) => latex_symbol(name), + Atom::Variable(var) => latex_text(&format!("{}{}", var.r#type.r#type_prefix(), var.index)), + } +} + +fn latex_symbol(input: &str) -> String { + input.replace('_', "\\_") +} + +fn latex_text(input: &str) -> String { + let escaped = input + .replace('\\', "\\textbackslash{}") + .replace('{', "\\{") + .replace('}', "\\}") + .replace('$', "\\$") + .replace('_', "\\_") + .replace('&', "\\&") + .replace('#', "\\#") + .replace('%', "\\%") + .replace('^', "\\^") + .replace('~', "\\~{}"); + format!("\\text{{{}}}", escaped) +} + +trait VariablePrefix { + fn r#type_prefix(&self) -> &'static str; +} + +impl VariablePrefix for crate::sexpr::VariableType { + fn r#type_prefix(&self) -> &'static str { + match self { + crate::sexpr::VariableType::Number => "$", + crate::sexpr::VariableType::Expr => "@", + crate::sexpr::VariableType::NonNumberExpr => "!", + crate::sexpr::VariableType::Ellipsis => "..", + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::sexpr::{Number, Variable, VariableType}; + + fn num(n: f64) -> Expr { + Expr::Atom(Atom::Number(Number(n))) + } + + fn sym(s: &str) -> Expr { + Expr::Atom(Atom::Builtin(s.to_string())) + } + + fn app(items: Vec) -> Expr { + Expr::Application(items) + } + + #[test] + fn renders_basic_arithmetic() { + let expr = app(vec![sym("+"), num(2.0), app(vec![sym("*"), num(3.0), sym("x")])]); + assert_eq!(to_latex(&expr), "2 + 3 \\cdot x"); + } + + #[test] + fn renders_derivative_and_power() { + let expr = app(vec![sym("d/dx"), app(vec![sym("^2"), sym("x")])]); + assert_eq!(to_latex(&expr), "\\frac{d}{dx}\\left(x^{2}\\right)"); + } + + #[test] + fn renders_ordered_list() { + let expr = Expr::OrderedList(vec![sym("x"), num(1.0)]); + assert_eq!(to_latex(&expr), "\\left[x, 1\\right]"); + } + + #[test] + fn renders_variable_atom() { + let var = Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::NonNumberExpr, + index: "a".to_string(), + })); + assert_eq!(to_latex(&var), "\\text{!a}"); + } +} diff --git a/src/lib.rs b/src/lib.rs index 1bb3e41..fa9fcb2 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,3 +1,5 @@ -pub mod sexpr; pub mod def_rules; +pub mod format_latex; pub mod match_rules; +pub mod sexpr; +pub mod simplify; diff --git a/src/main.rs b/src/main.rs index 911948d..aae5122 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,29 +1,129 @@ -use difx::sexpr::Parser; +use difx::def_rules::Rule; use difx::match_rules::{self, Bindings}; +use difx::simplify; +use difx::sexpr::{Expr, Parser}; +use rustyline::DefaultEditor; +use rustyline::error::ReadlineError; + +const MAX_REWRITE_ITERS: usize = 100; + +fn parse_expr(input: &str) -> Option { + std::panic::catch_unwind(|| { + let mut parser = Parser::new(input); + parser.parse_one() + }) + .ok() +} + +fn rewrite_node(rules: &[Rule], expr: &Expr) -> Option<(usize, Expr)> { + for (idx, rule) in rules.iter().enumerate() { + let mut bindings = Bindings::default(); + if match_rules::matches(&rule.lhs, expr, &mut bindings) { + let rewritten = match_rules::substitute(&rule.rhs, &bindings); + return Some((idx, rewritten)); + } + } + + None +} + +fn rewrite_top_down_once(rules: &[Rule], expr: &Expr) -> Option<(usize, Expr)> { + if let Some(hit) = rewrite_node(rules, expr) { + return Some(hit); + } + + match expr { + Expr::Atom(_) => None, + Expr::Application(items) => { + for (idx, item) in items.iter().enumerate() { + if let Some((rule_idx, rewritten_child)) = rewrite_top_down_once(rules, item) { + let mut rewritten_items = items.clone(); + rewritten_items[idx] = rewritten_child; + return Some((rule_idx, Expr::Application(rewritten_items))); + } + } + + None + } + Expr::OrderedList(items) => { + for (idx, item) in items.iter().enumerate() { + if let Some((rule_idx, rewritten_child)) = rewrite_top_down_once(rules, item) { + let mut rewritten_items = items.clone(); + rewritten_items[idx] = rewritten_child; + return Some((rule_idx, Expr::OrderedList(rewritten_items))); + } + } + + None + } + } +} fn main() { - let input = include_str!("../test_rule"); + let input = include_str!("../base.rules"); let rules = difx::def_rules::parse_def_rules(input.to_string()); - // println!("{:#?}", rules); - let pat = rules[0].clone().lhs; - println!("{:?}", pat); + println!("Loaded {} rule(s) from base.rules", rules.len()); + + let mut rl = DefaultEditor::new().expect("failed to initialize line editor"); loop { - let mut input = String::new(); - std::io::stdin() - .read_line(&mut input) - .unwrap(); + let line = match rl.readline("> ") { + Ok(line) => line, + Err(ReadlineError::Interrupted) => { + println!("^C"); + continue; + } + Err(ReadlineError::Eof) => { + println!(); + break; + } + Err(err) => { + eprintln!("Readline error: {err}"); + break; + } + }; - let mut parser = Parser::new(&input); - let parsed = parser.parse_one(); + let line = line.trim(); + if line.is_empty() { + continue; + } - let mut bindings = Bindings::default(); + if let Err(err) = rl.add_history_entry(line) { + eprintln!("history error: {err}"); + } + + let Some(mut expr) = parse_expr(line) else { + println!("Parse error"); + continue; + }; - if match_rules::matches(&pat, &parsed, &mut bindings) { - println!("{:?}", &bindings); - println!("MATCH"); + let mut applied = 0usize; + while applied < MAX_REWRITE_ITERS { + let Some((idx, rewritten_expr)) = rewrite_top_down_once(&rules, &expr) else { + break; + }; + let next_expr = simplify::simplify(rewritten_expr); + + let rule = &rules[idx]; + println!("{}.", applied + 1); + println!("```lisp\n{} => {}\n```", rule.lhs, rule.rhs); + println!("> {}", next_expr); + + expr = next_expr; + applied += 1; + } + + if applied == 0 { + println!("NO MATCH"); + } else { + println!("final: {}", expr); + } + + let stopped = applied == MAX_REWRITE_ITERS; + if applied == MAX_REWRITE_ITERS { + let stop_msg = format!("Stopped after {} iterations (safety limit).", MAX_REWRITE_ITERS); + println!("{}", stop_msg); } } } - diff --git a/src/match_rules.rs b/src/match_rules.rs index 938ee08..fe58119 100644 --- a/src/match_rules.rs +++ b/src/match_rules.rs @@ -1,20 +1,20 @@ use std::collections::HashMap; -use crate::sexpr::{Atom, Expr, Variable, VariableType}; +use crate::sexpr::{Atom, Expr, Number, Variable, VariableType}; -#[derive(Debug, PartialEq, Eq, PartialOrd, Ord)] +#[derive(Debug, PartialEq, Eq, PartialOrd, Ord, Clone)] pub enum BindValue { - Integer(i64), + Number(Number), Expr(Expr), Ellipsis(Vec) } -#[derive(Default)] +#[derive(Default, Clone)] pub struct Bindings<'b>(HashMap<&'b str, BindValue>); impl<'b> Bindings<'b> { - pub fn check_or_insert(&mut self, index: &'b String, bind_value: BindValue) -> bool { - if let Some(existing) = self.0.get(index.as_str()) { + pub fn check_or_insert(&mut self, index: &'b str, bind_value: BindValue) -> bool { + if let Some(existing) = self.0.get(index) { existing == &bind_value } else { self.0.insert(index, bind_value); @@ -29,45 +29,152 @@ impl std::fmt::Debug for Bindings<'_> { } } -pub fn args_match<'b>(p_args: &'b [Expr], args: &'b [Expr], bindings: &mut Bindings<'b>) -> bool { - const ELLIPSIS: Expr = Expr::Atom(Atom::Ellipsis); +pub fn substitute(template: &Expr, bindings: &Bindings<'_>) -> Expr { + match template { + Expr::Atom(Atom::Variable(variable)) => { + if let Some(value) = bindings.0.get(variable.index.as_str()) { + match (&variable.r#type, value) { + (VariableType::Number, BindValue::Number(i)) => Expr::Atom(Atom::Number(*i)), + (VariableType::Number, BindValue::Expr(Expr::Atom(Atom::Number(i)))) => { + Expr::Atom(Atom::Number(*i)) + } + (VariableType::Expr | VariableType::NonNumberExpr, BindValue::Expr(expr)) => { + expr.clone() + } + _ => template.clone(), + } + } else { + template.clone() + } + } + Expr::Atom(_) => template.clone(), + Expr::Application(items) => { + let mut rewritten = Vec::new(); + + for item in items { + if let Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Ellipsis, + index, + })) = item + { + if let Some(BindValue::Ellipsis(exprs)) = bindings.0.get(index.as_str()) { + rewritten.extend(exprs.clone()); + } + } else { + rewritten.push(substitute(item, bindings)); + } + } + + Expr::Application(rewritten) + } + Expr::OrderedList(items) => { + Expr::OrderedList(items.iter().map(|item| substitute(item, bindings)).collect()) + } + } +} + +fn match_with_ellipsis<'b>( + patterns: &[&'b Expr], + remaining: Vec<&'b Expr>, + bindings: &Bindings<'b>, +) -> Option<(Bindings<'b>, Vec<&'b Expr>)> { + if patterns.is_empty() { + return Some((bindings.clone(), remaining)); + } - if p_args.contains(&ELLIPSIS) { - let mut sorted_p: Vec<&Expr> = p_args.iter() - .filter(|p| **p != ELLIPSIS) - .collect(); - - let mut sorted_a: Vec<&Expr> = args.iter().collect(); + let p = patterns[0]; - if sorted_p.len() != sorted_a.len() { - return false; + for (idx, a) in remaining.iter().enumerate() { + let mut trial_bindings = bindings.clone(); + + if matches(p, a, &mut trial_bindings) { + let mut next_remaining = remaining.clone(); + next_remaining.remove(idx); + + if let Some((final_bindings, final_remaining)) = + match_with_ellipsis(&patterns[1..], next_remaining, &trial_bindings) + { + return Some((final_bindings, final_remaining)); + } } + } - sorted_p.sort(); - sorted_a.sort(); + None +} - std::iter::zip(sorted_p, sorted_a).all(|(p, a)| { - matches(p, a, bindings) - }) +fn is_ellipsis_pattern(p: &Expr) -> Option<&str> { + if let Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Ellipsis, + index, + })) = p + { + Some(index.as_str()) } else { - if p_args.len() != args.len() { - return false; + None + } +} + +fn ellipsis_name(p_args: &[Expr]) -> Result, ()> { + let mut found: Option<&str> = None; + + for p in p_args { + if let Some(name) = is_ellipsis_pattern(p) { + if found.is_some() { + return Err(()); + } + found = Some(name); + } + } + + Ok(found) +} + +pub fn args_match<'b>(p_args: &'b [Expr], args: &'b [Expr], bindings: &mut Bindings<'b>) -> bool { + match ellipsis_name(p_args) { + Err(()) => false, + Ok(Some(name)) => { + let patterns: Vec<&Expr> = p_args.iter().filter(|p| is_ellipsis_pattern(p).is_none()).collect(); + let remaining: Vec<&Expr> = args.iter().collect(); + + if patterns.len() > remaining.len() { + return false; + } + + if let Some((mut final_bindings, final_remaining)) = + match_with_ellipsis(&patterns, remaining, bindings) + { + final_bindings.0.insert( + name, + BindValue::Ellipsis(final_remaining.into_iter().cloned().collect()), + ); + *bindings = final_bindings; + true + } else { + false + } + } + Ok(None) => { + if p_args.len() != args.len() { + return false; + } + std::iter::zip(p_args.iter(), args.iter()).all(|(a, b)| { + matches(a, b, bindings) + }) } - std::iter::zip(p_args.iter(), args.iter()).all(|(a, b)| { - matches(a, b, bindings) - }) } } pub fn matches<'b>(p: &'b Expr, expr: &'b Expr, bindings: &mut Bindings<'b>) -> bool { match expr { - Expr::Atom(Atom::Int(i)) => { + Expr::Atom(Atom::Number(i)) => { match p { - Expr::Atom(Atom::Int(pi)) => i == pi, - Expr::Atom(Atom::Variable(Variable { r#type: e @ (VariableType::Integer | VariableType::Expr), index })) => { - bindings.check_or_insert(index, match e { - VariableType::Integer => BindValue::Integer(*i), - VariableType::Expr => BindValue::Expr(Expr::Atom(Atom::Int(*i))) + Expr::Atom(Atom::Number(pi)) => i == pi, + Expr::Atom(Atom::Variable(Variable { r#type: e, index })) => { + bindings.check_or_insert(index.as_str(), match e { + VariableType::Number => BindValue::Number(*i), + VariableType::Expr => BindValue::Expr(Expr::Atom(Atom::Number(*i))), + VariableType::NonNumberExpr => return false, + VariableType::Ellipsis => return false, }) } _ => false @@ -76,8 +183,11 @@ pub fn matches<'b>(p: &'b Expr, expr: &'b Expr, bindings: &mut Bindings<'b>) -> Expr::Atom(Atom::Builtin(s)) => { match p { Expr::Atom(Atom::Builtin(ps)) => s == ps, - Expr::Atom(Atom::Variable(Variable { r#type: VariableType::Expr, index })) => { - bindings.check_or_insert(index, BindValue::Expr(Expr::Atom(Atom::Builtin(s.to_string())))) + Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Expr | VariableType::NonNumberExpr, + index, + })) => { + bindings.check_or_insert(index.as_str(), BindValue::Expr(Expr::Atom(Atom::Builtin(s.to_string())))) } _ => false } @@ -89,15 +199,263 @@ pub fn matches<'b>(p: &'b Expr, expr: &'b Expr, bindings: &mut Bindings<'b>) -> f == pf && args_match(&p_exprs[1..], &exprs[1..], bindings) } else if exprs.len() == 1 { matches(p, &exprs[0], bindings) - } else if let Expr::Atom(Atom::Variable(Variable { r#type: VariableType::Expr, index })) = p { - bindings.check_or_insert(index, BindValue::Expr(expr.clone())) + } else if let Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Expr | VariableType::NonNumberExpr, + index, + })) = p { + bindings.check_or_insert(index.as_str(), BindValue::Expr(expr.clone())) } else { false } }, - Expr::Atom(Atom::Variable(_)) | Expr::Atom(Atom::Ellipsis) => { + Expr::OrderedList(exprs) => { + if let Expr::OrderedList(p_exprs) = p { + if p_exprs.len() != exprs.len() { + return false; + } + + std::iter::zip(p_exprs.iter(), exprs.iter()).all(|(p_expr, expr)| { + matches(p_expr, expr, bindings) + }) + } else if let Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Expr | VariableType::NonNumberExpr, + index, + })) = p { + bindings.check_or_insert(index.as_str(), BindValue::Expr(expr.clone())) + } else { + false + } + } + Expr::Atom(Atom::Variable(_)) => { // These should only appear in patterns, not in expressions being matched false } } } + +#[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 var_expr(name: &str) -> Expr { + Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Expr, + index: name.to_string(), + })) + } + + fn var_int(name: &str) -> Expr { + Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Number, + index: name.to_string(), + })) + } + + fn var_ellipsis(name: &str) -> Expr { + Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::Ellipsis, + index: name.to_string(), + })) + } + + fn var_non_number(name: &str) -> Expr { + Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::NonNumberExpr, + index: name.to_string(), + })) + } + + fn app(items: Vec) -> Expr { + Expr::Application(items) + } + + #[test] + fn ellipsis_allows_extra_arguments() { + let pattern = app(vec![ + builtin("+"), + num(3.0), + app(vec![builtin("*"), var_int("i"), var_expr("a")]), + var_ellipsis("rest"), + ]); + + let expr = app(vec![ + builtin("+"), + num(3.0), + app(vec![builtin("*"), num(7.0), builtin("foo")]), + builtin("bar"), + builtin("baz"), + ]); + + let mut bindings = Bindings::default(); + assert!(matches(&pattern, &expr, &mut bindings)); + } + + #[test] + fn ellipsis_matching_backtracks_for_bindings() { + let pattern = app(vec![ + builtin("f"), + var_expr("x"), + app(vec![builtin("g"), var_expr("x")]), + var_ellipsis("rest"), + ]); + + let expr = app(vec![ + builtin("f"), + app(vec![builtin("g"), builtin("a")]), + builtin("a"), + builtin("tail"), + ]); + + let mut bindings = Bindings::default(); + assert!(matches(&pattern, &expr, &mut bindings)); + } + + #[test] + fn substitute_replaces_expression_and_integer_variables() { + let pattern = app(vec![builtin("*"), var_int("i"), var_expr("x")]); + let expr = app(vec![builtin("*"), num(9.0), builtin("foo")]); + let mut bindings = Bindings::default(); + + assert!(matches(&pattern, &expr, &mut bindings)); + + let rhs = app(vec![builtin("+"), var_int("i"), var_expr("x")]); + let rewritten = substitute(&rhs, &bindings); + + assert_eq!(rewritten, app(vec![builtin("+"), num(9.0), builtin("foo")])); + } + + #[test] + fn substitute_splices_ellipsis_tail() { + let pattern = app(vec![ + builtin("+"), + app(vec![builtin("*"), var_int("i"), var_expr("a")]), + app(vec![builtin("*"), var_int("j"), var_expr("a")]), + var_ellipsis("rest"), + ]); + + let expr = app(vec![ + builtin("+"), + app(vec![builtin("*"), num(3.0), builtin("a")]), + app(vec![builtin("*"), num(2.0), builtin("a")]), + num(7.0), + ]); + + let mut bindings = Bindings::default(); + assert!(matches(&pattern, &expr, &mut bindings)); + + let rhs = app(vec![ + builtin("+"), + app(vec![ + builtin("*"), + app(vec![builtin("+"), var_int("i"), var_int("j")]), + var_expr("a"), + ]), + var_ellipsis("rest"), + ]); + + let rewritten = substitute(&rhs, &bindings); + + assert_eq!( + rewritten, + app(vec![ + builtin("+"), + app(vec![ + builtin("*"), + app(vec![builtin("+"), num(3.0), num(2.0)]), + builtin("a"), + ]), + num(7.0), + ]) + ); + } + + #[test] + fn named_ellipsis_captures_unmatched_arguments() { + let pattern = app(vec![ + builtin("+"), + app(vec![builtin("*"), var_int("i"), var_expr("a")]), + app(vec![builtin("*"), var_int("j"), var_expr("a")]), + var_ellipsis("rest"), + ]); + + let expr = app(vec![ + builtin("+"), + builtin("x"), + app(vec![builtin("*"), num(3.0), builtin("a")]), + builtin("y"), + app(vec![builtin("*"), num(2.0), builtin("a")]), + num(7.0), + ]); + + let mut bindings = Bindings::default(); + assert!(matches(&pattern, &expr, &mut bindings)); + + let rhs = app(vec![ + builtin("+"), + app(vec![ + builtin("*"), + app(vec![builtin("+"), var_int("i"), var_int("j")]), + var_expr("a"), + ]), + var_ellipsis("rest"), + ]); + + let rewritten = substitute(&rhs, &bindings); + assert_eq!( + rewritten, + app(vec![ + builtin("+"), + app(vec![ + builtin("*"), + app(vec![builtin("+"), num(3.0), num(2.0)]), + builtin("a"), + ]), + builtin("x"), + builtin("y"), + num(7.0), + ]) + ); + } + + #[test] + fn duplicate_named_ellipsis_is_rejected() { + let pattern = app(vec![builtin("+"), var_ellipsis("rest"), var_ellipsis("rest")]); + let expr = app(vec![builtin("+"), num(1.0), num(2.0)]); + + let mut bindings = Bindings::default(); + assert!(!matches(&pattern, &expr, &mut bindings)); + } + + #[test] + fn non_number_var_matches_symbol_and_application() { + let symbol_pattern = app(vec![builtin("*"), var_non_number("a"), var_non_number("a")]); + let symbol_expr = app(vec![builtin("*"), builtin("x"), builtin("x")]); + let mut bindings = Bindings::default(); + assert!(matches(&symbol_pattern, &symbol_expr, &mut bindings)); + + let app_expr = app(vec![ + builtin("*"), + app(vec![builtin("+"), builtin("x"), num(1.0)]), + app(vec![builtin("+"), builtin("x"), num(1.0)]), + ]); + let mut bindings = Bindings::default(); + assert!(matches(&symbol_pattern, &app_expr, &mut bindings)); + } + + #[test] + fn non_number_var_does_not_match_number_atom() { + let pattern = app(vec![builtin("*"), var_non_number("a"), var_non_number("a")]); + let expr = app(vec![builtin("*"), num(2.0), num(2.0)]); + + let mut bindings = Bindings::default(); + assert!(!matches(&pattern, &expr, &mut bindings)); + } +} diff --git a/src/sexpr.rs b/src/sexpr.rs index 557bc5d..dbea2f8 100644 --- a/src/sexpr.rs +++ b/src/sexpr.rs @@ -1,7 +1,47 @@ use std::cmp::Ordering; +use std::hash::{Hash, Hasher}; use std::iter::Peekable; use std::str::{Chars, FromStr}; +#[derive(Debug, Clone, Copy)] +pub struct Number(pub f64); + +impl PartialEq for Number { + fn eq(&self, other: &Self) -> bool { + self.0.to_bits() == other.0.to_bits() + } +} + +impl Eq for Number {} + +impl PartialOrd for Number { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } +} + +impl Ord for Number { + fn cmp(&self, other: &Self) -> Ordering { + self.0.total_cmp(&other.0) + } +} + +impl Hash for Number { + fn hash(&self, state: &mut H) { + self.0.to_bits().hash(state); + } +} + +impl std::fmt::Display for Number { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + if self.0.fract() == 0.0 { + write!(f, "{:.0}", self.0) + } else { + write!(f, "{}", self.0) + } + } +} + #[derive(Debug, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)] pub struct Variable { pub r#type: VariableType, @@ -11,9 +51,31 @@ pub struct Variable { impl FromStr for Variable { type Err = (); fn from_str(s: &str) -> Result { + if let Some(index) = s.strip_prefix("..") { + if index.is_empty() { + return Err(()); + } + + return Ok(Variable { + r#type: VariableType::Ellipsis, + index: index.to_string(), + }); + } + + if let Some(index) = s.strip_prefix('!') { + if index.is_empty() { + return Err(()); + } + + return Ok(Variable { + r#type: VariableType::NonNumberExpr, + index: index.to_string(), + }); + } + match s.chars().nth(0) { Some('$') => Ok(Variable { - r#type: VariableType::Integer, + r#type: VariableType::Number, index: s[1..].to_string() }), Some('@') => Ok(Variable { @@ -27,16 +89,17 @@ impl FromStr for Variable { #[derive(Debug, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)] pub enum VariableType { - Integer, - Expr + Number, + Expr, + NonNumberExpr, + Ellipsis, } #[derive(Debug, PartialEq, Clone, Eq, Hash)] pub enum Atom { - Int(i64), + Number(Number), Builtin(String), Variable(Variable), - Ellipsis, } impl PartialOrd for Atom { @@ -48,22 +111,17 @@ impl PartialOrd for Atom { impl Ord for Atom { fn cmp(&self, other: &Self) -> Ordering { match (self, other) { - (Atom::Int(i1), Atom::Int(i2)) => i1.cmp(i2), - (Atom::Int(_), _) => Ordering::Less, - (_, Atom::Int(_)) => Ordering::Greater, + (Atom::Number(i1), Atom::Number(i2)) => i1.cmp(i2), + (Atom::Number(_), _) => Ordering::Less, + (_, Atom::Number(_)) => Ordering::Greater, (Atom::Builtin(b1), Atom::Builtin(b2)) => b1.cmp(b2), (Atom::Builtin(_), Atom::Variable(_)) => Ordering::Less, - (Atom::Builtin(_), Atom::Ellipsis) => Ordering::Less, - + (Atom::Variable(v1), Atom::Variable(v2)) => { v1.index.cmp(&v2.index) } (Atom::Variable(_), Atom::Builtin(_)) => Ordering::Greater, - (Atom::Variable(_), Atom::Ellipsis) => Ordering::Less, - - (Atom::Ellipsis, Atom::Ellipsis) => Ordering::Equal, - (Atom::Ellipsis, _) => Ordering::Greater, } } } @@ -71,16 +129,17 @@ impl Ord for Atom { impl std::fmt::Display for Atom { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { - Atom::Int(i) => write!(f, "{}", i), + Atom::Number(i) => write!(f, "{}", i), Atom::Builtin(b) => write!(f, "{}", b), Atom::Variable(v) => write!(f, "{}{}", match v.r#type { - VariableType::Integer => "$", + VariableType::Number => "$", VariableType::Expr => "@", + VariableType::NonNumberExpr => "!", + VariableType::Ellipsis => "..", }, v.index ), - Atom::Ellipsis => write!(f, "."), } } } @@ -89,6 +148,7 @@ impl std::fmt::Display for Atom { pub enum Expr { Atom(Atom), Application(Vec), + OrderedList(Vec), } impl PartialOrd for Expr { @@ -100,13 +160,19 @@ impl PartialOrd for Expr { impl Ord for Expr { fn cmp(&self, other: &Self) -> Ordering { match (self, other) { - (Expr::Atom(_), Expr::Application(_)) => Ordering::Less, - (Expr::Application(_), Expr::Atom(_)) => Ordering::Greater, + (Expr::Atom(_), Expr::Application(_) | Expr::OrderedList(_)) => Ordering::Less, + (Expr::Application(_) | Expr::OrderedList(_), Expr::Atom(_)) => Ordering::Greater, + (Expr::Application(_), Expr::OrderedList(_)) => Ordering::Less, + (Expr::OrderedList(_), Expr::Application(_)) => Ordering::Greater, (Expr::Atom(a1), Expr::Atom(a2)) => a1.cmp(a2), (Expr::Application(v1), Expr::Application(v2)) => { v1.len().cmp(&v2.len()) .then_with(|| v1.cmp(v2)) } + (Expr::OrderedList(v1), Expr::OrderedList(v2)) => { + v1.len().cmp(&v2.len()) + .then_with(|| v1.cmp(v2)) + } } } } @@ -119,6 +185,10 @@ impl std::fmt::Display for Expr { let pieces: Vec = args.into_iter().map(ToString::to_string).collect(); write!(f, "({})", pieces.join(" ")) } + Expr::OrderedList(args) => { + let pieces: Vec = args.into_iter().map(ToString::to_string).collect(); + write!(f, "[{}]", pieces.join(" ")) + } } } } @@ -133,6 +203,71 @@ pub struct Parser<'a> { input: Peekable>, } +#[cfg(test)] +mod tests { + use super::{Atom, Expr, Number, Parser, Variable, VariableType}; + + fn builtin(name: &str) -> Expr { + Expr::Atom(Atom::Builtin(name.to_string())) + } + + fn number(i: f64) -> Expr { + Expr::Atom(Atom::Number(Number(i))) + } + + #[test] + fn parses_ordered_list() { + let mut parser = Parser::new("[a b 3]"); + let parsed = parser.parse_one(); + + assert_eq!( + parsed, + Expr::OrderedList(vec![builtin("a"), builtin("b"), number(3.0)]) + ); + } + + #[test] + fn parses_nested_ordered_list_and_application() { + let mut parser = Parser::new("(+ [a b] (* [1 2] 3))"); + let parsed = parser.parse_one(); + + assert_eq!( + parsed, + Expr::Application(vec![ + builtin("+"), + Expr::OrderedList(vec![builtin("a"), builtin("b")]), + Expr::Application(vec![ + builtin("*"), + Expr::OrderedList(vec![number(1.0), number(2.0)]), + number(3.0), + ]), + ]) + ); + } + + #[test] + fn parses_decimal_number() { + let mut parser = Parser::new("1.25"); + let parsed = parser.parse_one(); + + assert_eq!(parsed, number(1.25)); + } + + #[test] + fn parses_non_number_variable() { + let mut parser = Parser::new("!term"); + let parsed = parser.parse_one(); + + assert_eq!( + parsed, + Expr::Atom(Atom::Variable(Variable { + r#type: VariableType::NonNumberExpr, + index: "term".to_string(), + })) + ); + } +} + impl<'a> Parser<'a> { pub fn new(input: &'a str) -> Self { Parser { @@ -156,6 +291,7 @@ impl<'a> Parser<'a> { match self.input.peek() { Some('(') => ParseResult::Expr(self.parse_list()), + Some('[') => ParseResult::Expr(self.parse_ordered_list()), Some(_) => ParseResult::Expr(self.parse_atom()), None => ParseResult::Eof, } @@ -206,11 +342,41 @@ impl<'a> Parser<'a> { Expr::Application(expressions) } + fn parse_ordered_list(&mut self) -> Expr { + if self.input.next() != Some('[') { + panic!("Parser error: Expected '['"); + } + + let mut expressions = Vec::new(); + + loop { + self.skip_whitespace(); + + match self.input.peek() { + Some(']') => { + self.input.next(); + break; + } + None => panic!("Parser error: Unexpected EOF, unclosed ordered list"), + Some(_) => { + match self.parse_one_optional() { + ParseResult::Expr(expr) => expressions.push(expr), + ParseResult::Eof => panic!( + "Parser error: Unexpected EOF after space inside ordered list" + ), + } + } + } + } + + Expr::OrderedList(expressions) + } + fn parse_atom(&mut self) -> Expr { let mut buffer = String::new(); while let Some(&c) = self.input.peek() { - if c.is_whitespace() || c == '(' || c == ')' { + if c.is_whitespace() || c == '(' || c == ')' || c == '[' || c == ']' { break; } buffer.push(c); @@ -221,19 +387,14 @@ impl<'a> Parser<'a> { panic!("Parser error: Tried to parse an atom, but it was empty."); } - // Check for ellipsis - if buffer == "." { - return Expr::Atom(Atom::Ellipsis); - } - // Check for variable if let Ok(v) = buffer.parse::() { return Expr::Atom(Atom::Variable(v)); } - // Check for integer - if let Ok(i) = buffer.parse::() { - return Expr::Atom(Atom::Int(i)); + // Check for number + if let Ok(i) = buffer.parse::() { + return Expr::Atom(Atom::Number(Number(i))); } // Default to builtin 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 = 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::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); + } +} -- cgit v1.3.1