use symplex::prelude::*;
fn assert_ftc(integrand: &Ex, var: &Ex, label: &str) {
let anti = integrand.integrate(var);
let s = format!("{anti}");
assert!(
!s.contains("Integral"),
"{label}: integration returned unevaluated: {s}"
);
let deriv = anti.diff(var);
let test_point = integrand.context().rational(7, 10);
if let (Ok(o), Ok(d)) = (
integrand.subs(var, &test_point).eval_f64(),
deriv.subs(var, &test_point).eval_f64(),
) && o.is_finite()
&& d.is_finite()
{
let diff = (o - d).abs();
let tol = 1e-7 * o.abs().max(1.0);
assert!(
diff < tol,
"{label}: FTC failed — integrand={o:.12}, deriv={d:.12}, diff={diff:.2e}",
);
}
}
#[test]
fn integrate_tan_squared() {
let ctx = Context::new();
let x = ctx.symbol("x");
let integrand = x.tan().powi(2);
assert_ftc(&integrand, &x, "∫tan²(x)dx");
}
#[test]
fn integrate_tan_squared_numeric() {
let ctx = Context::new();
let x = ctx.symbol("x");
let anti = x.tan().powi(2).integrate(&x);
let s = format!("{anti}");
assert!(!s.contains("Integral"), "tan² should be integrated: {s}");
let deriv = anti.diff(&x);
let pt = ctx.rational(7, 10);
let integrand_val = x.tan().powi(2).subs(&x, &pt).eval_f64().unwrap();
let deriv_val = deriv.subs(&x, &pt).eval_f64().unwrap();
let err = (integrand_val - deriv_val).abs();
assert!(err < 1e-8, "FTC for tan²: {integrand_val} vs {deriv_val}");
}
#[test]
fn integrate_sech_squared() {
let ctx = Context::new();
let x = ctx.symbol("x");
let integrand = x.cosh().powi(-2);
let anti = integrand.integrate(&x);
let s = format!("{anti}");
assert!(!s.contains("Integral"), "sech² should be integrated: {s}");
let pt = ctx.rational(7, 10);
let anti_val = anti.subs(&x, &pt).eval_f64().unwrap();
let tanh_val = (0.7_f64).tanh();
assert!(
(anti_val - tanh_val).abs() < 1e-10,
"∫sech²(x)dx at x=0.7: got {anti_val}, expected tanh(0.7)={tanh_val}"
);
}
#[test]
fn integrate_sech_squared_ftc() {
let ctx = Context::new();
let x = ctx.symbol("x");
let integrand = x.cosh().powi(-2);
assert_ftc(&integrand, &x, "∫sech²(x)dx");
}
#[test]
fn integrate_sinh_squared() {
let ctx = Context::new();
let x = ctx.symbol("x");
let integrand = x.sinh().powi(2);
let anti = integrand.integrate(&x);
let s = format!("{anti}");
assert!(!s.contains("Integral"), "sinh² should be integrated: {s}");
assert_ftc(&integrand, &x, "∫sinh²(x)dx");
}
#[test]
fn integrate_sinh_squared_numeric() {
let ctx = Context::new();
let x = ctx.symbol("x");
let anti = x.sinh().powi(2).integrate(&x);
let val = anti.subs(&x, &ctx.rational(7, 10)).eval_f64().unwrap();
let expected = (1.4_f64).sinh() / 4.0 - 0.7 / 2.0;
assert!(
(val - expected).abs() < 1e-10,
"∫sinh²(x)dx at x=0.7: got {val}, expected {expected}"
);
}
#[test]
fn integrate_cosh_squared() {
let ctx = Context::new();
let x = ctx.symbol("x");
let integrand = x.cosh().powi(2);
let anti = integrand.integrate(&x);
let s = format!("{anti}");
assert!(!s.contains("Integral"), "cosh² should be integrated: {s}");
assert_ftc(&integrand, &x, "∫cosh²(x)dx");
}
#[test]
fn integrate_tanh_squared() {
let ctx = Context::new();
let x = ctx.symbol("x");
let integrand = x.tanh().powi(2);
assert_ftc(&integrand, &x, "∫tanh²(x)dx");
}
#[test]
fn integrate_ln_x_squared() {
let ctx = Context::new();
let x = ctx.symbol("x");
let integrand = x.ln().powi(2);
let anti = integrand.integrate(&x);
let s = format!("{anti}");
assert!(!s.contains("Integral"), "ln(x)² should be integrated: {s}");
let deriv = anti.diff(&x);
let pt = ctx.int(2);
let orig_val = integrand.subs(&x, &pt).eval_f64().unwrap();
let deriv_val = deriv.subs(&x, &pt).eval_f64().unwrap();
let err = (orig_val - deriv_val).abs();
assert!(
err < 1e-8,
"FTC for ∫ln(x)²dx: integrand={orig_val}, deriv={deriv_val}, err={err}"
);
}
#[test]
fn integrate_ln_x_squared_numeric() {
let ctx = Context::new();
let x = ctx.symbol("x");
let anti = x.ln().powi(2).integrate(&x);
let e_val = std::f64::consts::E;
let e_expr = ctx.symbol("__e_placeholder");
let val = anti.subs(&x, &ctx.int(2)).eval_f64().unwrap();
let ln2 = 2.0_f64.ln();
let expected = 2.0 * ln2 * ln2 - 2.0 * 2.0 * ln2 + 2.0 * 2.0;
let _ = e_val;
let _ = e_expr;
assert!(
(val - expected).abs() < 1e-8,
"∫ln(x)²dx at x=2: got {val}, expected {expected}"
);
}
#[test]
fn integrate_ln_x_cubed() {
let ctx = Context::new();
let x = ctx.symbol("x");
let integrand = x.ln().powi(3);
let anti = integrand.integrate(&x);
let s = format!("{anti}");
assert!(!s.contains("Integral"), "ln(x)³ should be integrated: {s}");
let deriv = anti.diff(&x);
let pt = ctx.int(2);
let orig_val = integrand.subs(&x, &pt).eval_f64().unwrap();
let deriv_val = deriv.subs(&x, &pt).eval_f64().unwrap();
let err = (orig_val - deriv_val).abs();
assert!(
err < 1e-7,
"FTC for ∫ln(x)³dx: integrand={orig_val}, deriv={deriv_val}, err={err}"
);
}
#[test]
fn ode_y_double_prime_plus_y_eq_0_uses_trig() {
let ctx = Context::new();
let x = ctx.symbol("x");
let y = ctx.symbol("y");
let dy = y.formal_diff(&x);
let d2y = dy.formal_diff(&x);
let ode = &d2y + &y;
let sol = ode.try_solve_ode(&y, &x).expect("should solve y'' + y = 0");
let s = format!("{sol}");
assert!(s.contains("C1"), "should contain C1: {s}");
assert!(s.contains("C2"), "should contain C2: {s}");
assert!(
s.contains("cos") && s.contains("sin"),
"should use cos and sin (not complex exp): {s}"
);
}
#[test]
fn ode_y_double_prime_plus_y_eq_0_trig_solution_correct() {
let ctx = Context::new();
let x = ctx.symbol("x");
let y = ctx.symbol("y");
let dy = y.formal_diff(&x);
let d2y = dy.formal_diff(&x);
let ode = &d2y + &y;
let sol = ode.try_solve_ode(&y, &x).expect("should solve");
let c1 = ctx.symbol("C1");
let c2 = ctx.symbol("C2");
let sol_c = sol.subs(&c1, &ctx.int(1)).subs(&c2, &ctx.int(0));
let sol_dd = sol_c.diff(&x).diff(&x);
let check = &sol_dd + &sol_c;
let pt = ctx.rational(7, 10);
if let Ok(val) = check.subs(&x, &pt).eval_f64() {
assert!(
val.abs() < 1e-8,
"y'' + y should be 0 for y=cos(x): got {val}"
);
}
}
#[test]
fn ode_y_double_prime_plus_y_eq_sin_x() {
let ctx = Context::new();
let x = ctx.symbol("x");
let y = ctx.symbol("y");
let dy = y.formal_diff(&x);
let d2y = dy.formal_diff(&x);
let sin_x = x.sin();
let ode = &(&d2y + &y) - &sin_x;
let sol = ode.solve_ode(&y, &x);
assert!(!sol.has_unevaluated(), "should solve y'' + y = sin(x)");
let s = format!("{sol}");
assert!(
s.contains("C1") && s.contains("C2"),
"should have two constants: {s}"
);
assert!(
s.contains("cos") && s.contains("sin"),
"should use trig form: {s}"
);
}
#[test]
fn ode_y_double_prime_plus_y_eq_sin_x_verifies() {
let ctx = Context::new();
let x = ctx.symbol("x");
let y = ctx.symbol("y");
let dy = y.formal_diff(&x);
let d2y = dy.formal_diff(&x);
let sin_x = x.sin();
let ode = &(&d2y + &y) - &sin_x;
let sol = ode.solve_ode(&y, &x);
if !sol.has_unevaluated() {
let c1 = ctx.symbol("C1");
let c2 = ctx.symbol("C2");
let sol_specific = sol.subs(&c1, &ctx.int(0)).subs(&c2, &ctx.int(0));
let sol_dd = sol_specific.diff(&x).diff(&x);
let check = &(&sol_dd + &sol_specific) - &sin_x;
let pt = ctx.rational(7, 10);
if let Ok(val) = check.subs(&x, &pt).eval_f64() {
assert!(
val.abs() < 1e-6,
"y'' + y - sin(x) should be ≈0 for particular solution: got {val}"
);
}
}
}
#[test]
fn ode_first_order_linear_exp_rhs() {
let ctx = Context::new();
let x = ctx.symbol("x");
let y = ctx.symbol("y");
let dy = y.formal_diff(&x);
let two_y = &y * 2;
let neg_x = (&x * -1).exp();
let ode = &(&dy + &two_y) - &neg_x;
let sol = ode.solve_ode(&y, &x);
assert!(!sol.has_unevaluated(), "should solve y' + 2y = exp(-x)");
let s = format!("{sol}");
assert!(
!s.contains("Integral"),
"solution should not have unevaluated Integral: {s}"
);
assert!(s.contains("C1"), "should have constant C1: {s}");
}
#[test]
fn ode_y_double_prime_minus_y_eq_0_real_exp() {
let ctx = Context::new();
let x = ctx.symbol("x");
let y = ctx.symbol("y");
let dy = y.formal_diff(&x);
let d2y = dy.formal_diff(&x);
let ode = &d2y - &y;
let sol = ode.try_solve_ode(&y, &x).expect("should solve y'' - y = 0");
let s = format!("{sol}");
assert!(
s.contains("C1") && s.contains("C2"),
"should have 2 constants: {s}"
);
assert!(s.contains("exp"), "should use exp: {s}");
}
#[test]
fn ode_damped_oscillator() {
let ctx = Context::new();
let x = ctx.symbol("x");
let y = ctx.symbol("y");
let dy = y.formal_diff(&x);
let d2y = dy.formal_diff(&x);
let ode = &(&d2y + &(&dy * 2)) + &(&y * 5);
let sol = ode.solve_ode(&y, &x);
assert!(!sol.has_unevaluated(), "should solve y'' + 2y' + 5y = 0");
let s = format!("{sol}");
assert!(
s.contains("cos") && s.contains("sin"),
"damped oscillator should use trig form: {s}"
);
assert!(s.contains("exp"), "should have exponential decay: {s}");
}
#[test]
fn no_exp_zero_artifact() {
let ctx = Context::new();
let x = ctx.symbol("x");
let integrand = x.sin();
let anti = integrand.integrate(&x);
let s = format!("{anti}");
assert!(!s.contains("exp(0)"), "should not contain exp(0): {s}");
}
#[test]
fn no_exp_zero_in_linear_sub() {
let ctx = Context::new();
let x = ctx.symbol("x");
let two_x = &x * 2;
let integrand = two_x.sin();
let anti = integrand.integrate(&x);
let s = format!("{anti}");
assert!(!s.contains("exp(0)"), "should not contain exp(0): {s}");
}