use crate::Inner;
use crate::node::Node;
use std::cmp::Ordering;
const MAX_DEPTH: u32 = 10_000;
fn family_of(inner: &Inner, n: &Node) -> u8 {
match n {
Node::Int(_) | Node::Rat(_) => 0,
Node::Float { .. } => 1,
Node::Sym(_) => 2,
Node::Pow { base, .. } => match &inner.nodes[*base as usize] {
Node::Sym(_) => 2,
_ => 4,
},
Node::Fn { .. } => 3,
Node::Mul { .. } => 5,
Node::Add { .. } => 6,
}
}
fn is_exact(n: &Node) -> bool {
matches!(n, Node::Int(_) | Node::Rat(_))
}
fn cmp_exact(a: &Node, b: &Node) -> Ordering {
match (a, b) {
(Node::Int(x), Node::Int(y)) => x.cmp(y),
(Node::Rat(x), Node::Rat(y)) => x.cmp(y),
(Node::Int(x), Node::Rat(y)) => y.cmp_integer(x).reverse(),
(Node::Rat(x), Node::Int(y)) => x.cmp_integer(y),
_ => unreachable!("精确数配对"),
}
}
impl Inner {
pub(crate) fn cmp_ids(&self, a: u32, b: u32) -> Ordering {
self.cmp_at(a, b, 0)
}
fn cmp_at(&self, a: u32, b: u32, depth: u32) -> Ordering {
assert!(depth <= MAX_DEPTH, "表达式嵌套超过 {MAX_DEPTH} 层");
let na = &self.nodes[a as usize];
let nb = &self.nodes[b as usize];
if is_exact(na) && is_exact(nb) {
return cmp_exact(na, nb);
}
if is_exact(na) {
return Ordering::Less;
}
if is_exact(nb) {
return Ordering::Greater;
}
let (fa, fb) = (family_of(self, na), family_of(self, nb));
if fa != fb {
return fa.cmp(&fb);
}
match (na, nb) {
(Node::Float { bits: x, .. }, Node::Float { bits: y, .. }) => {
f64::from_bits(*x).total_cmp(&f64::from_bits(*y))
}
(Node::Sym(x), Node::Sym(y)) => {
self.sym_names[*x as usize].cmp(&self.sym_names[*y as usize])
}
(Node::Sym(x), Node::Pow { base, .. }) => match &self.nodes[*base as usize] {
Node::Sym(t) => self.sym_names[*x as usize]
.cmp(&self.sym_names[*t as usize])
.then(Ordering::Less),
_ => unreachable!("复合底幂不在符号桶"),
},
(Node::Pow { base, .. }, Node::Sym(y)) => match &self.nodes[*base as usize] {
Node::Sym(t) => self.sym_names[*t as usize]
.cmp(&self.sym_names[*y as usize])
.then(Ordering::Greater),
_ => unreachable!("复合底幂不在符号桶"),
},
(Node::Pow { base: b1, exp: e1 }, Node::Pow { base: b2, exp: e2 }) => {
self.cmp_at(*b1, *b2, depth + 1)
.then_with(|| self.cmp_at(*e1, *e2, depth + 1))
}
(Node::Fn { head: h, args: sa }, Node::Fn { head: k, args: sb }) => self.fn_names
[*h as usize]
.cmp(&self.fn_names[*k as usize])
.then_with(|| self.cmp_seq(self.node_args(*sa), self.node_args(*sb), depth)),
(Node::Mul { args: sa }, Node::Mul { args: sb })
| (Node::Add { args: sa }, Node::Add { args: sb }) => {
self.cmp_seq(self.node_args(*sa), self.node_args(*sb), depth)
}
_ => unreachable!("同桶必有同类配对"),
}
}
fn cmp_seq(&self, a: &[u32], b: &[u32], depth: u32) -> Ordering {
for (&x, &y) in a.iter().zip(b.iter()) {
let o = self.cmp_at(x, y, depth + 1);
if o != Ordering::Equal {
return o;
}
}
a.len().cmp(&b.len())
}
}
pub(crate) fn deep_eq(x: &Inner, a: u32, y: &Inner, b: u32) -> bool {
deep_eq_at(x, a, y, b, 0)
}
fn deep_eq_at(x: &Inner, a: u32, y: &Inner, b: u32, depth: u32) -> bool {
assert!(depth <= MAX_DEPTH, "表达式嵌套超过 {MAX_DEPTH} 层");
match (&x.nodes[a as usize], &y.nodes[b as usize]) {
(Node::Int(u), Node::Int(v)) => u == v,
(Node::Rat(u), Node::Rat(v)) => u == v,
(Node::Float { bits: u, .. }, Node::Float { bits: v, .. }) => u == v,
(Node::Sym(u), Node::Sym(v)) => x.sym_names[*u as usize] == y.sym_names[*v as usize],
(Node::Fn { head: h, args: sa }, Node::Fn { head: k, args: sb }) => {
x.fn_names[*h as usize] == y.fn_names[*k as usize]
&& seq_eq(x, y, x.node_args(*sa), y.node_args(*sb), depth)
}
(Node::Pow { base: b1, exp: e1 }, Node::Pow { base: b2, exp: e2 }) => {
deep_eq_at(x, *b1, y, *b2, depth + 1) && deep_eq_at(x, *e1, y, *e2, depth + 1)
}
(Node::Mul { args: sa }, Node::Mul { args: sb })
| (Node::Add { args: sa }, Node::Add { args: sb }) => {
seq_eq(x, y, x.node_args(*sa), y.node_args(*sb), depth)
}
_ => false,
}
}
fn seq_eq(x: &Inner, y: &Inner, a: &[u32], b: &[u32], depth: u32) -> bool {
a.len() == b.len()
&& a.iter()
.zip(b.iter())
.all(|(&p, &q)| deep_eq_at(x, p, y, q, depth + 1))
}