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 [crate::ast::AssignmentTarget],
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            implicit: true,
232            expr: Box::new(Expr::Literal(value)),
233            ty,
234        },
235        None => Expr::Literal(value),
236    }
237}
238
239struct RuleSourceResolver<'a> {
240    old_columns: &'a BTreeMap<String, InternalColumnRef>,
241    new_columns: &'a BTreeMap<String, InternalColumnRef>,
242    old_row: InternalColumnRef,
243    new_row: InternalColumnRef,
244}
245
246impl VariableResolver for RuleSourceResolver<'_> {
247    fn resolve_name(&mut self, _name: &str) -> Result<Option<ResolvedVariable>, SQLError> {
248        Ok(None)
249    }
250
251    fn resolve_qualified(
252        &mut self,
253        _qualifier: &str,
254        _column: &str,
255    ) -> Result<Option<ResolvedVariable>, SQLError> {
256        Ok(None)
257    }
258
259    fn resolve_param(&mut self, _index: usize) -> Result<Option<ResolvedVariable>, SQLError> {
260        Ok(None)
261    }
262
263    fn rewrite_name(&mut self, name: &str) -> Result<Option<Expr>, SQLError> {
264        Ok(if name.eq_ignore_ascii_case("old") {
265            Some(Expr::InternalColumn(self.old_row))
266        } else if name.eq_ignore_ascii_case("new") {
267            Some(Expr::InternalColumn(self.new_row))
268        } else {
269            None
270        })
271    }
272
273    fn rewrite_qualified(
274        &mut self,
275        qualifier: &str,
276        column: &str,
277    ) -> Result<Option<Expr>, SQLError> {
278        let columns = if qualifier.eq_ignore_ascii_case("old") {
279            self.old_columns
280        } else if qualifier.eq_ignore_ascii_case("new") {
281            self.new_columns
282        } else {
283            return Ok(None);
284        };
285        let source_column = columns
286            .get(column)
287            .ok_or_else(|| SQLError::UnknownColumn(format!("{qualifier}.{column}")))?;
288        Ok(Some(Expr::InternalColumn(*source_column)))
289    }
290
291    fn rewrite_qualified_whole_row(&mut self, qualifier: &str) -> Result<Option<Expr>, SQLError> {
292        self.rewrite_name(qualifier)
293    }
294}
295
296fn attach_rule_row_source(
297    statement: &mut Statement,
298    source: FromClause,
299    relation: InternalRelationId,
300) -> Result<(), SQLError> {
301    match statement {
302        Statement::Select(select) => attach_select_rule_source(select, &source, relation),
303        Statement::Insert(insert) => {
304            let select = insert.select_source.as_mut().ok_or_else(|| {
305                SQLError::Internal("set-oriented rule INSERT action has no SELECT source".into())
306            })?;
307            attach_select_rule_source(select, &source, relation);
308        }
309        Statement::Update(update) => {
310            update.from = Some(prepend_rule_row_source(
311                update.from.take(),
312                source,
313                relation,
314            ));
315        }
316        Statement::Delete(delete) => {
317            delete.using = Some(prepend_rule_row_source(
318                delete.using.take(),
319                source,
320                relation,
321            ));
322        }
323        _ => {
324            return Err(SQLError::Internal(
325                "validated rewrite-rule action changed statement kind".into(),
326            ))
327        }
328    }
329    Ok(())
330}
331
332fn attach_select_rule_source(
333    select: &mut SelectStmt,
334    source: &FromClause,
335    relation: InternalRelationId,
336) {
337    if let Some(set_op) = select.set_op.as_mut() {
338        if let Some(left) = set_op.left.as_mut() {
339            attach_select_rule_source(left, source, relation);
340        } else {
341            select.from = Some(prepend_rule_row_source(
342                select.from.take(),
343                source.clone(),
344                relation,
345            ));
346        }
347        attach_select_rule_source(&mut set_op.right, source, relation);
348    } else {
349        select.from = Some(prepend_rule_row_source(
350            select.from.take(),
351            source.clone(),
352            relation,
353        ));
354    }
355}
356
357fn prepend_rule_row_source(
358    existing: Option<FromClause>,
359    source: FromClause,
360    relation: InternalRelationId,
361) -> FromClause {
362    let Some(existing) = existing else {
363        return source;
364    };
365    let lateral = from_references_internal_relation(&existing, relation);
366    FromClause::Join {
367        left: Box::new(source),
368        right: Box::new(existing),
369        kind: JoinKind::Cross,
370        on: None,
371        using: None,
372        natural: false,
373        alias: None,
374        column_aliases: Vec::new(),
375        lateral,
376    }
377}
378
379struct InternalReferenceResolver {
380    relation: InternalRelationId,
381    referenced: bool,
382}
383
384impl VariableResolver for InternalReferenceResolver {
385    fn resolve_name(&mut self, _name: &str) -> Result<Option<ResolvedVariable>, SQLError> {
386        Ok(None)
387    }
388
389    fn resolve_qualified(
390        &mut self,
391        _qualifier: &str,
392        _column: &str,
393    ) -> Result<Option<ResolvedVariable>, SQLError> {
394        Ok(None)
395    }
396
397    fn resolve_param(&mut self, _index: usize) -> Result<Option<ResolvedVariable>, SQLError> {
398        Ok(None)
399    }
400
401    fn rewrite_internal(&mut self, column: InternalColumnRef) -> Result<Option<Expr>, SQLError> {
402        if column.relation() == self.relation {
403            self.referenced = true;
404        }
405        Ok(None)
406    }
407}
408
409fn expr_references_internal_relation(expr: &Expr, relation: InternalRelationId) -> bool {
410    let mut resolver = InternalReferenceResolver {
411        relation,
412        referenced: false,
413    };
414    let _ = bind_expr(expr, &mut resolver);
415    resolver.referenced
416}
417
418fn select_references_internal_relation(select: &SelectStmt, relation: InternalRelationId) -> bool {
419    let mut resolver = InternalReferenceResolver {
420        relation,
421        referenced: false,
422    };
423    let _ = bind_select(select, &mut resolver);
424    resolver.referenced
425}
426
427fn from_references_internal_relation(from: &FromClause, relation: InternalRelationId) -> bool {
428    match from {
429        FromClause::Table { .. } => false,
430        FromClause::Join {
431            left, right, on, ..
432        } => {
433            from_references_internal_relation(left, relation)
434                || from_references_internal_relation(right, relation)
435                || on
436                    .as_ref()
437                    .is_some_and(|expr| expr_references_internal_relation(expr, relation))
438        }
439        FromClause::Values { rows, .. } => rows
440            .iter()
441            .flatten()
442            .any(|expr| expr_references_internal_relation(expr, relation)),
443        FromClause::Function { args, .. } => args
444            .iter()
445            .any(|expr| expr_references_internal_relation(expr, relation)),
446        FromClause::FunctionGroup { functions, .. } => functions.iter().any(|function| {
447            function
448                .args
449                .iter()
450                .any(|expr| expr_references_internal_relation(expr, relation))
451        }),
452        FromClause::Subquery { body, .. } => select_references_internal_relation(body, relation),
453    }
454}