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, ResolvedUnwind,
11    ResolvedWith,
12};
13
14pub struct Planner {
15    nodes: Vec<LogicalOp>,
16}
17
18impl Default for Planner {
19    fn default() -> Self {
20        Self::new()
21    }
22}
23
24impl Planner {
25    pub fn new() -> Self {
26        Self { nodes: Vec::new() }
27    }
28
29    pub(crate) fn push(&mut self, op: LogicalOp) -> PlanNodeId {
30        let id = self.nodes.len();
31        self.nodes.push(op);
32        id
33    }
34
35    pub fn plan(&mut self, query: &ResolvedQuery) -> LogicalPlan {
36        let root = self.plan_query(query);
37
38        LogicalPlan {
39            root,
40            nodes: std::mem::take(&mut self.nodes),
41        }
42    }
43
44    fn plan_query(&mut self, query: &ResolvedQuery) -> PlanNodeId {
45        let mut input = None;
46
47        for clause in &query.clauses {
48            input = Some(match clause {
49                ResolvedClause::Match(m) => self.plan_match(input, m),
50
51                ResolvedClause::Unwind(u) => {
52                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
53                    self.plan_unwind(upstream, u)
54                }
55
56                ResolvedClause::Create(c) => {
57                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
58                    self.plan_create(upstream, c)
59                }
60
61                ResolvedClause::Merge(m) => {
62                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
63                    self.plan_merge(upstream, m)
64                }
65
66                ResolvedClause::Delete(d) => {
67                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
68                    self.plan_delete(upstream, d)
69                }
70
71                ResolvedClause::Set(s) => {
72                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
73                    self.plan_set(upstream, s)
74                }
75
76                ResolvedClause::Remove(rm) => {
77                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
78                    self.plan_remove(upstream, rm)
79                }
80
81                ResolvedClause::Foreach(f) => {
82                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
83                    self.plan_foreach(upstream, f)
84                }
85
86                ResolvedClause::With(w) => {
87                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
88                    self.plan_with(upstream, w)
89                }
90
91                ResolvedClause::Return(r) => {
92                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
93                    self.plan_return(upstream, r)
94                }
95
96                ResolvedClause::CallSubquery(c) => {
97                    let upstream = input.unwrap_or_else(|| self.plan_unit_input());
98                    self.plan_call_subquery(upstream, c)
99                }
100            });
101        }
102
103        input.unwrap_or_else(|| self.plan_unit_input())
104    }
105
106    /// Plan a `CALL { ... }` subquery: build the inner plan starting
107    /// from a fresh `Argument`, then wrap it with `CallSubquery` to
108    /// drive it per outer row.
109    fn plan_call_subquery(&mut self, input: PlanNodeId, call: &ResolvedCallSubquery) -> PlanNodeId {
110        let inner_query = ResolvedQuery {
111            clauses: call.clauses.clone(),
112            unions: Vec::new(),
113        };
114        let inner = self.plan_query(&inner_query);
115        self.push(LogicalOp::CallSubquery(crate::logical::CallSubquery {
116            input,
117            inner,
118            new_vars: call.return_vars.clone(),
119        }))
120    }
121
122    fn plan_match(&mut self, input: Option<PlanNodeId>, m: &ResolvedMatch) -> PlanNodeId {
123        if let (true, Some(upstream)) = (m.optional, input) {
124            // OPTIONAL MATCH: build the inner sub-plan that reads from Argument,
125            // then wrap it in an OptionalMatch node that provides null-extension.
126
127            // Collect variables introduced by this pattern (for null-extension).
128            let new_vars = collect_pattern_vars(&m.pattern);
129
130            // Build inner match plan WITHOUT the upstream input — the executor
131            // will inject each upstream row individually.
132            let mut pattern_planner = PatternPlanner::new(self);
133            let mut inner = pattern_planner.plan_pattern(None, &m.pattern);
134
135            if let Some(pred) = &m.where_ {
136                inner = self.push(LogicalOp::Filter(Filter {
137                    input: inner,
138                    predicate: pred.clone(),
139                }));
140            }
141
142            self.push(LogicalOp::OptionalMatch(OptionalMatch {
143                input: upstream,
144                inner,
145                new_vars,
146            }))
147        } else {
148            let mut pattern_planner = PatternPlanner::new(self);
149            let mut node = pattern_planner.plan_pattern(input, &m.pattern);
150
151            if let Some(pred) = &m.where_ {
152                node = self.push(LogicalOp::Filter(Filter {
153                    input: node,
154                    predicate: pred.clone(),
155                }));
156            }
157
158            node
159        }
160    }
161
162    fn plan_unwind(&mut self, input: PlanNodeId, u: &ResolvedUnwind) -> PlanNodeId {
163        self.push(LogicalOp::Unwind(Unwind {
164            input,
165            expr: u.expr.clone(),
166            alias: u.alias,
167        }))
168    }
169
170    fn plan_create(&mut self, input: PlanNodeId, c: &ResolvedCreate) -> PlanNodeId {
171        self.push(LogicalOp::Create(crate::Create {
172            input,
173            pattern: c.pattern.clone(),
174        }))
175    }
176
177    fn plan_merge(&mut self, input: PlanNodeId, m: &ResolvedMerge) -> PlanNodeId {
178        self.push(LogicalOp::Merge(crate::Merge {
179            input,
180            pattern_part: m.pattern_part.clone(),
181            actions: m.actions.clone(),
182        }))
183    }
184
185    fn plan_delete(&mut self, input: PlanNodeId, d: &ResolvedDelete) -> PlanNodeId {
186        self.push(LogicalOp::Delete(crate::Delete {
187            input,
188            detach: d.detach,
189            expressions: d.expressions.clone(),
190        }))
191    }
192
193    fn plan_set(&mut self, input: PlanNodeId, s: &ResolvedSet) -> PlanNodeId {
194        self.push(LogicalOp::Set(crate::Set {
195            input,
196            items: s.items.clone(),
197        }))
198    }
199
200    fn plan_remove(&mut self, input: PlanNodeId, r: &ResolvedRemove) -> PlanNodeId {
201        self.push(LogicalOp::Remove(crate::Remove {
202            input,
203            items: r.items.clone(),
204        }))
205    }
206
207    fn plan_foreach(&mut self, input: PlanNodeId, f: &ResolvedForeach) -> PlanNodeId {
208        self.push(LogicalOp::Foreach(crate::Foreach {
209            input,
210            variable: f.variable,
211            list: f.list.clone(),
212            body: f.body.clone(),
213        }))
214    }
215
216    fn plan_with(&mut self, input: PlanNodeId, with: &ResolvedWith) -> PlanNodeId {
217        let mut node = input;
218
219        // Sort before projection so sort expressions can access original variables.
220        if !with.order.is_empty() {
221            node = self.push(LogicalOp::Sort(Sort {
222                input: node,
223                items: with.order.clone(),
224                top_k: None,
225            }));
226        }
227
228        if with.skip.is_some() || with.limit.is_some() {
229            node = self.push(LogicalOp::Limit(Limit {
230                input: node,
231                skip: with.skip.clone(),
232                limit: with.limit.clone(),
233            }));
234        }
235
236        node = self.plan_projection_or_aggregation(
237            node,
238            &with.items,
239            with.distinct,
240            with.include_existing,
241        );
242
243        if let Some(pred) = &with.where_ {
244            node = self.push(LogicalOp::Filter(Filter {
245                input: node,
246                predicate: pred.clone(),
247            }));
248        }
249
250        node
251    }
252
253    fn plan_return(&mut self, input: PlanNodeId, ret: &ResolvedReturn) -> PlanNodeId {
254        let mut node = input;
255
256        // Sort must happen BEFORE projection so that the sort expressions
257        // can access the original variables (e.g. n.name) which are not
258        // available after projection replaces the row with output VarIds.
259        if !ret.order.is_empty() {
260            node = self.push(LogicalOp::Sort(Sort {
261                input: node,
262                items: ret.order.clone(),
263                top_k: None,
264            }));
265        }
266
267        if ret.skip.is_some() || ret.limit.is_some() {
268            node = self.push(LogicalOp::Limit(Limit {
269                input: node,
270                skip: ret.skip.clone(),
271                limit: ret.limit.clone(),
272            }));
273        }
274
275        node = self.plan_projection_or_aggregation(
276            node,
277            &ret.items,
278            ret.distinct,
279            ret.include_existing,
280        );
281
282        node
283    }
284
285    /// If any projection item contains an aggregate function, emit an
286    /// Aggregation node followed by a Projection. Otherwise emit a plain
287    /// Projection.
288    fn plan_projection_or_aggregation(
289        &mut self,
290        input: PlanNodeId,
291        items: &[ResolvedProjection],
292        distinct: bool,
293        include_existing: bool,
294    ) -> PlanNodeId {
295        let has_aggregates = items.iter().any(|item| expr_contains_aggregate(&item.expr));
296
297        if !has_aggregates {
298            return self.push(LogicalOp::Projection(Projection {
299                input,
300                distinct,
301                items: items.to_vec(),
302                include_existing,
303            }));
304        }
305
306        // Split items into group-by keys and aggregate expressions.
307        let mut group_by = Vec::new();
308        let mut aggregates = Vec::new();
309
310        for item in items {
311            if expr_contains_aggregate(&item.expr) {
312                aggregates.push(item.clone());
313            } else {
314                group_by.push(item.clone());
315            }
316        }
317
318        let node = self.push(LogicalOp::Aggregation(Aggregation {
319            input,
320            group_by: group_by.clone(),
321            aggregates: aggregates.clone(),
322        }));
323
324        // After aggregation the row already contains the right VarIds and names,
325        // but we still emit a Projection to handle DISTINCT and to ensure the
326        // final column order matches the original item list. The projection uses
327        // include_existing=true so it picks up the aggregation output, and each
328        // item just reads its own output variable.
329        //
330        // However, since the aggregation node already produces correctly-named
331        // rows, we can skip the extra projection when not needed.
332        if distinct {
333            // For DISTINCT we still need the dedup pass in exec_projection.
334            let passthrough_items: Vec<ResolvedProjection> = items
335                .iter()
336                .map(|item| ResolvedProjection {
337                    expr: ResolvedExpr::Variable(item.output),
338                    output: item.output,
339                    name: item.name.clone(),
340                    explicit_alias: item.explicit_alias,
341                    span: item.span,
342                })
343                .collect();
344            self.push(LogicalOp::Projection(Projection {
345                input: node,
346                distinct: true,
347                items: passthrough_items,
348                include_existing: false,
349            }))
350        } else {
351            node
352        }
353    }
354
355    fn plan_unit_input(&mut self) -> PlanNodeId {
356        self.push(LogicalOp::Argument(Argument))
357    }
358}
359
360/// Collect all VarIds introduced by a pattern (node vars, relationship vars).
361fn collect_pattern_vars(pattern: &ResolvedPattern) -> Vec<VarId> {
362    let mut vars = Vec::new();
363    for part in &pattern.parts {
364        if let Some(v) = part.binding {
365            vars.push(v);
366        }
367        match &part.element {
368            ResolvedPatternElement::Node { var, .. } => {
369                if let Some(v) = var {
370                    vars.push(*v);
371                }
372            }
373            ResolvedPatternElement::ShortestPath { head, chain, .. }
374            | ResolvedPatternElement::NodeChain { head, chain } => {
375                if let Some(v) = head.var {
376                    vars.push(v);
377                }
378                for step in chain {
379                    if let Some(v) = step.rel.var {
380                        vars.push(v);
381                    }
382                    if let Some(v) = step.node.var {
383                        vars.push(v);
384                    }
385                }
386            }
387        }
388    }
389    vars
390}
391
392fn expr_contains_aggregate(expr: &ResolvedExpr) -> bool {
393    match expr {
394        ResolvedExpr::Function { function, args, .. } => {
395            if function.is_aggregate() {
396                return true;
397            }
398            args.iter().any(expr_contains_aggregate)
399        }
400        ResolvedExpr::Property { expr, .. } => expr_contains_aggregate(expr),
401        ResolvedExpr::Binary { lhs, rhs, .. } => {
402            expr_contains_aggregate(lhs) || expr_contains_aggregate(rhs)
403        }
404        ResolvedExpr::Unary { expr, .. } => expr_contains_aggregate(expr),
405        ResolvedExpr::List(items) => items.iter().any(expr_contains_aggregate),
406        ResolvedExpr::Map(items) => items.iter().any(|(_, v)| expr_contains_aggregate(v)),
407        ResolvedExpr::Case {
408            input,
409            alternatives,
410            else_expr,
411        } => {
412            input.as_ref().is_some_and(|e| expr_contains_aggregate(e))
413                || alternatives
414                    .iter()
415                    .any(|(w, t)| expr_contains_aggregate(w) || expr_contains_aggregate(t))
416                || else_expr
417                    .as_ref()
418                    .is_some_and(|e| expr_contains_aggregate(e))
419        }
420        _ => false,
421    }
422}