use crate::ast::{ColumnClause, ColumnClauseKind, ColumnDeclaration};
use crate::SQLError;
#[derive(Clone, Copy)]
pub struct ColumnDeclarationTarget<'a> {
pub table: &'a str,
pub partitioned: bool,
pub partition: bool,
}
pub fn check_serial_array(declaration: &ColumnDeclaration) -> Result<(), SQLError> {
if declaration.serial_array {
return Err(error("0A000", "array of serial is not implemented".into()));
}
Ok(())
}
pub fn check_column_declaration(
declaration: &ColumnDeclaration,
column: &str,
target: ColumnDeclarationTarget<'_>,
) -> Result<bool, SQLError> {
let deferrable_key = check_constraint_attributes(&declaration.clauses)?;
ClauseConflicts::new(declaration, column, target, false).check()?;
Ok(deferrable_key)
}
pub fn check_foreign_column_declaration(
declaration: &ColumnDeclaration,
column: &str,
table: &str,
) -> Result<(), SQLError> {
check_constraint_attributes(&declaration.clauses)?;
ClauseConflicts::new(
declaration,
column,
ColumnDeclarationTarget {
table,
partitioned: false,
partition: false,
},
true,
)
.check()
}
pub(crate) fn foreign_table_constraint_error(kind: &str) -> SQLError {
SQLError::Unsupported(format!(
"{kind} constraints are not supported on foreign tables"
))
}
fn check_constraint_attributes(clauses: &[ColumnClause]) -> Result<bool, SQLError> {
use ColumnClauseKind::{
Check, Deferrable, Enforced, ForeignKey, InitiallyDeferred, InitiallyImmediate,
NotDeferrable, NotEnforced, PrimaryKey, Unique,
};
let mut last = None;
let mut deferrable = false;
let mut initially_deferred = false;
let mut saw_deferrability = false;
let mut saw_initially = false;
let mut saw_enforced = false;
let mut deferrable_key = false;
let syntax = |message: &str| Err(error("42601", message.into()));
for clause in clauses {
let supports_timing = matches!(last, Some(PrimaryKey | Unique | ForeignKey));
match clause.kind {
Deferrable | NotDeferrable => {
let written = if clause.kind == Deferrable {
"DEFERRABLE"
} else {
"NOT DEFERRABLE"
};
if !supports_timing {
return syntax(&format!("misplaced {written} clause"));
}
if saw_deferrability {
return syntax("multiple DEFERRABLE/NOT DEFERRABLE clauses not allowed");
}
saw_deferrability = true;
deferrable = clause.kind == Deferrable;
if !deferrable && saw_initially && initially_deferred {
return syntax("constraint declared INITIALLY DEFERRED must be DEFERRABLE");
}
}
InitiallyDeferred | InitiallyImmediate => {
let written = if clause.kind == InitiallyDeferred {
"INITIALLY DEFERRED"
} else {
"INITIALLY IMMEDIATE"
};
if !supports_timing {
return syntax(&format!("misplaced {written} clause"));
}
if saw_initially {
return syntax("multiple INITIALLY IMMEDIATE/DEFERRED clauses not allowed");
}
saw_initially = true;
initially_deferred = clause.kind == InitiallyDeferred;
if initially_deferred {
if !saw_deferrability {
deferrable = true;
} else if !deferrable {
return syntax("constraint declared INITIALLY DEFERRED must be DEFERRABLE");
}
}
}
Enforced | NotEnforced => {
let written = if clause.kind == Enforced {
"ENFORCED"
} else {
"NOT ENFORCED"
};
if !matches!(last, Some(Check | ForeignKey)) {
return syntax(&format!("misplaced {written} clause"));
}
if saw_enforced {
return syntax("multiple ENFORCED/NOT ENFORCED clauses not allowed");
}
saw_enforced = true;
}
kind => {
last = Some(kind);
deferrable = false;
initially_deferred = false;
saw_deferrability = false;
saw_initially = false;
saw_enforced = false;
}
}
deferrable_key |= deferrable && matches!(last, Some(PrimaryKey | Unique));
}
Ok(deferrable_key)
}
static SERIAL_DEFAULT: ColumnClause = ColumnClause {
kind: ColumnClauseKind::Default,
name: None,
no_inherit: false,
};
#[derive(Clone, Copy, PartialEq, Eq)]
enum Nullability {
Unspecified,
Null,
NotNull,
}
#[derive(Default)]
struct SeenClauses {
default: bool,
identity: bool,
generated: bool,
}
struct SeenNotNull<'a> {
name: Option<&'a str>,
no_inherit: bool,
}
struct ClauseConflicts<'a> {
declaration: &'a ColumnDeclaration,
column: &'a str,
target: ColumnDeclarationTarget<'a>,
foreign: bool,
need_not_null: bool,
disallow_no_inherit: bool,
nullability: Nullability,
not_null: Option<SeenNotNull<'a>>,
seen: SeenClauses,
}
impl<'a> ClauseConflicts<'a> {
fn new(
declaration: &'a ColumnDeclaration,
column: &'a str,
target: ColumnDeclarationTarget<'a>,
foreign: bool,
) -> Self {
let disallow_no_inherit = declaration.serial
|| declaration.clauses.iter().any(|clause| {
matches!(
clause.kind,
ColumnClauseKind::Identity | ColumnClauseKind::PrimaryKey
)
});
Self {
declaration,
column,
target,
foreign,
need_not_null: declaration.serial,
disallow_no_inherit,
nullability: Nullability::Unspecified,
not_null: None,
seen: SeenClauses::default(),
}
}
fn check(mut self) -> Result<(), SQLError> {
let serial = self.declaration.serial.then_some(&SERIAL_DEFAULT);
let declaration = self.declaration;
for clause in declaration.clauses.iter().chain(serial) {
self.read(clause)?;
self.check_combinations()?;
}
Ok(())
}
fn read(&mut self, clause: &'a ColumnClause) -> Result<(), SQLError> {
match clause.kind {
ColumnClauseKind::Null => {
if self.nullability == Nullability::NotNull || self.need_not_null {
return Err(self.conflicting_nullability());
}
self.nullability = Nullability::Null;
}
ColumnClauseKind::NotNull => self.read_not_null(clause)?,
ColumnClauseKind::Default => {
if self.seen.default {
return Err(self.column_error("multiple default values specified"));
}
self.seen.default = true;
}
ColumnClauseKind::Identity => {
if self.target.partition {
return Err(error(
"0A000",
"identity columns are not supported on partitions".into(),
));
}
if self.seen.identity {
return Err(self.column_error("multiple identity specifications"));
}
self.seen.identity = true;
match self.nullability {
Nullability::Unspecified => self.need_not_null = true,
Nullability::Null => return Err(self.conflicting_nullability()),
Nullability::NotNull => {}
}
}
ColumnClauseKind::Generated => {
if self.seen.generated {
return Err(self.column_error("multiple generation clauses specified"));
}
self.seen.generated = true;
}
ColumnClauseKind::PrimaryKey => {
if self.nullability == Nullability::Null {
return Err(self.conflicting_nullability());
}
if self.foreign {
return Err(foreign_table_constraint_error("primary key"));
}
self.need_not_null = true;
}
ColumnClauseKind::Unique if self.foreign => {
return Err(foreign_table_constraint_error("unique"));
}
ColumnClauseKind::ForeignKey if self.foreign => {
return Err(foreign_table_constraint_error("foreign key"));
}
_ => {}
}
Ok(())
}
fn read_not_null(&mut self, clause: &'a ColumnClause) -> Result<(), SQLError> {
if self.target.partitioned && clause.no_inherit {
return Err(error(
"0A000",
"not-null constraints on partitioned tables cannot be NO INHERIT".into(),
));
}
if self.nullability == Nullability::Null {
return Err(self.conflicting_nullability());
}
if self.disallow_no_inherit && clause.no_inherit {
return Err(self.conflicting_no_inherit());
}
if self.nullability != Nullability::NotNull {
self.nullability = Nullability::NotNull;
self.need_not_null = false;
self.not_null = Some(SeenNotNull {
name: clause.name.as_deref(),
no_inherit: clause.no_inherit,
});
} else if let Some(previous) = &mut self.not_null {
if let (Some(previous_name), Some(name)) = (previous.name, clause.name.as_deref()) {
if previous_name != name {
return Err(error(
"XX000",
format!(
"conflicting not-null constraint names \"{previous_name}\" and \"{name}\""
),
));
}
}
if previous.no_inherit != clause.no_inherit {
return Err(self.conflicting_no_inherit());
}
if previous.name.is_none() {
previous.name = clause.name.as_deref();
}
}
Ok(())
}
fn check_combinations(&self) -> Result<(), SQLError> {
if self.seen.default && self.seen.identity {
return Err(self.column_error("both default and identity specified"));
}
if self.seen.default && self.seen.generated {
return Err(self.column_error("both default and generation expression specified"));
}
if self.seen.identity && self.seen.generated {
return Err(self.column_error("both identity and generation expression specified"));
}
Ok(())
}
fn column_error(&self, what: &str) -> SQLError {
error(
"42601",
format!(
"{what} for column \"{}\" of table \"{}\"",
self.column, self.target.table
),
)
}
fn conflicting_nullability(&self) -> SQLError {
error(
"42601",
format!(
"conflicting NULL/NOT NULL declarations for column \"{}\" of table \"{}\"",
self.column, self.target.table
),
)
}
fn conflicting_no_inherit(&self) -> SQLError {
error(
"42601",
format!(
"conflicting NO INHERIT declarations for not-null constraints on column \"{}\"",
self.column
),
)
}
}
fn error(sqlstate: &str, message: String) -> SQLError {
SQLError::Routine {
sqlstate: sqlstate.into(),
message,
}
}