use std::ops::Deref;
use std::ops::DerefMut;
use pumpkin_checking::InferenceChecker;
use super::Domains;
use super::LocalId;
use super::Propagator;
use super::PropagatorId;
use super::PropagatorVarId;
#[cfg(doc)]
use crate::Solver;
use crate::basic_types::PredicateId;
use crate::basic_types::RefOrOwned;
use crate::engine::Assignments;
use crate::engine::State;
use crate::engine::TrailedValues;
use crate::engine::notifications::Watchers;
#[cfg(doc)]
use crate::engine::variables::AffineView;
#[cfg(doc)]
use crate::engine::variables::DomainId;
use crate::predicates::Predicate;
use crate::proof::InferenceCode;
#[cfg(doc)]
use crate::propagation::DomainEvent;
use crate::propagation::DomainEvents;
use crate::propagators::reified_propagator::ReifiedChecker;
use crate::variables::IntegerVariable;
use crate::variables::Literal;
pub trait PropagatorConstructor {
type PropagatorImpl: Propagator + Clone;
fn add_inference_checkers(&self, _checkers: InferenceCheckers<'_>) {}
fn create(self, context: PropagatorConstructorContext) -> Self::PropagatorImpl;
}
#[derive(Debug)]
pub struct InferenceCheckers<'state> {
state: &'state mut State,
reification_literal: Option<Literal>,
}
impl<'state> InferenceCheckers<'state> {
#[cfg(feature = "check-propagations")]
pub(crate) fn new(state: &'state mut State) -> Self {
InferenceCheckers {
state,
reification_literal: None,
}
}
}
impl InferenceCheckers<'_> {
pub fn add_inference_checker(
&mut self,
inference_code: InferenceCode,
checker: Box<dyn InferenceChecker<Predicate>>,
) {
if let Some(reification_literal) = self.reification_literal {
let reification_checker = ReifiedChecker {
inner: checker.into(),
reification_literal,
};
self.state
.add_inference_checker(inference_code, Box::new(reification_checker));
} else {
self.state.add_inference_checker(inference_code, checker);
}
}
pub fn with_reification_literal(&mut self, literal: Literal) {
self.reification_literal = Some(literal)
}
}
#[derive(Debug)]
pub struct PropagatorConstructorContext<'a> {
state: &'a mut State,
pub(crate) propagator_id: PropagatorId,
next_local_id: RefOrOwned<'a, LocalId>,
did_register: RefOrOwned<'a, bool>,
}
impl PropagatorConstructorContext<'_> {
pub(crate) fn new<'a>(
propagator_id: PropagatorId,
state: &'a mut State,
) -> PropagatorConstructorContext<'a> {
PropagatorConstructorContext {
next_local_id: RefOrOwned::Owned(LocalId::from(0)),
propagator_id,
state,
did_register: RefOrOwned::Owned(false),
}
}
pub fn will_not_register_any_events(&mut self) {
*self.did_register = true;
}
pub fn domains(&mut self) -> Domains<'_> {
Domains::new(&self.state.assignments, &mut self.state.trailed_values)
}
pub fn register(
&mut self,
var: impl IntegerVariable,
domain_events: DomainEvents,
local_id: LocalId,
) {
self.will_not_register_any_events();
let propagator_var = PropagatorVarId {
propagator: self.propagator_id,
variable: local_id,
};
self.update_next_local_id(local_id);
let mut watchers = Watchers::new(propagator_var, &mut self.state.notification_engine);
var.watch_all(&mut watchers, domain_events.events());
}
pub fn register_predicate(&mut self, predicate: Predicate) -> PredicateId {
self.will_not_register_any_events();
self.state.notification_engine.watch_predicate(
predicate,
self.propagator_id,
&mut self.state.trailed_values,
&self.state.assignments,
)
}
pub fn register_backtrack<Var: IntegerVariable>(
&mut self,
var: Var,
domain_events: DomainEvents,
local_id: LocalId,
) {
let propagator_var = PropagatorVarId {
propagator: self.propagator_id,
variable: local_id,
};
self.update_next_local_id(local_id);
let mut watchers = Watchers::new(propagator_var, &mut self.state.notification_engine);
var.watch_all_backtrack(&mut watchers, domain_events.events());
}
pub(crate) fn get_next_local_id(&self) -> LocalId {
*self.next_local_id.deref()
}
pub fn reborrow(&mut self) -> PropagatorConstructorContext<'_> {
PropagatorConstructorContext {
propagator_id: self.propagator_id,
next_local_id: self.next_local_id.reborrow(),
did_register: self.did_register.reborrow(),
state: self.state,
}
}
pub fn add_inference_checker(
&mut self,
inference_code: InferenceCode,
checker: Box<dyn InferenceChecker<Predicate>>,
) {
self.state.add_inference_checker(inference_code, checker);
}
fn update_next_local_id(&mut self, local_id: LocalId) {
let next_local_id = (*self.next_local_id.deref()).max(LocalId::from(local_id.unpack() + 1));
*self.next_local_id.deref_mut() = next_local_id;
}
}
impl Drop for PropagatorConstructorContext<'_> {
fn drop(&mut self) {
if std::thread::panicking() {
return;
}
let did_register = match self.did_register {
RefOrOwned::Ref(_) => return,
RefOrOwned::Owned(did_register) => did_register,
};
if !did_register {
panic!(
"Propagator did not register to be enqueued. If this is intentional, call PropagatorConstructorContext::will_not_register_any_events()."
);
}
}
}
mod private {
use super::*;
use crate::propagation::HasAssignments;
impl HasAssignments for PropagatorConstructorContext<'_> {
fn assignments(&self) -> &Assignments {
&self.state.assignments
}
fn trailed_values(&self) -> &TrailedValues {
&self.state.trailed_values
}
fn trailed_values_mut(&mut self) -> &mut TrailedValues {
&mut self.state.trailed_values
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::variables::DomainId;
#[test]
#[should_panic]
fn panic_when_no_registration_happened() {
let mut state = State::default();
state.notification_engine.grow();
let _c1 = PropagatorConstructorContext::new(PropagatorId(0), &mut state);
}
#[test]
fn do_not_panic_if_told_no_registration_will_happen() {
let mut state = State::default();
state.notification_engine.grow();
let mut ctx = PropagatorConstructorContext::new(PropagatorId(0), &mut state);
ctx.will_not_register_any_events();
}
#[test]
fn do_not_panic_if_no_registration_happens_in_reborrowed() {
let mut state = State::default();
state.notification_engine.grow();
let mut ctx = PropagatorConstructorContext::new(PropagatorId(0), &mut state);
let ctx2 = ctx.reborrow();
drop(ctx2);
ctx.will_not_register_any_events();
}
#[test]
fn reborrowing_remembers_next_local_id() {
let mut state = State::default();
state.notification_engine.grow();
let mut c1 = PropagatorConstructorContext::new(PropagatorId(0), &mut state);
c1.will_not_register_any_events();
let mut c2 = c1.reborrow();
c2.register(DomainId::new(0), DomainEvents::ANY_INT, LocalId::from(1));
drop(c2);
assert_eq!(LocalId::from(2), c1.get_next_local_id());
}
}