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 [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}