use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use super::ODESolution;
use super::util::collect_terms;
pub(crate) fn solve_linear_system<'a>(
ctx: &'a AtomArena<'a>,
equations: &[Atom<'a>],
funcs: &[Atom<'a>],
var: Symbol,
) -> Option<ODESolution<'a>> {
if equations.len() != 2 || funcs.len() != 2 {
return None;
}
let x = ctx.var(var.as_str());
let a11 = extract_coeff(ctx, equations[0], funcs[0], var)?;
let a12 = extract_coeff(ctx, equations[0], funcs[1], var)?;
let a21 = extract_coeff(ctx, equations[1], funcs[0], var)?;
let a22 = extract_coeff(ctx, equations[1], funcs[1], var)?;
let trace = a11 + a22;
let det = a11 * a22 - a12 * a21;
let disc = trace * trace - 4 * det;
if disc < 0 {
if trace % 2 != 0 {
return None;
}
let beta_sq = -disc;
let sb = isqrt_i64(beta_sq);
if sb * sb != beta_sq || sb % 2 != 0 {
return None;
}
let alpha = trace / 2;
let beta = sb / 2;
return Some(solve_complex_2x2(ctx, x, a11, a12, a21, alpha, beta));
}
let sd = isqrt_i64(disc);
if sd * sd != disc || (trace + sd) % 2 != 0 {
return None; }
let l1 = (trace + sd) / 2;
let l2 = (trace - sd) / 2;
let c1 = ctx.var("C1");
let c2 = ctx.var("C2");
let e1 = ctx.fun("exp", &[ctx.mul(&[ctx.num(l1), x])]);
if l1 == l2 {
let e_lx = e1;
if a11 == l1 && a12 == 0 && a21 == 0 && a22 == l1 {
let y1 = ctx.mul(&[c1, e_lx]);
let y2 = ctx.mul(&[c2, e_lx]);
return Some(ODESolution::System(ctx.slice(&[y1, y2])));
}
let (v1, v2) = eigenvector(a11, a12, a21, a22, l1)?;
let (w1, w2) = generalized_eigenvector(a11, a12, a21, a22, l1, v1, v2)?;
let xv1_w1 = collect_terms(ctx, ctx.add(&[ctx.mul(&[x, ctx.num(v1)]), ctx.num(w1)]));
let xv2_w2 = collect_terms(ctx, ctx.add(&[ctx.mul(&[x, ctx.num(v2)]), ctx.num(w2)]));
let y1 = collect_terms(
ctx,
ctx.mul(&[
e_lx,
ctx.add(&[ctx.mul(&[c1, ctx.num(v1)]), ctx.mul(&[c2, xv1_w1])]),
]),
);
let y2 = collect_terms(
ctx,
ctx.mul(&[
e_lx,
ctx.add(&[ctx.mul(&[c1, ctx.num(v2)]), ctx.mul(&[c2, xv2_w2])]),
]),
);
return Some(ODESolution::System(ctx.slice(&[y1, y2])));
}
let e2 = ctx.fun("exp", &[ctx.mul(&[ctx.num(l2), x])]);
let (v11, v12) = eigenvector(a11, a12, a21, a22, l1)?;
let (v21, v22) = eigenvector(a11, a12, a21, a22, l2)?;
let y1 = collect_terms(
ctx,
ctx.add(&[
ctx.mul(&[c1, ctx.num(v11), e1]),
ctx.mul(&[c2, ctx.num(v21), e2]),
]),
);
let y2 = collect_terms(
ctx,
ctx.add(&[
ctx.mul(&[c1, ctx.num(v12), e1]),
ctx.mul(&[c2, ctx.num(v22), e2]),
]),
);
Some(ODESolution::System(ctx.slice(&[y1, y2])))
}
fn solve_complex_2x2<'a>(
ctx: &'a AtomArena<'a>,
x: Atom<'a>,
a11: i64,
a12: i64,
a21: i64,
alpha: i64,
beta: i64,
) -> ODESolution<'a> {
let p1 = a12;
let p2 = alpha - a11;
let q1 = 0i64;
let q2 = beta;
let _ = a21;
let c1 = ctx.var("C1");
let c2 = ctx.var("C2");
let bx = ctx.mul(&[ctx.num(beta), x]);
let cos_bx = ctx.fun("cos", &[bx]);
let sin_bx = ctx.fun("sin", &[bx]);
let e_ax = if alpha == 0 {
ctx.num(1)
} else {
ctx.fun("exp", &[ctx.mul(&[ctx.num(alpha), x])])
};
let y1_inner = collect_terms(
ctx,
ctx.add(&[
ctx.mul(&[
c1,
ctx.add(&[
ctx.mul(&[ctx.num(p1), cos_bx]),
ctx.mul(&[ctx.num(-q1), sin_bx]),
]),
]),
ctx.mul(&[
c2,
ctx.add(&[
ctx.mul(&[ctx.num(p1), sin_bx]),
ctx.mul(&[ctx.num(q1), cos_bx]),
]),
]),
]),
);
let y2_inner = collect_terms(
ctx,
ctx.add(&[
ctx.mul(&[
c1,
ctx.add(&[
ctx.mul(&[ctx.num(p2), cos_bx]),
ctx.mul(&[ctx.num(-q2), sin_bx]),
]),
]),
ctx.mul(&[
c2,
ctx.add(&[
ctx.mul(&[ctx.num(p2), sin_bx]),
ctx.mul(&[ctx.num(q2), cos_bx]),
]),
]),
]),
);
let y1 = if alpha == 0 {
y1_inner
} else {
collect_terms(ctx, ctx.mul(&[e_ax, y1_inner]))
};
let y2 = if alpha == 0 {
y2_inner
} else {
collect_terms(ctx, ctx.mul(&[e_ax, y2_inner]))
};
ODESolution::System(ctx.slice(&[y1, y2]))
}
fn extract_coeff<'a>(
ctx: &'a AtomArena<'a>,
eq: Atom<'a>,
func: Atom<'a>,
var: Symbol,
) -> Option<i64> {
let func_str = func.to_string();
let dy_str = ctx
.fun("Derivative", &[func, ctx.var(var.as_str())])
.to_string();
let _ = dy_str;
let terms: Vec<Atom<'a>> = match eq.node() {
AtomNode::Add(args) => args.to_vec(),
_ => vec![eq],
};
let mut total: i64 = 0;
for term in terms {
let factors: Vec<Atom<'a>> = match term.node() {
AtomNode::Mul(args) => args.to_vec(),
_ => vec![term],
};
let has_func = factors.iter().any(|f| f.to_string() == func_str);
if has_func {
let coeff: i64 = factors
.iter()
.filter(|f| f.to_string() != func_str)
.map(|f| match f.node() {
AtomNode::Num(n) => *n,
_ => 0,
})
.product::<i64>()
.max(i64::MIN + 1); let coeff = if factors.len() == 1 { 1 } else { coeff };
total += -coeff;
}
}
Some(total)
}
fn eigenvector(a11: i64, a12: i64, a21: i64, a22: i64, lambda: i64) -> Option<(i64, i64)> {
let m11 = a11 - lambda;
let m22 = a22 - lambda;
if a12 != 0 {
return Some((a12, -m11));
}
if a21 != 0 {
return Some((-m22, a21));
}
None
}
fn generalized_eigenvector(
a11: i64,
a12: i64,
a21: i64,
a22: i64,
lambda: i64,
v1: i64,
v2: i64,
) -> Option<(i64, i64)> {
let m11 = a11 - lambda;
let m22 = a22 - lambda;
if a12 != 0 {
if v1 % a12 != 0 {
let w1 = 1;
let num = v1 - m11 * w1;
if num % a12 == 0 {
return Some((w1, num / a12));
}
return None;
}
return Some((0, v1 / a12));
}
if a21 != 0 {
if v2 % a21 == 0 {
return Some((v2 / a21, 0));
}
let w2 = 1;
let num = v2 - m22 * w2;
if num % a21 == 0 {
return Some((num / a21, w2));
}
return None;
}
None
}
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
}