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