use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use ocas_rewrite::rules::default_rules;
use ocas_rewrite::simplify::simplify;
use super::ODESolution;
use super::util::{contains_func, substitute_solution};
use crate::derivative::diff;
use crate::integral::integrate;
fn is_atom_zero(expr: Atom<'_>) -> bool {
matches!(expr.node(), AtomNode::Num(0))
}
pub(crate) fn solve_separable<'a>(
ctx: &'a AtomArena<'a>,
ode: super::ODE<'a>,
) -> Option<ODESolution<'a>> {
let super::ODE {
equation,
func,
var,
} = ode;
let func_name = match func.node() {
AtomNode::Fun(name, _) => *name,
_ => return None,
};
let y = ctx.var(func_name.as_str());
let _dy = ctx.fun("Derivative", &[func, ctx.var(var.as_str())]);
let (y_terms, x_terms) = separate_by_func(ctx, equation, func, var)?;
if is_atom_zero(x_terms) && is_atom_zero(y_terms) {
return None;
}
let x_integral = integrate(ctx, x_terms, var);
let y_in_x = substitute_var(ctx, y_terms, var, y);
let y_integral = integrate(ctx, y_in_x, func_name);
let implicit = ctx.add(&[y_integral, ctx.mul(&[ctx.num(-1), x_integral])]);
Some(ODESolution::Implicit(implicit))
}
pub(crate) fn solve_linear_first<'a>(
ctx: &'a AtomArena<'a>,
ode: super::ODE<'a>,
) -> Option<ODESolution<'a>> {
let super::ODE {
equation,
func,
var,
} = ode;
let _x = ctx.var(var.as_str());
let _y_sym = match func.node() {
AtomNode::Fun(name, _) => ctx.var(name.as_str()),
_ => return None,
};
let (p, q) = extract_linear_coeffs(ctx, equation, func, var)?;
let p_integral = integrate(ctx, p, var);
let mu = ctx.fun("exp", &[p_integral]);
let mu_q = ctx.mul(&[mu, q]);
let int_mu_q = integrate(ctx, mu_q, var);
let neg_p_integral = ctx.mul(&[ctx.num(-1), p_integral]);
let exp_neg_p = ctx.fun("exp", &[neg_p_integral]);
let particular = ctx.mul(&[exp_neg_p, int_mu_q]);
let c1 = ctx.var("C1");
let homogeneous = ctx.mul(&[c1, exp_neg_p]);
let solution = ctx.add(&[particular, homogeneous]);
Some(ODESolution::Explicit(solution))
}
pub(crate) fn solve_bernoulli<'a>(
ctx: &'a AtomArena<'a>,
ode: super::ODE<'a>,
) -> Option<ODESolution<'a>> {
let super::ODE {
equation,
func,
var,
} = ode;
let _x = ctx.var(var.as_str());
let y_name = match func.node() {
AtomNode::Fun(name, _) => name,
_ => return None,
};
let _y = ctx.var(y_name.as_str());
let n = find_bernoulli_power(equation, func, var)?;
if n == 0 || n == 1 {
return None;
}
let (p, q) = extract_linear_coeffs(ctx, equation, func, var)?;
let one_minus_n = 1 - n;
let new_p = ctx.mul(&[ctx.num(one_minus_n), p]);
let new_q = ctx.mul(&[ctx.num(one_minus_n), q]);
let new_p_integral = integrate(ctx, new_p, var);
let mu = ctx.fun("exp", &[new_p_integral]);
let mu_new_q = ctx.mul(&[mu, new_q]);
let int_mu_new_q = integrate(ctx, mu_new_q, var);
let neg_new_p_int = ctx.mul(&[ctx.num(-1), new_p_integral]);
let exp_neg = ctx.fun("exp", &[neg_new_p_int]);
let v_particular = ctx.mul(&[exp_neg, int_mu_new_q]);
let c1 = ctx.var("C1");
let v_homogeneous = ctx.mul(&[c1, exp_neg]);
let v_solution = ctx.add(&[v_particular, v_homogeneous]);
let power = if one_minus_n == 1 {
v_solution
} else {
let exponent = ctx.pow(ctx.num(one_minus_n), ctx.num(-1));
ctx.pow(v_solution, exponent)
};
Some(ODESolution::Explicit(power))
}
pub(crate) fn solve_exact<'a>(
ctx: &'a AtomArena<'a>,
ode: super::ODE<'a>,
) -> Option<ODESolution<'a>> {
let super::ODE {
equation,
func,
var,
} = ode;
let (m, n) = super::classify::split_mn(ctx, equation, func, var)?;
let y_name = match func.node() {
AtomNode::Fun(name, _) => *name,
_ => return None,
};
let y_var = ctx.var(y_name.as_str());
let m_sub = super::classify::replace_atom(ctx, m, func, y_var);
let n_sub = super::classify::replace_atom(ctx, n, func, y_var);
let (m_eff, n_eff) = if partials_equal(ctx, m_sub, n_sub, y_name, var) {
(m_sub, n_sub)
} else {
find_integrating_factor(ctx, m_sub, n_sub, y_name, var)?
};
let f_partial = integrate(ctx, m_eff, var);
let df_dy = diff(ctx, f_partial, y_name);
let correction =
super::util::collect_terms(ctx, ctx.add(&[n_eff, ctx.mul(&[ctx.num(-1), df_dy])]));
if !contains_x(correction, var) {
let g_y = integrate(ctx, correction, y_name);
let solution = ctx.add(&[f_partial, g_y]);
let solution = super::classify::replace_atom(ctx, solution, y_var, func);
return Some(ODESolution::Implicit(solution));
}
None
}
fn partials_equal<'a>(
ctx: &'a AtomArena<'a>,
m: Atom<'a>,
n: Atom<'a>,
y_sym: Symbol,
var: Symbol,
) -> bool {
let dm_dy = diff(ctx, m, y_sym);
let dn_dx = diff(ctx, n, var);
let dm_norm = ocas_atom::normalize::normalize(ctx, dm_dy);
let dn_norm = ocas_atom::normalize::normalize(ctx, dn_dx);
if dm_norm.to_string() == dn_norm.to_string() {
return true;
}
let difference =
super::util::collect_terms(ctx, ctx.add(&[dm_dy, ctx.mul(&[ctx.num(-1), dn_dx])]));
matches!(difference.node(), AtomNode::Num(0))
}
fn find_integrating_factor<'a>(
ctx: &'a AtomArena<'a>,
m: Atom<'a>,
n: Atom<'a>,
y_sym: Symbol,
var: Symbol,
) -> Option<(Atom<'a>, Atom<'a>)> {
let dm_dy = diff(ctx, m, y_sym);
let dn_dx = diff(ctx, n, var);
let rules = default_rules(ctx, &crate::pattern_alloc::VecAlloc);
let diff1 = simplify(
ctx,
ctx.add(&[dm_dy, ctx.mul(&[ctx.num(-1), dn_dx])]),
&rules,
20,
);
let ratio1 = ocas_atom::normalize::normalize(
ctx,
simplify(ctx, ctx.mul(&[diff1, ctx.pow(n, ctx.num(-1))]), &rules, 20),
);
if !ratio1.to_string().contains(y_sym.as_str()) && !contains_fun_named(ratio1, y_sym) {
let exponent = integrate(ctx, ratio1, var);
let mu = exp_simplify(ctx, exponent);
let m_new = ctx.mul(&[mu, m]);
let n_new = ctx.mul(&[mu, n]);
if partials_equal(ctx, m_new, n_new, y_sym, var) {
return Some((m_new, n_new));
}
}
let diff2 = simplify(
ctx,
ctx.add(&[dn_dx, ctx.mul(&[ctx.num(-1), dm_dy])]),
&rules,
20,
);
let ratio2 = ocas_atom::normalize::normalize(
ctx,
simplify(ctx, ctx.mul(&[diff2, ctx.pow(m, ctx.num(-1))]), &rules, 20),
);
if !contains_x(ratio2, var) {
let exponent = integrate(ctx, ratio2, y_sym);
let mu = exp_simplify(ctx, exponent);
let m_new = ctx.mul(&[mu, m]);
let n_new = ctx.mul(&[mu, n]);
if partials_equal(ctx, m_new, n_new, y_sym, var) {
return Some((m_new, n_new));
}
}
None
}
fn exp_simplify<'a>(ctx: &'a AtomArena<'a>, exponent: Atom<'a>) -> Atom<'a> {
super::util::exp_simplify(ctx, exponent)
}
fn contains_fun_named<'a>(expr: Atom<'a>, name: Symbol) -> bool {
match expr.node() {
AtomNode::Num(_) | AtomNode::Var(_) => false,
AtomNode::Add(args) | AtomNode::Mul(args) => {
args.iter().any(|a| contains_fun_named(*a, name))
}
AtomNode::Pow(base, exp) => {
contains_fun_named(*base, name) || contains_fun_named(*exp, name)
}
AtomNode::Fun(n, args) => *n == name || args.iter().any(|a| contains_fun_named(*a, name)),
}
}
pub(crate) fn solve_homogeneous<'a>(
ctx: &'a AtomArena<'a>,
ode: super::ODE<'a>,
) -> Option<ODESolution<'a>> {
let super::ODE {
equation,
func,
var,
} = ode;
let x = ctx.var(var.as_str());
let y_name = match func.node() {
AtomNode::Fun(name, _) => name,
_ => return None,
};
let y = ctx.var(y_name.as_str());
let v = ctx.var("v");
let y_replacement = ctx.mul(&[v, x]);
let substituted = substitute_solution(ctx, equation, func, y_replacement, var);
let _dv_symbol = ctx.var("dv");
let rules = default_rules(ctx, &crate::pattern_alloc::VecAlloc);
let simplified = simplify(ctx, substituted, &rules, 20);
let (v_part, x_part) = separate_by_var(ctx, simplified, var, v)?;
let x_integral = integrate(ctx, x_part, var);
let v_sym = Symbol::new("v");
let v_integral = integrate(ctx, v_part, v_sym);
let v_in_yx = ctx.mul(&[y, ctx.pow(x, ctx.num(-1))]);
let v_sol = substitute_var(ctx, v_integral, v_sym, v_in_yx);
let implicit = ctx.add(&[v_sol, ctx.mul(&[ctx.num(-1), x_integral])]);
Some(ODESolution::Implicit(implicit))
}
fn separate_by_func<'a>(
ctx: &'a AtomArena<'a>,
equation: Atom<'a>,
func: Atom<'a>,
var: Symbol,
) -> Option<(Atom<'a>, Atom<'a>)> {
match equation.node() {
AtomNode::Add(args) => {
let mut y_terms = Vec::new();
let mut x_terms = Vec::new();
for a in args.iter() {
if contains_func(*a, func, var) {
y_terms.push(*a);
} else {
x_terms.push(*a);
}
}
if y_terms.is_empty() || x_terms.is_empty() {
return None;
}
let y_sum = ctx.add(&y_terms);
let x_sum = ctx.add(&x_terms);
Some((y_sum, ctx.mul(&[ctx.num(-1), x_sum])))
}
_ => {
if contains_func(equation, func, var) {
Some((equation, ctx.num(0)))
} else {
None
}
}
}
}
fn extract_linear_coeffs<'a>(
ctx: &'a AtomArena<'a>,
equation: Atom<'a>,
func: Atom<'a>,
var: Symbol,
) -> Option<(Atom<'a>, Atom<'a>)> {
let x = ctx.var(var.as_str());
let dy = ctx.fun("Derivative", &[func, x]);
let terms = flatten_add(equation);
let mut p_coeff: Option<Atom<'a>> = None; let mut q_neg: Vec<Atom<'a>> = Vec::new();
for term in &terms {
let s = term.to_string();
if s == dy.to_string() {
continue;
}
if s == func.to_string() {
p_coeff = Some(ctx.num(1));
continue;
}
if contains_func(*term, func, var) && !is_derivative(*term, func, var) {
if let AtomNode::Mul(args) = term.node() {
let factors: Vec<_> = args
.iter()
.filter(|a| a.to_string() != func.to_string())
.copied()
.collect();
if factors.is_empty() {
p_coeff = Some(ctx.num(1));
} else if factors.len() == 1 {
p_coeff = Some(factors[0]);
} else {
p_coeff = Some(ctx.mul(&factors));
}
} else if term.to_string() == func.to_string() {
p_coeff = Some(ctx.num(1));
}
} else if !contains_func(*term, func, var) {
q_neg.push(*term);
}
}
let p = p_coeff.unwrap_or(ctx.num(0));
let q = if q_neg.is_empty() {
ctx.num(0)
} else if q_neg.len() == 1 {
ctx.mul(&[ctx.num(-1), q_neg[0]])
} else {
ctx.mul(&[ctx.num(-1), ctx.add(&q_neg)])
};
Some((p, q))
}
fn find_bernoulli_power<'a>(equation: Atom<'a>, func: Atom<'a>, var: Symbol) -> Option<i64> {
find_power_inner(equation, func, var)
}
fn find_power_inner<'a>(expr: Atom<'a>, func: Atom<'a>, var: Symbol) -> Option<i64> {
match expr.node() {
AtomNode::Add(args) => args.iter().find_map(|a| find_power_inner(*a, func, var)),
AtomNode::Mul(args) => args.iter().find_map(|a| find_power_inner(*a, func, var)),
AtomNode::Pow(base, exp) => {
if contains_func(*base, func, var)
&& let AtomNode::Num(n) = exp.node()
&& *n >= 2
{
return Some(*n);
}
find_power_inner(*base, func, var).or_else(|| find_power_inner(*exp, func, var))
}
_ => None,
}
}
fn separate_by_var<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
_x_var: Symbol,
dep_var: Atom<'a>,
) -> Option<(Atom<'a>, Atom<'a>)> {
match expr.node() {
AtomNode::Add(args) => {
let mut dep_terms = Vec::new();
let mut free_terms = Vec::new();
for a in args.iter() {
if contains_var_atom(*a, dep_var) {
dep_terms.push(*a);
} else {
free_terms.push(*a);
}
}
if dep_terms.is_empty() || free_terms.is_empty() {
return None;
}
Some((ctx.add(&dep_terms), ctx.add(&free_terms)))
}
_ => None,
}
}
fn flatten_add<'a>(expr: Atom<'a>) -> Vec<Atom<'a>> {
match expr.node() {
AtomNode::Add(args) => args.to_vec(),
_ => vec![expr],
}
}
fn is_derivative<'a>(expr: Atom<'a>, func: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Fun(name, args) => {
*name == Symbol::new("Derivative")
&& args.len() >= 2
&& args[0].to_string() == func.to_string()
&& args[1].to_string() == var.as_str()
}
_ => false,
}
}
fn contains_x<'a>(expr: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Num(_) => false,
AtomNode::Var(v) => *v == var,
AtomNode::Add(args) | AtomNode::Mul(args) => args.iter().any(|a| contains_x(*a, var)),
AtomNode::Pow(base, exp) => contains_x(*base, var) || contains_x(*exp, var),
AtomNode::Fun(_, args) => args.iter().any(|a| contains_x(*a, var)),
}
}
fn contains_var_atom<'a>(expr: Atom<'a>, target: Atom<'a>) -> bool {
let target_str = target.to_string();
contains_str(expr, &target_str)
}
fn contains_str<'a>(expr: Atom<'a>, target: &str) -> bool {
if expr.to_string() == target {
return true;
}
match expr.node() {
AtomNode::Num(_) | AtomNode::Var(_) => false,
AtomNode::Add(args) | AtomNode::Mul(args) => args.iter().any(|a| contains_str(*a, target)),
AtomNode::Pow(base, exp) => contains_str(*base, target) || contains_str(*exp, target),
AtomNode::Fun(_, args) => args.iter().any(|a| contains_str(*a, target)),
}
}
fn substitute_var<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
var_sym: Symbol,
replacement: Atom<'a>,
) -> Atom<'a> {
crate::series::substitute(ctx, expr, var_sym, replacement)
}