summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--Cargo.lock266
-rw-r--r--Cargo.toml1
-rw-r--r--base.rules35
-rw-r--r--src/format_latex.rs194
-rw-r--r--src/lib.rs4
-rw-r--r--src/main.rs132
-rw-r--r--src/match_rules.rs432
-rw-r--r--src/sexpr.rs217
-rw-r--r--src/simplify.rs101
-rw-r--r--test_rule1
10 files changed, 1297 insertions, 86 deletions
diff --git a/Cargo.lock b/Cargo.lock
index 4e6784c..6130586 100644
--- a/Cargo.lock
+++ b/Cargo.lock
@@ -3,19 +3,128 @@
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"
source = "registry+https://github.com/rust-lang/crates.io-index"
@@ -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<String> = 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<String> = 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<String> = args.iter().map(|arg| render_expr(arg, PREC_ADD)).collect();
+ (pieces.join(" + "), PREC_ADD)
+ }
+ "*" => {
+ let pieces: Vec<String> = args.iter().map(|arg| render_expr(arg, PREC_MUL)).collect();
+ (pieces.join(" \\cdot "), PREC_MUL)
+ }
+ "=" => {
+ let pieces: Vec<String> = 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<String> = 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 {
+ 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<Expr> {
+ 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<Expr>)
}
-#[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<Option<&str>, ()> {
+ 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 {
+ 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<Ordering> {
+ 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<H: Hasher>(&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<Self, Self::Err> {
+ 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<Expr>),
+ OrderedList(Vec<Expr>),
}
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<String> = args.into_iter().map(ToString::to_string).collect();
write!(f, "({})", pieces.join(" "))
}
+ Expr::OrderedList(args) => {
+ let pieces: Vec<String> = args.into_iter().map(ToString::to_string).collect();
+ write!(f, "[{}]", pieces.join(" "))
+ }
}
}
}
@@ -133,6 +203,71 @@ pub struct Parser<'a> {
input: Peekable<Chars<'a>>,
}
+#[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::<Variable>() {
return Expr::Atom(Atom::Variable(v));
}
- // Check for integer
- if let Ok(i) = buffer.parse::<i64>() {
- return Expr::Atom(Atom::Int(i));
+ // Check for number
+ if let Ok(i) = buffer.parse::<f64>() {
+ 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<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);
+ }
+}
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) .) _)