use crate::cnf::CnfFormula;
use core::cmp::Ordering;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum SatResult {
Sat(Vec<bool>),
Unsat,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct LearnedStep {
pub clause: Vec<i32>,
pub antecedents: Vec<usize>,
}
type ILit = u32;
const NO_REASON: usize = usize::MAX;
const RESTART_BASE: u64 = 64;
const VAR_DECAY: f64 = 0.95;
fn ilit(l: i32) -> ILit {
debug_assert!(l != 0, "DIMACS literals are nonzero");
((l.unsigned_abs() - 1) << 1) | u32::from(l < 0)
}
fn ivar(l: ILit) -> usize {
(l >> 1) as usize
}
fn dimacs(l: ILit) -> i32 {
let v = (l >> 1) as i32 + 1;
if l & 1 == 1 { -v } else { v }
}
fn luby(i: u64) -> u64 {
let (mut size, mut seq, mut x) = (1u64, 0u32, i);
while size < x + 1 {
seq += 1;
size = 2 * size + 1;
}
while size - 1 != x {
size = (size - 1) / 2;
seq -= 1;
x %= size;
}
1 << seq
}
#[derive(Debug, Default)]
struct VarOrder {
heap: Vec<u32>,
pos: Vec<usize>,
}
fn better(act: &[f64], a: u32, b: u32) -> bool {
match act[a as usize].partial_cmp(&act[b as usize]) {
Some(Ordering::Greater) => true,
Some(Ordering::Equal) => a < b,
_ => false,
}
}
impl VarOrder {
fn reset(&mut self, n: usize) {
self.heap = (0..n as u32).collect();
self.pos = (0..n).collect();
}
fn insert(&mut self, act: &[f64], v: u32) {
if self.pos[v as usize] != usize::MAX {
return;
}
self.pos[v as usize] = self.heap.len();
self.heap.push(v);
self.sift_up(act, self.heap.len() - 1);
}
fn update(&mut self, act: &[f64], v: u32) {
let i = self.pos[v as usize];
if i != usize::MAX {
self.sift_up(act, i);
}
}
fn pop(&mut self, act: &[f64]) -> Option<u32> {
let top = *self.heap.first()?;
self.pos[top as usize] = usize::MAX;
let last = self.heap.pop().expect("heap is non-empty");
if !self.heap.is_empty() {
self.heap[0] = last;
self.pos[last as usize] = 0;
self.sift_down(act, 0);
}
Some(top)
}
fn sift_up(&mut self, act: &[f64], mut i: usize) {
while i > 0 {
let parent = (i - 1) / 2;
if !better(act, self.heap[i], self.heap[parent]) {
break;
}
self.swap(i, parent);
i = parent;
}
}
fn sift_down(&mut self, act: &[f64], mut i: usize) {
loop {
let (l, r) = (2 * i + 1, 2 * i + 2);
let mut best = i;
if l < self.heap.len() && better(act, self.heap[l], self.heap[best]) {
best = l;
}
if r < self.heap.len() && better(act, self.heap[r], self.heap[best]) {
best = r;
}
if best == i {
break;
}
self.swap(i, best);
i = best;
}
}
fn swap(&mut self, i: usize, j: usize) {
self.heap.swap(i, j);
self.pos[self.heap[i] as usize] = i;
self.pos[self.heap[j] as usize] = j;
}
}
#[derive(Debug, Default)]
pub struct SatSolver {
clauses: Vec<Vec<ILit>>,
n_orig: usize,
watches: Vec<Vec<usize>>,
assign: Vec<Option<bool>>,
level: Vec<u32>,
reason: Vec<usize>,
trail: Vec<ILit>,
trail_lim: Vec<usize>,
qhead: usize,
activity: Vec<f64>,
var_inc: f64,
order: VarOrder,
saved_phase: Vec<bool>,
seen: Vec<bool>,
mark: Vec<bool>,
trace: Vec<LearnedStep>,
restarts: u64,
conflicts_since_restart: u64,
}
impl SatSolver {
pub fn new() -> Self {
Self::default()
}
pub fn solve(&mut self, formula: &CnfFormula) -> SatResult {
self.run(formula, None, None)
.expect("an unbounded solve always reaches a verdict")
}
pub fn solve_with_budget(
&mut self,
formula: &CnfFormula,
max_conflicts: u64,
) -> Option<SatResult> {
self.run(formula, Some(max_conflicts), None)
}
pub fn solve_with_deadline(
&mut self,
formula: &CnfFormula,
deadline: std::time::Instant,
) -> Option<SatResult> {
self.run(formula, None, Some(deadline))
}
fn run(
&mut self,
formula: &CnfFormula,
budget: Option<u64>,
deadline: Option<std::time::Instant>,
) -> Option<SatResult> {
let n = formula
.clauses
.iter()
.flatten()
.map(|l| l.unsigned_abs())
.max()
.unwrap_or(0)
.max(formula.num_vars) as usize;
self.reset(n);
if !self.load(formula) {
return Some(SatResult::Unsat);
}
if let Some(confl) = self.propagate() {
self.record_root_conflict(confl);
return Some(SatResult::Unsat);
}
self.search(budget, deadline)
}
pub fn proof_trace(&self) -> &[LearnedStep] {
&self.trace
}
fn reset(&mut self, n: usize) {
self.clauses.clear();
self.n_orig = 0;
self.watches = vec![Vec::new(); 2 * n];
self.assign = vec![None; n];
self.level = vec![0; n];
self.reason = vec![NO_REASON; n];
self.trail.clear();
self.trail_lim.clear();
self.qhead = 0;
self.activity = vec![0.0; n];
self.var_inc = 1.0;
self.order.reset(n);
self.saved_phase = vec![false; n];
self.seen = vec![false; n];
self.mark = vec![false; n];
self.trace.clear();
self.restarts = 0;
self.conflicts_since_restart = 0;
}
fn load(&mut self, formula: &CnfFormula) -> bool {
self.n_orig = formula.clauses.len();
for input in &formula.clauses {
let mut lits: Vec<ILit> = input.iter().map(|&l| ilit(l)).collect();
lits.sort_unstable();
lits.dedup();
let tautology = lits.windows(2).any(|w| w[0] ^ 1 == w[1]);
let idx = self.clauses.len();
self.clauses.push(lits);
if tautology {
continue;
}
match self.clauses[idx].len() {
0 => {
self.trace.push(LearnedStep {
clause: Vec::new(),
antecedents: vec![idx],
});
return false;
}
1 => {
let l = self.clauses[idx][0];
match self.lit_value(l) {
Some(true) => {}
Some(false) => {
self.record_root_conflict(idx);
return false;
}
None => self.enqueue(l, idx),
}
}
_ => {
let (w0, w1) = (self.clauses[idx][0], self.clauses[idx][1]);
self.watches[w0 as usize].push(idx);
self.watches[w1 as usize].push(idx);
}
}
}
true
}
fn search(
&mut self,
budget: Option<u64>,
deadline: Option<std::time::Instant>,
) -> Option<SatResult> {
let mut conflicts = 0u64;
loop {
if let Some(confl) = self.propagate() {
if self.trail_lim.is_empty() {
self.record_root_conflict(confl);
return Some(SatResult::Unsat);
}
self.conflicts_since_restart += 1;
conflicts += 1;
if budget.is_some_and(|max| conflicts > max) {
return None;
}
if deadline.is_some_and(|d| std::time::Instant::now() >= d) {
return None;
}
let (learnt, bt_level, antecedents) = self.analyze(confl);
self.trace.push(LearnedStep {
clause: learnt.iter().map(|&l| dimacs(l)).collect(),
antecedents,
});
self.backtrack(bt_level);
let ci = self.clauses.len();
if learnt.len() >= 2 {
self.watches[learnt[0] as usize].push(ci);
self.watches[learnt[1] as usize].push(ci);
}
let asserting = learnt[0];
self.clauses.push(learnt);
self.enqueue(asserting, ci);
self.var_inc /= VAR_DECAY;
} else if self.conflicts_since_restart >= RESTART_BASE * luby(self.restarts) {
self.restarts += 1;
self.conflicts_since_restart = 0;
self.backtrack(0);
} else {
match self.pick_branch() {
None => return Some(SatResult::Sat(self.model())),
Some(l) => {
self.trail_lim.push(self.trail.len());
self.enqueue(l, NO_REASON);
}
}
}
}
}
fn lit_value(&self, l: ILit) -> Option<bool> {
self.assign[ivar(l)].map(|b| b == (l & 1 == 0))
}
fn enqueue(&mut self, l: ILit, reason: usize) {
let v = ivar(l);
debug_assert!(self.assign[v].is_none());
self.assign[v] = Some(l & 1 == 0);
self.level[v] = self.trail_lim.len() as u32;
self.reason[v] = reason;
self.trail.push(l);
}
fn propagate(&mut self) -> Option<usize> {
while self.qhead < self.trail.len() {
let p = self.trail[self.qhead];
self.qhead += 1;
let fl = p ^ 1; let mut ws = core::mem::take(&mut self.watches[fl as usize]);
let mut i = 0;
'clauses: while i < ws.len() {
let ci = ws[i];
if self.clauses[ci][0] == fl {
self.clauses[ci].swap(0, 1);
}
debug_assert_eq!(self.clauses[ci][1], fl);
let first = self.clauses[ci][0];
if self.lit_value(first) == Some(true) {
i += 1;
continue;
}
for k in 2..self.clauses[ci].len() {
let q = self.clauses[ci][k];
if self.lit_value(q) != Some(false) {
self.clauses[ci].swap(1, k);
self.watches[q as usize].push(ci);
ws.swap_remove(i);
continue 'clauses;
}
}
if self.lit_value(first) == Some(false) {
self.watches[fl as usize] = ws;
return Some(ci);
}
self.enqueue(first, ci);
i += 1;
}
self.watches[fl as usize] = ws;
}
None
}
fn analyze(&mut self, confl: usize) -> (Vec<ILit>, usize, Vec<usize>) {
let cur_level = self.trail_lim.len() as u32;
let mut learnt: Vec<ILit> = vec![0]; let mut chain: Vec<usize> = Vec::new(); let mut level0_seeds: Vec<usize> = Vec::new();
let mut to_clear: Vec<usize> = Vec::new();
let mut counter = 0usize; let mut idx = self.trail.len();
let mut c = confl;
let mut skip_first = false;
let uip;
loop {
let mut k = usize::from(skip_first);
while k < self.clauses[c].len() {
let q = self.clauses[c][k];
k += 1;
let v = ivar(q);
if self.seen[v] {
continue;
}
if self.level[v] == 0 {
if !self.mark[v] {
self.mark[v] = true;
level0_seeds.push(v);
}
continue;
}
self.seen[v] = true;
to_clear.push(v);
self.bump(v);
if self.level[v] == cur_level {
counter += 1;
} else {
learnt.push(q);
}
}
loop {
idx -= 1;
if self.seen[ivar(self.trail[idx])] {
break;
}
}
let pivot = self.trail[idx];
counter -= 1;
if counter == 0 {
uip = pivot;
break;
}
c = self.reason[ivar(pivot)];
chain.push(c);
skip_first = true;
}
learnt[0] = uip ^ 1;
for v in to_clear {
self.seen[v] = false;
}
let bt_level = if learnt.len() == 1 {
0
} else {
let mut deepest = 1;
for k in 2..learnt.len() {
if self.level[ivar(learnt[k])] > self.level[ivar(learnt[deepest])] {
deepest = k;
}
}
learnt.swap(1, deepest);
self.level[ivar(learnt[1])] as usize
};
let mut antecedents = self.root_support(&level0_seeds);
chain.reverse();
antecedents.extend(chain);
antecedents.push(confl);
(learnt, bt_level, antecedents)
}
fn root_support(&mut self, seeds: &[usize]) -> Vec<usize> {
if seeds.is_empty() {
return Vec::new();
}
let root_end = self.trail_lim.first().copied().unwrap_or(self.trail.len());
let mut support = Vec::new();
for i in (0..root_end).rev() {
let v = ivar(self.trail[i]);
if self.mark[v] {
self.mark[v] = false;
let r = self.reason[v];
debug_assert_ne!(r, NO_REASON);
support.push(r);
for k in 1..self.clauses[r].len() {
self.mark[ivar(self.clauses[r][k])] = true;
}
}
}
support.reverse();
support
}
fn record_root_conflict(&mut self, confl: usize) {
debug_assert!(self.trail_lim.is_empty());
let mut seeds = Vec::new();
for k in 0..self.clauses[confl].len() {
let v = ivar(self.clauses[confl][k]);
if !self.mark[v] {
self.mark[v] = true;
seeds.push(v);
}
}
let mut antecedents = self.root_support(&seeds);
antecedents.push(confl);
self.trace.push(LearnedStep {
clause: Vec::new(),
antecedents,
});
}
fn backtrack(&mut self, target_level: usize) {
if self.trail_lim.len() <= target_level {
return;
}
let end = self.trail_lim[target_level];
for i in (end..self.trail.len()).rev() {
let l = self.trail[i];
let v = ivar(l);
self.saved_phase[v] = l & 1 == 0;
self.assign[v] = None;
self.order.insert(&self.activity, v as u32);
}
self.trail.truncate(end);
self.trail_lim.truncate(target_level);
self.qhead = end;
}
fn bump(&mut self, v: usize) {
self.activity[v] += self.var_inc;
if self.activity[v] > 1e100 {
for a in &mut self.activity {
*a *= 1e-100;
}
self.var_inc *= 1e-100;
}
self.order.update(&self.activity, v as u32);
}
fn pick_branch(&mut self) -> Option<ILit> {
while let Some(v) = self.order.pop(&self.activity) {
let v = v as usize;
if self.assign[v].is_none() {
return Some(((v as u32) << 1) | u32::from(!self.saved_phase[v]));
}
}
None
}
fn model(&self) -> Vec<bool> {
self.assign
.iter()
.map(|a| a.expect("full assignment at Sat"))
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn formula(num_vars: u32, clauses: &[&[i32]]) -> CnfFormula {
CnfFormula {
num_vars,
clauses: clauses.iter().map(|c| c.to_vec()).collect(),
}
}
fn brute_sat(f: &CnfFormula) -> bool {
let n = f.num_vars as usize;
assert!(n <= 20);
(0u64..1 << n).any(|bits| {
let assignment: Vec<bool> = (0..n).map(|i| bits >> i & 1 == 1).collect();
f.eval(&assignment)
})
}
fn assert_sat_with_model(f: &CnfFormula) -> Vec<bool> {
match SatSolver::new().solve(f) {
SatResult::Sat(model) => {
assert!(model.len() >= f.num_vars as usize);
assert!(f.eval(&model), "returned model must satisfy the formula");
model
}
SatResult::Unsat => panic!("expected Sat"),
}
}
#[test]
fn empty_formula_is_sat() {
assert_sat_with_model(&formula(0, &[]));
let model = assert_sat_with_model(&formula(3, &[]));
assert_eq!(model.len(), 3);
}
#[test]
fn empty_clause_is_unsat() {
let f = formula(2, &[&[1, 2], &[], &[-1]]);
let mut solver = SatSolver::new();
assert_eq!(solver.solve(&f), SatResult::Unsat);
assert_eq!(solver.proof_trace().len(), 1);
assert!(solver.proof_trace()[0].clause.is_empty());
assert_eq!(solver.proof_trace()[0].antecedents, vec![1]);
}
#[test]
fn chained_unit_propagation() {
let mut clauses: Vec<Vec<i32>> = vec![vec![1]];
for v in 1..10 {
clauses.push(vec![-v, v + 1]);
}
let f = CnfFormula {
num_vars: 10,
clauses,
};
let model = assert_sat_with_model(&f);
assert!(model.iter().all(|&b| b), "the chain forces all-true");
let mut clauses = f.clauses.clone();
clauses.push(vec![-10]);
let g = CnfFormula {
num_vars: 10,
clauses,
};
let mut solver = SatSolver::new();
assert_eq!(solver.solve(&g), SatResult::Unsat);
let trace = solver.proof_trace();
assert!(!trace.is_empty());
assert!(trace.last().unwrap().clause.is_empty());
}
#[test]
fn tautology_only_formula_is_sat() {
let f = formula(3, &[&[1, -1], &[2, 3, -3, 2], &[-2, 1, 2, -1]]);
assert_sat_with_model(&f);
}
struct XorShift(u64);
impl XorShift {
fn next(&mut self) -> u64 {
self.0 ^= self.0 << 13;
self.0 ^= self.0 >> 7;
self.0 ^= self.0 << 17;
self.0
}
}
#[test]
fn random_3sat_agrees_with_brute_force() {
let mut rng = XorShift(0x9E3779B97F4A7C15);
let (mut sats, mut unsats) = (0u32, 0u32);
for round in 0..150 {
let n = 8 + (rng.next() % 9) as u32; let ratio = 3.0 + (rng.next() % 21) as f64 / 10.0; let m = (ratio * f64::from(n)).round() as usize;
let mut clauses = Vec::with_capacity(m);
for _ in 0..m {
let mut lits: Vec<i32> = Vec::with_capacity(3);
while lits.len() < 3 {
let v = 1 + (rng.next() % u64::from(n)) as i32;
if lits.iter().any(|l| l.abs() == v) {
continue;
}
lits.push(if rng.next().is_multiple_of(2) { v } else { -v });
}
clauses.push(lits);
}
let f = CnfFormula {
num_vars: n,
clauses,
};
let mut solver = SatSolver::new();
match solver.solve(&f) {
SatResult::Sat(model) => {
assert!(
brute_sat(&f),
"round {round}: solver said Sat, oracle Unsat"
);
assert!(f.eval(&model), "round {round}: model must satisfy");
sats += 1;
}
SatResult::Unsat => {
assert!(
!brute_sat(&f),
"round {round}: solver said Unsat, oracle Sat"
);
let trace = solver.proof_trace();
assert!(!trace.is_empty(), "round {round}: Unsat needs a proof");
assert!(trace.last().unwrap().clause.is_empty());
unsats += 1;
}
}
}
assert!(sats > 0 && unsats > 0, "sats={sats} unsats={unsats}");
}
fn pigeonhole(pigeons: u32, holes: u32) -> CnfFormula {
let var = |p: u32, h: u32| (p * holes + h + 1) as i32;
let mut clauses: Vec<Vec<i32>> = Vec::new();
for p in 0..pigeons {
clauses.push((0..holes).map(|h| var(p, h)).collect());
}
for h in 0..holes {
for p1 in 0..pigeons {
for p2 in p1 + 1..pigeons {
clauses.push(vec![-var(p1, h), -var(p2, h)]);
}
}
}
CnfFormula {
num_vars: pigeons * holes,
clauses,
}
}
#[test]
fn pigeonhole_4_3_is_unsat() {
let f = pigeonhole(4, 3);
assert_eq!(SatSolver::new().solve(&f), SatResult::Unsat);
assert_sat_with_model(&pigeonhole(3, 3));
}
#[test]
fn budget_matches_unbounded_on_trivially_decidable() {
let sat = formula(3, &[&[1], &[-1, 2], &[-2, 3]]);
assert_eq!(
SatSolver::new().solve_with_budget(&sat, 1_000_000),
Some(SatSolver::new().solve(&sat)),
);
let unsat = pigeonhole(4, 3);
assert_eq!(
SatSolver::new().solve_with_budget(&unsat, 1_000_000),
Some(SatResult::Unsat),
);
}
#[test]
fn budget_exhaustion_returns_none() {
let f = pigeonhole(4, 3);
assert_eq!(SatSolver::new().solve_with_budget(&f, 0), None);
}
#[test]
fn budget_zero_still_decides_pure_propagation() {
let mut clauses: Vec<Vec<i32>> = vec![vec![1]];
for v in 1..10 {
clauses.push(vec![-v, v + 1]);
}
let f = CnfFormula {
num_vars: 10,
clauses,
};
match SatSolver::new().solve_with_budget(&f, 0) {
Some(SatResult::Sat(model)) => assert!(model.iter().all(|&b| b)),
other => panic!("expected Sat within a zero budget, got {other:?}"),
}
}
#[test]
fn unsat_proof_trace_derives_the_empty_clause() {
let f = pigeonhole(4, 3);
let mut solver = SatSolver::new();
assert_eq!(solver.solve(&f), SatResult::Unsat);
let trace = solver.proof_trace();
assert!(!trace.is_empty());
assert!(
trace.last().unwrap().clause.is_empty(),
"the final step derives the empty clause"
);
let n_orig = f.clauses.len();
for (k, step) in trace.iter().enumerate() {
assert!(!step.antecedents.is_empty());
for &a in &step.antecedents {
assert!(a < n_orig + k, "step {k} references future clause {a}");
}
}
}
}