use num_traits::Signed;
use crate::base::arena::Arena;
use crate::base::node::{ExprId, ExprNode};
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ExtensionKind {
Logarithmic,
Exponential,
}
#[derive(Clone, Debug)]
pub struct ExtensionLevel {
pub kind: ExtensionKind,
pub ext_var: ExprId,
pub argument: ExprId,
pub derivative: ExprId,
}
#[derive(Clone, Debug)]
pub struct DifferentialExtension {
pub base_var: ExprId,
pub levels: Vec<ExtensionLevel>,
pub integrand: ExprId,
pub current_level: isize,
}
impl DifferentialExtension {
pub fn new(base_var: ExprId) -> Self {
DifferentialExtension {
base_var,
levels: Vec::new(),
integrand: base_var, current_level: -1,
}
}
pub fn depth(&self) -> usize {
self.levels.len()
}
pub fn is_base_level(&self) -> bool {
self.levels.is_empty() || self.current_level < 0
}
pub fn current_kind(&self) -> Option<&ExtensionKind> {
if self.current_level >= 0 && (self.current_level as usize) < self.levels.len() {
Some(&self.levels[self.current_level as usize].kind)
} else {
None
}
}
pub fn current_level_ext(&self) -> Option<&ExtensionLevel> {
if self.current_level >= 0 && (self.current_level as usize) < self.levels.len() {
Some(&self.levels[self.current_level as usize])
} else {
None
}
}
pub fn push_logarithmic(&mut self, ext_var: ExprId, argument: ExprId, derivative: ExprId) {
self.levels.push(ExtensionLevel {
kind: ExtensionKind::Logarithmic,
ext_var,
argument,
derivative,
});
self.current_level = (self.levels.len() - 1) as isize;
}
pub fn push_exponential(&mut self, ext_var: ExprId, argument: ExprId, derivative: ExprId) {
self.levels.push(ExtensionLevel {
kind: ExtensionKind::Exponential,
ext_var,
argument,
derivative,
});
self.current_level = (self.levels.len() - 1) as isize;
}
#[cfg_attr(not(test), allow(dead_code))]
pub fn decrement_level(&mut self) {
self.current_level -= 1;
}
#[cfg_attr(not(test), allow(dead_code))]
pub fn increment_level(&mut self) {
if self.current_level < self.levels.len() as isize - 1 {
self.current_level += 1;
}
}
}
pub fn build_tower(
arena: &mut Arena,
expr: ExprId,
var: ExprId,
) -> Result<DifferentialExtension, String> {
let mut exp_args: Vec<ExprId> = Vec::new();
let mut ln_args: Vec<ExprId> = Vec::new();
collect_transcendentals(arena, expr, var, &mut exp_args, &mut ln_args);
exp_args.sort_by_key(|e| e.0);
exp_args.dedup();
ln_args.sort_by_key(|e| e.0);
ln_args.dedup();
let has_exp = !exp_args.is_empty();
let has_ln = !ln_args.is_empty();
if !has_exp && !has_ln {
let mut de = DifferentialExtension::new(var);
de.integrand = expr;
return Ok(de);
}
if has_exp && has_ln {
return Err("Mixed exp and ln extensions not yet supported".into());
}
if has_exp {
if exp_args.len() > 1 {
match find_integer_multiples(arena, &exp_args, var) {
Some((base_arg, multiples)) => {
return build_exp_tower_multi(arena, expr, var, base_arg, &multiples);
}
None => {
return Err("Multiple independent exp arguments not yet supported".into());
}
}
}
build_exp_tower_multi(arena, expr, var, exp_args[0], &[(exp_args[0], 1)])
} else {
if ln_args.len() > 1 {
let base_arg = ln_args[0];
for &arg in &ln_args[1..] {
if arg != base_arg {
return Err("Multiple independent ln arguments not yet supported".into());
}
}
}
build_ln_tower(arena, expr, var, ln_args[0])
}
}
fn collect_transcendentals(
arena: &Arena,
expr: ExprId,
var: ExprId,
exp_args: &mut Vec<ExprId>,
ln_args: &mut Vec<ExprId>,
) {
let mut stack = vec![expr];
let mut visited = std::collections::HashSet::new();
while let Some(id) = stack.pop() {
if !visited.insert(id) {
continue;
}
match arena.node(id).clone() {
ExprNode::Exp(inner) => {
if crate::base::walk::contains(arena, inner, var) {
exp_args.push(inner);
}
stack.push(inner);
}
ExprNode::Ln(inner) => {
if crate::base::walk::contains(arena, inner, var) {
ln_args.push(inner);
}
stack.push(inner);
}
other => {
for &child in other.children().iter() {
stack.push(child);
}
}
}
}
}
fn find_integer_multiples(
arena: &Arena,
args: &[ExprId],
var: ExprId,
) -> Option<(ExprId, Vec<(ExprId, i64)>)> {
if args.is_empty() {
return None;
}
if args.len() == 1 {
return Some((args[0], vec![(args[0], 1)]));
}
let polys: Vec<crate::poly::dense::Poly> = args
.iter()
.filter_map(|&arg| crate::poly::polybridge::expr_to_poly(arena, arg, var))
.collect();
if polys.len() != args.len() {
return None; }
let ref_poly = &polys[0];
if ref_poly.is_zero() {
return None;
}
let mut ratios: Vec<num_rational::Ratio<num_bigint::BigInt>> = Vec::new();
ratios.push(num_rational::Ratio::from_integer(num_bigint::BigInt::from(
1,
)));
for poly in &polys[1..] {
let (quot, rem) = poly.div_rem(ref_poly);
if !rem.is_zero() {
return None; }
if !quot.is_constant() {
return None; }
let k = quot.coeff(0);
if !k.is_positive() {
return None; }
ratios.push(k);
}
let min_ratio = ratios.iter().min().cloned()?;
let int_multiples: Vec<num_rational::Ratio<num_bigint::BigInt>> =
ratios.iter().map(|r| r / &min_ratio).collect();
for m in &int_multiples {
if !m.is_integer() || !m.is_positive() {
return None;
}
}
let base_idx = ratios.iter().position(|r| *r == min_ratio);
let base_idx = base_idx?;
let base_arg = args[base_idx];
let multiples: Vec<(ExprId, i64)> = args
.iter()
.zip(int_multiples.iter())
.map(|(&arg, m)| {
let k = i64::try_from(m.to_integer()).unwrap_or(0);
(arg, k)
})
.collect();
if multiples.iter().any(|&(_, k)| k <= 0) {
return None;
}
Some((base_arg, multiples))
}
fn build_exp_tower_multi(
arena: &mut Arena,
expr: ExprId,
var: ExprId,
base_u: ExprId,
multiples: &[(ExprId, i64)],
) -> Result<DifferentialExtension, String> {
let theta = arena.symbol("__t0");
let du = crate::transforms::diff::diff(arena, base_u, var);
let d_theta = arena.mul(&[du, theta]);
let mut rewritten = expr;
for &(arg, k) in multiples {
let exp_arg = arena.exp(arg);
let replacement = if k == 1 {
theta
} else {
let k_id = arena.int(k);
arena.pow(theta, k_id)
};
rewritten = crate::transforms::subs::subs(arena, rewritten, exp_arg, replacement);
}
let mut de = DifferentialExtension::new(var);
de.push_exponential(theta, base_u, d_theta);
de.integrand = rewritten;
Ok(de)
}
fn build_ln_tower(
arena: &mut Arena,
expr: ExprId,
var: ExprId,
u: ExprId,
) -> Result<DifferentialExtension, String> {
let theta = arena.symbol("__t0");
let du = crate::transforms::diff::diff(arena, u, var);
let d_theta = arena.div(du, u);
let ln_u = arena.ln(u);
let rewritten = crate::transforms::subs::subs(arena, expr, ln_u, theta);
let mut de = DifferentialExtension::new(var);
de.push_logarithmic(theta, u, d_theta);
de.integrand = rewritten;
Ok(de)
}
pub fn extract_poly_in_ext(
arena: &Arena,
expr: ExprId,
ext_var: ExprId,
) -> Option<Vec<(usize, ExprId)>> {
let mut coeffs: std::collections::BTreeMap<usize, Vec<ExprId>> =
std::collections::BTreeMap::new();
extract_terms(arena, expr, ext_var, &mut coeffs)?;
let result: Vec<(usize, ExprId)> = coeffs
.into_iter()
.map(|(power, terms)| {
let coeff = if terms.len() == 1 {
terms[0]
} else {
terms[0] };
(power, coeff)
})
.collect();
if result.is_empty() {
None
} else {
Some(result)
}
}
fn extract_terms(
arena: &Arena,
expr: ExprId,
ext_var: ExprId,
coeffs: &mut std::collections::BTreeMap<usize, Vec<ExprId>>,
) -> Option<()> {
if !crate::base::walk::contains(arena, expr, ext_var) {
coeffs.entry(0).or_default().push(expr);
return Some(());
}
if expr == ext_var {
coeffs.entry(1).or_default().push(arena.one());
return Some(());
}
match arena.node(expr).clone() {
ExprNode::Add(ref children) => {
for &child in children.iter() {
extract_terms(arena, child, ext_var, coeffs)?;
}
Some(())
}
ExprNode::Neg(inner) => {
if !crate::base::walk::contains(arena, inner, ext_var) {
coeffs.entry(0).or_default().push(expr);
return Some(());
}
None
}
ExprNode::Mul(ref children) => {
let mut theta_power: usize = 0;
let mut coeff_factors: Vec<ExprId> = Vec::new();
for &child in children.iter() {
if !crate::base::walk::contains(arena, child, ext_var) {
coeff_factors.push(child);
} else if child == ext_var {
theta_power += 1;
} else if let ExprNode::Pow(base, exp) = arena.node(child).clone() {
if base == ext_var {
{
let r = arena.as_num(exp)?;
if r.is_integer() && !(*r).is_negative() {
if let Ok(n) = usize::try_from(r.to_integer()) {
theta_power += n;
} else {
return None;
}
} else {
return None; }
}
} else {
return None; }
} else {
return None; }
}
let coeff = if coeff_factors.is_empty() {
arena.one()
} else if coeff_factors.len() == 1 {
coeff_factors[0]
} else {
return None;
};
coeffs.entry(theta_power).or_default().push(coeff);
Some(())
}
ExprNode::Pow(base, exp) => {
if base == ext_var
&& let Some(r) = arena.as_num(exp)
&& r.is_integer()
&& !(*r).is_negative()
&& let Ok(n) = usize::try_from(r.to_integer())
{
coeffs.entry(n).or_default().push(arena.one());
return Some(());
}
None
}
_ => None,
}
}
pub fn extract_poly_in_ext_mut(
arena: &mut Arena,
expr: ExprId,
ext_var: ExprId,
) -> Option<Vec<(usize, ExprId)>> {
if let Some(result) = extract_poly_in_ext(arena, expr, ext_var) {
return Some(result);
}
let expanded = crate::transforms::expand::expand(arena, expr);
let evaled = crate::transforms::eval::eval(arena, expanded);
extract_poly_in_ext(arena, evaled, ext_var)
}
#[cfg_attr(not(test), allow(dead_code))]
pub fn derivation(arena: &mut Arena, expr: ExprId, de: &DifferentialExtension) -> ExprId {
if de.levels.is_empty() {
return crate::transforms::diff::diff(arena, expr, de.base_var);
}
let mut result = crate::transforms::diff::diff(arena, expr, de.base_var);
for level in &de.levels {
let df_dtheta = crate::transforms::diff::diff(arena, expr, level.ext_var);
if df_dtheta == arena.zero() {
continue;
}
let contribution = arena.mul(&[df_dtheta, level.derivative]);
result = arena.add(&[result, contribution]);
}
crate::transforms::eval::eval(arena, result)
}
#[cfg(test)]
mod tests {
use super::*;
fn sym(arena: &mut Arena, name: &str) -> ExprId {
arena.symbol(name)
}
fn display(arena: &Arena, id: ExprId) -> String {
arena.display(id).to_string()
}
#[test]
fn empty_tower() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let de = DifferentialExtension::new(x);
assert_eq!(de.depth(), 0);
assert!(de.is_base_level());
assert!(de.current_kind().is_none());
}
#[test]
fn push_logarithmic_extension() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let theta = sym(&mut arena, "t0");
let d_theta = arena.div(arena.one(), x);
let mut de = DifferentialExtension::new(x);
de.push_logarithmic(theta, x, d_theta);
assert_eq!(de.depth(), 1);
assert!(!de.is_base_level());
assert_eq!(de.current_kind(), Some(&ExtensionKind::Logarithmic));
}
#[test]
fn push_exponential_extension() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let theta = sym(&mut arena, "t0");
let mut de = DifferentialExtension::new(x);
de.push_exponential(theta, x, theta);
assert_eq!(de.depth(), 1);
assert_eq!(de.current_kind(), Some(&ExtensionKind::Exponential));
}
#[test]
fn level_navigation() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let t0 = sym(&mut arena, "t0");
let t1 = sym(&mut arena, "t1");
let one = arena.one();
let dt0 = arena.div(one, x);
let mut de = DifferentialExtension::new(x);
de.push_logarithmic(t0, x, dt0);
de.push_exponential(t1, x, t1);
assert_eq!(de.depth(), 2);
assert_eq!(de.current_level, 1);
assert_eq!(de.current_kind(), Some(&ExtensionKind::Exponential));
de.decrement_level();
assert_eq!(de.current_level, 0);
assert_eq!(de.current_kind(), Some(&ExtensionKind::Logarithmic));
de.decrement_level();
assert_eq!(de.current_level, -1);
assert!(de.is_base_level());
de.increment_level();
assert_eq!(de.current_level, 0);
}
#[test]
fn derivation_base_level_polynomial() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let two = arena.int(2);
let one = arena.one();
let x_sq = arena.pow(x, two);
let expr = arena.add(&[x_sq, one]);
let de = DifferentialExtension::new(x);
let result = derivation(&mut arena, expr, &de);
let s = display(&arena, result);
assert_eq!(s, "2*x", "D(x²+1) = 2x, got: {s}");
}
#[test]
fn derivation_exp_extension_theta() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let theta = sym(&mut arena, "t0");
let mut de = DifferentialExtension::new(x);
de.push_exponential(theta, x, theta);
let result = derivation(&mut arena, theta, &de);
let s = display(&arena, result);
assert_eq!(s, "t0", "D(θ) = θ for exp extension, got: {s}");
}
#[test]
fn derivation_exp_extension_theta_squared() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let theta = sym(&mut arena, "t0");
let two = arena.int(2);
let theta_sq = arena.pow(theta, two);
let mut de = DifferentialExtension::new(x);
de.push_exponential(theta, x, theta);
let result = derivation(&mut arena, theta_sq, &de);
let s = display(&arena, result);
assert!(
s.contains("t0") && s.contains("2"),
"D(θ²) = 2θ² for exp extension, got: {s}"
);
}
#[test]
fn derivation_exp_extension_x_times_theta() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let theta = sym(&mut arena, "t0");
let x_theta = arena.mul(&[x, theta]);
let mut de = DifferentialExtension::new(x);
de.push_exponential(theta, x, theta);
let result = derivation(&mut arena, x_theta, &de);
let s = display(&arena, result);
assert!(
s.contains("t0") && s.contains("x"),
"D(x·θ) should contain both x and θ, got: {s}"
);
}
#[test]
fn derivation_ln_extension_theta() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let theta = sym(&mut arena, "t0");
let one = arena.one();
let d_theta = arena.div(one, x);
let mut de = DifferentialExtension::new(x);
de.push_logarithmic(theta, x, d_theta);
let result = derivation(&mut arena, theta, &de);
let s = display(&arena, result);
assert!(s.contains("x"), "D(θ) = 1/x for ln extension, got: {s}");
}
#[test]
fn derivation_ln_extension_constant() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let theta = sym(&mut arena, "t0");
let one = arena.one();
let d_theta = arena.div(one, x);
let five = arena.int(5);
let mut de = DifferentialExtension::new(x);
de.push_logarithmic(theta, x, d_theta);
let result = derivation(&mut arena, five, &de);
assert_eq!(result, arena.zero(), "D(5) = 0");
}
#[test]
fn extract_poly_constant() {
let mut arena = Arena::new();
let theta = sym(&mut arena, "t0");
let five = arena.int(5);
let result = extract_poly_in_ext(&arena, five, theta).unwrap();
assert_eq!(result.len(), 1);
assert_eq!(result[0].0, 0); assert_eq!(result[0].1, five); }
#[test]
fn extract_poly_theta_itself() {
let mut arena = Arena::new();
let theta = sym(&mut arena, "t0");
let result = extract_poly_in_ext(&arena, theta, theta).unwrap();
assert_eq!(result.len(), 1);
assert_eq!(result[0].0, 1); }
#[test]
fn extract_poly_theta_squared() {
let mut arena = Arena::new();
let theta = sym(&mut arena, "t0");
let two = arena.int(2);
let theta_sq = arena.pow(theta, two);
let result = extract_poly_in_ext(&arena, theta_sq, theta).unwrap();
assert_eq!(result.len(), 1);
assert_eq!(result[0].0, 2); }
#[test]
fn extract_poly_theta_plus_constant() {
let mut arena = Arena::new();
let theta = sym(&mut arena, "t0");
let five = arena.int(5);
let expr = arena.add(&[theta, five]);
let result = extract_poly_in_ext(&arena, expr, theta).unwrap();
assert_eq!(result.len(), 2);
let powers: Vec<usize> = result.iter().map(|&(p, _)| p).collect();
assert!(powers.contains(&0));
assert!(powers.contains(&1));
}
#[test]
fn extract_poly_x_times_theta() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let theta = sym(&mut arena, "t0");
let expr = arena.mul(&[x, theta]);
let result = extract_poly_in_ext(&arena, expr, theta).unwrap();
assert_eq!(result.len(), 1);
assert_eq!(result[0].0, 1); assert_eq!(result[0].1, x); }
#[test]
fn extract_poly_no_theta() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let theta = sym(&mut arena, "t0");
let two = arena.int(2);
let three = arena.int(3);
let one = arena.one();
let x_sq = arena.pow(x, two);
let three_x = arena.mul(&[three, x]);
let expr = arena.add(&[x_sq, three_x, one]);
let result = extract_poly_in_ext(&arena, expr, theta).unwrap();
assert_eq!(result.len(), 1);
assert_eq!(result[0].0, 0); assert_eq!(result[0].1, expr); }
#[test]
fn extract_poly_reconstruct_linear() {
let mut arena = Arena::new();
let theta = sym(&mut arena, "t0");
let three = arena.int(3);
let five = arena.int(5);
let three_theta = arena.mul(&[three, theta]);
let expr = arena.add(&[three_theta, five]);
let terms = extract_poly_in_ext(&arena, expr, theta).unwrap();
let mut reconstructed_parts: Vec<ExprId> = Vec::new();
for &(power, coeff) in &terms {
if power == 0 {
reconstructed_parts.push(coeff);
} else if power == 1 {
let term = arena.mul(&[coeff, theta]);
reconstructed_parts.push(term);
} else {
let exp = arena.int(power as i64);
let theta_k = arena.pow(theta, exp);
let term = arena.mul(&[coeff, theta_k]);
reconstructed_parts.push(term);
}
}
let reconstructed = if reconstructed_parts.len() == 1 {
reconstructed_parts[0]
} else {
arena.add(&reconstructed_parts)
};
let seven = arena.int(7);
let orig_val = crate::transforms::subs::subs(&mut arena, expr, theta, seven);
let orig_eval = crate::transforms::eval::eval(&mut arena, orig_val);
let recon_val = crate::transforms::subs::subs(&mut arena, reconstructed, theta, seven);
let recon_eval = crate::transforms::eval::eval(&mut arena, recon_val);
assert_eq!(
orig_eval,
recon_eval,
"reconstruction at θ=7: orig={}, recon={}",
display(&arena, orig_eval),
display(&arena, recon_eval)
);
}
#[test]
fn extract_poly_reconstruct_x_theta() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let theta = sym(&mut arena, "t0");
let expr = arena.mul(&[x, theta]);
let terms = extract_poly_in_ext(&arena, expr, theta).unwrap();
assert_eq!(terms.len(), 1);
assert_eq!(terms[0].0, 1);
assert_eq!(terms[0].1, x);
let coeff = terms[0].1;
let reconstructed = arena.mul(&[coeff, theta]);
let three = arena.int(3);
let five = arena.int(5);
let r1 = crate::transforms::subs::subs(&mut arena, reconstructed, x, three);
let r2 = crate::transforms::subs::subs(&mut arena, r1, theta, five);
let result = crate::transforms::eval::eval(&mut arena, r2);
assert_eq!(display(&arena, result), "15");
}
#[test]
fn extract_poly_reconstruct_constant_only() {
let mut arena = Arena::new();
let theta = sym(&mut arena, "t0");
let forty_two = arena.int(42);
let terms = extract_poly_in_ext(&arena, forty_two, theta).unwrap();
assert_eq!(terms.len(), 1);
assert_eq!(terms[0].0, 0);
assert_eq!(terms[0].1, forty_two);
}
#[test]
fn extract_poly_reconstruct_quadratic() {
let mut arena = Arena::new();
let theta = sym(&mut arena, "t0");
let two = arena.int(2);
let one = arena.one();
let theta_sq = arena.pow(theta, two);
let expr = arena.add(&[theta_sq, one]);
let terms = extract_poly_in_ext(&arena, expr, theta).unwrap();
let powers: Vec<usize> = terms.iter().map(|t| t.0).collect();
assert!(powers.contains(&0), "should have θ⁰ term");
assert!(powers.contains(&2), "should have θ² term");
let three = arena.int(3);
let orig_val = crate::transforms::subs::subs(&mut arena, expr, theta, three);
let orig_eval = crate::transforms::eval::eval(&mut arena, orig_val);
assert_eq!(display(&arena, orig_eval), "10");
}
#[test]
fn build_tower_exp_x_equivalence() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let expr = arena.exp(x);
let de = build_tower(&mut arena, expr, x).unwrap();
let ext_var = de.levels[0].ext_var;
let _arg = de.levels[0].argument;
let one = arena.int(1);
let orig_at_1 = crate::transforms::subs::subs(&mut arena, expr, x, one);
let integrand_at_1 = crate::transforms::subs::subs(&mut arena, de.integrand, x, one);
let exp_1 = arena.exp(one);
let integrand_back =
crate::transforms::subs::subs(&mut arena, integrand_at_1, ext_var, exp_1);
let integrand_eval = crate::transforms::eval::eval(&mut arena, integrand_back);
let orig_eval = crate::transforms::eval::eval(&mut arena, orig_at_1);
assert_eq!(
orig_eval,
integrand_eval,
"tower equivalence at x=1: orig={}, rewritten={}",
display(&arena, orig_eval),
display(&arena, integrand_eval)
);
}
#[test]
fn build_tower_ln_x_equivalence() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let expr = arena.ln(x);
let de = build_tower(&mut arena, expr, x).unwrap();
let ext_var = de.levels[0].ext_var;
let two = arena.int(2);
let orig_at_2 = crate::transforms::subs::subs(&mut arena, expr, x, two);
let integrand_at_2 = crate::transforms::subs::subs(&mut arena, de.integrand, x, two);
let ln_2 = arena.ln(two);
let integrand_back =
crate::transforms::subs::subs(&mut arena, integrand_at_2, ext_var, ln_2);
let integrand_eval = crate::transforms::eval::eval(&mut arena, integrand_back);
let orig_eval = crate::transforms::eval::eval(&mut arena, orig_at_2);
assert_eq!(
orig_eval,
integrand_eval,
"tower equivalence at x=2: orig={}, rewritten={}",
display(&arena, orig_eval),
display(&arena, integrand_eval)
);
}
#[test]
fn build_tower_exp_2x_plus_exp_x_equivalence() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let two = arena.int(2);
let two_x = arena.mul(&[two, x]);
let exp_x = arena.exp(x);
let exp_2x = arena.exp(two_x);
let expr = arena.add(&[exp_2x, exp_x]);
let de = build_tower(&mut arena, expr, x).unwrap();
let ext_var = de.levels[0].ext_var;
let zero = arena.int(0);
let orig_at_0 = crate::transforms::subs::subs(&mut arena, expr, x, zero);
let orig_eval = crate::transforms::eval::eval(&mut arena, orig_at_0);
let int_at_0 = crate::transforms::subs::subs(&mut arena, de.integrand, x, zero);
let exp_0 = arena.exp(zero);
let exp_0_eval = crate::transforms::eval::eval(&mut arena, exp_0);
let int_back = crate::transforms::subs::subs(&mut arena, int_at_0, ext_var, exp_0_eval);
let int_eval = crate::transforms::eval::eval(&mut arena, int_back);
assert_eq!(
display(&arena, orig_eval),
display(&arena, int_eval),
"tower equivalence at x=0: orig={}, rewritten={}",
display(&arena, orig_eval),
display(&arena, int_eval)
);
}
#[test]
fn build_tower_pure_rational() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let two = arena.int(2);
let one = arena.one();
let x_sq = arena.pow(x, two);
let expr = arena.add(&[x_sq, one]);
let de = build_tower(&mut arena, expr, x).unwrap();
assert_eq!(de.depth(), 0);
assert!(de.is_base_level());
}
#[test]
fn build_tower_single_exp() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let expr = arena.exp(x);
let de = build_tower(&mut arena, expr, x).unwrap();
assert_eq!(de.depth(), 1);
assert_eq!(de.current_kind(), Some(&ExtensionKind::Exponential));
let s = display(&arena, de.integrand);
assert!(s.contains("__t0"), "integrand should use __t0, got: {s}");
}
#[test]
fn build_tower_exp_over_1_plus_exp() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let exp_x = arena.exp(x);
let one = arena.one();
let denom = arena.add(&[one, exp_x]);
let expr = arena.div(exp_x, denom);
let de = build_tower(&mut arena, expr, x).unwrap();
assert_eq!(de.depth(), 1);
assert_eq!(de.current_kind(), Some(&ExtensionKind::Exponential));
let s = display(&arena, de.integrand);
assert!(s.contains("__t0"), "integrand should use __t0, got: {s}");
}
#[test]
fn build_tower_single_ln() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let expr = arena.ln(x);
let de = build_tower(&mut arena, expr, x).unwrap();
assert_eq!(de.depth(), 1);
assert_eq!(de.current_kind(), Some(&ExtensionKind::Logarithmic));
let s = display(&arena, de.integrand);
assert!(s.contains("__t0"), "integrand should use __t0, got: {s}");
}
#[test]
fn build_tower_exp_2x_plus_exp_x() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let two = arena.int(2);
let two_x = arena.mul(&[two, x]);
let exp_x = arena.exp(x);
let exp_2x = arena.exp(two_x);
let expr = arena.add(&[exp_2x, exp_x]);
let de = build_tower(&mut arena, expr, x).unwrap();
assert_eq!(de.depth(), 1);
assert_eq!(de.current_kind(), Some(&ExtensionKind::Exponential));
let s = display(&arena, de.integrand);
assert!(s.contains("__t0"), "integrand should use __t0, got: {s}");
assert!(
s.contains("__t0^2") || s.contains("__t0"),
"integrand should have powers of __t0, got: {s}"
);
}
#[test]
fn build_tower_exp_3x_plus_exp_x() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let three = arena.int(3);
let three_x = arena.mul(&[three, x]);
let exp_x = arena.exp(x);
let exp_3x = arena.exp(three_x);
let expr = arena.add(&[exp_3x, exp_x]);
let de = build_tower(&mut arena, expr, x).unwrap();
assert_eq!(de.depth(), 1);
let s = display(&arena, de.integrand);
assert!(s.contains("__t0"), "integrand should use __t0, got: {s}");
}
#[test]
fn build_tower_independent_exps_fails() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let two = arena.int(2);
let x_sq = arena.pow(x, two);
let exp_x = arena.exp(x);
let exp_x_sq = arena.exp(x_sq);
let expr = arena.add(&[exp_x, exp_x_sq]);
let result = build_tower(&mut arena, expr, x);
assert!(
result.is_err(),
"independent exps should fail: {:?}",
result
);
}
#[test]
fn integer_multiples_basic() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let two = arena.int(2);
let two_x = arena.mul(&[two, x]);
let result = find_integer_multiples(&arena, &[x, two_x], x);
assert!(result.is_some(), "x and 2x should be integer multiples");
let (base, mults) = result.unwrap();
assert_eq!(base, x);
assert_eq!(mults.len(), 2);
assert_eq!(mults[0].1, 1);
assert_eq!(mults[1].1, 2);
}
#[test]
fn integer_multiples_independent() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let two = arena.int(2);
let x_sq = arena.pow(x, two);
let result = find_integer_multiples(&arena, &[x, x_sq], x);
assert!(result.is_none(), "x and x² should not be integer multiples");
}
#[test]
fn integer_multiples_three_args() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let two = arena.int(2);
let two_x = arena.mul(&[two, x]);
let three = arena.int(3);
let three_x = arena.mul(&[three, x]);
let result = find_integer_multiples(&arena, &[x, two_x, three_x], x);
assert!(result.is_some());
let (_, mults) = result.unwrap();
let ks: Vec<i64> = mults.iter().map(|m| m.1).collect();
assert_eq!(ks, vec![1, 2, 3]);
}
#[test]
fn build_tower_one_over_x_ln_x() {
let mut arena = Arena::new();
let x = sym(&mut arena, "x");
let ln_x = arena.ln(x);
let x_ln_x = arena.mul(&[x, ln_x]);
let one = arena.one();
let expr = arena.div(one, x_ln_x);
let de = build_tower(&mut arena, expr, x).unwrap();
assert_eq!(de.depth(), 1);
assert_eq!(de.current_kind(), Some(&ExtensionKind::Logarithmic));
}
}