pumpkin_constraints/constraints/arithmetic/
equality.rs1use 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
21pub 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
36pub 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
57pub 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
68pub 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}