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