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