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, ResolvedUnwind,
11 ResolvedWith,
12};
13
14pub struct Planner {
15 nodes: Vec<LogicalOp>,
16}
17
18impl Default for Planner {
19 fn default() -> Self {
20 Self::new()
21 }
22}
23
24impl Planner {
25 pub fn new() -> Self {
26 Self { nodes: Vec::new() }
27 }
28
29 pub(crate) fn push(&mut self, op: LogicalOp) -> PlanNodeId {
30 let id = self.nodes.len();
31 self.nodes.push(op);
32 id
33 }
34
35 pub fn plan(&mut self, query: &ResolvedQuery) -> LogicalPlan {
36 let root = self.plan_query(query);
37
38 LogicalPlan {
39 root,
40 nodes: std::mem::take(&mut self.nodes),
41 }
42 }
43
44 fn plan_query(&mut self, query: &ResolvedQuery) -> PlanNodeId {
45 let mut input = None;
46
47 for clause in &query.clauses {
48 input = Some(match clause {
49 ResolvedClause::Match(m) => self.plan_match(input, m),
50
51 ResolvedClause::Unwind(u) => {
52 let upstream = input.unwrap_or_else(|| self.plan_unit_input());
53 self.plan_unwind(upstream, u)
54 }
55
56 ResolvedClause::Create(c) => {
57 let upstream = input.unwrap_or_else(|| self.plan_unit_input());
58 self.plan_create(upstream, c)
59 }
60
61 ResolvedClause::Merge(m) => {
62 let upstream = input.unwrap_or_else(|| self.plan_unit_input());
63 self.plan_merge(upstream, m)
64 }
65
66 ResolvedClause::Delete(d) => {
67 let upstream = input.unwrap_or_else(|| self.plan_unit_input());
68 self.plan_delete(upstream, d)
69 }
70
71 ResolvedClause::Set(s) => {
72 let upstream = input.unwrap_or_else(|| self.plan_unit_input());
73 self.plan_set(upstream, s)
74 }
75
76 ResolvedClause::Remove(rm) => {
77 let upstream = input.unwrap_or_else(|| self.plan_unit_input());
78 self.plan_remove(upstream, rm)
79 }
80
81 ResolvedClause::Foreach(f) => {
82 let upstream = input.unwrap_or_else(|| self.plan_unit_input());
83 self.plan_foreach(upstream, f)
84 }
85
86 ResolvedClause::With(w) => {
87 let upstream = input.unwrap_or_else(|| self.plan_unit_input());
88 self.plan_with(upstream, w)
89 }
90
91 ResolvedClause::Return(r) => {
92 let upstream = input.unwrap_or_else(|| self.plan_unit_input());
93 self.plan_return(upstream, r)
94 }
95
96 ResolvedClause::CallSubquery(c) => {
97 let upstream = input.unwrap_or_else(|| self.plan_unit_input());
98 self.plan_call_subquery(upstream, c)
99 }
100 });
101 }
102
103 input.unwrap_or_else(|| self.plan_unit_input())
104 }
105
106 fn plan_call_subquery(&mut self, input: PlanNodeId, call: &ResolvedCallSubquery) -> PlanNodeId {
110 let inner_query = ResolvedQuery {
111 clauses: call.clauses.clone(),
112 unions: Vec::new(),
113 };
114 let inner = self.plan_query(&inner_query);
115 self.push(LogicalOp::CallSubquery(crate::logical::CallSubquery {
116 input,
117 inner,
118 new_vars: call.return_vars.clone(),
119 }))
120 }
121
122 fn plan_match(&mut self, input: Option<PlanNodeId>, m: &ResolvedMatch) -> PlanNodeId {
123 if let (true, Some(upstream)) = (m.optional, input) {
124 let new_vars = collect_pattern_vars(&m.pattern);
129
130 let mut pattern_planner = PatternPlanner::new(self);
133 let mut inner = pattern_planner.plan_pattern(None, &m.pattern);
134
135 if let Some(pred) = &m.where_ {
136 inner = self.push(LogicalOp::Filter(Filter {
137 input: inner,
138 predicate: pred.clone(),
139 }));
140 }
141
142 self.push(LogicalOp::OptionalMatch(OptionalMatch {
143 input: upstream,
144 inner,
145 new_vars,
146 }))
147 } else {
148 let mut pattern_planner = PatternPlanner::new(self);
149 let mut node = pattern_planner.plan_pattern(input, &m.pattern);
150
151 if let Some(pred) = &m.where_ {
152 node = self.push(LogicalOp::Filter(Filter {
153 input: node,
154 predicate: pred.clone(),
155 }));
156 }
157
158 node
159 }
160 }
161
162 fn plan_unwind(&mut self, input: PlanNodeId, u: &ResolvedUnwind) -> PlanNodeId {
163 self.push(LogicalOp::Unwind(Unwind {
164 input,
165 expr: u.expr.clone(),
166 alias: u.alias,
167 }))
168 }
169
170 fn plan_create(&mut self, input: PlanNodeId, c: &ResolvedCreate) -> PlanNodeId {
171 self.push(LogicalOp::Create(crate::Create {
172 input,
173 pattern: c.pattern.clone(),
174 }))
175 }
176
177 fn plan_merge(&mut self, input: PlanNodeId, m: &ResolvedMerge) -> PlanNodeId {
178 self.push(LogicalOp::Merge(crate::Merge {
179 input,
180 pattern_part: m.pattern_part.clone(),
181 actions: m.actions.clone(),
182 }))
183 }
184
185 fn plan_delete(&mut self, input: PlanNodeId, d: &ResolvedDelete) -> PlanNodeId {
186 self.push(LogicalOp::Delete(crate::Delete {
187 input,
188 detach: d.detach,
189 expressions: d.expressions.clone(),
190 }))
191 }
192
193 fn plan_set(&mut self, input: PlanNodeId, s: &ResolvedSet) -> PlanNodeId {
194 self.push(LogicalOp::Set(crate::Set {
195 input,
196 items: s.items.clone(),
197 }))
198 }
199
200 fn plan_remove(&mut self, input: PlanNodeId, r: &ResolvedRemove) -> PlanNodeId {
201 self.push(LogicalOp::Remove(crate::Remove {
202 input,
203 items: r.items.clone(),
204 }))
205 }
206
207 fn plan_foreach(&mut self, input: PlanNodeId, f: &ResolvedForeach) -> PlanNodeId {
208 self.push(LogicalOp::Foreach(crate::Foreach {
209 input,
210 variable: f.variable,
211 list: f.list.clone(),
212 body: f.body.clone(),
213 }))
214 }
215
216 fn plan_with(&mut self, input: PlanNodeId, with: &ResolvedWith) -> PlanNodeId {
217 let mut node = input;
218
219 if !with.order.is_empty() {
221 node = self.push(LogicalOp::Sort(Sort {
222 input: node,
223 items: with.order.clone(),
224 top_k: None,
225 }));
226 }
227
228 if with.skip.is_some() || with.limit.is_some() {
229 node = self.push(LogicalOp::Limit(Limit {
230 input: node,
231 skip: with.skip.clone(),
232 limit: with.limit.clone(),
233 }));
234 }
235
236 node = self.plan_projection_or_aggregation(
237 node,
238 &with.items,
239 with.distinct,
240 with.include_existing,
241 );
242
243 if let Some(pred) = &with.where_ {
244 node = self.push(LogicalOp::Filter(Filter {
245 input: node,
246 predicate: pred.clone(),
247 }));
248 }
249
250 node
251 }
252
253 fn plan_return(&mut self, input: PlanNodeId, ret: &ResolvedReturn) -> PlanNodeId {
254 let mut node = input;
255
256 if !ret.order.is_empty() {
260 node = self.push(LogicalOp::Sort(Sort {
261 input: node,
262 items: ret.order.clone(),
263 top_k: None,
264 }));
265 }
266
267 if ret.skip.is_some() || ret.limit.is_some() {
268 node = self.push(LogicalOp::Limit(Limit {
269 input: node,
270 skip: ret.skip.clone(),
271 limit: ret.limit.clone(),
272 }));
273 }
274
275 node = self.plan_projection_or_aggregation(
276 node,
277 &ret.items,
278 ret.distinct,
279 ret.include_existing,
280 );
281
282 node
283 }
284
285 fn plan_projection_or_aggregation(
289 &mut self,
290 input: PlanNodeId,
291 items: &[ResolvedProjection],
292 distinct: bool,
293 include_existing: bool,
294 ) -> PlanNodeId {
295 let has_aggregates = items.iter().any(|item| expr_contains_aggregate(&item.expr));
296
297 if !has_aggregates {
298 return self.push(LogicalOp::Projection(Projection {
299 input,
300 distinct,
301 items: items.to_vec(),
302 include_existing,
303 }));
304 }
305
306 let mut group_by = Vec::new();
308 let mut aggregates = Vec::new();
309
310 for item in items {
311 if expr_contains_aggregate(&item.expr) {
312 aggregates.push(item.clone());
313 } else {
314 group_by.push(item.clone());
315 }
316 }
317
318 let node = self.push(LogicalOp::Aggregation(Aggregation {
319 input,
320 group_by: group_by.clone(),
321 aggregates: aggregates.clone(),
322 }));
323
324 if distinct {
333 let passthrough_items: Vec<ResolvedProjection> = items
335 .iter()
336 .map(|item| ResolvedProjection {
337 expr: ResolvedExpr::Variable(item.output),
338 output: item.output,
339 name: item.name.clone(),
340 explicit_alias: item.explicit_alias,
341 span: item.span,
342 })
343 .collect();
344 self.push(LogicalOp::Projection(Projection {
345 input: node,
346 distinct: true,
347 items: passthrough_items,
348 include_existing: false,
349 }))
350 } else {
351 node
352 }
353 }
354
355 fn plan_unit_input(&mut self) -> PlanNodeId {
356 self.push(LogicalOp::Argument(Argument))
357 }
358}
359
360fn collect_pattern_vars(pattern: &ResolvedPattern) -> Vec<VarId> {
362 let mut vars = Vec::new();
363 for part in &pattern.parts {
364 if let Some(v) = part.binding {
365 vars.push(v);
366 }
367 match &part.element {
368 ResolvedPatternElement::Node { var, .. } => {
369 if let Some(v) = var {
370 vars.push(*v);
371 }
372 }
373 ResolvedPatternElement::ShortestPath { head, chain, .. }
374 | ResolvedPatternElement::NodeChain { head, chain } => {
375 if let Some(v) = head.var {
376 vars.push(v);
377 }
378 for step in chain {
379 if let Some(v) = step.rel.var {
380 vars.push(v);
381 }
382 if let Some(v) = step.node.var {
383 vars.push(v);
384 }
385 }
386 }
387 }
388 }
389 vars
390}
391
392fn expr_contains_aggregate(expr: &ResolvedExpr) -> bool {
393 match expr {
394 ResolvedExpr::Function { function, args, .. } => {
395 if function.is_aggregate() {
396 return true;
397 }
398 args.iter().any(expr_contains_aggregate)
399 }
400 ResolvedExpr::Property { expr, .. } => expr_contains_aggregate(expr),
401 ResolvedExpr::Binary { lhs, rhs, .. } => {
402 expr_contains_aggregate(lhs) || expr_contains_aggregate(rhs)
403 }
404 ResolvedExpr::Unary { expr, .. } => expr_contains_aggregate(expr),
405 ResolvedExpr::List(items) => items.iter().any(expr_contains_aggregate),
406 ResolvedExpr::Map(items) => items.iter().any(|(_, v)| expr_contains_aggregate(v)),
407 ResolvedExpr::Case {
408 input,
409 alternatives,
410 else_expr,
411 } => {
412 input.as_ref().is_some_and(|e| expr_contains_aggregate(e))
413 || alternatives
414 .iter()
415 .any(|(w, t)| expr_contains_aggregate(w) || expr_contains_aggregate(t))
416 || else_expr
417 .as_ref()
418 .is_some_and(|e| expr_contains_aggregate(e))
419 }
420 _ => false,
421 }
422}