use crate::constraint::{Constraint, PropagationResult};
use crate::model::domain::TrailedDomains;
use crate::model::variable::VariableId;
use std::collections::{HashMap, HashSet, VecDeque};
use std::sync::atomic::{AtomicU32, Ordering};
const REGIN_INTERVAL: u32 = 4;
#[derive(Debug)]
pub struct AllDifferent {
scope: Vec<VariableId>,
call_count: AtomicU32,
}
impl Clone for AllDifferent {
fn clone(&self) -> Self {
Self {
scope: self.scope.clone(),
call_count: AtomicU32::new(0),
}
}
}
impl AllDifferent {
pub fn new(variables: impl IntoIterator<Item = VariableId>) -> Self {
Self {
scope: variables.into_iter().collect(),
call_count: AtomicU32::new(0),
}
}
}
fn prune_fixed_values(
domains: &mut TrailedDomains,
scope: &[VariableId],
) -> Result<bool, PropagationResult> {
let mut changed = false;
let mut fixed_values = HashSet::new();
for &var in scope {
if let Some(domain) = domains.get(&var)
&& domain.len() == 1
&& let Some(val) = domain.min()
&& !fixed_values.insert(val)
{
return Err(PropagationResult::Conflict);
}
}
if fixed_values.is_empty() {
return Ok(false);
}
for &var in scope {
if !domains.get(&var).is_some_and(|d| d.len() > 1) {
continue;
}
let did_change = domains
.mutate(var, |domain| {
let mut any = false;
for &val in &fixed_values {
if domain.remove(val) {
any = true;
}
}
any
})
.unwrap_or(false);
if did_change {
changed = true;
}
if domains.get(&var).is_some_and(|d| d.is_empty()) {
return Err(PropagationResult::Conflict);
}
}
Ok(changed)
}
fn try_augment(
var: usize,
var_candidates: &[Vec<usize>],
match_var: &mut [Option<usize>],
match_val: &mut [Option<usize>],
visited: &mut [bool],
) -> bool {
for &val in &var_candidates[var] {
if visited[val] {
continue;
}
visited[val] = true;
match match_val[val] {
None => {
match_val[val] = Some(var);
match_var[var] = Some(val);
return true;
}
Some(other_var) => {
if try_augment(other_var, var_candidates, match_var, match_val, visited) {
match_val[val] = Some(var);
match_var[var] = Some(val);
return true;
}
}
}
}
false
}
fn compute_scc(adj: &[Vec<usize>]) -> Vec<usize> {
const UNVISITED: usize = usize::MAX;
struct State {
counter: usize,
indices: Vec<usize>,
low_links: Vec<usize>,
on_stack: Vec<bool>,
stack: Vec<usize>,
scc_id: Vec<usize>,
next_scc_id: usize,
}
fn strongconnect(node: usize, adj: &[Vec<usize>], state: &mut State) {
state.indices[node] = state.counter;
state.low_links[node] = state.counter;
state.counter += 1;
state.stack.push(node);
state.on_stack[node] = true;
for &next in &adj[node] {
if state.indices[next] == UNVISITED {
strongconnect(next, adj, state);
state.low_links[node] = state.low_links[node].min(state.low_links[next]);
} else if state.on_stack[next] {
state.low_links[node] = state.low_links[node].min(state.indices[next]);
}
}
if state.low_links[node] == state.indices[node] {
let scc_id = state.next_scc_id;
state.next_scc_id += 1;
loop {
let w = state.stack.pop().expect("node's own SCC root is on stack");
state.on_stack[w] = false;
state.scc_id[w] = scc_id;
if w == node {
break;
}
}
}
}
let total = adj.len();
let mut state = State {
counter: 0,
indices: vec![UNVISITED; total],
low_links: vec![0; total],
on_stack: vec![false; total],
stack: Vec::new(),
scc_id: vec![UNVISITED; total],
next_scc_id: 0,
};
for node in 0..total {
if state.indices[node] == UNVISITED {
strongconnect(node, adj, &mut state);
}
}
state.scc_id
}
fn reaches_any(sources: impl IntoIterator<Item = usize>, adj: &[Vec<usize>]) -> Vec<bool> {
let total = adj.len();
let mut rev_adj: Vec<Vec<usize>> = vec![Vec::new(); total];
for (from, tos) in adj.iter().enumerate() {
for &to in tos {
rev_adj[to].push(from);
}
}
let mut reached = vec![false; total];
let mut queue: VecDeque<usize> = VecDeque::new();
for source in sources {
if !reached[source] {
reached[source] = true;
queue.push_back(source);
}
}
while let Some(node) = queue.pop_front() {
for &pred in &rev_adj[node] {
if !reached[pred] {
reached[pred] = true;
queue.push_back(pred);
}
}
}
reached
}
impl Constraint for AllDifferent {
fn name(&self) -> &str {
"AllDifferent"
}
fn scope(&self) -> &[VariableId] {
&self.scope
}
fn is_satisfied(&self, assignment: &HashMap<VariableId, i64>) -> bool {
let mut seen = HashSet::new();
for var in &self.scope {
if let Some(&val) = assignment.get(var)
&& !seen.insert(val)
{
return false; }
}
true
}
fn propagate(&self, domains: &mut TrailedDomains) -> PropagationResult {
let n = self.scope.len();
if n == 0 {
return PropagationResult::Success { changed: false };
}
let mut changed = match prune_fixed_values(domains, &self.scope) {
Ok(changed) => changed,
Err(conflict) => return conflict,
};
let call_index = self.call_count.fetch_add(1, Ordering::Relaxed);
if !call_index.is_multiple_of(REGIN_INTERVAL) {
return PropagationResult::Success { changed };
}
let mut val_index: HashMap<i64, usize> = HashMap::new();
let mut values: Vec<i64> = Vec::new();
let mut var_candidates: Vec<Vec<usize>> = Vec::with_capacity(n);
for &var in &self.scope {
let mut candidates = Vec::new();
if let Some(domain) = domains.get(&var) {
for val in domain.values() {
let idx = *val_index.entry(val).or_insert_with(|| {
values.push(val);
values.len() - 1
});
candidates.push(idx);
}
}
var_candidates.push(candidates);
}
let m = values.len();
let mut match_var: Vec<Option<usize>> = vec![None; n];
let mut match_val: Vec<Option<usize>> = vec![None; m];
let mut visited = vec![false; m];
for var in 0..n {
visited.iter_mut().for_each(|v| *v = false);
if !try_augment(
var,
&var_candidates,
&mut match_var,
&mut match_val,
&mut visited,
) {
return PropagationResult::Conflict;
}
}
let total = n + m;
let mut adj: Vec<Vec<usize>> = vec![Vec::new(); total];
for var in 0..n {
let matched_val = match_var[var].expect("every variable was matched above");
for &val in &var_candidates[var] {
let val_node = n + val;
if val == matched_val {
adj[val_node].push(var);
} else {
adj[var].push(val_node);
}
}
}
let scc_id = compute_scc(&adj);
let free_value_nodes = (0..m)
.filter(|&val| match_val[val].is_none())
.map(|val| n + val);
let reaches_free = reaches_any(free_value_nodes, &adj);
for var in 0..n {
let matched_val = match_var[var].expect("every variable was matched above");
let var_scc = scc_id[var];
let to_remove: Vec<i64> = var_candidates[var]
.iter()
.copied()
.filter(|&val| {
if val == matched_val {
return false;
}
let val_node = n + val;
scc_id[val_node] != var_scc && !reaches_free[val_node]
})
.map(|val| values[val])
.collect();
if to_remove.is_empty() {
continue;
}
let did_change = domains
.mutate(self.scope[var], |domain| {
let mut any = false;
for &val in &to_remove {
if domain.remove(val) {
any = true;
}
}
any
})
.unwrap_or(false);
if did_change {
changed = true;
}
}
PropagationResult::Success { changed }
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::domain::Domain;
use crate::model::variable::Variable;
use crate::propagation::graph::ConstraintGraph;
use crate::solver::{BacktrackingSolver, SolverOptions};
use std::sync::Arc;
#[test]
fn test_propagate_prunes_hall_set_beyond_pairwise_filtering() {
let mut domains = HashMap::new();
let a = VariableId(0);
let b = VariableId(1);
let c = VariableId(2);
domains.insert(a, Domain::range(1, 2));
domains.insert(b, Domain::range(1, 2));
domains.insert(c, Domain::range(1, 3));
let mut trailed = TrailedDomains::new(domains);
let constraint = AllDifferent::new([a, b, c]);
let result = constraint.propagate(&mut trailed);
assert_eq!(result, PropagationResult::Success { changed: true });
assert_eq!(trailed.get(&a).unwrap().values(), vec![1, 2]);
assert_eq!(trailed.get(&b).unwrap().values(), vec![1, 2]);
assert_eq!(
trailed.get(&c).unwrap().values(),
vec![3],
"c can never take 1 or 2: a,b (a Hall set) exhaust both between them"
);
}
#[test]
fn test_propagate_throttles_full_regin_pass_to_every_interval_th_call() {
let a = VariableId(0);
let b = VariableId(1);
let c = VariableId(2);
let constraint = AllDifferent::new([a, b, c]);
let fresh_domains = || {
let mut domains = HashMap::new();
domains.insert(a, Domain::range(1, 2));
domains.insert(b, Domain::range(1, 2));
domains.insert(c, Domain::range(1, 3));
TrailedDomains::new(domains)
};
let mut pruned_c_on_call = Vec::new();
for call in 1..=(REGIN_INTERVAL as usize + 1) {
let mut trailed = fresh_domains();
constraint.propagate(&mut trailed);
if trailed.get(&c).unwrap().values() == vec![3] {
pruned_c_on_call.push(call);
}
}
assert_eq!(
pruned_c_on_call,
vec![1, REGIN_INTERVAL as usize + 1],
"full Régin pass (and thus the Hall-set pruning) should fire on call 1 and call \
REGIN_INTERVAL+1, not the calls in between"
);
}
#[test]
fn test_propagate_keeps_consistent_values() {
let mut domains = HashMap::new();
let x = VariableId(0);
let y = VariableId(1);
let z = VariableId(2);
domains.insert(x, Domain::range(1, 3));
domains.insert(y, Domain::range(1, 3));
domains.insert(z, Domain::range(1, 3));
let mut trailed = TrailedDomains::new(domains);
let constraint = AllDifferent::new([x, y, z]);
let result = constraint.propagate(&mut trailed);
assert_eq!(result, PropagationResult::Success { changed: false });
for &var in &[x, y, z] {
assert_eq!(trailed.get(&var).unwrap().values(), vec![1, 2, 3]);
}
}
#[test]
fn test_propagate_detects_pigeonhole_conflict() {
let mut domains = HashMap::new();
let vars: Vec<VariableId> = (0..4).map(VariableId).collect();
for &v in &vars {
domains.insert(v, Domain::range(1, 3));
}
let mut trailed = TrailedDomains::new(domains);
let constraint = AllDifferent::new(vars);
let result = constraint.propagate(&mut trailed);
assert_eq!(result, PropagationResult::Conflict);
}
#[test]
fn test_solve_simple_all_different() {
let mut graph = ConstraintGraph::new();
let v1 = VariableId(1);
let v2 = VariableId(2);
let v3 = VariableId(3);
graph.add_variable(Variable::new(v1, "x"), Domain::range(1, 3));
graph.add_variable(Variable::new(v2, "y"), Domain::range(1, 3));
graph.add_variable(Variable::new(v3, "z"), Domain::range(1, 3));
graph.add_constraint(Arc::new(AllDifferent::new([v1, v2, v3])));
let graph = graph.finalize().unwrap();
let solver = BacktrackingSolver::new();
let outcome = solver.solve(&graph, &SolverOptions::default());
let solution = outcome.solution.expect("should find a feasible solution");
let mut values: Vec<i64> = solution.assignment.values().copied().collect();
values.sort_unstable();
assert_eq!(values, vec![1, 2, 3]);
}
}