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 let new_vars = collect_pattern_vars(&m.pattern);
102
103 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 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 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 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 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 if distinct {
297 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
343fn 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}