Skip to main content

uqa_sql/semantics/rules/binding/
actions.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Set-oriented rewrite-rule action binding and internal OLD/NEW row sources.
8
9use std::collections::BTreeMap;
10
11use crate::ast::{
12    ColumnType, Expr, FromClause, InsertStmt, InternalColumnRef, InternalRelationId, JoinKind,
13    OnConflict, Projection, ReturningAliases, SelectStmt, Statement,
14};
15use crate::plpgsql::{bind_expr, bind_select, ResolvedVariable, VariableResolver};
16use crate::SQLError;
17use uqa_core::Value;
18
19use super::{RuleColumnMetadata, RuleRowValues, RuntimeRuleResolver};
20
21type BindRuleAction<'a> = dyn Fn(&mut dyn VariableResolver) -> Result<Statement, SQLError> + 'a;
22
23pub fn bind_insert_values_action(
24    matching_rows: &[usize],
25    rows: &[impl RuleRowValues],
26    columns: &BTreeMap<String, RuleColumnMetadata>,
27    bind_action: &BindRuleAction<'_>,
28) -> Result<Statement, SQLError> {
29    if matching_rows.is_empty() {
30        let mut bound = bind_action(&mut RuntimeRuleResolver {
31            old: None,
32            new: None,
33            old_doc_id: None,
34            new_doc_id: None,
35            columns,
36        })?;
37        let Statement::Insert(insert) = &mut bound else {
38            return Err(SQLError::Internal(
39                "rewrite rule INSERT VALUES action changed statement kind".into(),
40            ));
41        };
42        insert.rows.clear();
43        return Ok(bound);
44    }
45    let mut combined = None;
46    for row_index in matching_rows {
47        let row = rows.get(*row_index).ok_or_else(|| {
48            SQLError::Internal("rewrite rule lost its qualified row image".into())
49        })?;
50        let bound = bind_action(&mut runtime_rule_resolver(row, columns))?;
51        let Statement::Insert(mut insert) = bound else {
52            return Err(SQLError::Internal(
53                "rewrite rule INSERT VALUES action changed statement kind".into(),
54            ));
55        };
56        if let Some(Statement::Insert(existing)) = combined.as_mut() {
57            if BoundInsertContract::from(&*existing) != BoundInsertContract::from(&insert) {
58                return Err(SQLError::Internal(
59                    "rewrite rule INSERT VALUES action produced row-dependent statement clauses"
60                        .into(),
61                ));
62            }
63            existing.rows.append(&mut insert.rows);
64        } else {
65            combined = Some(Statement::Insert(insert));
66        }
67    }
68    combined.ok_or_else(|| SQLError::Internal("rewrite rule action lost its row source".into()))
69}
70
71#[derive(PartialEq)]
72struct BoundInsertContract<'a> {
73    table: &'a str,
74    target_relation_bound: bool,
75    target_qualifier: &'a str,
76    include_descendants: bool,
77    columns: &'a [String],
78    with: &'a [crate::ast::CTE],
79    select_source: Option<&'a SelectStmt>,
80    on_conflict: Option<&'a OnConflict>,
81    returning: &'a [Projection],
82    returning_aliases: &'a ReturningAliases,
83}
84
85impl<'a> From<&'a InsertStmt> for BoundInsertContract<'a> {
86    fn from(insert: &'a InsertStmt) -> Self {
87        Self {
88            table: &insert.table,
89            target_relation_bound: insert.target_relation_bound,
90            target_qualifier: &insert.target_qualifier,
91            include_descendants: insert.include_descendants,
92            columns: &insert.columns,
93            with: &insert.with,
94            select_source: insert.select_source.as_deref(),
95            on_conflict: insert.on_conflict.as_ref(),
96            returning: &insert.returning,
97            returning_aliases: &insert.returning_aliases,
98        }
99    }
100}
101
102fn runtime_rule_resolver<'a>(
103    row: &'a dyn RuleRowValues,
104    columns: &'a BTreeMap<String, RuleColumnMetadata>,
105) -> RuntimeRuleResolver<'a> {
106    RuntimeRuleResolver {
107        old: row.old_row(),
108        new: row.new_row(),
109        old_doc_id: row.old_doc_id(),
110        new_doc_id: row.new_doc_id(),
111        columns,
112    }
113}
114
115struct RuleRowSource {
116    clause: FromClause,
117    relation: InternalRelationId,
118    old_columns: BTreeMap<String, InternalColumnRef>,
119    new_columns: BTreeMap<String, InternalColumnRef>,
120    old_row: InternalColumnRef,
121    new_row: InternalColumnRef,
122    source_index: InternalColumnRef,
123}
124
125pub struct BoundSetOrientedAction {
126    pub statement: Statement,
127    pub source_index: Expr,
128}
129
130pub fn bind_set_oriented_action(
131    matching_rows: &[usize],
132    rows: &[impl RuleRowValues],
133    columns: &BTreeMap<String, RuleColumnMetadata>,
134    bind_action: &BindRuleAction<'_>,
135) -> Result<BoundSetOrientedAction, SQLError> {
136    let source = rule_row_source(matching_rows, rows, columns)?;
137    let source_index = Expr::InternalColumn(source.source_index);
138    let mut bound = bind_action(&mut RuleSourceResolver {
139        old_columns: &source.old_columns,
140        new_columns: &source.new_columns,
141        old_row: source.old_row,
142        new_row: source.new_row,
143    })?;
144    attach_rule_row_source(&mut bound, source.clause, source.relation)?;
145    Ok(BoundSetOrientedAction {
146        statement: bound,
147        source_index,
148    })
149}
150
151fn rule_row_source(
152    matching_rows: &[usize],
153    rows: &[impl RuleRowValues],
154    columns: &BTreeMap<String, RuleColumnMetadata>,
155) -> Result<RuleRowSource, SQLError> {
156    let relation = InternalRelationId::allocate();
157    let mut internal_column_types = Vec::with_capacity(columns.len() * 2 + 3);
158    let mut old_columns = BTreeMap::new();
159    let mut new_columns = BTreeMap::new();
160    for (index, (column, metadata)) in columns.iter().enumerate() {
161        let old = relation.column(index * 2);
162        let new = relation.column(index * 2 + 1);
163        old_columns.insert(column.clone(), old);
164        new_columns.insert(column.clone(), new);
165        internal_column_types.push(Some(metadata.ty.clone()));
166        internal_column_types.push(Some(metadata.ty.clone()));
167    }
168    let old_row = relation.column(internal_column_types.len());
169    internal_column_types.push(Some(ColumnType::Record));
170    let new_row = relation.column(internal_column_types.len());
171    internal_column_types.push(Some(ColumnType::Record));
172    let source_index = relation.column(internal_column_types.len());
173    internal_column_types.push(Some(ColumnType::BigInteger));
174    let values = matching_rows
175        .iter()
176        .map(|row_index| {
177            let row = rows.get(*row_index).ok_or_else(|| {
178                SQLError::Internal("rewrite rule lost its qualified row image".into())
179            })?;
180            let resolver = runtime_rule_resolver(row, columns);
181            let mut values = Vec::with_capacity(internal_column_types.len());
182            for column in columns.keys() {
183                values.push(resolved_variable_expr(resolver.record_field(
184                    row.old_row(),
185                    row.old_doc_id(),
186                    column,
187                )?));
188                values.push(resolved_variable_expr(resolver.record_field(
189                    row.new_row(),
190                    row.new_doc_id(),
191                    column,
192                )?));
193            }
194            values.push(resolved_variable_expr(
195                resolver.record(row.old_row(), row.old_doc_id())?,
196            ));
197            values.push(resolved_variable_expr(
198                resolver.record(row.new_row(), row.new_doc_id())?,
199            ));
200            let row_index = i64::try_from(*row_index).map_err(|_| {
201                SQLError::Internal("rewrite rule event row index exceeds BIGINT".into())
202            })?;
203            values.push(Expr::Literal(Value::Int(row_index)));
204            Ok(values)
205        })
206        .collect::<Result<Vec<_>, SQLError>>()?;
207    Ok(RuleRowSource {
208        clause: FromClause::Values {
209            rows: values,
210            alias: None,
211            column_aliases: Vec::new(),
212            internal_relation: Some(relation),
213            internal_column_types,
214        },
215        relation,
216        old_columns,
217        new_columns,
218        old_row,
219        new_row,
220        source_index,
221    })
222}
223
224fn resolved_variable_expr(variable: ResolvedVariable) -> Expr {
225    let ResolvedVariable {
226        value,
227        declared_type,
228    } = variable;
229    match declared_type {
230        Some(ty) => Expr::Cast {
231            expr: Box::new(Expr::Literal(value)),
232            ty,
233        },
234        None => Expr::Literal(value),
235    }
236}
237
238struct RuleSourceResolver<'a> {
239    old_columns: &'a BTreeMap<String, InternalColumnRef>,
240    new_columns: &'a BTreeMap<String, InternalColumnRef>,
241    old_row: InternalColumnRef,
242    new_row: InternalColumnRef,
243}
244
245impl VariableResolver for RuleSourceResolver<'_> {
246    fn resolve_name(&mut self, _name: &str) -> Result<Option<ResolvedVariable>, SQLError> {
247        Ok(None)
248    }
249
250    fn resolve_qualified(
251        &mut self,
252        _qualifier: &str,
253        _column: &str,
254    ) -> Result<Option<ResolvedVariable>, SQLError> {
255        Ok(None)
256    }
257
258    fn resolve_param(&mut self, _index: usize) -> Result<Option<ResolvedVariable>, SQLError> {
259        Ok(None)
260    }
261
262    fn rewrite_name(&mut self, name: &str) -> Result<Option<Expr>, SQLError> {
263        Ok(if name.eq_ignore_ascii_case("old") {
264            Some(Expr::InternalColumn(self.old_row))
265        } else if name.eq_ignore_ascii_case("new") {
266            Some(Expr::InternalColumn(self.new_row))
267        } else {
268            None
269        })
270    }
271
272    fn rewrite_qualified(
273        &mut self,
274        qualifier: &str,
275        column: &str,
276    ) -> Result<Option<Expr>, SQLError> {
277        let columns = if qualifier.eq_ignore_ascii_case("old") {
278            self.old_columns
279        } else if qualifier.eq_ignore_ascii_case("new") {
280            self.new_columns
281        } else {
282            return Ok(None);
283        };
284        let source_column = columns
285            .get(column)
286            .ok_or_else(|| SQLError::UnknownColumn(format!("{qualifier}.{column}")))?;
287        Ok(Some(Expr::InternalColumn(*source_column)))
288    }
289
290    fn rewrite_qualified_whole_row(&mut self, qualifier: &str) -> Result<Option<Expr>, SQLError> {
291        self.rewrite_name(qualifier)
292    }
293}
294
295fn attach_rule_row_source(
296    statement: &mut Statement,
297    source: FromClause,
298    relation: InternalRelationId,
299) -> Result<(), SQLError> {
300    match statement {
301        Statement::Select(select) => attach_select_rule_source(select, &source, relation),
302        Statement::Insert(insert) => {
303            let select = insert.select_source.as_mut().ok_or_else(|| {
304                SQLError::Internal("set-oriented rule INSERT action has no SELECT source".into())
305            })?;
306            attach_select_rule_source(select, &source, relation);
307        }
308        Statement::Update(update) => {
309            update.from = Some(prepend_rule_row_source(
310                update.from.take(),
311                source,
312                relation,
313            ));
314        }
315        Statement::Delete(delete) => {
316            delete.using = Some(prepend_rule_row_source(
317                delete.using.take(),
318                source,
319                relation,
320            ));
321        }
322        _ => {
323            return Err(SQLError::Internal(
324                "validated rewrite-rule action changed statement kind".into(),
325            ))
326        }
327    }
328    Ok(())
329}
330
331fn attach_select_rule_source(
332    select: &mut SelectStmt,
333    source: &FromClause,
334    relation: InternalRelationId,
335) {
336    if let Some(set_op) = select.set_op.as_mut() {
337        if let Some(left) = set_op.left.as_mut() {
338            attach_select_rule_source(left, source, relation);
339        } else {
340            select.from = Some(prepend_rule_row_source(
341                select.from.take(),
342                source.clone(),
343                relation,
344            ));
345        }
346        attach_select_rule_source(&mut set_op.right, source, relation);
347    } else {
348        select.from = Some(prepend_rule_row_source(
349            select.from.take(),
350            source.clone(),
351            relation,
352        ));
353    }
354}
355
356fn prepend_rule_row_source(
357    existing: Option<FromClause>,
358    source: FromClause,
359    relation: InternalRelationId,
360) -> FromClause {
361    let Some(existing) = existing else {
362        return source;
363    };
364    let lateral = from_references_internal_relation(&existing, relation);
365    FromClause::Join {
366        left: Box::new(source),
367        right: Box::new(existing),
368        kind: JoinKind::Cross,
369        on: None,
370        using: None,
371        natural: false,
372        alias: None,
373        column_aliases: Vec::new(),
374        lateral,
375    }
376}
377
378struct InternalReferenceResolver {
379    relation: InternalRelationId,
380    referenced: bool,
381}
382
383impl VariableResolver for InternalReferenceResolver {
384    fn resolve_name(&mut self, _name: &str) -> Result<Option<ResolvedVariable>, SQLError> {
385        Ok(None)
386    }
387
388    fn resolve_qualified(
389        &mut self,
390        _qualifier: &str,
391        _column: &str,
392    ) -> Result<Option<ResolvedVariable>, SQLError> {
393        Ok(None)
394    }
395
396    fn resolve_param(&mut self, _index: usize) -> Result<Option<ResolvedVariable>, SQLError> {
397        Ok(None)
398    }
399
400    fn rewrite_internal(&mut self, column: InternalColumnRef) -> Result<Option<Expr>, SQLError> {
401        if column.relation() == self.relation {
402            self.referenced = true;
403        }
404        Ok(None)
405    }
406}
407
408fn expr_references_internal_relation(expr: &Expr, relation: InternalRelationId) -> bool {
409    let mut resolver = InternalReferenceResolver {
410        relation,
411        referenced: false,
412    };
413    let _ = bind_expr(expr, &mut resolver);
414    resolver.referenced
415}
416
417fn select_references_internal_relation(select: &SelectStmt, relation: InternalRelationId) -> bool {
418    let mut resolver = InternalReferenceResolver {
419        relation,
420        referenced: false,
421    };
422    let _ = bind_select(select, &mut resolver);
423    resolver.referenced
424}
425
426fn from_references_internal_relation(from: &FromClause, relation: InternalRelationId) -> bool {
427    match from {
428        FromClause::Table { .. } => false,
429        FromClause::Join {
430            left, right, on, ..
431        } => {
432            from_references_internal_relation(left, relation)
433                || from_references_internal_relation(right, relation)
434                || on
435                    .as_ref()
436                    .is_some_and(|expr| expr_references_internal_relation(expr, relation))
437        }
438        FromClause::Values { rows, .. } => rows
439            .iter()
440            .flatten()
441            .any(|expr| expr_references_internal_relation(expr, relation)),
442        FromClause::Function { args, .. } => args
443            .iter()
444            .any(|expr| expr_references_internal_relation(expr, relation)),
445        FromClause::FunctionGroup { functions, .. } => functions.iter().any(|function| {
446            function
447                .args
448                .iter()
449                .any(|expr| expr_references_internal_relation(expr, relation))
450        }),
451        FromClause::Subquery { body, .. } => select_references_internal_relation(body, relation),
452    }
453}