use std::cmp::Reverse;
use std::cmp::min;
use pumpkin_core::asserts::pumpkin_assert_simple;
use pumpkin_core::containers::StorageKey;
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::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::state::propagator_conflict;
use pumpkin_core::variables::IntegerVariable;
use super::disjunctive_task::ArgDisjunctiveTask;
use super::disjunctive_task::DisjunctiveTask;
use super::theta_lambda_tree::ThetaLambdaTree;
use crate::disjunctive::checker::DisjunctiveEdgeFindingChecker;
use crate::propagators::disjunctive::DisjunctiveEdgeFinding;
#[derive(Debug, Clone)]
pub struct DisjunctivePropagator<Var: IntegerVariable> {
tasks: Box<[DisjunctiveTask<Var>]>,
sorted_tasks: Vec<DisjunctiveTask<Var>>,
theta_lambda_tree: ThetaLambdaTree<Var>,
inference_code: InferenceCode,
}
#[derive(Debug)]
pub struct DisjunctiveConstructor<Var> {
constraint_tag: ConstraintTag,
tasks: Vec<ArgDisjunctiveTask<Var>>,
}
impl<Var> DisjunctiveConstructor<Var> {
pub fn new(
tasks: impl IntoIterator<Item = ArgDisjunctiveTask<Var>>,
constraint_tag: ConstraintTag,
) -> Self {
Self {
constraint_tag,
tasks: tasks.into_iter().collect(),
}
}
}
impl<Var: IntegerVariable + 'static> PropagatorConstructor for DisjunctiveConstructor<Var> {
type PropagatorImpl = DisjunctivePropagator<Var>;
fn create(self, _: PropagatorConstructorContext) -> PropagatorSpec<Self::PropagatorImpl> {
let tasks = self
.tasks
.into_iter()
.enumerate()
.map(|(index, task)| DisjunctiveTask {
start_time: task.start_time.clone(),
processing_time: task.processing_time,
id: LocalId::from(index as u32),
})
.collect::<Vec<_>>();
let theta_lambda_tree = ThetaLambdaTree::new(&tasks);
let mut registration = EventsToRegister::builder();
for task in tasks.iter() {
registration = registration.add(&task.start_time, DomainEvents::BOUNDS, task.id);
}
let mut checkers = RuntimeCheckers::builder();
let inference_code = checkers.add_inference_checker(
self.constraint_tag,
DisjunctiveEdgeFinding,
DisjunctiveEdgeFindingChecker {
tasks: tasks
.iter()
.map(|task| ArgDisjunctiveTask {
start_time: task.start_time.clone(),
processing_time: task.processing_time,
})
.collect(),
},
);
let propagator = DisjunctivePropagator {
tasks: tasks.clone().into_boxed_slice(),
sorted_tasks: tasks,
theta_lambda_tree,
inference_code,
};
PropagatorSpec {
registration: registration.build(),
checkers: checkers.build(),
propagator,
}
}
}
impl<Var: IntegerVariable + 'static> Propagator for DisjunctivePropagator<Var> {
fn name(&self) -> &str {
"DisjunctiveStrict"
}
fn propagate(&mut self, mut context: PropagationContext) -> PropagationStatusCP {
edge_finding(
&mut self.theta_lambda_tree,
&mut context,
&self.tasks,
&mut self.sorted_tasks,
&self.inference_code,
)
}
fn propagate_from_scratch(&self, mut context: PropagationContext) -> PropagationStatusCP {
let mut sorted_tasks = self.sorted_tasks.clone();
let mut theta_lambda_tree = self.theta_lambda_tree.clone();
edge_finding(
&mut theta_lambda_tree,
&mut context,
&self.tasks,
&mut sorted_tasks,
&self.inference_code,
)
}
}
fn edge_finding<Var: IntegerVariable, SortedTaskVar: IntegerVariable>(
theta_lambda_tree: &mut ThetaLambdaTree<Var>,
context: &mut PropagationContext,
tasks: &[DisjunctiveTask<Var>],
sorted_tasks: &mut [DisjunctiveTask<SortedTaskVar>],
inference_code: &InferenceCode,
) -> PropagationStatusCP {
theta_lambda_tree.update(context.domains());
for task in tasks.iter() {
theta_lambda_tree.add_to_theta(task, context.domains());
}
sorted_tasks
.sort_by_key(|task| Reverse(context.upper_bound(&task.start_time) + task.processing_time));
let mut index = 0;
let mut j = &sorted_tasks[index];
let mut lct_j = context.upper_bound(&j.start_time) + j.processing_time;
while index < tasks.len() - 1 {
if theta_lambda_tree.ect() > lct_j {
return propagator_conflict(
create_conflict_explanation(theta_lambda_tree, context, lct_j),
inference_code,
);
}
theta_lambda_tree.remove_from_theta(j);
theta_lambda_tree.add_to_lambda(j, context.domains());
index += 1;
j = &sorted_tasks[index];
lct_j = context.upper_bound(&j.start_time) + j.processing_time;
while theta_lambda_tree.ect_bar() > lct_j {
if let Some(i) = theta_lambda_tree.responsible_ect_bar() {
let new_bound = theta_lambda_tree.ect();
if new_bound > context.lower_bound(&tasks[i.index()].start_time) {
let propagated_variable = &tasks[i.index()].start_time;
let propagated_predicate = predicate!(propagated_variable >= new_bound);
context.post(
propagated_predicate,
(
create_propagation_explanation(
tasks,
i,
theta_lambda_tree,
context,
new_bound,
lct_j,
),
inference_code,
),
)?;
}
theta_lambda_tree.remove_from_lambda(&tasks[i.index()]);
} else {
break;
}
}
}
Ok(())
}
fn create_conflict_explanation<Var: IntegerVariable>(
theta_lambda_tree: &mut ThetaLambdaTree<Var>,
context: &PropagationContext,
lct: i32,
) -> PropositionalConjunction {
let theta = theta_lambda_tree.get_theta();
pumpkin_assert_simple!(!theta.is_empty());
let mut est = context.lower_bound(&theta[0].start_time);
let mut p_omega = theta_lambda_tree.sum_of_processing_times();
let mut delta = p_omega - (lct - est) - 1;
let mut i = 0;
while i < theta.len() - 1 {
if delta >= 0 {
break;
}
let task = &theta[i];
p_omega -= task.processing_time;
est = context.lower_bound(&theta[i + 1].start_time);
delta = p_omega - (lct - est) - 1;
i += 1;
}
let offset_left = (delta as f64 / 2.0).floor() as i32;
let offset_right = (delta as f64 / 2.0).ceil() as i32;
let mut explanation = Vec::new();
for task in theta.iter().skip(i) {
explanation.push(predicate!(task.start_time >= est - offset_left));
explanation.push(predicate!(
task.start_time <= lct + offset_right - task.processing_time
))
}
explanation.into()
}
fn create_propagation_explanation<'a, Var: IntegerVariable>(
original_tasks: &'a [DisjunctiveTask<Var>],
propagated_task_id: LocalId,
theta_lambda_tree: &mut ThetaLambdaTree<Var>,
context: &'a PropagationContext,
new_bound: i32,
lct_j: i32,
) -> PropositionalConjunction {
let theta = theta_lambda_tree.get_theta();
pumpkin_assert_simple!(!theta.is_empty());
let propagated_task = &original_tasks[propagated_task_id.index()];
let est_propagated = context.lower_bound(&propagated_task.start_time);
let mut p_omega = theta_lambda_tree.sum_of_processing_times();
let mut i = 0;
let mut delta = min(est_propagated, context.lower_bound(&theta[i].start_time))
+ propagated_task.processing_time
+ p_omega;
while i < theta.len() && delta <= lct_j {
p_omega -= theta[i].processing_time;
i += 1;
if i == theta.len() {
break;
}
delta = min(est_propagated, context.lower_bound(&theta[i].start_time))
+ p_omega
+ propagated_task.processing_time;
}
pumpkin_assert_simple!(i < theta.len());
let mut j = i;
let mut p_omega_prime = p_omega;
while j < theta.len() {
if theta_lambda_tree.ect() == context.lower_bound(&theta[j].start_time) + p_omega_prime {
break;
}
p_omega_prime -= theta[j].processing_time;
j += 1;
}
pumpkin_assert_simple!(j < theta.len());
let mut explanation = Vec::new();
let r = min(est_propagated, context.lower_bound(&theta[i].start_time));
for (task_index, task) in theta.iter().enumerate().skip(i) {
if task_index < j {
explanation.push(predicate!(task.start_time >= r));
} else {
explanation.push(predicate!(task.start_time >= new_bound - p_omega_prime));
}
explanation.push(predicate!(
task.start_time
<= r + p_omega + propagated_task.processing_time - 1 - task.processing_time
))
}
explanation.push(predicate!(propagated_task.start_time >= r));
explanation.into()
}
#[cfg(test)]
mod tests {
use pumpkin_core::state::State;
use crate::disjunctive::ArgDisjunctiveTask;
use crate::disjunctive::DisjunctiveConstructor;
#[test]
fn propagator_propagates_lower_bound() {
let mut state = State::default();
let c = state.new_interval_variable(4, 26, None);
let d = state.new_interval_variable(13, 13, None);
let e = state.new_interval_variable(5, 10, None);
let f = state.new_interval_variable(5, 10, None);
let constraint_tag = state.new_constraint_tag();
let _ = state.add_propagator(DisjunctiveConstructor::new(
[
ArgDisjunctiveTask {
start_time: c,
processing_time: 4,
},
ArgDisjunctiveTask {
start_time: d,
processing_time: 5,
},
ArgDisjunctiveTask {
start_time: e,
processing_time: 3,
},
ArgDisjunctiveTask {
start_time: f,
processing_time: 3,
},
],
constraint_tag,
));
state.propagate_to_fixed_point().expect("No conflict");
assert_eq!(state.lower_bound(c), 18);
}
}