use ocas_atom::{Atom, AtomArena, AtomNode, Symbol};
use super::ODE;
use super::util::{contains_func, is_linear_in, ode_order};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ODEType {
Separable,
LinearFirst,
Bernoulli,
Exact,
Homogeneous,
LinearConstantCoeff,
CauchyEuler,
ReductionOfOrder,
PowerSeries,
}
pub fn classify_ode<'a>(ctx: &'a AtomArena<'a>, ode: ODE<'a>) -> Vec<ODEType> {
let ODE {
equation,
func,
var,
} = ode;
let order = ode_order(equation, func, var);
let mut types = Vec::new();
if order == 0 {
return types;
}
if order == 1 {
let linear = is_first_order_linear(ctx, equation, func, var);
if linear {
types.push(ODEType::LinearFirst);
}
if is_bernoulli(ctx, equation, func, var) {
types.push(ODEType::Bernoulli);
}
if is_separable(ctx, equation, func, var) {
types.push(ODEType::Separable);
}
if is_exact(ctx, equation, func, var) {
types.push(ODEType::Exact);
}
if is_homogeneous(ctx, equation, func, var) {
types.push(ODEType::Homogeneous);
}
} else {
if is_linear_in(equation, func, var) {
if is_constant_coeff_linear(ctx, equation, func, var) {
types.push(ODEType::LinearConstantCoeff);
}
if is_cauchy_euler(ctx, equation, func, var) {
types.push(ODEType::CauchyEuler);
}
if types.is_empty() {
types.push(ODEType::LinearConstantCoeff);
}
if ode_order(equation, func, var) == 2 {
types.push(ODEType::ReductionOfOrder);
}
}
}
if is_linear_in(equation, func, var) && ode_order(equation, func, var) >= 1 {
types.push(ODEType::PowerSeries);
}
types
}
fn is_first_order_linear<'a>(
_ctx: &'a AtomArena<'a>,
equation: Atom<'a>,
func: Atom<'a>,
var: Symbol,
) -> bool {
if !is_linear_in(equation, func, var) {
return false;
}
if !contains_func(equation, func, var) {
return false;
}
ode_order(equation, func, var) == 1
}
fn is_bernoulli<'a>(
_ctx: &'a AtomArena<'a>,
equation: Atom<'a>,
func: Atom<'a>,
var: Symbol,
) -> bool {
if ode_order(equation, func, var) != 1 {
return false;
}
has_nonlinear_func_term(equation, func, var)
}
fn has_nonlinear_func_term<'a>(expr: Atom<'a>, func: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Num(_) | AtomNode::Var(_) => false,
AtomNode::Add(args) => args.iter().any(|a| has_nonlinear_func_term(*a, func, var)),
AtomNode::Mul(args) => args.iter().any(|a| has_nonlinear_func_term(*a, func, var)),
AtomNode::Pow(base, exp) => {
if contains_func(*base, func, var) {
if let AtomNode::Num(n) = exp.node() {
*n >= 2
} else {
true }
} else {
has_nonlinear_func_term(*exp, func, var)
}
}
AtomNode::Fun(_, args) => args.iter().any(|a| has_nonlinear_func_term(*a, func, var)),
}
}
fn is_separable<'a>(
_ctx: &'a AtomArena<'a>,
equation: Atom<'a>,
func: Atom<'a>,
var: Symbol,
) -> bool {
if ode_order(equation, func, var) != 1 {
return false;
}
let _dy = derivative_atom(func, var);
let has_free_terms = has_func_free_additive_terms(equation, func, var);
let has_func_terms = contains_func(equation, func, var);
has_free_terms && has_func_terms && ode_order(equation, func, var) == 1
}
fn is_exact<'a>(ctx: &'a AtomArena<'a>, equation: Atom<'a>, func: Atom<'a>, var: Symbol) -> bool {
use ocas_rewrite::rules::default_rules;
use ocas_rewrite::simplify::simplify;
if ode_order(equation, func, var) != 1 {
return false;
}
let Some(y_sym) = func_symbol(func) else {
return false;
};
let Some((m, n)) = split_mn(ctx, equation, func, var) else {
return false;
};
let y_var = ctx.var(y_sym.as_str());
let m_sub = replace_atom(ctx, m, func, y_var);
let n_sub = replace_atom(ctx, n, func, y_var);
let dm_dy = crate::derivative::diff(ctx, m_sub, y_sym);
let dn_dx = crate::derivative::diff(ctx, n_sub, 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 rules = default_rules(ctx, &crate::pattern_alloc::VecAlloc);
let difference = simplify(
ctx,
ctx.add(&[dm_dy, ctx.mul(&[ctx.num(-1), dn_dx])]),
&rules,
20,
);
let difference = ocas_atom::normalize::normalize(ctx, difference);
matches!(difference.node(), AtomNode::Num(0))
}
pub(crate) fn split_mn<'a>(
ctx: &'a AtomArena<'a>,
equation: Atom<'a>,
func: Atom<'a>,
var: Symbol,
) -> Option<(Atom<'a>, Atom<'a>)> {
let dy_str = derivative_atom(func, var);
let mut m_terms = Vec::new();
let mut n_terms = Vec::new();
let terms: Vec<Atom<'a>> = match equation.node() {
AtomNode::Add(args) => args.to_vec(),
_ => vec![equation],
};
for term in terms {
let s = term.to_string();
if s == dy_str {
n_terms.push(ctx.num(1));
continue;
}
match term.node() {
AtomNode::Mul(args) => {
let dy_count = args.iter().filter(|a| a.to_string() == dy_str).count();
if dy_count == 1 {
let rest: Vec<_> = args
.iter()
.filter(|a| a.to_string() != dy_str)
.copied()
.collect();
if rest.iter().any(|a| contains_derivative_str(*a, &dy_str)) {
return None;
}
n_terms.push(if rest.is_empty() {
ctx.num(1)
} else {
ctx.mul(&rest)
});
} else if dy_count == 0 {
if contains_derivative_str(term, &dy_str) {
return None;
}
m_terms.push(term);
} else {
return None;
}
}
_ => {
if contains_derivative_str(term, &dy_str) {
return None;
}
m_terms.push(term);
}
}
}
if m_terms.is_empty() || n_terms.is_empty() {
return None;
}
let m = if m_terms.len() == 1 {
m_terms[0]
} else {
ctx.add(&m_terms)
};
let n = if n_terms.len() == 1 {
n_terms[0]
} else {
ctx.add(&n_terms)
};
Some((m, n))
}
pub(crate) fn contains_derivative_str<'a>(expr: Atom<'a>, dy_str: &str) -> bool {
if expr.to_string() == dy_str {
return true;
}
match expr.node() {
AtomNode::Num(_) | AtomNode::Var(_) => false,
AtomNode::Add(args) | AtomNode::Mul(args) => {
args.iter().any(|a| contains_derivative_str(*a, dy_str))
}
AtomNode::Pow(base, exp) => {
contains_derivative_str(*base, dy_str) || contains_derivative_str(*exp, dy_str)
}
AtomNode::Fun(_, args) => args.iter().any(|a| contains_derivative_str(*a, dy_str)),
}
}
pub(crate) fn replace_atom<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
target: Atom<'a>,
replacement: Atom<'a>,
) -> Atom<'a> {
if expr.to_string() == target.to_string() {
return replacement;
}
match expr.node() {
AtomNode::Num(_) | AtomNode::Var(_) => expr,
AtomNode::Add(args) => {
let mapped: Vec<_> = args
.iter()
.map(|a| replace_atom(ctx, *a, target, replacement))
.collect();
ctx.add(&mapped)
}
AtomNode::Mul(args) => {
let mapped: Vec<_> = args
.iter()
.map(|a| replace_atom(ctx, *a, target, replacement))
.collect();
ctx.mul(&mapped)
}
AtomNode::Pow(base, exp) => {
let b = replace_atom(ctx, *base, target, replacement);
let e = replace_atom(ctx, *exp, target, replacement);
ctx.pow(b, e)
}
AtomNode::Fun(name, args) => {
let mapped: Vec<_> = args
.iter()
.map(|a| replace_atom(ctx, *a, target, replacement))
.collect();
ctx.fun(name.as_str(), &mapped)
}
}
}
fn is_homogeneous<'a>(
_ctx: &'a AtomArena<'a>,
equation: Atom<'a>,
func: Atom<'a>,
var: Symbol,
) -> bool {
if ode_order(equation, func, var) != 1 {
return false;
}
if !is_linear_in(equation, func, var) {
return false;
}
let has_free_terms = has_func_free_additive_terms(equation, func, var);
all_terms_homogeneous_degree(equation, func, var) && has_free_terms
}
fn all_terms_homogeneous_degree<'a>(expr: Atom<'a>, func: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Add(args) => {
let degrees: Vec<i64> = args.iter().map(|a| term_degree(*a, func, var)).collect();
if degrees.is_empty() {
return true;
}
let first = degrees[0];
first > 0 && degrees.iter().all(|&d| d == first)
}
_ => {
let d = term_degree(expr, func, var);
d > 0
}
}
}
fn term_degree<'a>(expr: Atom<'a>, func: Atom<'a>, var: Symbol) -> i64 {
match expr.node() {
AtomNode::Num(_) => 0,
AtomNode::Var(v) => {
if *v == var {
1
} else {
0
}
}
AtomNode::Mul(args) => args.iter().map(|a| term_degree(*a, func, var)).sum(),
AtomNode::Pow(base, exp) => {
if let AtomNode::Num(n) = exp.node() {
term_degree(*base, func, var) * *n
} else {
1 }
}
AtomNode::Fun(name, args) => {
if *name == Symbol::new("Derivative") && args.len() >= 2 {
if args[0].to_string() == func.to_string() {
1
} else {
0
}
} else {
if expr.to_string() == func.to_string() {
1
} else {
0
}
}
}
AtomNode::Add(_) => 1, }
}
fn has_func_free_additive_terms<'a>(expr: Atom<'a>, func: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Add(args) => args.iter().any(|a| !contains_func(*a, func, var)),
_ => !contains_func(expr, func, var),
}
}
fn is_constant_coeff_linear<'a>(
_ctx: &'a AtomArena<'a>,
equation: Atom<'a>,
func: Atom<'a>,
var: Symbol,
) -> bool {
if !is_linear_in(equation, func, var) {
return false;
}
coefficients_free_of_var(equation, func, var)
}
fn coefficients_free_of_var<'a>(expr: Atom<'a>, func: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Add(args) => args.iter().all(|a| coefficient_free_of_var(*a, func, var)),
AtomNode::Mul(args) => {
let func_factor_count = args
.iter()
.filter(|a| is_func_or_derivative(**a, func, var))
.count();
if func_factor_count == 1 {
args.iter().all(|a| {
if is_func_or_derivative(*a, func, var) {
true
} else {
!contains_var(*a, var)
}
})
} else if func_factor_count == 0 {
true
} else {
false
}
}
AtomNode::Fun(name, args) if *name == Symbol::new("Derivative") && args.len() >= 2 => {
args[0].to_string() == func.to_string()
}
_ => true,
}
}
fn coefficient_free_of_var<'a>(expr: Atom<'a>, func: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Mul(args) => {
let func_factor = args.iter().find(|a| is_func_or_derivative(**a, func, var));
if func_factor.is_some() {
args.iter().all(|a| {
if is_func_or_derivative(*a, func, var) {
true
} else {
!contains_var(*a, var)
}
})
} else {
!contains_var(expr, var)
}
}
AtomNode::Fun(name, args) => {
if *name == Symbol::new("Derivative") && args.len() >= 2 {
args[0].to_string() == func.to_string()
} else if expr.to_string() == func.to_string() {
true
} else {
!contains_var(expr, var)
}
}
_ => !contains_var(expr, var),
}
}
fn is_cauchy_euler<'a>(
_ctx: &'a AtomArena<'a>,
equation: Atom<'a>,
func: Atom<'a>,
var: Symbol,
) -> bool {
if !is_linear_in(equation, func, var) {
return false;
}
let order = ode_order(equation, func, var);
if order < 2 {
return false;
}
is_cauchy_euler_pattern(equation, func, var)
}
fn is_cauchy_euler_pattern<'a>(expr: Atom<'a>, func: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Add(args) => args.iter().all(|a| is_ce_term(*a, func, var)),
_ => is_ce_term(expr, func, var),
}
}
fn is_ce_term<'a>(expr: Atom<'a>, func: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Mul(args) => {
let func_factor = args.iter().find(|a| is_func_or_derivative(**a, func, var));
if let Some(ff) = func_factor {
let order = func_derivative_order(*ff, func, var);
let coeff_factors: Vec<_> = args
.iter()
.filter(|a| !is_func_or_derivative(**a, func, var))
.copied()
.collect();
match coeff_factors.len() {
0 => order == 0, 1 => {
if order == 0 {
!contains_var(coeff_factors[0], var)
} else {
is_x_power(coeff_factors[0], var, order as i64)
}
}
_ => {
let has_x_power = coeff_factors
.iter()
.any(|c| is_x_power(*c, var, order as i64));
let rest_const = coeff_factors
.iter()
.filter(|c| !is_x_power(**c, var, order as i64))
.all(|c| !contains_var(*c, var));
has_x_power && rest_const
}
}
} else {
!contains_var(expr, var)
}
}
AtomNode::Fun(name, args) => {
if *name == Symbol::new("Derivative") && args.len() >= 2 {
args[0].to_string() == func.to_string()
} else {
expr.to_string() == func.to_string()
}
}
_ => !contains_var(expr, var),
}
}
fn is_x_power<'a>(expr: Atom<'a>, var: Symbol, n: i64) -> bool {
match expr.node() {
AtomNode::Var(v) => *v == var && n == 1,
AtomNode::Pow(base, exp) => {
if let AtomNode::Num(e) = exp.node() {
if let AtomNode::Var(v) = base.node() {
*v == var && *e == n
} else {
false
}
} else {
false
}
}
AtomNode::Num(1) => n == 0,
_ => false,
}
}
fn is_func_or_derivative<'a>(expr: Atom<'a>, func: Atom<'a>, var: Symbol) -> bool {
match expr.node() {
AtomNode::Fun(name, args) => {
if *name == Symbol::new("Derivative") && args.len() >= 2 {
args[0].to_string() == func.to_string() && args[1].to_string() == var.as_str()
} else {
expr.to_string() == func.to_string()
}
}
_ => expr.to_string() == func.to_string(),
}
}
fn func_derivative_order<'a>(expr: Atom<'a>, _func: Atom<'a>, _var: Symbol) -> usize {
match expr.node() {
AtomNode::Fun(name, args) if *name == Symbol::new("Derivative") && args.len() >= 2 => {
args.len() - 1
}
_ => 0,
}
}
fn contains_var<'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_var(*a, var)),
AtomNode::Pow(base, exp) => contains_var(*base, var) || contains_var(*exp, var),
AtomNode::Fun(_, args) => args.iter().any(|a| contains_var(*a, var)),
}
}
fn derivative_atom<'a>(func: Atom<'a>, var: Symbol) -> String {
format!("Derivative({}, {})", func, var.as_str())
}
fn func_symbol<'a>(func: Atom<'a>) -> Option<Symbol> {
match func.node() {
AtomNode::Fun(name, _) => Some(*name),
AtomNode::Var(v) => Some(*v),
_ => None,
}
}