Skip to main content

uqa_sql/semantics/rules/
returning.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Rule RETURNING image requirements, positional contracts, and action rewriting.
8use crate::{
9    ast::{Expr, Projection, ReturningAliases, RuleEvent, Statement},
10    plan::{ProjectionPlan, QueryPlan, RelationalPlan, SourcePlan},
11    plpgsql::{ResolvedVariable, VariableResolver},
12    SQLError, ScalarExpr,
13};
14use std::collections::BTreeSet;
15use uqa_core::Value;
16#[derive(Clone, Copy, Default)]
17pub struct RuleReturningRequest {
18    capture: bool,
19    images: RuleReturningImages,
20}
21
22#[derive(Clone, Copy, Default)]
23struct RuleReturningImages(u8);
24
25impl RuleReturningImages {
26    const CURRENT: Self = Self(1 << 0);
27    const OLD: Self = Self(1 << 1);
28    const NEW: Self = Self(1 << 2);
29
30    fn insert(&mut self, image: Self) {
31        self.0 |= image.0;
32    }
33
34    const fn contains(self, image: Self) -> bool {
35        self.0 & image.0 != 0
36    }
37}
38
39impl RuleReturningRequest {
40    pub fn from_plan(
41        returning: &[ProjectionPlan],
42        aliases: &ReturningAliases,
43        subqueries: &[QueryPlan],
44    ) -> Self {
45        if returning.is_empty() {
46            return Self::default();
47        }
48        let mut request = Self {
49            capture: true,
50            ..Self::default()
51        };
52        let shadowed = BTreeSet::new();
53        for projection in returning {
54            let ids = request.inspect_expression(&projection.expr, aliases, &shadowed);
55            for id in ids {
56                if let Some(query) = subqueries.get(id) {
57                    request.inspect_query(query, aliases, &shadowed);
58                }
59            }
60        }
61        request
62    }
63
64    fn inspect_expression(
65        &mut self,
66        expression: &ScalarExpr,
67        aliases: &ReturningAliases,
68        shadowed: &BTreeSet<String>,
69    ) -> Vec<usize> {
70        let mut expression = expression.clone();
71        let mut subqueries = Vec::new();
72        crate::plan::rewrite_scalar_expression(&mut expression, &mut |node| match node {
73            ScalarExpr::Star | ScalarExpr::Column(_) | ScalarExpr::Position(_) => {
74                self.images.insert(RuleReturningImages::CURRENT);
75            }
76            ScalarExpr::QualifiedStar(qualifier)
77            | ScalarExpr::QualifiedColumn { qualifier, .. }
78                if !shadowed.contains(&qualifier.to_ascii_lowercase()) =>
79            {
80                if qualifier.eq_ignore_ascii_case(&aliases.old) {
81                    self.images.insert(RuleReturningImages::OLD);
82                } else if qualifier.eq_ignore_ascii_case(&aliases.new) {
83                    self.images.insert(RuleReturningImages::NEW);
84                } else {
85                    self.images.insert(RuleReturningImages::CURRENT);
86                }
87            }
88            ScalarExpr::ScalarSubquery(id)
89            | ScalarExpr::Exists { subquery: id, .. }
90            | ScalarExpr::InSubquery { subquery: id, .. } => subqueries.push(*id),
91            _ => {}
92        });
93        subqueries
94    }
95
96    fn inspect_cte_body(
97        &mut self,
98        body: &crate::plan::CtePlanBody,
99        aliases: &ReturningAliases,
100        inherited: &BTreeSet<String>,
101    ) {
102        match body {
103            crate::plan::CtePlanBody::Query(query) => self.inspect_query(query, aliases, inherited),
104            crate::plan::CtePlanBody::Command(command) => {
105                let mut scope = inherited.clone();
106                if let Some(qualifier) = command.target_qualifier() {
107                    scope.insert(qualifier.to_ascii_lowercase());
108                }
109                if let Some(source) = command.source_input() {
110                    collect_source_qualifiers(source, &mut scope);
111                    self.inspect_source(source, aliases, inherited);
112                }
113                if let Some(aliases) = command.returning_aliases() {
114                    scope.insert(aliases.old.to_ascii_lowercase());
115                    scope.insert(aliases.new.to_ascii_lowercase());
116                }
117                for cte in command.ctes() {
118                    self.inspect_cte_body(&cte.body, aliases, &scope);
119                }
120                for query in command.query_inputs() {
121                    self.inspect_query(query, aliases, &scope);
122                }
123                for expression in command.expressions() {
124                    let _ = self.inspect_expression(expression, aliases, &scope);
125                }
126            }
127        }
128    }
129
130    fn inspect_query(
131        &mut self,
132        query: &QueryPlan,
133        aliases: &ReturningAliases,
134        inherited: &BTreeSet<String>,
135    ) {
136        for cte in &query.ctes {
137            self.inspect_cte_body(&cte.body, aliases, inherited);
138        }
139        match &query.root {
140            RelationalPlan::QueryBlock(block) => {
141                let mut scope = inherited.clone();
142                if let Some(source) = &block.from {
143                    collect_source_qualifiers(source, &mut scope);
144                    self.inspect_source(source, aliases, inherited);
145                }
146                for expression in block
147                    .projections
148                    .iter()
149                    .map(|projection| &projection.expr)
150                    .chain(block.r#where.iter())
151                    .chain(block.group_by.iter())
152                    .chain(block.grouping_sets.iter().flatten())
153                    .chain(block.having.iter())
154                    .chain(block.order_by.iter().map(|order| &order.expr))
155                    .chain(block.limit.iter())
156                    .chain(block.offset.iter())
157                    .chain(block.distinct_on.iter())
158                {
159                    let _ = self.inspect_expression(expression, aliases, &scope);
160                }
161                for subquery in &block.subqueries {
162                    self.inspect_query(subquery, aliases, &scope);
163                }
164            }
165            RelationalPlan::SetOp {
166                left,
167                right,
168                order_by,
169                limit,
170                offset,
171                subqueries,
172                ..
173            } => {
174                self.inspect_query(left, aliases, inherited);
175                self.inspect_query(right, aliases, inherited);
176                for expression in order_by
177                    .iter()
178                    .map(|order| &order.expr)
179                    .chain(limit.iter().map(Box::as_ref))
180                    .chain(offset.iter().map(Box::as_ref))
181                {
182                    let _ = self.inspect_expression(expression, aliases, inherited);
183                }
184                for subquery in subqueries {
185                    self.inspect_query(subquery, aliases, inherited);
186                }
187            }
188            RelationalPlan::Values { rows, subqueries } => {
189                for expression in rows.iter().flatten() {
190                    let _ = self.inspect_expression(expression, aliases, inherited);
191                }
192                for subquery in subqueries {
193                    self.inspect_query(subquery, aliases, inherited);
194                }
195            }
196        }
197    }
198
199    fn inspect_source(
200        &mut self,
201        source: &SourcePlan,
202        aliases: &ReturningAliases,
203        inherited: &BTreeSet<String>,
204    ) {
205        match source {
206            SourcePlan::Table { .. } => {}
207            SourcePlan::Join {
208                left,
209                right,
210                on,
211                lateral,
212                ..
213            } => {
214                self.inspect_source(left, aliases, inherited);
215                let mut right_scope = inherited.clone();
216                if *lateral {
217                    collect_source_qualifiers(left, &mut right_scope);
218                }
219                self.inspect_source(right, aliases, &right_scope);
220                if let Some(on) = on {
221                    let mut scope = inherited.clone();
222                    collect_source_qualifiers(left, &mut scope);
223                    collect_source_qualifiers(right, &mut scope);
224                    let _ = self.inspect_expression(on, aliases, &scope);
225                }
226            }
227            SourcePlan::Values { rows, .. } => {
228                for expression in rows.iter().flatten() {
229                    let _ = self.inspect_expression(expression, aliases, inherited);
230                }
231            }
232            SourcePlan::Function { args, .. } => {
233                for expression in args {
234                    let _ = self.inspect_expression(expression, aliases, inherited);
235                }
236            }
237            SourcePlan::FunctionGroup { functions, .. } => {
238                for expression in functions.iter().flat_map(|function| &function.args) {
239                    let _ = self.inspect_expression(expression, aliases, inherited);
240                }
241            }
242            SourcePlan::Subquery { body, .. } => self.inspect_query(body, aliases, inherited),
243        }
244    }
245
246    pub const fn captures(self) -> bool {
247        self.capture
248    }
249}
250
251fn collect_source_qualifiers(source: &SourcePlan, output: &mut BTreeSet<String>) {
252    match source {
253        SourcePlan::Join {
254            left, right, alias, ..
255        } => {
256            if let Some(alias) = alias {
257                output.insert(alias.to_ascii_lowercase());
258            } else {
259                collect_source_qualifiers(left, output);
260                collect_source_qualifiers(right, output);
261            }
262        }
263        _ => {
264            if let Some(qualifier) = source.visible_qualifier() {
265                output.insert(qualifier.to_ascii_lowercase());
266            }
267        }
268    }
269}
270
271pub fn validate_rule_returning_provider_width(
272    provider_width: usize,
273    event_width: usize,
274) -> Result<(), SQLError> {
275    if provider_width < event_width {
276        return Err(SQLError::Internal(format!(
277            "could not find replacement targetlist entry for attno {}",
278            provider_width + 1
279        )));
280    }
281    if provider_width > event_width {
282        return Err(SQLError::Internal(format!(
283            "rewrite-rule RETURNING provider produced {provider_width} columns, expected {event_width}"
284        )));
285    }
286    Ok(())
287}
288
289pub fn augment_rule_returning_action(
290    statement: &mut Statement,
291    source_index: Option<Expr>,
292    event_width: usize,
293    request: RuleReturningRequest,
294    target_columns: &BTreeSet<String>,
295) -> Result<(), SQLError> {
296    let (target_qualifier, aliases, returning) = match statement {
297        Statement::Insert(action) => (
298            action.target_qualifier.clone(),
299            action.returning_aliases.clone(),
300            action.returning.clone(),
301        ),
302        Statement::Update(action) => (
303            action.target_qualifier.clone(),
304            action.returning_aliases.clone(),
305            action.returning.clone(),
306        ),
307        Statement::Delete(action) => (
308            action.target_qualifier.clone(),
309            action.returning_aliases.clone(),
310            action.returning.clone(),
311        ),
312        _ => return Ok(()),
313    };
314    if returning.is_empty() {
315        return Ok(());
316    }
317    let mut target_columns = target_columns.clone();
318    target_columns.insert(crate::semantics::DOC_ID_COLUMN.into());
319    let provider_event = match statement {
320        Statement::Insert(_) => RuleEvent::Insert,
321        Statement::Update(_) => RuleEvent::Update,
322        Statement::Delete(_) => RuleEvent::Delete,
323        _ => unreachable!("validated rule provider changed statement kind"),
324    };
325    let current = if request.images.contains(RuleReturningImages::CURRENT) {
326        returning.clone()
327    } else {
328        null_rule_returning_image(event_width)
329    };
330    let old = if request.images.contains(RuleReturningImages::OLD)
331        && provider_event != RuleEvent::Insert
332    {
333        rewrite_rule_returning_image(
334            &returning,
335            &target_qualifier,
336            &target_columns,
337            &aliases.old,
338            &aliases.new,
339            &aliases.old,
340        )?
341    } else {
342        null_rule_returning_image(event_width)
343    };
344    let new = if request.images.contains(RuleReturningImages::NEW)
345        && provider_event != RuleEvent::Delete
346    {
347        rewrite_rule_returning_image(
348            &returning,
349            &target_qualifier,
350            &target_columns,
351            &aliases.old,
352            &aliases.new,
353            &aliases.new,
354        )?
355    } else {
356        null_rule_returning_image(event_width)
357    };
358    let output = match statement {
359        Statement::Insert(action) => &mut action.returning,
360        Statement::Update(action) => &mut action.returning,
361        Statement::Delete(action) => &mut action.returning,
362        _ => unreachable!("validated rule provider changed statement kind"),
363    };
364    *output = current;
365    output.extend(old);
366    output.extend(new);
367    if let Some(expr) = source_index {
368        output.push(Projection { expr, alias: None });
369    }
370    Ok(())
371}
372
373fn null_rule_returning_image(width: usize) -> Vec<Projection> {
374    (0..width)
375        .map(|_| Projection {
376            expr: Expr::Literal(Value::Null),
377            alias: None,
378        })
379        .collect()
380}
381
382fn rewrite_rule_returning_image(
383    returning: &[Projection],
384    target_qualifier: &str,
385    target_columns: &BTreeSet<String>,
386    old_qualifier: &str,
387    new_qualifier: &str,
388    image_qualifier: &str,
389) -> Result<Vec<Projection>, SQLError> {
390    let mut resolver = ReturningImageResolver {
391        target_qualifier,
392        target_columns,
393        old_qualifier,
394        new_qualifier,
395        image_qualifier,
396    };
397    returning
398        .iter()
399        .map(|projection| {
400            let expr = match &projection.expr {
401                Expr::Star => Expr::QualifiedStar(image_qualifier.to_string()),
402                Expr::QualifiedStar(qualifier) if resolver.retargets_qualifier(qualifier) => {
403                    Expr::QualifiedStar(image_qualifier.to_string())
404                }
405                expr => super::action_binding::bind_rule_expr_scoped(
406                    expr,
407                    &mut resolver,
408                    &BTreeSet::new(),
409                )?,
410            };
411            Ok(Projection {
412                expr,
413                alias: projection.alias.clone(),
414            })
415        })
416        .collect()
417}
418
419struct ReturningImageResolver<'a> {
420    target_qualifier: &'a str,
421    target_columns: &'a BTreeSet<String>,
422    old_qualifier: &'a str,
423    new_qualifier: &'a str,
424    image_qualifier: &'a str,
425}
426
427impl ReturningImageResolver<'_> {
428    fn retargets_qualifier(&self, qualifier: &str) -> bool {
429        qualifier.eq_ignore_ascii_case(self.target_qualifier)
430            || qualifier.eq_ignore_ascii_case(self.old_qualifier)
431            || qualifier.eq_ignore_ascii_case(self.new_qualifier)
432    }
433}
434
435impl VariableResolver for ReturningImageResolver<'_> {
436    fn resolve_name(&mut self, _name: &str) -> Result<Option<ResolvedVariable>, SQLError> {
437        Ok(None)
438    }
439
440    fn resolve_qualified(
441        &mut self,
442        _qualifier: &str,
443        _column: &str,
444    ) -> Result<Option<ResolvedVariable>, SQLError> {
445        Ok(None)
446    }
447
448    fn resolve_param(&mut self, _index: usize) -> Result<Option<ResolvedVariable>, SQLError> {
449        Ok(None)
450    }
451
452    fn rewrite_name(&mut self, name: &str) -> Result<Option<Expr>, SQLError> {
453        Ok(self
454            .target_columns
455            .contains(name)
456            .then(|| Expr::qualified_column(self.image_qualifier, name)))
457    }
458
459    fn rewrite_qualified(
460        &mut self,
461        qualifier: &str,
462        column: &str,
463    ) -> Result<Option<Expr>, SQLError> {
464        Ok(self
465            .retargets_qualifier(qualifier)
466            .then(|| Expr::qualified_column(self.image_qualifier, column)))
467    }
468}