use pumpkin_checking::AtomicConstraint;
use pumpkin_checking::CheckerVariable;
use pumpkin_checking::InferenceChecker;
use pumpkin_checking::IntExt;
use pumpkin_core::asserts::pumpkin_assert_simple;
use pumpkin_core::conjunction;
use pumpkin_core::declare_inference_label;
use pumpkin_core::predicate;
use pumpkin_core::proof::ConstraintTag;
use pumpkin_core::proof::InferenceCode;
use pumpkin_core::propagation::DomainEvents;
use pumpkin_core::propagation::InferenceCheckers;
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::ReadDomains;
use pumpkin_core::state::PropagationStatusCP;
use pumpkin_core::variables::IntegerVariable;
#[derive(Clone, Debug)]
pub struct DivisionArgs<VA, VB, VC> {
pub numerator: VA,
pub denominator: VB,
pub rhs: VC,
pub constraint_tag: ConstraintTag,
}
const ID_NUMERATOR: LocalId = LocalId::from(0);
const ID_DENOMINATOR: LocalId = LocalId::from(1);
const ID_RHS: LocalId = LocalId::from(2);
declare_inference_label!(Division);
impl<VA, VB, VC> PropagatorConstructor for DivisionArgs<VA, VB, VC>
where
VA: IntegerVariable + 'static,
VB: IntegerVariable + 'static,
VC: IntegerVariable + 'static,
{
type PropagatorImpl = DivisionPropagator<VA, VB, VC>;
fn create(self, mut context: PropagatorConstructorContext) -> Self::PropagatorImpl {
let DivisionArgs {
numerator,
denominator,
rhs,
constraint_tag,
} = self;
pumpkin_assert_simple!(
!context.contains(&denominator, 0),
"Denominator cannot contain 0"
);
context.register(numerator.clone(), DomainEvents::BOUNDS, ID_NUMERATOR);
context.register(denominator.clone(), DomainEvents::BOUNDS, ID_DENOMINATOR);
context.register(rhs.clone(), DomainEvents::BOUNDS, ID_RHS);
let inference_code = InferenceCode::new(constraint_tag, Division);
DivisionPropagator {
numerator,
denominator,
rhs,
inference_code,
}
}
fn add_inference_checkers(&self, mut checkers: InferenceCheckers<'_>) {
checkers.add_inference_checker(
InferenceCode::new(self.constraint_tag, Division),
Box::new(IntegerDivisionChecker {
numerator: self.numerator.clone(),
denominator: self.denominator.clone(),
rhs: self.rhs.clone(),
}),
);
}
}
#[derive(Clone, Debug)]
pub struct DivisionPropagator<VA, VB, VC> {
numerator: VA,
denominator: VB,
rhs: VC,
inference_code: InferenceCode,
}
impl<VA: 'static, VB: 'static, VC: 'static> Propagator for DivisionPropagator<VA, VB, VC>
where
VA: IntegerVariable,
VB: IntegerVariable,
VC: IntegerVariable,
{
fn priority(&self) -> Priority {
Priority::High
}
fn name(&self) -> &str {
"Division"
}
fn propagate_from_scratch(&self, context: PropagationContext) -> PropagationStatusCP {
perform_propagation(
context,
&self.numerator,
&self.denominator,
&self.rhs,
&self.inference_code,
)
}
}
fn perform_propagation<VA: IntegerVariable, VB: IntegerVariable, VC: IntegerVariable>(
mut context: PropagationContext,
numerator: &VA,
denominator: &VB,
rhs: &VC,
inference_code: &InferenceCode,
) -> PropagationStatusCP {
if context.lower_bound(denominator) < 0 && context.upper_bound(denominator) > 0 {
return Ok(());
}
let mut negated_numerator = &numerator.scaled(-1);
let mut numerator = &numerator.scaled(1);
let mut negated_denominator = &denominator.scaled(-1);
let mut denominator = &denominator.scaled(1);
if context.upper_bound(denominator) < 0 {
std::mem::swap(&mut numerator, &mut negated_numerator);
std::mem::swap(&mut denominator, &mut negated_denominator);
}
let negated_rhs = &rhs.scaled(-1);
propagate_signs(&mut context, numerator, denominator, rhs, inference_code)?;
if context.upper_bound(numerator) >= 0 && context.upper_bound(rhs) >= 0 {
propagate_upper_bounds(&mut context, numerator, denominator, rhs, inference_code)?;
}
if context.upper_bound(negated_numerator) >= 0 && context.upper_bound(negated_rhs) >= 0 {
propagate_upper_bounds(
&mut context,
negated_numerator,
denominator,
negated_rhs,
inference_code,
)?;
}
if context.lower_bound(numerator) >= 0 && context.lower_bound(rhs) >= 0 {
propagate_positive_domains(&mut context, numerator, denominator, rhs, inference_code)?;
}
if context.lower_bound(negated_numerator) >= 0 && context.lower_bound(negated_rhs) >= 0 {
propagate_positive_domains(
&mut context,
negated_numerator,
denominator,
negated_rhs,
inference_code,
)?;
}
Ok(())
}
fn propagate_positive_domains<VA: IntegerVariable, VB: IntegerVariable, VC: IntegerVariable>(
context: &mut PropagationContext,
numerator: &VA,
denominator: &VB,
rhs: &VC,
inference_code: &InferenceCode,
) -> PropagationStatusCP {
let rhs_min = context.lower_bound(rhs);
let rhs_max = context.upper_bound(rhs);
let numerator_min = context.lower_bound(numerator);
let numerator_max = context.upper_bound(numerator);
let denominator_min = context.lower_bound(denominator);
let denominator_max = context.upper_bound(denominator);
let new_min_rhs = numerator_min / denominator_max;
if rhs_min < new_min_rhs {
context.post(
predicate![rhs >= new_min_rhs],
(
conjunction!(
[numerator >= numerator_min]
& [denominator <= denominator_max]
& [denominator >= 1]
),
inference_code,
),
)?;
}
let new_min_numerator = denominator_min * rhs_min;
if numerator_min < new_min_numerator {
context.post(
predicate![numerator >= new_min_numerator],
(
conjunction!([denominator >= denominator_min] & [rhs >= rhs_min]),
inference_code,
),
)?;
}
if rhs_min > 0 {
let new_max_denominator = numerator_max / rhs_min;
if denominator_max > new_max_denominator {
context.post(
predicate![denominator <= new_max_denominator],
(
conjunction!(
[numerator <= numerator_max]
& [numerator >= 0]
& [rhs >= rhs_min]
& [denominator >= 1]
),
inference_code,
),
)?;
}
}
let new_min_denominator = {
let dividend = numerator_min + 1;
let positive_divisor = rhs_max + 1;
let result = dividend / positive_divisor;
let adjust = result * positive_divisor < dividend;
result + adjust as i32
};
if denominator_min < new_min_denominator {
context.post(
predicate![denominator >= new_min_denominator],
(
conjunction!(
[numerator >= numerator_min]
& [rhs <= rhs_max]
& [rhs >= 0]
& [denominator >= 1]
),
inference_code,
),
)?;
}
Ok(())
}
fn propagate_upper_bounds<VA: IntegerVariable, VB: IntegerVariable, VC: IntegerVariable>(
context: &mut PropagationContext,
numerator: &VA,
denominator: &VB,
rhs: &VC,
inference_code: &InferenceCode,
) -> PropagationStatusCP {
let rhs_max = context.upper_bound(rhs);
let numerator_max = context.upper_bound(numerator);
let denominator_min = context.lower_bound(denominator);
let denominator_max = context.upper_bound(denominator);
let new_max_rhs = numerator_max / denominator_min;
if rhs_max > new_max_rhs {
context.post(
predicate![rhs <= new_max_rhs],
(
conjunction!([numerator <= numerator_max] & [denominator >= denominator_min]),
inference_code,
),
)?;
}
let new_max_numerator = (rhs_max + 1) * denominator_max - 1;
if numerator_max > new_max_numerator {
context.post(
predicate![numerator <= new_max_numerator],
(
conjunction!(
[denominator <= denominator_max] & [denominator >= 1] & [rhs <= rhs_max]
),
inference_code,
),
)?;
}
Ok(())
}
fn propagate_signs<VA: IntegerVariable, VB: IntegerVariable, VC: IntegerVariable>(
context: &mut PropagationContext,
numerator: &VA,
denominator: &VB,
rhs: &VC,
inference_code: &InferenceCode,
) -> PropagationStatusCP {
let rhs_min = context.lower_bound(rhs);
let rhs_max = context.upper_bound(rhs);
let numerator_min = context.lower_bound(numerator);
let numerator_max = context.upper_bound(numerator);
if numerator_min >= 0 && rhs_min < 0 {
context.post(
predicate![rhs >= 0],
(
conjunction!([numerator >= 0] & [denominator >= 1]),
inference_code,
),
)?;
}
if numerator_min <= 0 && rhs_min > 0 {
context.post(
predicate![numerator >= 1],
(
conjunction!([rhs >= 1] & [denominator >= 1]),
inference_code,
),
)?;
}
if numerator_max <= 0 && rhs_max > 0 {
context.post(
predicate![rhs <= 0],
(
conjunction!([numerator <= 0] & [denominator >= 1]),
inference_code,
),
)?;
}
if numerator_max >= 0 && rhs_max < 0 {
context.post(
predicate![numerator <= -1],
(
conjunction!([rhs <= -1] & [denominator >= 1]),
inference_code,
),
)?;
}
Ok(())
}
#[derive(Clone, Debug)]
pub struct IntegerDivisionChecker<VA, VB, VC> {
pub numerator: VA,
pub denominator: VB,
pub rhs: VC,
}
impl<VA, VB, VC, Atomic> InferenceChecker<Atomic> for IntegerDivisionChecker<VA, VB, VC>
where
Atomic: AtomicConstraint,
VA: CheckerVariable<Atomic>,
VB: CheckerVariable<Atomic>,
VC: CheckerVariable<Atomic>,
{
fn check(
&self,
state: pumpkin_checking::VariableState<Atomic>,
_premises: &[Atomic],
_consequent: Option<&Atomic>,
) -> bool {
let x1 = self.numerator.induced_lower_bound(&state);
let x2 = self.numerator.induced_upper_bound(&state);
let y1 = self.denominator.induced_lower_bound(&state);
let y2 = self.denominator.induced_upper_bound(&state);
assert!(
y2 < 0 || y1 > 0,
"Currentl, the checker does not contain inferences where the denominator spans 0"
);
let computed_c_lower: IntExt = *[
x1.div_ceil(y1),
x1.div_ceil(y2),
x2.div_ceil(y1),
x2.div_ceil(y2),
]
.iter()
.flatten()
.min()
.expect("Expected at least one element to be defined");
let computed_c_upper: IntExt = *[
x1.div_floor(y1),
x1.div_floor(y2),
x2.div_floor(y1),
x2.div_floor(y2),
]
.iter()
.flatten()
.min()
.expect("Expected at least one element to be defined");
let c_lower = self.rhs.induced_lower_bound(&state);
let c_upper = self.rhs.induced_upper_bound(&state);
computed_c_upper < c_lower || computed_c_lower > c_upper
}
}
#[cfg(test)]
mod tests {
use pumpkin_core::state::State;
use super::*;
#[test]
fn detects_conflicts() {
let mut state = State::default();
let numerator = state.new_interval_variable(1, 1, None);
let denominator = state.new_interval_variable(2, 2, None);
let rhs = state.new_interval_variable(2, 2, None);
let constraint_tag = state.new_constraint_tag();
let _ = state.add_propagator(DivisionArgs {
numerator,
denominator,
rhs,
constraint_tag,
});
let _ = state.propagate_to_fixed_point().unwrap_err();
}
}