use crate::grobner::buchberger::{Relation, SatResult};
use crate::polynomial::{Polynomial, Var};
use crate::prelude::FxHashMap;
use num_rational::BigRational;
use num_traits::{Signed, Zero};
const FM_MAX_ATOMS: usize = 20_000;
const MAX_DISEQUALITY_BRANCHES: u32 = 12;
#[derive(Debug, Clone)]
struct LinearAtom {
coeffs: FxHashMap<Var, BigRational>,
constant: BigRational,
strict: bool,
}
impl LinearAtom {
fn from_polynomial(poly: &Polynomial, relation: Relation) -> Self {
debug_assert!(matches!(
relation,
Relation::Greater | Relation::GreaterEqual | Relation::Less | Relation::LessEqual
));
let negate = matches!(relation, Relation::Less | Relation::LessEqual);
let strict = matches!(relation, Relation::Greater | Relation::Less);
let mut coeffs: FxHashMap<Var, BigRational> = FxHashMap::default();
let mut constant = BigRational::zero();
for term in poly.terms() {
if term.monomial.is_unit() {
constant += term.coeff.clone();
continue;
}
let vp = term.monomial.vars()[0];
let entry = coeffs.entry(vp.var).or_insert_with(BigRational::zero);
*entry += term.coeff.clone();
}
if negate {
constant = -constant;
for c in coeffs.values_mut() {
*c = -core::mem::replace(c, BigRational::zero());
}
}
Self {
coeffs,
constant,
strict,
}
}
fn coeff(&self, var: Var) -> BigRational {
self.coeffs
.get(&var)
.cloned()
.unwrap_or_else(BigRational::zero)
}
}
fn fm_eliminate_var(atoms: Vec<LinearAtom>, var: Var) -> Option<Vec<LinearAtom>> {
let mut no_var = Vec::new();
let mut lower = Vec::new(); let mut upper = Vec::new();
for atom in atoms {
let c = atom.coeff(var);
if c.is_zero() {
no_var.push(atom);
} else if c.is_positive() {
lower.push(atom);
} else {
upper.push(atom);
}
}
if no_var.len() + lower.len() * upper.len() > FM_MAX_ATOMS {
return None;
}
for l in &lower {
let c_v = l.coeff(var); for u in &upper {
let d_v = u.coeff(var); let pos_l = -d_v;
let pos_u = c_v.clone();
let mut coeffs: FxHashMap<Var, BigRational> = FxHashMap::default();
for (&w, c) in &l.coeffs {
if w == var {
continue;
}
*coeffs.entry(w).or_insert_with(BigRational::zero) += &pos_l * c;
}
for (&w, c) in &u.coeffs {
if w == var {
continue;
}
*coeffs.entry(w).or_insert_with(BigRational::zero) += &pos_u * c;
}
coeffs.retain(|_, c| !c.is_zero());
let constant = &pos_l * &l.constant + &pos_u * &u.constant;
let strict = l.strict || u.strict;
no_var.push(LinearAtom {
coeffs,
constant,
strict,
});
}
}
Some(no_var)
}
fn fm_decide(mut atoms: Vec<LinearAtom>, vars: &[Var]) -> SatResult {
for &var in vars {
match fm_eliminate_var(atoms, var) {
Some(next) => atoms = next,
None => return SatResult::Unknown,
}
}
for atom in &atoms {
let ok = if atom.strict {
atom.constant.is_positive()
} else {
!atom.constant.is_negative()
};
if !ok {
return SatResult::Unsat;
}
}
SatResult::Sat
}
pub(crate) fn decide_linear_arms(constraints: &[(&Polynomial, Relation)]) -> SatResult {
let mut definite: Vec<(&Polynomial, Relation)> = Vec::new();
let mut disequalities: Vec<&Polynomial> = Vec::new();
for (poly, relation) in constraints {
match relation {
Relation::NotEqual => disequalities.push(poly),
Relation::Equal => {
definite.push((poly, Relation::GreaterEqual));
definite.push((poly, Relation::LessEqual));
}
other => definite.push((poly, *other)),
}
}
if disequalities.len() as u32 > MAX_DISEQUALITY_BRANCHES {
return SatResult::Unknown;
}
let branch_count = 1u32 << disequalities.len();
let mut any_unknown = false;
for branch in 0..branch_count {
let mut arms = definite.clone();
for (bit, poly) in disequalities.iter().enumerate() {
let relation = if (branch >> bit) & 1 == 0 {
Relation::Greater
} else {
Relation::Less
};
arms.push((*poly, relation));
}
match linear_system_sat(&arms) {
SatResult::Sat => return SatResult::Sat,
SatResult::Unknown => any_unknown = true,
SatResult::Unsat => {}
}
}
if any_unknown {
SatResult::Unknown
} else {
SatResult::Unsat
}
}
fn linear_system_sat(arms: &[(&Polynomial, Relation)]) -> SatResult {
let mut vars: Vec<Var> = Vec::new();
for (poly, _) in arms {
for term in poly.terms() {
for vp in term.monomial.vars() {
if !vars.contains(&vp.var) {
vars.push(vp.var);
}
}
}
}
vars.sort_unstable();
let atoms: Vec<LinearAtom> = arms
.iter()
.map(|(poly, relation)| LinearAtom::from_polynomial(poly, *relation))
.collect();
fm_decide(atoms, &vars)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_fm_eliminate_var_combines_strict_and_non_strict_correctly() {
let ge = LinearAtom::from_polynomial(&Polynomial::from_var(0), Relation::GreaterEqual);
let lt = LinearAtom::from_polynomial(&Polynomial::from_var(0), Relation::Less);
let resolved = fm_eliminate_var(vec![ge, lt], 0).expect("small system stays under cap");
assert_eq!(resolved.len(), 1);
assert!(resolved[0].coeffs.is_empty());
assert!(
resolved[0].strict,
"resolvent of a strict parent must be strict"
);
assert!(
!resolved[0].constant.is_positive() && !resolved[0].constant.is_negative(),
"0 > 0 must reduce to a false constant atom (constant == 0, strict)"
);
}
}