use symplex::prelude::*;
fn at_minus_two(e: &Ex) -> f64 {
let ctx = e.context();
let m2 = ctx.int(-2);
e.subs(&ctx.symbol("x"), &m2)
.subs(&ctx.symbol("y"), &m2)
.eval_f64()
.unwrap_or_else(|err| panic!("{e} at -2 did not evaluate to a real f64: {err}"))
}
fn assert_not_collapsed(e: &Ex, collapsed: &Ex, reference: f64) {
let s = e.simplify();
assert_ne!(
s.id(),
collapsed.id(),
"{e} must not collapse to {collapsed}; got {s}"
);
let orig = at_minus_two(e);
let simp = at_minus_two(&s);
assert!(
(orig - reference).abs() <= 1e-12 * reference.abs().max(1.0),
"{e} at -2: got {orig}, sympy says {reference}"
);
assert!(
(simp - reference).abs() <= 1e-12 * reference.abs().max(1.0),
"simplify({e}) = {s} at -2: got {simp}, sympy says {reference}"
);
}
#[test]
fn pow_pow_x2_to_3_2_stays() {
let ctx = Context::new();
let x = ctx.symbol("x");
let e = x.powi(2).pow(&ctx.rational(3, 2));
assert_not_collapsed(&e, &x.powi(3), 8.0);
}
#[test]
fn pow_pow_sqrt_of_inverse_square_stays() {
let ctx = Context::new();
let x = ctx.symbol("x");
let e = (ctx.int(1) / x.powi(2)).sqrt();
assert_not_collapsed(&e, &(ctx.int(1) / &x), 0.5);
}
#[test]
fn pow_pow_x2_to_1_4_stays() {
let ctx = Context::new();
let x = ctx.symbol("x");
let e = x.powi(2).pow(&ctx.rational(1, 4));
assert_not_collapsed(&e, &x.sqrt(), std::f64::consts::SQRT_2);
}
#[test]
fn pow_pow_cbrt_of_square_stays() {
let ctx = Context::new();
let y = ctx.symbol("y");
let e = y.powi(2).cbrt();
assert_not_collapsed(&e, &y.pow(&ctx.rational(2, 3)), 1.5874010519681996);
}
#[test]
fn pow_pow_cbrt_of_cube_stays() {
let ctx = Context::new();
let x = ctx.symbol("x");
let e = x.powi(3).cbrt();
let s = e.simplify();
assert_ne!(s.id(), x.id(), "cbrt(x^3) must not collapse to x; got {s}");
assert_eq!(
s.id(),
e.id(),
"cbrt(x^3) should be left unchanged; got {s}"
);
}
#[test]
fn pow_pow_cbrt_of_fourth_power_stays() {
let ctx = Context::new();
let x = ctx.symbol("x");
let e = x.powi(4).cbrt();
assert_not_collapsed(&e, &x.pow(&ctx.rational(4, 3)), 2.5198420997897464);
}
#[test]
fn pow_pow_cbrt_of_shifted_square_stays() {
let ctx = Context::new();
let x = ctx.symbol("x");
let base = &x + ctx.int(1);
let e = base.powi(2).cbrt();
assert_not_collapsed(&e, &base.pow(&ctx.rational(2, 3)), 1.0);
}
#[test]
fn pow_pow_constant_cbrt_abs_sin6_squared_stays_real() {
let ctx = Context::new();
let e = ctx.int(6).sin().powi(2).abs().cbrt();
let s = e.simplify();
let v = s
.eval_f64()
.unwrap_or_else(|err| panic!("simplify({e}) = {s} is not real: {err}"));
assert!(
(v - 0.4273991566043774).abs() <= 1e-12,
"simplify({e}) = {s} evaluates to {v}, sympy says 0.4273991566043774"
);
}
#[test]
fn pow_pow_sqrt_squared_collapses() {
let ctx = Context::new();
let x = ctx.symbol("x");
let e = x.pow(&ctx.rational(1, 2)).powi(2);
assert_eq!(e.simplify().id(), x.id(), "(x^(1/2))^2 should be x");
}
#[test]
fn pow_pow_cbrt_cubed_collapses() {
let ctx = Context::new();
let x = ctx.symbol("x");
let e = x.pow(&ctx.rational(1, 3)).powi(3);
assert_eq!(e.simplify().id(), x.id(), "(x^(1/3))^3 should be x");
}
#[test]
fn pow_pow_integer_integer_collapses() {
let ctx = Context::new();
let x = ctx.symbol("x");
let e = x.powi(2).powi(3);
assert_eq!(e.simplify().id(), x.powi(6).id(), "(x^2)^3 should be x^6");
}
#[test]
fn pow_pow_inner_in_unit_interval_collapses() {
let ctx = Context::new();
let x = ctx.symbol("x");
let e = x.pow(&ctx.rational(1, 2)).pow(&ctx.rational(1, 2));
assert_eq!(
e.simplify().id(),
x.pow(&ctx.rational(1, 4)).id(),
"(x^(1/2))^(1/2) should be x^(1/4)"
);
}
#[test]
fn pow_pow_positive_base_collapses() {
let ctx = Context::new();
let xp = ctx.symbol_with("xp", &[Assumption::Positive]);
let e = xp.powi(2).pow(&ctx.rational(3, 2));
assert_eq!(
e.simplify().id(),
xp.powi(3).id(),
"(xp^2)^(3/2) should be xp^3"
);
let e = (ctx.int(1) / xp.powi(2)).sqrt();
assert_eq!(
e.simplify().id(),
(ctx.int(1) / &xp).id(),
"sqrt(1/xp^2) should be 1/xp"
);
let e = xp.powi(2).pow(&ctx.rational(1, 4));
assert_eq!(
e.simplify().id(),
xp.sqrt().id(),
"(xp^2)^(1/4) should be sqrt(xp)"
);
}
#[test]
fn pow_pow_negative_base_does_not_give_x_cubed() {
let ctx = Context::new();
let xn = ctx.symbol_with("xn", &[Assumption::Negative]);
let e = xn.powi(2).pow(&ctx.rational(3, 2));
let s = e.simplify();
assert_ne!(
s.id(),
xn.powi(3).id(),
"(xn^2)^(3/2) must not be xn^3 for xn < 0"
);
let v = s
.subs(&xn, &ctx.int(-2))
.eval_f64()
.unwrap_or_else(|err| panic!("{s} at xn = -2 did not evaluate: {err}"));
assert!(
(v - 8.0).abs() <= 1e-12,
"simplify({e}) = {s} at xn = -2 gives {v}, sympy says 8"
);
}