pub mod diophantine;
pub mod homotopy;
pub mod polyhedral;
pub mod regular_chains;
pub mod transcendental;
mod verify;
pub use transcendental::{solve_transcendental, TranscendentalOutcome};
pub use regular_chains::{
extract_regular_chain_from_basis, main_variable_recursive, triangularize, RegularChain,
};
pub use homotopy::{solve_numerical, CertifiedPoint, HomotopyError, HomotopyOpts};
pub use diophantine::{diophantine, DiophantineError, DiophantineSolution};
use crate::errors::AlkahestError;
use crate::kernel::{ExprData, ExprId, ExprPool};
use crate::poly::collect_free_vars;
use crate::poly::groebner::{GbPoly, GroebnerBasis, MonomialOrder};
use rug::ops::Pow;
use rug::Rational;
use std::collections::{BTreeMap, BTreeSet};
use std::fmt;
pub type Solution = Vec<ExprId>;
pub enum SolutionSet {
Finite(Vec<Solution>),
Parametric(GroebnerBasis),
NoSolution,
}
#[derive(Debug, Clone)]
pub enum SolverError {
NotPolynomial(String),
HighDegree(usize),
ShapeMismatch,
}
impl fmt::Display for SolverError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
SolverError::NotPolynomial(s) => write!(f, "not a polynomial: {s}"),
SolverError::HighDegree(d) => write!(
f,
"back-substitution requires solving a degree-{d} univariate polynomial \
(only degree ≤ 2 is currently supported)"
),
SolverError::ShapeMismatch => write!(
f,
"number of equations must equal number of variables for zero-dimensional solving"
),
}
}
}
impl std::error::Error for SolverError {}
impl AlkahestError for SolverError {
fn code(&self) -> &'static str {
match self {
SolverError::NotPolynomial(_) => "E-SOLVE-001",
SolverError::HighDegree(_) => "E-SOLVE-002",
SolverError::ShapeMismatch => "E-SOLVE-003",
}
}
fn remediation(&self) -> Option<&'static str> {
match self {
SolverError::NotPolynomial(_) => Some(
"ensure all equations are polynomial in the declared variables; \
transcendental functions are not supported",
),
SolverError::HighDegree(_) => Some(
"degree > 2 univariate solving is not yet implemented symbolically; \
retry with numeric=True or method=\"homotopy\"",
),
SolverError::ShapeMismatch => {
Some("provide one equation per variable for zero-dimensional system solving")
}
}
}
}
pub fn expr_to_gbpoly(
expr: ExprId,
vars: &[ExprId],
pool: &ExprPool,
) -> Result<GbPoly, SolverError> {
let n = vars.len();
expr_to_gbpoly_rec(expr, vars, n, pool)
}
fn expr_to_gbpoly_rec(
expr: ExprId,
vars: &[ExprId],
n_vars: usize,
pool: &ExprPool,
) -> Result<GbPoly, SolverError> {
if let Some(idx) = vars.iter().position(|&v| v == expr) {
let mut exp = vec![0u32; n_vars];
exp[idx] = 1;
let mut terms = BTreeMap::new();
terms.insert(exp, rug::Rational::from(1));
return Ok(GbPoly { terms, n_vars });
}
enum Node {
Var(usize),
IntConst(rug::Integer),
RatConst(Rational),
FloatConst(f64),
FreeSymbol(String),
Add(Vec<ExprId>),
Mul(Vec<ExprId>),
Pow(ExprId, ExprId),
Func(String),
Other,
}
let node = pool.with(expr, |data| match data {
ExprData::Integer(n) => Node::IntConst(n.0.clone()),
ExprData::Rational(r) => Node::RatConst(r.0.clone()),
ExprData::Float(f) => Node::FloatConst(f.inner.to_f64()),
ExprData::Symbol { name, .. } => {
if let Some(idx) = vars.iter().position(|&v| v == expr) {
Node::Var(idx)
} else {
Node::FreeSymbol(name.clone())
}
}
ExprData::Add(args) => Node::Add(args.clone()),
ExprData::Mul(args) => Node::Mul(args.clone()),
ExprData::Pow { base, exp } => Node::Pow(*base, *exp),
ExprData::Func { name, .. } => Node::Func(name.clone()),
_ => Node::Other,
});
match node {
Node::Var(idx) => {
let mut exp = vec![0u32; n_vars];
exp[idx] = 1;
let mut terms = BTreeMap::new();
terms.insert(exp, Rational::from(1));
Ok(GbPoly { terms, n_vars })
}
Node::IntConst(n) => Ok(GbPoly::constant(Rational::from(n), n_vars)),
Node::RatConst(r) => Ok(GbPoly::constant(r, n_vars)),
Node::FloatConst(v) => {
let r = Rational::from_f64(v).unwrap_or_else(|| Rational::from(0));
Ok(GbPoly::constant(r, n_vars))
}
Node::FreeSymbol(name) => Err(SolverError::NotPolynomial(format!(
"free symbol '{name}' not in variable list"
))),
Node::Add(args) => {
let mut result = GbPoly::zero(n_vars);
for a in args {
let p = expr_to_gbpoly_rec(a, vars, n_vars, pool)?;
result = result.add(&p);
}
Ok(result)
}
Node::Mul(args) => {
let mut result = GbPoly::constant(Rational::from(1), n_vars);
for a in args {
let p = expr_to_gbpoly_rec(a, vars, n_vars, pool)?;
result = result.mul(&p);
}
Ok(result)
}
Node::Pow(base, exp_id) => {
let exp_node = pool.with(exp_id, |d| match d {
ExprData::Integer(n) => Some(n.0.clone()),
_ => None,
});
match exp_node {
Some(n) => {
let n_val = n.to_i64().unwrap_or(-1);
if n_val < 0 {
return Err(SolverError::NotPolynomial(format!(
"negative exponent {n_val} in polynomial"
)));
}
let base_poly = expr_to_gbpoly_rec(base, vars, n_vars, pool)?;
let mut result = GbPoly::constant(Rational::from(1), n_vars);
let mut cur = base_poly;
let mut rem = n_val as u64;
while rem > 0 {
if rem & 1 == 1 {
result = result.mul(&cur);
}
let cur2 = cur.clone();
cur = cur.mul(&cur2);
rem >>= 1;
}
Ok(result)
}
None => Err(SolverError::NotPolynomial(
"symbolic or non-integer exponent".to_string(),
)),
}
}
Node::Func(name) => Err(SolverError::NotPolynomial(format!(
"function '{name}' is not a polynomial"
))),
Node::Other => Err(SolverError::NotPolynomial(
"unsupported expression node".to_string(),
)),
}
}
fn rational_to_expr(r: &Rational, pool: &ExprPool) -> ExprId {
let (num, den) = r.clone().into_numer_denom();
if den == 1 {
pool.integer(num)
} else {
pool.rational(num, den)
}
}
fn neg_expr(e: ExprId, pool: &ExprPool) -> ExprId {
let neg_one = pool.integer(rug::Integer::from(-1));
pool.mul(vec![neg_one, e])
}
fn div_expr(num: ExprId, den: ExprId, pool: &ExprPool) -> ExprId {
let neg_one = pool.integer(rug::Integer::from(-1));
let inv_den = pool.pow(den, neg_one);
pool.mul(vec![num, inv_den])
}
fn is_zero_value(e: ExprId, pool: &ExprPool) -> bool {
is_certain_zero(e, pool) || rational_value(e, pool).is_some_and(|v| v == 0)
}
fn rational_value(expr: ExprId, pool: &ExprPool) -> Option<Rational> {
match pool.get(expr) {
ExprData::Integer(n) => Some(Rational::from(n.0.clone())),
ExprData::Rational(r) => Some(r.0.clone()),
ExprData::Add(args) => args.iter().try_fold(Rational::from(0), |acc, &a| {
Some(acc + rational_value(a, pool)?)
}),
ExprData::Mul(args) => args.iter().try_fold(Rational::from(1), |acc, &a| {
Some(acc * rational_value(a, pool)?)
}),
ExprData::Pow { base, exp } => {
let ExprData::Integer(k) = pool.get(exp) else {
return None;
};
let k = k.0.to_i32()?;
let b = rational_value(base, pool)?;
if k < 0 && b == 0 {
return None;
}
Some(b.pow(k))
}
_ => None,
}
}
fn is_certain_zero(e: ExprId, pool: &ExprPool) -> bool {
match pool.get(e) {
ExprData::Integer(n) => n.0 == 0,
ExprData::Rational(r) => r.0 == 0,
ExprData::Add(args) => args.iter().all(|&a| is_certain_zero(a, pool)),
ExprData::Pow { base, exp } => {
let positive = matches!(pool.get(exp), ExprData::Integer(k) if k.0 > 0);
positive && is_certain_zero(base, pool)
}
ExprData::Func { name, args } if name == "sqrt" && args.len() == 1 => {
is_certain_zero(args[0], pool)
}
_ => false,
}
}
fn extract_coeff_in_var(
poly: &GbPoly,
var_idx: usize,
k: u32,
vars: &[ExprId],
assigned: &[Option<ExprId>],
pool: &ExprPool,
) -> ExprId {
let mut sum_terms: Vec<ExprId> = Vec::new();
for (exp, coeff) in &poly.terms {
let e_k = exp.get(var_idx).copied().unwrap_or(0);
if e_k != k {
continue;
}
let mut factors: Vec<ExprId> = Vec::new();
if *coeff != 1 {
factors.push(rational_to_expr(coeff, pool));
}
for (i, &e) in exp.iter().enumerate() {
if i == var_idx || e == 0 {
continue;
}
let base = assigned
.get(i)
.and_then(|o| o.as_ref())
.copied()
.unwrap_or(vars[i]);
if e == 1 {
factors.push(base);
} else {
let exp_id = pool.integer(rug::Integer::from(e));
factors.push(pool.pow(base, exp_id));
}
}
let term = match factors.len() {
0 => pool.integer(rug::Integer::from(1)),
1 => factors[0],
_ => pool.mul(factors),
};
let signed = if *coeff == 1 {
term
} else {
term
};
sum_terms.push(signed);
}
match sum_terms.len() {
0 => pool.integer(rug::Integer::from(0)),
1 => sum_terms[0],
_ => pool.add(sum_terms),
}
}
fn solve_univariate_symbolic(
coeffs: &[ExprId],
pool: &ExprPool,
) -> Result<Vec<ExprId>, SolverError> {
let mut degree = 0usize;
for (i, &c) in coeffs.iter().enumerate() {
if !is_zero_value(c, pool) {
degree = i;
}
}
match degree {
0 => {
Ok(vec![])
}
1 => {
let a = coeffs[1];
let b = coeffs[0];
let neg_b = neg_expr(b, pool);
Ok(vec![div_expr(neg_b, a, pool)])
}
2 => {
let a = coeffs[2];
let b = coeffs[1];
let c = coeffs[0];
let two = pool.integer(rug::Integer::from(2));
let four = pool.integer(rug::Integer::from(4));
let b2 = pool.pow(b, two);
let four_ac = pool.mul(vec![four, a, c]);
let neg_four_ac = neg_expr(four_ac, pool);
let disc = pool.add(vec![b2, neg_four_ac]);
let two_b = pool.integer(rug::Integer::from(2));
let two_a = pool.mul(vec![two_b, a]);
let neg_b = neg_expr(b, pool);
if is_zero_value(disc, pool) {
return Ok(vec![div_expr(neg_b, two_a, pool)]);
}
let sqrt_disc = pool.func("sqrt", vec![disc]);
let root_plus = div_expr(pool.add(vec![neg_b, sqrt_disc]), two_a, pool);
let neg_sqrt = neg_expr(sqrt_disc, pool);
let root_minus = div_expr(pool.add(vec![neg_b, neg_sqrt]), two_a, pool);
Ok(vec![root_plus, root_minus])
}
d => Err(SolverError::HighDegree(d)),
}
}
fn max_degree_in_var(poly: &GbPoly, var_idx: usize) -> u32 {
poly.terms
.keys()
.map(|e| e.get(var_idx).copied().unwrap_or(0))
.max()
.unwrap_or(0)
}
fn active_solve_vars(poly: &GbPoly, n_solve: usize) -> Vec<usize> {
(0..n_solve)
.filter(|&i| {
poly.terms
.keys()
.any(|e| e.get(i).copied().unwrap_or(0) > 0)
})
.collect()
}
thread_local! {
static ASSUMED_NONZERO: std::cell::RefCell<Vec<ExprId>> =
const { std::cell::RefCell::new(Vec::new()) };
}
fn assume_nonzero(lead: ExprId) {
ASSUMED_NONZERO.with(|c| {
let mut v = c.borrow_mut();
if !v.contains(&lead) {
v.push(lead);
}
});
}
pub fn take_solve_side_conditions() -> Vec<crate::deriv::log::SideCondition> {
ASSUMED_NONZERO.with(|c| {
std::mem::take(&mut *c.borrow_mut())
.into_iter()
.map(crate::deriv::log::SideCondition::NonZero)
.collect()
})
}
fn leading_is_reliable(lead: ExprId, pool: &ExprPool) -> LeadStatus {
if let Some(v) = rational_value(lead, pool) {
return if v != 0 {
LeadStatus::Nonzero
} else {
LeadStatus::Unusable
};
}
match verify::CBallEval::default().eval(lead, pool) {
Ok(ball) => {
if ball.excludes_zero() {
LeadStatus::Nonzero
} else {
LeadStatus::Unusable
}
}
Err(verify::VerifyGap::Unsupported) => LeadStatus::AssumedNonzero,
Err(verify::VerifyGap::Undefined) => LeadStatus::Unusable,
}
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
enum LeadStatus {
Nonzero,
AssumedNonzero,
Unusable,
}
fn find_step(
gens: &[GbPoly],
partial: &[Option<ExprId>],
vars: &[ExprId],
n_solve: usize,
pool: &ExprPool,
) -> Result<Option<(usize, Vec<ExprId>)>, SolverError> {
let mut best: Option<(usize, Vec<ExprId>, u32, Option<ExprId>)> = None;
let mut blocked_by_degree: Option<u32> = None;
for g in gens {
let unassigned: Vec<usize> = active_solve_vars(g, n_solve)
.into_iter()
.filter(|&i| partial[i].is_none())
.collect();
let [var_idx] = unassigned[..] else {
continue;
};
let deg = max_degree_in_var(g, var_idx);
if deg == 0 {
continue;
}
if deg > 2 {
blocked_by_degree = Some(blocked_by_degree.map_or(deg, |d: u32| d.min(deg)));
continue;
}
if best.as_ref().is_some_and(|(_, _, bd, _)| *bd <= deg) {
continue;
}
let coeffs: Vec<ExprId> = (0..=deg)
.map(|k| extract_coeff_in_var(g, var_idx, k, vars, partial, pool))
.collect();
let lead = coeffs[deg as usize];
let assumed = match leading_is_reliable(lead, pool) {
LeadStatus::Unusable => continue,
LeadStatus::Nonzero => None,
LeadStatus::AssumedNonzero => Some(lead),
};
best = Some((var_idx, coeffs, deg, assumed));
}
match best {
Some((var_idx, coeffs, _, assumed)) => {
if let Some(lead) = assumed {
assume_nonzero(lead);
}
Ok(Some((var_idx, coeffs)))
}
None => match blocked_by_degree {
Some(d) => Err(SolverError::HighDegree(d as usize)),
None => Ok(None),
},
}
}
enum BacksolveOutcome {
Finite(Vec<Solution>),
Stuck,
NoSolution,
}
fn try_backsolve_generators(
gens: &[GbPoly],
vars: &[ExprId],
n_solve: usize,
pool: &ExprPool,
) -> Result<BacksolveOutcome, SolverError> {
let n_vars = vars.len();
debug_assert!(n_solve <= n_vars);
let mut initial = vec![None; n_vars];
for i in n_solve..n_vars {
initial[i] = Some(vars[i]);
}
let mut partials: Vec<Vec<Option<ExprId>>> = vec![initial];
for _ in 0..n_solve {
let mut new_partials = Vec::new();
let mut high_degree: Option<SolverError> = None;
for partial in &partials {
let step = match find_step(gens, partial, vars, n_solve, pool) {
Ok(s) => s,
Err(e) => {
high_degree = Some(e);
continue;
}
};
let Some((var_idx, coeffs)) = step else {
if partial_is_refuted(gens, partial, n_solve, n_vars, pool) {
continue;
}
return Ok(BacksolveOutcome::Stuck);
};
for root in solve_univariate_symbolic(&coeffs, pool)? {
let mut np = partial.clone();
np[var_idx] = Some(root);
new_partials.push(np);
}
}
if let Some(e) = high_degree {
return Err(e);
}
partials = new_partials;
if partials.is_empty() {
return Ok(BacksolveOutcome::NoSolution);
}
}
let solutions: Vec<Solution> = partials
.into_iter()
.map(|p| {
p.into_iter()
.take(n_solve)
.map(|o| o.expect("all solve vars assigned"))
.collect()
})
.collect();
Ok(BacksolveOutcome::Finite(solutions))
}
fn partial_is_refuted(
gens: &[GbPoly],
partial: &[Option<ExprId>],
n_solve: usize,
n_vars: usize,
pool: &ExprPool,
) -> bool {
if n_solve != n_vars {
return false; }
let mut evaluator = verify::CBallEval::default();
let mut values: Vec<Option<verify::CBall>> = Vec::with_capacity(n_vars);
for slot in partial.iter().take(n_vars) {
values.push(match slot {
Some(v) => evaluator.eval(*v, pool).ok(),
None => None,
});
}
gens.iter()
.any(|g| verify::poly_residual_partial(g, &values).is_some_and(|r| r.excludes_zero()))
}
fn refine_solutions(
solutions: Vec<Solution>,
orig_polys: &[GbPoly],
n_vars: usize,
pool: &ExprPool,
) -> Vec<Solution> {
let mut kept: Vec<Solution> = Vec::new();
let mut kept_values: Vec<Vec<verify::CBall>> = Vec::new();
let mut evaluator = verify::CBallEval::default();
for sol in solutions {
let mut values: Vec<verify::CBall> = Vec::with_capacity(n_vars);
let mut gap = None;
for &v in &sol {
match evaluator.eval(v, pool) {
Ok(b) => values.push(b),
Err(g) => {
gap = Some(g);
break;
}
}
}
match gap {
Some(verify::VerifyGap::Unsupported) => {
kept.push(sol);
continue;
}
Some(verify::VerifyGap::Undefined) => continue,
None => {}
}
if values.len() < n_vars {
kept.push(sol);
continue;
}
if verify::is_refuted(orig_polys, &values) {
continue;
}
if kept_values
.iter()
.any(|prev| verify::same_point(prev, &values))
{
continue;
}
kept_values.push(values);
kept.push(sol);
}
kept
}
fn collect_parameters(equations: &[ExprId], vars: &[ExprId], pool: &ExprPool) -> Vec<ExprId> {
let declared: BTreeSet<ExprId> = vars.iter().copied().collect();
let mut params = BTreeSet::new();
for &eq in equations {
for v in collect_free_vars(eq, pool) {
if !declared.contains(&v) {
params.insert(v);
}
}
}
params.into_iter().collect()
}
pub fn solve_polynomial_system(
equations: Vec<ExprId>,
vars: Vec<ExprId>,
pool: &ExprPool,
) -> Result<SolutionSet, SolverError> {
let _ = take_solve_side_conditions();
let n_solve = vars.len();
let params = collect_parameters(&equations, &vars, pool);
let mut all_vars = vars;
all_vars.extend(params);
let n_vars = all_vars.len();
let mut polys: Vec<GbPoly> = Vec::with_capacity(equations.len());
for eq in &equations {
polys.push(expr_to_gbpoly(*eq, &all_vars, pool)?);
}
let gb = GroebnerBasis::compute(polys.clone(), MonomialOrder::Lex);
let gens = gb.generators();
if gens.len() == 1
&& gens[0].terms.len() == 1
&& gens[0].leading_exp(MonomialOrder::Lex) == Some(vec![0u32; n_vars])
{
return Ok(SolutionSet::NoSolution);
}
let finish = |solutions: Vec<Solution>| -> Option<SolutionSet> {
let had_candidates = !solutions.is_empty();
let refined = refine_solutions(solutions, &polys, n_vars, pool);
if had_candidates && refined.is_empty() {
return None;
}
Some(SolutionSet::Finite(refined))
};
match try_backsolve_generators(gens, &all_vars, n_solve, pool)? {
BacksolveOutcome::Finite(solutions) => {
if let Some(set) = finish(solutions) {
return Ok(set);
}
}
BacksolveOutcome::NoSolution => return Ok(SolutionSet::NoSolution),
BacksolveOutcome::Stuck => {}
}
let chain = extract_regular_chain_from_basis(gens, n_vars, MonomialOrder::Lex);
if !chain.polys.is_empty() {
if let BacksolveOutcome::Finite(solutions) =
try_backsolve_generators(&chain.polys, &all_vars, n_solve, pool)?
{
if let Some(set) = finish(solutions) {
return Ok(set);
}
}
}
Ok(SolutionSet::Parametric(gb))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::jit::eval_interp;
use crate::kernel::{Domain, ExprPool};
use std::collections::HashMap;
fn eval_no_env(e: ExprId, pool: &ExprPool) -> f64 {
eval_interp(e, &HashMap::new(), pool).expect("numeric eval")
}
fn has_numeric_pair(sols: &[Solution], pool: &ExprPool, expected: &[(f64, f64)]) -> bool {
let tol = 1e-10;
expected.iter().all(|(ex, ey)| {
sols.iter().any(|s| {
let x = eval_no_env(s[0], pool);
let y = eval_no_env(s[1], pool);
(x - ex).abs() < tol && (y - ey).abs() < tol
})
})
}
#[test]
fn linear_system() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let neg_one = pool.integer(-1_i32);
let eq1 = pool.add(vec![x, y, neg_one]);
let eq2 = pool.add(vec![x, pool.mul(vec![neg_one, y])]);
let result = solve_polynomial_system(vec![eq1, eq2], vec![x, y], &pool).unwrap();
if let SolutionSet::Finite(sols) = result {
assert!(has_numeric_pair(&sols, &pool, &[(0.5, 0.5)]));
} else {
panic!("expected finite solution set");
}
}
#[test]
fn univariate_quadratic() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let neg_one = pool.integer(-1_i32);
let x2 = pool.pow(x, pool.integer(2_i32));
let eq = pool.add(vec![x2, neg_one]);
let result = solve_polynomial_system(vec![eq], vec![x], &pool).unwrap();
if let SolutionSet::Finite(sols) = result {
let vals: Vec<f64> = sols.iter().map(|s| eval_no_env(s[0], &pool)).collect();
assert!(vals.iter().any(|v| (v - 1.0).abs() < 1e-10));
assert!(vals.iter().any(|v| (v + 1.0).abs() < 1e-10));
} else {
panic!("expected finite solution set");
}
}
#[test]
fn circle_line_intersection() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let neg_one = pool.integer(-1_i32);
let two = pool.integer(2_i32);
let x2 = pool.pow(x, two);
let y2 = pool.pow(y, two);
let eq1 = pool.add(vec![x2, y2, neg_one]);
let eq2 = pool.add(vec![y, pool.mul(vec![neg_one, x])]);
let result = solve_polynomial_system(vec![eq1, eq2], vec![x, y], &pool).unwrap();
if let SolutionSet::Finite(sols) = result {
assert_eq!(
sols.len(),
2,
"expected exactly 2 solutions, got {}",
sols.len()
);
let root = (0.5_f64).sqrt(); assert!(has_numeric_pair(
&sols,
&pool,
&[(root, root), (-root, -root)]
));
} else {
panic!("expected finite solution set");
}
}
#[test]
fn no_solution_inconsistent() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let neg_one = pool.integer(-1_i32);
let eq1 = x; let eq2 = pool.add(vec![x, neg_one]); let result = solve_polynomial_system(vec![eq1, eq2], vec![x], &pool).unwrap();
assert!(matches!(result, SolutionSet::NoSolution));
}
#[test]
fn parabola_and_line() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let neg_one = pool.integer(-1_i32);
let two = pool.integer(2_i32);
let x2 = pool.pow(x, two);
let eq1 = pool.add(vec![y, pool.mul(vec![neg_one, x2])]);
let eq2 = pool.add(vec![y, pool.mul(vec![neg_one, x])]);
let result = solve_polynomial_system(vec![eq1, eq2], vec![x, y], &pool).unwrap();
if let SolutionSet::Finite(sols) = result {
assert_eq!(sols.len(), 2);
assert!(has_numeric_pair(&sols, &pool, &[(0.0, 0.0), (1.0, 1.0)]));
} else {
panic!("expected finite solution set");
}
}
fn powk(pool: &ExprPool, base: ExprId, k: i32) -> ExprId {
pool.pow(base, pool.integer(k))
}
fn finite(eqs: Vec<ExprId>, vars: Vec<ExprId>, pool: &ExprPool) -> Vec<Solution> {
match solve_polynomial_system(eqs, vars, pool).expect("solve") {
SolutionSet::Finite(s) => s,
other => panic!(
"expected a finite solution set, got {}",
match other {
SolutionSet::NoSolution => "NoSolution",
_ => "Parametric",
}
),
}
}
#[test]
fn spurious_tuple_is_refuted() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let neg_one = pool.integer(-1_i32);
let eq1 = pool.add(vec![powk(&pool, x, 2), pool.mul(vec![neg_one, x, y])]);
let eq2 = pool.add(vec![pool.mul(vec![x, y]), pool.mul(vec![neg_one, y])]);
let sols = finite(vec![eq1, eq2], vec![x, y], &pool);
assert!(has_numeric_pair(&sols, &pool, &[(0.0, 0.0), (1.0, 1.0)]));
assert_eq!(sols.len(), 2, "exactly two points, got {sols:?}");
}
#[test]
fn vanishing_leading_coefficient_branch_is_kept() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let eq1 = pool.add(vec![
pool.mul(vec![pool.integer(-3_i32), x]),
pool.mul(vec![pool.integer(-2_i32), x, y]),
]);
let eq2 = pool.add(vec![
pool.mul(vec![pool.integer(-3_i32), y]),
pool.mul(vec![pool.integer(-1_i32), powk(&pool, x, 2)]),
]);
let sols = finite(vec![eq1, eq2], vec![x, y], &pool);
let r = (4.5_f64).sqrt();
assert!(has_numeric_pair(
&sols,
&pool,
&[(0.0, 0.0), (r, -1.5), (-r, -1.5)]
));
assert_eq!(sols.len(), 3, "exactly three points, got {sols:?}");
}
#[test]
fn undefined_coordinate_is_not_a_solution() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let eq1 = pool.add(vec![
pool.mul(vec![x, y]),
pool.mul(vec![pool.integer(-1_i32), y]),
]);
let eq2 = pool.add(vec![
y,
pool.mul(vec![pool.integer(-2_i32), powk(&pool, x, 2)]),
]);
let sols = finite(vec![eq1, eq2], vec![x, y], &pool);
assert!(has_numeric_pair(&sols, &pool, &[(0.0, 0.0), (1.0, 2.0)]));
assert_eq!(sols.len(), 2, "exactly two points, got {sols:?}");
}
#[test]
fn unfolded_vanishing_discriminant_is_recognised() {
let pool = ExprPool::new();
let zero = pool.integer(0_i32);
let b2 = pool.pow(zero, pool.integer(2_i32));
let four_ac = pool.mul(vec![pool.integer(4_i32), pool.integer(1_i32), zero]);
let disc = pool.add(vec![b2, pool.mul(vec![pool.integer(-1_i32), four_ac])]);
assert!(is_zero_value(disc, &pool), "0² − 4·1·0 = 0");
let b2 = pool.pow(pool.integer(-2_i32), pool.integer(2_i32));
let four_ac = pool.mul(vec![
pool.integer(4_i32),
pool.integer(1_i32),
pool.integer(1_i32),
]);
let disc = pool.add(vec![b2, pool.mul(vec![pool.integer(-1_i32), four_ac])]);
assert!(is_zero_value(disc, &pool), "(−2)² − 4·1·1 = 0");
assert!(!is_zero_value(pool.integer(8_i32), &pool));
assert!(!is_zero_value(pool.symbol("a", Domain::Real), &pool));
assert!(is_zero_value(pool.func("sqrt", vec![zero]), &pool));
}
#[test]
fn conjugate_roots_behind_a_nested_radical_both_survive() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let eq1 = pool.mul(vec![x, y]);
let eq2 = pool.add(vec![
powk(&pool, x, 2),
pool.mul(vec![pool.integer(-1_i32), y]),
pool.integer(1_i32),
]);
let sols = finite(vec![eq1, eq2], vec![x, y], &pool);
assert_eq!(sols.len(), 3, "(0,1) and (±i,0), got {sols:?}");
}
#[test]
fn repeated_root_is_one_solution() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let sols = finite(vec![powk(&pool, x, 2)], vec![x], &pool);
assert_eq!(sols.len(), 1, "{sols:?}");
assert!(eval_no_env(sols[0][0], &pool).abs() < 1e-12);
}
#[test]
fn repeated_roots_do_not_multiply_across_variables() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let z = pool.symbol("z", Domain::Real);
let sols = finite(
vec![powk(&pool, x, 2), powk(&pool, y, 2), powk(&pool, z, 2)],
vec![x, y, z],
&pool,
);
assert_eq!(sols.len(), 1, "{sols:?}");
}
#[test]
fn shifted_double_root_is_one_solution() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let shifted = pool.add(vec![x, pool.integer(-1_i32)]);
let sols = finite(vec![powk(&pool, shifted, 2)], vec![x], &pool);
assert_eq!(sols.len(), 1, "{sols:?}");
assert!((eval_no_env(sols[0][0], &pool) - 1.0).abs() < 1e-12);
}
#[test]
fn parametric_quadratic_free_rhs() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let two = pool.integer(2_i32);
let x2 = pool.pow(x, two);
let eq = pool.add(vec![x2, pool.mul(vec![pool.integer(-1_i32), y])]);
let result = solve_polynomial_system(vec![eq], vec![x], &pool).unwrap();
let SolutionSet::Finite(sols) = result else {
panic!("expected finite parametric solutions");
};
assert_eq!(sols.len(), 2);
let mut env = HashMap::new();
env.insert(y, 4.0);
let vals: Vec<f64> = sols
.iter()
.map(|s| eval_interp(s[0], &env, &pool).expect("eval"))
.collect();
assert!(vals.iter().any(|v| (v - 2.0).abs() < 1e-10));
assert!(vals.iter().any(|v| (v + 2.0).abs() < 1e-10));
}
#[test]
fn parametric_linear_affine() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let a = pool.symbol("a", Domain::Real);
let b = pool.symbol("b", Domain::Real);
let eq = pool.add(vec![
pool.mul(vec![a, x]),
pool.mul(vec![pool.integer(-1_i32), b]),
]);
let result = solve_polynomial_system(vec![eq], vec![x], &pool).unwrap();
let SolutionSet::Finite(sols) = result else {
panic!("expected finite parametric solution");
};
assert_eq!(sols.len(), 1);
let mut env = HashMap::new();
env.insert(a, 2.0);
env.insert(b, 6.0);
let val = eval_interp(sols[0][0], &env, &pool).expect("eval");
assert!((val - 3.0).abs() < 1e-10);
}
#[test]
fn a_parametric_division_states_its_non_vanishing_hypothesis() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let a = pool.symbol("a", Domain::Real);
let b = pool.symbol("b", Domain::Real);
let eq = pool.add(vec![
pool.mul(vec![a, x]),
pool.mul(vec![pool.integer(-1_i32), b]),
]);
let _ = solve_polynomial_system(vec![eq], vec![x], &pool).unwrap();
let conds = take_solve_side_conditions();
assert_eq!(conds.len(), 1, "{conds:?}");
let crate::deriv::log::SideCondition::NonZero(id) = conds[0] else {
panic!("expected a non-vanishing hypothesis, got {:?}", conds[0]);
};
assert_eq!(id, a);
assert!(take_solve_side_conditions().is_empty());
}
#[test]
fn a_solve_that_proves_its_divisors_states_nothing() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let b = pool.symbol("b", Domain::Real);
let eq = pool.add(vec![
pool.mul(vec![pool.integer(2_i32), x]),
pool.mul(vec![pool.integer(-1_i32), b]),
]);
let result = solve_polynomial_system(vec![eq], vec![x], &pool).unwrap();
assert!(matches!(result, SolutionSet::Finite(ref s) if s.len() == 1));
assert!(take_solve_side_conditions().is_empty());
}
#[test]
fn parametric_system_line_with_parameter() {
let pool = ExprPool::new();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let c = pool.symbol("c", Domain::Real);
let neg_one = pool.integer(-1_i32);
let eq1 = pool.add(vec![x, y, pool.mul(vec![neg_one, c])]);
let eq2 = pool.add(vec![x, pool.mul(vec![neg_one, y])]);
let result = solve_polynomial_system(vec![eq1, eq2], vec![x, y], &pool).unwrap();
let SolutionSet::Finite(sols) = result else {
panic!("expected finite parametric solutions");
};
assert_eq!(sols.len(), 1);
let mut env = HashMap::new();
env.insert(c, 4.0);
let xv = eval_interp(sols[0][0], &env, &pool).expect("eval x");
let yv = eval_interp(sols[0][1], &env, &pool).expect("eval y");
assert!((xv - 2.0).abs() < 1e-10);
assert!((yv - 2.0).abs() < 1e-10);
}
}