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