extern crate alloc;
use alloc::collections::{BTreeMap, BTreeSet};
use alloc::vec::Vec;
use spg_storage::row_header::{RelId, RowId};
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub enum LockMode {
KeyShare,
Share,
NoKeyUpdate,
Exclusive,
}
impl LockMode {
#[must_use]
pub fn conflicts_with(self, requested: LockMode) -> bool {
use LockMode::{Exclusive, KeyShare, NoKeyUpdate, Share};
match self {
KeyShare => matches!(requested, Exclusive),
Share => matches!(requested, NoKeyUpdate | Exclusive),
NoKeyUpdate => !matches!(requested, KeyShare),
Exclusive => true,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum WaitPolicy {
Wait,
NoWait,
SkipLocked,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum LockOutcome {
Granted,
WouldBlock { on: Vec<u64> },
Skip,
NotAvailable,
Deadlock { victim: u64 },
}
#[derive(Debug, Default, Clone)]
struct LockEntry {
holders: Vec<(u64, LockMode)>,
waiters: Vec<u64>,
}
#[derive(Debug, Default, Clone)]
pub struct LockTable {
entries: BTreeMap<(RelId, RowId), LockEntry>,
wait_for: BTreeMap<u64, BTreeSet<u64>>,
}
impl LockTable {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn acquire(
&mut self,
rel: RelId,
row: RowId,
mode: LockMode,
version: u64,
policy: WaitPolicy,
) -> LockOutcome {
let entry = self.entries.entry((rel, row)).or_default();
let mut blockers: Vec<u64> = Vec::new();
for &(hv, hmode) in &entry.holders {
if hv != version && hmode.conflicts_with(mode) && !blockers.contains(&hv) {
blockers.push(hv);
}
}
if blockers.is_empty() {
if !entry
.holders
.iter()
.any(|&(hv, hm)| hv == version && hm == mode)
{
entry.holders.push((version, mode));
}
entry.waiters.retain(|&w| w != version);
self.wait_for.remove(&version);
return LockOutcome::Granted;
}
match policy {
WaitPolicy::NoWait => LockOutcome::NotAvailable,
WaitPolicy::SkipLocked => LockOutcome::Skip,
WaitPolicy::Wait => {
if !entry.waiters.contains(&version) {
entry.waiters.push(version);
}
let edges = self.wait_for.entry(version).or_default();
for &b in &blockers {
edges.insert(b);
}
if let Some(cycle) = self.find_cycle(version) {
let victim = cycle.into_iter().max().unwrap_or(version);
return LockOutcome::Deadlock { victim };
}
LockOutcome::WouldBlock { on: blockers }
}
}
}
pub fn release_all(&mut self, version: u64) {
self.entries.retain(|_, e| {
e.holders.retain(|&(hv, _)| hv != version);
e.waiters.retain(|&w| w != version);
!(e.holders.is_empty() && e.waiters.is_empty())
});
self.wait_for.remove(&version);
for edges in self.wait_for.values_mut() {
edges.remove(&version);
}
}
#[must_use]
pub fn locked_row_count(&self) -> usize {
self.entries.len()
}
fn find_cycle(&self, start: u64) -> Option<BTreeSet<u64>> {
let mut stack: Vec<u64> = Vec::new();
let mut on_path: BTreeSet<u64> = BTreeSet::new();
let mut visited: BTreeSet<u64> = BTreeSet::new();
stack.push(start);
self.dfs_cycle(start, start, &mut on_path, &mut visited, &mut stack)
}
fn dfs_cycle(
&self,
start: u64,
node: u64,
on_path: &mut BTreeSet<u64>,
visited: &mut BTreeSet<u64>,
path: &mut Vec<u64>,
) -> Option<BTreeSet<u64>> {
on_path.insert(node);
visited.insert(node);
if let Some(edges) = self.wait_for.get(&node) {
for &next in edges {
if next == start {
let mut cyc: BTreeSet<u64> = on_path.iter().copied().collect();
cyc.insert(start);
return Some(cyc);
}
if !on_path.contains(&next) {
path.push(next);
if let Some(c) = self.dfs_cycle(start, next, on_path, visited, path) {
return Some(c);
}
path.pop();
}
}
}
on_path.remove(&node);
None
}
}
impl crate::Engine {
pub(crate) fn run_locking_prepass(
&mut self,
stmt: &spg_sql::ast::SelectStatement,
) -> Result<(), crate::EngineError> {
use spg_sql::ast::{LockStrength as LS, LockWait as LW};
let Some(lock) = &stmt.locking else {
return Ok(());
};
let Some(from) = &stmt.from else {
return Ok(()); };
let derived = from.primary.lateral_subquery.is_some()
|| from.primary.unnest_expr.is_some()
|| from.primary.generate_series_args.is_some()
|| from.primary.table_fn_call.is_some();
if !from.joins.is_empty() || derived {
self.notice(alloc::format!(
"{} over a join or subquery is accepted but NOT enforced by SPG yet; \
rows are returned unlocked",
lock_verb(lock.strength)
));
return Ok(());
}
let tname = from.primary.name.clone();
let Some(table) = self.active_catalog().get(&tname) else {
return Ok(()); };
let mode = match lock.strength {
LS::KeyShare => LockMode::KeyShare,
LS::Share => LockMode::Share,
LS::NoKeyUpdate => LockMode::NoKeyUpdate,
LS::Update => LockMode::Exclusive,
};
let policy = match lock.policy {
LW::Wait => WaitPolicy::Wait,
LW::NoWait => WaitPolicy::NoWait,
LW::SkipLocked => WaitPolicy::SkipLocked,
};
let version = self
.current_tx
.and_then(|tx| self.tx_writer_versions.get(&tx).copied())
.unwrap_or(0);
let rel = table.rel_id();
let snap = self.current_snapshot();
let cols = table.schema().columns.clone();
let alias = from.primary.alias.clone();
let ctx = crate::eval::EvalContext::new(&cols, alias.as_deref())
.with_catalog(self.active_catalog());
let mut picked: alloc::vec::Vec<(usize, spg_storage::Row<'static>)> =
alloc::vec::Vec::new();
for (idx, row) in table.scan_visible(&snap) {
if let Some(pred) = &stmt.where_ {
let keep =
crate::eval::eval_expr(pred, row, &ctx).map_err(crate::EngineError::Eval)?;
if !matches!(keep, spg_storage::Value::Bool(true)) {
continue;
}
}
picked.push((idx, row.clone()));
}
if !stmt.order_by.is_empty() {
let descs: alloc::vec::Vec<bool> = stmt.order_by.iter().map(|o| o.desc).collect();
let mut tagged: alloc::vec::Vec<(alloc::vec::Vec<crate::orderby::OrderKey>, usize)> =
alloc::vec::Vec::with_capacity(picked.len());
for (idx, row) in &picked {
tagged.push((
crate::orderby::build_order_keys(&stmt.order_by, row, &ctx)?,
*idx,
));
}
tagged.sort_by(|a, b| crate::orderby::cmp_multi_key(&a.0, &b.0, &descs));
let order: alloc::vec::Vec<usize> = tagged.into_iter().map(|(_, i)| i).collect();
picked = order
.into_iter()
.map(|i| (i, spg_storage::Row::new(alloc::vec::Vec::new())))
.collect();
}
let offset = stmt.offset_literal().unwrap_or(0) as usize;
let limit = stmt.limit_literal().map(|n| n as usize);
let want = limit.map(|n| n.saturating_add(offset));
let mut skipped: alloc::collections::BTreeSet<usize> = alloc::collections::BTreeSet::new();
let mut taken = 0usize;
for (idx, _) in &picked {
if want.is_some_and(|w| taken >= w) {
break;
}
let outcome = self.acquire_row_lock(
rel,
spg_storage::row_header::RowId(*idx as u64),
mode,
version,
policy,
);
match outcome {
LockOutcome::Granted => taken += 1,
LockOutcome::Skip => {
skipped.insert(*idx);
}
LockOutcome::NotAvailable => {
return Err(crate::EngineError::Unsupported(alloc::format!(
"could not obtain lock on row in relation \"{tname}\""
)));
}
LockOutcome::WouldBlock { .. } => {
return Err(crate::EngineError::LockWouldBlock);
}
LockOutcome::Deadlock { victim } if victim == version => {
return Err(crate::EngineError::LockDeadlock);
}
LockOutcome::Deadlock { .. } => {
return Err(crate::EngineError::LockWouldBlock);
}
}
}
self.lock_skip_rows = Some((tname, skipped));
Ok(())
}
}
const fn lock_verb(s: spg_sql::ast::LockStrength) -> &'static str {
use spg_sql::ast::LockStrength as LS;
match s {
LS::Update => "FOR UPDATE",
LS::NoKeyUpdate => "FOR NO KEY UPDATE",
LS::Share => "FOR SHARE",
LS::KeyShare => "FOR KEY SHARE",
}
}
#[cfg(test)]
mod tests {
use super::*;
const R: RelId = RelId(1);
fn row(n: u64) -> RowId {
RowId(n)
}
#[test]
fn conflict_matrix_matches_pg() {
use LockMode::{Exclusive, KeyShare, NoKeyUpdate, Share};
assert!(!KeyShare.conflicts_with(KeyShare));
assert!(!KeyShare.conflicts_with(Share));
assert!(!KeyShare.conflicts_with(NoKeyUpdate));
assert!(KeyShare.conflicts_with(Exclusive));
assert!(!Share.conflicts_with(KeyShare));
assert!(!Share.conflicts_with(Share));
assert!(Share.conflicts_with(NoKeyUpdate));
assert!(Share.conflicts_with(Exclusive));
assert!(!NoKeyUpdate.conflicts_with(KeyShare)); assert!(NoKeyUpdate.conflicts_with(Share));
assert!(NoKeyUpdate.conflicts_with(NoKeyUpdate));
assert!(NoKeyUpdate.conflicts_with(Exclusive));
assert!(Exclusive.conflicts_with(KeyShare));
assert!(Exclusive.conflicts_with(Share));
assert!(Exclusive.conflicts_with(NoKeyUpdate));
assert!(Exclusive.conflicts_with(Exclusive));
}
#[test]
fn compatible_locks_both_granted() {
let mut t = LockTable::new();
assert_eq!(
t.acquire(R, row(1), LockMode::KeyShare, 10, WaitPolicy::Wait),
LockOutcome::Granted
);
assert_eq!(
t.acquire(R, row(1), LockMode::NoKeyUpdate, 20, WaitPolicy::Wait),
LockOutcome::Granted
);
}
#[test]
fn exclusive_blocks_and_nowait_skiplocked_report() {
let mut t = LockTable::new();
assert_eq!(
t.acquire(R, row(1), LockMode::Exclusive, 10, WaitPolicy::Wait),
LockOutcome::Granted
);
match t.acquire(R, row(1), LockMode::Exclusive, 20, WaitPolicy::Wait) {
LockOutcome::WouldBlock { on } => assert_eq!(on, alloc::vec![10]),
other => panic!("expected WouldBlock, got {other:?}"),
}
assert_eq!(
t.acquire(R, row(1), LockMode::Exclusive, 30, WaitPolicy::NoWait),
LockOutcome::NotAvailable
);
assert_eq!(
t.acquire(R, row(1), LockMode::Exclusive, 40, WaitPolicy::SkipLocked),
LockOutcome::Skip
);
}
#[test]
fn release_lets_a_waiter_in() {
let mut t = LockTable::new();
t.acquire(R, row(1), LockMode::Exclusive, 10, WaitPolicy::Wait);
t.acquire(R, row(1), LockMode::Exclusive, 20, WaitPolicy::Wait);
t.release_all(10);
assert_eq!(
t.acquire(R, row(1), LockMode::Exclusive, 20, WaitPolicy::Wait),
LockOutcome::Granted
);
assert_eq!(t.locked_row_count(), 1);
t.release_all(20);
assert_eq!(t.locked_row_count(), 0);
}
#[test]
fn deadlock_cycle_aborts_youngest() {
let mut t = LockTable::new();
t.acquire(R, row(1), LockMode::Exclusive, 10, WaitPolicy::Wait);
t.acquire(R, row(2), LockMode::Exclusive, 20, WaitPolicy::Wait);
assert!(matches!(
t.acquire(R, row(2), LockMode::Exclusive, 10, WaitPolicy::Wait),
LockOutcome::WouldBlock { .. }
));
assert_eq!(
t.acquire(R, row(1), LockMode::Exclusive, 20, WaitPolicy::Wait),
LockOutcome::Deadlock { victim: 20 }
);
}
#[test]
fn relock_same_version_is_idempotent() {
let mut t = LockTable::new();
assert_eq!(
t.acquire(R, row(1), LockMode::Exclusive, 10, WaitPolicy::Wait),
LockOutcome::Granted
);
assert_eq!(
t.acquire(R, row(1), LockMode::Exclusive, 10, WaitPolicy::Wait),
LockOutcome::Granted
);
}
}