use crate::{
assignment::columns::AssignmentColumnCatalog,
ast::{AutoIncrementKind, OverridingKind},
SQLError,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Generation {
IdentityAlways,
IdentityByDefault,
Expression,
}
pub struct GeneratedValueColumns {
columns: Vec<(String, Generation)>,
}
impl GeneratedValueColumns {
pub fn of(catalog: &dyn AssignmentColumnCatalog, table: &str) -> Result<Self, SQLError> {
let columns = catalog
.try_describe_table(table)
.map_err(|error| SQLError::Internal(format!("read generated columns: {error}")))?
.unwrap_or_default()
.into_iter()
.filter_map(|column| {
let generation = match column.auto_increment.as_ref() {
Some(provenance) if provenance.kind == AutoIncrementKind::IdentityAlways => {
Generation::IdentityAlways
}
Some(provenance) if provenance.is_identity() => Generation::IdentityByDefault,
_ if column.generated.is_some() => Generation::Expression,
_ => return None,
};
Some((column.name, generation))
})
.collect();
Ok(Self { columns })
}
#[must_use]
pub fn contains(&self, column: &str) -> bool {
self.columns.iter().any(|(name, generation)| {
name == column
&& matches!(
generation,
Generation::IdentityAlways | Generation::IdentityByDefault
)
})
}
pub fn validate_insert<'a>(
&self,
targets: impl IntoIterator<Item = (&'a str, bool)>,
overriding: Option<OverridingKind>,
) -> Result<(), SQLError> {
let supplied = targets
.into_iter()
.filter_map(|(column, supplied)| supplied.then_some(column))
.collect::<std::collections::BTreeSet<_>>();
for (column, generation) in &self.columns {
if !supplied.contains(column.as_str()) {
continue;
}
match generation {
Generation::Expression => return Err(generated_column_insert_error(column)),
Generation::IdentityAlways if overriding.is_none() => {
return Err(generated_always_insert_error(column));
}
Generation::IdentityAlways | Generation::IdentityByDefault => {}
}
}
Ok(())
}
pub fn validate_update<'a>(
&self,
assignments: impl IntoIterator<Item = (&'a str, bool)>,
) -> Result<(), SQLError> {
let assigned = assignments
.into_iter()
.filter_map(|(column, default)| (!default).then_some(column))
.collect::<std::collections::BTreeSet<_>>();
for (column, generation) in &self.columns {
if !assigned.contains(column.as_str()) {
continue;
}
match generation {
Generation::Expression => return Err(generated_column_update_error(column)),
Generation::IdentityAlways => return Err(generated_always_update_error(column)),
Generation::IdentityByDefault => {}
}
}
Ok(())
}
}
#[must_use]
pub fn generated_always_insert_error(column: &str) -> SQLError {
SQLError::Diagnostic {
sqlstate: "428C9".into(),
message: format!("cannot insert a non-DEFAULT value into column \"{column}\""),
detail: Some(format!(
"Column \"{column}\" is an identity column defined as GENERATED ALWAYS."
)),
hint: Some("Use OVERRIDING SYSTEM VALUE to override.".into()),
}
}
#[must_use]
pub fn generated_always_update_error(column: &str) -> SQLError {
SQLError::Diagnostic {
sqlstate: "428C9".into(),
message: format!("column \"{column}\" can only be updated to DEFAULT"),
detail: Some(format!(
"Column \"{column}\" is an identity column defined as GENERATED ALWAYS."
)),
hint: None,
}
}
#[must_use]
pub fn generated_column_insert_error(column: &str) -> SQLError {
SQLError::Diagnostic {
sqlstate: "428C9".into(),
message: format!("cannot insert a non-DEFAULT value into column \"{column}\""),
detail: Some(format!("Column \"{column}\" is a generated column.")),
hint: None,
}
}
#[must_use]
pub fn generated_column_update_error(column: &str) -> SQLError {
SQLError::Diagnostic {
sqlstate: "428C9".into(),
message: format!("column \"{column}\" can only be updated to DEFAULT"),
detail: Some(format!("Column \"{column}\" is a generated column.")),
hint: None,
}
}