use inillucent_base::limits::Limits;
use crate::ast::{ReferentialAction, TriggerTime};
use crate::catalog_view::{
ForeignKeyInfo, ForeignKeyTrigger, TableInfo, TableKind, TriggerEventInfo, TriggerInfo,
};
use crate::parser::parse_next_statement;
pub const VIOLATION_MESSAGE: &str = "FOREIGN KEY constraint failed";
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ForeignKeyEvent {
ChildInsert,
ChildUpdate,
ParentDelete,
ParentUpdate,
}
impl ForeignKeyEvent {}
pub fn parent_columns(key: &ForeignKeyInfo, parent: &TableInfo) -> Option<Vec<Vec<u8>>> {
if !key.parent_columns.is_empty() {
return Some(key.parent_columns.clone());
}
let primary = parent.primary_key();
if primary.is_empty() {
return None;
}
let mut names = Vec::with_capacity(primary.len());
for position in primary {
names.push(parent.columns.get(usize::from(position))?.name.clone());
}
Some(names)
}
fn child_columns(key: &ForeignKeyInfo, child: &TableInfo) -> Option<Vec<Vec<u8>>> {
let mut names = Vec::with_capacity(key.columns.len());
for position in &key.columns {
names.push(child.columns.get(usize::from(*position))?.name.clone());
}
Some(names)
}
fn quoted(name: &[u8], out: &mut String) {
out.push('"');
for byte in name {
if *byte == b'"' {
out.push('"');
}
out.push(char::from(*byte));
}
out.push('"');
}
fn quote(name: &[u8]) -> String {
let mut out = String::new();
quoted(name, &mut out);
out
}
fn qualified(database: &[u8], table: &[u8]) -> String {
let mut out = quote(database);
out.push('.');
quoted(table, &mut out);
out
}
fn conjunction(parts: &[String]) -> String {
if parts.is_empty() {
return "1".to_string();
}
parts.join(" AND ")
}
fn children_of(child: &[Vec<u8>], parent: &[Vec<u8>], row: &str) -> String {
let mut parts = Vec::with_capacity(child.len());
for (near, far) in child.iter().zip(parent.iter()) {
parts.push(format!("{} = {row}.{}", quote(near), quote(far)));
}
conjunction(&parts)
}
fn trigger_name(child: &TableInfo, key: &ForeignKeyInfo, event: ForeignKeyEvent) -> Vec<u8> {
let suffix = match event {
ForeignKeyEvent::ChildInsert => "ci",
ForeignKeyEvent::ChildUpdate => "cu",
ForeignKeyEvent::ParentDelete => "pd",
ForeignKeyEvent::ParentUpdate => "pu",
};
let mut name = b"sqlite_fk_".to_vec();
name.extend_from_slice(&child.folded);
name.push(b'_');
name.extend_from_slice(key.id.to_string().as_bytes());
name.push(b'_');
name.extend_from_slice(suffix.as_bytes());
name
}
pub fn trigger_for(
child: &TableInfo,
parent: &TableInfo,
key: &ForeignKeyInfo,
event: ForeignKeyEvent,
database: &[u8],
deferred: bool,
limits: &Limits,
) -> Option<TriggerInfo> {
let near = child_columns(key, child)?;
let far = parent_columns(key, parent)?;
if near.len() != far.len() || near.is_empty() {
return None;
}
let sql = match event {
ForeignKeyEvent::ChildInsert | ForeignKeyEvent::ChildUpdate => {
if deferred {
return None;
}
child_check(child, parent, key, event, database, &near, &far)
}
ForeignKeyEvent::ParentDelete | ForeignKeyEvent::ParentUpdate => {
parent_action(child, parent, key, event, database, &near, &far, deferred)?
}
};
build(&sql, trigger_name(child, key, event), limits)
}
fn build(sql: &str, name: Vec<u8>, limits: &Limits) -> Option<TriggerInfo> {
let parsed = parse_next_statement(sql.as_bytes(), 0, limits).ok()?;
let crate::ast::Statement::CreateTrigger {
time,
event,
when,
body,
..
} = &parsed.statement
else {
return None;
};
let event = match event {
crate::ast::TriggerEvent::Insert => TriggerEventInfo::Insert,
crate::ast::TriggerEvent::Delete => TriggerEventInfo::Delete,
crate::ast::TriggerEvent::Update(columns) => TriggerEventInfo::Update(
columns
.iter()
.map(|column| parsed.ast.folded(*column).to_vec())
.collect(),
),
};
Some(TriggerInfo {
folded: name.to_ascii_lowercase(),
name,
time: time.unwrap_or(TriggerTime::Before),
event,
when: *when,
body: body.clone(),
ast: parsed.ast,
table_database: None,
})
}
fn child_check(
child: &TableInfo,
parent: &TableInfo,
key: &ForeignKeyInfo,
event: ForeignKeyEvent,
database: &[u8],
near: &[Vec<u8>],
far: &[Vec<u8>],
) -> String {
let mut guards: Vec<String> = near
.iter()
.map(|column| format!("NEW.{} IS NOT NULL", quote(column)))
.collect();
let lookup = children_of(far, near, "NEW");
guards.push(format!(
"NOT EXISTS (SELECT 1 FROM {} WHERE {lookup})",
qualified(database, &parent.name)
));
let fires = match event {
ForeignKeyEvent::ChildUpdate => format!("BEFORE UPDATE OF {} ON", column_list(near)),
_ => "BEFORE INSERT ON".to_string(),
};
format!(
"CREATE TRIGGER {} {fires} {} BEGIN SELECT RAISE(ABORT, '{VIOLATION_MESSAGE}') WHERE {}; END",
quote(&trigger_name(child, key, event)),
quote(&child.name),
conjunction(&guards)
)
}
fn parent_action(
child: &TableInfo,
parent: &TableInfo,
key: &ForeignKeyInfo,
event: ForeignKeyEvent,
database: &[u8],
near: &[Vec<u8>],
far: &[Vec<u8>],
deferred: bool,
) -> Option<String> {
let action = match event {
ForeignKeyEvent::ParentDelete => key.on_delete,
_ => key.on_update,
};
let matching = children_of(near, far, "OLD");
let target = qualified(database, &child.name);
let body = match action {
ReferentialAction::NoAction | ReferentialAction::Restrict => {
if deferred && action == ReferentialAction::NoAction {
return None;
}
format!(
"SELECT RAISE(ABORT, '{VIOLATION_MESSAGE}') WHERE EXISTS (SELECT 1 FROM {target} WHERE {matching});"
)
}
ReferentialAction::Cascade => match event {
ForeignKeyEvent::ParentDelete => {
format!("DELETE FROM {target} WHERE {matching};")
}
_ => {
let sets: Vec<String> = near
.iter()
.zip(far.iter())
.map(|(child_column, parent_column)| {
format!("{} = NEW.{}", quote(child_column), quote(parent_column))
})
.collect();
format!("UPDATE {target} SET {} WHERE {matching};", sets.join(", "))
}
},
ReferentialAction::SetNull => {
let sets: Vec<String> = near
.iter()
.map(|column| format!("{} = NULL", quote(column)))
.collect();
format!("UPDATE {target} SET {} WHERE {matching};", sets.join(", "))
}
ReferentialAction::SetDefault => {
let mut sets = Vec::with_capacity(near.len());
for (position, column) in key.columns.iter().zip(near.iter()) {
let default = child
.columns
.get(usize::from(*position))
.and_then(|info| info.default_sql.clone())
.unwrap_or_else(|| b"NULL".to_vec());
sets.push(format!(
"{} = ({})",
quote(column),
String::from_utf8_lossy(&default)
));
}
format!("UPDATE {target} SET {} WHERE {matching};", sets.join(", "))
}
};
let time = if action == ReferentialAction::Restrict {
"BEFORE"
} else {
"AFTER"
};
let fires = match event {
ForeignKeyEvent::ParentDelete => format!("{time} DELETE ON"),
_ => format!("{time} UPDATE OF {} ON", column_list(far)),
};
let guard = match event {
ForeignKeyEvent::ParentUpdate => {
let changed: Vec<String> = far
.iter()
.map(|column| {
let name = quote(column);
format!("OLD.{name} IS NOT NEW.{name}")
})
.collect();
format!(" WHEN {}", changed.join(" OR "))
}
_ => String::new(),
};
Some(format!(
"CREATE TRIGGER {} {fires} {}{guard} BEGIN {body} END",
quote(&trigger_name(child, key, event)),
quote(&parent.name)
))
}
fn column_list(columns: &[Vec<u8>]) -> String {
columns
.iter()
.map(|column| quote(column))
.collect::<Vec<_>>()
.join(", ")
}
pub fn plan_schema(tables: &mut [TableInfo], database: &[u8], limits: &Limits) {
mark_cycles(tables);
let snapshot: Vec<TableInfo> = tables.to_vec();
for table in tables.iter_mut() {
if table.kind != TableKind::Table {
continue;
}
table.foreign_key_triggers = plan_table(table, &snapshot, database, limits);
}
}
fn mark_cycles(tables: &mut [TableInfo]) {
let edges: Vec<(Vec<u8>, Vec<u8>)> = tables
.iter()
.flat_map(|table| {
table
.foreign_keys
.iter()
.map(|key| (table.folded.clone(), key.parent_folded.clone()))
})
.collect();
for table in tables.iter_mut() {
for key in &mut table.foreign_keys {
key.cyclic = reaches(&edges, &key.parent_folded, &table.folded);
}
}
}
fn reaches(edges: &[(Vec<u8>, Vec<u8>)], from: &[u8], wanted: &[u8]) -> bool {
let mut seen: Vec<Vec<u8>> = Vec::new();
let mut pending: Vec<Vec<u8>> = vec![from.to_vec()];
while let Some(table) = pending.pop() {
if table == wanted {
return true;
}
if seen.contains(&table) {
continue;
}
seen.push(table.clone());
for (child, parent) in edges {
if *child == table {
pending.push(parent.clone());
}
}
}
false
}
pub fn sweep_statement(
child: &TableInfo,
parent: &TableInfo,
key: &ForeignKeyInfo,
database: &[u8],
) -> Option<String> {
let near = child_columns(key, child)?;
let far = parent_columns(key, parent)?;
if near.len() != far.len() || near.is_empty() {
return None;
}
let outer = quote(&child.name);
let mut guards: Vec<String> = near
.iter()
.map(|column| format!("{outer}.{} IS NOT NULL", quote(column)))
.collect();
let lookup: Vec<String> = far
.iter()
.zip(near.iter())
.map(|(parent_column, child_column)| {
format!(
"p.{} = {outer}.{}",
quote(parent_column),
quote(child_column)
)
})
.collect();
guards.push(format!(
"NOT EXISTS (SELECT 1 FROM {} AS p WHERE {})",
qualified(database, &parent.name),
conjunction(&lookup)
));
let target = qualified(database, &child.name);
let where_clause = conjunction(&guards);
match key.on_delete {
ReferentialAction::Cascade => Some(format!("DELETE FROM {target} WHERE {where_clause}")),
ReferentialAction::SetNull => {
let sets: Vec<String> = near
.iter()
.map(|column| format!("{} = NULL", quote(column)))
.collect();
Some(format!(
"UPDATE {target} SET {} WHERE {where_clause}",
sets.join(", ")
))
}
ReferentialAction::SetDefault => {
let mut sets = Vec::with_capacity(near.len());
for (position, column) in key.columns.iter().zip(near.iter()) {
let default = child
.columns
.get(usize::from(*position))
.and_then(|info| info.default_sql.clone())
.unwrap_or_else(|| b"NULL".to_vec());
sets.push(format!(
"{} = ({})",
quote(column),
String::from_utf8_lossy(&default)
));
}
Some(format!(
"UPDATE {target} SET {} WHERE {where_clause}",
sets.join(", ")
))
}
ReferentialAction::NoAction | ReferentialAction::Restrict => None,
}
}
fn plan_table(
table: &TableInfo,
tables: &[TableInfo],
database: &[u8],
limits: &Limits,
) -> Vec<ForeignKeyTrigger> {
let mut planned = Vec::new();
for key in &table.foreign_keys {
let parent = tables
.iter()
.find(|candidate| candidate.folded == key.parent_folded);
let Some(parent) = parent else {
planned.push(unusable(
key,
format!(
"no such table: {}.{}",
String::from_utf8_lossy(database),
String::from_utf8_lossy(&key.parent)
),
true,
key.parent_folded == table.folded,
));
continue;
};
if !parent_key_is_unique(parent, key) {
planned.push(unusable(
key,
mismatch(table, parent),
true,
parent.folded == table.folded,
));
continue;
}
for event in [ForeignKeyEvent::ChildInsert, ForeignKeyEvent::ChildUpdate] {
if let Some(trigger) = trigger_for(table, parent, key, event, database, false, limits) {
planned.push(ForeignKeyTrigger {
is_check: true,
deferred: key.is_deferred(),
trigger: Some(trigger),
fault: Vec::new(),
self_referencing: parent.folded == table.folded,
});
}
}
}
for child in tables {
if child.kind != TableKind::Table {
continue;
}
for key in &child.foreign_keys {
if key.parent_folded != table.folded {
continue;
}
if !parent_key_is_unique(table, key) {
planned.push(unusable(
key,
mismatch(child, table),
false,
child.folded == table.folded,
));
continue;
}
for event in [ForeignKeyEvent::ParentDelete, ForeignKeyEvent::ParentUpdate] {
let Some(trigger) = trigger_for(child, table, key, event, database, false, limits)
else {
continue;
};
let action = match event {
ForeignKeyEvent::ParentDelete => key.on_delete,
_ => key.on_update,
};
planned.push(ForeignKeyTrigger {
is_check: action == ReferentialAction::NoAction,
deferred: key.is_deferred(),
trigger: Some(trigger),
fault: Vec::new(),
self_referencing: child.folded == table.folded,
});
}
}
}
planned
}
fn mismatch(child: &TableInfo, parent: &TableInfo) -> String {
format!(
"foreign key mismatch - \"{}\" referencing \"{}\"",
String::from_utf8_lossy(&child.name),
String::from_utf8_lossy(&parent.name)
)
}
fn unusable(
key: &ForeignKeyInfo,
message: String,
is_check: bool,
self_referencing: bool,
) -> ForeignKeyTrigger {
ForeignKeyTrigger {
is_check,
deferred: key.is_deferred(),
trigger: None,
fault: message.into_bytes(),
self_referencing,
}
}
pub fn parent_key_is_unique(parent: &TableInfo, key: &ForeignKeyInfo) -> bool {
let Some(wanted) = parent_columns(key, parent) else {
return false;
};
let folded: Vec<Vec<u8>> = wanted
.iter()
.map(|name| name.to_ascii_lowercase())
.collect();
if folded.len() == 1 {
if let Some(alias) = parent.rowid_alias {
if let Some(column) = parent.columns.get(usize::from(alias)) {
if folded.first() == Some(&column.folded) {
return true;
}
}
}
}
let primary = parent.primary_key();
if !primary.is_empty() && primary.len() == folded.len() {
let names: Vec<Vec<u8>> = primary
.iter()
.filter_map(|position| parent.columns.get(usize::from(*position)))
.map(|column| column.folded.clone())
.collect();
if same_set(&names, &folded) {
return true;
}
}
parent.indexes.iter().any(|index| {
index.unique && index.columns.len() == folded.len() && {
let names: Vec<Vec<u8>> = index
.columns
.iter()
.filter_map(|key| key.column)
.filter_map(|position| parent.columns.get(usize::from(position)))
.map(|column| column.folded.clone())
.collect();
same_set(&names, &folded)
}
})
}
fn same_set(left: &[Vec<u8>], right: &[Vec<u8>]) -> bool {
left.len() == right.len() && right.iter().all(|name| left.contains(name))
}
pub fn violation_query(
child: &TableInfo,
parent: &TableInfo,
key: &ForeignKeyInfo,
database: &[u8],
) -> Option<String> {
let near = child_columns(key, child)?;
let far = parent_columns(key, parent)?;
if near.len() != far.len() || near.is_empty() {
return None;
}
let mut guards: Vec<String> = near
.iter()
.map(|column| format!("c.{} IS NOT NULL", quote(column)))
.collect();
let lookup: Vec<String> = far
.iter()
.zip(near.iter())
.map(|(parent_column, child_column)| {
format!("p.{} = c.{}", quote(parent_column), quote(child_column))
})
.collect();
guards.push(format!(
"NOT EXISTS (SELECT 1 FROM {} AS p WHERE {})",
qualified(database, &parent.name),
conjunction(&lookup)
));
let identity = if child.without_rowid {
"NULL"
} else {
"c.rowid"
};
Some(format!(
"SELECT {identity} FROM {} AS c WHERE {}",
qualified(database, &child.name),
conjunction(&guards)
))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::catalog_view::{ColumnInfo, TableKind};
use inillucent_value::Affinity;
fn table(name: &[u8], columns: &[&[u8]]) -> TableInfo {
TableInfo {
name: name.to_vec(),
folded: name.to_ascii_lowercase(),
database: 0,
root: 2,
columns: columns
.iter()
.map(|column| ColumnInfo {
name: column.to_vec(),
folded: column.to_ascii_lowercase(),
declared_type: Vec::new(),
affinity: Affinity::Blob,
collation: b"binary".to_vec(),
not_null: false,
not_null_conflict: None,
primary_key_conflict: None,
default_sql: None,
primary_key_position: None,
hidden: false,
generated: false,
stored: false,
generated_sql: None,
})
.collect(),
rowid_alias: None,
without_rowid: false,
strict: false,
autoincrement: false,
kind: TableKind::Table,
create_sql: Vec::new(),
indexes: Vec::new(),
view: None,
triggers: Vec::new(),
analysed_rows: None,
checks: Vec::new(),
foreign_keys: Vec::new(),
foreign_key_triggers: Vec::new(),
module: None,
}
}
fn key(on_delete: ReferentialAction, on_update: ReferentialAction) -> ForeignKeyInfo {
ForeignKeyInfo {
id: 0,
columns: vec![1],
parent: b"p".to_vec(),
parent_folded: b"p".to_vec(),
parent_columns: vec![b"id".to_vec()],
on_delete,
on_update,
match_clause: Vec::new(),
deferrable: false,
initially_deferred: false,
cyclic: false,
}
}
#[test]
fn every_generated_trigger_parses() {
let child = table(b"c", &[b"id", b"pid"]);
let parent = table(b"p", &[b"id"]);
let limits = Limits::default();
let actions = [
ReferentialAction::NoAction,
ReferentialAction::Restrict,
ReferentialAction::Cascade,
ReferentialAction::SetNull,
ReferentialAction::SetDefault,
];
let events = [
ForeignKeyEvent::ChildInsert,
ForeignKeyEvent::ChildUpdate,
ForeignKeyEvent::ParentDelete,
ForeignKeyEvent::ParentUpdate,
];
for action in actions {
let key = key(action, action);
for event in events {
let built = trigger_for(&child, &parent, &key, event, b"main", false, &limits);
assert!(
built.is_some(),
"{action:?} on {event:?} produced no trigger"
);
}
}
}
#[test]
fn the_child_check_reads_as_it_should() {
let child = table(b"c", &[b"id", b"pid"]);
let parent = table(b"p", &[b"id"]);
let key = key(ReferentialAction::NoAction, ReferentialAction::NoAction);
let sql = child_check(
&child,
&parent,
&key,
ForeignKeyEvent::ChildInsert,
b"main",
&[b"pid".to_vec()],
&[b"id".to_vec()],
);
assert!(sql.contains("BEFORE INSERT ON \"c\""), "{sql}");
assert!(sql.contains("NEW.\"pid\" IS NOT NULL"), "{sql}");
assert!(sql.contains("NOT EXISTS"), "{sql}");
assert!(sql.contains("FOREIGN KEY constraint failed"), "{sql}");
}
#[test]
fn restrict_fires_before_and_no_action_after() {
let child = table(b"c", &[b"id", b"pid"]);
let parent = table(b"p", &[b"id"]);
let limits = Limits::default();
for (action, expected) in [
(ReferentialAction::Restrict, TriggerTime::Before),
(ReferentialAction::NoAction, TriggerTime::After),
] {
let key = key(action, action);
let built = trigger_for(
&child,
&parent,
&key,
ForeignKeyEvent::ParentDelete,
b"main",
false,
&limits,
)
.expect("the trigger is generated");
assert_eq!(built.time, expected, "{action:?}");
}
}
#[test]
fn a_deferred_key_defers_only_its_checks() {
let child = table(b"c", &[b"id", b"pid"]);
let parent = table(b"p", &[b"id"]);
let limits = Limits::default();
let deferred = key(ReferentialAction::NoAction, ReferentialAction::NoAction);
assert!(trigger_for(
&child,
&parent,
&deferred,
ForeignKeyEvent::ChildInsert,
b"main",
true,
&limits
)
.is_none());
assert!(trigger_for(
&child,
&parent,
&deferred,
ForeignKeyEvent::ParentDelete,
b"main",
true,
&limits
)
.is_none());
let restrict = key(ReferentialAction::Restrict, ReferentialAction::Restrict);
assert!(trigger_for(
&child,
&parent,
&restrict,
ForeignKeyEvent::ParentDelete,
b"main",
true,
&limits
)
.is_some());
let cascade = key(ReferentialAction::Cascade, ReferentialAction::Cascade);
assert!(trigger_for(
&child,
&parent,
&cascade,
ForeignKeyEvent::ParentDelete,
b"main",
true,
&limits
)
.is_some());
}
#[test]
fn a_parent_update_guards_on_the_key_changing() {
let child = table(b"c", &[b"id", b"pid"]);
let parent = table(b"p", &[b"id"]);
let key = key(ReferentialAction::Cascade, ReferentialAction::Cascade);
let sql = parent_action(
&child,
&parent,
&key,
ForeignKeyEvent::ParentUpdate,
b"main",
&[b"pid".to_vec()],
&[b"id".to_vec()],
false,
)
.expect("the trigger is generated");
assert!(sql.contains("AFTER UPDATE OF \"id\""), "{sql}");
assert!(sql.contains("WHEN OLD.\"id\" IS NOT NEW.\"id\""), "{sql}");
assert!(sql.contains("SET \"pid\" = NEW.\"id\""), "{sql}");
assert!(sql.contains("WHERE \"pid\" = OLD.\"id\""), "{sql}");
}
#[test]
fn an_awkward_identifier_is_quoted() {
assert_eq!(quote(b"we\"ird"), "\"we\"\"ird\"");
let child = table(b"we\"ird", &[b"id", b"pid"]);
let parent = table(b"p", &[b"id"]);
let key = key(ReferentialAction::Cascade, ReferentialAction::Cascade);
let limits = Limits::default();
assert!(trigger_for(
&child,
&parent,
&key,
ForeignKeyEvent::ChildInsert,
b"main",
false,
&limits
)
.is_some());
}
#[test]
fn a_composite_key_compares_every_column() {
let matching = children_of(
&[b"a".to_vec(), b"b".to_vec()],
&[b"x".to_vec(), b"y".to_vec()],
"OLD",
);
assert_eq!(matching, "\"a\" = OLD.\"x\" AND \"b\" = OLD.\"y\"");
}
}