Skip to main content

pumpkin_core/engine/variables/
affine_view.rs

1use 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/// Models the constraint `y = ax + b`, by expressing the domain of `y` as a transformation of the
22/// domain of `x`.
23#[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    /// Apply the inverse transformation of this view on a value, to go from the value in the domain
45    /// of `self` to a value in the domain of `self.inner`.
46    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 the translated value does not divide by scale, then the original value is not in the
158        // domain of this affine view.
159        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}