use std::any::{Any, TypeId};
use std::marker::PhantomData;
use std::ops::{Bound, RangeBounds};
use std::sync::Arc;
use std::sync::atomic::Ordering;
use crate::core::{
Encodable, Error, Index, IndexKey, IsolationLevel, Result, Timestamp, TxnId, Versioned,
Visibility,
};
use crossbeam_epoch::{Guard, Owned, Shared};
use crate::engine::ssi::TxnState;
use crate::engine::store::{Claim, Database, Slot, Table, Version};
trait WriteOp<'db> {
fn commit(&self, ts: Timestamp, gc: Timestamp, guard: &Guard);
fn abort(&self);
fn conflicting_readers(&self, txn: TxnId, gc: Timestamp) -> Vec<Arc<TxnState>>;
}
struct SlotWrite<'db, T: Versioned> {
table: &'db Table<T>,
slot: &'db Slot<T>,
installed: *const Version<T>,
replaced: *const Version<T>,
txn: TxnId,
}
impl<'db, T: Versioned> WriteOp<'db> for SlotWrite<'db, T> {
fn commit(&self, ts: Timestamp, gc: Timestamp, guard: &Guard) {
unsafe {
if let Some(prev) = self.replaced.as_ref() {
prev.end.store(ts.raw(), Ordering::Release);
}
(*self.installed).begin.store(ts.raw(), Ordering::Release);
}
self.slot.prune(gc, guard);
self.slot.unlock(self.txn);
}
fn abort(&self) {
let guard = &crossbeam_epoch::pin();
self.slot
.latest
.store(Shared::from(self.replaced), Ordering::Release);
unsafe {
guard.defer_destroy(Shared::from(self.installed));
}
self.slot.unlock(self.txn);
}
fn conflicting_readers(&self, txn: TxnId, gc: Timestamp) -> Vec<Arc<TxnState>> {
let mut readers = self.slot.readers.lock().others(txn, gc);
unsafe {
if let Some(value) = (*self.installed).value.as_ref() {
readers.extend(self.table.predicate_readers_of(value, txn, gc));
}
if let Some(previous) = self.replaced.as_ref().and_then(|v| v.value.as_ref()) {
readers.extend(self.table.predicate_readers_of(previous, txn, gc));
}
}
readers
}
}
trait UniqueClaim {
fn release(&self);
}
struct IndexClaim<'db, T: Versioned> {
table: &'db Table<T>,
position: usize,
key: IndexKey,
txn: TxnId,
}
impl<T: Versioned> UniqueClaim for IndexClaim<'_, T> {
fn release(&self) {
self.table
.release_unique(self.position, &self.key, self.txn);
}
}
trait ReadValidation<'db> {
fn revalidate(&self, now: Timestamp, guard: &Guard) -> Revalidation;
}
#[derive(Default)]
struct Revalidation {
changed: bool,
writers: Vec<Arc<TxnState>>,
}
impl Revalidation {
fn unchanged() -> Self {
Revalidation::default()
}
fn changed_by(writers: Vec<Arc<TxnState>>) -> Self {
Revalidation {
changed: true,
writers,
}
}
}
struct SlotRead<'db, T> {
slot: &'db Slot<T>,
observed: Option<usize>,
}
impl<'db, T: Send + Sync> ReadValidation<'db> for SlotRead<'db, T> {
fn revalidate(&self, now: Timestamp, guard: &Guard) -> Revalidation {
let current = self.slot.read(now, TxnId::NONE, guard);
let unchanged = match (self.observed, current) {
(None, None) => true,
(Some(a), Some(b)) => a == b as *const Version<T> as usize,
_ => false,
};
if unchanged {
Revalidation::unchanged()
} else {
Revalidation::changed_by(current.and_then(|v| v.writer.clone()).into_iter().collect())
}
}
}
struct PredicateRead<'db, T: Versioned> {
table: &'db Table<T>,
predicate: Arc<dyn Fn(&T) -> bool + Send + Sync>,
observed: Vec<(T::Key, usize)>,
}
impl<'db, T: Versioned> ReadValidation<'db> for PredicateRead<'db, T> {
fn revalidate(&self, now: Timestamp, guard: &Guard) -> Revalidation {
let current = self
.table
.matching(now, TxnId::NONE, &*self.predicate, guard);
writers_of_change(self.table, &self.observed, ¤t, now, guard)
}
}
struct IndexRangeRead<'db, T: Versioned> {
table: &'db Table<T>,
position: usize,
lo: Bound<IndexKey>,
hi: Bound<IndexKey>,
observed: Vec<(T::Key, usize)>,
}
impl<'db, T: Versioned> ReadValidation<'db> for IndexRangeRead<'db, T> {
fn revalidate(&self, now: Timestamp, guard: &Guard) -> Revalidation {
let current = self.table.matching_in_index(
self.position,
&self.lo,
&self.hi,
now,
TxnId::NONE,
guard,
);
writers_of_change(self.table, &self.observed, ¤t, now, guard)
}
}
fn writers_of_change<T: Versioned>(
table: &Table<T>,
observed: &[(T::Key, usize)],
current: &[(T::Key, &Version<T>)],
now: Timestamp,
guard: &Guard,
) -> Revalidation {
let addr = |v: &Version<T>| v as *const Version<T> as usize;
let identical = observed.len() == current.len()
&& observed
.iter()
.zip(current)
.all(|((ko, vo), (kc, vc))| ko == kc && *vo == addr(vc));
if identical {
return Revalidation::unchanged();
}
let mut writers = Vec::new();
for (key, version) in current {
let same_as_observed = observed
.iter()
.any(|(ko, vo)| ko == key && *vo == addr(version));
if !same_as_observed {
writers.extend(version.writer.clone());
}
}
for (key, _) in observed {
if current.iter().any(|(kc, _)| kc == key) {
continue;
}
if let Some(slot) = table.slot(key, guard)
&& let Some(version) = slot.read(now, TxnId::NONE, guard)
{
writers.extend(version.writer.clone());
}
}
Revalidation::changed_by(writers)
}
pub struct Transaction<'db, I: IsolationLevel> {
db: &'db Database,
id: TxnId,
snapshot: Timestamp,
state: Option<Arc<TxnState>>,
guard: Guard,
writes: Vec<Box<dyn WriteOp<'db> + 'db>>,
reads: Vec<Box<dyn ReadValidation<'db> + 'db>>,
claims: Vec<Box<dyn UniqueClaim + 'db>>,
table_cache: Vec<(TypeId, &'db (dyn Any + Send + Sync))>,
done: bool,
_level: PhantomData<fn() -> I>,
}
impl<'db, I: IsolationLevel> Transaction<'db, I> {
pub(crate) fn new(db: &'db Database) -> Self {
let id = db.oracle().next_txn_id();
let snapshot = db.oracle().begin_snapshot(id);
Transaction {
id,
snapshot,
state: I::VALIDATES_READS.then(|| TxnState::new(id)),
db,
writes: Vec::new(),
reads: Vec::new(),
claims: Vec::new(),
table_cache: Vec::new(),
done: false,
_level: PhantomData,
guard: crossbeam_epoch::pin(),
}
}
pub fn id(&self) -> TxnId {
self.id
}
pub fn snapshot(&self) -> Timestamp {
self.snapshot
}
fn statement_snapshot(&self) -> Timestamp {
if I::REFRESH_SNAPSHOT_PER_STATEMENT {
self.db.oracle().statement_snapshot()
} else {
self.snapshot
}
}
fn ssi(&self) -> &Arc<TxnState> {
self.state
.as_ref()
.expect("SSI state is only touched when I::VALIDATES_READS")
}
fn table<T: Versioned>(&mut self) -> Result<&'db Table<T>> {
let type_id = TypeId::of::<T>();
if let Some((_, erased)) = self.table_cache.iter().find(|(id, _)| *id == type_id) {
return Ok(erased
.downcast_ref::<Table<T>>()
.expect("cache is keyed by TypeId"));
}
let erased = self.db.table_erased(type_id, T::TABLE_NAME)?;
self.table_cache.push((type_id, erased));
Ok(erased
.downcast_ref::<Table<T>>()
.expect("registry is keyed by TypeId"))
}
fn ensure_live(&self) -> Result<()> {
if self.done {
Err(Error::Aborted)
} else {
Ok(())
}
}
pub fn get<T: Versioned>(&mut self, key: &T::Key) -> Result<Option<Ref<'_, T>>> {
self.ensure_live()?;
let table = self.table::<T>()?;
let snapshot = self.statement_snapshot();
let Some(slot) = table.slot(key, &self.guard) else {
if I::VALIDATES_READS {
let slot = table.slot_or_create(key, &self.guard);
let state = self
.state
.as_ref()
.expect("SSI state is only touched when I::VALIDATES_READS");
slot.readers.lock().register(state);
self.reads.push(Box::new(SlotRead::<T> {
slot,
observed: None,
}));
}
return Ok(None);
};
let version = slot.read(snapshot, self.id, &self.guard);
if I::VALIDATES_READS {
let state = self
.state
.as_ref()
.expect("SSI state is only touched when I::VALIDATES_READS");
slot.readers.lock().register(state);
self.reads.push(Box::new(SlotRead::<T> {
slot,
observed: version.map(|v| v as *const Version<T> as usize),
}));
}
Ok(version
.filter(|v| v.value.is_some())
.map(|v| Ref { version: v }))
}
fn write<T: Versioned>(
&mut self,
table: &'db Table<T>,
key: &T::Key,
value: Option<T>,
fresh: bool,
) -> Result<()> {
let slot = table.slot_or_create(key, &self.guard);
if !slot.try_lock(self.id) {
return Err(Error::WriteConflict {
table: T::TABLE_NAME,
});
}
let replaced = {
let shared = slot.latest.load(Ordering::Acquire, &self.guard);
unsafe { shared.as_ref() }
};
if I::FIRST_COMMITTER_WINS
&& let Some(current) = &replaced
&& let Visibility::CommittedAt(ts) =
Visibility::decode(current.begin.load(Ordering::Acquire))
&& ts > self.snapshot
{
slot.unlock(self.id);
return Err(Error::WriteConflict {
table: T::TABLE_NAME,
});
}
if fresh
&& slot
.read_committed_now(self.id, &self.guard)
.is_some_and(|v| v.value.is_some())
{
slot.unlock(self.id);
return Err(Error::DuplicateKey {
table: T::TABLE_NAME,
index: "primary_key",
});
}
if let Some(record) = &value {
table.index_record(record, replaced.and_then(|v| v.value.as_ref()));
}
let installed = Owned::new(Version {
begin: std::sync::atomic::AtomicU64::new(self.id.tagged()),
end: std::sync::atomic::AtomicU64::new(Timestamp::MAX.raw()),
prev: crossbeam_epoch::Atomic::null(),
value,
writer: self.state.clone(),
});
if let Some(prev) = replaced {
installed
.prev
.store(Shared::from(prev as *const Version<T>), Ordering::Relaxed);
}
let installed = installed.into_shared(&self.guard);
slot.latest.store(installed, Ordering::Release);
self.writes.push(Box::new(SlotWrite {
table,
slot,
installed: installed.as_raw(),
replaced: replaced
.map(|v| v as *const Version<T>)
.unwrap_or(std::ptr::null()),
txn: self.id,
}));
Ok(())
}
fn claim_unique_keys<T: Versioned>(
&self,
table: &'db Table<T>,
record: &T,
previous: Option<&T>,
) -> Result<Vec<Box<dyn UniqueClaim + 'db>>> {
let mut claimed: Vec<Box<dyn UniqueClaim + 'db>> = Vec::new();
let own_key = record.key();
for (position, desc) in T::indexes().iter().enumerate() {
if !desc.unique {
continue;
}
let index_key = (desc.extract)(record);
if previous.is_some_and(|prev| (desc.extract)(prev) == index_key) {
continue;
}
let outcome = match table.try_claim_unique(position, &index_key, self.id) {
Claim::Contended => Err(Error::WriteConflict {
table: T::TABLE_NAME,
}),
claim => {
if claim == Claim::Acquired {
claimed.push(Box::new(IndexClaim {
table,
position,
key: index_key.clone(),
txn: self.id,
}));
}
if table.unique_key_taken(position, &index_key, &own_key, self.id, &self.guard)
{
Err(Error::DuplicateKey {
table: T::TABLE_NAME,
index: desc.name,
})
} else {
Ok(())
}
}
};
if let Err(e) = outcome {
for claim in &claimed {
claim.release();
}
return Err(e);
}
}
Ok(claimed)
}
pub fn insert<T: Versioned>(&mut self, value: T) -> Result<()> {
self.ensure_live()?;
let table = self.table::<T>()?;
let key = value.key();
let snapshot = self.statement_snapshot();
if let Some(slot) = table.slot(&key, &self.guard)
&& slot
.read(snapshot, self.id, &self.guard)
.is_some_and(|v| v.value.is_some())
{
return Err(Error::DuplicateKey {
table: T::TABLE_NAME,
index: "primary_key",
});
}
let claimed = self.claim_unique_keys(table, &value, None)?;
self.claims.extend(claimed);
self.write(table, &key, Some(value), true)
}
pub fn update<T: Versioned>(&mut self, key: &T::Key, f: impl FnOnce(&mut T)) -> Result<bool> {
self.ensure_live()?;
let table = self.table::<T>()?;
let snapshot = self.statement_snapshot();
let Some(slot) = table.slot(key, &self.guard) else {
return Ok(false);
};
let Some(version) = slot.read(snapshot, self.id, &self.guard) else {
return Ok(false);
};
let Some(current) = version.value.as_ref() else {
return Ok(false);
};
let mut next = current.clone();
f(&mut next);
if next.key() != *key {
return Err(Error::PrimaryKeyChanged {
table: T::TABLE_NAME,
});
}
let claimed = self.claim_unique_keys(table, &next, Some(current))?;
self.claims.extend(claimed);
self.write(table, key, Some(next), false)?;
Ok(true)
}
pub fn delete<T: Versioned>(&mut self, key: &T::Key) -> Result<bool> {
self.ensure_live()?;
let table = self.table::<T>()?;
let snapshot = self.statement_snapshot();
let Some(slot) = table.slot(key, &self.guard) else {
return Ok(false);
};
if slot
.read(snapshot, self.id, &self.guard)
.is_none_or(|v| v.value.is_none())
{
return Ok(false);
}
self.write::<T>(table, key, None, false)?;
Ok(true)
}
pub fn scan<T: Versioned>(&mut self) -> Result<Vec<Ref<'_, T>>> {
self.scan_where::<T, _>(|_| true)
}
pub fn scan_where<T, P>(&mut self, predicate: P) -> Result<Vec<Ref<'_, T>>>
where
T: Versioned,
P: Fn(&T) -> bool + Send + Sync + 'static,
{
self.ensure_live()?;
let table = self.table::<T>()?;
let snapshot = self.statement_snapshot();
let predicate: Arc<dyn Fn(&T) -> bool + Send + Sync> = Arc::new(predicate);
let matched = if I::VALIDATES_READS {
let (matched, committed) =
table.matching_pair(snapshot, self.id, &*predicate, &self.guard);
let observed: Vec<(T::Key, usize)> = committed
.iter()
.map(|(k, v)| (k.clone(), *v as *const Version<T> as usize))
.collect();
table.register_predicate(self.ssi(), Arc::clone(&predicate));
self.reads.push(Box::new(PredicateRead {
table,
predicate,
observed,
}));
matched
} else {
table.matching(snapshot, self.id, &*predicate, &self.guard)
};
Ok(matched
.into_iter()
.map(|(_, version)| Ref { version })
.collect())
}
pub fn scan_index<T, K, R>(&mut self, index: Index<T, K>, range: R) -> Result<Vec<Ref<'_, T>>>
where
T: Versioned,
K: Encodable,
R: RangeBounds<K>,
{
self.ensure_live()?;
let table = self.table::<T>()?;
let snapshot = self.statement_snapshot();
let position = index.position;
let encode_bound = |b: Bound<&K>| match b {
Bound::Included(k) => Bound::Included(k.encode()),
Bound::Excluded(k) => Bound::Excluded(k.encode()),
Bound::Unbounded => Bound::Unbounded,
};
let lo = encode_bound(range.start_bound());
let hi = encode_bound(range.end_bound());
let matched = table.matching_in_index(position, &lo, &hi, snapshot, self.id, &self.guard);
if I::VALIDATES_READS {
let observed: Vec<(T::Key, usize)> = table
.matching_in_index(position, &lo, &hi, snapshot, TxnId::NONE, &self.guard)
.iter()
.map(|(k, v)| (k.clone(), *v as *const Version<T> as usize))
.collect();
let extract = T::indexes()[position].extract;
let (plo, phi) = (lo.clone(), hi.clone());
table.register_predicate(
self.ssi(),
Arc::new(move |record: &T| {
let k = extract(record);
let lo_ok = match &plo {
Bound::Included(b) => k >= *b,
Bound::Excluded(b) => k > *b,
Bound::Unbounded => true,
};
let hi_ok = match &phi {
Bound::Included(b) => k <= *b,
Bound::Excluded(b) => k < *b,
Bound::Unbounded => true,
};
lo_ok && hi_ok
}),
);
self.reads.push(Box::new(IndexRangeRead {
table,
position,
lo,
hi,
observed,
}));
}
Ok(matched
.into_iter()
.map(|(_, version)| Ref { version })
.collect())
}
fn detect_conflicts(&self) -> bool {
if self.ssi().is_aborted() {
return false;
}
let now = self.db.oracle().statement_snapshot();
for read in &self.reads {
let outcome = read.revalidate(now, &self.guard);
if !outcome.changed {
continue;
}
self.ssi().set_out_conflict();
for writer in outcome.writers {
writer.set_in_conflict();
if writer.is_pivot() && writer.is_committed() {
return false;
}
}
}
if self.ssi().is_pivot() {
return false;
}
if self.ssi().has_out_conflict() {
let gc = self.db.oracle().gc_watermark();
for write in &self.writes {
for reader in write.conflicting_readers(self.id, gc) {
reader.set_out_conflict();
self.ssi().set_in_conflict();
}
}
}
!self.ssi().is_pivot()
}
pub fn commit(mut self) -> Result<()> {
self.ensure_live()?;
let needs_validation = I::VALIDATES_READS && !self.writes.is_empty();
if needs_validation && !self.detect_conflicts() {
self.rollback();
return Err(Error::SerializationFailure);
}
{
let _decision = needs_validation.then(|| self.db.commit_lock());
if needs_validation && self.ssi().is_pivot() {
drop(_decision);
self.rollback();
return Err(Error::SerializationFailure);
}
let gc = self.db.gc_hint();
let ts = self.db.oracle().begin_commit();
for write in &self.writes {
write.commit(ts, gc, &self.guard);
}
self.db.oracle().publish(ts);
self.db.refresh_gc_hint(ts, self.id);
if let Some(state) = &self.state {
state.mark_committed(ts);
}
}
self.release_claims();
self.db.oracle().release_snapshot(self.id, self.snapshot);
self.done = true;
Ok(())
}
fn release_claims(&mut self) {
for claim in self.claims.drain(..) {
claim.release();
}
}
pub fn abort(mut self) {
self.rollback();
}
fn rollback(&mut self) {
if self.done {
return;
}
for write in self.writes.iter().rev() {
write.abort();
}
self.release_claims();
if let Some(state) = &self.state {
state.mark_aborted();
}
self.db.oracle().release_snapshot(self.id, self.snapshot);
self.done = true;
}
}
impl<'db, I: IsolationLevel> Drop for Transaction<'db, I> {
fn drop(&mut self) {
self.rollback();
}
}
pub struct Ref<'txn, T> {
version: &'txn Version<T>,
}
impl<T> std::ops::Deref for Ref<'_, T> {
type Target = T;
fn deref(&self) -> &T {
self.version
.value
.as_ref()
.expect("a Ref is only constructed for a version with a value")
}
}
impl<T: std::fmt::Debug> std::fmt::Debug for Ref<'_, T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
(**self).fmt(f)
}
}
impl<T: Clone> Ref<'_, T> {
pub fn to_owned(&self) -> T {
(**self).clone()
}
}