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 parameters: Default::default(),
176 };
177 let outer_bound = self.bound.clone();
180 let inner = self.plan_query(&inner_query);
181 self.bound = outer_bound;
182 self.push(LogicalOp::CallSubquery(crate::logical::CallSubquery {
183 input,
184 inner,
185 new_vars: call.return_vars.clone(),
186 }))
187 }
188
189 fn plan_match(&mut self, input: Option<PlanNodeId>, m: &ResolvedMatch) -> PlanNodeId {
190 if m.optional {
191 let upstream = input.unwrap_or_else(|| self.plan_unit_input());
198
199 let new_vars = pattern_binders(&m.pattern);
201
202 let mut pattern_planner = PatternPlanner::new(self);
207 let (inner, residual) =
208 pattern_planner.plan_pattern_with_where(None, &m.pattern, m.where_.as_ref());
209 let inner = self.push_conjuncts(inner, residual);
210
211 self.push(LogicalOp::OptionalMatch(OptionalMatch {
212 input: upstream,
213 inner,
214 new_vars,
215 }))
216 } else {
217 let mut pattern_planner = PatternPlanner::new(self);
218 let (node, residual) =
219 pattern_planner.plan_pattern_with_where(input, &m.pattern, m.where_.as_ref());
220 self.push_conjuncts(node, residual)
221 }
222 }
223
224 fn push_conjuncts(&mut self, input: PlanNodeId, conjuncts: Vec<ResolvedExpr>) -> PlanNodeId {
227 let predicate = conjuncts
228 .into_iter()
229 .reduce(|acc, next| ResolvedExpr::Binary {
230 lhs: Box::new(acc),
231 op: lora_ast::BinaryOp::And,
232 rhs: Box::new(next),
233 });
234 match predicate {
235 Some(predicate) => self.push(LogicalOp::Filter(Filter { input, predicate })),
236 None => input,
237 }
238 }
239
240 fn plan_unwind(&mut self, input: PlanNodeId, u: &ResolvedUnwind) -> PlanNodeId {
241 self.push(LogicalOp::Unwind(Unwind {
242 input,
243 expr: u.expr.clone(),
244 alias: u.alias,
245 }))
246 }
247
248 fn plan_create(&mut self, input: PlanNodeId, c: &ResolvedCreate) -> PlanNodeId {
249 self.push(LogicalOp::Create(crate::Create {
250 input,
251 pattern: c.pattern.clone(),
252 }))
253 }
254
255 fn plan_merge(&mut self, input: PlanNodeId, m: &ResolvedMerge) -> PlanNodeId {
256 self.push(LogicalOp::Merge(crate::Merge {
257 input,
258 pattern_part: m.pattern_part.clone(),
259 actions: m.actions.clone(),
260 }))
261 }
262
263 fn plan_delete(&mut self, input: PlanNodeId, d: &ResolvedDelete) -> PlanNodeId {
264 self.push(LogicalOp::Delete(crate::Delete {
265 input,
266 detach: d.detach,
267 expressions: d.expressions.clone(),
268 }))
269 }
270
271 fn plan_set(&mut self, input: PlanNodeId, s: &ResolvedSet) -> PlanNodeId {
272 self.push(LogicalOp::Set(crate::Set {
273 input,
274 items: s.items.clone(),
275 }))
276 }
277
278 fn plan_remove(&mut self, input: PlanNodeId, r: &ResolvedRemove) -> PlanNodeId {
279 self.push(LogicalOp::Remove(crate::Remove {
280 input,
281 items: r.items.clone(),
282 }))
283 }
284
285 fn plan_foreach(&mut self, input: PlanNodeId, f: &ResolvedForeach) -> PlanNodeId {
286 self.push(LogicalOp::Foreach(crate::Foreach {
287 input,
288 variable: f.variable,
289 list: f.list.clone(),
290 body: f.body.clone(),
291 }))
292 }
293
294 fn plan_with(&mut self, input: PlanNodeId, with: &ResolvedWith) -> PlanNodeId {
295 let mut node = self.plan_projection_sort_limit(
296 input,
297 &with.items,
298 &with.lifted_aggregates,
299 with.distinct,
300 with.include_existing,
301 &with.order,
302 &with.skip,
303 &with.limit,
304 );
305
306 if let Some(pred) = &with.where_ {
307 node = self.push(LogicalOp::Filter(Filter {
308 input: node,
309 predicate: pred.clone(),
310 }));
311 }
312
313 node
314 }
315
316 fn plan_return(&mut self, input: PlanNodeId, ret: &ResolvedReturn) -> PlanNodeId {
317 self.plan_projection_sort_limit(
318 input,
319 &ret.items,
320 &ret.lifted_aggregates,
321 ret.distinct,
322 ret.include_existing,
323 &ret.order,
324 &ret.skip,
325 &ret.limit,
326 )
327 }
328
329 #[allow(clippy::too_many_arguments)]
347 fn plan_projection_sort_limit(
348 &mut self,
349 input: PlanNodeId,
350 items: &[ResolvedProjection],
351 lifted: &[ResolvedProjection],
352 distinct: bool,
353 include_existing: bool,
354 order: &[ResolvedSortItem],
355 skip: &Option<ResolvedExpr>,
356 limit: &Option<ResolvedExpr>,
357 ) -> PlanNodeId {
358 let aggregates =
359 !lifted.is_empty() || items.iter().any(|item| expr_contains_aggregate(&item.expr));
360 let has_order = !order.is_empty();
361 let has_limit = skip.is_some() || limit.is_some();
362
363 if aggregates || distinct {
364 let sort_reads: BTreeSet<VarId> = order
367 .iter()
368 .flat_map(|key| {
369 let mut reads = BTreeSet::new();
370 key.expr.collect_vars(&mut reads);
371 reads
372 })
373 .collect();
374 let kept: Vec<ResolvedProjection> = lifted
375 .iter()
376 .filter(|p| sort_reads.contains(&p.output))
377 .cloned()
378 .collect();
379 let mut node = self.plan_projection_or_aggregation(
380 input,
381 items,
382 lifted,
383 &kept,
384 distinct,
385 include_existing,
386 );
387 if has_order {
388 node = self.push(LogicalOp::Sort(Sort {
389 input: node,
390 items: sort_keys_on_outputs(order, items),
391 top_k: None,
392 limit: None,
393 }));
394 }
395 if has_limit {
396 node = self.push(LogicalOp::Limit(Limit {
397 input: node,
398 skip: skip.clone(),
399 limit: limit.clone(),
400 }));
401 }
402 if !kept.is_empty() {
403 node = self.push(LogicalOp::Projection(Projection {
404 input: node,
405 distinct: false,
406 items: passthrough_items(items),
407 include_existing: false,
408 }));
409 }
410 return node;
411 }
412
413 if !has_order {
414 let mut node = input;
417 if has_limit {
418 node = self.push(LogicalOp::Limit(Limit {
419 input: node,
420 skip: skip.clone(),
421 limit: limit.clone(),
422 }));
423 }
424 return self.plan_projection_or_aggregation(
425 node,
426 items,
427 lifted,
428 &[],
429 false,
430 include_existing,
431 );
432 }
433
434 let mut node = self.push(LogicalOp::Projection(Projection {
435 input,
436 distinct: false,
437 items: items.to_vec(),
438 include_existing: true,
439 }));
440 node = self.push(LogicalOp::Sort(Sort {
441 input: node,
442 items: order.to_vec(),
443 top_k: None,
444 limit: None,
445 }));
446 if has_limit {
447 node = self.push(LogicalOp::Limit(Limit {
448 input: node,
449 skip: skip.clone(),
450 limit: limit.clone(),
451 }));
452 }
453 if include_existing {
454 return node;
456 }
457 self.push(LogicalOp::Projection(Projection {
458 input: node,
459 distinct: false,
460 items: passthrough_items(items),
461 include_existing: false,
462 }))
463 }
464
465 fn plan_projection_or_aggregation(
474 &mut self,
475 input: PlanNodeId,
476 items: &[ResolvedProjection],
477 lifted: &[ResolvedProjection],
478 kept: &[ResolvedProjection],
479 distinct: bool,
480 include_existing: bool,
481 ) -> PlanNodeId {
482 if !lifted.is_empty() {
483 return self.plan_lifted_aggregation(input, items, lifted, kept, distinct);
484 }
485
486 let has_aggregates = items.iter().any(|item| expr_contains_aggregate(&item.expr));
487
488 if !has_aggregates {
489 return self.push(LogicalOp::Projection(Projection {
490 input,
491 distinct,
492 items: items.to_vec(),
493 include_existing,
494 }));
495 }
496
497 let mut group_by = Vec::new();
499 let mut aggregates = Vec::new();
500
501 for item in items {
502 if expr_contains_aggregate(&item.expr) {
503 aggregates.push(item.clone());
504 } else {
505 group_by.push(item.clone());
506 }
507 }
508
509 let node = self.push(LogicalOp::Aggregation(Aggregation {
510 input,
511 group_by: group_by.clone(),
512 aggregates: aggregates.clone(),
513 }));
514
515 if distinct {
524 self.push(LogicalOp::Projection(Projection {
526 input: node,
527 distinct: true,
528 items: passthrough_items(items),
529 include_existing: false,
530 }))
531 } else {
532 node
533 }
534 }
535
536 fn plan_lifted_aggregation(
541 &mut self,
542 input: PlanNodeId,
543 items: &[ResolvedProjection],
544 lifted: &[ResolvedProjection],
545 kept: &[ResolvedProjection],
546 distinct: bool,
547 ) -> PlanNodeId {
548 let lifted_outputs: BTreeSet<VarId> = lifted.iter().map(|p| p.output).collect();
549 let reads_lifted = |item: &ResolvedProjection| {
550 let mut reads = BTreeSet::new();
551 item.expr.collect_vars(&mut reads);
552 !reads.is_disjoint(&lifted_outputs)
553 };
554
555 let mut group_by = Vec::new();
556 let mut aggregates = Vec::new();
557 let mut projected = Vec::with_capacity(items.len());
558 for item in items {
559 if reads_lifted(item) {
560 projected.push(item.clone());
561 continue;
562 }
563 if expr_contains_aggregate(&item.expr) {
564 aggregates.push(item.clone());
565 } else {
566 group_by.push(item.clone());
567 }
568 projected.extend(passthrough_items(std::slice::from_ref(item)));
569 }
570 aggregates.extend(lifted.iter().cloned());
571 projected.extend(passthrough_items(kept));
572
573 let node = self.push(LogicalOp::Aggregation(Aggregation {
574 input,
575 group_by,
576 aggregates,
577 }));
578 self.push(LogicalOp::Projection(Projection {
579 input: node,
580 distinct,
581 items: projected,
582 include_existing: false,
583 }))
584 }
585
586 fn plan_unit_input(&mut self) -> PlanNodeId {
587 self.push(LogicalOp::Argument(Argument))
588 }
589}
590
591fn passthrough_items(items: &[ResolvedProjection]) -> Vec<ResolvedProjection> {
593 items
594 .iter()
595 .map(|item| ResolvedProjection {
596 expr: ResolvedExpr::Variable(item.output),
597 output: item.output,
598 name: item.name.clone(),
599 explicit_alias: item.explicit_alias,
600 span: item.span,
601 })
602 .collect()
603}
604
605fn sort_keys_on_outputs(
614 order: &[ResolvedSortItem],
615 items: &[ResolvedProjection],
616) -> Vec<ResolvedSortItem> {
617 let projected: Vec<(String, VarId)> = items
618 .iter()
619 .map(|item| (format!("{:?}", item.expr), item.output))
620 .collect();
621 order
622 .iter()
623 .map(|key| {
624 let shape = format!("{:?}", key.expr);
625 match projected.iter().find(|(expr, _)| *expr == shape) {
626 Some((_, output)) => ResolvedSortItem {
627 expr: ResolvedExpr::Variable(*output),
628 direction: key.direction,
629 },
630 None => key.clone(),
631 }
632 })
633 .collect()
634}
635
636fn pattern_binders(pattern: &ResolvedPattern) -> Vec<VarId> {
639 let mut vars = Vec::new();
640 for part in &pattern.parts {
641 if let Some(v) = part.binding {
642 vars.push(v);
643 }
644 match &part.element {
645 ResolvedPatternElement::Node { var, .. } => {
646 if let Some(v) = var {
647 vars.push(*v);
648 }
649 }
650 ResolvedPatternElement::ShortestPath { head, chain, .. }
651 | ResolvedPatternElement::NodeChain { head, chain } => {
652 if let Some(v) = head.var {
653 vars.push(v);
654 }
655 for step in chain {
656 if let Some(v) = step.rel.var {
657 vars.push(v);
658 }
659 if let Some(v) = step.node.var {
660 vars.push(v);
661 }
662 }
663 }
664 }
665 }
666 vars
667}
668
669fn expr_contains_aggregate(expr: &ResolvedExpr) -> bool {
670 match expr {
671 ResolvedExpr::Function { function, args, .. } => {
672 if function.is_aggregate() {
673 return true;
674 }
675 args.iter().any(expr_contains_aggregate)
676 }
677 ResolvedExpr::Property { expr, .. } => expr_contains_aggregate(expr),
678 ResolvedExpr::Binary { lhs, rhs, .. } => {
679 expr_contains_aggregate(lhs) || expr_contains_aggregate(rhs)
680 }
681 ResolvedExpr::Unary { expr, .. } => expr_contains_aggregate(expr),
682 ResolvedExpr::List(items) => items.iter().any(expr_contains_aggregate),
683 ResolvedExpr::Map(items) => items.iter().any(|(_, v)| expr_contains_aggregate(v)),
684 ResolvedExpr::Case {
685 input,
686 alternatives,
687 else_expr,
688 } => {
689 input.as_ref().is_some_and(|e| expr_contains_aggregate(e))
690 || alternatives
691 .iter()
692 .any(|(w, t)| expr_contains_aggregate(w) || expr_contains_aggregate(t))
693 || else_expr
694 .as_ref()
695 .is_some_and(|e| expr_contains_aggregate(e))
696 }
697 _ => false,
698 }
699}