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(collect_pattern_vars(&m.pattern)),
90 ResolvedClause::Create(c) => self.bound.extend(collect_pattern_vars(&c.pattern)),
91 ResolvedClause::Merge(m) => {
92 let pattern = ResolvedPattern {
93 parts: vec![m.pattern_part.clone()],
94 };
95 self.bound.extend(collect_pattern_vars(&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 = collect_pattern_vars(&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 }));
362 }
363 if has_limit {
364 node = self.push(LogicalOp::Limit(Limit {
365 input: node,
366 skip: skip.clone(),
367 limit: limit.clone(),
368 }));
369 }
370 return node;
371 }
372
373 if !has_order {
374 let mut node = input;
377 if has_limit {
378 node = self.push(LogicalOp::Limit(Limit {
379 input: node,
380 skip: skip.clone(),
381 limit: limit.clone(),
382 }));
383 }
384 return self.plan_projection_or_aggregation(node, items, false, include_existing);
385 }
386
387 let mut node = self.push(LogicalOp::Projection(Projection {
388 input,
389 distinct: false,
390 items: items.to_vec(),
391 include_existing: true,
392 }));
393 node = self.push(LogicalOp::Sort(Sort {
394 input: node,
395 items: order.to_vec(),
396 top_k: None,
397 }));
398 if has_limit {
399 node = self.push(LogicalOp::Limit(Limit {
400 input: node,
401 skip: skip.clone(),
402 limit: limit.clone(),
403 }));
404 }
405 if include_existing {
406 return node;
408 }
409 self.push(LogicalOp::Projection(Projection {
410 input: node,
411 distinct: false,
412 items: passthrough_items(items),
413 include_existing: false,
414 }))
415 }
416
417 fn plan_projection_or_aggregation(
421 &mut self,
422 input: PlanNodeId,
423 items: &[ResolvedProjection],
424 distinct: bool,
425 include_existing: bool,
426 ) -> PlanNodeId {
427 let has_aggregates = items.iter().any(|item| expr_contains_aggregate(&item.expr));
428
429 if !has_aggregates {
430 return self.push(LogicalOp::Projection(Projection {
431 input,
432 distinct,
433 items: items.to_vec(),
434 include_existing,
435 }));
436 }
437
438 let mut group_by = Vec::new();
440 let mut aggregates = Vec::new();
441
442 for item in items {
443 if expr_contains_aggregate(&item.expr) {
444 aggregates.push(item.clone());
445 } else {
446 group_by.push(item.clone());
447 }
448 }
449
450 let node = self.push(LogicalOp::Aggregation(Aggregation {
451 input,
452 group_by: group_by.clone(),
453 aggregates: aggregates.clone(),
454 }));
455
456 if distinct {
465 self.push(LogicalOp::Projection(Projection {
467 input: node,
468 distinct: true,
469 items: passthrough_items(items),
470 include_existing: false,
471 }))
472 } else {
473 node
474 }
475 }
476
477 fn plan_unit_input(&mut self) -> PlanNodeId {
478 self.push(LogicalOp::Argument(Argument))
479 }
480}
481
482fn passthrough_items(items: &[ResolvedProjection]) -> Vec<ResolvedProjection> {
484 items
485 .iter()
486 .map(|item| ResolvedProjection {
487 expr: ResolvedExpr::Variable(item.output),
488 output: item.output,
489 name: item.name.clone(),
490 explicit_alias: item.explicit_alias,
491 span: item.span,
492 })
493 .collect()
494}
495
496fn sort_keys_on_outputs(
505 order: &[ResolvedSortItem],
506 items: &[ResolvedProjection],
507) -> Vec<ResolvedSortItem> {
508 let projected: Vec<(String, VarId)> = items
509 .iter()
510 .map(|item| (format!("{:?}", item.expr), item.output))
511 .collect();
512 order
513 .iter()
514 .map(|key| {
515 let shape = format!("{:?}", key.expr);
516 match projected.iter().find(|(expr, _)| *expr == shape) {
517 Some((_, output)) => ResolvedSortItem {
518 expr: ResolvedExpr::Variable(*output),
519 direction: key.direction,
520 },
521 None => key.clone(),
522 }
523 })
524 .collect()
525}
526
527fn collect_pattern_vars(pattern: &ResolvedPattern) -> Vec<VarId> {
529 let mut vars = Vec::new();
530 for part in &pattern.parts {
531 if let Some(v) = part.binding {
532 vars.push(v);
533 }
534 match &part.element {
535 ResolvedPatternElement::Node { var, .. } => {
536 if let Some(v) = var {
537 vars.push(*v);
538 }
539 }
540 ResolvedPatternElement::ShortestPath { head, chain, .. }
541 | ResolvedPatternElement::NodeChain { head, chain } => {
542 if let Some(v) = head.var {
543 vars.push(v);
544 }
545 for step in chain {
546 if let Some(v) = step.rel.var {
547 vars.push(v);
548 }
549 if let Some(v) = step.node.var {
550 vars.push(v);
551 }
552 }
553 }
554 }
555 }
556 vars
557}
558
559fn expr_contains_aggregate(expr: &ResolvedExpr) -> bool {
560 match expr {
561 ResolvedExpr::Function { function, args, .. } => {
562 if function.is_aggregate() {
563 return true;
564 }
565 args.iter().any(expr_contains_aggregate)
566 }
567 ResolvedExpr::Property { expr, .. } => expr_contains_aggregate(expr),
568 ResolvedExpr::Binary { lhs, rhs, .. } => {
569 expr_contains_aggregate(lhs) || expr_contains_aggregate(rhs)
570 }
571 ResolvedExpr::Unary { expr, .. } => expr_contains_aggregate(expr),
572 ResolvedExpr::List(items) => items.iter().any(expr_contains_aggregate),
573 ResolvedExpr::Map(items) => items.iter().any(|(_, v)| expr_contains_aggregate(v)),
574 ResolvedExpr::Case {
575 input,
576 alternatives,
577 else_expr,
578 } => {
579 input.as_ref().is_some_and(|e| expr_contains_aggregate(e))
580 || alternatives
581 .iter()
582 .any(|(w, t)| expr_contains_aggregate(w) || expr_contains_aggregate(t))
583 || else_expr
584 .as_ref()
585 .is_some_and(|e| expr_contains_aggregate(e))
586 }
587 _ => false,
588 }
589}