use symplex::prelude::*;
#[test]
fn evalf_f64_integer() {
let ctx = Context::new();
let five = ctx.int(5);
let val = five.eval_f64().unwrap();
assert!((val - 5.0).abs() < 1e-10);
}
#[test]
fn evalf_f64_pi() {
let ctx = Context::new();
let val = ctx.pi().eval_f64().unwrap();
assert!((val - std::f64::consts::PI).abs() < 1e-10);
}
#[test]
fn evalf_f64_expression() {
let ctx = Context::new();
let x = ctx.symbol("x");
let val = x.powi(2).subs_i64(&x, 3).eval_f64().unwrap();
assert!((val - 9.0).abs() < 1e-10);
}
#[test]
fn evalf_f64_free_symbol_errors() {
let ctx = Context::new();
let x = ctx.symbol("x");
assert!(x.eval_f64().is_err());
}
#[test]
fn evalf_f64_rational() {
let ctx = Context::new();
let half = ctx.rational(1, 3);
let val = half.eval_f64().unwrap();
assert!((val - 1.0 / 3.0).abs() < 1e-10);
}
#[test]
fn assume_positive() {
let ctx = Context::new();
let t = ctx.symbol("t").assume(Assumption::Positive);
assert_eq!(t.is_positive(), Some(true));
}
#[test]
fn assume_chained() {
let ctx = Context::new();
let t = ctx
.symbol("t")
.assume(Assumption::Positive)
.assume(Assumption::Real);
assert_eq!(t.is_positive(), Some(true));
assert_eq!(t.is_real(), Some(true));
}
#[test]
fn assume_integer() {
let ctx = Context::new();
let n = ctx.symbol("n").assume(Assumption::Integer);
assert_eq!(n.is_integer(), Some(true));
assert_eq!(n.is_real(), Some(true));
}
#[test]
fn assume_on_non_symbol_ignored() {
let ctx = Context::new();
let five = ctx.int(5);
let result = five.assume(Assumption::Positive);
assert_eq!(result.is_positive(), Some(true));
let y = ctx.symbol("y");
assert_eq!(
y.is_positive(),
None,
"bare symbol should not be known positive"
);
let y_pos = y.assume(Assumption::Positive);
assert_eq!(
y_pos.is_positive(),
Some(true),
"assumed-positive symbol should be positive"
);
let neg = ctx.int(-3);
assert_eq!(neg.is_positive(), Some(false), "-3 should not be positive");
}
#[test]
fn sum_of_integers() {
let ctx = Context::new();
let terms: Vec<Ex> = (1..=4).map(|n| ctx.int(n)).collect();
let total = Ex::sum_of(&ctx, terms);
assert_eq!(format!("{total}"), "10");
}
#[test]
fn product_of_integers() {
let ctx = Context::new();
let factors: Vec<Ex> = (1..=4).map(|n| ctx.int(n)).collect();
let total = Ex::product_of(&ctx, factors);
assert_eq!(format!("{total}"), "24");
}
#[test]
fn sum_of_empty() {
let ctx = Context::new();
let total = Ex::sum_of(&ctx, vec![]);
assert_eq!(format!("{total}"), "0");
}
#[test]
fn product_of_empty() {
let ctx = Context::new();
let total = Ex::product_of(&ctx, vec![]);
assert_eq!(format!("{total}"), "1");
}
#[test]
fn sum_of_expressions() {
let ctx = Context::new();
let x = ctx.symbol("x");
let terms = vec![x.powi(2), x.clone(), ctx.int(1)];
let total = Ex::sum_of(&ctx, terms);
let s = format!("{total}");
assert!(
s.contains("x^2") && s.contains("x"),
"should be polynomial: {s}"
);
}
#[test]
fn integrate_x_sin_x() {
let ctx = Context::new();
let x = ctx.symbol("x");
let expr = &x * &x.sin();
let result = expr.integrate(&x);
let s = format!("{result}");
assert!(s.contains("sin(x)"), "should contain sin(x): {s}");
assert!(s.contains("cos(x)"), "should contain cos(x): {s}");
let f_at_2 = result
.subs_i64(&x, 2)
.eval()
.eval_f64()
.expect("F(2) should evaluate");
let f_at_1 = result
.subs_i64(&x, 1)
.eval()
.eval_f64()
.expect("F(1) should evaluate");
let ftc_value = f_at_2 - f_at_1;
let expected = 2.0_f64.sin() - 2.0 * 2.0_f64.cos() - 1.0_f64.sin() + 1.0_f64.cos();
assert!(
(ftc_value - expected).abs() < 1e-9,
"FTC check failed: F(2)-F(1) = {ftc_value}, expected {expected}"
);
}
#[test]
fn integrate_x_exp_x() {
let ctx = Context::new();
let x = ctx.symbol("x");
let expr = &x * &x.exp();
let result = expr.integrate(&x);
let s = format!("{result}");
assert!(s.contains("exp(x)"), "should contain exp(x): {s}");
}
#[test]
fn integrate_x_cos_x() {
let ctx = Context::new();
let x = ctx.symbol("x");
let expr = &x * &x.cos();
let result = expr.integrate(&x);
let s = format!("{result}");
assert!(s.contains("sin(x)"), "should contain sin(x): {s}");
assert!(s.contains("cos(x)"), "should contain cos(x): {s}");
}
#[test]
fn args_of_add() {
let ctx = Context::new();
let x = ctx.symbol("x");
let expr = &x + 1;
let children = expr.args();
assert_eq!(children.len(), 2);
}
#[test]
fn args_of_atom() {
let ctx = Context::new();
let x = ctx.symbol("x");
let children = x.args();
assert_eq!(children.len(), 0);
}
#[test]
fn args_of_function() {
let ctx = Context::new();
let x = ctx.symbol("x");
let expr = x.sin();
let children = expr.args();
assert_eq!(children.len(), 1);
}
#[test]
fn diff_n_third_derivative() {
let ctx = Context::new();
let x = ctx.symbol("x");
let f = x.powi(4);
let d3 = f.diff_n(&x, 3);
assert_eq!(format!("{d3}"), "24*x");
}
#[test]
fn diff_n_zero() {
let ctx = Context::new();
let x = ctx.symbol("x");
let f = x.powi(2);
let d0 = f.diff_n(&x, 0);
assert_eq!(format!("{d0}"), "x^2");
}
#[test]
fn log_base_2() {
let ctx = Context::new();
let result = ctx.int(8).log(&ctx.int(2));
let val = result
.eval()
.eval_f64()
.expect("log_2(8) should evaluate to a float");
assert!(
(val - 3.0).abs() < 1e-9,
"log_2(8) should equal 3, got {val}"
);
let s = format!("{}", result.simplify());
assert!(
s == "3" || s.contains("ln"),
"log should simplify to 3 or produce ln expressions: {s}"
);
}