pumpkin_core/engine/variables/
affine_view.rs1use std::cmp::Ordering;
2
3use enumset::EnumSet;
4use pumpkin_checking::CheckerVariable;
5use pumpkin_checking::IntExt;
6
7use super::TransformableVariable;
8use crate::engine::Assignments;
9use crate::engine::notifications::DomainEvent;
10use crate::engine::notifications::OpaqueDomainEvent;
11use crate::engine::notifications::Watchers;
12use crate::engine::predicates::predicate::Predicate;
13use crate::engine::predicates::predicate_constructor::PredicateConstructor;
14use crate::engine::variables::DomainId;
15use crate::engine::variables::IntegerVariable;
16use crate::math::num_ext::NumExt;
17use crate::propagation::EventDispatcher;
18use crate::propagation::EventTarget;
19use crate::propagation::LocalId;
20
21#[derive(Clone, Copy, Hash, Eq, PartialEq)]
24pub struct AffineView<Inner> {
25 pub(crate) inner: Inner,
26 pub(crate) scale: i32,
27 pub(crate) offset: i32,
28}
29
30impl<Inner> AffineView<Inner> {
31 pub fn new(inner: Inner, scale: i32, offset: i32) -> Self {
32 assert_ne!(scale, 0, "Multiplication by zero is not invertable");
33 AffineView {
34 inner,
35 scale,
36 offset,
37 }
38 }
39
40 pub fn inner(&self) -> &Inner {
41 &self.inner
42 }
43
44 fn invert(&self, value: i32, rounding: Rounding) -> i32 {
47 let inverted_translation = value - self.offset;
48
49 match rounding {
50 Rounding::Up => <i32 as NumExt>::div_ceil(inverted_translation, self.scale),
51 Rounding::Down => <i32 as NumExt>::div_floor(inverted_translation, self.scale),
52 }
53 }
54
55 fn map(&self, value: i32) -> i32 {
56 self.scale * value + self.offset
57 }
58}
59
60impl<Inner: EventTarget> EventTarget for AffineView<Inner> {
61 fn register(
62 &self,
63 registration: &mut impl EventDispatcher,
64 mut events: EnumSet<DomainEvent>,
65 local_id: LocalId,
66 ) {
67 let bound = DomainEvent::LowerBound | DomainEvent::UpperBound;
68 let intersection = events.intersection(bound);
69 if intersection.len() == 1 && self.scale.is_negative() {
70 events = events.symmetric_difference(bound);
71 }
72 self.inner.register(registration, events, local_id);
73 }
74}
75
76impl<Var: IntegerVariable> CheckerVariable<Predicate> for AffineView<Var> {
77 fn does_atomic_constrain_self(&self, atomic: &Predicate) -> bool {
78 self.inner.does_atomic_constrain_self(atomic)
79 }
80
81 fn atomic_less_than(&self, value: i32) -> Predicate {
82 use crate::predicate;
83
84 predicate![self <= value]
85 }
86
87 fn atomic_greater_than(&self, value: i32) -> Predicate {
88 use crate::predicate;
89
90 predicate![self >= value]
91 }
92
93 fn atomic_equal(&self, value: i32) -> Predicate {
94 use crate::predicate;
95
96 predicate![self == value]
97 }
98
99 fn atomic_not_equal(&self, value: i32) -> Predicate {
100 use crate::predicate;
101
102 predicate![self != value]
103 }
104
105 fn induced_lower_bound(
106 &self,
107 variable_state: &pumpkin_checking::VariableState<Predicate>,
108 ) -> IntExt {
109 if self.scale.is_positive() {
110 match self.inner.induced_lower_bound(variable_state) {
111 IntExt::Int(value) => IntExt::Int(self.map(value)),
112 bound => bound,
113 }
114 } else {
115 match self.inner.induced_upper_bound(variable_state) {
116 IntExt::Int(value) => IntExt::Int(self.map(value)),
117 IntExt::NegativeInf => IntExt::PositiveInf,
118 IntExt::PositiveInf => IntExt::NegativeInf,
119 }
120 }
121 }
122
123 fn induced_upper_bound(
124 &self,
125 variable_state: &pumpkin_checking::VariableState<Predicate>,
126 ) -> IntExt {
127 if self.scale.is_positive() {
128 match self.inner.induced_upper_bound(variable_state) {
129 IntExt::Int(value) => IntExt::Int(self.map(value)),
130 bound => bound,
131 }
132 } else {
133 match self.inner.induced_lower_bound(variable_state) {
134 IntExt::Int(value) => IntExt::Int(self.map(value)),
135 IntExt::NegativeInf => IntExt::PositiveInf,
136 IntExt::PositiveInf => IntExt::NegativeInf,
137 }
138 }
139 }
140
141 fn induced_fixed_value(
142 &self,
143 variable_state: &pumpkin_checking::VariableState<Predicate>,
144 ) -> Option<i32> {
145 self.inner
146 .induced_fixed_value(variable_state)
147 .map(|value| self.map(value))
148 }
149
150 fn induced_domain_contains(
151 &self,
152 variable_state: &pumpkin_checking::VariableState<Predicate>,
153 value: i32,
154 ) -> bool {
155 let translated_value = value - self.offset;
156
157 if translated_value % self.scale != 0 {
160 return false;
161 }
162
163 let unscaled_value = translated_value / self.scale;
164
165 self.inner
166 .induced_domain_contains(variable_state, unscaled_value)
167 }
168
169 fn induced_holes<'this, 'state>(
170 &'this self,
171 variable_state: &'state pumpkin_checking::VariableState<Predicate>,
172 ) -> impl Iterator<Item = i32> + 'state
173 where
174 'this: 'state,
175 {
176 if self.scale == 1 || self.scale == -1 {
177 return self
178 .inner
179 .induced_holes(variable_state)
180 .map(|value| self.map(value));
181 }
182
183 todo!("how to iterate holes of a scaled domain");
184 }
185
186 fn iter_induced_domain<'this, 'state>(
187 &'this self,
188 variable_state: &'state pumpkin_checking::VariableState<Predicate>,
189 ) -> Option<impl Iterator<Item = i32> + 'state>
190 where
191 'this: 'state,
192 {
193 self.inner
194 .iter_induced_domain(variable_state)
195 .map(|iter| iter.map(|value| self.map(value)))
196 }
197}
198
199impl<View> IntegerVariable for AffineView<View>
200where
201 View: IntegerVariable,
202{
203 type AffineView = Self;
204
205 fn lower_bound(&self, assignment: &Assignments) -> i32 {
206 if self.scale < 0 {
207 self.map(self.inner.upper_bound(assignment))
208 } else {
209 self.map(self.inner.lower_bound(assignment))
210 }
211 }
212
213 fn lower_bound_at_trail_position(
214 &self,
215 assignment: &Assignments,
216 trail_position: usize,
217 ) -> i32 {
218 if self.scale < 0 {
219 self.map(
220 self.inner
221 .upper_bound_at_trail_position(assignment, trail_position),
222 )
223 } else {
224 self.map(
225 self.inner
226 .lower_bound_at_trail_position(assignment, trail_position),
227 )
228 }
229 }
230
231 fn upper_bound(&self, assignment: &Assignments) -> i32 {
232 if self.scale < 0 {
233 self.map(self.inner.lower_bound(assignment))
234 } else {
235 self.map(self.inner.upper_bound(assignment))
236 }
237 }
238
239 fn upper_bound_at_trail_position(
240 &self,
241 assignment: &Assignments,
242 trail_position: usize,
243 ) -> i32 {
244 if self.scale < 0 {
245 self.map(
246 self.inner
247 .lower_bound_at_trail_position(assignment, trail_position),
248 )
249 } else {
250 self.map(
251 self.inner
252 .upper_bound_at_trail_position(assignment, trail_position),
253 )
254 }
255 }
256
257 fn contains(&self, assignment: &Assignments, value: i32) -> bool {
258 if (value - self.offset) % self.scale == 0 {
259 let inverted = self.invert(value, Rounding::Up);
260 self.inner.contains(assignment, inverted)
261 } else {
262 false
263 }
264 }
265
266 fn contains_at_trail_position(
267 &self,
268 assignment: &Assignments,
269 value: i32,
270 trail_position: usize,
271 ) -> bool {
272 if (value - self.offset) % self.scale == 0 {
273 let inverted = self.invert(value, Rounding::Up);
274 self.inner
275 .contains_at_trail_position(assignment, inverted, trail_position)
276 } else {
277 false
278 }
279 }
280
281 fn iterate_domain(&self, assignment: &Assignments) -> impl Iterator<Item = i32> {
282 self.inner
283 .iterate_domain(assignment)
284 .map(|value| self.map(value))
285 }
286
287 fn unwatch_all(&self, watchers: &mut Watchers<'_>) {
288 self.inner.unwatch_all(watchers);
289 }
290
291 fn watch_all_backtrack(&self, watchers: &mut Watchers<'_>, mut events: EnumSet<DomainEvent>) {
292 let bound = DomainEvent::LowerBound | DomainEvent::UpperBound;
293 let intersection = events.intersection(bound);
294 if intersection.len() == 1 && self.scale.is_negative() {
295 events = events.symmetric_difference(bound);
296 }
297 self.inner.watch_all_backtrack(watchers, events);
298 }
299
300 fn unpack_event(&self, event: OpaqueDomainEvent) -> DomainEvent {
301 if self.scale.is_negative() {
302 match self.inner.unpack_event(event) {
303 DomainEvent::LowerBound => DomainEvent::UpperBound,
304 DomainEvent::UpperBound => DomainEvent::LowerBound,
305 event => event,
306 }
307 } else {
308 self.inner.unpack_event(event)
309 }
310 }
311
312 fn get_holes_at_current_checkpoint(
313 &self,
314 assignments: &Assignments,
315 ) -> impl Iterator<Item = i32> {
316 self.inner
317 .get_holes_at_current_checkpoint(assignments)
318 .map(|value| self.map(value))
319 }
320
321 fn get_holes(&self, assignments: &Assignments) -> impl Iterator<Item = i32> {
322 self.inner
323 .get_holes(assignments)
324 .map(|value| self.map(value))
325 }
326}
327
328impl<View> TransformableVariable<AffineView<View>> for AffineView<View>
329where
330 View: IntegerVariable,
331{
332 fn scaled(&self, scale: i32) -> AffineView<View> {
333 let mut result = self.clone();
334 result.scale *= scale;
335 result.offset *= scale;
336 result
337 }
338
339 fn offset(&self, offset: i32) -> AffineView<View> {
340 let mut result = self.clone();
341 result.offset += offset;
342 result
343 }
344}
345
346impl<Var: std::fmt::Debug> std::fmt::Debug for AffineView<Var> {
347 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
348 if self.scale == -1 {
349 write!(f, "-")?;
350 } else if self.scale != 1 {
351 write!(f, "{} * ", self.scale)?;
352 }
353
354 write!(f, "({:?})", self.inner)?;
355
356 match self.offset.cmp(&0) {
357 Ordering::Less => write!(f, " - {}", -self.offset)?,
358 Ordering::Equal => {}
359 Ordering::Greater => write!(f, " + {}", self.offset)?,
360 }
361
362 Ok(())
363 }
364}
365
366impl<Var: PredicateConstructor<Value = i32>> PredicateConstructor for AffineView<Var> {
367 type Value = Var::Value;
368
369 fn lower_bound_predicate(&self, bound: Self::Value) -> Predicate {
370 if self.scale < 0 {
371 let inverted_bound = self.invert(bound, Rounding::Down);
372 self.inner.upper_bound_predicate(inverted_bound)
373 } else {
374 let inverted_bound = self.invert(bound, Rounding::Up);
375 self.inner.lower_bound_predicate(inverted_bound)
376 }
377 }
378
379 fn upper_bound_predicate(&self, bound: Self::Value) -> Predicate {
380 if self.scale < 0 {
381 let inverted_bound = self.invert(bound, Rounding::Up);
382 self.inner.lower_bound_predicate(inverted_bound)
383 } else {
384 let inverted_bound = self.invert(bound, Rounding::Down);
385 self.inner.upper_bound_predicate(inverted_bound)
386 }
387 }
388
389 fn equality_predicate(&self, bound: Self::Value) -> Predicate {
390 if (bound - self.offset) % self.scale == 0 {
391 let inverted_bound = self.invert(bound, Rounding::Up);
392 self.inner.equality_predicate(inverted_bound)
393 } else {
394 Predicate::trivially_false()
395 }
396 }
397
398 fn disequality_predicate(&self, bound: Self::Value) -> Predicate {
399 if (bound - self.offset) % self.scale == 0 {
400 let inverted_bound = self.invert(bound, Rounding::Up);
401 self.inner.disequality_predicate(inverted_bound)
402 } else {
403 Predicate::trivially_true()
404 }
405 }
406}
407
408impl From<DomainId> for AffineView<DomainId> {
409 fn from(value: DomainId) -> Self {
410 AffineView::new(value, 1, 0)
411 }
412}
413
414enum Rounding {
415 Up,
416 Down,
417}
418
419#[cfg(test)]
420mod tests {
421 use super::*;
422 use crate::predicate;
423
424 #[test]
425 fn scaling_an_affine_view() {
426 let view = AffineView::new(DomainId::new(0), 3, 4);
427 assert_eq!(3, view.scale);
428 assert_eq!(4, view.offset);
429 let scaled_view = view.scaled(6);
430 assert_eq!(18, scaled_view.scale);
431 assert_eq!(24, scaled_view.offset);
432 }
433
434 #[test]
435 fn offsetting_an_affine_view() {
436 let view = AffineView::new(DomainId::new(0), 3, 4);
437 assert_eq!(3, view.scale);
438 assert_eq!(4, view.offset);
439 let scaled_view = view.offset(6);
440 assert_eq!(3, scaled_view.scale);
441 assert_eq!(10, scaled_view.offset);
442 }
443
444 #[test]
445 fn affine_view_obtaining_a_bound_should_round_optimistically_in_inner_domain() {
446 let domain = DomainId::new(0);
447 let view = AffineView::new(domain, 2, 0);
448
449 assert_eq!(predicate!(domain >= 1), predicate!(view >= 1));
450 assert_eq!(predicate!(domain >= -1), predicate!(view >= -3));
451 assert_eq!(predicate!(domain <= 0), predicate!(view <= 1));
452 assert_eq!(predicate!(domain <= -3), predicate!(view <= -5));
453 }
454
455 #[test]
456 fn test_negated_variable_has_bounds_rounded_correctly() {
457 let domain = DomainId::new(0);
458 let view = AffineView::new(domain, -2, 0);
459
460 assert_eq!(predicate!(view <= -3), predicate!(domain >= 2));
461 assert_eq!(predicate!(view >= 5), predicate!(domain <= -3));
462 }
463}