uqa_sql/routines/
merge_columns.rs1use crate::{
10 ast::{ColumnDef, ColumnType, Expr, MergeStmt, MergeTargetColumnBinding, MergeWhen, Statement},
11 SQLError,
12};
13use std::collections::{BTreeMap, BTreeSet};
14
15pub trait StoredMergeColumnCatalog {
16 fn stored_merge_target_definitions(&self, table: &str) -> Option<Vec<ColumnDef>>;
17}
18
19pub fn bind_stored_merge_target_columns(
20 catalog: &dyn StoredMergeColumnCatalog,
21 statement: &mut Statement,
22) -> Result<bool, SQLError> {
23 let mut changed = false;
24 crate::catalog::stored_ast::visit_stored_statement_merges(statement, &mut |merge| {
25 let Some(definitions) = catalog.stored_merge_target_definitions(&merge.target) else {
26 return Ok(());
27 };
28 let previous = merge.target_column_bindings.clone();
29 let mut targets = BTreeSet::new();
30 let mut coerced_targets = BTreeSet::new();
31 for action in &mut merge.when_clauses {
32 match action {
33 MergeWhen::UpdateMatched { assignments, .. }
34 | MergeWhen::UpdateNotMatchedBySource { assignments, .. } => {
35 targets.extend(assignments.iter().map(|(name, _)| name.clone()));
36 coerced_targets.extend(
37 assignments
38 .iter()
39 .filter(|(_, expression)| !matches!(expression, Expr::Default))
40 .map(|(name, _)| name.clone()),
41 );
42 }
43 MergeWhen::InsertNotMatched {
44 columns, values, ..
45 } => {
46 if columns.is_empty() && !values.is_empty() {
47 *columns = definitions
48 .iter()
49 .take(values.len())
50 .map(|column| column.name.clone())
51 .collect();
52 changed = true;
53 }
54 targets.extend(columns.iter().cloned());
55 coerced_targets.extend(
56 columns
57 .iter()
58 .zip(values.iter())
59 .filter(|(_, expression)| !matches!(expression, Expr::Default))
60 .map(|(name, _)| name.clone()),
61 );
62 }
63 _ => {}
64 }
65 }
66 for name in targets {
67 if let Some(column) = definitions.iter().find(|column| column.name == name) {
68 if let Some(object_id) = column.object_id {
69 let mut domain_dependencies = BTreeSet::new();
70 if coerced_targets.contains(&name) {
71 collect_target_domains(&column.ty, &mut domain_dependencies);
72 }
73 merge
74 .target_column_bindings
75 .entry(name)
76 .or_insert(MergeTargetColumnBinding {
77 object_id,
78 domain_dependencies,
79 });
80 }
81 }
82 }
83 changed |= merge.target_column_bindings != previous;
84 Ok(())
85 })?;
86 Ok(changed)
87}
88
89pub fn dropped_stored_merge_targets(
90 catalog: &dyn StoredMergeColumnCatalog,
91 merge: &MergeStmt,
92) -> BTreeSet<String> {
93 if merge.target_column_bindings.is_empty() {
94 return BTreeSet::new();
95 }
96 let live = catalog
97 .stored_merge_target_definitions(&merge.target)
98 .unwrap_or_default()
99 .into_iter()
100 .filter_map(|column| column.object_id)
101 .collect::<BTreeSet<_>>();
102 merge
103 .target_column_bindings
104 .iter()
105 .filter(|(_, binding)| !live.contains(&binding.object_id))
106 .map(|(name, _)| name.clone())
107 .collect()
108}
109
110pub fn normalize_stored_merge_target_columns(
111 catalog: &dyn StoredMergeColumnCatalog,
112 statement: &mut Statement,
113) -> Result<(), SQLError> {
114 crate::catalog::stored_ast::visit_stored_statement_merges(statement, &mut |merge| {
115 if dropped_stored_merge_targets(catalog, merge).is_empty() {
116 return Ok(());
117 }
118 let current = catalog
119 .stored_merge_target_definitions(&merge.target)
120 .unwrap_or_default()
121 .into_iter()
122 .filter_map(|column| column.object_id.map(|id| (id, column.name)))
123 .collect::<BTreeMap<_, _>>();
124 let name_for = |name: &str| match merge.target_column_bindings.get(name) {
125 Some(binding) => current.get(&binding.object_id).cloned(),
126 None => Some(name.to_string()),
127 };
128 for action in &mut merge.when_clauses {
129 match action {
130 MergeWhen::UpdateMatched { assignments, .. }
131 | MergeWhen::UpdateNotMatchedBySource { assignments, .. } => {
132 assignments.retain_mut(|(name, _)| {
133 if let Some(current) = name_for(name) {
134 *name = current;
135 true
136 } else {
137 false
138 }
139 });
140 }
141 MergeWhen::InsertNotMatched {
142 columns, values, ..
143 } => {
144 let mut surviving = Vec::new();
145 let mut expressions = Vec::new();
146 for (column, expression) in std::mem::take(columns)
147 .into_iter()
148 .zip(std::mem::take(values))
149 {
150 if let Some(current) = name_for(&column) {
151 surviving.push(current);
152 expressions.push(expression);
153 }
154 }
155 *columns = surviving;
156 *values = expressions;
157 }
158 _ => {}
159 }
160 }
161 Ok(())
162 })
163}
164
165pub fn collect_target_domains(ty: &ColumnType, dependencies: &mut BTreeSet<u32>) {
166 match ty {
167 ColumnType::Domain { oid, base, .. } => {
168 dependencies.insert(*oid);
169 collect_target_domains(base, dependencies);
170 }
171 ColumnType::Array(element) => collect_target_domains(element, dependencies),
172 _ => {}
173 }
174}
175
176pub fn statement_has_removed_merge_target(
177 catalog: &dyn StoredMergeColumnCatalog,
178 statement: &Statement,
179) -> Result<bool, SQLError> {
180 let mut changed = false;
181 crate::catalog::stored_ast::visit_stored_statement_merges(
182 &mut statement.clone(),
183 &mut |merge| {
184 changed |= !dropped_stored_merge_targets(catalog, merge).is_empty();
185 Ok(())
186 },
187 )?;
188 Ok(changed)
189}
190
191pub fn routine_has_removed_merge_target(
192 catalog: &dyn StoredMergeColumnCatalog,
193 definition: &crate::ast::CreateFunction,
194) -> Result<bool, SQLError> {
195 let crate::ast::FunctionBody::Statements(statements) = &definition.body else {
196 return Ok(false);
197 };
198 let mut changed = false;
199 for statement in statements {
200 changed |= statement_has_removed_merge_target(catalog, statement)?;
201 }
202 Ok(changed)
203}