use std::fmt;
use std::hash::Hasher;
use smallvec::SmallVec;
use crate::base::node::{ExprId, ExprNode, NumId, SymbolId};
const RANK_NUM: u8 = 0;
const RANK_SYMBOL: u8 = 20;
const RANK_POW: u8 = 40;
const RANK_MUL: u8 = 60;
const RANK_ADD: u8 = 80;
const RANK_FUNCTION: u8 = 100;
const RANK_RELATIONAL: u8 = 110;
const RANK_AND: u8 = 112;
const RANK_OR: u8 = 114;
const RANK_NOT: u8 = 116;
const RANK_PIECEWISE: u8 = 118;
const RANK_DERIVATIVE: u8 = 120;
const RANK_MIN: u8 = 130;
const RANK_MAX: u8 = 132;
const RANK_INTEGRAL: u8 = 140;
const RANK_DEFINITE_INTEGRAL: u8 = 142;
const RANK_SUM: u8 = 150;
const RANK_PRODUCT: u8 = 152;
const RANK_LIMIT: u8 = 154;
const RANK_SERIES: u8 = 155;
const RANK_LAPLACE_TRANSFORM: u8 = 156;
const RANK_INV_LAPLACE_TRANSFORM: u8 = 157;
const RANK_RESIDUE: u8 = 158;
const RANK_ROOTOF: u8 = 160;
const RANK_DSOLVE: u8 = 162;
const RANK_ROOTSUM: u8 = 163;
const RANK_CONDITION_SET: u8 = 164;
const RANK_CONSTANT: u8 = 170;
const RANK_SET: u8 = 190;
const RANK_SPECIAL: u8 = 210;
const FN_SIN: u8 = 0;
const FN_COS: u8 = 1;
const FN_TAN: u8 = 2;
const FN_EXP: u8 = 3;
const FN_LN: u8 = 4;
const FN_ABS: u8 = 6;
const FN_ASIN: u8 = 8;
const FN_ACOS: u8 = 9;
const FN_ATAN: u8 = 10;
const FN_SINH: u8 = 11;
const FN_COSH: u8 = 12;
const FN_TANH: u8 = 13;
const FN_ASINH: u8 = 14;
const FN_ACOSH: u8 = 15;
const FN_ATANH: u8 = 16;
const FN_APPLY: u8 = 17;
const FN_SIGN: u8 = 18;
const FN_ATAN2: u8 = 19;
const FN_FLOOR: u8 = 20;
const FN_CEILING: u8 = 21;
const FN_GAMMA: u8 = 22;
const FN_LOG_GAMMA: u8 = 23;
const FN_DIGAMMA: u8 = 24;
const FN_ERF: u8 = 25;
const FN_ERFC: u8 = 26;
const FN_BETA: u8 = 27;
const FN_HEAVISIDE: u8 = 28;
const FN_DIRAC_DELTA: u8 = 29;
const FN_LAMBERT_W: u8 = 30;
const FN_RE: u8 = 31;
const FN_IM: u8 = 32;
const FN_CONJUGATE: u8 = 33;
const FN_ARG: u8 = 34;
const FN_SI: u8 = 35;
const FN_CI: u8 = 36;
const FN_EI: u8 = 37;
const FN_LI: u8 = 38;
const FN_ZETA: u8 = 39;
const FN_POLYGAMMA: u8 = 40;
const FN_KRONECKER_DELTA: u8 = 41;
const CONST_PI: u8 = 0;
const CONST_E: u8 = 1;
const CONST_IMAGINARY_UNIT: u8 = 2;
const CONST_PHYSICAL: u8 = 3;
const CONST_EULER_GAMMA: u8 = 4;
const CONST_CATALAN: u8 = 5;
const CONST_GOLDEN_RATIO: u8 = 6;
const CONST_BOOL_TRUE: u8 = 10;
const CONST_BOOL_FALSE: u8 = 11;
const SPECIAL_INFINITY: u8 = 0;
const SPECIAL_NEG_INFINITY: u8 = 1;
const SPECIAL_COMPLEX_INFINITY: u8 = 2;
const SPECIAL_NAN: u8 = 3;
const SPECIAL_NEG: u8 = 4;
const SET_EMPTY: u8 = 0;
const SET_UNIVERSAL: u8 = 1;
const SET_INTERVAL: u8 = 2;
const SET_FINITE_SET: u8 = 3;
const SET_UNION: u8 = 4;
const SET_INTERSECTION: u8 = 5;
const SET_COMPLEMENT: u8 = 6;
pub const MAX_KEY_BYTES: usize = 256;
#[derive(Clone, PartialEq, Eq, Hash)]
pub struct SortKey(SmallVec<[u8; 24]>);
impl SortKey {
#[inline]
fn new() -> Self {
SortKey(SmallVec::new())
}
#[inline]
fn push(&mut self, byte: u8) {
self.0.push(byte);
}
#[inline]
fn extend(&mut self, bytes: &[u8]) {
self.0.extend_from_slice(bytes);
}
#[inline]
pub fn as_bytes(&self) -> &[u8] {
&self.0
}
#[inline]
pub fn is_truncated(&self) -> bool {
self.0.len() > MAX_KEY_BYTES
}
fn bound(mut self) -> Self {
if self.0.len() > MAX_KEY_BYTES {
let mut h = rustc_hash::FxHasher::default();
h.write(&self.0);
let digest = h.finish().to_be_bytes();
self.0.truncate(MAX_KEY_BYTES);
self.0.extend_from_slice(&digest);
}
self
}
}
impl PartialOrd for SortKey {
#[inline]
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for SortKey {
#[inline]
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.0.as_slice().cmp(other.0.as_slice())
}
}
impl fmt::Debug for SortKey {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "SortKey({:?})", self.0.as_slice())
}
}
pub fn compute_sort_key(
node: &ExprNode,
get_key: impl Fn(ExprId) -> SortKey,
get_num_bytes: impl Fn(NumId) -> Vec<u8>,
get_sym_name: impl Fn(SymbolId) -> String,
) -> SortKey {
let mut key = SortKey::new();
match node {
ExprNode::Num(id) => {
key.push(RANK_NUM);
key.extend(&get_num_bytes(*id));
}
ExprNode::Symbol(id) => {
key.push(RANK_SYMBOL);
key.extend(get_sym_name(*id).as_bytes());
}
ExprNode::Add(children) => {
key.push(RANK_ADD);
for &child in children {
key.extend(get_key(child).as_bytes());
}
}
ExprNode::Mul(children) => {
key.push(RANK_MUL);
for &child in children {
key.extend(get_key(child).as_bytes());
}
}
ExprNode::Pow(base, exp) => {
key.push(RANK_POW);
key.extend(get_key(*base).as_bytes());
key.extend(get_key(*exp).as_bytes());
}
ExprNode::Derivative(body, var) => {
key.push(RANK_DERIVATIVE);
key.extend(get_key(*body).as_bytes());
key.extend(get_key(*var).as_bytes());
}
ExprNode::Integral(body, var) => {
key.push(RANK_INTEGRAL);
key.extend(get_key(*body).as_bytes());
key.extend(get_key(*var).as_bytes());
}
ExprNode::DefiniteIntegral(body, var, lo, hi) => {
key.push(RANK_DEFINITE_INTEGRAL);
key.extend(get_key(*body).as_bytes());
key.extend(get_key(*var).as_bytes());
key.extend(get_key(*lo).as_bytes());
key.extend(get_key(*hi).as_bytes());
}
ExprNode::Sin(x) => {
key.push(RANK_FUNCTION);
key.push(FN_SIN);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Cos(x) => {
key.push(RANK_FUNCTION);
key.push(FN_COS);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Tan(x) => {
key.push(RANK_FUNCTION);
key.push(FN_TAN);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Exp(x) => {
key.push(RANK_FUNCTION);
key.push(FN_EXP);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Ln(x) => {
key.push(RANK_FUNCTION);
key.push(FN_LN);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Abs(x) => {
key.push(RANK_FUNCTION);
key.push(FN_ABS);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Asin(x) => {
key.push(RANK_FUNCTION);
key.push(FN_ASIN);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Acos(x) => {
key.push(RANK_FUNCTION);
key.push(FN_ACOS);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Atan(x) => {
key.push(RANK_FUNCTION);
key.push(FN_ATAN);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Atan2(y, x) => {
key.push(RANK_FUNCTION);
key.push(FN_ATAN2);
key.extend(get_key(*y).as_bytes());
key.extend(get_key(*x).as_bytes());
}
ExprNode::Sinh(x) => {
key.push(RANK_FUNCTION);
key.push(FN_SINH);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Cosh(x) => {
key.push(RANK_FUNCTION);
key.push(FN_COSH);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Tanh(x) => {
key.push(RANK_FUNCTION);
key.push(FN_TANH);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Asinh(x) => {
key.push(RANK_FUNCTION);
key.push(FN_ASINH);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Acosh(x) => {
key.push(RANK_FUNCTION);
key.push(FN_ACOSH);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Atanh(x) => {
key.push(RANK_FUNCTION);
key.push(FN_ATANH);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Sign(x) => {
key.push(RANK_FUNCTION);
key.push(FN_SIGN);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Heaviside(x) => {
key.push(RANK_FUNCTION);
key.push(FN_HEAVISIDE);
key.extend(get_key(*x).as_bytes());
}
ExprNode::DiracDelta(x) => {
key.push(RANK_FUNCTION);
key.push(FN_DIRAC_DELTA);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Gamma(x) => {
key.push(RANK_FUNCTION);
key.push(FN_GAMMA);
key.extend(get_key(*x).as_bytes());
}
ExprNode::LogGamma(x) => {
key.push(RANK_FUNCTION);
key.push(FN_LOG_GAMMA);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Digamma(x) => {
key.push(RANK_FUNCTION);
key.push(FN_DIGAMMA);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Erf(x) => {
key.push(RANK_FUNCTION);
key.push(FN_ERF);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Erfc(x) => {
key.push(RANK_FUNCTION);
key.push(FN_ERFC);
key.extend(get_key(*x).as_bytes());
}
ExprNode::LambertW(x) => {
key.push(RANK_FUNCTION);
key.push(FN_LAMBERT_W);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Re(x) => {
key.push(RANK_FUNCTION);
key.push(FN_RE);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Im(x) => {
key.push(RANK_FUNCTION);
key.push(FN_IM);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Conjugate(x) => {
key.push(RANK_FUNCTION);
key.push(FN_CONJUGATE);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Arg(x) => {
key.push(RANK_FUNCTION);
key.push(FN_ARG);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Si(x) => {
key.push(RANK_FUNCTION);
key.push(FN_SI);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Ci(x) => {
key.push(RANK_FUNCTION);
key.push(FN_CI);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Ei(x) => {
key.push(RANK_FUNCTION);
key.push(FN_EI);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Li(x) => {
key.push(RANK_FUNCTION);
key.push(FN_LI);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Zeta(x) => {
key.push(RANK_FUNCTION);
key.push(FN_ZETA);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Polygamma(n, x) => {
key.push(RANK_FUNCTION);
key.push(FN_POLYGAMMA);
key.extend(get_key(*n).as_bytes());
key.extend(get_key(*x).as_bytes());
}
ExprNode::KroneckerDelta(i, j) => {
key.push(RANK_FUNCTION);
key.push(FN_KRONECKER_DELTA);
key.extend(get_key(*i).as_bytes());
key.extend(get_key(*j).as_bytes());
}
ExprNode::Beta(a, b) => {
key.push(RANK_FUNCTION);
key.push(FN_BETA);
key.extend(get_key(*a).as_bytes());
key.extend(get_key(*b).as_bytes());
}
ExprNode::Floor(x) => {
key.push(RANK_FUNCTION);
key.push(FN_FLOOR);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Ceiling(x) => {
key.push(RANK_FUNCTION);
key.push(FN_CEILING);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Min(children) => {
key.push(RANK_MIN);
for &child in children {
key.extend(get_key(child).as_bytes());
}
}
ExprNode::Max(children) => {
key.push(RANK_MAX);
for &child in children {
key.extend(get_key(child).as_bytes());
}
}
ExprNode::Sum(body, var, lo, hi) => {
key.push(RANK_SUM);
key.extend(get_key(*body).as_bytes());
key.extend(get_key(*var).as_bytes());
key.extend(get_key(*lo).as_bytes());
key.extend(get_key(*hi).as_bytes());
}
ExprNode::Product_(body, var, lo, hi) => {
key.push(RANK_PRODUCT);
key.extend(get_key(*body).as_bytes());
key.extend(get_key(*var).as_bytes());
key.extend(get_key(*lo).as_bytes());
key.extend(get_key(*hi).as_bytes());
}
ExprNode::Apply(sym, args) => {
key.push(RANK_FUNCTION);
key.push(FN_APPLY);
key.extend(get_sym_name(*sym).as_bytes());
key.push(0x00);
for &arg in args {
key.extend(get_key(arg).as_bytes());
}
}
ExprNode::Pi => {
key.push(RANK_CONSTANT);
key.push(CONST_PI);
}
ExprNode::E => {
key.push(RANK_CONSTANT);
key.push(CONST_E);
}
ExprNode::ImaginaryUnit => {
key.push(RANK_CONSTANT);
key.push(CONST_IMAGINARY_UNIT);
}
ExprNode::EulerGamma => {
key.push(RANK_CONSTANT);
key.push(CONST_EULER_GAMMA);
}
ExprNode::Catalan => {
key.push(RANK_CONSTANT);
key.push(CONST_CATALAN);
}
ExprNode::GoldenRatio => {
key.push(RANK_CONSTANT);
key.push(CONST_GOLDEN_RATIO);
}
ExprNode::PhysicalConstant(name_id, _) => {
key.push(RANK_CONSTANT);
key.push(CONST_PHYSICAL);
key.extend(get_sym_name(*name_id).as_bytes());
}
ExprNode::Infinity => {
key.push(RANK_SPECIAL);
key.push(SPECIAL_INFINITY);
}
ExprNode::NegInfinity => {
key.push(RANK_SPECIAL);
key.push(SPECIAL_NEG_INFINITY);
}
ExprNode::ComplexInfinity => {
key.push(RANK_SPECIAL);
key.push(SPECIAL_COMPLEX_INFINITY);
}
ExprNode::NaN => {
key.push(RANK_SPECIAL);
key.push(SPECIAL_NAN);
}
ExprNode::Neg(x) => {
key.push(RANK_SPECIAL);
key.push(SPECIAL_NEG);
key.extend(get_key(*x).as_bytes());
}
ExprNode::Factorial(x) => {
key.push(RANK_FUNCTION);
key.push(FN_ABS + 1); key.extend(get_key(*x).as_bytes());
}
ExprNode::Binomial(n, k) => {
key.push(RANK_FUNCTION);
key.push(FN_ABS + 2);
key.extend(get_key(*n).as_bytes());
key.extend(get_key(*k).as_bytes());
}
ExprNode::BoolTrue => {
key.push(RANK_CONSTANT);
key.push(CONST_BOOL_TRUE);
}
ExprNode::BoolFalse => {
key.push(RANK_CONSTANT);
key.push(CONST_BOOL_FALSE);
}
ExprNode::Gt(lhs, rhs) => {
key.push(RANK_RELATIONAL);
key.push(0); key.extend(get_key(*lhs).as_bytes());
key.extend(get_key(*rhs).as_bytes());
}
ExprNode::Ge(lhs, rhs) => {
key.push(RANK_RELATIONAL);
key.push(1);
key.extend(get_key(*lhs).as_bytes());
key.extend(get_key(*rhs).as_bytes());
}
ExprNode::Eq_(lhs, rhs) => {
key.push(RANK_RELATIONAL);
key.push(2);
key.extend(get_key(*lhs).as_bytes());
key.extend(get_key(*rhs).as_bytes());
}
ExprNode::Ne(lhs, rhs) => {
key.push(RANK_RELATIONAL);
key.push(3);
key.extend(get_key(*lhs).as_bytes());
key.extend(get_key(*rhs).as_bytes());
}
ExprNode::And(children) => {
key.push(RANK_AND);
for &c in children {
key.extend(get_key(c).as_bytes());
}
}
ExprNode::Or(children) => {
key.push(RANK_OR);
for &c in children {
key.extend(get_key(c).as_bytes());
}
}
ExprNode::Not(inner) => {
key.push(RANK_NOT);
key.extend(get_key(*inner).as_bytes());
}
ExprNode::Piecewise(children) => {
key.push(RANK_PIECEWISE);
for &(val, cond) in children {
key.extend(get_key(val).as_bytes());
key.extend(get_key(cond).as_bytes());
}
}
ExprNode::EmptySet => {
key.push(RANK_SET);
key.push(SET_EMPTY);
}
ExprNode::UniversalSet => {
key.push(RANK_SET);
key.push(SET_UNIVERSAL);
}
ExprNode::Interval(start, end, flags) => {
key.push(RANK_SET);
key.push(SET_INTERVAL);
key.push(*flags);
key.extend(get_key(*start).as_bytes());
key.extend(get_key(*end).as_bytes());
}
ExprNode::FiniteSet(elems) => {
key.push(RANK_SET);
key.push(SET_FINITE_SET);
for &elem in elems {
key.extend(get_key(elem).as_bytes());
}
}
ExprNode::SetUnion(sets) => {
key.push(RANK_SET);
key.push(SET_UNION);
for &s in sets {
key.extend(get_key(s).as_bytes());
}
}
ExprNode::SetIntersection(sets) => {
key.push(RANK_SET);
key.push(SET_INTERSECTION);
for &s in sets {
key.extend(get_key(s).as_bytes());
}
}
ExprNode::SetComplement(a, b) => {
key.push(RANK_SET);
key.push(SET_COMPLEMENT);
key.extend(get_key(*a).as_bytes());
key.extend(get_key(*b).as_bytes());
}
ExprNode::Limit(body, var, point) => {
key.push(RANK_LIMIT);
key.extend(get_key(*body).as_bytes());
key.extend(get_key(*var).as_bytes());
key.extend(get_key(*point).as_bytes());
}
ExprNode::Series(body, var, point, order) => {
key.push(RANK_SERIES);
key.extend(get_key(*body).as_bytes());
key.extend(get_key(*var).as_bytes());
key.extend(get_key(*point).as_bytes());
key.extend(get_key(*order).as_bytes());
}
ExprNode::LaplaceTransform(body, t, s) => {
key.push(RANK_LAPLACE_TRANSFORM);
key.extend(get_key(*body).as_bytes());
key.extend(get_key(*t).as_bytes());
key.extend(get_key(*s).as_bytes());
}
ExprNode::InverseLaplaceTransform(body, s, t) => {
key.push(RANK_INV_LAPLACE_TRANSFORM);
key.extend(get_key(*body).as_bytes());
key.extend(get_key(*s).as_bytes());
key.extend(get_key(*t).as_bytes());
}
ExprNode::Residue(body, var, point) => {
key.push(RANK_RESIDUE);
key.extend(get_key(*body).as_bytes());
key.extend(get_key(*var).as_bytes());
key.extend(get_key(*point).as_bytes());
}
ExprNode::RootOf(poly, index) => {
key.push(RANK_ROOTOF);
key.extend(get_key(*poly).as_bytes());
key.extend(get_key(*index).as_bytes());
}
ExprNode::DSolve(expr, func, var) => {
key.push(RANK_DSOLVE);
key.extend(get_key(*expr).as_bytes());
key.extend(get_key(*func).as_bytes());
key.extend(get_key(*var).as_bytes());
}
ExprNode::RootSum(poly, body, sumvar) => {
key.push(RANK_ROOTSUM);
key.extend(get_key(*poly).as_bytes());
key.extend(get_key(*body).as_bytes());
key.extend(get_key(*sumvar).as_bytes());
}
ExprNode::ConditionSet(var, cond) => {
key.push(RANK_CONDITION_SET);
key.extend(get_key(*var).as_bytes());
key.extend(get_key(*cond).as_bytes());
}
}
key.bound()
}
#[cfg(test)]
mod tests {
use super::*;
use smallvec::smallvec;
fn atom_key(node: &ExprNode) -> SortKey {
compute_sort_key(
node,
|_| unreachable!("atom should not query children"),
|_| vec![0x42],
|id| format!("s{}", id.0),
)
}
#[test]
fn numbers_sort_before_symbols() {
let num_key = atom_key(&ExprNode::Num(NumId(0)));
let sym_key = atom_key(&ExprNode::Symbol(SymbolId(0)));
assert!(num_key < sym_key, "Num should sort before Symbol");
}
#[test]
fn symbols_sort_alphabetically() {
let key_a = compute_sort_key(
&ExprNode::Symbol(SymbolId(0)),
|_| unreachable!(),
|_| unreachable!(),
|_| "alpha".to_string(),
);
let key_b = compute_sort_key(
&ExprNode::Symbol(SymbolId(1)),
|_| unreachable!(),
|_| unreachable!(),
|_| "beta".to_string(),
);
assert!(key_a < key_b);
}
#[test]
fn constants_sort_after_composites() {
let add_key = compute_sort_key(
&ExprNode::Add(smallvec![ExprId(0)]),
|_| SortKey(SmallVec::from_slice(&[RANK_NUM, 0x01])),
|_| unreachable!(),
|_| unreachable!(),
);
let pi_key = atom_key(&ExprNode::Pi);
assert!(add_key < pi_key, "Add should sort before Pi");
}
#[test]
fn constant_sub_ranks_are_distinct() {
let pi = atom_key(&ExprNode::Pi);
let e = atom_key(&ExprNode::E);
let i = atom_key(&ExprNode::ImaginaryUnit);
assert!(pi < e);
assert!(e < i);
}
#[test]
fn special_values_sort_last() {
let pi_key = atom_key(&ExprNode::Pi);
let inf_key = atom_key(&ExprNode::Infinity);
assert!(pi_key < inf_key, "Constants should sort before specials");
}
#[test]
fn function_discriminants_differ() {
let sin_key = compute_sort_key(
&ExprNode::Sin(ExprId(0)),
|_| SortKey(SmallVec::from_slice(&[RANK_NUM, 0x01])),
|_| unreachable!(),
|_| unreachable!(),
);
let cos_key = compute_sort_key(
&ExprNode::Cos(ExprId(0)),
|_| SortKey(SmallVec::from_slice(&[RANK_NUM, 0x01])),
|_| unreachable!(),
|_| unreachable!(),
);
assert_ne!(sin_key, cos_key);
assert!(
sin_key < cos_key,
"Sin (discriminant 0) < Cos (discriminant 1)"
);
}
#[test]
fn sort_key_ord_is_consistent() {
let a = SortKey(SmallVec::from_slice(&[10, 20]));
let b = SortKey(SmallVec::from_slice(&[10, 30]));
let c = SortKey(SmallVec::from_slice(&[20]));
assert!(a < b);
assert!(b < c);
assert!(a < c);
}
#[test]
fn neg_is_special_rank() {
let neg_key = compute_sort_key(
&ExprNode::Neg(ExprId(0)),
|_| SortKey(SmallVec::from_slice(&[RANK_NUM, 0x01])),
|_| unreachable!(),
|_| unreachable!(),
);
assert_eq!(neg_key.as_bytes()[0], RANK_SPECIAL);
assert_eq!(neg_key.as_bytes()[1], SPECIAL_NEG);
}
#[test]
fn named_constants_have_distinct_constant_subranks() {
let keys = [
atom_key(&ExprNode::Pi),
atom_key(&ExprNode::E),
atom_key(&ExprNode::ImaginaryUnit),
atom_key(&ExprNode::EulerGamma),
atom_key(&ExprNode::Catalan),
atom_key(&ExprNode::GoldenRatio),
];
for k in &keys {
assert_eq!(k.as_bytes()[0], RANK_CONSTANT);
}
for i in 0..keys.len() {
for j in (i + 1)..keys.len() {
assert!(keys[i] < keys[j], "constant sub-ranks must be ordered");
}
}
}
#[test]
fn long_keys_are_bounded_and_deterministic() {
let big = SortKey(SmallVec::from_slice(&vec![7u8; MAX_KEY_BYTES]));
let get = |_| big.clone();
let add = ExprNode::Add(smallvec![ExprId(0), ExprId(1)]);
let k1 = compute_sort_key(&add, get, |_| unreachable!(), |_| unreachable!());
assert!(k1.is_truncated());
assert_eq!(k1.as_bytes().len(), MAX_KEY_BYTES + 8);
let k2 = compute_sort_key(&add, get, |_| unreachable!(), |_| unreachable!());
assert_eq!(k1, k2);
let add3 = ExprNode::Add(smallvec![ExprId(0), ExprId(1), ExprId(2)]);
let k3 = compute_sort_key(&add3, get, |_| unreachable!(), |_| unreachable!());
assert!(k3.is_truncated());
assert_ne!(k1, k3);
assert_eq!(
k1.as_bytes()[..MAX_KEY_BYTES],
k3.as_bytes()[..MAX_KEY_BYTES]
);
let exact = SortKey(SmallVec::from_slice(&k1.as_bytes()[..MAX_KEY_BYTES]));
assert!(!exact.is_truncated());
assert!(exact < k1);
let small = atom_key(&ExprNode::Pi);
assert!(!small.is_truncated());
assert_eq!(small.as_bytes().len(), 2);
}
#[test]
fn new_function_discriminants_are_distinct() {
let child = |_| SortKey(SmallVec::from_slice(&[RANK_NUM, 0x01]));
let nodes = [
ExprNode::Re(ExprId(0)),
ExprNode::Im(ExprId(0)),
ExprNode::Conjugate(ExprId(0)),
ExprNode::Arg(ExprId(0)),
ExprNode::Si(ExprId(0)),
ExprNode::Ci(ExprId(0)),
ExprNode::Ei(ExprId(0)),
ExprNode::Li(ExprId(0)),
ExprNode::Zeta(ExprId(0)),
ExprNode::Polygamma(ExprId(0), ExprId(0)),
ExprNode::KroneckerDelta(ExprId(0), ExprId(0)),
ExprNode::LambertW(ExprId(0)),
];
let keys: Vec<SortKey> = nodes
.iter()
.map(|n| compute_sort_key(n, child, |_| unreachable!(), |_| unreachable!()))
.collect();
for k in &keys {
assert_eq!(k.as_bytes()[0], RANK_FUNCTION);
}
for i in 0..keys.len() {
for j in (i + 1)..keys.len() {
assert_ne!(keys[i], keys[j], "{:?} vs {:?}", nodes[i], nodes[j]);
}
}
}
}