Skip to main content

uqa_sql/plan/
rewrite.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Recursive scalar-expression traversal and in-place plan rewriting.
8
9use super::{
10    AssignmentPlan, CommandPlan, ConflictActionPlan, MergeWhenPlan, OrderPlan, ProjectionPlan,
11    QueryPlan, RelationalPlan, ScalarExpr, ScalarFrameBound, SourcePlan,
12};
13
14pub(super) fn rewrite_query_scalars(
15    query: &mut QueryPlan,
16    rewrite: &mut dyn FnMut(&mut ScalarExpr),
17) {
18    for cte in &mut query.ctes {
19        match &mut cte.body {
20            super::CtePlanBody::Query(query) => rewrite_query_scalars(query, rewrite),
21            super::CtePlanBody::Command(command) => rewrite_command_scalars(command, rewrite),
22        }
23    }
24    match &mut query.root {
25        RelationalPlan::QueryBlock(block) => {
26            if let Some(source) = &mut block.from {
27                rewrite_source_scalars(source, rewrite);
28            }
29            rewrite_optional_scalar(&mut block.r#where, rewrite);
30            for projection in &mut block.projections {
31                rewrite_scalar(&mut projection.expr, rewrite);
32            }
33            for expression in &mut block.group_by {
34                rewrite_scalar(expression, rewrite);
35            }
36            for set in &mut block.grouping_sets {
37                for expression in set {
38                    rewrite_scalar(expression, rewrite);
39                }
40            }
41            rewrite_optional_scalar(&mut block.having, rewrite);
42            rewrite_orders(&mut block.order_by, rewrite);
43            rewrite_optional_scalar(&mut block.limit, rewrite);
44            rewrite_optional_scalar(&mut block.offset, rewrite);
45            for expression in &mut block.distinct_on {
46                rewrite_scalar(expression, rewrite);
47            }
48            for subquery in &mut block.subqueries {
49                rewrite_query_scalars(subquery, rewrite);
50            }
51        }
52        RelationalPlan::SetOp {
53            left,
54            right,
55            order_by,
56            limit,
57            offset,
58            subqueries,
59            ..
60        } => {
61            rewrite_query_scalars(left, rewrite);
62            rewrite_query_scalars(right, rewrite);
63            rewrite_orders(order_by, rewrite);
64            if let Some(limit) = limit {
65                rewrite_scalar(limit, rewrite);
66            }
67            if let Some(offset) = offset {
68                rewrite_scalar(offset, rewrite);
69            }
70            for subquery in subqueries {
71                rewrite_query_scalars(subquery, rewrite);
72            }
73        }
74        RelationalPlan::Values { rows, subqueries } => {
75            for row in rows {
76                for expression in row {
77                    rewrite_scalar(expression, rewrite);
78                }
79            }
80            for subquery in subqueries {
81                rewrite_query_scalars(subquery, rewrite);
82            }
83        }
84    }
85}
86
87pub(super) fn rewrite_source_scalars(
88    source: &mut SourcePlan,
89    rewrite: &mut dyn FnMut(&mut ScalarExpr),
90) {
91    match source {
92        SourcePlan::Table { .. } => {}
93        SourcePlan::Join {
94            left, right, on, ..
95        } => {
96            rewrite_source_scalars(left, rewrite);
97            rewrite_source_scalars(right, rewrite);
98            rewrite_optional_scalar(on, rewrite);
99        }
100        SourcePlan::Values { rows, .. } => {
101            for row in rows {
102                for expression in row {
103                    rewrite_scalar(expression, rewrite);
104                }
105            }
106        }
107        SourcePlan::Function { args, .. } => {
108            for expression in args {
109                rewrite_scalar(expression, rewrite);
110            }
111        }
112        SourcePlan::FunctionGroup { functions, .. } => {
113            for function in functions {
114                for expression in &mut function.args {
115                    rewrite_scalar(expression, rewrite);
116                }
117            }
118        }
119        SourcePlan::Subquery { body, .. } => rewrite_query_scalars(body, rewrite),
120    }
121}
122
123#[expect(
124    clippy::too_many_lines,
125    reason = "optimizer rewrite preserves exhaustive variants and fixed-point order"
126)]
127pub(super) fn rewrite_command_scalars(
128    command: &mut CommandPlan,
129    rewrite: &mut dyn FnMut(&mut ScalarExpr),
130) {
131    match command {
132        CommandPlan::Insert(plan) => {
133            for expression in plan
134                .columns
135                .iter_mut()
136                .flat_map(crate::ast::AssignmentTarget::expressions_mut)
137            {
138                rewrite_scalar(expression, rewrite);
139            }
140            for cte in &mut plan.ctes {
141                match &mut cte.body {
142                    super::CtePlanBody::Query(query) => rewrite_query_scalars(query, rewrite),
143                    super::CtePlanBody::Command(command) => {
144                        rewrite_command_scalars(command, rewrite);
145                    }
146                }
147            }
148            for row in &mut plan.rows {
149                for expression in row {
150                    rewrite_scalar(expression, rewrite);
151                }
152            }
153            if let Some(source) = &mut plan.source {
154                rewrite_query_scalars(source, rewrite);
155            }
156            if let Some(conflict) = &mut plan.on_conflict {
157                for expression in &mut conflict.expressions {
158                    rewrite_scalar(expression, rewrite);
159                }
160                if let Some(predicate) = &mut conflict.predicate {
161                    rewrite_scalar(predicate, rewrite);
162                }
163                if let ConflictActionPlan::Update {
164                    assignments,
165                    predicate,
166                } = &mut conflict.action
167                {
168                    rewrite_assignments(assignments, rewrite);
169                    if let Some(predicate) = predicate {
170                        rewrite_scalar(predicate, rewrite);
171                    }
172                }
173            }
174            rewrite_projections(&mut plan.returning, rewrite);
175            for check in &mut plan.view_checks {
176                rewrite_scalar(&mut check.predicate, rewrite);
177            }
178            rewrite_subqueries(&mut plan.subqueries, rewrite);
179        }
180        CommandPlan::Update(plan) => {
181            for cte in &mut plan.ctes {
182                match &mut cte.body {
183                    super::CtePlanBody::Query(query) => rewrite_query_scalars(query, rewrite),
184                    super::CtePlanBody::Command(command) => {
185                        rewrite_command_scalars(command, rewrite);
186                    }
187                }
188            }
189            if let Some(source) = &mut plan.source {
190                rewrite_source_scalars(source, rewrite);
191            }
192            rewrite_assignments(&mut plan.assignments, rewrite);
193            rewrite_optional_scalar(&mut plan.predicate, rewrite);
194            rewrite_projections(&mut plan.returning, rewrite);
195            for check in &mut plan.view_checks {
196                rewrite_scalar(&mut check.predicate, rewrite);
197            }
198            rewrite_subqueries(&mut plan.subqueries, rewrite);
199        }
200        CommandPlan::Delete(plan) => {
201            for cte in &mut plan.ctes {
202                match &mut cte.body {
203                    super::CtePlanBody::Query(query) => rewrite_query_scalars(query, rewrite),
204                    super::CtePlanBody::Command(command) => {
205                        rewrite_command_scalars(command, rewrite);
206                    }
207                }
208            }
209            if let Some(source) = &mut plan.source {
210                rewrite_source_scalars(source, rewrite);
211            }
212            rewrite_optional_scalar(&mut plan.predicate, rewrite);
213            rewrite_projections(&mut plan.returning, rewrite);
214            rewrite_subqueries(&mut plan.subqueries, rewrite);
215        }
216        CommandPlan::Merge(plan) => {
217            for cte in &mut plan.ctes {
218                match &mut cte.body {
219                    super::CtePlanBody::Query(query) => rewrite_query_scalars(query, rewrite),
220                    super::CtePlanBody::Command(command) => {
221                        rewrite_command_scalars(command, rewrite);
222                    }
223                }
224            }
225            rewrite_source_scalars(&mut plan.source, rewrite);
226            rewrite_optional_scalar(&mut plan.target_predicate, rewrite);
227            rewrite_scalar(&mut plan.join_condition, rewrite);
228            for clause in &mut plan.when_clauses {
229                match clause {
230                    MergeWhenPlan::UpdateMatched {
231                        condition,
232                        assignments,
233                    }
234                    | MergeWhenPlan::UpdateNotMatchedBySource {
235                        condition,
236                        assignments,
237                    } => {
238                        rewrite_optional_scalar(condition, rewrite);
239                        rewrite_assignments(assignments, rewrite);
240                    }
241                    MergeWhenPlan::DeleteMatched { condition }
242                    | MergeWhenPlan::DeleteNotMatchedBySource { condition }
243                    | MergeWhenPlan::NothingMatched { condition }
244                    | MergeWhenPlan::NothingNotMatched { condition }
245                    | MergeWhenPlan::NothingNotMatchedBySource { condition } => {
246                        rewrite_optional_scalar(condition, rewrite);
247                    }
248                    MergeWhenPlan::InsertNotMatched {
249                        condition,
250                        columns,
251                        values,
252                        ..
253                    } => {
254                        rewrite_optional_scalar(condition, rewrite);
255                        for expression in columns
256                            .iter_mut()
257                            .flat_map(crate::ast::AssignmentTarget::expressions_mut)
258                        {
259                            rewrite_scalar(expression, rewrite);
260                        }
261                        for value in values {
262                            rewrite_scalar(value, rewrite);
263                        }
264                    }
265                }
266            }
267            rewrite_projections(&mut plan.returning, rewrite);
268            for check in &mut plan.view_checks {
269                rewrite_scalar(&mut check.predicate, rewrite);
270            }
271            rewrite_subqueries(&mut plan.subqueries, rewrite);
272        }
273        CommandPlan::CreateView { query, .. }
274        | CommandPlan::CreateMaterializedView { query, .. }
275        | CommandPlan::CreateTableAs { query, .. }
276        | CommandPlan::DeclareCursor { query, .. } => {
277            rewrite_query_scalars(query, rewrite);
278        }
279        CommandPlan::Explain { body, .. } | CommandPlan::Prepare { body, .. } => {
280            body.rewrite_scalar_expressions(rewrite);
281        }
282        CommandPlan::Execute { params, .. } | CommandPlan::Call { args: params, .. } => {
283            for expression in params {
284                rewrite_scalar(&mut expression.scalar, rewrite);
285                rewrite_subqueries(&mut expression.subqueries, rewrite);
286            }
287        }
288        CommandPlan::CreateTable(_)
289        | CommandPlan::CreateTableIfNotExists(_)
290        | CommandPlan::CreateIndex(_)
291        | CommandPlan::RenameIndex(_)
292        | CommandPlan::Drop(_)
293        | CommandPlan::AlterTable(_)
294        | CommandPlan::AlterForeignTable(_)
295        | CommandPlan::AlterView(_)
296        | CommandPlan::RefreshMaterializedView { .. }
297        | CommandPlan::CreateSchema { .. }
298        | CommandPlan::AlterSchemaOwner { .. }
299        | CommandPlan::Notify { .. }
300        | CommandPlan::Listen { .. }
301        | CommandPlan::Unlisten { .. }
302        | CommandPlan::SetVariable { .. }
303        | CommandPlan::ResetVariable { .. }
304        | CommandPlan::ResetAllVariables
305        | CommandPlan::SetConstraints { .. }
306        | CommandPlan::ShowVariable { .. }
307        | CommandPlan::Discard { .. }
308        | CommandPlan::Load { .. }
309        | CommandPlan::Analyze { .. }
310        | CommandPlan::Vacuum(_)
311        | CommandPlan::LockTable(_)
312        | CommandPlan::Truncate { .. }
313        | CommandPlan::Transaction(_)
314        | CommandPlan::FetchCursor(_)
315        | CommandPlan::CloseCursor { .. }
316        | CommandPlan::CreateSequence(_)
317        | CommandPlan::CreateDomain(_)
318        | CommandPlan::AlterSequence(_)
319        | CommandPlan::Deallocate { .. }
320        | CommandPlan::CreateForeignServer(_)
321        | CommandPlan::CreateForeignTable(_)
322        | CommandPlan::CreateForeignTableIfNotExists(_)
323        | CommandPlan::CreateFunction(_)
324        | CommandPlan::DropFunction(_)
325        | CommandPlan::AlterRoutine(_)
326        | CommandPlan::AlterRoutineOwner(_)
327        | CommandPlan::RenameRoutine(_)
328        | CommandPlan::GrantRoutine(_)
329        | CommandPlan::GrantTable(_)
330        | CommandPlan::GrantSequence(_)
331        | CommandPlan::GrantDatabase(_)
332        | CommandPlan::GrantSchema(_)
333        | CommandPlan::GrantRole(_)
334        | CommandPlan::CreateRole(_)
335        | CommandPlan::AlterRole(_)
336        | CommandPlan::RenameRole(_)
337        | CommandPlan::DropRole(_)
338        | CommandPlan::CreateTrigger(_)
339        | CommandPlan::DropTrigger(_)
340        | CommandPlan::CreateRule(_)
341        | CommandPlan::DropRule(_)
342        | CommandPlan::DoBlock { .. } => {}
343    }
344}
345
346pub(super) fn rewrite_assignments(
347    assignments: &mut [AssignmentPlan],
348    rewrite: &mut dyn FnMut(&mut ScalarExpr),
349) {
350    for assignment in assignments {
351        for expression in assignment.expressions_mut() {
352            rewrite_scalar(expression, rewrite);
353        }
354    }
355}
356
357pub(super) fn rewrite_projections(
358    projections: &mut [ProjectionPlan],
359    rewrite: &mut dyn FnMut(&mut ScalarExpr),
360) {
361    for projection in projections {
362        rewrite_scalar(&mut projection.expr, rewrite);
363    }
364}
365
366pub(super) fn rewrite_orders(orders: &mut [OrderPlan], rewrite: &mut dyn FnMut(&mut ScalarExpr)) {
367    for order in orders {
368        rewrite_scalar(&mut order.expr, rewrite);
369    }
370}
371
372pub(super) fn rewrite_subqueries(
373    subqueries: &mut [QueryPlan],
374    rewrite: &mut dyn FnMut(&mut ScalarExpr),
375) {
376    for subquery in subqueries {
377        rewrite_query_scalars(subquery, rewrite);
378    }
379}
380
381pub(super) fn rewrite_optional_scalar(
382    expression: &mut Option<ScalarExpr>,
383    rewrite: &mut dyn FnMut(&mut ScalarExpr),
384) {
385    if let Some(expression) = expression {
386        rewrite_scalar(expression, rewrite);
387    }
388}
389
390pub(super) fn rewrite_scalar(
391    expression: &mut ScalarExpr,
392    rewrite: &mut dyn FnMut(&mut ScalarExpr),
393) {
394    match expression {
395        ScalarExpr::Func {
396            args,
397            order_by,
398            filter,
399            ..
400        } => {
401            for argument in args {
402                rewrite_scalar(argument, rewrite);
403            }
404            for order in order_by {
405                rewrite_scalar(&mut order.expr, rewrite);
406            }
407            if let Some(filter) = filter {
408                rewrite_scalar(filter, rewrite);
409            }
410        }
411        ScalarExpr::Array(items)
412        | ScalarExpr::Row(items)
413        | ScalarExpr::And(items)
414        | ScalarExpr::Or(items) => {
415            for item in items {
416                rewrite_scalar(item, rewrite);
417            }
418        }
419        ScalarExpr::Binary { lhs, rhs, .. } => {
420            rewrite_scalar(lhs, rewrite);
421            rewrite_scalar(rhs, rewrite);
422        }
423        ScalarExpr::UnaryMinus(inner)
424        | ScalarExpr::Not(inner)
425        | ScalarExpr::IsNull { expr: inner, .. }
426        | ScalarExpr::Cast { expr: inner, .. } => rewrite_scalar(inner, rewrite),
427        ScalarExpr::Between { expr, low, high } => {
428            rewrite_scalar(expr, rewrite);
429            rewrite_scalar(low, rewrite);
430            rewrite_scalar(high, rewrite);
431        }
432        ScalarExpr::InList { expr, list, .. } => {
433            rewrite_scalar(expr, rewrite);
434            for item in list {
435                rewrite_scalar(item, rewrite);
436            }
437        }
438        ScalarExpr::WindowCall { args, spec, .. } => {
439            for argument in args {
440                rewrite_scalar(argument, rewrite);
441            }
442            for expression in &mut spec.partition_by {
443                rewrite_scalar(expression, rewrite);
444            }
445            for order in &mut spec.order_by {
446                rewrite_scalar(&mut order.expr, rewrite);
447            }
448            if let Some(frame) = &mut spec.frame {
449                rewrite_frame_bound(&mut frame.start, rewrite);
450                rewrite_frame_bound(&mut frame.end, rewrite);
451            }
452        }
453        ScalarExpr::Case {
454            base,
455            when,
456            else_branch,
457        } => {
458            if let Some(base) = base {
459                rewrite_scalar(base, rewrite);
460            }
461            for (condition, result) in when {
462                rewrite_scalar(condition, rewrite);
463                rewrite_scalar(result, rewrite);
464            }
465            if let Some(branch) = else_branch {
466                rewrite_scalar(branch, rewrite);
467            }
468        }
469        ScalarExpr::InSubquery { expr, .. } => rewrite_scalar(expr, rewrite),
470        ScalarExpr::Default
471        | ScalarExpr::Star
472        | ScalarExpr::QualifiedStar(_)
473        | ScalarExpr::Column(_)
474        | ScalarExpr::Position(_)
475        | ScalarExpr::InternalColumn(_)
476        | ScalarExpr::QualifiedColumn { .. }
477        | ScalarExpr::Literal(_)
478        | ScalarExpr::TypedLiteral { .. }
479        | ScalarExpr::Param(_)
480        | ScalarExpr::ScalarSubquery(_)
481        | ScalarExpr::Exists { .. } => {}
482    }
483    rewrite(expression);
484}
485
486/// Visit one scalar-expression tree in post-order and rewrite each node once.
487///
488/// Query-owned callers that need scope-sensitive rewriting can use this entry
489/// point without duplicating the exhaustive [`ScalarExpr`] traversal.
490pub fn rewrite_scalar_expression(
491    expression: &mut ScalarExpr,
492    rewrite: &mut dyn FnMut(&mut ScalarExpr),
493) {
494    rewrite_scalar(expression, rewrite);
495}
496
497pub(super) fn rewrite_frame_bound(
498    bound: &mut ScalarFrameBound,
499    rewrite: &mut dyn FnMut(&mut ScalarExpr),
500) {
501    match bound {
502        ScalarFrameBound::Preceding(expression) | ScalarFrameBound::Following(expression) => {
503            rewrite_scalar(expression, rewrite);
504        }
505        ScalarFrameBound::UnboundedPreceding
506        | ScalarFrameBound::UnboundedFollowing
507        | ScalarFrameBound::CurrentRow => {}
508    }
509}