Skip to main content

lora_compiler/
planner.rs

1use crate::pattern::PatternPlanner;
2use crate::{
3    Aggregation, Argument, Filter, Limit, LogicalOp, LogicalPlan, OptionalMatch, PlanNodeId,
4    Projection, Sort, Unwind,
5};
6use lora_analyzer::symbols::VarId;
7use lora_analyzer::{
8    ResolvedCallSubquery, ResolvedClause, ResolvedCreate, ResolvedDelete, ResolvedExpr,
9    ResolvedForeach, ResolvedMatch, ResolvedMerge, ResolvedPattern, ResolvedPatternElement,
10    ResolvedProjection, ResolvedQuery, ResolvedRemove, ResolvedReturn, ResolvedSet,
11    ResolvedSortItem, ResolvedUnwind, ResolvedWith,
12};
13use lora_store::GraphStats;
14use std::collections::BTreeSet;
15
16pub struct Planner {
17    nodes: Vec<LogicalOp>,
18    /// Cardinalities used to pick which end of a pattern to start from.
19    stats: GraphStats,
20    /// Variables bound by the clauses planned so far. Only a hint for
21    /// choosing a pattern's starting point: a wrong entry costs speed,
22    /// never correctness, because scans re-check bound variables.
23    bound: BTreeSet<VarId>,
24}
25
26impl Default for Planner {
27    fn default() -> Self {
28        Self::new()
29    }
30}
31
32impl Planner {
33    pub fn new() -> Self {
34        Self::with_stats(&GraphStats::default())
35    }
36
37    pub fn with_stats(stats: &GraphStats) -> Self {
38        Self {
39            nodes: Vec::new(),
40            stats: stats.clone(),
41            bound: BTreeSet::new(),
42        }
43    }
44
45    pub(crate) fn push(&mut self, op: LogicalOp) -> PlanNodeId {
46        let id = self.nodes.len();
47        self.nodes.push(op);
48        id
49    }
50
51    pub(crate) fn stats(&self) -> &GraphStats {
52        &self.stats
53    }
54
55    pub(crate) fn is_bound(&self, var: VarId) -> bool {
56        self.bound.contains(&var)
57    }
58
59    fn bind_projection(&mut self, items: &[ResolvedProjection], include_existing: bool) {
60        if !include_existing {
61            self.bound.clear();
62        }
63        self.bound.extend(items.iter().map(|item| item.output));
64    }
65
66    pub fn plan(&mut self, query: &ResolvedQuery) -> LogicalPlan {
67        let root = self.plan_query(query);
68
69        LogicalPlan {
70            root,
71            nodes: std::mem::take(&mut self.nodes),
72        }
73    }
74
75    fn plan_query(&mut self, query: &ResolvedQuery) -> PlanNodeId {
76        let mut input = None;
77
78        for clause in &query.clauses {
79            input = Some(self.plan_clause(input, clause));
80            self.track_bindings(clause);
81        }
82
83        input.unwrap_or_else(|| self.plan_unit_input())
84    }
85
86    /// Record the variables `clause` leaves bound for the clauses after it.
87    fn track_bindings(&mut self, clause: &ResolvedClause) {
88        match clause {
89            ResolvedClause::Match(m) => self.bound.extend(pattern_binders(&m.pattern)),
90            ResolvedClause::Create(c) => self.bound.extend(pattern_binders(&c.pattern)),
91            ResolvedClause::Merge(m) => {
92                let pattern = ResolvedPattern {
93                    parts: vec![m.pattern_part.clone()],
94                };
95                self.bound.extend(pattern_binders(&pattern));
96            }
97            ResolvedClause::Unwind(u) => {
98                self.bound.insert(u.alias);
99            }
100            ResolvedClause::With(w) => self.bind_projection(&w.items, w.include_existing),
101            ResolvedClause::Return(r) => self.bind_projection(&r.items, r.include_existing),
102            ResolvedClause::CallSubquery(c) => self.bound.extend(c.return_vars.iter().copied()),
103            ResolvedClause::Delete(_)
104            | ResolvedClause::Set(_)
105            | ResolvedClause::Remove(_)
106            | ResolvedClause::Foreach(_) => {}
107        }
108    }
109
110    fn plan_clause(&mut self, input: Option<PlanNodeId>, clause: &ResolvedClause) -> PlanNodeId {
111        {
112            match clause {
113                ResolvedClause::Match(m) => self.plan_match(input, m),
114
115                ResolvedClause::Unwind(u) => {
116                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
117                    self.plan_unwind(upstream, u)
118                }
119
120                ResolvedClause::Create(c) => {
121                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
122                    self.plan_create(upstream, c)
123                }
124
125                ResolvedClause::Merge(m) => {
126                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
127                    self.plan_merge(upstream, m)
128                }
129
130                ResolvedClause::Delete(d) => {
131                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
132                    self.plan_delete(upstream, d)
133                }
134
135                ResolvedClause::Set(s) => {
136                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
137                    self.plan_set(upstream, s)
138                }
139
140                ResolvedClause::Remove(rm) => {
141                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
142                    self.plan_remove(upstream, rm)
143                }
144
145                ResolvedClause::Foreach(f) => {
146                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
147                    self.plan_foreach(upstream, f)
148                }
149
150                ResolvedClause::With(w) => {
151                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
152                    self.plan_with(upstream, w)
153                }
154
155                ResolvedClause::Return(r) => {
156                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
157                    self.plan_return(upstream, r)
158                }
159
160                ResolvedClause::CallSubquery(c) => {
161                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
162                    self.plan_call_subquery(upstream, c)
163                }
164            }
165        }
166    }
167
168    /// Plan a `CALL { ... }` subquery: build the inner plan starting
169    /// from a fresh `Argument`, then wrap it with `CallSubquery` to
170    /// drive it per outer row.
171    fn plan_call_subquery(&mut self, input: PlanNodeId, call: &ResolvedCallSubquery) -> PlanNodeId {
172        let inner_query = ResolvedQuery {
173            clauses: call.clauses.clone(),
174            unions: Vec::new(),
175            parameters: Default::default(),
176        };
177        // The inner query sees the outer row; its own WITH / RETURN must
178        // not change what the outer query considers bound.
179        let outer_bound = self.bound.clone();
180        let inner = self.plan_query(&inner_query);
181        self.bound = outer_bound;
182        self.push(LogicalOp::CallSubquery(crate::logical::CallSubquery {
183            input,
184            inner,
185            new_vars: call.return_vars.clone(),
186        }))
187    }
188
189    fn plan_match(&mut self, input: Option<PlanNodeId>, m: &ResolvedMatch) -> PlanNodeId {
190        if m.optional {
191            // OPTIONAL MATCH: build the inner sub-plan that reads from Argument,
192            // then wrap it in an OptionalMatch node that provides null-extension.
193            //
194            // A leading OPTIONAL MATCH has no upstream clause; openCypher
195            // drives it from the single empty unit row, so an empty match
196            // still yields one row with the pattern's variables null.
197            let upstream = input.unwrap_or_else(|| self.plan_unit_input());
198
199            // Collect variables introduced by this pattern (for null-extension).
200            let new_vars = pattern_binders(&m.pattern);
201
202            // Build inner match plan WITHOUT the upstream input — the executor
203            // will inject each upstream row individually. The WHERE belongs
204            // to the OPTIONAL MATCH, so its conjuncts are only ever placed
205            // inside this inner plan, never above the OptionalMatch.
206            let mut pattern_planner = PatternPlanner::new(self);
207            let (inner, residual) =
208                pattern_planner.plan_pattern_with_where(None, &m.pattern, m.where_.as_ref());
209            let inner = self.push_conjuncts(inner, residual);
210
211            self.push(LogicalOp::OptionalMatch(OptionalMatch {
212                input: upstream,
213                inner,
214                new_vars,
215            }))
216        } else {
217            let mut pattern_planner = PatternPlanner::new(self);
218            let (node, residual) =
219                pattern_planner.plan_pattern_with_where(input, &m.pattern, m.where_.as_ref());
220            self.push_conjuncts(node, residual)
221        }
222    }
223
224    /// A `Filter` over `input` holding `conjuncts` ANDed in order, or
225    /// `input` itself when there are none.
226    fn push_conjuncts(&mut self, input: PlanNodeId, conjuncts: Vec<ResolvedExpr>) -> PlanNodeId {
227        let predicate = conjuncts
228            .into_iter()
229            .reduce(|acc, next| ResolvedExpr::Binary {
230                lhs: Box::new(acc),
231                op: lora_ast::BinaryOp::And,
232                rhs: Box::new(next),
233            });
234        match predicate {
235            Some(predicate) => self.push(LogicalOp::Filter(Filter { input, predicate })),
236            None => input,
237        }
238    }
239
240    fn plan_unwind(&mut self, input: PlanNodeId, u: &ResolvedUnwind) -> PlanNodeId {
241        self.push(LogicalOp::Unwind(Unwind {
242            input,
243            expr: u.expr.clone(),
244            alias: u.alias,
245        }))
246    }
247
248    fn plan_create(&mut self, input: PlanNodeId, c: &ResolvedCreate) -> PlanNodeId {
249        self.push(LogicalOp::Create(crate::Create {
250            input,
251            pattern: c.pattern.clone(),
252        }))
253    }
254
255    fn plan_merge(&mut self, input: PlanNodeId, m: &ResolvedMerge) -> PlanNodeId {
256        self.push(LogicalOp::Merge(crate::Merge {
257            input,
258            pattern_part: m.pattern_part.clone(),
259            actions: m.actions.clone(),
260        }))
261    }
262
263    fn plan_delete(&mut self, input: PlanNodeId, d: &ResolvedDelete) -> PlanNodeId {
264        self.push(LogicalOp::Delete(crate::Delete {
265            input,
266            detach: d.detach,
267            expressions: d.expressions.clone(),
268        }))
269    }
270
271    fn plan_set(&mut self, input: PlanNodeId, s: &ResolvedSet) -> PlanNodeId {
272        self.push(LogicalOp::Set(crate::Set {
273            input,
274            items: s.items.clone(),
275        }))
276    }
277
278    fn plan_remove(&mut self, input: PlanNodeId, r: &ResolvedRemove) -> PlanNodeId {
279        self.push(LogicalOp::Remove(crate::Remove {
280            input,
281            items: r.items.clone(),
282        }))
283    }
284
285    fn plan_foreach(&mut self, input: PlanNodeId, f: &ResolvedForeach) -> PlanNodeId {
286        self.push(LogicalOp::Foreach(crate::Foreach {
287            input,
288            variable: f.variable,
289            list: f.list.clone(),
290            body: f.body.clone(),
291        }))
292    }
293
294    fn plan_with(&mut self, input: PlanNodeId, with: &ResolvedWith) -> PlanNodeId {
295        let mut node = self.plan_projection_sort_limit(
296            input,
297            &with.items,
298            &with.lifted_aggregates,
299            with.distinct,
300            with.include_existing,
301            &with.order,
302            &with.skip,
303            &with.limit,
304        );
305
306        if let Some(pred) = &with.where_ {
307            node = self.push(LogicalOp::Filter(Filter {
308                input: node,
309                predicate: pred.clone(),
310            }));
311        }
312
313        node
314    }
315
316    fn plan_return(&mut self, input: PlanNodeId, ret: &ResolvedReturn) -> PlanNodeId {
317        self.plan_projection_sort_limit(
318            input,
319            &ret.items,
320            &ret.lifted_aggregates,
321            ret.distinct,
322            ret.include_existing,
323            &ret.order,
324            &ret.skip,
325            &ret.limit,
326        )
327    }
328
329    /// Plan `items [ORDER BY ...] [SKIP ...] [LIMIT ...]` for WITH / RETURN.
330    ///
331    /// Cypher evaluates projection (with aggregation and DISTINCT) first,
332    /// then ORDER BY, then SKIP / LIMIT. Sort keys may name projected
333    /// aliases (`RETURN p.name AS name ORDER BY name`) and, unless the
334    /// projection aggregates or deduplicates, pre-projection variables too
335    /// (`RETURN p.name AS name ORDER BY p.age`).
336    ///
337    /// * Aggregation or DISTINCT changes the row set, so both must run
338    ///   before sorting and limiting; otherwise `LIMIT 1` would cut the
339    ///   input to one row before counting, and DISTINCT would dedupe an
340    ///   already-truncated list. Sort keys that restate a projected
341    ///   expression are pointed at that expression's output column.
342    /// * A plain projection maps rows 1:1. It first projects while keeping
343    ///   the input bindings, so keys can use both aliases and original
344    ///   variables, then sorts and limits, then trims the row down to the
345    ///   projected columns.
346    #[allow(clippy::too_many_arguments)]
347    fn plan_projection_sort_limit(
348        &mut self,
349        input: PlanNodeId,
350        items: &[ResolvedProjection],
351        lifted: &[ResolvedProjection],
352        distinct: bool,
353        include_existing: bool,
354        order: &[ResolvedSortItem],
355        skip: &Option<ResolvedExpr>,
356        limit: &Option<ResolvedExpr>,
357    ) -> PlanNodeId {
358        let aggregates =
359            !lifted.is_empty() || items.iter().any(|item| expr_contains_aggregate(&item.expr));
360        let has_order = !order.is_empty();
361        let has_limit = skip.is_some() || limit.is_some();
362
363        if aggregates || distinct {
364            // Hidden aggregate columns an ORDER BY key reads (the analyzer
365            // lifted them out of it) live until the sort, then go.
366            let sort_reads: BTreeSet<VarId> = order
367                .iter()
368                .flat_map(|key| {
369                    let mut reads = BTreeSet::new();
370                    key.expr.collect_vars(&mut reads);
371                    reads
372                })
373                .collect();
374            let kept: Vec<ResolvedProjection> = lifted
375                .iter()
376                .filter(|p| sort_reads.contains(&p.output))
377                .cloned()
378                .collect();
379            let mut node = self.plan_projection_or_aggregation(
380                input,
381                items,
382                lifted,
383                &kept,
384                distinct,
385                include_existing,
386            );
387            if has_order {
388                node = self.push(LogicalOp::Sort(Sort {
389                    input: node,
390                    items: sort_keys_on_outputs(order, items),
391                    top_k: None,
392                    limit: None,
393                }));
394            }
395            if has_limit {
396                node = self.push(LogicalOp::Limit(Limit {
397                    input: node,
398                    skip: skip.clone(),
399                    limit: limit.clone(),
400                }));
401            }
402            if !kept.is_empty() {
403                node = self.push(LogicalOp::Projection(Projection {
404                    input: node,
405                    distinct: false,
406                    items: passthrough_items(items),
407                    include_existing: false,
408                }));
409            }
410            return node;
411        }
412
413        if !has_order {
414            // LIMIT on a 1:1 projection can run first and saves projecting
415            // rows that would be dropped.
416            let mut node = input;
417            if has_limit {
418                node = self.push(LogicalOp::Limit(Limit {
419                    input: node,
420                    skip: skip.clone(),
421                    limit: limit.clone(),
422                }));
423            }
424            return self.plan_projection_or_aggregation(
425                node,
426                items,
427                lifted,
428                &[],
429                false,
430                include_existing,
431            );
432        }
433
434        let mut node = self.push(LogicalOp::Projection(Projection {
435            input,
436            distinct: false,
437            items: items.to_vec(),
438            include_existing: true,
439        }));
440        node = self.push(LogicalOp::Sort(Sort {
441            input: node,
442            items: order.to_vec(),
443            top_k: None,
444            limit: None,
445        }));
446        if has_limit {
447            node = self.push(LogicalOp::Limit(Limit {
448                input: node,
449                skip: skip.clone(),
450                limit: limit.clone(),
451            }));
452        }
453        if include_existing {
454            // WITH * / RETURN * keep every binding anyway.
455            return node;
456        }
457        self.push(LogicalOp::Projection(Projection {
458            input: node,
459            distinct: false,
460            items: passthrough_items(items),
461            include_existing: false,
462        }))
463    }
464
465    /// If any projection item contains an aggregate function, emit an
466    /// Aggregation node followed by a Projection. Otherwise emit a plain
467    /// Projection.
468    ///
469    /// `lifted` are the aggregate calls the analyzer took out of larger
470    /// items (`size(collect(x))`) or ORDER BY keys: they aggregate as hidden
471    /// columns, and a Projection then evaluates those items on the grouped
472    /// rows, re-emitting the hidden columns in `kept` for a sort to read.
473    fn plan_projection_or_aggregation(
474        &mut self,
475        input: PlanNodeId,
476        items: &[ResolvedProjection],
477        lifted: &[ResolvedProjection],
478        kept: &[ResolvedProjection],
479        distinct: bool,
480        include_existing: bool,
481    ) -> PlanNodeId {
482        if !lifted.is_empty() {
483            return self.plan_lifted_aggregation(input, items, lifted, kept, distinct);
484        }
485
486        let has_aggregates = items.iter().any(|item| expr_contains_aggregate(&item.expr));
487
488        if !has_aggregates {
489            return self.push(LogicalOp::Projection(Projection {
490                input,
491                distinct,
492                items: items.to_vec(),
493                include_existing,
494            }));
495        }
496
497        // Split items into group-by keys and aggregate expressions.
498        let mut group_by = Vec::new();
499        let mut aggregates = Vec::new();
500
501        for item in items {
502            if expr_contains_aggregate(&item.expr) {
503                aggregates.push(item.clone());
504            } else {
505                group_by.push(item.clone());
506            }
507        }
508
509        let node = self.push(LogicalOp::Aggregation(Aggregation {
510            input,
511            group_by: group_by.clone(),
512            aggregates: aggregates.clone(),
513        }));
514
515        // After aggregation the row already contains the right VarIds and names,
516        // but we still emit a Projection to handle DISTINCT and to ensure the
517        // final column order matches the original item list. The projection uses
518        // include_existing=true so it picks up the aggregation output, and each
519        // item just reads its own output variable.
520        //
521        // However, since the aggregation node already produces correctly-named
522        // rows, we can skip the extra projection when not needed.
523        if distinct {
524            // For DISTINCT we still need the dedup pass in exec_projection.
525            self.push(LogicalOp::Projection(Projection {
526                input: node,
527                distinct: true,
528                items: passthrough_items(items),
529                include_existing: false,
530            }))
531        } else {
532            node
533        }
534    }
535
536    /// Aggregate by the items that read no lifted column and aggregate
537    /// nothing, fold the bare aggregate items and the lifted calls, then
538    /// project every item in order: keys and bare aggregates re-emit their
539    /// column, the rest evaluate on the grouped row.
540    fn plan_lifted_aggregation(
541        &mut self,
542        input: PlanNodeId,
543        items: &[ResolvedProjection],
544        lifted: &[ResolvedProjection],
545        kept: &[ResolvedProjection],
546        distinct: bool,
547    ) -> PlanNodeId {
548        let lifted_outputs: BTreeSet<VarId> = lifted.iter().map(|p| p.output).collect();
549        let reads_lifted = |item: &ResolvedProjection| {
550            let mut reads = BTreeSet::new();
551            item.expr.collect_vars(&mut reads);
552            !reads.is_disjoint(&lifted_outputs)
553        };
554
555        let mut group_by = Vec::new();
556        let mut aggregates = Vec::new();
557        let mut projected = Vec::with_capacity(items.len());
558        for item in items {
559            if reads_lifted(item) {
560                projected.push(item.clone());
561                continue;
562            }
563            if expr_contains_aggregate(&item.expr) {
564                aggregates.push(item.clone());
565            } else {
566                group_by.push(item.clone());
567            }
568            projected.extend(passthrough_items(std::slice::from_ref(item)));
569        }
570        aggregates.extend(lifted.iter().cloned());
571        projected.extend(passthrough_items(kept));
572
573        let node = self.push(LogicalOp::Aggregation(Aggregation {
574            input,
575            group_by,
576            aggregates,
577        }));
578        self.push(LogicalOp::Projection(Projection {
579            input: node,
580            distinct,
581            items: projected,
582            include_existing: false,
583        }))
584    }
585
586    fn plan_unit_input(&mut self) -> PlanNodeId {
587        self.push(LogicalOp::Argument(Argument))
588    }
589}
590
591/// Projection items that re-emit each item's own output column.
592fn passthrough_items(items: &[ResolvedProjection]) -> Vec<ResolvedProjection> {
593    items
594        .iter()
595        .map(|item| ResolvedProjection {
596            expr: ResolvedExpr::Variable(item.output),
597            output: item.output,
598            name: item.name.clone(),
599            explicit_alias: item.explicit_alias,
600            span: item.span,
601        })
602        .collect()
603}
604
605/// Rewrite sort keys for a sort that runs after aggregation / DISTINCT,
606/// where only the projected columns exist. A key that restates a
607/// projected expression (`RETURN n.v AS v, count(*) AS c ORDER BY
608/// count(*)`) is pointed at that expression's output column. Alias keys
609/// already resolve to output columns in the analyzer.
610///
611/// Expressions are compared by their derived `Debug` form: structural,
612/// deterministic, and only paid once at planning time.
613fn sort_keys_on_outputs(
614    order: &[ResolvedSortItem],
615    items: &[ResolvedProjection],
616) -> Vec<ResolvedSortItem> {
617    let projected: Vec<(String, VarId)> = items
618        .iter()
619        .map(|item| (format!("{:?}", item.expr), item.output))
620        .collect();
621    order
622        .iter()
623        .map(|key| {
624            let shape = format!("{:?}", key.expr);
625            match projected.iter().find(|(expr, _)| *expr == shape) {
626                Some((_, output)) => ResolvedSortItem {
627                    expr: ResolvedExpr::Variable(*output),
628                    direction: key.direction,
629                },
630                None => key.clone(),
631            }
632        })
633        .collect()
634}
635
636/// The VarIds a pattern binds (path, node and relationship variables); unlike
637/// `ResolvedPattern::collect_vars`, not the variables its expressions read.
638fn pattern_binders(pattern: &ResolvedPattern) -> Vec<VarId> {
639    let mut vars = Vec::new();
640    for part in &pattern.parts {
641        if let Some(v) = part.binding {
642            vars.push(v);
643        }
644        match &part.element {
645            ResolvedPatternElement::Node { var, .. } => {
646                if let Some(v) = var {
647                    vars.push(*v);
648                }
649            }
650            ResolvedPatternElement::ShortestPath { head, chain, .. }
651            | ResolvedPatternElement::NodeChain { head, chain } => {
652                if let Some(v) = head.var {
653                    vars.push(v);
654                }
655                for step in chain {
656                    if let Some(v) = step.rel.var {
657                        vars.push(v);
658                    }
659                    if let Some(v) = step.node.var {
660                        vars.push(v);
661                    }
662                }
663            }
664        }
665    }
666    vars
667}
668
669fn expr_contains_aggregate(expr: &ResolvedExpr) -> bool {
670    match expr {
671        ResolvedExpr::Function { function, args, .. } => {
672            if function.is_aggregate() {
673                return true;
674            }
675            args.iter().any(expr_contains_aggregate)
676        }
677        ResolvedExpr::Property { expr, .. } => expr_contains_aggregate(expr),
678        ResolvedExpr::Binary { lhs, rhs, .. } => {
679            expr_contains_aggregate(lhs) || expr_contains_aggregate(rhs)
680        }
681        ResolvedExpr::Unary { expr, .. } => expr_contains_aggregate(expr),
682        ResolvedExpr::List(items) => items.iter().any(expr_contains_aggregate),
683        ResolvedExpr::Map(items) => items.iter().any(|(_, v)| expr_contains_aggregate(v)),
684        ResolvedExpr::Case {
685            input,
686            alternatives,
687            else_expr,
688        } => {
689            input.as_ref().is_some_and(|e| expr_contains_aggregate(e))
690                || alternatives
691                    .iter()
692                    .any(|(w, t)| expr_contains_aggregate(w) || expr_contains_aggregate(t))
693                || else_expr
694                    .as_ref()
695                    .is_some_and(|e| expr_contains_aggregate(e))
696        }
697        _ => false,
698    }
699}