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