use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use super::ODESolution;
use super::util::ode_order;
use crate::derivative::diff;
pub(crate) fn solve_power_series<'a>(
ctx: &'a AtomArena<'a>,
ode: super::ODE<'a>,
x0: Atom<'a>,
n_terms: usize,
) -> Option<ODESolution<'a>> {
let super::ODE {
equation,
func,
var,
} = ode;
let order = ode_order(equation, func, var);
if order == 0 {
return None;
}
let x = ctx.var(var.as_str());
let h = ctx.add(&[x, ctx.mul(&[ctx.num(-1), x0])]);
let max_coeff = n_terms;
let mut a_syms: Vec<Atom<'a>> = Vec::with_capacity(max_coeff);
for i in 0..max_coeff {
a_syms.push(ctx.var(&format!("a{i}")));
}
let (y_series, y1, y2) = build_series_triple(ctx, &a_syms, h);
let mut residual = substitute_series(ctx, equation, func, var, y_series, y1, y2);
residual = super::util::collect_terms(ctx, residual);
let mut solved: Vec<(Atom<'a>, Atom<'a>)> = Vec::new();
for k in 0..n_terms.saturating_sub(order) {
if k > 0 {
residual = super::util::collect_terms(ctx, diff(ctx, residual, var));
}
let mut cond = super::classify::replace_atom(ctx, residual, x, x0);
cond = super::util::collect_terms(ctx, cond);
for (sym, val) in &solved {
cond = super::classify::replace_atom(ctx, cond, *sym, *val);
}
cond = super::util::collect_terms(ctx, cond);
let target = a_syms.iter().skip(order).rev().find(|s| {
cond.to_string().contains(&s.to_string())
&& !solved
.iter()
.any(|(sym, _)| sym.to_string() == s.to_string())
});
match target {
Some(&target) => {
if let Some(value) = solve_linear_coeff(ctx, cond, target) {
solved.push((target, value));
}
}
None => {
if !matches!(cond.node(), AtomNode::Num(0)) {
return None;
}
}
}
}
if solved.is_empty() {
return None;
}
let mut series_expr = y_series;
for (sym, val) in &solved {
series_expr = super::classify::replace_atom(ctx, series_expr, *sym, *val);
}
series_expr = super::util::collect_terms(ctx, series_expr);
Some(ODESolution::Series(series_expr, n_terms))
}
fn solve_linear_coeff<'a>(
ctx: &'a AtomArena<'a>,
cond: Atom<'a>,
target: Atom<'a>,
) -> Option<Atom<'a>> {
let terms: Vec<Atom<'a>> = match cond.node() {
AtomNode::Add(args) => args.to_vec(),
_ => vec![cond],
};
let target_str = target.to_string();
let mut coeff_terms: Vec<Atom<'a>> = Vec::new();
let mut rest_terms: Vec<Atom<'a>> = Vec::new();
for term in terms {
let factors: Vec<Atom<'a>> = match term.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![term],
};
let target_count = factors
.iter()
.filter(|f| f.to_string() == target_str)
.count();
if target_count == 1 {
let rest: Vec<_> = factors
.iter()
.filter(|f| f.to_string() != target_str)
.copied()
.collect();
if rest.iter().any(|f| f.to_string().contains(&target_str)) {
return None;
}
coeff_terms.push(if rest.is_empty() {
ctx.num(1)
} else {
ctx.mul(&rest)
});
} else if target_count == 0 {
if term.to_string().contains(&target_str) {
return None; }
rest_terms.push(term);
} else {
return None; }
}
if coeff_terms.is_empty() {
return None;
}
let coeff = super::util::collect_terms(ctx, ctx.add(&coeff_terms));
if matches!(coeff.node(), AtomNode::Num(0)) {
return None;
}
let rest = if rest_terms.is_empty() {
ctx.num(0)
} else {
super::util::collect_terms(ctx, ctx.add(&rest_terms))
};
let value = super::util::collect_terms(
ctx,
ctx.mul(&[ctx.num(-1), rest, ctx.pow(coeff, ctx.num(-1))]),
);
Some(value)
}
pub(crate) fn solve_frobenius<'a>(
ctx: &'a AtomArena<'a>,
ode: super::ODE<'a>,
x0: Atom<'a>,
n_terms: usize,
) -> Option<ODESolution<'a>> {
let super::ODE {
equation,
func,
var,
} = ode;
let order = ode_order(equation, func, var);
if order != 2 {
return None;
}
if !matches!(x0.node(), AtomNode::Num(0)) {
return None;
}
let x = ctx.var(var.as_str());
let (a, b, c, forcing) =
super::second_order::extract_second_order_coeffs(ctx, equation, func, var)?;
if !matches!(
super::util::collect_terms(ctx, forcing).node(),
AtomNode::Num(0)
) {
return None;
}
let u = ctx.var("XR");
let r = ctx.var("r");
let a_syms: Vec<Atom<'a>> = (0..n_terms).map(|i| ctx.var(&format!("a{i}"))).collect();
let (s, s1, s2) = build_series_triple(ctx, &a_syms, x);
let y = ctx.mul(&[u, s]);
let x_inv = ctx.pow(x, ctx.num(-1));
let x_inv2 = ctx.pow(x, ctx.num(-2));
let y1 = ctx.add(&[ctx.mul(&[r, x_inv, u, s]), ctx.mul(&[u, s1])]);
let y2 = ctx.add(&[
ctx.mul(&[r, ctx.add(&[r, ctx.num(-1)]), x_inv2, u, s]),
ctx.mul(&[ctx.num(2), r, x_inv, u, s1]),
ctx.mul(&[u, s2]),
]);
let residual = super::util::collect_terms(
ctx,
ctx.add(&[ctx.mul(&[a, y2]), ctx.mul(&[b, y1]), ctx.mul(&[c, y])]),
);
let terms: Vec<Atom<'a>> = match residual.node() {
AtomNode::Add(args) => args.to_vec(),
_ => vec![residual],
};
let mut groups: Vec<(i64, Vec<Atom<'a>>)> = Vec::new();
for term in terms {
let (power, stripped) = strip_x_and_u(ctx, term, x, u)?;
if let Some(g) = groups.iter_mut().find(|(p, _)| *p == power) {
g.1.push(stripped);
} else {
groups.push((power, vec![stripped]));
}
}
groups.sort_by_key(|(p, _)| *p);
let (_, lowest_terms) = groups.first()?.clone();
let (ca, cb, cc) = indicial_coeffs(ctx, &lowest_terms, r, a_syms[0])?;
let disc = cb * cb - 4 * ca * cc;
if disc < 0 || ca == 0 {
return None;
}
let sd = isqrt_i64(disc);
if sd * sd != disc {
return None; }
let r1_num = -cb + sd;
let r1_den = 2 * ca;
if r1_den == 0 {
return None;
}
let g = gcd_i64(r1_num.unsigned_abs() as i64, r1_den.unsigned_abs() as i64).max(1);
let (rn, rd) = (r1_num / g, r1_den / g);
let r_val = if rd == 1 {
ctx.num(rn)
} else {
ctx.mul(&[ctx.num(rn), ctx.pow(ctx.num(rd), ctx.num(-1))])
};
let mut solved: Vec<(Atom<'a>, Atom<'a>)> = Vec::new();
for (idx, (_, g_terms)) in groups.iter().enumerate() {
let mut cond = ctx.add(g_terms);
cond = super::classify::replace_atom(ctx, cond, r, r_val);
cond = super::util::collect_terms(ctx, cond);
for (sym, val) in &solved {
cond = super::classify::replace_atom(ctx, cond, *sym, *val);
}
cond = super::util::collect_terms(ctx, cond);
if idx == 0 {
if !matches!(cond.node(), AtomNode::Num(0)) {
return None;
}
continue;
}
let target = a_syms.iter().skip(1).rev().find(|s| {
cond.to_string().contains(&s.to_string())
&& !solved
.iter()
.any(|(sym, _)| sym.to_string() == s.to_string())
});
let Some(&target) = target else {
continue;
};
let value = solve_linear_coeff(ctx, cond, target)?;
solved.push((target, value));
}
if solved.is_empty() {
return None;
}
let mut series = s;
for (sym, val) in &solved {
series = super::classify::replace_atom(ctx, series, *sym, *val);
}
series = super::util::collect_terms(ctx, series);
let y_sol = super::util::collect_terms(ctx, ctx.mul(&[ctx.pow(x, r_val), series]));
Some(ODESolution::Series(y_sol, n_terms))
}
fn build_series_triple<'a>(
ctx: &'a AtomArena<'a>,
a_syms: &[Atom<'a>],
x: Atom<'a>,
) -> (Atom<'a>, Atom<'a>, Atom<'a>) {
let mut s_terms = Vec::with_capacity(a_syms.len());
let mut s1_terms = Vec::new();
let mut s2_terms = Vec::new();
for (n, an) in a_syms.iter().enumerate() {
if n == 0 {
s_terms.push(*an);
} else {
s_terms.push(ctx.mul(&[*an, ctx.pow(x, ctx.num(n as i64))]));
let d1 = if n == 1 {
ctx.mul(&[ctx.num(n as i64), *an])
} else {
ctx.mul(&[ctx.num(n as i64), *an, ctx.pow(x, ctx.num(n as i64 - 1))])
};
s1_terms.push(d1);
if n >= 2 {
let d2 = if n == 2 {
ctx.mul(&[ctx.num((n * (n - 1)) as i64), *an])
} else {
ctx.mul(&[
ctx.num((n * (n - 1)) as i64),
*an,
ctx.pow(x, ctx.num(n as i64 - 2)),
])
};
s2_terms.push(d2);
}
}
}
let s = ctx.add(&s_terms);
let s1 = if s1_terms.is_empty() {
ctx.num(0)
} else {
ctx.add(&s1_terms)
};
let s2 = if s2_terms.is_empty() {
ctx.num(0)
} else {
ctx.add(&s2_terms)
};
(s, s1, s2)
}
fn strip_x_and_u<'a>(
ctx: &'a AtomArena<'a>,
term: Atom<'a>,
x: Atom<'a>,
u: Atom<'a>,
) -> Option<(i64, Atom<'a>)> {
let factors: Vec<Atom<'a>> = match term.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![term],
};
let mut power: i64 = 0;
let mut rest: Vec<Atom<'a>> = Vec::new();
for f in factors {
if f.to_string() == u.to_string() {
continue; }
match f.node() {
AtomNode::Var(v) if *v == Symbol::new("x") && f.to_string() == x.to_string() => {
power += 1;
}
AtomNode::Pow(base, exp) => {
if base.to_string() == x.to_string() {
if let AtomNode::Num(n) = exp.node() {
power += n;
} else {
return None;
}
} else {
rest.push(f);
}
}
_ => rest.push(f),
}
}
let stripped = if rest.is_empty() {
ctx.num(1)
} else {
ctx.mul(&rest)
};
Some((power, stripped))
}
fn indicial_coeffs<'a>(
ctx: &'a AtomArena<'a>,
terms: &[Atom<'a>],
r: Atom<'a>,
a0: Atom<'a>,
) -> Option<(i64, i64, i64)> {
let r_str = r.to_string();
let a0_str = a0.to_string();
let mut ca = 0i64;
let mut cb = 0i64;
let mut cc = 0i64;
for term in terms {
let factors: Vec<Atom<'a>> = match term.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![*term],
};
let mut r_pow = 0i64;
let mut num: i64 = 1;
for f in factors {
match f.node() {
AtomNode::Num(n) => num *= n,
AtomNode::Var(_) if f.to_string() == r_str => r_pow += 1,
AtomNode::Pow(base, exp) => {
if base.to_string() == r_str {
if let AtomNode::Num(n) = exp.node() {
r_pow += n;
} else {
return None;
}
} else {
return None;
}
}
AtomNode::Add(args) => {
let mut has_r = false;
let mut has_m1 = false;
for aa in args.iter() {
if aa.to_string() == r_str {
has_r = true;
} else if matches!(aa.node(), AtomNode::Num(-1)) {
has_m1 = true;
}
}
if has_r && has_m1 && args.len() == 2 {
return expand_indicial(ctx, terms, r, a0);
}
return None;
}
_ => {
if f.to_string() != a0_str {
return None;
}
}
}
}
match r_pow {
2 => ca += num,
1 => cb += num,
0 => cc += num,
_ => return None,
}
}
Some((ca, cb, cc))
}
fn expand_indicial<'a>(
ctx: &'a AtomArena<'a>,
terms: &[Atom<'a>],
r: Atom<'a>,
a0: Atom<'a>,
) -> Option<(i64, i64, i64)> {
let mut expanded: Vec<Atom<'a>> = Vec::new();
for term in terms {
let factors: Vec<Atom<'a>> = match term.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![*term],
};
let mut current: Vec<Atom<'a>> = vec![ctx.num(1)];
for f in factors {
if let AtomNode::Add(args) = f.node() {
let mut next: Vec<Atom<'a>> = Vec::new();
for c in ¤t {
for aa in args.iter() {
next.push(ctx.mul(&[*c, *aa]));
}
}
current = next;
} else {
for c in &mut current {
*c = ctx.mul(&[*c, f]);
}
}
}
expanded.extend(current);
}
let expanded_sum = super::util::collect_terms(ctx, ctx.add(&expanded));
let new_terms: Vec<Atom<'a>> = match expanded_sum.node() {
AtomNode::Add(args) => args.to_vec(),
_ => vec![expanded_sum],
};
let r_str = r.to_string();
let a0_str = a0.to_string();
let mut ca = 0i64;
let mut cb = 0i64;
let mut cc = 0i64;
for term in new_terms {
let factors: Vec<Atom<'a>> = match term.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![term],
};
let mut r_pow = 0i64;
let mut num: i64 = 1;
for f in factors {
match f.node() {
AtomNode::Num(n) => num *= n,
AtomNode::Var(_) if f.to_string() == r_str => r_pow += 1,
_ => {
if f.to_string() != a0_str {
return None;
}
}
}
}
match r_pow {
2 => ca += num,
1 => cb += num,
0 => cc += num,
_ => return None,
}
}
Some((ca, cb, cc))
}
fn isqrt_i64(n: i64) -> i64 {
if n <= 0 {
return 0;
}
let mut x = n;
let mut y = x / 2 + 1;
while y < x {
x = y;
y = (x + n / x) / 2;
}
x
}
fn gcd_i64(mut a: i64, mut b: i64) -> i64 {
while b != 0 {
let t = b;
b = a % b;
a = t;
}
a.abs()
}
fn substitute_series<'a>(
ctx: &'a AtomArena<'a>,
equation: Atom<'a>,
func: Atom<'a>,
var: Symbol,
y_val: Atom<'a>,
y1_val: Atom<'a>,
y2_val: Atom<'a>,
) -> Atom<'a> {
substitute_series_inner(ctx, equation, func, var, y_val, y1_val, y2_val)
}
fn substitute_series_inner<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
func: Atom<'a>,
var: Symbol,
y_val: Atom<'a>,
y1_val: Atom<'a>,
y2_val: Atom<'a>,
) -> Atom<'a> {
match expr.node() {
AtomNode::Num(_) | AtomNode::Var(_) => expr,
AtomNode::Add(args) => {
let mapped: Vec<_> = args
.iter()
.map(|a| substitute_series_inner(ctx, *a, func, var, y_val, y1_val, y2_val))
.collect();
ctx.add(&mapped)
}
AtomNode::Mul(args) => {
let mapped: Vec<_> = args
.iter()
.map(|a| substitute_series_inner(ctx, *a, func, var, y_val, y1_val, y2_val))
.collect();
ctx.mul(&mapped)
}
AtomNode::Pow(base, exp) => {
let b = substitute_series_inner(ctx, *base, func, var, y_val, y1_val, y2_val);
let e = substitute_series_inner(ctx, *exp, func, var, y_val, y1_val, y2_val);
ctx.pow(b, e)
}
AtomNode::Fun(name, args) => {
if *name == Symbol::new("Derivative") && args.len() >= 2 {
let is_func =
args[0].to_string() == func.to_string() && args[1].to_string() == var.as_str();
if is_func {
let deriv_order = args.len() - 1;
return match deriv_order {
1 => y1_val,
2 => y2_val,
_ => {
let mut result = y_val;
for _ in 0..deriv_order {
result = diff(ctx, result, var);
}
result
}
};
}
}
if expr.to_string() == func.to_string() {
return y_val;
}
let mapped: Vec<_> = args
.iter()
.map(|a| substitute_series_inner(ctx, *a, func, var, y_val, y1_val, y2_val))
.collect();
ctx.fun(name.as_str(), &mapped)
}
}
}