1use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet};
4
5use xlog_core::{Result, XlogError};
6
7use crate::logsumexp::{
8 circuit_logsumexp, validate_circuit_log_weight_pair, validate_circuit_value,
9};
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
12pub enum DdnnfNodeKind {
13 Or,
14 And,
15 True,
16 False,
17}
18
19#[derive(Debug, Clone)]
20pub struct DdnnfNode {
21 pub kind: DdnnfNodeKind,
22}
23
24#[derive(Debug, Clone)]
25pub struct DdnnfEdge {
26 pub from: u32,
27 pub to: u32,
28 pub lits: Vec<i32>,
29}
30
31#[derive(Debug, Clone)]
32pub struct DecisionDnnf {
33 root: u32,
34 nodes: BTreeMap<u32, DdnnfNode>,
35 edges: Vec<DdnnfEdge>,
36 outgoing: BTreeMap<u32, Vec<usize>>,
37 max_var: u32,
38}
39
40impl DecisionDnnf {
41 pub fn root(&self) -> u32 {
42 self.root
43 }
44
45 pub fn max_var(&self) -> u32 {
46 self.max_var
47 }
48
49 pub fn node_kind(&self, node_id: u32) -> Option<DdnnfNodeKind> {
50 self.nodes.get(&node_id).map(|n| n.kind)
51 }
52
53 pub fn outgoing_edge_indices(&self, node_id: u32) -> Option<&[usize]> {
54 self.outgoing.get(&node_id).map(|v| v.as_slice())
55 }
56
57 pub fn edge(&self, edge_idx: usize) -> Option<&DdnnfEdge> {
58 self.edges.get(edge_idx)
59 }
60
61 pub fn parse_str(input: &str) -> Result<Self> {
62 let mut nodes: BTreeMap<u32, DdnnfNode> = BTreeMap::new();
63 let mut edges: Vec<DdnnfEdge> = Vec::new();
64 let mut targets: HashSet<u32> = HashSet::new();
65 let mut max_var: u32 = 0;
66
67 for (line_no, raw_line) in input.lines().enumerate() {
68 let line = raw_line.trim();
69 if line.is_empty() {
70 continue;
71 }
72
73 let mut tokens: Vec<&str> = line.split_whitespace().collect();
74 if tokens.is_empty() {
75 continue;
76 }
77
78 if tokens.last() != Some(&"0") {
79 return Err(XlogError::Compilation(format!(
80 "Decision-DNNF parse error at line {}: missing 0 terminator",
81 line_no + 1
82 )));
83 }
84 tokens.pop();
85 if tokens.is_empty() {
86 return Err(XlogError::Compilation(format!(
87 "Decision-DNNF parse error at line {}: empty record before terminator",
88 line_no + 1
89 )));
90 }
91
92 match tokens[0] {
93 "o" | "a" | "t" | "f" => {
94 if tokens.len() < 2 {
95 return Err(XlogError::Compilation(format!(
96 "Decision-DNNF parse error at line {}: node record missing id",
97 line_no + 1
98 )));
99 }
100 let id: u32 = tokens[1].parse().map_err(|_| {
101 XlogError::Compilation(format!(
102 "Decision-DNNF parse error at line {}: invalid node id '{}'",
103 line_no + 1,
104 tokens[1]
105 ))
106 })?;
107
108 let kind = match tokens[0] {
109 "o" => DdnnfNodeKind::Or,
110 "a" => DdnnfNodeKind::And,
111 "t" => DdnnfNodeKind::True,
112 "f" => DdnnfNodeKind::False,
113 _ => unreachable!(),
114 };
115
116 if nodes.insert(id, DdnnfNode { kind }).is_some() {
117 return Err(XlogError::Compilation(format!(
118 "Decision-DNNF parse error at line {}: duplicate node id {}",
119 line_no + 1,
120 id
121 )));
122 }
123 }
124 _ => {
125 if tokens.len() < 2 {
126 return Err(XlogError::Compilation(format!(
127 "Decision-DNNF parse error at line {}: edge record missing dst",
128 line_no + 1
129 )));
130 }
131 let from: u32 = tokens[0].parse().map_err(|_| {
132 XlogError::Compilation(format!(
133 "Decision-DNNF parse error at line {}: invalid edge src '{}'",
134 line_no + 1,
135 tokens[0]
136 ))
137 })?;
138 let to: u32 = tokens[1].parse().map_err(|_| {
139 XlogError::Compilation(format!(
140 "Decision-DNNF parse error at line {}: invalid edge dst '{}'",
141 line_no + 1,
142 tokens[1]
143 ))
144 })?;
145
146 let mut lits: Vec<i32> = Vec::new();
147 for &tok in &tokens[2..] {
148 let lit: i32 = tok.parse().map_err(|_| {
149 XlogError::Compilation(format!(
150 "Decision-DNNF parse error at line {}: invalid literal '{}'",
151 line_no + 1,
152 tok
153 ))
154 })?;
155 if lit == 0 {
156 return Err(XlogError::Compilation(format!(
157 "Decision-DNNF parse error at line {}: literal cannot be 0",
158 line_no + 1
159 )));
160 }
161 max_var = max_var.max(lit.unsigned_abs());
162 lits.push(lit);
163 }
164
165 let edge_id = edges.len();
166 edges.push(DdnnfEdge { from, to, lits });
167 targets.insert(to);
168
169 let _ = edge_id;
171 }
172 }
173 }
174
175 if nodes.is_empty() {
176 return Err(XlogError::Compilation(
177 "Decision-DNNF parse error: no nodes found".to_string(),
178 ));
179 }
180
181 for edge in &edges {
182 let from_kind = nodes.get(&edge.from).ok_or_else(|| {
183 XlogError::Compilation(format!(
184 "Decision-DNNF parse error: edge references unknown src node {}",
185 edge.from
186 ))
187 })?;
188 let _to_kind = nodes.get(&edge.to).ok_or_else(|| {
189 XlogError::Compilation(format!(
190 "Decision-DNNF parse error: edge references unknown dst node {}",
191 edge.to
192 ))
193 })?;
194
195 match from_kind.kind {
196 DdnnfNodeKind::Or | DdnnfNodeKind::And => {}
197 DdnnfNodeKind::True | DdnnfNodeKind::False => {
198 return Err(XlogError::Compilation(format!(
199 "Decision-DNNF parse error: leaf node {} cannot have outgoing edges",
200 edge.from
201 )));
202 }
203 }
204 }
205
206 let declared: BTreeSet<u32> = nodes.keys().copied().collect();
207 let target_set: BTreeSet<u32> = targets.into_iter().collect();
208 let roots: Vec<u32> = declared.difference(&target_set).copied().collect();
209 let root = match roots.as_slice() {
210 [only] => *only,
211 [] => {
212 return Err(XlogError::Compilation(
213 "Decision-DNNF parse error: could not infer root (no root candidates)"
214 .to_string(),
215 ))
216 }
217 many => {
218 return Err(XlogError::Compilation(format!(
219 "Decision-DNNF parse error: could not infer unique root (candidates: {:?})",
220 many
221 )))
222 }
223 };
224
225 let mut outgoing: BTreeMap<u32, Vec<usize>> = BTreeMap::new();
226 for (idx, edge) in edges.iter().enumerate() {
227 outgoing.entry(edge.from).or_default().push(idx);
228 }
229
230 Self::check_acyclic(root, &nodes, &edges, &outgoing)?;
232
233 Ok(Self {
234 root,
235 nodes,
236 edges,
237 outgoing,
238 max_var,
239 })
240 }
241
242 fn check_acyclic(
243 root: u32,
244 nodes: &BTreeMap<u32, DdnnfNode>,
245 edges: &[DdnnfEdge],
246 outgoing: &BTreeMap<u32, Vec<usize>>,
247 ) -> Result<()> {
248 let mut visiting: HashSet<u32> = HashSet::new();
249 let mut visited: HashSet<u32> = HashSet::new();
250
251 fn dfs(
252 node_id: u32,
253 nodes: &BTreeMap<u32, DdnnfNode>,
254 edges: &[DdnnfEdge],
255 outgoing: &BTreeMap<u32, Vec<usize>>,
256 visiting: &mut HashSet<u32>,
257 visited: &mut HashSet<u32>,
258 ) -> Result<()> {
259 if visited.contains(&node_id) {
260 return Ok(());
261 }
262 if !visiting.insert(node_id) {
263 return Err(XlogError::Compilation(format!(
264 "Decision-DNNF parse error: cycle detected at node {}",
265 node_id
266 )));
267 }
268
269 let node = nodes.get(&node_id).ok_or_else(|| {
270 XlogError::Compilation(format!(
271 "Decision-DNNF parse error: unknown node {} during cycle check",
272 node_id
273 ))
274 })?;
275
276 match node.kind {
277 DdnnfNodeKind::True | DdnnfNodeKind::False => {}
278 DdnnfNodeKind::Or | DdnnfNodeKind::And => {
279 if let Some(out) = outgoing.get(&node_id) {
280 for &edge_idx in out {
281 let edge = &edges[edge_idx];
282 dfs(edge.to, nodes, edges, outgoing, visiting, visited)?;
283 }
284 }
285 }
286 }
287
288 visiting.remove(&node_id);
289 visited.insert(node_id);
290 Ok(())
291 }
292
293 dfs(root, nodes, edges, outgoing, &mut visiting, &mut visited)
294 }
295
296 pub fn eval_log_wmc<F>(&self, var_log_weights: F) -> Result<f64>
297 where
298 F: Fn(u32) -> (f64, f64),
299 {
300 let mut memo: HashMap<u32, f64> = HashMap::new();
301
302 fn eval_node<F>(
303 node_id: u32,
304 ddnnf: &DecisionDnnf,
305 memo: &mut HashMap<u32, f64>,
306 var_log_weights: &F,
307 ) -> Result<f64>
308 where
309 F: Fn(u32) -> (f64, f64),
310 {
311 if let Some(&v) = memo.get(&node_id) {
312 return Ok(v);
313 }
314
315 let node = ddnnf.nodes.get(&node_id).ok_or_else(|| {
316 XlogError::Compilation(format!(
317 "Decision-DNNF eval error: unknown node {}",
318 node_id
319 ))
320 })?;
321
322 let value = match node.kind {
323 DdnnfNodeKind::True => 0.0,
324 DdnnfNodeKind::False => f64::NEG_INFINITY,
325 DdnnfNodeKind::And => {
326 let out = ddnnf.outgoing.get(&node_id).ok_or_else(|| {
327 XlogError::Compilation(format!(
328 "Decision-DNNF eval error: AND node {} has no children",
329 node_id
330 ))
331 })?;
332
333 let mut acc = 0.0;
334 for &edge_idx in out {
335 let edge = &ddnnf.edges[edge_idx];
336 let child = eval_node(edge.to, ddnnf, memo, var_log_weights)?;
337 let mut lit_sum = 0.0;
338 for &lit in &edge.lits {
339 let var = lit.unsigned_abs();
340 let (t, f) = validate_circuit_log_weight_pair(var_log_weights(var))?;
341 lit_sum += if lit > 0 { t } else { f };
342 }
343 acc += lit_sum + child;
344 }
345 acc
346 }
347 DdnnfNodeKind::Or => {
348 let out = ddnnf.outgoing.get(&node_id).ok_or_else(|| {
349 XlogError::Compilation(format!(
350 "Decision-DNNF eval error: OR node {} has no children",
351 node_id
352 ))
353 })?;
354
355 let mut branch_vals: Vec<f64> = Vec::with_capacity(out.len());
356 for &edge_idx in out {
357 let edge = &ddnnf.edges[edge_idx];
358 let child = eval_node(edge.to, ddnnf, memo, var_log_weights)?;
359 let mut lit_sum = 0.0;
360 for &lit in &edge.lits {
361 let var = lit.unsigned_abs();
362 let (t, f) = validate_circuit_log_weight_pair(var_log_weights(var))?;
363 lit_sum += if lit > 0 { t } else { f };
364 }
365 branch_vals.push(lit_sum + child);
366 }
367 circuit_logsumexp(&branch_vals)?
368 }
369 };
370
371 let value = validate_circuit_value(value)?;
372 memo.insert(node_id, value);
373 Ok(value)
374 }
375
376 eval_node(self.root, self, &mut memo, &var_log_weights)
377 }
378}
379
380#[cfg(test)]
381mod tests {
382 use super::*;
383
384 #[test]
385 fn test_parse_and_eval_identity_variable() {
386 let nnf = r#"
388o 1 0
389t 2 0
390f 3 0
3911 2 1 0
3921 3 -1 0
393"#;
394
395 let ddnnf = DecisionDnnf::parse_str(nnf).unwrap();
396 assert_eq!(ddnnf.root(), 1);
397 assert_eq!(ddnnf.max_var(), 1);
398
399 let p = 0.3_f64;
400 let log_wmc = ddnnf
401 .eval_log_wmc(|var| match var {
402 1 => (p.ln(), (1.0 - p).ln()),
403 _ => panic!("unexpected var {}", var),
404 })
405 .unwrap();
406
407 assert!((log_wmc - p.ln()).abs() < 1e-9, "log_wmc={}", log_wmc);
408 }
409
410 #[test]
411 fn test_parse_detects_missing_terminator() {
412 let nnf = "t 1";
413 let err = DecisionDnnf::parse_str(nnf).unwrap_err();
414 let msg = err.to_string();
415 assert!(msg.contains("terminator"), "msg={}", msg);
416 }
417}