use crate::{
ast::{ColumnDef, ColumnType, Expr, MergeStmt, MergeTargetColumnBinding, MergeWhen, Statement},
SQLError,
};
use std::collections::{BTreeMap, BTreeSet};
pub trait StoredMergeColumnCatalog {
fn stored_merge_target_definitions(&self, table: &str) -> Option<Vec<ColumnDef>>;
}
pub fn bind_stored_merge_target_columns(
catalog: &dyn StoredMergeColumnCatalog,
statement: &mut Statement,
) -> Result<bool, SQLError> {
let mut changed = false;
crate::catalog::stored_ast::visit_stored_statement_merges(statement, &mut |merge| {
let Some(definitions) = catalog.stored_merge_target_definitions(&merge.target) else {
return Ok(());
};
let previous = merge.target_column_bindings.clone();
let mut targets = BTreeSet::new();
let mut coerced_targets = BTreeSet::new();
for action in &mut merge.when_clauses {
match action {
MergeWhen::UpdateMatched { assignments, .. }
| MergeWhen::UpdateNotMatchedBySource { assignments, .. } => {
targets.extend(assignments.iter().flat_map(|(group, _)| {
group.targets().iter().map(|target| target.column.clone())
}));
coerced_targets.extend(
assignments
.iter()
.filter(|(_, expression)| !matches!(expression, Expr::Default))
.flat_map(|(group, _)| {
group.targets().iter().map(|target| target.column.clone())
}),
);
}
MergeWhen::InsertNotMatched {
columns, values, ..
} => {
if columns.is_empty() && !values.is_empty() {
*columns = definitions
.iter()
.take(values.len())
.map(|column| column.name.clone().into())
.collect();
changed = true;
}
targets.extend(columns.iter().map(|target| target.column.clone()));
coerced_targets.extend(
columns
.iter()
.zip(values.iter())
.filter(|(_, expression)| !matches!(expression, Expr::Default))
.map(|(target, _)| target.column.clone()),
);
}
_ => {}
}
}
for name in targets {
if let Some((position, column)) = definitions
.iter()
.enumerate()
.find(|(_, column)| column.name == name)
{
if let Some(object_id) = column.object_id {
let mut domain_dependencies = BTreeSet::new();
if coerced_targets.contains(&name) {
collect_target_domains(&column.ty, &mut domain_dependencies);
}
merge
.target_column_bindings
.entry(name)
.or_insert(MergeTargetColumnBinding {
object_id,
attribute_number: Some(
crate::catalog::relation_attributes::column_number(
column, position,
)?,
),
domain_dependencies,
});
}
}
}
changed |= merge.target_column_bindings != previous;
Ok(())
})?;
Ok(changed)
}
pub fn dropped_stored_merge_targets(
catalog: &dyn StoredMergeColumnCatalog,
merge: &MergeStmt,
) -> BTreeSet<String> {
if merge.target_column_bindings.is_empty() {
return BTreeSet::new();
}
let live = catalog
.stored_merge_target_definitions(&merge.target)
.unwrap_or_default()
.into_iter()
.filter_map(|column| column.object_id)
.collect::<BTreeSet<_>>();
merge
.target_column_bindings
.iter()
.filter(|(_, binding)| !live.contains(&binding.object_id))
.map(|(name, _)| name.clone())
.collect()
}
pub fn normalize_stored_merge_target_columns(
catalog: &dyn StoredMergeColumnCatalog,
statement: &mut Statement,
) -> Result<(), SQLError> {
crate::catalog::stored_ast::visit_stored_statement_merges(statement, &mut |merge| {
if dropped_stored_merge_targets(catalog, merge).is_empty() {
return Ok(());
}
let current = catalog
.stored_merge_target_definitions(&merge.target)
.unwrap_or_default()
.into_iter()
.filter_map(|column| column.object_id.map(|id| (id, column.name)))
.collect::<BTreeMap<_, _>>();
let name_for = |name: &str| match merge.target_column_bindings.get(name) {
Some(binding) => current.get(&binding.object_id).cloned(),
None => Some(name.to_string()),
};
for action in &mut merge.when_clauses {
match action {
MergeWhen::UpdateMatched { assignments, .. }
| MergeWhen::UpdateNotMatchedBySource { assignments, .. } => {
assignments.retain_mut(|(group, _)| match group {
crate::ast::AssignmentTargets::Single(name) => {
if let Some(current) = name_for(&name.column) {
name.column = current;
true
} else {
false
}
}
crate::ast::AssignmentTargets::Multiple(targets) => {
let mut positions = Vec::new();
let mut surviving = Vec::new();
for (mut target, position) in std::mem::take(&mut targets.targets)
.into_iter()
.zip(std::mem::take(&mut targets.source_positions))
{
if let Some(current) = name_for(&target.column) {
target.column = current;
surviving.push(target);
positions.push(position);
}
}
targets.targets = surviving;
targets.source_positions = positions;
!targets.targets.is_empty()
}
});
}
MergeWhen::InsertNotMatched {
columns, values, ..
} => {
let mut surviving = Vec::new();
let mut expressions = Vec::new();
for (mut column, expression) in std::mem::take(columns)
.into_iter()
.zip(std::mem::take(values))
{
if let Some(current) = name_for(&column.column) {
column.column = current;
surviving.push(column);
expressions.push(expression);
}
}
*columns = surviving;
*values = expressions;
}
_ => {}
}
}
Ok(())
})
}
pub fn render_stored_merge_target_columns(
catalog: &dyn StoredMergeColumnCatalog,
statement: &mut Statement,
) -> Result<(), SQLError> {
crate::catalog::stored_ast::visit_stored_statement_merges(statement, &mut |merge| {
let current = catalog
.stored_merge_target_definitions(&merge.target)
.unwrap_or_default()
.into_iter()
.filter_map(|column| column.object_id.map(|id| (id, column.name)))
.collect::<BTreeMap<_, _>>();
let rename = |target: &mut crate::ast::AssignmentTarget| {
if let Some(binding) = merge.target_column_bindings.get(&target.column) {
if let Some(name) = current.get(&binding.object_id) {
target.column.clone_from(name);
} else if let Some(number) = binding.attribute_number {
target.column =
crate::catalog::composite_type::StoredCompositeAttribute::dropped_name(
number,
);
}
}
};
for action in &mut merge.when_clauses {
match action {
MergeWhen::UpdateMatched { assignments, .. }
| MergeWhen::UpdateNotMatchedBySource { assignments, .. } => {
for (targets, _) in assignments {
for target in targets.targets_mut() {
rename(target);
}
}
}
MergeWhen::InsertNotMatched { columns, .. } => {
for target in columns {
rename(target);
}
}
_ => {}
}
}
Ok(())
})
}
pub fn collect_target_domains(ty: &ColumnType, dependencies: &mut BTreeSet<u32>) {
match ty {
ColumnType::Domain { oid, base, .. } => {
dependencies.insert(*oid);
collect_target_domains(base, dependencies);
}
ColumnType::Array(element) => collect_target_domains(element, dependencies),
_ => {}
}
}
pub fn statement_has_removed_merge_target(
catalog: &dyn StoredMergeColumnCatalog,
statement: &Statement,
) -> Result<bool, SQLError> {
let mut changed = false;
crate::catalog::stored_ast::visit_stored_statement_merges(
&mut statement.clone(),
&mut |merge| {
changed |= !dropped_stored_merge_targets(catalog, merge).is_empty();
Ok(())
},
)?;
Ok(changed)
}
#[cfg(test)]
mod tests;