use std::ops::Deref;
use crate::conflict_resolving::ConflictAnalysisContext;
use crate::predicates::Predicate;
#[derive(Clone, Debug)]
pub(crate) struct LearnedNogood {
pub(crate) predicates: Vec<Predicate>,
pub(crate) backtrack_level: usize,
}
impl Deref for LearnedNogood {
type Target = [Predicate];
fn deref(&self) -> &Self::Target {
&self.predicates
}
}
impl LearnedNogood {
pub(crate) fn create_from_vec(
mut clean_nogood: Vec<Predicate>,
context: &ConflictAnalysisContext,
uses_cpip: bool,
) -> Self {
if clean_nogood.is_empty() {
return Self {
predicates: vec![],
backtrack_level: 0,
};
}
let propagating_domain = if uses_cpip {
Some(
clean_nogood
.iter()
.find(|predicate| {
context
.state
.get_checkpoint_for_predicate(**predicate)
.is_some_and(|checkpoint| checkpoint == context.state.get_checkpoint())
})
.copied()
.expect("Expected at least one element to be from the current decision level"),
)
} else {
None
};
let mut index = 1;
let mut highest_level_below_current = 0;
while index < clean_nogood.len() {
let predicate = clean_nogood[index];
let dl = context
.state
.get_checkpoint_for_predicate(predicate)
.unwrap();
if dl == context.state.get_checkpoint() {
clean_nogood.swap(0, index);
if propagating_domain.is_none()
|| clean_nogood[0].get_domain() != clean_nogood[index].get_domain()
{
index -= 1;
}
} else if dl > highest_level_below_current
&& (propagating_domain.is_none()
|| predicate.get_domain() != propagating_domain.unwrap().get_domain())
{
highest_level_below_current = dl;
clean_nogood.swap(1, index);
}
index += 1;
}
let mut backjump_level = if clean_nogood.len() > 1 {
context
.state
.get_checkpoint_for_predicate(clean_nogood[1])
.unwrap()
} else {
0
};
if clean_nogood.len() > 1 && propagating_domain.is_some() {
let propagating_domain = clean_nogood[0].get_domain();
backjump_level = clean_nogood
.iter()
.filter(|predicate| predicate.get_domain() != propagating_domain)
.map(|predicate| {
context
.state
.get_checkpoint_for_predicate(*predicate)
.unwrap()
})
.max()
.unwrap_or(0);
}
Self {
predicates: clean_nogood,
backtrack_level: backjump_level,
}
}
}