1use 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}