use num_bigint::BigInt;
use num_rational::Ratio;
use super::RischResult;
use super::rde::{self, RdeResult};
use super::tower::{DifferentialExtension, ExtensionKind};
use super::tower_integrate::{tower_hermite_reduce, tower_logarithmic_part};
use crate::base::arena::Arena;
use crate::base::node::{ExprId, ExprNode};
use crate::poly::dense::Poly;
use crate::poly::generic::GenPoly;
use crate::poly::ratfn::RationalFn;
use crate::poly::traits::Ring;
pub fn risch_integrate(arena: &mut Arena, de: &mut DifferentialExtension) -> RischResult {
if de.is_base_level() {
let numer_poly = crate::poly::polybridge::expr_to_poly(arena, de.integrand, de.base_var);
if let Some(numer) = numer_poly {
return integrate_rational(&numer, &Poly::from_int(1));
}
let (n_id, d_id) = crate::poly::polybridge::as_numer_denom(arena, de.integrand);
let n_poly = crate::poly::polybridge::expr_to_poly(arena, n_id, de.base_var);
let d_poly = crate::poly::polybridge::expr_to_poly(arena, d_id, de.base_var);
if let (Some(n), Some(d)) = (n_poly, d_poly) {
return integrate_rational(&n, &d);
}
return RischResult::Failed("integrand is not a rational function at base level".into());
}
match de.current_kind().cloned() {
Some(ExtensionKind::Logarithmic) => integrate_primitive(arena, de),
Some(ExtensionKind::Exponential) => integrate_hyperexponential(arena, de),
None => RischResult::Failed("no extension kind at current level".into()),
}
}
fn integrate_rational(a: &Poly, d: &Poly) -> RischResult {
if d.is_zero() {
return RischResult::Failed("zero denominator".into());
}
if d.is_constant() {
let scaled = if let Some(lc) = d.leading_coeff() {
let inv = Ratio::new(lc.denom().clone(), lc.numer().clone());
a.scale(&inv)
} else {
a.clone()
};
let integral = integrate_poly(&scaled);
return RischResult::Elementary {
rational_numer: integral,
rational_denom: Poly::from_int(1),
log_terms: vec![],
arena_expr: None,
};
}
let hr = super::hermite::hermite_reduce(a, d);
let log_result = if hr.h_numer.is_zero() {
super::rothstein_trager::LogPartResult { terms: vec![] }
} else {
super::rothstein_trager::logarithmic_part(&hr.h_numer, &hr.h_denom)
};
RischResult::Elementary {
rational_numer: hr.g_numer,
rational_denom: hr.g_denom,
log_terms: log_result.terms,
arena_expr: None,
}
}
fn try_tower_rational_path(
arena: &mut Arena,
integrand: ExprId,
ext_var: ExprId,
base_var: ExprId,
) -> RischResult {
let (n_id, d_id) = crate::poly::polybridge::as_numer_denom(arena, integrand);
let n_gp = match arena_to_genpoly_ratfn(arena, n_id, ext_var, base_var) {
Some(gp) => gp,
None => {
return RischResult::Failed("Cannot convert numerator to GenPoly<RationalFn>".into());
}
};
let d_gp = match arena_to_genpoly_ratfn(arena, d_id, ext_var, base_var) {
Some(gp) => gp,
None => {
return RischResult::Failed("Cannot convert denominator to GenPoly<RationalFn>".into());
}
};
if d_gp.is_zero() {
return RischResult::Failed("zero denominator in tower rational path".into());
}
let hr = tower_hermite_reduce(&n_gp, &d_gp);
if !hr.h_numer.is_zero() {
let rt = tower_logarithmic_part(&hr.h_numer, &hr.h_denom);
if rt.is_non_elementary {
return RischResult::NonElementary;
}
}
RischResult::Failed(
"Tower HR+RT proved elementary but arena conversion not yet implemented".into(),
)
}
#[allow(dead_code)]
fn arena_to_genpoly_ratfn(
arena: &mut Arena,
expr: ExprId,
ext_var: ExprId,
base_var: ExprId,
) -> Option<GenPoly<RationalFn>> {
let terms = super::tower::extract_poly_in_ext_mut(arena, expr, ext_var)?;
let mut coeffs_map: std::collections::BTreeMap<usize, RationalFn> =
std::collections::BTreeMap::new();
for &(power, coeff_expr) in &terms {
let (cn, cd) = crate::poly::polybridge::as_numer_denom(arena, coeff_expr);
let cn_exp = crate::transforms::expand::expand(arena, cn);
let cn_eval = crate::transforms::eval::eval(arena, cn_exp);
let cd_exp = crate::transforms::expand::expand(arena, cd);
let cd_eval = crate::transforms::eval::eval(arena, cd_exp);
let cn_poly = crate::poly::polybridge::expr_to_poly(arena, cn_eval, base_var)?;
let cd_poly = crate::poly::polybridge::expr_to_poly(arena, cd_eval, base_var)?;
let rf = RationalFn::new(cn_poly, cd_poly);
coeffs_map.insert(power, rf);
}
let max_power = coeffs_map.keys().max().copied().unwrap_or(0);
let mut coeffs = Vec::with_capacity(max_power + 1);
for i in 0..=max_power {
coeffs.push(coeffs_map.remove(&i).unwrap_or_else(Ring::zero));
}
Some(GenPoly::from_coeffs(coeffs))
}
fn integrate_poly(p: &Poly) -> Poly {
if p.is_zero() {
return Poly::zero();
}
let coeffs = p.coeffs();
let mut result = vec![<Ratio<BigInt> as num_traits::Zero>::zero()]; for (k, c) in coeffs.iter().enumerate() {
let k_plus_1 = Ratio::from_integer(BigInt::from((k + 1) as i64));
result.push(c / &k_plus_1);
}
Poly::from_coeffs(result)
}
fn integrate_primitive(arena: &mut Arena, de: &mut DifferentialExtension) -> RischResult {
let level = match de.current_level_ext() {
Some(l) => l.clone(),
None => return RischResult::Failed("no current level".into()),
};
let ext_var = level.ext_var;
let base_var = de.base_var;
let poly_terms = match super::tower::extract_poly_in_ext_mut(arena, de.integrand, ext_var) {
Some(terms) => terms,
None => {
return try_tower_rational_path(arena, de.integrand, ext_var, base_var);
}
};
let max_power = poly_terms.iter().map(|&(p, _)| p).max().unwrap_or(0);
let mut coeff_map: std::collections::BTreeMap<usize, ExprId> =
std::collections::BTreeMap::new();
for &(power, coeff) in &poly_terms {
coeff_map.insert(power, coeff);
}
let d_theta = level.derivative;
let mut b_coeffs: std::collections::BTreeMap<usize, ExprId> = std::collections::BTreeMap::new();
for k in (0..=max_power).rev() {
let a_k = coeff_map.get(&k).copied().unwrap_or_else(|| arena.zero());
let rhs = if k < max_power {
if let Some(&b_next) = b_coeffs.get(&(k + 1)) {
let k_plus_1 = arena.int((k + 1) as i64);
let correction = arena.mul(&[k_plus_1, b_next, d_theta]);
let correction_eval = crate::transforms::eval::eval(arena, correction);
let diff = arena.sub(a_k, correction_eval);
crate::transforms::eval::eval(arena, diff)
} else {
a_k
}
} else {
a_k
};
if rhs == arena.zero() {
b_coeffs.insert(k, arena.zero());
continue;
}
let b_k = crate::transforms::integrate::integrate(arena, rhs, base_var);
if matches!(arena.node(b_k), ExprNode::Integral(_, _)) || b_k == rhs {
return RischResult::Failed(
"Recursive integration failed in logarithmic polynomial coefficient matching"
.into(),
);
}
b_coeffs.insert(k, b_k);
}
let ln_u = arena.ln(level.argument);
let mut result_terms: Vec<ExprId> = Vec::new();
for (&k, &b_k) in &b_coeffs {
if b_k == arena.zero() {
continue;
}
let term = if k == 0 {
b_k
} else if k == 1 {
let theta_sub = ln_u;
arena.mul(&[b_k, theta_sub])
} else {
let exp_k = arena.int(k as i64);
let theta_k = arena.pow(ln_u, exp_k);
arena.mul(&[b_k, theta_k])
};
result_terms.push(term);
}
let result_expr = if result_terms.is_empty() {
arena.zero()
} else if result_terms.len() == 1 {
result_terms[0]
} else {
arena.add(&result_terms)
};
let result_eval = crate::transforms::eval::eval(arena, result_expr);
RischResult::Elementary {
rational_numer: Poly::zero(),
rational_denom: Poly::from_int(1),
log_terms: vec![],
arena_expr: Some(result_eval),
}
}
fn integrate_hyperexponential(arena: &mut Arena, de: &mut DifferentialExtension) -> RischResult {
let level = match de.current_level_ext() {
Some(l) => l.clone(),
None => return RischResult::Failed("no current level".into()),
};
let ext_var = level.ext_var;
let base_var = de.base_var;
let poly_terms = match super::tower::extract_poly_in_ext_mut(arena, de.integrand, ext_var) {
Some(terms) => terms,
None => {
return try_tower_rational_path(arena, de.integrand, ext_var, base_var);
}
};
let max_power = poly_terms.iter().map(|&(p, _)| p).max().unwrap_or(0);
let mut coeff_map: std::collections::BTreeMap<usize, ExprId> =
std::collections::BTreeMap::new();
for &(power, coeff) in &poly_terms {
coeff_map.insert(power, coeff);
}
let du = crate::transforms::diff::diff(arena, level.argument, base_var);
let mut b_coeffs: std::collections::BTreeMap<usize, ExprId> = std::collections::BTreeMap::new();
for k in (0..=max_power).rev() {
let a_k = coeff_map.get(&k).copied().unwrap_or_else(|| arena.zero());
if a_k == arena.zero() {
b_coeffs.insert(k, arena.zero());
continue;
}
if k == 0 {
let b_0 = crate::transforms::integrate::integrate(arena, a_k, base_var);
if matches!(arena.node(b_0), ExprNode::Integral(_, _)) {
return RischResult::Failed(
"Recursive integration failed for k=0 coefficient in exponential case".into(),
);
}
b_coeffs.insert(0, b_0);
} else {
let k_expr = arena.int(k as i64);
let f_expr = arena.mul(&[k_expr, du]);
let f_eval = crate::transforms::eval::eval(arena, f_expr);
let (f_n_id, f_d_id) = crate::poly::polybridge::as_numer_denom(arena, f_eval);
let (g_n_id, g_d_id) = crate::poly::polybridge::as_numer_denom(arena, a_k);
let f_n = crate::poly::polybridge::expr_to_poly(arena, f_n_id, base_var);
let f_d = crate::poly::polybridge::expr_to_poly(arena, f_d_id, base_var);
let g_n = crate::poly::polybridge::expr_to_poly(arena, g_n_id, base_var);
let g_d = crate::poly::polybridge::expr_to_poly(arena, g_d_id, base_var);
match (f_n, f_d, g_n, g_d) {
(Some(fn_p), Some(fd_p), Some(gn_p), Some(gd_p)) => {
let rde_result = rde::solve_risch_de_rational(&fn_p, &fd_p, &gn_p, &gd_p);
match rde_result {
RdeResult::Solution { numer, denom } => {
let n_id =
crate::poly::polybridge::poly_to_expr(arena, &numer, base_var);
let d_id =
crate::poly::polybridge::poly_to_expr(arena, &denom, base_var);
let b_k = if d_id == arena.one() {
n_id
} else {
arena.div(n_id, d_id)
};
b_coeffs.insert(k, b_k);
}
RdeResult::NoSolution => {
tracing::info!(
"Risch: proved non-elementary — RDE B_{k}' + {k}·u'·B_{k} = a_{k} has no solution"
);
return RischResult::NonElementary;
}
RdeResult::NotImplemented(msg) => {
return RischResult::Failed(format!("RDE solver: {msg}"));
}
}
}
_ => {
return RischResult::Failed(format!(
"Cannot convert RDE coefficients to polynomials for k={k}"
));
}
}
}
}
let exp_u = arena.exp(level.argument);
let mut result_terms: Vec<ExprId> = Vec::new();
for (&k, &b_k) in &b_coeffs {
if b_k == arena.zero() {
continue;
}
let term = if k == 0 {
b_k
} else if k == 1 {
arena.mul(&[b_k, exp_u])
} else {
let exp_k = arena.int(k as i64);
let theta_k = arena.pow(exp_u, exp_k);
arena.mul(&[b_k, theta_k])
};
result_terms.push(term);
}
let result_expr = if result_terms.is_empty() {
arena.zero()
} else if result_terms.len() == 1 {
result_terms[0]
} else {
arena.add(&result_terms)
};
let result_eval = crate::transforms::eval::eval(arena, result_expr);
RischResult::Elementary {
rational_numer: Poly::zero(),
rational_denom: Poly::from_int(1),
log_terms: vec![],
arena_expr: Some(result_eval),
}
}
#[cfg(test)]
mod tests {
use super::*;
fn rat(n: i64, d: i64) -> Ratio<BigInt> {
Ratio::new(BigInt::from(n), BigInt::from(d))
}
#[test]
fn integrate_rational_polynomial() {
let a = Poly::from_coeffs(vec![rat(1, 1), rat(2, 1), rat(3, 1)]);
let d = Poly::from_int(1);
match integrate_rational(&a, &d) {
RischResult::Elementary {
rational_numer,
rational_denom,
log_terms,
..
} => {
assert!(
log_terms.is_empty(),
"polynomial integral should have no log terms"
);
assert_eq!(rational_denom.degree().unwrap_or(0), 0, "denom should be 1");
assert_eq!(rational_numer.degree(), Some(3));
}
other => panic!("expected Elementary, got {:?}", other),
}
}
#[test]
fn integrate_rational_one_over_x() {
let a = Poly::from_int(1);
let d = Poly::x();
match integrate_rational(&a, &d) {
RischResult::Elementary {
rational_numer,
log_terms,
..
} => {
assert!(
rational_numer.is_zero() || rational_numer.degree().unwrap_or(0) == 0,
"1/x should have zero rational part"
);
assert!(!log_terms.is_empty(), "should have log terms for 1/x");
}
other => panic!("expected Elementary, got {:?}", other),
}
}
#[test]
fn integrate_rational_one_over_x_squared() {
let a = Poly::from_int(1);
let d = Poly::from_coeffs(vec![rat(0, 1), rat(0, 1), rat(1, 1)]); match integrate_rational(&a, &d) {
RischResult::Elementary {
rational_numer,
log_terms,
..
} => {
assert!(log_terms.is_empty(), "1/x² should have no log terms");
assert!(
!rational_numer.is_zero(),
"1/x² should have a nonzero rational part"
);
}
other => panic!("expected Elementary, got {:?}", other),
}
}
#[test]
fn integrate_rational_partial_fractions() {
let a = Poly::from_int(1);
let d = Poly::from_coeffs(vec![rat(-1, 1), rat(0, 1), rat(1, 1)]); match integrate_rational(&a, &d) {
RischResult::Elementary { log_terms, .. } => {
assert!(!log_terms.is_empty(), "1/(x²-1) should have log terms");
}
other => panic!("expected Elementary, got {:?}", other),
}
}
#[test]
fn integrate_rational_hermite_plus_log() {
let a = Poly::from_coeffs(vec![rat(1, 1), rat(2, 1)]); let d = Poly::from_coeffs(vec![rat(1, 1), rat(2, 1), rat(1, 1)]);
match integrate_rational(&a, &d) {
RischResult::Elementary { .. } => {
}
other => panic!("expected Elementary, got {:?}", other),
}
}
#[test]
fn integrate_log_extension_theta() {
let mut arena = crate::base::arena::Arena::new();
let x = arena.symbol("x");
let theta = arena.symbol("__t0");
let one = arena.one();
let d_theta = arena.div(one, x);
let mut de = DifferentialExtension::new(x);
de.push_logarithmic(theta, x, d_theta);
de.integrand = theta;
let result = risch_integrate(&mut arena, &mut de);
match result {
RischResult::Elementary { .. } => {
}
RischResult::Failed(msg) => {
eprintln!("integrate_log(θ): {msg}");
}
other => panic!("unexpected: {:?}", other),
}
}
#[test]
fn integrate_exp_polynomial_theta() {
let mut arena = crate::base::arena::Arena::new();
let x = arena.symbol("x");
let theta = arena.symbol("__t0");
let mut de = DifferentialExtension::new(x);
de.push_exponential(theta, x, theta);
de.integrand = theta;
let result = risch_integrate(&mut arena, &mut de);
match result {
RischResult::Elementary { .. } => {
}
RischResult::Failed(msg) => {
eprintln!("integrate_exp(θ): {msg}");
}
other => panic!("unexpected: {:?}", other),
}
}
#[test]
fn integrate_exp_nonelementary_exp_neg_x_squared() {
let mut arena = crate::base::arena::Arena::new();
let x = arena.symbol("x");
let theta = arena.symbol("__t0");
let two = arena.int(2);
let x_sq = arena.pow(x, two);
let neg_x_sq = arena.neg(x_sq);
let neg_two = arena.int(-2);
let neg_2x = arena.mul(&[neg_two, x]);
let d_theta = arena.mul(&[neg_2x, theta]);
let mut de = DifferentialExtension::new(x);
de.push_exponential(theta, neg_x_sq, d_theta);
de.integrand = theta;
let result = risch_integrate(&mut arena, &mut de);
match result {
RischResult::NonElementary => {
}
RischResult::Failed(msg) => {
eprintln!("exp(-x²) returned Failed instead of NonElementary: {msg}");
}
RischResult::Elementary { .. } => {
panic!("∫ exp(-x²) dx should be NonElementary, got Elementary");
}
}
}
#[test]
fn integrate_exp_x_times_exp_x_squared() {
let mut arena = crate::base::arena::Arena::new();
let x = arena.symbol("x");
let theta = arena.symbol("__t0");
let two = arena.int(2);
let x_sq = arena.pow(x, two);
let two_x = arena.mul(&[two, x]);
let d_theta = arena.mul(&[two_x, theta]);
let integrand = arena.mul(&[x, theta]);
let mut de = DifferentialExtension::new(x);
de.push_exponential(theta, x_sq, d_theta);
de.integrand = integrand;
let result = risch_integrate(&mut arena, &mut de);
match result {
RischResult::Elementary { .. } => {
}
other => {
panic!("∫ x·exp(x²) dx should be Elementary, got {:?}", other);
}
}
}
#[test]
fn integrate_poly_basic() {
let p = Poly::from_coeffs(vec![rat(1, 1), rat(2, 1)]);
let result = integrate_poly(&p);
assert_eq!(result.coeff(0), rat(0, 1));
assert_eq!(result.coeff(1), rat(1, 1));
assert_eq!(result.coeff(2), rat(1, 1));
}
#[test]
fn integrate_poly_zero() {
let result = integrate_poly(&Poly::zero());
assert!(result.is_zero());
}
#[test]
fn integrate_poly_constant() {
let result = integrate_poly(&Poly::from_int(5));
assert_eq!(result.coeff(1), rat(5, 1));
}
}