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        };
176        // The inner query sees the outer row; its own WITH / RETURN must
177        // not change what the outer query considers bound.
178        let outer_bound = self.bound.clone();
179        let inner = self.plan_query(&inner_query);
180        self.bound = outer_bound;
181        self.push(LogicalOp::CallSubquery(crate::logical::CallSubquery {
182            input,
183            inner,
184            new_vars: call.return_vars.clone(),
185        }))
186    }
187
188    fn plan_match(&mut self, input: Option<PlanNodeId>, m: &ResolvedMatch) -> PlanNodeId {
189        if m.optional {
190            // OPTIONAL MATCH: build the inner sub-plan that reads from Argument,
191            // then wrap it in an OptionalMatch node that provides null-extension.
192            //
193            // A leading OPTIONAL MATCH has no upstream clause; openCypher
194            // drives it from the single empty unit row, so an empty match
195            // still yields one row with the pattern's variables null.
196            let upstream = input.unwrap_or_else(|| self.plan_unit_input());
197
198            // Collect variables introduced by this pattern (for null-extension).
199            let new_vars = pattern_binders(&m.pattern);
200
201            // Build inner match plan WITHOUT the upstream input — the executor
202            // will inject each upstream row individually. The WHERE belongs
203            // to the OPTIONAL MATCH, so its conjuncts are only ever placed
204            // inside this inner plan, never above the OptionalMatch.
205            let mut pattern_planner = PatternPlanner::new(self);
206            let (inner, residual) =
207                pattern_planner.plan_pattern_with_where(None, &m.pattern, m.where_.as_ref());
208            let inner = self.push_conjuncts(inner, residual);
209
210            self.push(LogicalOp::OptionalMatch(OptionalMatch {
211                input: upstream,
212                inner,
213                new_vars,
214            }))
215        } else {
216            let mut pattern_planner = PatternPlanner::new(self);
217            let (node, residual) =
218                pattern_planner.plan_pattern_with_where(input, &m.pattern, m.where_.as_ref());
219            self.push_conjuncts(node, residual)
220        }
221    }
222
223    /// A `Filter` over `input` holding `conjuncts` ANDed in order, or
224    /// `input` itself when there are none.
225    fn push_conjuncts(&mut self, input: PlanNodeId, conjuncts: Vec<ResolvedExpr>) -> PlanNodeId {
226        let predicate = conjuncts
227            .into_iter()
228            .reduce(|acc, next| ResolvedExpr::Binary {
229                lhs: Box::new(acc),
230                op: lora_ast::BinaryOp::And,
231                rhs: Box::new(next),
232            });
233        match predicate {
234            Some(predicate) => self.push(LogicalOp::Filter(Filter { input, predicate })),
235            None => input,
236        }
237    }
238
239    fn plan_unwind(&mut self, input: PlanNodeId, u: &ResolvedUnwind) -> PlanNodeId {
240        self.push(LogicalOp::Unwind(Unwind {
241            input,
242            expr: u.expr.clone(),
243            alias: u.alias,
244        }))
245    }
246
247    fn plan_create(&mut self, input: PlanNodeId, c: &ResolvedCreate) -> PlanNodeId {
248        self.push(LogicalOp::Create(crate::Create {
249            input,
250            pattern: c.pattern.clone(),
251        }))
252    }
253
254    fn plan_merge(&mut self, input: PlanNodeId, m: &ResolvedMerge) -> PlanNodeId {
255        self.push(LogicalOp::Merge(crate::Merge {
256            input,
257            pattern_part: m.pattern_part.clone(),
258            actions: m.actions.clone(),
259        }))
260    }
261
262    fn plan_delete(&mut self, input: PlanNodeId, d: &ResolvedDelete) -> PlanNodeId {
263        self.push(LogicalOp::Delete(crate::Delete {
264            input,
265            detach: d.detach,
266            expressions: d.expressions.clone(),
267        }))
268    }
269
270    fn plan_set(&mut self, input: PlanNodeId, s: &ResolvedSet) -> PlanNodeId {
271        self.push(LogicalOp::Set(crate::Set {
272            input,
273            items: s.items.clone(),
274        }))
275    }
276
277    fn plan_remove(&mut self, input: PlanNodeId, r: &ResolvedRemove) -> PlanNodeId {
278        self.push(LogicalOp::Remove(crate::Remove {
279            input,
280            items: r.items.clone(),
281        }))
282    }
283
284    fn plan_foreach(&mut self, input: PlanNodeId, f: &ResolvedForeach) -> PlanNodeId {
285        self.push(LogicalOp::Foreach(crate::Foreach {
286            input,
287            variable: f.variable,
288            list: f.list.clone(),
289            body: f.body.clone(),
290        }))
291    }
292
293    fn plan_with(&mut self, input: PlanNodeId, with: &ResolvedWith) -> PlanNodeId {
294        let mut node = self.plan_projection_sort_limit(
295            input,
296            &with.items,
297            with.distinct,
298            with.include_existing,
299            &with.order,
300            &with.skip,
301            &with.limit,
302        );
303
304        if let Some(pred) = &with.where_ {
305            node = self.push(LogicalOp::Filter(Filter {
306                input: node,
307                predicate: pred.clone(),
308            }));
309        }
310
311        node
312    }
313
314    fn plan_return(&mut self, input: PlanNodeId, ret: &ResolvedReturn) -> PlanNodeId {
315        self.plan_projection_sort_limit(
316            input,
317            &ret.items,
318            ret.distinct,
319            ret.include_existing,
320            &ret.order,
321            &ret.skip,
322            &ret.limit,
323        )
324    }
325
326    /// Plan `items [ORDER BY ...] [SKIP ...] [LIMIT ...]` for WITH / RETURN.
327    ///
328    /// Cypher evaluates projection (with aggregation and DISTINCT) first,
329    /// then ORDER BY, then SKIP / LIMIT. Sort keys may name projected
330    /// aliases (`RETURN p.name AS name ORDER BY name`) and, unless the
331    /// projection aggregates or deduplicates, pre-projection variables too
332    /// (`RETURN p.name AS name ORDER BY p.age`).
333    ///
334    /// * Aggregation or DISTINCT changes the row set, so both must run
335    ///   before sorting and limiting; otherwise `LIMIT 1` would cut the
336    ///   input to one row before counting, and DISTINCT would dedupe an
337    ///   already-truncated list. Sort keys that restate a projected
338    ///   expression are pointed at that expression's output column.
339    /// * A plain projection maps rows 1:1. It first projects while keeping
340    ///   the input bindings, so keys can use both aliases and original
341    ///   variables, then sorts and limits, then trims the row down to the
342    ///   projected columns.
343    #[allow(clippy::too_many_arguments)]
344    fn plan_projection_sort_limit(
345        &mut self,
346        input: PlanNodeId,
347        items: &[ResolvedProjection],
348        distinct: bool,
349        include_existing: bool,
350        order: &[ResolvedSortItem],
351        skip: &Option<ResolvedExpr>,
352        limit: &Option<ResolvedExpr>,
353    ) -> PlanNodeId {
354        let aggregates = items.iter().any(|item| expr_contains_aggregate(&item.expr));
355        let has_order = !order.is_empty();
356        let has_limit = skip.is_some() || limit.is_some();
357
358        if aggregates || distinct {
359            let mut node =
360                self.plan_projection_or_aggregation(input, items, distinct, include_existing);
361            if has_order {
362                node = self.push(LogicalOp::Sort(Sort {
363                    input: node,
364                    items: sort_keys_on_outputs(order, items),
365                    top_k: None,
366                    limit: None,
367                }));
368            }
369            if has_limit {
370                node = self.push(LogicalOp::Limit(Limit {
371                    input: node,
372                    skip: skip.clone(),
373                    limit: limit.clone(),
374                }));
375            }
376            return node;
377        }
378
379        if !has_order {
380            // LIMIT on a 1:1 projection can run first and saves projecting
381            // rows that would be dropped.
382            let mut node = input;
383            if has_limit {
384                node = self.push(LogicalOp::Limit(Limit {
385                    input: node,
386                    skip: skip.clone(),
387                    limit: limit.clone(),
388                }));
389            }
390            return self.plan_projection_or_aggregation(node, items, false, include_existing);
391        }
392
393        let mut node = self.push(LogicalOp::Projection(Projection {
394            input,
395            distinct: false,
396            items: items.to_vec(),
397            include_existing: true,
398        }));
399        node = self.push(LogicalOp::Sort(Sort {
400            input: node,
401            items: order.to_vec(),
402            top_k: None,
403            limit: None,
404        }));
405        if has_limit {
406            node = self.push(LogicalOp::Limit(Limit {
407                input: node,
408                skip: skip.clone(),
409                limit: limit.clone(),
410            }));
411        }
412        if include_existing {
413            // WITH * / RETURN * keep every binding anyway.
414            return node;
415        }
416        self.push(LogicalOp::Projection(Projection {
417            input: node,
418            distinct: false,
419            items: passthrough_items(items),
420            include_existing: false,
421        }))
422    }
423
424    /// If any projection item contains an aggregate function, emit an
425    /// Aggregation node followed by a Projection. Otherwise emit a plain
426    /// Projection.
427    fn plan_projection_or_aggregation(
428        &mut self,
429        input: PlanNodeId,
430        items: &[ResolvedProjection],
431        distinct: bool,
432        include_existing: bool,
433    ) -> PlanNodeId {
434        let has_aggregates = items.iter().any(|item| expr_contains_aggregate(&item.expr));
435
436        if !has_aggregates {
437            return self.push(LogicalOp::Projection(Projection {
438                input,
439                distinct,
440                items: items.to_vec(),
441                include_existing,
442            }));
443        }
444
445        // Split items into group-by keys and aggregate expressions.
446        let mut group_by = Vec::new();
447        let mut aggregates = Vec::new();
448
449        for item in items {
450            if expr_contains_aggregate(&item.expr) {
451                aggregates.push(item.clone());
452            } else {
453                group_by.push(item.clone());
454            }
455        }
456
457        let node = self.push(LogicalOp::Aggregation(Aggregation {
458            input,
459            group_by: group_by.clone(),
460            aggregates: aggregates.clone(),
461        }));
462
463        // After aggregation the row already contains the right VarIds and names,
464        // but we still emit a Projection to handle DISTINCT and to ensure the
465        // final column order matches the original item list. The projection uses
466        // include_existing=true so it picks up the aggregation output, and each
467        // item just reads its own output variable.
468        //
469        // However, since the aggregation node already produces correctly-named
470        // rows, we can skip the extra projection when not needed.
471        if distinct {
472            // For DISTINCT we still need the dedup pass in exec_projection.
473            self.push(LogicalOp::Projection(Projection {
474                input: node,
475                distinct: true,
476                items: passthrough_items(items),
477                include_existing: false,
478            }))
479        } else {
480            node
481        }
482    }
483
484    fn plan_unit_input(&mut self) -> PlanNodeId {
485        self.push(LogicalOp::Argument(Argument))
486    }
487}
488
489/// Projection items that re-emit each item's own output column.
490fn passthrough_items(items: &[ResolvedProjection]) -> Vec<ResolvedProjection> {
491    items
492        .iter()
493        .map(|item| ResolvedProjection {
494            expr: ResolvedExpr::Variable(item.output),
495            output: item.output,
496            name: item.name.clone(),
497            explicit_alias: item.explicit_alias,
498            span: item.span,
499        })
500        .collect()
501}
502
503/// Rewrite sort keys for a sort that runs after aggregation / DISTINCT,
504/// where only the projected columns exist. A key that restates a
505/// projected expression (`RETURN n.v AS v, count(*) AS c ORDER BY
506/// count(*)`) is pointed at that expression's output column. Alias keys
507/// already resolve to output columns in the analyzer.
508///
509/// Expressions are compared by their derived `Debug` form: structural,
510/// deterministic, and only paid once at planning time.
511fn sort_keys_on_outputs(
512    order: &[ResolvedSortItem],
513    items: &[ResolvedProjection],
514) -> Vec<ResolvedSortItem> {
515    let projected: Vec<(String, VarId)> = items
516        .iter()
517        .map(|item| (format!("{:?}", item.expr), item.output))
518        .collect();
519    order
520        .iter()
521        .map(|key| {
522            let shape = format!("{:?}", key.expr);
523            match projected.iter().find(|(expr, _)| *expr == shape) {
524                Some((_, output)) => ResolvedSortItem {
525                    expr: ResolvedExpr::Variable(*output),
526                    direction: key.direction,
527                },
528                None => key.clone(),
529            }
530        })
531        .collect()
532}
533
534/// The VarIds a pattern binds (path, node and relationship variables); unlike
535/// `ResolvedPattern::collect_vars`, not the variables its expressions read.
536fn pattern_binders(pattern: &ResolvedPattern) -> Vec<VarId> {
537    let mut vars = Vec::new();
538    for part in &pattern.parts {
539        if let Some(v) = part.binding {
540            vars.push(v);
541        }
542        match &part.element {
543            ResolvedPatternElement::Node { var, .. } => {
544                if let Some(v) = var {
545                    vars.push(*v);
546                }
547            }
548            ResolvedPatternElement::ShortestPath { head, chain, .. }
549            | ResolvedPatternElement::NodeChain { head, chain } => {
550                if let Some(v) = head.var {
551                    vars.push(v);
552                }
553                for step in chain {
554                    if let Some(v) = step.rel.var {
555                        vars.push(v);
556                    }
557                    if let Some(v) = step.node.var {
558                        vars.push(v);
559                    }
560                }
561            }
562        }
563    }
564    vars
565}
566
567fn expr_contains_aggregate(expr: &ResolvedExpr) -> bool {
568    match expr {
569        ResolvedExpr::Function { function, args, .. } => {
570            if function.is_aggregate() {
571                return true;
572            }
573            args.iter().any(expr_contains_aggregate)
574        }
575        ResolvedExpr::Property { expr, .. } => expr_contains_aggregate(expr),
576        ResolvedExpr::Binary { lhs, rhs, .. } => {
577            expr_contains_aggregate(lhs) || expr_contains_aggregate(rhs)
578        }
579        ResolvedExpr::Unary { expr, .. } => expr_contains_aggregate(expr),
580        ResolvedExpr::List(items) => items.iter().any(expr_contains_aggregate),
581        ResolvedExpr::Map(items) => items.iter().any(|(_, v)| expr_contains_aggregate(v)),
582        ResolvedExpr::Case {
583            input,
584            alternatives,
585            else_expr,
586        } => {
587            input.as_ref().is_some_and(|e| expr_contains_aggregate(e))
588                || alternatives
589                    .iter()
590                    .any(|(w, t)| expr_contains_aggregate(w) || expr_contains_aggregate(t))
591                || else_expr
592                    .as_ref()
593                    .is_some_and(|e| expr_contains_aggregate(e))
594        }
595        _ => false,
596    }
597}