Skip to main content

pumpkin_constraints/constraints/arithmetic/
equality.rs

1use pumpkin_core::Solver;
2use pumpkin_core::constraints::Constraint;
3use pumpkin_core::constraints::NegatableConstraint;
4use pumpkin_core::options::ReifiedPropagatorArgs;
5use pumpkin_core::proof::ConstraintTag;
6use pumpkin_core::variables::IntegerVariable;
7use pumpkin_core::variables::Literal;
8use pumpkin_core::variables::TransformableVariable;
9use pumpkin_propagators::arithmetic::BinaryEqualsPropagatorArgs;
10use pumpkin_propagators::arithmetic::BinaryNotEqualsPropagatorArgs;
11use pumpkin_propagators::arithmetic::LinearNotEqualPropagatorArgs;
12
13use super::less_than_or_equals;
14
15struct EqualConstraint<Var> {
16    terms: Box<[Var]>,
17    rhs: i32,
18    constraint_tag: ConstraintTag,
19}
20
21/// Creates the [`NegatableConstraint`] `∑ terms_i = rhs`.
22///
23/// Its negation is [`not_equals`].
24pub fn equals<Var: IntegerVariable + Clone + 'static>(
25    terms: impl Into<Box<[Var]>>,
26    rhs: i32,
27    constraint_tag: ConstraintTag,
28) -> impl NegatableConstraint {
29    EqualConstraint {
30        terms: terms.into(),
31        rhs,
32        constraint_tag,
33    }
34}
35
36/// Creates the [`NegatableConstraint`] `lhs = rhs`.
37///
38/// Its negation is [`binary_not_equals`].
39pub fn binary_equals<Var: IntegerVariable + 'static>(
40    lhs: Var,
41    rhs: Var,
42    constraint_tag: ConstraintTag,
43) -> impl NegatableConstraint {
44    EqualConstraint {
45        terms: [lhs.scaled(1), rhs.scaled(-1)].into(),
46        rhs: 0,
47        constraint_tag,
48    }
49}
50
51struct NotEqualConstraint<Var> {
52    terms: Box<[Var]>,
53    rhs: i32,
54    constraint_tag: ConstraintTag,
55}
56
57/// Create the [`NegatableConstraint`] `∑ terms_i != rhs`.
58///
59/// Its negation is [`equals`].
60pub fn not_equals<Var: IntegerVariable + Clone + 'static>(
61    terms: impl Into<Box<[Var]>>,
62    rhs: i32,
63    constraint_tag: ConstraintTag,
64) -> impl NegatableConstraint {
65    equals(terms, rhs, constraint_tag).negation()
66}
67
68/// Creates the [`NegatableConstraint`] `lhs != rhs`.
69///
70/// Its negation is [`binary_equals`].
71pub fn binary_not_equals<Var: IntegerVariable + 'static>(
72    lhs: Var,
73    rhs: Var,
74    constraint_tag: ConstraintTag,
75) -> impl NegatableConstraint {
76    NotEqualConstraint {
77        terms: [lhs.scaled(1), rhs.scaled(-1)].into(),
78        rhs: 0,
79        constraint_tag,
80    }
81}
82
83impl<Var> Constraint for EqualConstraint<Var>
84where
85    Var: IntegerVariable + Clone + 'static,
86{
87    fn post(self, solver: &mut Solver) {
88        if self.terms.len() == 2 && !solver.is_logging_proof() {
89            let _ = solver.add_propagator(BinaryEqualsPropagatorArgs {
90                a: self.terms[0].clone(),
91                b: self.terms[1].scaled(-1).offset(self.rhs),
92                constraint_tag: self.constraint_tag,
93            });
94        } else {
95            less_than_or_equals(self.terms.clone(), self.rhs, self.constraint_tag).post(solver);
96
97            let negated = self
98                .terms
99                .iter()
100                .map(|var| var.scaled(-1))
101                .collect::<Box<[_]>>();
102            less_than_or_equals(negated, -self.rhs, self.constraint_tag).post(solver);
103        }
104    }
105
106    fn implied_by(self, solver: &mut Solver, reification_literal: Literal) {
107        if self.terms.len() == 2 && !solver.is_logging_proof() {
108            let _ = solver.add_propagator(ReifiedPropagatorArgs {
109                propagator: BinaryEqualsPropagatorArgs {
110                    a: self.terms[0].clone(),
111                    b: self.terms[1].scaled(-1).offset(self.rhs),
112                    constraint_tag: self.constraint_tag,
113                },
114                reification_literal,
115            });
116        } else {
117            less_than_or_equals(self.terms.clone(), self.rhs, self.constraint_tag)
118                .implied_by(solver, reification_literal);
119
120            let negated = self
121                .terms
122                .iter()
123                .map(|var| var.scaled(-1))
124                .collect::<Box<[_]>>();
125            less_than_or_equals(negated, -self.rhs, self.constraint_tag)
126                .implied_by(solver, reification_literal);
127        }
128    }
129}
130
131impl<Var> NegatableConstraint for EqualConstraint<Var>
132where
133    Var: IntegerVariable + Clone + 'static,
134{
135    type NegatedConstraint = NotEqualConstraint<Var>;
136
137    fn negation(&self) -> Self::NegatedConstraint {
138        NotEqualConstraint {
139            terms: self.terms.clone(),
140            rhs: self.rhs,
141            constraint_tag: self.constraint_tag,
142        }
143    }
144}
145
146impl<Var> Constraint for NotEqualConstraint<Var>
147where
148    Var: IntegerVariable + Clone + 'static,
149{
150    fn post(self, solver: &mut Solver) {
151        let NotEqualConstraint {
152            terms,
153            rhs,
154            constraint_tag,
155        } = self;
156
157        if terms.len() == 2 {
158            let _ = solver.add_propagator(BinaryNotEqualsPropagatorArgs {
159                a: terms[0].clone(),
160                b: terms[1].scaled(-1).offset(self.rhs),
161                constraint_tag: self.constraint_tag,
162            });
163        } else {
164            LinearNotEqualPropagatorArgs {
165                terms: terms.into(),
166                rhs,
167                constraint_tag,
168            }
169            .post(solver)
170        }
171    }
172
173    fn implied_by(self, solver: &mut Solver, reification_literal: Literal) {
174        let NotEqualConstraint {
175            terms,
176            rhs,
177            constraint_tag,
178        } = self;
179
180        if terms.len() == 2 {
181            let _ = solver.add_propagator(ReifiedPropagatorArgs {
182                propagator: BinaryNotEqualsPropagatorArgs {
183                    a: terms[0].clone(),
184                    b: terms[1].scaled(-1).offset(self.rhs),
185                    constraint_tag: self.constraint_tag,
186                },
187                reification_literal,
188            });
189        } else {
190            LinearNotEqualPropagatorArgs {
191                terms: terms.into(),
192                rhs,
193                constraint_tag,
194            }
195            .implied_by(solver, reification_literal)
196        }
197    }
198}
199
200impl<Var> NegatableConstraint for NotEqualConstraint<Var>
201where
202    Var: IntegerVariable + Clone + 'static,
203{
204    type NegatedConstraint = EqualConstraint<Var>;
205
206    fn negation(&self) -> Self::NegatedConstraint {
207        EqualConstraint {
208            terms: self.terms.clone(),
209            rhs: self.rhs,
210            constraint_tag: self.constraint_tag,
211        }
212    }
213}