1use crate::ast::{InternalColumnRef, InternalRelationId, RuleEvent};
10use crate::ir::ScalarExpr;
11use crate::plan::ExpressionPlan;
12use serde::{Deserialize, Serialize};
13
14#[derive(Debug, Clone, Serialize, Deserialize)]
15pub struct RuleConditionBinding {
16 old_relation: Option<InternalRelationId>,
17 new_relation: Option<InternalRelationId>,
18 columns: Vec<String>,
19}
20
21impl RuleConditionBinding {
22 pub fn for_event(columns: &[String], event: RuleEvent) -> Self {
23 let old_relation = matches!(event, RuleEvent::Update | RuleEvent::Delete)
24 .then(InternalRelationId::allocate);
25 let new_relation = matches!(event, RuleEvent::Insert | RuleEvent::Update)
26 .then(InternalRelationId::allocate);
27 Self {
28 old_relation,
29 new_relation,
30 columns: columns.to_vec(),
31 }
32 }
33
34 pub const fn old_relation(&self) -> Option<InternalRelationId> {
35 self.old_relation
36 }
37
38 pub const fn new_relation(&self) -> Option<InternalRelationId> {
39 self.new_relation
40 }
41
42 pub fn old_column(&self, name: &str) -> Option<InternalColumnRef> {
43 self.column(self.old_relation, name)
44 }
45
46 pub fn new_column(&self, name: &str) -> Option<InternalColumnRef> {
47 self.column(self.new_relation, name)
48 }
49
50 pub fn column_name(&self, column: InternalColumnRef) -> Option<&str> {
51 if Some(column.relation()) != self.old_relation
52 && Some(column.relation()) != self.new_relation
53 {
54 return None;
55 }
56 self.columns.get(column.attribute()).map(String::as_str)
57 }
58
59 pub fn reallocate_plan_relations(&self, plan: &mut ExpressionPlan) -> Self {
61 let rebound = Self {
62 old_relation: self.old_relation.map(|_| InternalRelationId::allocate()),
63 new_relation: self.new_relation.map(|_| InternalRelationId::allocate()),
64 columns: self.columns.clone(),
65 };
66 self.remap_plan(plan, &rebound);
67 rebound
68 }
69
70 pub fn refresh_columns(&self, columns: Vec<String>, plan: &mut ExpressionPlan) -> Self {
72 if self.columns == columns {
73 return self.clone();
74 }
75 let rebound = Self {
76 old_relation: self.old_relation.map(|_| InternalRelationId::allocate()),
77 new_relation: self.new_relation.map(|_| InternalRelationId::allocate()),
78 columns,
79 };
80 self.remap_plan(plan, &rebound);
81 rebound
82 }
83
84 pub fn whole_row_expression(&self, qualifier: &str, table: &str) -> Option<ScalarExpr> {
86 let relation = match qualifier {
87 "old" => self.old_relation?,
88 "new" => self.new_relation?,
89 _ => return None,
90 };
91 Some(ScalarExpr::Cast {
92 implicit: false,
93 expr: Box::new(ScalarExpr::Row(
94 (0..self.columns.len())
95 .map(|position| ScalarExpr::InternalColumn(relation.column(position)))
96 .collect(),
97 )),
98 ty: table.to_string(),
99 })
100 }
101
102 pub fn referenced_columns(
103 &self,
104 plan: &ExpressionPlan,
105 ) -> std::collections::BTreeSet<InternalColumnRef> {
106 let mut required = std::collections::BTreeSet::new();
107 let mut collect = |expression: &ScalarExpr| {
108 expression.visit(&mut |node| {
109 if let ScalarExpr::InternalColumn(column) = node {
110 if self.column_name(*column).is_some() {
111 required.insert(*column);
112 }
113 }
114 });
115 };
116 collect(&plan.scalar);
117 for query in &plan.subqueries {
118 query.visit_scalar_expressions(&mut collect);
119 }
120 required
121 }
122
123 pub fn row_schema(&self, columns: &[(String, crate::ast::ColumnType)]) -> crate::RowSchema {
124 let mut names = Vec::with_capacity(columns.len() * 2);
125 let mut identities = Vec::with_capacity(columns.len() * 2);
126 let mut types = Vec::with_capacity(columns.len() * 2);
127 let mut internal = Vec::with_capacity(columns.len() * 2);
128 for (side, relation) in [("old", self.old_relation()), ("new", self.new_relation())] {
129 let Some(relation) = relation else {
130 continue;
131 };
132 for (attribute, (name, ty)) in columns.iter().enumerate() {
133 let slot = names.len();
134 names.push(name.clone());
135 identities.push(crate::ColumnIdentity::qualified(side, name));
136 types.push(Some(ty.clone()));
137 internal.push((relation.column(attribute), slot, Some(ty.clone())));
138 }
139 }
140 let schema = crate::RowSchema::with_identities(names, identities, types);
141 crate::RowSchema::with_physical_internal_aliases(&schema, &internal)
142 }
143
144 fn remap_plan(&self, plan: &mut ExpressionPlan, rebound: &Self) {
145 let mut rewrite = |expression: &mut ScalarExpr| {
146 let ScalarExpr::InternalColumn(column) = expression else {
147 return;
148 };
149 let relation = if Some(column.relation()) == self.old_relation {
150 rebound.old_relation
151 } else if Some(column.relation()) == self.new_relation {
152 rebound.new_relation
153 } else {
154 None
155 };
156 if let Some(relation) = relation {
157 if let Some(name) = self.column_name(*column) {
158 if let Some(replacement) = rebound.column(Some(relation), name) {
159 *column = replacement;
160 }
161 }
162 }
163 };
164 crate::plan::rewrite_scalar_expression(&mut plan.scalar, &mut rewrite);
165 for subquery in &mut plan.subqueries {
166 subquery.rewrite_scalar_expressions(&mut rewrite);
167 }
168 }
169
170 fn column(
171 &self,
172 relation: Option<InternalRelationId>,
173 name: &str,
174 ) -> Option<InternalColumnRef> {
175 let relation = relation?;
176 self.columns
177 .iter()
178 .position(|column| column == name)
179 .map(|position| relation.column(position))
180 }
181}
182
183#[cfg(test)]
184mod tests {
185 use super::*;
186
187 #[test]
188 fn deserialized_plan_relations_are_reallocated_and_rewritten_together() {
189 let binding = RuleConditionBinding {
190 old_relation: Some(InternalRelationId::from_raw(u64::MAX - 1)),
191 new_relation: Some(InternalRelationId::from_raw(u64::MAX)),
192 columns: vec!["id".into()],
193 };
194 let mut plan = ExpressionPlan {
195 scalar: ScalarExpr::InternalColumn(binding.old_column("id").unwrap()),
196 subqueries: Vec::new(),
197 };
198
199 let rebound = binding.reallocate_plan_relations(&mut plan);
200 let ScalarExpr::InternalColumn(column) = plan.scalar else {
201 panic!("condition plan lost its structural OLD column")
202 };
203 assert_eq!(Some(column.relation()), rebound.old_relation());
204 assert_eq!(rebound.column_name(column), Some("id"));
205 assert_ne!(rebound.old_relation(), binding.old_relation());
206 assert_ne!(rebound.new_relation(), binding.new_relation());
207 }
208 #[test]
209 fn refreshed_columns_keep_field_names_and_expand_the_live_whole_row() {
210 let binding = RuleConditionBinding::for_event(
211 &["id".into(), "value".into(), "unused".into()],
212 RuleEvent::Update,
213 );
214 let mut plan = ExpressionPlan {
215 scalar: ScalarExpr::Row(vec![
216 ScalarExpr::InternalColumn(binding.new_column("value").unwrap()),
217 ScalarExpr::Column("old".into()),
218 ]),
219 subqueries: Vec::new(),
220 };
221 let refreshed =
222 binding.refresh_columns(vec!["value".into(), "id".into(), "added".into()], &mut plan);
223 assert_eq!(
224 refreshed.referenced_columns(&plan),
225 std::collections::BTreeSet::from([refreshed.new_column("value").unwrap()])
226 );
227 let ScalarExpr::Row(items) = &plan.scalar else {
228 panic!("outer row");
229 };
230 assert_eq!(
231 items[0],
232 ScalarExpr::InternalColumn(refreshed.new_column("value").unwrap())
233 );
234 let whole = refreshed
235 .whole_row_expression("old", "public.event")
236 .unwrap();
237 let ScalarExpr::Cast { expr, ty, .. } = &whole else {
238 panic!("typed whole row");
239 };
240 assert_eq!(ty, "public.event");
241 assert_eq!(
242 expr.as_ref(),
243 &ScalarExpr::Row(
244 ["value", "id", "added"]
245 .map(|name| ScalarExpr::InternalColumn(refreshed.old_column(name).unwrap()))
246 .to_vec()
247 )
248 );
249 }
250}