1use crate::solver::Cmp;
4use crate::solver::poly::{Poly, Var, poly};
5
6pub fn prove(cmp: Cmp, lhs: &Poly, rhs: &Poly, nonneg: &dyn Fn(Var) -> bool) -> bool {
8 match cmp {
9 Cmp::Eq => lhs == rhs,
10 Cmp::Ne => (lhs.sub(rhs)).as_constant().is_some_and(|c| !c.is_zero()),
11 Cmp::Le => less_or_eq(lhs, rhs, nonneg),
12 Cmp::Lt => less_or_eq(&(lhs.add(&poly!(1))), rhs, nonneg),
13 Cmp::Ge => less_or_eq(rhs, lhs, nonneg),
14 Cmp::Gt => less_or_eq(&(rhs.add(&poly!(1))), lhs, nonneg),
15 }
16}
17
18fn less_or_eq(lhs: &Poly, rhs: &Poly, nonneg: &dyn Fn(Var) -> bool) -> bool {
19 lhs == rhs || rhs.sub(lhs).always_nonneg(nonneg)
20}
21
22#[cfg(test)]
23mod tests {
24 use super::{Cmp, Var, poly, prove};
25
26 fn all_but_a(v: Var) -> bool {
27 v != 'a' as Var
28 }
29
30 #[test]
31 fn nonnegative_differences_prove_leq() {
32 assert!(prove(Cmp::Le, &poly!(n), &poly!(n + 1), &all_but_a));
33 assert!(prove(Cmp::Le, &poly!(0), &poly!(n), &all_but_a));
34 assert!(prove(Cmp::Le, &poly!(n), &poly!(2 * n), &all_but_a));
35
36 assert!(!prove(Cmp::Lt, &poly!(0), &poly!(n), &all_but_a));
37 assert!(!prove(Cmp::Le, &poly!(1), &poly!(n), &all_but_a));
38 }
39
40 #[test]
41 fn rigid_sizes_not_nonneg() {
42 assert!(!prove(Cmp::Le, &poly!(0), &poly!(a), &all_but_a));
43 }
44}