1use crate::metadata::RirMeta;
4use crate::rir::RirNode;
5use xlog_core::RelId;
6
7#[derive(Debug, Clone)]
9pub struct Scc {
10 pub id: u32,
12 pub predicates: Vec<String>,
14 pub is_recursive: bool,
16}
17
18#[derive(Debug, Clone)]
20pub struct Stratum {
21 pub id: u32,
23 pub sccs: Vec<u32>,
25}
26
27#[derive(Debug, Clone)]
29pub struct CompiledRule {
30 pub head: String,
32 pub body: RirNode,
34 pub meta: RirMeta,
36}
37
38#[derive(Debug, Clone, PartialEq, Eq)]
40pub struct GeneratedQueryRuleProvenance {
41 pub query_index: usize,
43 pub scc_index: usize,
45 pub rule_index: usize,
47}
48
49#[derive(Debug, Clone)]
51pub struct ExecutionPlan {
52 pub sccs: Vec<Scc>,
54 pub strata: Vec<Stratum>,
56 pub rules_by_scc: Vec<Vec<CompiledRule>>,
58 pub generated_query_rules: Vec<GeneratedQueryRuleProvenance>,
62 pub est_memory_peak: u64,
64 pub rel_arities: std::collections::HashMap<RelId, usize>,
69}
70
71impl ExecutionPlan {
72 pub fn new(sccs: Vec<Scc>) -> Self {
74 Self {
75 sccs,
76 strata: vec![],
77 rules_by_scc: vec![],
78 generated_query_rules: vec![],
79 est_memory_peak: 0,
80 rel_arities: std::collections::HashMap::new(),
81 }
82 }
83
84 pub fn with_strata(mut self, strata: Vec<Stratum>) -> Self {
86 self.strata = strata;
87 self
88 }
89
90 pub fn recursive_scc_count(&self) -> usize {
92 self.sccs.iter().filter(|s| s.is_recursive).count()
93 }
94
95 pub fn dependency_closed_subplan(
103 &self,
104 root_sccs: &[usize],
105 defining_sccs: &std::collections::HashMap<RelId, usize>,
106 ) -> Option<Self> {
107 let scc_count = self.sccs.len();
108 if root_sccs.is_empty()
109 || self.rules_by_scc.len() != scc_count
110 || self
111 .sccs
112 .iter()
113 .enumerate()
114 .any(|(index, scc)| u32::try_from(index).ok() != Some(scc.id))
115 || self
116 .strata
117 .iter()
118 .enumerate()
119 .any(|(index, stratum)| u32::try_from(index).ok() != Some(stratum.id))
120 || root_sccs.iter().any(|root| *root >= scc_count)
121 || defining_sccs.values().any(|scc| *scc >= scc_count)
122 {
123 return None;
124 }
125
126 let mut stratum_membership = vec![0_u8; scc_count];
127 for stratum in &self.strata {
128 for scc in &stratum.sccs {
129 let scc_index = usize::try_from(*scc).ok()?;
130 let membership = stratum_membership.get_mut(scc_index)?;
131 *membership = membership.checked_add(1)?;
132 if *membership != 1 {
133 return None;
134 }
135 }
136 }
137 if stratum_membership.iter().any(|membership| *membership != 1) {
138 return None;
139 }
140
141 for query in &self.generated_query_rules {
142 self.rules_by_scc
143 .get(query.scc_index)?
144 .get(query.rule_index)?;
145 }
146
147 let mut retained = std::collections::BTreeSet::new();
148 let mut pending = root_sccs.to_vec();
149 while let Some(scc_index) = pending.pop() {
150 if !retained.insert(scc_index) {
151 continue;
152 }
153 for rule in &self.rules_by_scc[scc_index] {
154 for relation in rule.body.referenced_relations() {
155 let Some(dependency) = defining_sccs.get(&relation).copied() else {
156 continue;
157 };
158 if dependency >= scc_count {
159 return None;
160 }
161 if !retained.contains(&dependency) {
162 pending.push(dependency);
163 }
164 }
165 }
166 }
167 let mut remapped_sccs = vec![None; scc_count];
168 let mut sccs = Vec::with_capacity(retained.len());
169 let mut rules_by_scc = Vec::with_capacity(retained.len());
170 for old_index in 0..scc_count {
171 if !retained.contains(&old_index) {
172 continue;
173 }
174 let new_index = u32::try_from(sccs.len()).ok()?;
175 remapped_sccs[old_index] = Some(new_index);
176 let mut scc = self.sccs[old_index].clone();
177 scc.id = new_index;
178 sccs.push(scc);
179 rules_by_scc.push(self.rules_by_scc[old_index].clone());
180 }
181
182 let mut strata = Vec::new();
183 for original in &self.strata {
184 let mut remapped = Vec::new();
185 for scc in &original.sccs {
186 let old_index = usize::try_from(*scc).ok()?;
187 if let Some(new_index) = remapped_sccs.get(old_index).copied().flatten() {
188 remapped.push(new_index);
189 }
190 }
191 if remapped.is_empty() {
192 continue;
193 }
194 strata.push(Stratum {
195 id: u32::try_from(strata.len()).ok()?,
196 sccs: remapped,
197 });
198 }
199
200 let generated_query_rules = self
201 .generated_query_rules
202 .iter()
203 .map(|query| {
204 Some(GeneratedQueryRuleProvenance {
205 query_index: query.query_index,
206 scc_index: remapped_sccs[query.scc_index]? as usize,
207 rule_index: query.rule_index,
208 })
209 })
210 .collect::<Option<Vec<_>>>()?;
211
212 Some(Self {
213 sccs,
214 strata,
215 rules_by_scc,
216 generated_query_rules,
217 est_memory_peak: self.est_memory_peak,
218 rel_arities: self.rel_arities.clone(),
219 })
220 }
221
222 pub fn has_recursion(&self) -> bool {
224 self.sccs.iter().any(|s| s.is_recursive)
225 }
226}
227
228#[derive(Debug, Default)]
230pub struct PlanBuilder {
231 sccs: Vec<Scc>,
232 strata: Vec<Stratum>,
233 rules: Vec<Vec<CompiledRule>>,
234}
235
236impl PlanBuilder {
237 pub fn new() -> Self {
239 Self::default()
240 }
241
242 pub fn add_scc(&mut self, scc: Scc) -> &mut Self {
244 self.sccs.push(scc);
245 self.rules.push(vec![]);
246 self
247 }
248
249 pub fn add_rule(&mut self, scc_id: u32, rule: CompiledRule) -> &mut Self {
251 if let Some(rules) = self.rules.get_mut(scc_id as usize) {
252 rules.push(rule);
253 }
254 self
255 }
256
257 pub fn add_stratum(&mut self, stratum: Stratum) -> &mut Self {
259 self.strata.push(stratum);
260 self
261 }
262
263 pub fn build(self) -> ExecutionPlan {
265 ExecutionPlan {
266 sccs: self.sccs,
267 strata: self.strata,
268 rules_by_scc: self.rules,
269 generated_query_rules: vec![],
270 est_memory_peak: 0,
271 rel_arities: std::collections::HashMap::new(),
272 }
273 }
274}
275
276#[cfg(test)]
277mod tests {
278 use super::*;
279
280 #[test]
281 fn test_scc_ordering() {
282 let sccs = vec![
283 Scc {
284 id: 0,
285 predicates: vec!["edge".into()],
286 is_recursive: false,
287 },
288 Scc {
289 id: 1,
290 predicates: vec!["reach".into()],
291 is_recursive: true,
292 },
293 ];
294 let plan = ExecutionPlan::new(sccs);
295 assert_eq!(plan.sccs.len(), 2);
296 assert!(!plan.sccs[0].is_recursive);
297 assert!(plan.sccs[1].is_recursive);
298 }
299
300 #[test]
301 fn test_stratum_assignment() {
302 let strata = [
303 Stratum {
304 id: 0,
305 sccs: vec![0, 1],
306 },
307 Stratum {
308 id: 1,
309 sccs: vec![2],
310 },
311 ];
312 assert_eq!(strata[0].sccs.len(), 2);
313 }
314
315 #[test]
316 fn test_plan_builder() {
317 let mut builder = PlanBuilder::new();
318 builder.add_scc(Scc {
319 id: 0,
320 predicates: vec!["p".into()],
321 is_recursive: false,
322 });
323 builder.add_stratum(Stratum {
324 id: 0,
325 sccs: vec![0],
326 });
327 let plan = builder.build();
328 assert_eq!(plan.sccs.len(), 1);
329 assert_eq!(plan.strata.len(), 1);
330 }
331
332 fn compiled_rule(head: &str, body: RirNode) -> CompiledRule {
333 CompiledRule {
334 head: head.into(),
335 body,
336 meta: RirMeta::default(),
337 }
338 }
339
340 #[test]
341 fn dependency_closed_subplan_keeps_whole_sccs_and_remaps_plan_indices() {
342 let mut plan = ExecutionPlan {
343 sccs: vec![
344 Scc {
345 id: 0,
346 predicates: vec!["source".into()],
347 is_recursive: false,
348 },
349 Scc {
350 id: 1,
351 predicates: vec!["disconnected".into()],
352 is_recursive: false,
353 },
354 Scc {
355 id: 2,
356 predicates: vec!["reachable".into()],
357 is_recursive: true,
358 },
359 Scc {
360 id: 3,
361 predicates: vec!["audited".into()],
362 is_recursive: false,
363 },
364 Scc {
365 id: 4,
366 predicates: vec!["__xlog_constraint_0".into()],
367 is_recursive: false,
368 },
369 Scc {
370 id: 5,
371 predicates: vec!["__xlog_query_0".into()],
372 is_recursive: false,
373 },
374 ],
375 strata: vec![
376 Stratum {
377 id: 0,
378 sccs: vec![0],
379 },
380 Stratum {
381 id: 1,
382 sccs: vec![1],
383 },
384 Stratum {
385 id: 2,
386 sccs: vec![2, 3],
387 },
388 Stratum {
389 id: 3,
390 sccs: vec![4, 5],
391 },
392 ],
393 rules_by_scc: vec![
394 vec![compiled_rule("source", RirNode::Scan { rel: RelId(1) })],
395 vec![compiled_rule(
396 "disconnected",
397 RirNode::Scan { rel: RelId(2) },
398 )],
399 vec![compiled_rule(
400 "reachable",
401 RirNode::Union {
402 inputs: vec![
403 RirNode::Scan { rel: RelId(10) },
404 RirNode::Scan { rel: RelId(12) },
405 ],
406 },
407 )],
408 vec![compiled_rule("audited", RirNode::Scan { rel: RelId(10) })],
409 vec![compiled_rule(
410 "__xlog_constraint_0",
411 RirNode::Scan { rel: RelId(13) },
412 )],
413 vec![compiled_rule(
414 "__xlog_query_0",
415 RirNode::Scan { rel: RelId(12) },
416 )],
417 ],
418 generated_query_rules: vec![GeneratedQueryRuleProvenance {
419 query_index: 0,
420 scc_index: 5,
421 rule_index: 0,
422 }],
423 est_memory_peak: 123,
424 rel_arities: [(RelId(10), 1), (RelId(12), 1)].into_iter().collect(),
425 };
426 let original_heads = plan
427 .rules_by_scc
428 .iter()
429 .flatten()
430 .map(|rule| rule.head.clone())
431 .collect::<Vec<_>>();
432 let defining_sccs = [
433 (RelId(10), 0),
434 (RelId(11), 1),
435 (RelId(12), 2),
436 (RelId(13), 3),
437 (RelId(14), 4),
438 (RelId(15), 5),
439 ]
440 .into_iter()
441 .collect();
442
443 let projected = plan
444 .dependency_closed_subplan(&[4, 5], &defining_sccs)
445 .expect("consistent dependency closure");
446
447 assert_eq!(
448 projected
449 .rules_by_scc
450 .iter()
451 .flatten()
452 .map(|rule| rule.head.as_str())
453 .collect::<Vec<_>>(),
454 vec![
455 "source",
456 "reachable",
457 "audited",
458 "__xlog_constraint_0",
459 "__xlog_query_0",
460 ]
461 );
462 assert_eq!(
463 projected.sccs.iter().map(|scc| scc.id).collect::<Vec<_>>(),
464 vec![0, 1, 2, 3, 4]
465 );
466 assert_eq!(
467 projected
468 .strata
469 .iter()
470 .map(|stratum| (stratum.id, stratum.sccs.clone()))
471 .collect::<Vec<_>>(),
472 vec![(0, vec![0]), (1, vec![1, 2]), (2, vec![3, 4])]
473 );
474 assert_eq!(
475 projected.generated_query_rules,
476 vec![GeneratedQueryRuleProvenance {
477 query_index: 0,
478 scc_index: 4,
479 rule_index: 0,
480 }]
481 );
482 assert_eq!(projected.est_memory_peak, 123);
483 assert_eq!(projected.rel_arities, plan.rel_arities);
484 assert_eq!(
485 plan.rules_by_scc
486 .iter()
487 .flatten()
488 .map(|rule| rule.head.clone())
489 .collect::<Vec<_>>(),
490 original_heads
491 );
492
493 let mut missing_unretained_schedule_entry = plan.clone();
494 missing_unretained_schedule_entry.strata[1].sccs.clear();
495 assert!(missing_unretained_schedule_entry
496 .dependency_closed_subplan(&[4, 5], &defining_sccs)
497 .is_none());
498
499 plan.generated_query_rules[0].scc_index = 1;
500 assert!(plan
501 .dependency_closed_subplan(&[4, 5], &defining_sccs)
502 .is_none());
503 }
504
505 #[test]
506 fn dependency_closed_subplan_rejects_inconsistent_proof_without_pruning() {
507 let plan = ExecutionPlan {
508 sccs: vec![Scc {
509 id: 0,
510 predicates: vec!["output".into()],
511 is_recursive: false,
512 }],
513 strata: vec![Stratum {
514 id: 0,
515 sccs: vec![0],
516 }],
517 rules_by_scc: vec![vec![compiled_rule(
518 "output",
519 RirNode::Scan { rel: RelId(1) },
520 )]],
521 generated_query_rules: vec![],
522 est_memory_peak: 0,
523 rel_arities: Default::default(),
524 };
525
526 assert!(plan
527 .dependency_closed_subplan(&[], &Default::default())
528 .is_none());
529 assert!(plan
530 .dependency_closed_subplan(&[1], &Default::default())
531 .is_none());
532 assert!(plan
533 .dependency_closed_subplan(&[0], &[(RelId(1), 1)].into_iter().collect())
534 .is_none());
535
536 let mut invalid_scc_id = plan.clone();
537 invalid_scc_id.sccs[0].id = 7;
538 assert!(invalid_scc_id
539 .dependency_closed_subplan(&[0], &Default::default())
540 .is_none());
541
542 let mut invalid_stratum = plan;
543 invalid_stratum.strata[0].sccs[0] = 7;
544 assert!(invalid_stratum
545 .dependency_closed_subplan(&[0], &Default::default())
546 .is_none());
547
548 let mut invalid_stratum_id = invalid_scc_id;
549 invalid_stratum_id.sccs[0].id = 0;
550 invalid_stratum_id.strata[0].id = 7;
551 assert!(invalid_stratum_id
552 .dependency_closed_subplan(&[0], &Default::default())
553 .is_none());
554 }
555
556 #[test]
557 fn test_has_recursion() {
558 let non_recursive = ExecutionPlan::new(vec![Scc {
559 id: 0,
560 predicates: vec!["p".into()],
561 is_recursive: false,
562 }]);
563 assert!(!non_recursive.has_recursion());
564
565 let recursive = ExecutionPlan::new(vec![Scc {
566 id: 0,
567 predicates: vec!["reach".into()],
568 is_recursive: true,
569 }]);
570 assert!(recursive.has_recursion());
571 }
572}