logicaffeine_proof/
hornsat.rs1use std::collections::VecDeque;
12
13#[derive(Clone, Debug, PartialEq, Eq)]
16pub struct HornClause {
17 pub body: Vec<usize>,
19 pub head: Option<usize>,
21}
22
23impl HornClause {
24 pub fn rule(body: impl Into<Vec<usize>>, head: usize) -> Self {
26 HornClause { body: body.into(), head: Some(head) }
27 }
28 pub fn fact(head: usize) -> Self {
30 HornClause { body: Vec::new(), head: Some(head) }
31 }
32 pub fn goal(body: impl Into<Vec<usize>>) -> Self {
34 HornClause { body: body.into(), head: None }
35 }
36}
37
38#[derive(Clone, Debug, PartialEq, Eq)]
40pub enum HornOutcome {
41 Sat(Vec<bool>),
43 Unsat(Vec<usize>),
46}
47
48pub fn solve(clauses: &[HornClause], num_vars: usize) -> HornOutcome {
51 let mut val = vec![false; num_vars];
52 let mut forced_by = vec![usize::MAX; num_vars]; let mut remaining: Vec<usize> = clauses
55 .iter()
56 .map(|c| c.body.iter().filter(|&&b| b < num_vars).count())
57 .collect();
58 let mut in_body: Vec<Vec<usize>> = vec![Vec::new(); num_vars];
59 for (i, c) in clauses.iter().enumerate() {
60 for &b in &c.body {
61 if b < num_vars {
62 in_body[b].push(i);
63 }
64 }
65 }
66 let mut ready: VecDeque<usize> =
67 (0..clauses.len()).filter(|&i| remaining[i] == 0).collect();
68 while let Some(ci) = ready.pop_front() {
69 match clauses[ci].head {
70 None => {
71 return HornOutcome::Unsat(derivation(clauses, ci, &forced_by));
73 }
74 Some(h) => {
75 if h < num_vars && !val[h] {
76 val[h] = true;
77 forced_by[h] = ci;
78 for &cj in &in_body[h] {
79 remaining[cj] -= 1;
80 if remaining[cj] == 0 {
81 ready.push_back(cj);
82 }
83 }
84 }
85 }
86 }
87 }
88 HornOutcome::Sat(val)
89}
90
91fn derivation(clauses: &[HornClause], goal_ci: usize, forced_by: &[usize]) -> Vec<usize> {
94 let mut used = vec![goal_ci];
95 let mut seen_clause: std::collections::HashSet<usize> = std::iter::once(goal_ci).collect();
96 let mut seen_var: std::collections::HashSet<usize> = std::collections::HashSet::new();
97 let mut stack: Vec<usize> = clauses[goal_ci].body.clone();
98 while let Some(v) = stack.pop() {
99 if !seen_var.insert(v) {
100 continue;
101 }
102 let fc = forced_by.get(v).copied().unwrap_or(usize::MAX);
103 if fc != usize::MAX && seen_clause.insert(fc) {
104 used.push(fc);
105 stack.extend(clauses[fc].body.iter().copied());
106 }
107 }
108 used
109}
110
111pub fn satisfies(clauses: &[HornClause], assignment: &[bool]) -> bool {
114 clauses.iter().all(|c| {
115 let body_true = c.body.iter().all(|&b| b < assignment.len() && assignment[b]);
116 match c.head {
117 Some(h) => !body_true || (h < assignment.len() && assignment[h]),
118 None => !body_true,
119 }
120 })
121}
122
123pub fn is_refutation(clauses: &[HornClause], num_vars: usize, refutation: &[usize]) -> bool {
126 let mut val = vec![false; num_vars];
127 loop {
128 let mut changed = false;
129 for &ci in refutation {
130 let Some(c) = clauses.get(ci) else {
131 return false;
132 };
133 if let Some(h) = c.head {
134 let body_true = c.body.iter().all(|&b| b < num_vars && val[b]);
135 if body_true && h < num_vars && !val[h] {
136 val[h] = true;
137 changed = true;
138 }
139 }
140 }
141 if !changed {
142 break;
143 }
144 }
145 refutation.iter().any(|&ci| {
146 clauses
147 .get(ci)
148 .is_some_and(|c| c.head.is_none() && c.body.iter().all(|&b| b < num_vars && val[b]))
149 })
150}
151
152#[cfg(test)]
153mod tests {
154 use super::*;
155
156 #[test]
157 fn facts_and_rules_chain_to_a_least_model() {
158 let cs = vec![
160 HornClause::fact(0),
161 HornClause::fact(1),
162 HornClause::rule([0, 1], 2),
163 HornClause::rule([2], 3),
164 ];
165 match solve(&cs, 4) {
166 HornOutcome::Sat(m) => {
167 assert_eq!(m, vec![true, true, true, true]);
168 assert!(satisfies(&cs, &m));
169 }
170 o => panic!("expected Sat, got {o:?}"),
171 }
172 }
173
174 #[test]
175 fn least_model_leaves_unforced_variables_false() {
176 let cs = vec![HornClause::fact(0), HornClause::rule([1], 0)];
178 match solve(&cs, 2) {
179 HornOutcome::Sat(m) => {
180 assert_eq!(m, vec![true, false]);
181 assert!(satisfies(&cs, &m));
182 }
183 o => panic!("expected Sat, got {o:?}"),
184 }
185 }
186
187 #[test]
188 fn forced_goal_is_refuted_with_a_derivation() {
189 let cs = vec![
191 HornClause::fact(0),
192 HornClause::rule([0], 1),
193 HornClause::goal([0, 1]),
194 ];
195 match solve(&cs, 2) {
196 HornOutcome::Unsat(r) => {
197 assert!(is_refutation(&cs, 2, &r), "refutation must re-check: {r:?}");
198 }
199 o => panic!("expected Unsat, got {o:?}"),
200 }
201 }
202
203 #[test]
204 fn an_unforced_goal_is_satisfiable() {
205 let cs = vec![HornClause::fact(0), HornClause::goal([0, 1])];
207 match solve(&cs, 2) {
208 HornOutcome::Sat(m) => assert!(satisfies(&cs, &m)),
209 o => panic!("expected Sat, got {o:?}"),
210 }
211 }
212
213 #[test]
214 fn matches_brute_force_on_random_horn_systems() {
215 let mut s: u64 = 0xA0761D6478BD642F;
216 let mut next = || {
217 s ^= s << 13;
218 s ^= s >> 7;
219 s ^= s << 17;
220 s
221 };
222 for _ in 0..400 {
223 let num_vars = (next() % 6) as usize + 1;
224 let m = (next() % 8) as usize + 1;
225 let cs: Vec<HornClause> = (0..m)
226 .map(|_| {
227 let body: Vec<usize> = (0..num_vars).filter(|_| next() % 3 == 0).collect();
228 if next() % 4 == 0 {
230 HornClause::goal(body)
231 } else {
232 HornClause::rule(body, (next() as usize) % num_vars)
233 }
234 })
235 .collect();
236 let brute_sat = (0..(1u32 << num_vars)).any(|mask| {
237 let a: Vec<bool> = (0..num_vars).map(|i| (mask >> i) & 1 == 1).collect();
238 satisfies(&cs, &a)
239 });
240 match solve(&cs, num_vars) {
241 HornOutcome::Sat(m) => {
242 assert!(brute_sat, "we said SAT, brute force UNSAT: {cs:?}");
243 assert!(satisfies(&cs, &m), "least model is wrong: {m:?}");
244 }
245 HornOutcome::Unsat(r) => {
246 assert!(!brute_sat, "we said UNSAT, brute force SAT: {cs:?}");
247 assert!(is_refutation(&cs, num_vars, &r), "bogus refutation {r:?}");
248 }
249 }
250 }
251 }
252}