use crate::integrate::IntegrateResult;
use crate::lower::LoweredOp;
use crate::solve::SolveResult;
use std::sync::Arc;
#[derive(Debug, Clone, Copy)]
pub struct OdeForm {
pub x: usize,
pub y: usize,
pub dy: usize,
pub d2y: usize,
pub c_start: usize,
}
impl OdeForm {
pub fn new(n_existing_vars: usize) -> Self {
OdeForm {
x: 0,
y: 1,
dy: 2,
d2y: 3,
c_start: n_existing_vars.max(4),
}
}
fn c1(&self) -> Arc<LoweredOp> {
Arc::new(LoweredOp::Var(self.c_start))
}
fn c2(&self) -> Arc<LoweredOp> {
Arc::new(LoweredOp::Var(self.c_start + 1))
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum OdeSolution {
Explicit(LoweredOp),
Implicit(LoweredOp),
Unsolved,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum OdeKind {
Separable,
FirstOrderLinear,
Exact,
Bernoulli,
SecondOrderConstCoeff,
Unsolved,
}
pub fn dsolve(eq: &LoweredOp, form: &OdeForm) -> (OdeSolution, OdeKind) {
if let Some(sol) = try_second_order_const_coeff(eq, form) {
return (sol, OdeKind::SecondOrderConstCoeff);
}
if let Some(sol) = try_separable(eq, form) {
return (sol, OdeKind::Separable);
}
if let Some(sol) = try_first_order_linear(eq, form) {
return (sol, OdeKind::FirstOrderLinear);
}
if let Some(sol) = try_exact(eq, form) {
return (sol, OdeKind::Exact);
}
if let Some(sol) = try_bernoulli(eq, form) {
return (sol, OdeKind::Bernoulli);
}
(OdeSolution::Unsolved, OdeKind::Unsolved)
}
fn depends_on_var(expr: &LoweredOp, var: usize) -> bool {
expr.contains_var(var)
}
fn integrate_expr(expr: &LoweredOp, wrt: usize) -> Option<LoweredOp> {
match expr.integrate(wrt) {
IntegrateResult::Closed(anti) => Some(anti),
IntegrateResult::Unsupported => None,
}
}
fn solve_for_dy(eq: &LoweredOp, target_var: usize) -> Option<LoweredOp> {
match eq.solve_for(target_var, &LoweredOp::Const(0.0)) {
SolveResult::Closed(sol) => Some(sol),
SolveResult::Residual(_) => None,
}
}
fn recip(expr: LoweredOp) -> LoweredOp {
LoweredOp::Div(Arc::new(LoweredOp::Const(1.0)), Arc::new(expr))
}
fn simplify_exp_ln(expr: &LoweredOp) -> LoweredOp {
match expr {
LoweredOp::Exp(inner) => {
let inner_s = inner.as_ref().clone().simplify();
if let LoweredOp::Mul(a, b) = &inner_s {
if let LoweredOp::Const(c) = a.as_ref() {
if let LoweredOp::Ln(f) = b.as_ref() {
return LoweredOp::Pow(Arc::clone(f), Arc::new(LoweredOp::Const(*c)))
.simplify();
}
}
if let LoweredOp::Const(c) = b.as_ref() {
if let LoweredOp::Ln(f) = a.as_ref() {
return LoweredOp::Pow(Arc::clone(f), Arc::new(LoweredOp::Const(*c)))
.simplify();
}
}
}
if let LoweredOp::Ln(f) = &inner_s {
return f.as_ref().clone();
}
if let LoweredOp::Neg(inner2) = &inner_s {
if let LoweredOp::Ln(f) = inner2.as_ref() {
return recip(f.as_ref().clone());
}
}
expr.clone()
}
_ => expr.clone(),
}
}
fn scalar_mul(s: f64, expr: LoweredOp) -> LoweredOp {
if (s - 1.0).abs() < 1e-14 {
return expr;
}
LoweredOp::Mul(Arc::new(LoweredOp::Const(s)), Arc::new(expr))
}
fn exp_term(coeff: LoweredOp, r: f64, x: &LoweredOp) -> LoweredOp {
if r.abs() < 1e-14 {
return coeff;
}
let rx = scalar_mul(r, x.clone());
let exp_rx = LoweredOp::Exp(Arc::new(rx));
match &coeff {
LoweredOp::Const(v) if (*v - 1.0).abs() < 1e-14 => exp_rx,
_ => LoweredOp::Mul(Arc::new(coeff), Arc::new(exp_rx)),
}
}
fn structural_hash_value(expr: &LoweredOp) -> u64 {
use std::collections::hash_map::DefaultHasher;
use std::hash::Hasher;
let mut h = DefaultHasher::new();
expr.structural_hash(&mut h);
h.finish()
}
fn ops_struct_equal(a: &LoweredOp, b: &LoweredOp) -> bool {
structural_hash_value(a) == structural_hash_value(b) && format!("{a:?}") == format!("{b:?}")
}
fn ops_numeric_equal(a: &LoweredOp, b: &LoweredOp, x_slot: usize, y_slot: usize) -> bool {
let n_vars = x_slot.max(y_slot) + 1;
let mut vars = vec![0.0f64; n_vars];
let test_pts: &[(f64, f64)] = &[
(1.0, 1.0),
(-1.0, 2.0),
(0.5, -0.5),
(2.0, 3.0),
(-2.0, -1.0),
];
for &(xv, yv) in test_pts {
vars[x_slot] = xv;
vars[y_slot] = yv;
let va = a.eval(&vars);
let vb = b.eval(&vars);
if !va.is_finite() || !vb.is_finite() {
continue; }
let diff = (va - vb).abs();
let scale = va.abs().max(vb.abs()).max(1.0);
if diff / scale > 1e-8 {
return false;
}
}
true
}
fn only_x(expr: &LoweredOp, form: &OdeForm) -> bool {
!depends_on_var(expr, form.y) && !depends_on_var(expr, form.dy)
}
fn only_y(expr: &LoweredOp, form: &OdeForm) -> bool {
!depends_on_var(expr, form.x) && !depends_on_var(expr, form.dy)
}
fn separate_x_y(expr: &LoweredOp, form: &OdeForm) -> Option<(LoweredOp, LoweredOp)> {
if depends_on_var(expr, form.dy) {
return None;
}
let dep_x = depends_on_var(expr, form.x);
let dep_y = depends_on_var(expr, form.y);
if !dep_y {
return Some((expr.clone(), LoweredOp::Const(1.0)));
}
if !dep_x {
return Some((LoweredOp::Const(1.0), expr.clone()));
}
if let LoweredOp::Neg(inner) = expr {
if let Some((fx, gy)) = separate_x_y(inner, form) {
let neg_fx = LoweredOp::Neg(Arc::new(fx)).simplify();
return Some((neg_fx, gy));
}
return None;
}
if let LoweredOp::Mul(a, b) = expr {
if only_x(a, form) && only_y(b, form) {
return Some((a.as_ref().clone(), b.as_ref().clone()));
}
if only_y(a, form) && only_x(b, form) {
return Some((b.as_ref().clone(), a.as_ref().clone()));
}
if let Some((fa, ga)) = separate_x_y(a, form) {
if let Some((fb, gb)) = separate_x_y(b, form) {
let fx = LoweredOp::Mul(Arc::new(fa), Arc::new(fb)).simplify();
let gy = LoweredOp::Mul(Arc::new(ga), Arc::new(gb)).simplify();
return Some((fx, gy));
}
}
return None;
}
if let LoweredOp::Div(a, b) = expr {
if only_x(a, form) && only_y(b, form) {
return Some((a.as_ref().clone(), recip(b.as_ref().clone())));
}
if only_y(a, form) && only_x(b, form) {
return Some((recip(b.as_ref().clone()), a.as_ref().clone()));
}
if only_x(b, form) {
if let Some((fa, ga)) = separate_x_y(a, form) {
let fx = LoweredOp::Div(Arc::new(fa), Arc::clone(b)).simplify();
return Some((fx, ga));
}
}
if only_y(b, form) {
if let Some((fa, ga)) = separate_x_y(a, form) {
let gy = LoweredOp::Div(Arc::new(ga), Arc::clone(b)).simplify();
return Some((fa, gy));
}
}
return None;
}
None
}
fn try_separable(eq: &LoweredOp, form: &OdeForm) -> Option<OdeSolution> {
if depends_on_var(eq, form.d2y) {
return None;
}
let rhs = solve_for_dy(eq, form.dy)?;
let (fx, gy) = separate_x_y(&rhs, form)?;
let inv_gy = recip(gy);
let lhs_anti = integrate_expr(&inv_gy, form.y)?;
let rhs_anti = integrate_expr(&fx, form.x)?;
let c1 = form.c1();
let rhs_with_c = LoweredOp::Add(Arc::new(rhs_anti), c1);
let implicit_eq = LoweredOp::Sub(Arc::new(lhs_anti), Arc::new(rhs_with_c));
let implicit_eq = implicit_eq.simplify();
match implicit_eq.solve_for(form.y, &LoweredOp::Const(0.0)) {
SolveResult::Closed(y_sol) => Some(OdeSolution::Explicit(y_sol.simplify())),
SolveResult::Residual(_) => Some(OdeSolution::Implicit(implicit_eq)),
}
}
fn split_linear_in_y(expr: &LoweredOp, form: &OdeForm) -> Option<(LoweredOp, LoweredOp)> {
let grad_y = expr.grad(form.y);
if depends_on_var(&grad_y, form.y) {
return None;
}
let y_op = LoweredOp::Var(form.y);
let linear_term = LoweredOp::Mul(Arc::new(grad_y.clone()), Arc::new(y_op));
let b_raw = LoweredOp::Sub(Arc::new(expr.clone()), Arc::new(linear_term));
let b = b_raw.simplify();
if depends_on_var(&b, form.y) {
return None;
}
Some((grad_y.simplify(), b))
}
fn try_first_order_linear(eq: &LoweredOp, form: &OdeForm) -> Option<OdeSolution> {
if depends_on_var(eq, form.d2y) {
return None;
}
let rhs = solve_for_dy(eq, form.dy)?;
let (neg_p, q) = split_linear_in_y(&rhs, form)?;
let p = LoweredOp::Neg(Arc::new(neg_p)).simplify();
let int_p = integrate_expr(&p, form.x)?;
let mu_raw = LoweredOp::Exp(Arc::new(int_p)).simplify();
let mu = simplify_exp_ln(&mu_raw).simplify();
let mu_q = LoweredOp::Mul(Arc::new(mu.clone()), Arc::new(q)).simplify();
let int_mu_q = integrate_expr(&mu_q, form.x)?;
let c1 = form.c1();
let numerator = LoweredOp::Add(Arc::new(int_mu_q), c1);
let y_sol = LoweredOp::Div(Arc::new(numerator), Arc::new(mu)).simplify();
Some(OdeSolution::Explicit(y_sol))
}
fn subst_var(expr: &LoweredOp, slot: usize, replacement: &LoweredOp) -> LoweredOp {
match expr {
LoweredOp::Var(i) => {
if *i == slot {
replacement.clone()
} else {
expr.clone()
}
}
LoweredOp::Const(_) | LoweredOp::NamedConst(_) => expr.clone(),
LoweredOp::Add(a, b) => LoweredOp::Add(
Arc::new(subst_var(a, slot, replacement)),
Arc::new(subst_var(b, slot, replacement)),
),
LoweredOp::Sub(a, b) => LoweredOp::Sub(
Arc::new(subst_var(a, slot, replacement)),
Arc::new(subst_var(b, slot, replacement)),
),
LoweredOp::Mul(a, b) => LoweredOp::Mul(
Arc::new(subst_var(a, slot, replacement)),
Arc::new(subst_var(b, slot, replacement)),
),
LoweredOp::Div(a, b) => LoweredOp::Div(
Arc::new(subst_var(a, slot, replacement)),
Arc::new(subst_var(b, slot, replacement)),
),
LoweredOp::Neg(a) => LoweredOp::Neg(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Exp(a) => LoweredOp::Exp(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Ln(a) => LoweredOp::Ln(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Sin(a) => LoweredOp::Sin(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Cos(a) => LoweredOp::Cos(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Tan(a) => LoweredOp::Tan(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Sinh(a) => LoweredOp::Sinh(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Cosh(a) => LoweredOp::Cosh(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Tanh(a) => LoweredOp::Tanh(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Arcsin(a) => LoweredOp::Arcsin(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Arccos(a) => LoweredOp::Arccos(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Arctan(a) => LoweredOp::Arctan(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Arcsinh(a) => LoweredOp::Arcsinh(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Arccosh(a) => LoweredOp::Arccosh(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Arctanh(a) => LoweredOp::Arctanh(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Erf(a) => LoweredOp::Erf(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::LGamma(a) => LoweredOp::LGamma(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Digamma(a) => LoweredOp::Digamma(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Trigamma(a) => LoweredOp::Trigamma(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Ei(a) => LoweredOp::Ei(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Si(a) => LoweredOp::Si(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Ci(a) => LoweredOp::Ci(Arc::new(subst_var(a, slot, replacement))),
LoweredOp::Pow(a, b) => LoweredOp::Pow(
Arc::new(subst_var(a, slot, replacement)),
Arc::new(subst_var(b, slot, replacement)),
),
}
}
fn split_m_n(eq: &LoweredOp, form: &OdeForm) -> Option<(LoweredOp, LoweredOp)> {
let n_raw = eq.grad(form.dy);
let n = n_raw.simplify();
let zero = LoweredOp::Const(0.0);
let m = subst_var(eq, form.dy, &zero).simplify();
if depends_on_var(&m, form.dy) {
return None;
}
if depends_on_var(&n, form.dy) {
return None;
}
Some((m, n))
}
fn try_exact(eq: &LoweredOp, form: &OdeForm) -> Option<OdeSolution> {
if depends_on_var(eq, form.d2y) {
return None;
}
let (m, n) = split_m_n(eq, form)?;
let dm_dy = m.grad(form.y).simplify();
let dn_dx = n.grad(form.x).simplify();
if !ops_struct_equal(&dm_dy, &dn_dx) && !ops_numeric_equal(&dm_dy, &dn_dx, form.x, form.y) {
return None;
}
let f_partial = integrate_expr(&m, form.x)?;
let df_partial_dy = f_partial.grad(form.y).simplify();
let g_prime_raw = LoweredOp::Sub(Arc::new(n), Arc::new(df_partial_dy));
let g_prime = g_prime_raw.simplify();
let g_prime_has_x = if depends_on_var(&g_prime, form.x) {
!ops_numeric_equal(&g_prime, &LoweredOp::Const(0.0), form.x, form.y) && {
let n_slots = form.c_start.max(form.d2y + 1);
let mut vars = vec![0.0f64; n_slots];
vars[form.y] = 1.0;
vars[form.x] = 1.0;
let v1 = g_prime.eval(&vars);
vars[form.x] = 2.0;
let v2 = g_prime.eval(&vars);
v1.is_finite() && v2.is_finite() && (v1 - v2).abs() > 1e-8
}
} else {
false
};
if g_prime_has_x {
return None;
}
let g_prime_eff = if ops_numeric_equal(&g_prime, &LoweredOp::Const(0.0), form.x, form.y) {
LoweredOp::Const(0.0)
} else {
g_prime
};
let g = integrate_expr(&g_prime_eff, form.y).unwrap_or(LoweredOp::Const(0.0));
let f = LoweredOp::Add(Arc::new(f_partial), Arc::new(g)).simplify();
let c1 = form.c1();
let implicit_eq = LoweredOp::Sub(Arc::new(f), c1).simplify();
Some(OdeSolution::Implicit(implicit_eq))
}
fn detect_bernoulli_terms(rhs: &LoweredOp, form: &OdeForm) -> Option<(LoweredOp, LoweredOp, f64)> {
let rhs_simplified = rhs.clone().simplify();
let x0 = 1.5_f64; let mut vars = vec![0.0f64; form.c_start + 4];
vars[form.x] = x0;
let probe_y = [1.0_f64, 2.0, 4.0, 0.5];
let probe_rhs: Vec<f64> = probe_y
.iter()
.map(|&yv| {
vars[form.y] = yv;
rhs_simplified.eval(&vars)
})
.collect();
vars[form.y] = 0.0;
let rhs_at_0 = rhs_simplified.eval(&vars);
if rhs_at_0.is_finite() && rhs_at_0.abs() > 1e-6 {
return None; }
let g: Vec<f64> = probe_y
.iter()
.zip(&probe_rhs)
.map(|(&y, &r)| r / y)
.collect();
if !g[0].is_finite() || !g[1].is_finite() || !g[2].is_finite() {
return None;
}
let candidates = [2.0_f64, 3.0, -1.0, 0.5, 4.0, 1.0 / 3.0, -2.0];
let mut found_n = None;
for &n_cand in &candidates {
if (n_cand - 1.0).abs() < 1e-10 || n_cand.abs() < 1e-10 {
continue;
}
let exp_ratio =
if (probe_y[2].powf(n_cand - 1.0) - probe_y[0].powf(n_cand - 1.0)).abs() > 1e-10 {
(probe_y[1].powf(n_cand - 1.0) - probe_y[0].powf(n_cand - 1.0))
/ (probe_y[2].powf(n_cand - 1.0) - probe_y[0].powf(n_cand - 1.0))
} else {
continue;
};
let obs_ratio = if (g[2] - g[0]).abs() > 1e-10 {
(g[1] - g[0]) / (g[2] - g[0])
} else {
continue;
};
if (exp_ratio - obs_ratio).abs() < 1e-6 {
found_n = Some(n_cand);
break;
}
}
let n_exp = found_n?;
let grad_y = rhs_simplified.grad(form.y).simplify();
let n_int = n_exp.round() as i32;
if (n_exp - n_int as f64).abs() > 1e-10 || n_int < 2 {
return detect_bernoulli_numeric(rhs, form, n_exp);
}
let mut deriv = rhs_simplified.clone();
for _ in 0..n_int {
deriv = deriv.grad(form.y).simplify();
}
let n_factorial: f64 = (1..=n_int).map(|k| k as f64).product();
let q = LoweredOp::Div(Arc::new(deriv), Arc::new(LoweredOp::Const(n_factorial))).simplify();
if depends_on_var(&q, form.y) {
return None; }
let neg_p = subst_var(&grad_y, form.y, &LoweredOp::Const(0.0)).simplify();
if depends_on_var(&neg_p, form.y) {
return None;
}
let p = LoweredOp::Neg(Arc::new(neg_p)).simplify();
if depends_on_var(&p, form.y) {
return None;
}
Some((p, q, n_exp))
}
fn detect_bernoulli_numeric(
rhs: &LoweredOp,
form: &OdeForm,
n_exp: f64,
) -> Option<(LoweredOp, LoweredOp, f64)> {
let _ = (rhs, form, n_exp);
None
}
fn try_bernoulli(eq: &LoweredOp, form: &OdeForm) -> Option<OdeSolution> {
if depends_on_var(eq, form.d2y) {
return None;
}
let rhs = solve_for_dy(eq, form.dy)?;
let (p, q, n) = detect_bernoulli_terms(&rhs, form)?;
let one_minus_n = 1.0 - n;
let scale = LoweredOp::Const(one_minus_n);
let new_p = LoweredOp::Mul(Arc::new(scale.clone()), Arc::new(p));
let new_q = LoweredOp::Mul(Arc::new(scale), Arc::new(q));
let new_p_s = new_p.simplify();
let int_new_p = integrate_expr(&new_p_s, form.x)?;
let mu_raw = LoweredOp::Exp(Arc::new(int_new_p)).simplify();
let mu = simplify_exp_ln(&mu_raw).simplify();
let new_q_s = new_q.simplify();
let mu_new_q = LoweredOp::Mul(Arc::new(mu.clone()), Arc::new(new_q_s)).simplify();
let int_mu_new_q = integrate_expr(&mu_new_q, form.x)?;
let c1 = form.c1();
let v_numerator = LoweredOp::Add(Arc::new(int_mu_new_q), c1);
let v_sol = LoweredOp::Div(Arc::new(v_numerator), Arc::new(mu)).simplify();
let exp_back = LoweredOp::Const(1.0 / one_minus_n);
let y_sol = LoweredOp::Pow(Arc::new(v_sol), Arc::new(exp_back)).simplify();
Some(OdeSolution::Explicit(y_sol))
}
fn extract_const_coeff_2nd(eq: &LoweredOp, form: &OdeForm) -> Option<(f64, f64, f64)> {
let da = eq.grad(form.d2y).simplify();
let db = eq.grad(form.dy).simplify();
let dc = eq.grad(form.y).simplify();
let a = match &da {
LoweredOp::Const(v) => *v,
_ => return None,
};
let b = match &db {
LoweredOp::Const(v) => *v,
_ => return None,
};
let c = match &dc {
LoweredOp::Const(v) => *v,
_ => return None,
};
let recon = LoweredOp::Add(
Arc::new(LoweredOp::Add(
Arc::new(scalar_mul(a, LoweredOp::Var(form.d2y))),
Arc::new(scalar_mul(b, LoweredOp::Var(form.dy))),
)),
Arc::new(scalar_mul(c, LoweredOp::Var(form.y))),
);
let remainder = LoweredOp::Sub(Arc::new(eq.clone()), Arc::new(recon)).simplify();
match &remainder {
LoweredOp::Const(v) if v.abs() < 1e-12 => {}
_ => {
let test_pts: &[(f64, f64, f64, f64)] = &[
(1.0, 2.0, 3.0, 4.0),
(-1.0, 0.5, 2.0, -3.0),
(0.0, 1.0, -1.0, 0.5),
];
for &(xv, yv, dyv, d2yv) in test_pts {
let mut vars = vec![0.0f64; form.c_start + 4];
vars[form.x] = xv;
vars[form.y] = yv;
vars[form.dy] = dyv;
vars[form.d2y] = d2yv;
let eq_val = eq.eval(&vars);
let recon_val = a * d2yv + b * dyv + c * yv;
if (eq_val - recon_val).abs() > 1e-8 {
return None;
}
}
}
}
Some((a, b, c))
}
fn try_second_order_const_coeff(eq: &LoweredOp, form: &OdeForm) -> Option<OdeSolution> {
if !depends_on_var(eq, form.d2y) {
return None;
}
let (a, b, c) = extract_const_coeff_2nd(eq, form)?;
if a.abs() < 1e-14 {
return None;
}
let disc = b * b - 4.0 * a * c;
let x_op = LoweredOp::Var(form.x);
let c1 = form.c1();
let c2 = form.c2();
if disc > 1e-10 {
let sq = disc.sqrt();
let r1 = (-b - sq) / (2.0 * a);
let r2 = (-b + sq) / (2.0 * a);
let t1 = exp_term(c1.as_ref().clone(), r1, &x_op);
let t2 = exp_term(c2.as_ref().clone(), r2, &x_op);
let y_sol = LoweredOp::Add(Arc::new(t1), Arc::new(t2)).simplify();
Some(OdeSolution::Explicit(y_sol))
} else if disc.abs() <= 1e-10 {
let r = -b / (2.0 * a);
let c2_x = LoweredOp::Mul(Arc::clone(&c2), Arc::new(x_op.clone()));
let bracket = LoweredOp::Add(Arc::clone(&c1), Arc::new(c2_x));
let exp_rx = exp_term(LoweredOp::Const(1.0), r, &x_op);
let y_sol = LoweredOp::Mul(Arc::new(bracket), Arc::new(exp_rx)).simplify();
Some(OdeSolution::Explicit(y_sol))
} else {
let alpha = -b / (2.0 * a);
let beta = (-disc).sqrt() / (2.0 * a);
let exp_ax = exp_term(LoweredOp::Const(1.0), alpha, &x_op);
let bx = scalar_mul(beta, x_op.clone());
let cos_bx = LoweredOp::Cos(Arc::new(bx.clone()));
let sin_bx = LoweredOp::Sin(Arc::new(bx));
let c1_cos = LoweredOp::Mul(Arc::clone(&c1), Arc::new(cos_bx));
let c2_sin = LoweredOp::Mul(Arc::clone(&c2), Arc::new(sin_bx));
let trig_sum = LoweredOp::Add(Arc::new(c1_cos), Arc::new(c2_sin));
let y_sol = LoweredOp::Mul(Arc::new(exp_ax), Arc::new(trig_sum)).simplify();
Some(OdeSolution::Explicit(y_sol))
}
}
#[cfg(test)]
mod ode_tests {
use super::*;
fn c(v: f64) -> LoweredOp {
LoweredOp::Const(v)
}
fn var(i: usize) -> LoweredOp {
LoweredOp::Var(i)
}
fn mul(a: LoweredOp, b: LoweredOp) -> LoweredOp {
LoweredOp::Mul(Arc::new(a), Arc::new(b))
}
fn add(a: LoweredOp, b: LoweredOp) -> LoweredOp {
LoweredOp::Add(Arc::new(a), Arc::new(b))
}
fn sub(a: LoweredOp, b: LoweredOp) -> LoweredOp {
LoweredOp::Sub(Arc::new(a), Arc::new(b))
}
fn pow(base: LoweredOp, exp: LoweredOp) -> LoweredOp {
LoweredOp::Pow(Arc::new(base), Arc::new(exp))
}
fn default_form() -> OdeForm {
OdeForm {
x: 0,
y: 1,
dy: 2,
d2y: 3,
c_start: 10,
}
}
fn eval_solution(sol: &LoweredOp, form: &OdeForm, x_val: f64, c1_val: f64) -> f64 {
let mut vars = vec![0.0f64; form.c_start + 4];
vars[form.x] = x_val;
vars[form.c_start] = c1_val;
sol.eval(&vars)
}
#[test]
fn test_separable_y_prime_eq_xy() {
let form = default_form();
let eq = sub(var(form.dy), mul(var(form.x), var(form.y)));
let (sol, kind) = dsolve(&eq, &form);
assert_eq!(kind, OdeKind::Separable, "Should recognise separable ODE");
assert!(
matches!(sol, OdeSolution::Explicit(_) | OdeSolution::Implicit(_)),
"Should return a solution, got {sol:?}"
);
}
#[test]
fn test_separable_y_prime_eq_x() {
let form = default_form();
let eq = sub(var(form.dy), var(form.x));
let (sol, kind) = dsolve(&eq, &form);
assert_eq!(kind, OdeKind::Separable, "y′=x should be separable");
assert!(
matches!(sol, OdeSolution::Explicit(_) | OdeSolution::Implicit(_)),
"Should return a solution"
);
}
#[test]
fn test_first_order_linear_y_prime_plus_y_eq_x() {
let form = default_form();
let eq = sub(add(var(form.dy), var(form.y)), var(form.x));
let (sol, kind) = dsolve(&eq, &form);
assert_eq!(
kind,
OdeKind::FirstOrderLinear,
"Should recognise first-order linear, got {:?}",
kind
);
assert!(
matches!(sol, OdeSolution::Explicit(_)),
"Should return explicit solution"
);
if let OdeSolution::Explicit(y_expr) = &sol {
let y0 = eval_solution(y_expr, &form, 0.0, 1.0);
assert!(y0.is_finite(), "Solution should be finite at x=0");
}
}
#[test]
fn test_exact_2xy_dx_plus_x2_dy() {
let form = default_form();
let m = mul(mul(c(2.0), var(form.x)), var(form.y));
let x2 = mul(var(form.x), var(form.x));
let n_dy = mul(x2, var(form.dy));
let eq = add(m, n_dy);
let (sol, kind) = dsolve(&eq, &form);
assert!(
kind == OdeKind::Exact
|| kind == OdeKind::Separable
|| kind == OdeKind::FirstOrderLinear,
"Should recognise as exact, separable, or linear, got {kind:?}"
);
assert!(
matches!(sol, OdeSolution::Implicit(_) | OdeSolution::Explicit(_)),
"Should return a solution"
);
}
#[test]
fn test_second_order_cc_two_real_roots() {
let form = default_form();
let eq = add(
sub(var(form.d2y), mul(c(3.0), var(form.dy))),
mul(c(2.0), var(form.y)),
);
let (sol, kind) = dsolve(&eq, &form);
assert_eq!(kind, OdeKind::SecondOrderConstCoeff);
assert!(matches!(sol, OdeSolution::Explicit(_)));
}
#[test]
fn test_second_order_cc_complex_roots() {
let form = default_form();
let eq = add(var(form.d2y), var(form.y));
let (sol, kind) = dsolve(&eq, &form);
assert_eq!(kind, OdeKind::SecondOrderConstCoeff);
assert!(matches!(sol, OdeSolution::Explicit(_)));
}
#[test]
fn test_second_order_cc_repeated_root() {
let form = default_form();
let eq = add(sub(var(form.d2y), mul(c(2.0), var(form.dy))), var(form.y));
let (sol, kind) = dsolve(&eq, &form);
assert_eq!(kind, OdeKind::SecondOrderConstCoeff);
assert!(matches!(sol, OdeSolution::Explicit(_)));
}
#[test]
fn test_second_order_cc_pure_exponential() {
let form = default_form();
let eq = sub(var(form.d2y), var(form.y));
let (sol, kind) = dsolve(&eq, &form);
assert_eq!(kind, OdeKind::SecondOrderConstCoeff);
assert!(matches!(sol, OdeSolution::Explicit(_)));
}
#[test]
fn test_bernoulli_y_prime_plus_y_eq_y_squared() {
let form = default_form();
let y_sq = pow(var(form.y), c(2.0));
let eq = sub(add(var(form.dy), var(form.y)), y_sq);
let (sol, kind) = dsolve(&eq, &form);
assert!(
kind == OdeKind::Bernoulli || kind == OdeKind::Separable,
"Should recognise as Bernoulli or Separable (both valid), got {kind:?}"
);
assert!(
matches!(sol, OdeSolution::Explicit(_) | OdeSolution::Implicit(_)),
"Should return a solution"
);
}
#[test]
fn test_unrecognised_ode_returns_unsolved() {
let form = default_form();
let xy = mul(var(form.x), var(form.y));
let sin_xy = LoweredOp::Sin(Arc::new(xy));
let eq = sub(var(form.dy), sin_xy);
let (sol, kind) = dsolve(&eq, &form);
assert_eq!(kind, OdeKind::Unsolved);
assert!(matches!(sol, OdeSolution::Unsolved));
}
#[test]
fn test_ode_form_new() {
let form = OdeForm::new(10);
assert_eq!(form.x, 0);
assert_eq!(form.y, 1);
assert_eq!(form.dy, 2);
assert_eq!(form.d2y, 3);
assert_eq!(form.c_start, 10);
let form_small = OdeForm::new(2);
assert_eq!(form_small.c_start, 4); }
#[test]
fn test_separable_pure_y() {
let form = default_form();
let eq = sub(var(form.dy), var(form.y));
let (sol, kind) = dsolve(&eq, &form);
assert_eq!(kind, OdeKind::Separable, "y′=y should be separable");
assert!(matches!(sol, OdeSolution::Explicit(_)));
}
#[test]
fn test_first_order_linear_y_prime_eq_neg_y() {
let form = default_form();
let eq = add(var(form.dy), var(form.y));
let (sol, kind) = dsolve(&eq, &form);
assert!(
kind == OdeKind::Separable || kind == OdeKind::FirstOrderLinear,
"y′+y=0 should be separable or linear, got {kind:?}"
);
assert!(matches!(sol, OdeSolution::Explicit(_)));
}
}