use crate::model::variable::VariableId;
use std::collections::{BTreeSet, HashMap};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Domain {
Range { min: i64, max: i64 },
Explicit(BTreeSet<i64>),
}
impl Domain {
pub fn range(min: i64, max: i64) -> Self {
if min > max {
Domain::Range { min: 1, max: 0 } } else {
Domain::Range { min, max }
}
}
pub fn from_values<I: IntoIterator<Item = i64>>(values: I) -> Self {
let set: BTreeSet<i64> = values.into_iter().collect();
Domain::Explicit(set)
}
pub fn is_empty(&self) -> bool {
match self {
Domain::Range { min, max } => min > max,
Domain::Explicit(set) => set.is_empty(),
}
}
pub fn len(&self) -> usize {
match self {
Domain::Range { min, max } => {
if min > max {
0
} else {
let span = (*max as i128) - (*min as i128) + 1;
usize::try_from(span).unwrap_or(usize::MAX)
}
}
Domain::Explicit(set) => set.len(),
}
}
pub fn contains(&self, val: i64) -> bool {
match self {
Domain::Range { min, max } => val >= *min && val <= *max,
Domain::Explicit(set) => set.contains(&val),
}
}
pub fn min(&self) -> Option<i64> {
match self {
Domain::Range { min, max } => {
if min > max {
None
} else {
Some(*min)
}
}
Domain::Explicit(set) => set.iter().next().copied(),
}
}
pub fn max(&self) -> Option<i64> {
match self {
Domain::Range { min, max } => {
if min > max {
None
} else {
Some(*max)
}
}
Domain::Explicit(set) => set.iter().next_back().copied(),
}
}
pub fn remove(&mut self, val: i64) -> bool {
if !self.contains(val) {
return false;
}
match self {
Domain::Range { min, max } => {
if val == *min {
match min.checked_add(1) {
Some(new_min) => *min = new_min,
None => (*min, *max) = (1, 0),
}
true
} else if val == *max {
match max.checked_sub(1) {
Some(new_max) => *max = new_max,
None => (*min, *max) = (1, 0),
}
true
} else {
let set: BTreeSet<i64> = (*min..=*max).filter(|&v| v != val).collect();
*self = Domain::Explicit(set);
true
}
}
Domain::Explicit(set) => set.remove(&val),
}
}
pub fn remove_below(&mut self, min_val: i64) -> bool {
match self {
Domain::Range { min, .. } => {
if *min < min_val {
*min = min_val;
true
} else {
false
}
}
Domain::Explicit(set) => {
let to_remove: Vec<i64> = set.range(..min_val).copied().collect();
if to_remove.is_empty() {
false
} else {
for v in to_remove {
set.remove(&v);
}
true
}
}
}
}
pub fn remove_above(&mut self, max_val: i64) -> bool {
match self {
Domain::Range { min: _, max } => {
if *max > max_val {
*max = max_val;
true
} else {
false
}
}
Domain::Explicit(set) => {
let to_remove: Vec<i64> = match max_val.checked_add(1) {
Some(lower_bound) => set.range(lower_bound..).copied().collect(),
None => Vec::new(),
};
if to_remove.is_empty() {
false
} else {
for v in to_remove {
set.remove(&v);
}
true
}
}
}
}
pub fn assign(&mut self, val: i64) -> bool {
if !self.contains(val) {
*self = Domain::Range { min: 1, max: 0 }; return true;
}
if self.len() == 1 {
return false;
}
*self = Domain::Range { min: val, max: val };
true
}
pub fn values(&self) -> Vec<i64> {
match self {
Domain::Range { min, max } => {
if min > max {
Vec::new()
} else {
(*min..=*max).collect()
}
}
Domain::Explicit(set) => set.iter().copied().collect(),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct TrailedDomains {
domains: HashMap<VariableId, Domain>,
trail: Vec<(VariableId, Domain)>,
}
impl TrailedDomains {
pub fn new(domains: HashMap<VariableId, Domain>) -> Self {
Self {
domains,
trail: Vec::new(),
}
}
pub fn get_mut(&mut self, var: &VariableId) -> Option<&mut Domain> {
if let Some(current) = self.domains.get(var) {
self.trail.push((*var, current.clone()));
}
self.domains.get_mut(var)
}
pub fn mutate(
&mut self,
var: VariableId,
narrow: impl FnOnce(&mut Domain) -> bool,
) -> Option<bool> {
let before = self.domains.get(&var)?.clone();
let changed = narrow(self.domains.get_mut(&var)?);
if changed {
self.trail.push((var, before));
}
Some(changed)
}
pub fn checkpoint(&self) -> usize {
self.trail.len()
}
pub fn undo_to(&mut self, checkpoint: usize) {
while self.trail.len() > checkpoint {
let (var, previous) = self
.trail
.pop()
.expect("trail.len() > checkpoint implies non-empty");
self.domains.insert(var, previous);
}
}
pub fn changed_since(&self, checkpoint: usize) -> impl Iterator<Item = VariableId> + '_ {
let mut seen: Vec<VariableId> = Vec::new();
self.trail[checkpoint..].iter().filter_map(move |(var, _)| {
if seen.contains(var) {
None
} else {
seen.push(*var);
Some(*var)
}
})
}
}
impl std::ops::Deref for TrailedDomains {
type Target = HashMap<VariableId, Domain>;
fn deref(&self) -> &HashMap<VariableId, Domain> {
&self.domains
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_range_domain_basic() {
let mut d = Domain::range(1, 5);
assert_eq!(d.len(), 5);
assert_eq!(d.min(), Some(1));
assert_eq!(d.max(), Some(5));
assert!(d.remove(1));
assert_eq!(d.min(), Some(2));
assert_eq!(d.len(), 4);
assert!(d.remove(3)); assert_eq!(d.values(), vec![2, 4, 5]);
}
#[test]
fn test_explicit_domain_pruning() {
let mut d = Domain::from_values(vec![10, 20, 30, 40]);
assert_eq!(d.len(), 4);
assert!(d.remove_below(20));
assert_eq!(d.values(), vec![20, 30, 40]);
assert!(d.remove_above(30));
assert_eq!(d.values(), vec![20, 30]);
}
#[test]
fn test_len_near_i64_boundaries_does_not_overflow() {
let d = Domain::range(i64::MAX - 9, i64::MAX);
assert_eq!(d.len(), 10);
assert_eq!(d.min(), Some(i64::MAX - 9));
assert_eq!(d.max(), Some(i64::MAX));
let full = Domain::range(i64::MIN, i64::MAX);
assert_eq!(full.len(), usize::MAX);
assert!(!full.is_empty());
}
#[test]
fn test_remove_single_value_domain_at_i64_extremes_does_not_overflow() {
let mut at_max = Domain::range(i64::MAX, i64::MAX);
assert!(at_max.remove(i64::MAX));
assert!(at_max.is_empty());
let mut at_min = Domain::range(i64::MIN, i64::MIN);
assert!(at_min.remove(i64::MIN));
assert!(at_min.is_empty());
}
#[test]
fn test_remove_above_i64_max_on_explicit_domain_does_not_overflow() {
let mut d = Domain::from_values(vec![1, 2, i64::MAX]);
assert!(!d.remove_above(i64::MAX));
assert_eq!(d.values(), vec![1, 2, i64::MAX]);
}
#[test]
fn test_trailed_domains_undo_restores_single_mutation() {
let mut map = HashMap::new();
let v = VariableId(0);
map.insert(v, Domain::range(1, 10));
let mut trailed = TrailedDomains::new(map);
let checkpoint = trailed.checkpoint();
trailed.get_mut(&v).unwrap().remove_above(5);
assert_eq!(trailed.get(&v).unwrap().max(), Some(5));
trailed.undo_to(checkpoint);
assert_eq!(trailed.get(&v).unwrap(), &Domain::range(1, 10));
}
#[test]
fn test_trailed_domains_undo_restores_multiple_mutations_in_order() {
let mut map = HashMap::new();
let x = VariableId(0);
let y = VariableId(1);
map.insert(x, Domain::range(1, 10));
map.insert(y, Domain::range(1, 10));
let mut trailed = TrailedDomains::new(map);
let checkpoint = trailed.checkpoint();
trailed.get_mut(&x).unwrap().remove_above(8);
trailed.get_mut(&y).unwrap().remove_below(3);
trailed.get_mut(&x).unwrap().remove_below(2);
assert_eq!(trailed.get(&x).unwrap(), &Domain::range(2, 8));
assert_eq!(trailed.get(&y).unwrap(), &Domain::range(3, 10));
trailed.undo_to(checkpoint);
assert_eq!(trailed.get(&x).unwrap(), &Domain::range(1, 10));
assert_eq!(trailed.get(&y).unwrap(), &Domain::range(1, 10));
}
#[test]
fn test_trailed_domains_nested_checkpoints() {
let mut map = HashMap::new();
let v = VariableId(0);
map.insert(v, Domain::range(1, 10));
let mut trailed = TrailedDomains::new(map);
let outer = trailed.checkpoint();
trailed.get_mut(&v).unwrap().remove_above(8);
let inner = trailed.checkpoint();
trailed.get_mut(&v).unwrap().remove_above(5);
assert_eq!(trailed.get(&v).unwrap().max(), Some(5));
trailed.undo_to(inner);
assert_eq!(trailed.get(&v).unwrap().max(), Some(8));
trailed.undo_to(outer);
assert_eq!(trailed.get(&v).unwrap(), &Domain::range(1, 10));
}
#[test]
fn test_trailed_domains_deref_read_access() {
let mut map = HashMap::new();
let v = VariableId(0);
map.insert(v, Domain::range(1, 10));
let trailed = TrailedDomains::new(map);
assert!(trailed.contains_key(&v));
assert_eq!(trailed.len(), 1);
assert_eq!(trailed.get(&v), Some(&Domain::range(1, 10)));
}
#[test]
fn test_trailed_domains_mutate_no_op_does_not_grow_trail() {
let mut map = HashMap::new();
let v = VariableId(0);
map.insert(v, Domain::range(1, 10));
let mut trailed = TrailedDomains::new(map);
let checkpoint = trailed.checkpoint();
let changed = trailed.mutate(v, |d| d.remove_above(20));
assert_eq!(changed, Some(false));
assert_eq!(
trailed.checkpoint(),
checkpoint,
"no-op mutation must not grow the trail"
);
assert_eq!(trailed.get(&v).unwrap(), &Domain::range(1, 10));
}
#[test]
fn test_trailed_domains_mutate_confirmed_change_is_undoable() {
let mut map = HashMap::new();
let v = VariableId(0);
map.insert(v, Domain::range(1, 10));
let mut trailed = TrailedDomains::new(map);
let checkpoint = trailed.checkpoint();
let changed = trailed.mutate(v, |d| d.remove_above(5));
assert_eq!(changed, Some(true));
assert_eq!(trailed.get(&v).unwrap().max(), Some(5));
trailed.undo_to(checkpoint);
assert_eq!(trailed.get(&v).unwrap(), &Domain::range(1, 10));
}
#[test]
fn test_changed_since_reports_distinct_touched_variables() {
let mut map = HashMap::new();
let a = VariableId(0);
let b = VariableId(1);
let c = VariableId(2);
map.insert(a, Domain::range(1, 10));
map.insert(b, Domain::range(1, 10));
map.insert(c, Domain::range(1, 10));
let mut trailed = TrailedDomains::new(map);
let checkpoint = trailed.checkpoint();
assert_eq!(trailed.mutate(a, |d| d.remove_above(5)), Some(true));
assert_eq!(trailed.mutate(b, |d| d.remove_above(5)), Some(true));
assert_eq!(trailed.mutate(a, |d| d.remove_above(3)), Some(true));
let mut touched: Vec<VariableId> = trailed.changed_since(checkpoint).collect();
touched.sort_by_key(|v| v.0);
assert_eq!(touched, vec![a, b]);
}
#[test]
fn test_changed_since_empty_when_nothing_mutated_after_checkpoint() {
let mut map = HashMap::new();
let v = VariableId(0);
map.insert(v, Domain::range(1, 10));
let mut trailed = TrailedDomains::new(map);
let _ = trailed.mutate(v, |d| d.remove_above(5));
let checkpoint = trailed.checkpoint();
assert_eq!(trailed.mutate(v, |d| d.remove_above(5)), Some(false));
assert_eq!(trailed.changed_since(checkpoint).count(), 0);
}
}