antecedent_graph/
cpdag_completion.rs1use crate::cpdag::Cpdag;
12use crate::dag::Dag;
13use crate::error::GraphError;
14use crate::types::DenseNodeId;
15
16#[derive(Clone, Debug)]
18pub struct CpdagCompletion {
19 pub graph: Dag,
21 pub index: usize,
23}
24
25#[derive(Clone, Debug)]
27pub struct CpdagCompletionSampler {
28 base: Cpdag,
29 undirected: Vec<(DenseNodeId, DenseNodeId)>,
31 max_completions: usize,
32 next_index: usize,
33 assign: u64,
35}
36
37impl CpdagCompletionSampler {
38 pub fn new(cpdag: Cpdag, max_completions: usize) -> Result<Self, GraphError> {
44 if cpdag.conflict_edge_count() > 0 {
45 return Err(GraphError::InvalidEndpoints {
46 message: "CpdagCompletionSampler refuses conflict (x-x) edges",
47 });
48 }
49 let mut undirected = Vec::new();
50 for e in cpdag.edges() {
51 if e.is_undirected() {
52 let (a, b) = if e.a.raw() <= e.b.raw() { (e.a, e.b) } else { (e.b, e.a) };
53 undirected.push((a, b));
54 }
55 }
56 undirected.sort_by_key(|(a, b)| (a.raw(), b.raw()));
57 undirected.dedup();
58 if undirected.len() > 63 {
59 return Err(GraphError::InvalidEndpoints {
60 message: "too many undirected edges for CpdagCompletionSampler mask",
61 });
62 }
63 Ok(Self { base: cpdag, undirected, max_completions, next_index: 0, assign: 0 })
64 }
65
66 #[must_use]
68 pub fn max_completions(&self) -> usize {
69 self.max_completions
70 }
71
72 #[must_use]
74 pub fn n_undirected(&self) -> usize {
75 self.undirected.len()
76 }
77
78 fn build_completion(&self, mask: u64) -> Option<Dag> {
79 let mut g = self.base.clone();
80 for (i, &(a, b)) in self.undirected.iter().enumerate() {
81 let reverse = ((mask >> i) & 1) == 1;
82 let (from, to) = if reverse { (b, a) } else { (a, b) };
83 if g.orient_undirected(from, to).is_err() {
84 return None;
85 }
86 }
87 let dag = g.try_into_dag().ok()?;
88 if is_mec_member(&self.base, &dag) { Some(dag) } else { None }
89 }
90}
91
92#[must_use]
95pub fn is_mec_member(cpdag: &Cpdag, dag: &Dag) -> bool {
96 if cpdag.node_count() != dag.node_count() {
97 return false;
98 }
99 for e in cpdag.edges() {
101 if let Some((from, to)) = e.parent_child() {
102 if !dag.children(from).contains(&to) {
103 return false;
104 }
105 } else if e.is_undirected() {
106 let a = e.a;
107 let b = e.b;
108 let ab = dag.children(a).contains(&b);
109 let ba = dag.children(b).contains(&a);
110 if ab == ba {
111 return false;
113 }
114 } else if e.is_conflict() {
115 return false;
116 }
117 }
118 for e in dag.edges() {
120 if let Some((from, to)) = e.parent_child() {
121 if !cpdag.has_edge(from, to) {
122 return false;
123 }
124 }
125 }
126 let cpdag_colliders = unshielded_colliders_cpdag(cpdag);
128 let dag_colliders = unshielded_colliders_dag(dag);
129 cpdag_colliders == dag_colliders
130}
131
132fn unshielded_colliders_cpdag(g: &Cpdag) -> Vec<(u32, u32, u32)> {
133 let n = g.node_count();
134 let mut out = Vec::new();
135 for center_idx in 0..n {
136 let Ok(center_raw) = u32::try_from(center_idx) else {
137 break;
138 };
139 let center = DenseNodeId::from_raw(center_raw);
140 let parents = g.parents(center);
141 for left_i in 0..parents.len() {
142 for right_i in (left_i + 1)..parents.len() {
143 let left_parent = parents[left_i];
144 let right_parent = parents[right_i];
145 if !g.has_edge(left_parent, right_parent) {
146 let (lo, hi) = if left_parent.raw() <= right_parent.raw() {
147 (left_parent.raw(), right_parent.raw())
148 } else {
149 (right_parent.raw(), left_parent.raw())
150 };
151 out.push((lo, center.raw(), hi));
152 }
153 }
154 }
155 }
156 out.sort_unstable();
157 out.dedup();
158 out
159}
160
161fn unshielded_colliders_dag(g: &Dag) -> Vec<(u32, u32, u32)> {
162 let n = g.node_count();
163 let mut out = Vec::new();
164 for center_idx in 0..n {
165 let Ok(center_raw) = u32::try_from(center_idx) else {
166 break;
167 };
168 let center = DenseNodeId::from_raw(center_raw);
169 let parents = g.parents(center);
170 for left_i in 0..parents.len() {
171 for right_i in (left_i + 1)..parents.len() {
172 let left_parent = parents[left_i];
173 let right_parent = parents[right_i];
174 let adjacent = g.children(left_parent).contains(&right_parent)
175 || g.children(right_parent).contains(&left_parent);
176 if !adjacent {
177 let (lo, hi) = if left_parent.raw() <= right_parent.raw() {
178 (left_parent.raw(), right_parent.raw())
179 } else {
180 (right_parent.raw(), left_parent.raw())
181 };
182 out.push((lo, center.raw(), hi));
183 }
184 }
185 }
186 }
187 out.sort_unstable();
188 out.dedup();
189 out
190}
191
192impl Iterator for CpdagCompletionSampler {
193 type Item = CpdagCompletion;
194
195 fn next(&mut self) -> Option<Self::Item> {
196 if self.next_index >= self.max_completions {
197 return None;
198 }
199 let n = self.undirected.len();
200 let total = if n == 0 { 1u64 } else { 1u64 << n };
201 while self.assign < total {
202 let mask = self.assign;
203 self.assign += 1;
204 if let Some(graph) = self.build_completion(mask) {
205 let index = self.next_index;
206 self.next_index += 1;
207 return Some(CpdagCompletion { graph, index });
208 }
209 }
210 None
211 }
212}
213
214#[cfg(test)]
215mod tests {
216 use super::*;
217 use crate::cpdag::Cpdag;
218
219 #[test]
220 fn fully_oriented_yields_one() {
221 let mut g = Cpdag::with_variables(2);
222 g.insert_directed(DenseNodeId::from_raw(0), DenseNodeId::from_raw(1)).unwrap();
223 let collected: Vec<_> = CpdagCompletionSampler::new(g, 10).unwrap().collect();
224 assert_eq!(collected.len(), 1);
225 assert!(
226 collected[0]
227 .graph
228 .children(DenseNodeId::from_raw(0))
229 .contains(&DenseNodeId::from_raw(1))
230 );
231 }
232
233 #[test]
234 fn chain_undirected_has_three_mec_dags() {
235 let mut g = Cpdag::with_variables(3);
237 let a = DenseNodeId::from_raw(0);
238 let b = DenseNodeId::from_raw(1);
239 let c = DenseNodeId::from_raw(2);
240 g.insert_undirected(a, b).unwrap();
241 g.insert_undirected(b, c).unwrap();
242 let collected: Vec<_> = CpdagCompletionSampler::new(g.clone(), 16).unwrap().collect();
243 assert_eq!(g.undirected_edge_count(), 2);
244 assert_eq!(collected.len(), 3, "expected 3 MEC DAGs, got {}", collected.len());
245 for c in &collected {
246 assert!(is_mec_member(&g, &c.graph));
247 }
248 }
249
250 #[test]
251 fn respects_max_completions_bound() {
252 let mut g = Cpdag::with_variables(3);
253 g.insert_undirected(DenseNodeId::from_raw(0), DenseNodeId::from_raw(1)).unwrap();
254 g.insert_undirected(DenseNodeId::from_raw(1), DenseNodeId::from_raw(2)).unwrap();
255 let collected: Vec<_> = CpdagCompletionSampler::new(g, 2).unwrap().collect();
256 assert_eq!(collected.len(), 2);
257 }
258
259 #[test]
260 fn refuses_conflict_edges() {
261 let mut g = Cpdag::with_variables(2);
262 g.insert_undirected(DenseNodeId::from_raw(0), DenseNodeId::from_raw(1)).unwrap();
263 g.mark_conflict(DenseNodeId::from_raw(0), DenseNodeId::from_raw(1)).unwrap();
264 assert!(CpdagCompletionSampler::new(g, 4).is_err());
265 }
266
267 #[test]
268 fn compelled_collider_preserved() {
269 let mut g = Cpdag::with_variables(3);
271 let a = DenseNodeId::from_raw(0);
272 let b = DenseNodeId::from_raw(1);
273 let c = DenseNodeId::from_raw(2);
274 g.insert_directed(a, b).unwrap();
275 g.insert_directed(c, b).unwrap();
276 let collected: Vec<_> = CpdagCompletionSampler::new(g.clone(), 8).unwrap().collect();
277 assert_eq!(collected.len(), 1);
278 assert!(is_mec_member(&g, &collected[0].graph));
279 }
280}