Skip to main content

xlog_ir/
plan.rs

1//! Execution plan representation
2
3use crate::metadata::RirMeta;
4use crate::rir::RirNode;
5use xlog_core::RelId;
6
7/// Strongly Connected Component in the dependency graph
8#[derive(Debug, Clone)]
9pub struct Scc {
10    /// Unique SCC identifier
11    pub id: u32,
12    /// Predicate names in this SCC
13    pub predicates: Vec<String>,
14    /// Whether this SCC contains recursion
15    pub is_recursive: bool,
16}
17
18/// Stratum in stratified evaluation
19#[derive(Debug, Clone)]
20pub struct Stratum {
21    /// Stratum number (0 = base)
22    pub id: u32,
23    /// SCCs in this stratum (topologically ordered)
24    pub sccs: Vec<u32>,
25}
26
27/// Compiled rule ready for execution
28#[derive(Debug, Clone)]
29pub struct CompiledRule {
30    /// Head predicate name
31    pub head: String,
32    /// RIR tree for rule body
33    pub body: RirNode,
34    /// Metadata for cost estimation
35    pub meta: RirMeta,
36}
37
38/// Compiler-produced provenance for one desugared program query.
39#[derive(Debug, Clone, PartialEq, Eq)]
40pub struct GeneratedQueryRuleProvenance {
41    /// Zero-based position of the query in the authored program.
42    pub query_index: usize,
43    /// Position of the generated rule's SCC in [`ExecutionPlan::rules_by_scc`].
44    pub scc_index: usize,
45    /// Position of the generated rule within its SCC.
46    pub rule_index: usize,
47}
48
49/// Complete execution plan for a program
50#[derive(Debug, Clone)]
51pub struct ExecutionPlan {
52    /// SCCs in dependency order
53    pub sccs: Vec<Scc>,
54    /// Strata for negation ordering
55    pub strata: Vec<Stratum>,
56    /// Compiled rules grouped by SCC
57    pub rules_by_scc: Vec<Vec<CompiledRule>>,
58    /// Exact compiled rule positions originating from authored program queries.
59    /// Manually assembled plans leave this empty unless they explicitly model
60    /// generated-query semantics.
61    pub generated_query_rules: Vec<GeneratedQueryRuleProvenance>,
62    /// Total estimated memory peak (bytes)
63    pub est_memory_peak: u64,
64    /// Relation arities known at lowering time (every predicate the
65    /// lowerer assigned a RelId). Consumed by shape promoters that
66    /// must size Scan leaves without schema access (general Free Join
67    /// multiway promotion).
68    pub rel_arities: std::collections::HashMap<RelId, usize>,
69}
70
71impl ExecutionPlan {
72    /// Create a new execution plan from SCCs
73    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    /// Add strata to the plan
85    pub fn with_strata(mut self, strata: Vec<Stratum>) -> Self {
86        self.strata = strata;
87        self
88    }
89
90    /// Get the number of recursive SCCs
91    pub fn recursive_scc_count(&self) -> usize {
92        self.sccs.iter().filter(|s| s.is_recursive).count()
93    }
94
95    /// Return the dependency-closed subplan for a set of root SCC positions.
96    ///
97    /// `defining_sccs` maps derived relation IDs to the SCC that defines them.
98    /// Relations absent from that map are treated as extensional inputs. The
99    /// projection preserves source order, retains complete SCCs, and rewrites
100    /// every positional index in the plan. `None` means the supplied plan or
101    /// dependency proof is inconsistent, so callers must keep the original plan.
102    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    /// Check if this plan has any recursion
223    pub fn has_recursion(&self) -> bool {
224        self.sccs.iter().any(|s| s.is_recursive)
225    }
226}
227
228/// Builder for execution plans
229#[derive(Debug, Default)]
230pub struct PlanBuilder {
231    sccs: Vec<Scc>,
232    strata: Vec<Stratum>,
233    rules: Vec<Vec<CompiledRule>>,
234}
235
236impl PlanBuilder {
237    /// Create a new empty plan builder.
238    pub fn new() -> Self {
239        Self::default()
240    }
241
242    /// Append a strongly connected component to the plan.
243    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    /// Add a compiled rule to the given SCC (by index).
250    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    /// Append a stratum to the plan.
258    pub fn add_stratum(&mut self, stratum: Stratum) -> &mut Self {
259        self.strata.push(stratum);
260        self
261    }
262
263    /// Consume the builder and produce the final [`ExecutionPlan`].
264    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}