use pumpkin_checking::AtomicConstraint;
use pumpkin_checking::CheckerVariable;
use pumpkin_checking::InferenceChecker;
use pumpkin_checking::IntExt;
use pumpkin_core::conjunction;
use pumpkin_core::declare_inference_label;
use pumpkin_core::predicate;
use pumpkin_core::predicates::PropositionalConjunction;
use pumpkin_core::proof::ConstraintTag;
use pumpkin_core::proof::InferenceCode;
use pumpkin_core::propagation::DomainEvents;
use pumpkin_core::propagation::EventsToRegister;
use pumpkin_core::propagation::LocalId;
use pumpkin_core::propagation::Priority;
use pumpkin_core::propagation::PropagationContext;
use pumpkin_core::propagation::Propagator;
use pumpkin_core::propagation::PropagatorConstructor;
use pumpkin_core::propagation::PropagatorConstructorContext;
use pumpkin_core::propagation::PropagatorSpec;
use pumpkin_core::propagation::ReadDomains;
use pumpkin_core::propagation::RuntimeCheckers;
use pumpkin_core::state::PropagationStatusCP;
use pumpkin_core::variables::IntegerVariable;
#[derive(Clone, Debug)]
pub struct MaximumArgs<ElementVar, Rhs> {
pub array: Box<[ElementVar]>,
pub rhs: Rhs,
pub constraint_tag: ConstraintTag,
}
declare_inference_label!(Maximum);
impl<ElementVar, Rhs> PropagatorConstructor for MaximumArgs<ElementVar, Rhs>
where
ElementVar: IntegerVariable + 'static,
Rhs: IntegerVariable + 'static,
{
type PropagatorImpl = MaximumPropagator<ElementVar, Rhs>;
fn create(self, _: PropagatorConstructorContext) -> PropagatorSpec<Self::PropagatorImpl> {
let MaximumArgs {
array,
rhs,
constraint_tag,
} = self;
let mut registration = EventsToRegister::builder();
for (idx, var) in array.iter().enumerate() {
registration = registration.add(var, DomainEvents::BOUNDS, LocalId::from(idx as u32));
}
registration = registration.add(
&rhs,
DomainEvents::BOUNDS,
LocalId::from(array.len() as u32),
);
let mut checkers = RuntimeCheckers::builder();
let inference_code = checkers.add_inference_checker(
constraint_tag,
Maximum,
MaximumChecker {
array: array.clone(),
rhs: rhs.clone(),
},
);
let propagator = MaximumPropagator {
array,
rhs,
inference_code,
};
PropagatorSpec {
registration: registration.build(),
checkers: checkers.build(),
propagator,
}
}
}
#[derive(Clone, Debug)]
pub struct MaximumPropagator<ElementVar, Rhs> {
array: Box<[ElementVar]>,
rhs: Rhs,
inference_code: InferenceCode,
}
impl<ElementVar: IntegerVariable + 'static, Rhs: IntegerVariable + 'static> Propagator
for MaximumPropagator<ElementVar, Rhs>
{
fn priority(&self) -> Priority {
Priority::High
}
fn name(&self) -> &str {
"Maximum"
}
fn propagate_from_scratch(&self, mut context: PropagationContext) -> PropagationStatusCP {
let rhs_ub = context.upper_bound(&self.rhs);
let mut max_ub = context.upper_bound(&self.array[0]);
let mut max_lb = context.lower_bound(&self.array[0]);
let mut lb_reason = predicate![self.array[0] >= max_lb];
for var in self.array.iter() {
context.post(
predicate![var <= rhs_ub],
(conjunction!([self.rhs <= rhs_ub]), &self.inference_code),
)?;
let var_lb = context.lower_bound(var);
let var_ub = context.upper_bound(var);
if var_lb > max_lb {
max_lb = var_lb;
lb_reason = predicate![var >= var_lb];
}
if var_ub > max_ub {
max_ub = var_ub;
}
}
context.post(
predicate![self.rhs >= max_lb],
(
PropositionalConjunction::from(lb_reason),
&self.inference_code,
),
)?;
if rhs_ub > max_ub {
let ub_reason: PropositionalConjunction = self
.array
.iter()
.map(|var| predicate![var <= max_ub])
.collect();
context.post(
predicate![self.rhs <= max_ub],
(ub_reason, &self.inference_code),
)?;
}
let rhs_lb = context.lower_bound(&self.rhs);
let mut propagating_variable: Option<&ElementVar> = None;
let mut propagation_reason = PropositionalConjunction::default();
for var in self.array.iter() {
if context.upper_bound(var) >= rhs_lb {
if propagating_variable.is_none() {
propagating_variable = Some(var);
} else {
propagating_variable = None;
break;
}
} else {
propagation_reason.push(predicate![var <= rhs_lb - 1]);
}
}
if let Some(propagating_variable) = propagating_variable {
let var_lb = context.lower_bound(propagating_variable);
if var_lb < rhs_lb {
propagation_reason.push(predicate![self.rhs >= rhs_lb]);
context.post(
predicate![propagating_variable >= rhs_lb],
(propagation_reason, &self.inference_code),
)?;
}
}
Ok(())
}
}
#[derive(Clone, Debug)]
pub struct MaximumChecker<ElementVar, Rhs> {
pub array: Box<[ElementVar]>,
pub rhs: Rhs,
}
impl<ElementVar, Rhs, Atomic> InferenceChecker<Atomic> for MaximumChecker<ElementVar, Rhs>
where
Atomic: AtomicConstraint,
ElementVar: CheckerVariable<Atomic>,
Rhs: CheckerVariable<Atomic>,
{
fn check(
&self,
state: pumpkin_checking::VariableState<Atomic>,
_: &[Atomic],
_: Option<&Atomic>,
) -> bool {
let lowest_maximum = self
.array
.iter()
.map(|element| element.induced_lower_bound(&state))
.max()
.unwrap_or(IntExt::NegativeInf);
let highest_maximum = self
.array
.iter()
.map(|element| element.induced_upper_bound(&state))
.max()
.unwrap_or(IntExt::PositiveInf);
lowest_maximum > self.rhs.induced_upper_bound(&state)
|| highest_maximum < self.rhs.induced_lower_bound(&state)
}
}
#[cfg(test)]
mod tests {
use pumpkin_core::predicate;
use pumpkin_core::predicates::Predicate;
use pumpkin_core::predicates::PropositionalConjunction;
use pumpkin_core::propagation::CurrentNogood;
use pumpkin_core::state::State;
use super::*;
use crate::StateExt;
#[test]
fn upper_bound_of_rhs_matches_maximum_upper_bound_of_array_at_initialise() {
let mut state = State::default();
let a = state.new_interval_variable(1, 3, None);
let b = state.new_interval_variable(1, 4, None);
let c = state.new_interval_variable(1, 5, None);
let rhs = state.new_interval_variable(1, 10, None);
let constraint_tag = state.new_constraint_tag();
let _ = state.add_propagator(MaximumArgs {
array: [a, b, c].into(),
rhs,
constraint_tag,
});
state.propagate_to_fixed_point().expect("no empty domain");
state.assert_bounds(rhs, 1, 5);
let mut reason_buffer: Vec<Predicate> = vec![];
let _ = state.get_propagation_reason(
predicate![rhs <= 5],
&mut reason_buffer,
CurrentNogood::empty(),
);
let reason: PropositionalConjunction = reason_buffer.into();
assert_eq!(conjunction!([a <= 5] & [b <= 5] & [c <= 5]), reason);
}
#[test]
fn lower_bound_of_rhs_is_maximum_of_lower_bounds_in_array() {
let mut state = State::default();
let a = state.new_interval_variable(3, 10, None);
let b = state.new_interval_variable(4, 10, None);
let c = state.new_interval_variable(5, 10, None);
let rhs = state.new_interval_variable(1, 10, None);
let constraint_tag = state.new_constraint_tag();
let _ = state.add_propagator(MaximumArgs {
array: [a, b, c].into(),
rhs,
constraint_tag,
});
state.propagate_to_fixed_point().expect("no empty domain");
state.assert_bounds(rhs, 5, 10);
let mut reason_buffer: Vec<Predicate> = vec![];
let _ = state.get_propagation_reason(
predicate![rhs >= 5],
&mut reason_buffer,
CurrentNogood::empty(),
);
let reason: PropositionalConjunction = reason_buffer.into();
assert_eq!(conjunction!([c >= 5]), reason);
}
#[test]
fn upper_bound_of_all_array_elements_at_most_rhs_max_at_initialise() {
let mut state = State::default();
let array = (1..=5)
.map(|idx| state.new_interval_variable(1, 4 + idx, None))
.collect::<Box<_>>();
let rhs = state.new_interval_variable(1, 3, None);
let constraint_tag = state.new_constraint_tag();
let _ = state.add_propagator(MaximumArgs {
array: array.clone(),
rhs,
constraint_tag,
});
state.propagate_to_fixed_point().expect("no empty domain");
for var in array.iter() {
state.assert_bounds(*var, 1, 3);
let mut reason_buffer: Vec<Predicate> = vec![];
let _ = state.get_propagation_reason(
predicate![var <= 3],
&mut reason_buffer,
CurrentNogood::empty(),
);
let reason: PropositionalConjunction = reason_buffer.into();
assert_eq!(conjunction!([rhs <= 3]), reason);
}
}
#[test]
fn single_variable_propagate() {
let mut state = State::default();
let array = (1..=5)
.map(|idx| state.new_interval_variable(1, 1 + 10 * idx, None))
.collect::<Box<_>>();
let rhs = state.new_interval_variable(45, 60, None);
let constraint_tag = state.new_constraint_tag();
let _ = state.add_propagator(MaximumArgs {
array: array.clone(),
rhs,
constraint_tag,
});
state.propagate_to_fixed_point().expect("no empty domain");
state.assert_bounds(*array.last().unwrap(), 45, 51);
state.assert_bounds(rhs, 45, 51);
}
}