uqa_sql/semantics/view_rewrite/
rule_inputs.rs1use super::{
10 automatic_view_layer, has_instead_of_trigger, BTreeSet, SQLError, TriggerEvent,
11 ViewRewriteContext,
12};
13use crate::ast::RuleEvent;
14
15pub struct RuleInputRequirements {
16 pub columns: BTreeSet<String>,
17 pub requires_rows: bool,
18}
19
20pub fn rule_input_requirements(
21 services: ViewRewriteContext<'_>,
22 table: &str,
23 event: RuleEvent,
24) -> Result<Option<RuleInputRequirements>, SQLError> {
25 collect_requirements(services, table, event, &mut BTreeSet::new())
26}
27
28fn collect_requirements(
29 services: ViewRewriteContext<'_>,
30 table: &str,
31 event: RuleEvent,
32 visited: &mut BTreeSet<String>,
33) -> Result<Option<RuleInputRequirements>, SQLError> {
34 if !visited.insert(table.to_string()) {
35 return Err(SQLError::Internal(format!(
36 "cycle while resolving rewrite-rule inputs for `{table}`"
37 )));
38 }
39 let rules = services.catalog.rules_for(table, event)?;
40 let suppresses = rules
41 .iter()
42 .any(|rule| rule.definition.instead && rule.definition.condition.is_none());
43 let mut required = RuleInputRequirements {
44 columns: BTreeSet::new(),
45 requires_rows: rules
46 .iter()
47 .any(|rule| rule.definition.condition.is_some() || !rule.definition.actions.is_empty()),
48 };
49 for rule in &rules {
50 let Some(columns) = services.catalog.rule_new_row_columns(rule)? else {
51 return Ok(None);
52 };
53 required.columns.extend(columns);
54 }
55 if suppresses {
56 return Ok(Some(required));
57 }
58 let trigger = match event {
59 RuleEvent::Insert => TriggerEvent::Insert,
60 RuleEvent::Update => TriggerEvent::Update,
61 RuleEvent::Delete => TriggerEvent::Delete,
62 RuleEvent::Select => return Ok(None),
63 };
64 if has_instead_of_trigger(services, table, trigger)? {
65 return Ok(None);
66 }
67 let Some(layer) = automatic_view_layer(services, table)? else {
68 return Ok(None);
69 };
70 let Some(source) = collect_requirements(services, &layer.source_name, event, visited)? else {
71 return Ok(None);
72 };
73 required.requires_rows |= source.requires_rows;
74 for column in layer.columns {
75 if column
76 .writable_source_column
77 .as_ref()
78 .is_some_and(|name| source.columns.contains(name))
79 {
80 required.columns.insert(column.name);
81 }
82 }
83 Ok(Some(required))
84}