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 --- Cargo.lock | 266 ++++++++++++++++++++++++++++++++ Cargo.toml | 1 + base.rules | 35 ++++- 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 ++++++++++++ test_rule | 1 - 10 files changed, 1297 insertions(+), 86 deletions(-) create mode 100644 src/format_latex.rs create mode 100644 src/simplify.rs delete mode 100644 test_rule diff --git a/Cargo.lock b/Cargo.lock index 4e6784c..6130586 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,19 +2,128 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "bitflags" +version = "2.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4512299f36f043ab09a583e57bceb5a5aab7a73db1805848e8fef3c9e8c78b3" + +[[package]] +name = "cfg-if" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" + +[[package]] +name = "cfg_aliases" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" + +[[package]] +name = "clipboard-win" +version = "5.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bde03770d3df201d4fb868f2c9c59e66a3e4e2bd06692a0fe701e7103c7e84d4" +dependencies = [ + "error-code", +] + [[package]] name = "difx" version = "0.1.0" dependencies = [ "nom", + "rustyline", +] + +[[package]] +name = "endian-type" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c34f04666d835ff5d62e058c3995147c06f42fe86ff053337632bca83e42702d" + +[[package]] +name = "errno" +version = "0.3.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + +[[package]] +name = "error-code" +version = "3.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dea2df4cf52843e0452895c455a1a2cfbb842a1e7329671acf418fdc53ed4c59" + +[[package]] +name = "fd-lock" +version = "4.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ce92ff622d6dadf7349484f42c93271a0d49b7cc4d466a936405bacbe10aa78" +dependencies = [ + "cfg-if", + "rustix", + "windows-sys 0.59.0", ] +[[package]] +name = "home" +version = "0.5.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cc627f471c528ff0c4a49e1d5e60450c8f6461dd6d10ba9dcd3a61d3dff7728d" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "libc" +version = "0.2.186" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" + +[[package]] +name = "linux-raw-sys" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" + +[[package]] +name = "log" +version = "0.4.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" + [[package]] name = "memchr" version = "2.7.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f52b00d39961fc5b2736ea853c9cc86238e165017a493d1d5c8eac6bdc4cc273" +[[package]] +name = "nibble_vec" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77a5d83df9f36fe23f0c3648c6bbb8b0298bb5f1939c8f2704431371f4b84d43" +dependencies = [ + "smallvec", +] + +[[package]] +name = "nix" +version = "0.29.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "71e2746dc3a24dd78b3cfcb7be93368c6de9963d30f43a6a73998a9cf4b17b46" +dependencies = [ + "bitflags", + "cfg-if", + "cfg_aliases", + "libc", +] + [[package]] name = "nom" version = "8.0.0" @@ -23,3 +132,160 @@ checksum = "df9761775871bdef83bee530e60050f7e54b1105350d6884eb0fb4f46c2f9405" dependencies = [ "memchr", ] + +[[package]] +name = "radix_trie" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c069c179fcdc6a2fe24d8d18305cf085fdbd4f922c041943e203685d6a1c58fd" +dependencies = [ + "endian-type", + "nibble_vec", +] + +[[package]] +name = "rustix" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6fe4565b9518b83ef4f91bb47ce29620ca828bd32cb7e408f0062e9930ba190" +dependencies = [ + "bitflags", + "errno", + "libc", + "linux-raw-sys", + "windows-sys 0.61.2", +] + +[[package]] +name = "rustyline" +version = "15.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2ee1e066dc922e513bda599c6ccb5f3bb2b0ea5870a579448f2622993f0a9a2f" +dependencies = [ + "bitflags", + "cfg-if", + "clipboard-win", + "fd-lock", + "home", + "libc", + "log", + "memchr", + "nix", + "radix_trie", + "unicode-segmentation", + "unicode-width", + "utf8parse", + "windows-sys 0.59.0", +] + +[[package]] +name = "smallvec" +version = "1.15.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03" + +[[package]] +name = "unicode-segmentation" +version = "1.13.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9629274872b2bfaf8d66f5f15725007f635594914870f65218920345aa11aa8c" + +[[package]] +name = "unicode-width" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4ac048d71ede7ee76d585517add45da530660ef4390e49b098733c6e897f254" + +[[package]] +name = "utf8parse" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" + +[[package]] +name = "windows-link" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" + +[[package]] +name = "windows-sys" +version = "0.59.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e38bc4d79ed67fd075bcc251a1c39b32a1776bbe92e5bef1f0bf1f8c531853b" +dependencies = [ + "windows-targets", +] + +[[package]] +name = "windows-sys" +version = "0.61.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc" +dependencies = [ + "windows-link", +] + +[[package]] +name = "windows-targets" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" +dependencies = [ + "windows_aarch64_gnullvm", + "windows_aarch64_msvc", + "windows_i686_gnu", + "windows_i686_gnullvm", + "windows_i686_msvc", + "windows_x86_64_gnu", + "windows_x86_64_gnullvm", + "windows_x86_64_msvc", +] + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" + +[[package]] +name = "windows_i686_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" + +[[package]] +name = "windows_i686_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" + +[[package]] +name = "windows_i686_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" diff --git a/Cargo.toml b/Cargo.toml index a566ccb..c46ae24 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -5,3 +5,4 @@ edition = "2024" [dependencies] nom = "8.0.0" +rustyline = "15.0.0" diff --git a/base.rules b/base.rules index aa14298..f5363bc 100644 --- a/base.rules +++ b/base.rules @@ -1,5 +1,34 @@ +("Trivial Base Cases") (def-rule (+ @a) @a) (def-rule (* @a) @a) -(def-rule (+ @a @a .) (+ (* 2 @a) .)) -(def-rule (+ @a (* $i @a) .) (+ (* (+ $i 1) @a) .)) -(def-rule (+ (* $i @a) (* $j @a) .) (+ (* (+ $i $j) @a) .)) +(def-rule (+ 0 ..v) (+ ..v)) +(def-rule (* 1 ..v) (* ..v)) +(def-rule (* 0 ..v) 0) + +("Flatten argument lists") +(def-rule (+ (+ ..a) ..b) (+ ..a ..b)) +(def-rule (* (* ..a) ..b) (* ..a ..b)) + +("Pull out constant subexpressions") +(def-rule (+ $i $j ..a) (+ (+ $i $j) ..a)) +(def-rule (* $i $j ..a) (* (* $i $j) ..a)) + +("Algebraic simplifications") +(def-rule (+ @a (* $i @a) ..v) (+ (* (+ $i 1) @a) ..v)) +(def-rule (+ (* $i @a) (* $j @a) ..v) (+ (* (+ $i $j) @a) ..v)) +(def-rule (+ @a @a ..v) (+ (* 2 @a) ..v)) + +("Derivatives") +(def-rule (d/dx $i) 0) +(def-rule (d/dx x) 1) + +(def-rule (d/dx (sin @a)) (* (cos @a) (d/dx @a))) +(def-rule (d/dx (cos @a)) (* -1 (sin @a) (d/dx @a))) +(def-rule (d/dx (* -1 (sin @a))) (* -1 (cos @a) (d/dx @a))) +(def-rule (d/dx (* -1 (cos @a))) (* (sin x) (d/dx @a))) + +(def-rule (d/dx (* $i @a ..v)) (* $i (d/dx (* @a ..v)))) +(def-rule (d/dx (+ @a @b ..v)) (+ (d/dx @a) (d/dx @b) (d/dx (+ ..v)))) +(def-rule (d/dx (+ @a ..v)) (+ (d/dx @a) (d/dx (+ ..v)))) +(def-rule (d/dx (* @a @b)) (+ (* @a (d/dx @b)) (* @b (d/dx @a)))) +(def-rule (d/dx (* @a @b @c)) (+ (* (d/dx @a) @b @c) (* @a (d/dx @b) @c) (* @a @b (d/dx @c)))) 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); + } +} diff --git a/test_rule b/test_rule deleted file mode 100644 index 5658665..0000000 --- a/test_rule +++ /dev/null @@ -1 +0,0 @@ -(def-rule (+ 3 (* $i @a) .) _) -- cgit v1.3.1