use crate::deriv::{DerivationLog, DerivedExpr, RewriteStep};
use crate::flint::mpoly::FlintMPolyCtx;
use crate::flint::FlintPoly;
use crate::kernel::{ExprData, ExprId, ExprPool};
use crate::poly::error::ConversionError;
use crate::poly::multipoly::multi_to_flint_pub;
use crate::poly::multipoly::MultiPoly;
use crate::poly::unipoly::UniPoly;
use std::collections::{BTreeMap, BTreeSet};
use std::fmt;
use std::sync::Arc;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ResultantError {
NotAPolynomial(ConversionError),
FlintError,
}
impl From<ConversionError> for ResultantError {
fn from(e: ConversionError) -> Self {
ResultantError::NotAPolynomial(e)
}
}
impl fmt::Display for ResultantError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
ResultantError::NotAPolynomial(e) => write!(f, "not a polynomial: {e}"),
ResultantError::FlintError => {
write!(f, "FLINT resultant computation failed (E-RES-003)")
}
}
}
}
impl std::error::Error for ResultantError {}
impl crate::errors::AlkahestError for ResultantError {
fn code(&self) -> &'static str {
match self {
ResultantError::NotAPolynomial(_) => "E-RES-001",
ResultantError::FlintError => "E-RES-003",
}
}
fn remediation(&self) -> Option<&'static str> {
match self {
ResultantError::NotAPolynomial(_) => Some(
"ensure both arguments are polynomial expressions with integer \
coefficients in the given variable",
),
ResultantError::FlintError => None,
}
}
}
pub fn collect_free_vars(expr: ExprId, pool: &ExprPool) -> Vec<ExprId> {
let mut set = BTreeSet::new();
collect_vars_rec(expr, pool, &mut set);
set.into_iter().collect()
}
fn collect_vars_rec(expr: ExprId, pool: &ExprPool, out: &mut BTreeSet<ExprId>) {
let children: Vec<ExprId> = pool.with(expr, |data| match data {
ExprData::Symbol { .. } => {
out.insert(expr);
vec![]
}
ExprData::Integer(_) | ExprData::Rational(_) | ExprData::Float(_) => vec![],
ExprData::Add(args) | ExprData::Mul(args) => args.clone(),
ExprData::Pow { base, exp } => vec![*base, *exp],
ExprData::Func { args, .. } => args.clone(),
ExprData::Piecewise { branches, default } => {
let mut ids: Vec<ExprId> = branches.iter().flat_map(|(c, v)| [*c, *v]).collect();
ids.push(*default);
ids
}
ExprData::Predicate { args, .. } => args.clone(),
ExprData::Forall { var, body } | ExprData::Exists { var, body } => vec![*var, *body],
ExprData::BigO(arg) => vec![*arg],
ExprData::RootSum { poly, var, body } => vec![*poly, *var, *body],
});
for child in children {
collect_vars_rec(child, pool, out);
}
}
pub fn resultant(
p: ExprId,
q: ExprId,
var: ExprId,
pool: &ExprPool,
) -> Result<DerivedExpr<ExprId>, ResultantError> {
let mut all: BTreeSet<ExprId> = BTreeSet::new();
for v in collect_free_vars(p, pool) {
all.insert(v);
}
for v in collect_free_vars(q, pool) {
all.insert(v);
}
all.insert(var);
let vars: Vec<ExprId> = all.into_iter().collect();
let nvars = vars.len();
let var_idx = vars.iter().position(|&v| v == var).unwrap();
let mp = MultiPoly::from_symbolic(p, vars.clone(), pool)?;
let mq = MultiPoly::from_symbolic(q, vars.clone(), pool)?;
let ctx = FlintMPolyCtx::new(nvars.max(1));
let fp = multi_to_flint_pub(&mp, Arc::clone(&ctx));
let fq = multi_to_flint_pub(&mq, Arc::clone(&ctx));
let fr = fp
.resultant(&fq, var_idx)
.ok_or(ResultantError::FlintError)?;
let res_raw = fr.terms();
let remaining_vars: Vec<ExprId> = vars
.iter()
.enumerate()
.filter_map(|(i, &v)| if i == var_idx { None } else { Some(v) })
.collect();
let mut new_terms: BTreeMap<Vec<u32>, rug::Integer> = BTreeMap::new();
for (exp, coeff) in res_raw {
let mut new_exp: Vec<u32> = exp
.into_iter()
.enumerate()
.filter_map(|(i, e)| if i == var_idx { None } else { Some(e) })
.collect();
while new_exp.last() == Some(&0) {
new_exp.pop();
}
let entry = new_terms
.entry(new_exp)
.or_insert_with(|| rug::Integer::from(0));
*entry += &coeff;
}
new_terms.retain(|_, v| *v != 0);
let result_mp = MultiPoly {
vars: remaining_vars,
terms: new_terms,
};
let result_expr = result_mp.to_expr(pool);
let step = RewriteStep::simple("Resultant", p, result_expr);
Ok(DerivedExpr::with_step(result_expr, step))
}
pub fn subresultant_prs(
p: ExprId,
q: ExprId,
var: ExprId,
pool: &ExprPool,
) -> Result<DerivedExpr<Vec<ExprId>>, ResultantError> {
let mut up = UniPoly::from_symbolic(p, var, pool)?;
let mut uq = UniPoly::from_symbolic(q, var, pool)?;
if up.degree() < uq.degree() {
std::mem::swap(&mut up, &mut uq);
}
let prs_polys = sprs_inner(up, uq).ok_or(ResultantError::FlintError)?;
let exprs: Vec<ExprId> = prs_polys
.into_iter()
.map(|poly| poly.to_symbolic_expr(pool))
.collect();
let mut log = DerivationLog::new();
if let (Some(&first), Some(&last)) = (exprs.first(), exprs.last()) {
log.push(RewriteStep::simple("SubresultantPRS", first, last));
}
Ok(DerivedExpr::with_log(exprs, log))
}
type Coeffs = Vec<rug::Integer>;
fn trim(c: &mut Coeffs) {
while c.last().is_some_and(|t| *t == 0) {
c.pop();
}
}
fn rug_pow(base: &rug::Integer, exp: u32) -> rug::Integer {
if exp == 0 {
return rug::Integer::from(1);
}
let mut r = base.clone();
for _ in 1..exp {
r *= base;
}
r
}
fn scalar_mul(a: &Coeffs, c: &rug::Integer) -> Coeffs {
if *c == 0 {
return Coeffs::new();
}
a.iter().map(|t| rug::Integer::from(t * c)).collect()
}
fn scalar_div_exact(a: &Coeffs, c: &rug::Integer) -> Option<Coeffs> {
if *c == 0 {
return None;
}
let mut out = Coeffs::with_capacity(a.len());
for t in a {
if !t.is_divisible(c) {
return None;
}
out.push(rug::Integer::from(t / c));
}
trim(&mut out);
Some(out)
}
fn pseudo_remainder(a: &Coeffs, b: &Coeffs) -> Option<Coeffs> {
let db = b.len().checked_sub(1)?;
let lc_b = &b[db];
if a.len() <= db {
return Some(a.clone());
}
let delta = (a.len() - 1) - db;
let mut r = scalar_mul(a, &rug_pow(lc_b, delta as u32 + 1));
while r.len() > db {
let dr = r.len() - 1;
if !r[dr].is_divisible(lc_b) {
return None;
}
let quot = rug::Integer::from(&r[dr] / lc_b);
let shift = dr - db;
for (i, bi) in b.iter().enumerate() {
r[shift + i] -= rug::Integer::from(" * bi);
}
trim(&mut r);
if r.is_empty() {
break;
}
}
Some(r)
}
fn sprs_inner(p: UniPoly, q: UniPoly) -> Option<Vec<UniPoly>> {
let var = p.var;
let mut sequence = vec![p.clone(), q.clone()];
let mut pc = p.coefficients();
let mut qc = q.coefficients();
trim(&mut pc);
trim(&mut qc);
if qc.len() <= 1 || pc.is_empty() {
return Some(sequence);
}
let mut s = rug_pow(&qc[qc.len() - 1], (pc.len() - qc.len()) as u32);
let mut a = qc.clone();
let neg_q: Coeffs = qc.iter().map(|t| rug::Integer::from(-t)).collect();
let mut b = pseudo_remainder(&pc, &neg_q)?;
while !b.is_empty() {
let d = a.len() - 1;
let e = b.len() - 1;
let delta = d - e;
let c = if delta > 1 {
let scaled = scalar_mul(&b, &rug_pow(&b[e], delta as u32 - 1));
scalar_div_exact(&scaled, &rug_pow(&s, delta as u32 - 1))?
} else {
b.clone()
};
sequence.push(UniPoly {
var,
coeffs: FlintPoly::from_rug_coefficients(&c),
});
if e == 0 {
break;
}
let neg_b: Coeffs = b.iter().map(|t| rug::Integer::from(-t)).collect();
let rem = pseudo_remainder(&a, &neg_b)?;
let divisor = rug_pow(&s, delta as u32) * &a[d];
b = scalar_div_exact(&rem, &divisor)?;
a = c;
s = a[a.len() - 1].clone();
}
Some(sequence)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernel::{Domain, ExprPool};
fn pool_xy() -> (ExprPool, ExprId, ExprId) {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let y = p.symbol("y", Domain::Real);
(p, x, y)
}
#[test]
fn free_vars_constant() {
let p = ExprPool::new();
let five = p.integer(5_i32);
let vars = collect_free_vars(five, &p);
assert!(vars.is_empty());
}
#[test]
fn free_vars_symbol() {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let vars = collect_free_vars(x, &p);
assert_eq!(vars, vec![x]);
}
#[test]
fn free_vars_polynomial() {
let (p, x, y) = pool_xy();
let xsq = p.pow(x, p.integer(2_i32));
let expr = p.add(vec![xsq, y, p.integer(-1_i32)]);
let vars = collect_free_vars(expr, &p);
assert_eq!(vars.len(), 2);
assert!(vars.contains(&x));
assert!(vars.contains(&y));
}
#[test]
fn resultant_common_root() {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let xsq = p.pow(x, p.integer(2_i32));
let five_x = p.mul(vec![p.integer(-5_i32), x]);
let poly_p = p.add(vec![xsq, five_x, p.integer(6_i32)]);
let poly_q = p.add(vec![x, p.integer(-2_i32)]);
let dr = resultant(poly_p, poly_q, x, &p).unwrap();
match p.get(dr.value) {
ExprData::Integer(n) => assert_eq!(n.0, 0),
_ => panic!("expected integer 0, got {:?}", p.get(dr.value)),
}
assert_eq!(dr.log.len(), 1);
assert_eq!(dr.log.steps()[0].rule_name, "Resultant");
}
#[test]
fn resultant_coprime() {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let xsq = p.pow(x, p.integer(2_i32));
let poly_p = p.add(vec![xsq, p.integer(1_i32)]);
let poly_q = p.add(vec![x, p.integer(-1_i32)]);
let dr = resultant(poly_p, poly_q, x, &p).unwrap();
match p.get(dr.value) {
ExprData::Integer(n) => assert_eq!(n.0, 2),
_ => panic!("expected integer 2, got {:?}", p.get(dr.value)),
}
}
#[test]
fn resultant_linear_linear() {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let poly_p = p.add(vec![x, p.integer(-3_i32)]);
let poly_q = p.add(vec![x, p.integer(-7_i32)]);
let dr = resultant(poly_p, poly_q, x, &p).unwrap();
match p.get(dr.value) {
ExprData::Integer(n) => {
assert_eq!(
n.0.clone().abs(),
rug::Integer::from(4),
"magnitude should be 4"
);
}
_ => panic!("expected integer, got {:?}", p.get(dr.value)),
}
}
#[test]
fn resultant_bivariate_eliminates_var() {
let (p, x, y) = pool_xy();
let xsq = p.pow(x, p.integer(2_i32));
let ysq = p.pow(y, p.integer(2_i32));
let circle = p.add(vec![xsq, ysq, p.integer(-1_i32)]);
let line = p.add(vec![y, p.mul(vec![p.integer(-1_i32), x])]);
let dr = resultant(circle, line, y, &p).unwrap();
let res_expr = dr.value;
let res_poly = UniPoly::from_symbolic(res_expr, x, &p).unwrap();
assert_eq!(res_poly.degree(), 2, "expected degree-2 resultant in x");
let coeffs = res_poly.coefficients_i64();
assert_eq!(coeffs[0], -1, "constant term should be -1");
assert_eq!(coeffs[2], 2, "leading coefficient should be 2");
}
#[test]
fn resultant_implicitization_twisted_cubic() {
let pool = ExprPool::new();
let t = pool.symbol("t", Domain::Real);
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let t2 = pool.pow(t, pool.integer(2_i32));
let p1 = pool.add(vec![x, pool.mul(vec![pool.integer(-1_i32), t2])]);
let t3 = pool.pow(t, pool.integer(3_i32));
let p2 = pool.add(vec![y, pool.mul(vec![pool.integer(-1_i32), t3])]);
let dr = resultant(p1, p2, t, &pool).unwrap();
let res_expr = dr.value;
use crate::kernel::subs;
use std::collections::HashMap;
let one = pool.integer(1_i32);
let two = pool.integer(2_i32);
let four = pool.integer(4_i32);
let eight = pool.integer(8_i32);
let mut map_on = HashMap::new();
map_on.insert(x, four);
map_on.insert(y, eight);
let at_4_8 = subs(res_expr, &map_on, &pool);
let simplified_0 = crate::simplify::simplify(at_4_8, &pool);
match pool.get(simplified_0.value) {
ExprData::Integer(n) => assert_eq!(n.0, 0, "res at (4,8) should be 0"),
_ => {
panic!(
"expected integer 0 at (4,8), got {:?}",
pool.get(simplified_0.value)
)
}
}
let mut map_off = HashMap::new();
map_off.insert(x, one);
map_off.insert(y, two);
let at_1_2 = subs(res_expr, &map_off, &pool);
let simplified_nz = crate::simplify::simplify(at_1_2, &pool);
if let ExprData::Integer(n) = pool.get(simplified_nz.value) {
assert_ne!(n.0, 0, "res at (1,2) should be non-zero");
} }
#[test]
fn sprs_sequence_length() {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let xsq = p.pow(x, p.integer(2_i32));
let poly_p = p.add(vec![xsq, p.integer(1_i32)]);
let poly_q = p.add(vec![x, p.integer(-1_i32)]);
let dr = subresultant_prs(poly_p, poly_q, x, &p).unwrap();
let seq = &dr.value;
assert!(seq.len() >= 2, "sequence must have at least [p, q]");
let last_id = *seq.last().unwrap();
match p.get(last_id) {
ExprData::Integer(_) => {} _ => {
let last_poly = UniPoly::from_symbolic(last_id, x, &p).unwrap();
assert_eq!(last_poly.degree(), 0, "last PRS element should be degree 0");
}
}
}
#[test]
fn sprs_first_elements() {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let two = p.integer(2_i32);
let xsq = p.pow(x, p.integer(2_i32));
let poly_p_expr = p.add(vec![xsq, p.integer(-1_i32)]);
let two_x = p.mul(vec![two, x]);
let poly_q_expr = p.add(vec![two_x, p.integer(-2_i32)]);
let dr = subresultant_prs(poly_p_expr, poly_q_expr, x, &p).unwrap();
assert!(dr.value.len() >= 2);
}
#[test]
fn sprs_gcd_from_sequence() {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let xsq = p.pow(x, p.integer(2_i32));
let poly_p_expr = p.add(vec![xsq, p.integer(-1_i32)]);
let poly_q_expr = p.add(vec![x, p.integer(-1_i32)]);
let dr = subresultant_prs(poly_p_expr, poly_q_expr, x, &p).unwrap();
let seq = &dr.value;
assert!(seq.len() >= 2);
let last_id = *seq.last().unwrap();
let last_poly = UniPoly::from_symbolic(last_id, x, &p).unwrap();
assert_eq!(
last_poly.degree(),
1,
"last PRS element should be degree-1 (matching GCD)"
);
}
#[test]
fn sprs_sylvester_consistency() {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let poly_p_expr = p.add(vec![x, p.integer(-3_i32)]);
let poly_q_expr = p.add(vec![x, p.integer(-7_i32)]);
let dr_prs = subresultant_prs(poly_p_expr, poly_q_expr, x, &p).unwrap();
let dr_res = resultant(poly_p_expr, poly_q_expr, x, &p).unwrap();
let last = *dr_prs.value.last().unwrap();
match p.get(last) {
ExprData::Integer(n) => {
let res_n = match p.get(dr_res.value) {
ExprData::Integer(m) => m.0.clone(),
_ => panic!("resultant not integer"),
};
assert_eq!(n.0.clone().abs(), res_n.abs());
}
_ => {
}
}
}
#[test]
fn resultant_non_polynomial_error() {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let sin_x = p.func("sin", vec![x]);
let poly_q = p.add(vec![x, p.integer(-1_i32)]);
let err = resultant(sin_x, poly_q, x, &p);
assert!(
matches!(err, Err(ResultantError::NotAPolynomial(_))),
"expected NotAPolynomial error"
);
}
fn from_coeffs(p: &ExprPool, x: ExprId, c: &[i64]) -> ExprId {
let terms: Vec<ExprId> = c
.iter()
.enumerate()
.filter(|(_, &k)| k != 0)
.map(|(i, &k)| {
let xi = p.pow(x, p.integer(i as i64));
p.mul(vec![p.integer(k), xi])
})
.collect();
if terms.is_empty() {
p.integer(0_i32)
} else {
p.add(terms)
}
}
fn to_coeffs(p: &ExprPool, x: ExprId, e: ExprId) -> Vec<rug::Integer> {
let mut c = UniPoly::from_symbolic(e, x, p).unwrap().coefficients();
while c.last().is_some_and(|t| *t == 0) {
c.pop();
}
c
}
fn det_rational(mut m: Vec<Vec<rug::Rational>>) -> rug::Rational {
let n = m.len();
let mut d = rug::Rational::from(1);
for i in 0..n {
let Some(piv) = (i..n).find(|&r| m[r][i] != 0) else {
return rug::Rational::from(0);
};
if piv != i {
m.swap(i, piv);
d = -d;
}
let (head, tail) = m.split_at_mut(i + 1);
let pivot_row = &head[i];
d *= pivot_row[i].clone();
let inv = rug::Rational::from(1) / pivot_row[i].clone();
for row in tail.iter_mut() {
let f = row[i].clone() * inv.clone();
if f == 0 {
continue;
}
for (cell, pivot) in row[i..n].iter_mut().zip(pivot_row[i..n].iter()) {
*cell -= f.clone() * pivot.clone();
}
}
}
d
}
fn subresultant_by_determinant(f: &[i64], g: &[i64], j: usize) -> Vec<rug::Integer> {
let m = f.len() - 1;
let n = g.len() - 1;
let width = m + n - j; let row_of = |poly: &[i64], sh: usize| -> Vec<rug::Rational> {
(0..width)
.map(|c| {
let deg = width - 1 - c;
let k = deg.wrapping_sub(sh);
if deg >= sh && k < poly.len() {
rug::Rational::from(poly[k])
} else {
rug::Rational::from(0)
}
})
.collect()
};
let mut rows: Vec<Vec<rug::Rational>> = Vec::new();
for sh in (0..n - j).rev() {
rows.push(row_of(f, sh));
}
for sh in (0..m - j).rev() {
rows.push(row_of(g, sh));
}
let size = m + n - 2 * j;
assert_eq!(rows.len(), size);
let mut out: Vec<rug::Integer> = Vec::new();
for k in 0..=j {
let mut cols: Vec<usize> = (0..size - 1).collect();
cols.push(width - 1 - k);
let sub: Vec<Vec<rug::Rational>> = rows
.iter()
.map(|r| cols.iter().map(|&c| r[c].clone()).collect())
.collect();
let d = det_rational(sub);
assert_eq!(*d.denom(), 1);
out.push(d.numer().clone());
}
while out.last().is_some_and(|t| *t == 0) {
out.pop();
}
out
}
#[test]
fn sprs_matches_the_sylvester_determinants() {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let cases: &[(&[i64], &[i64])] = &[
(&[2, -3, 1], &[0, 2]),
(&[0, -1, 0, 3], &[-3, 2, -3]),
(&[1, 2, 2], &[1, 1, 2]),
(&[1, 0, 1], &[0, 2]),
(&[-2, 0, 0, 3, 2, -1], &[-3, 2, 0, -1, -1]),
(&[-2, -3, -1, 3, 3, -1], &[-3, -2, -2, 0, 2]),
(&[1, 1, 1, 1], &[2, 0, 3]),
(&[-5, 0, 0, 0, 7], &[1, -1, 1]),
];
for (f, g) in cases {
let pf = from_coeffs(&p, x, f);
let pg = from_coeffs(&p, x, g);
let seq = subresultant_prs(pf, pg, x, &p).unwrap().value;
for &elem in &seq[2..] {
let c = to_coeffs(&p, x, elem);
let j = c.len() - 1;
assert_eq!(
c,
subresultant_by_determinant(f, g, j),
"element of degree {j} is not S_{j} for f={f:?}, g={g:?}"
);
}
let last = to_coeffs(&p, x, *seq.last().unwrap());
if last.len() == 1 && seq.len() > 2 {
let r = resultant(pf, pg, x, &p).unwrap().value;
let expected = match p.get(r) {
ExprData::Integer(n) => n.0.clone(),
other => panic!("resultant was not an integer: {other:?}"),
};
assert_eq!(last[0], expected, "last PRS element ≠ resultant");
}
}
}
#[test]
fn sprs_survives_an_inexact_scaling_input() {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let f = from_coeffs(&p, x, &[1, 2, 2]);
let g = from_coeffs(&p, x, &[1, 1, 2]);
let seq = subresultant_prs(f, g, x, &p).unwrap().value;
assert_eq!(
to_coeffs(&p, x, *seq.last().unwrap()),
vec![rug::Integer::from(2)]
);
}
#[test]
fn subresultant_prs_non_polynomial_error() {
let p = ExprPool::new();
let x = p.symbol("x", Domain::Real);
let y = p.symbol("y", Domain::Real);
let poly_p = p.add(vec![x, y]);
let poly_q = p.add(vec![x, p.integer(-1_i32)]);
let err = subresultant_prs(poly_p, poly_q, x, &p);
assert!(
matches!(err, Err(ResultantError::NotAPolynomial(_))),
"expected NotAPolynomial error for multivariate input to subresultant_prs"
);
}
}