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