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 stats: GraphStats,
20 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 fn track_bindings(&mut self, clause: &ResolvedClause) {
88 match clause {
89 ResolvedClause::Match(m) => self.bound.extend(pattern_binders(&m.pattern)),
90 ResolvedClause::Create(c) => self.bound.extend(pattern_binders(&c.pattern)),
91 ResolvedClause::Merge(m) => {
92 let pattern = ResolvedPattern {
93 parts: vec![m.pattern_part.clone()],
94 };
95 self.bound.extend(pattern_binders(&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 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 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 m.optional {
190 let upstream = input.unwrap_or_else(|| self.plan_unit_input());
197
198 let new_vars = pattern_binders(&m.pattern);
200
201 let mut pattern_planner = PatternPlanner::new(self);
206 let (inner, residual) =
207 pattern_planner.plan_pattern_with_where(None, &m.pattern, m.where_.as_ref());
208 let inner = self.push_conjuncts(inner, residual);
209
210 self.push(LogicalOp::OptionalMatch(OptionalMatch {
211 input: upstream,
212 inner,
213 new_vars,
214 }))
215 } else {
216 let mut pattern_planner = PatternPlanner::new(self);
217 let (node, residual) =
218 pattern_planner.plan_pattern_with_where(input, &m.pattern, m.where_.as_ref());
219 self.push_conjuncts(node, residual)
220 }
221 }
222
223 fn push_conjuncts(&mut self, input: PlanNodeId, conjuncts: Vec<ResolvedExpr>) -> PlanNodeId {
226 let predicate = conjuncts
227 .into_iter()
228 .reduce(|acc, next| ResolvedExpr::Binary {
229 lhs: Box::new(acc),
230 op: lora_ast::BinaryOp::And,
231 rhs: Box::new(next),
232 });
233 match predicate {
234 Some(predicate) => self.push(LogicalOp::Filter(Filter { input, predicate })),
235 None => input,
236 }
237 }
238
239 fn plan_unwind(&mut self, input: PlanNodeId, u: &ResolvedUnwind) -> PlanNodeId {
240 self.push(LogicalOp::Unwind(Unwind {
241 input,
242 expr: u.expr.clone(),
243 alias: u.alias,
244 }))
245 }
246
247 fn plan_create(&mut self, input: PlanNodeId, c: &ResolvedCreate) -> PlanNodeId {
248 self.push(LogicalOp::Create(crate::Create {
249 input,
250 pattern: c.pattern.clone(),
251 }))
252 }
253
254 fn plan_merge(&mut self, input: PlanNodeId, m: &ResolvedMerge) -> PlanNodeId {
255 self.push(LogicalOp::Merge(crate::Merge {
256 input,
257 pattern_part: m.pattern_part.clone(),
258 actions: m.actions.clone(),
259 }))
260 }
261
262 fn plan_delete(&mut self, input: PlanNodeId, d: &ResolvedDelete) -> PlanNodeId {
263 self.push(LogicalOp::Delete(crate::Delete {
264 input,
265 detach: d.detach,
266 expressions: d.expressions.clone(),
267 }))
268 }
269
270 fn plan_set(&mut self, input: PlanNodeId, s: &ResolvedSet) -> PlanNodeId {
271 self.push(LogicalOp::Set(crate::Set {
272 input,
273 items: s.items.clone(),
274 }))
275 }
276
277 fn plan_remove(&mut self, input: PlanNodeId, r: &ResolvedRemove) -> PlanNodeId {
278 self.push(LogicalOp::Remove(crate::Remove {
279 input,
280 items: r.items.clone(),
281 }))
282 }
283
284 fn plan_foreach(&mut self, input: PlanNodeId, f: &ResolvedForeach) -> PlanNodeId {
285 self.push(LogicalOp::Foreach(crate::Foreach {
286 input,
287 variable: f.variable,
288 list: f.list.clone(),
289 body: f.body.clone(),
290 }))
291 }
292
293 fn plan_with(&mut self, input: PlanNodeId, with: &ResolvedWith) -> PlanNodeId {
294 let mut node = self.plan_projection_sort_limit(
295 input,
296 &with.items,
297 with.distinct,
298 with.include_existing,
299 &with.order,
300 &with.skip,
301 &with.limit,
302 );
303
304 if let Some(pred) = &with.where_ {
305 node = self.push(LogicalOp::Filter(Filter {
306 input: node,
307 predicate: pred.clone(),
308 }));
309 }
310
311 node
312 }
313
314 fn plan_return(&mut self, input: PlanNodeId, ret: &ResolvedReturn) -> PlanNodeId {
315 self.plan_projection_sort_limit(
316 input,
317 &ret.items,
318 ret.distinct,
319 ret.include_existing,
320 &ret.order,
321 &ret.skip,
322 &ret.limit,
323 )
324 }
325
326 #[allow(clippy::too_many_arguments)]
344 fn plan_projection_sort_limit(
345 &mut self,
346 input: PlanNodeId,
347 items: &[ResolvedProjection],
348 distinct: bool,
349 include_existing: bool,
350 order: &[ResolvedSortItem],
351 skip: &Option<ResolvedExpr>,
352 limit: &Option<ResolvedExpr>,
353 ) -> PlanNodeId {
354 let aggregates = items.iter().any(|item| expr_contains_aggregate(&item.expr));
355 let has_order = !order.is_empty();
356 let has_limit = skip.is_some() || limit.is_some();
357
358 if aggregates || distinct {
359 let mut node =
360 self.plan_projection_or_aggregation(input, items, distinct, include_existing);
361 if has_order {
362 node = self.push(LogicalOp::Sort(Sort {
363 input: node,
364 items: sort_keys_on_outputs(order, items),
365 top_k: None,
366 limit: None,
367 }));
368 }
369 if has_limit {
370 node = self.push(LogicalOp::Limit(Limit {
371 input: node,
372 skip: skip.clone(),
373 limit: limit.clone(),
374 }));
375 }
376 return node;
377 }
378
379 if !has_order {
380 let mut node = input;
383 if has_limit {
384 node = self.push(LogicalOp::Limit(Limit {
385 input: node,
386 skip: skip.clone(),
387 limit: limit.clone(),
388 }));
389 }
390 return self.plan_projection_or_aggregation(node, items, false, include_existing);
391 }
392
393 let mut node = self.push(LogicalOp::Projection(Projection {
394 input,
395 distinct: false,
396 items: items.to_vec(),
397 include_existing: true,
398 }));
399 node = self.push(LogicalOp::Sort(Sort {
400 input: node,
401 items: order.to_vec(),
402 top_k: None,
403 limit: None,
404 }));
405 if has_limit {
406 node = self.push(LogicalOp::Limit(Limit {
407 input: node,
408 skip: skip.clone(),
409 limit: limit.clone(),
410 }));
411 }
412 if include_existing {
413 return node;
415 }
416 self.push(LogicalOp::Projection(Projection {
417 input: node,
418 distinct: false,
419 items: passthrough_items(items),
420 include_existing: false,
421 }))
422 }
423
424 fn plan_projection_or_aggregation(
428 &mut self,
429 input: PlanNodeId,
430 items: &[ResolvedProjection],
431 distinct: bool,
432 include_existing: bool,
433 ) -> PlanNodeId {
434 let has_aggregates = items.iter().any(|item| expr_contains_aggregate(&item.expr));
435
436 if !has_aggregates {
437 return self.push(LogicalOp::Projection(Projection {
438 input,
439 distinct,
440 items: items.to_vec(),
441 include_existing,
442 }));
443 }
444
445 let mut group_by = Vec::new();
447 let mut aggregates = Vec::new();
448
449 for item in items {
450 if expr_contains_aggregate(&item.expr) {
451 aggregates.push(item.clone());
452 } else {
453 group_by.push(item.clone());
454 }
455 }
456
457 let node = self.push(LogicalOp::Aggregation(Aggregation {
458 input,
459 group_by: group_by.clone(),
460 aggregates: aggregates.clone(),
461 }));
462
463 if distinct {
472 self.push(LogicalOp::Projection(Projection {
474 input: node,
475 distinct: true,
476 items: passthrough_items(items),
477 include_existing: false,
478 }))
479 } else {
480 node
481 }
482 }
483
484 fn plan_unit_input(&mut self) -> PlanNodeId {
485 self.push(LogicalOp::Argument(Argument))
486 }
487}
488
489fn passthrough_items(items: &[ResolvedProjection]) -> Vec<ResolvedProjection> {
491 items
492 .iter()
493 .map(|item| ResolvedProjection {
494 expr: ResolvedExpr::Variable(item.output),
495 output: item.output,
496 name: item.name.clone(),
497 explicit_alias: item.explicit_alias,
498 span: item.span,
499 })
500 .collect()
501}
502
503fn sort_keys_on_outputs(
512 order: &[ResolvedSortItem],
513 items: &[ResolvedProjection],
514) -> Vec<ResolvedSortItem> {
515 let projected: Vec<(String, VarId)> = items
516 .iter()
517 .map(|item| (format!("{:?}", item.expr), item.output))
518 .collect();
519 order
520 .iter()
521 .map(|key| {
522 let shape = format!("{:?}", key.expr);
523 match projected.iter().find(|(expr, _)| *expr == shape) {
524 Some((_, output)) => ResolvedSortItem {
525 expr: ResolvedExpr::Variable(*output),
526 direction: key.direction,
527 },
528 None => key.clone(),
529 }
530 })
531 .collect()
532}
533
534fn pattern_binders(pattern: &ResolvedPattern) -> Vec<VarId> {
537 let mut vars = Vec::new();
538 for part in &pattern.parts {
539 if let Some(v) = part.binding {
540 vars.push(v);
541 }
542 match &part.element {
543 ResolvedPatternElement::Node { var, .. } => {
544 if let Some(v) = var {
545 vars.push(*v);
546 }
547 }
548 ResolvedPatternElement::ShortestPath { head, chain, .. }
549 | ResolvedPatternElement::NodeChain { head, chain } => {
550 if let Some(v) = head.var {
551 vars.push(v);
552 }
553 for step in chain {
554 if let Some(v) = step.rel.var {
555 vars.push(v);
556 }
557 if let Some(v) = step.node.var {
558 vars.push(v);
559 }
560 }
561 }
562 }
563 }
564 vars
565}
566
567fn expr_contains_aggregate(expr: &ResolvedExpr) -> bool {
568 match expr {
569 ResolvedExpr::Function { function, args, .. } => {
570 if function.is_aggregate() {
571 return true;
572 }
573 args.iter().any(expr_contains_aggregate)
574 }
575 ResolvedExpr::Property { expr, .. } => expr_contains_aggregate(expr),
576 ResolvedExpr::Binary { lhs, rhs, .. } => {
577 expr_contains_aggregate(lhs) || expr_contains_aggregate(rhs)
578 }
579 ResolvedExpr::Unary { expr, .. } => expr_contains_aggregate(expr),
580 ResolvedExpr::List(items) => items.iter().any(expr_contains_aggregate),
581 ResolvedExpr::Map(items) => items.iter().any(|(_, v)| expr_contains_aggregate(v)),
582 ResolvedExpr::Case {
583 input,
584 alternatives,
585 else_expr,
586 } => {
587 input.as_ref().is_some_and(|e| expr_contains_aggregate(e))
588 || alternatives
589 .iter()
590 .any(|(w, t)| expr_contains_aggregate(w) || expr_contains_aggregate(t))
591 || else_expr
592 .as_ref()
593 .is_some_and(|e| expr_contains_aggregate(e))
594 }
595 _ => false,
596 }
597}