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