use std::sync::atomic::{AtomicU64, Ordering};
use parking_lot::Mutex;
use crate::core::{Timestamp, TxnId};
const RING: usize = 4096;
const RING_MASK: u64 = RING as u64 - 1;
const SHARDS: usize = 16;
fn thread_slot() -> usize {
use std::cell::Cell;
static NEXT: AtomicU64 = AtomicU64::new(0);
thread_local! {
static SLOT: Cell<Option<usize>> = const { Cell::new(None) };
}
SLOT.with(|slot| match slot.get() {
Some(s) => s,
None => {
let s = (NEXT.fetch_add(1, Ordering::Relaxed) as usize) % SHARDS;
slot.set(Some(s));
s
}
})
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
#[non_exhaustive]
pub enum OracleConfig {
#[default]
Centralised,
}
#[repr(align(64))] struct IdShard {
next: AtomicU64,
}
#[repr(align(64))] struct ActiveShard {
snapshots: Mutex<Vec<u64>>,
}
pub(crate) struct Oracle {
next_ts: AtomicU64,
next_id: Box<[IdShard]>,
read_watermark: AtomicU64,
completed: Box<[AtomicU64]>,
active: Box<[ActiveShard]>,
}
impl Oracle {
pub(crate) fn new(_config: OracleConfig) -> Self {
Oracle {
next_ts: AtomicU64::new(1),
next_id: (0..SHARDS)
.map(|_| IdShard {
next: AtomicU64::new(0),
})
.collect(),
read_watermark: AtomicU64::new(0),
completed: (0..RING).map(|_| AtomicU64::new(0)).collect(),
active: (0..SHARDS)
.map(|_| ActiveShard {
snapshots: Mutex::new(Vec::new()),
})
.collect(),
}
}
pub(crate) fn next_txn_id(&self) -> TxnId {
let slot = thread_slot();
let n = self.next_id[slot].next.fetch_add(1, Ordering::Relaxed);
TxnId(1 + slot as u64 + SHARDS as u64 * n)
}
pub(crate) fn begin_snapshot(&self, id: TxnId) -> Timestamp {
self.advance();
let ts = Timestamp(self.read_watermark.load(Ordering::Acquire));
self.shard(id).snapshots.lock().push(ts.raw());
ts
}
pub(crate) fn statement_snapshot(&self) -> Timestamp {
Timestamp(self.read_watermark.load(Ordering::Acquire))
}
pub(crate) fn release_snapshot(&self, id: TxnId, ts: Timestamp) {
let mut shard = self.shard(id).snapshots.lock();
if let Some(at) = shard.iter().position(|&s| s == ts.raw()) {
shard.swap_remove(at);
}
}
pub(crate) fn begin_commit(&self) -> Timestamp {
let ts = self.next_ts.fetch_add(1, Ordering::Relaxed);
while ts.saturating_sub(self.read_watermark.load(Ordering::Acquire)) >= RING as u64 {
self.advance();
std::hint::spin_loop();
}
Timestamp(ts)
}
pub(crate) fn publish(&self, ts: Timestamp) {
self.completed[(ts.raw() & RING_MASK) as usize].store(ts.raw(), Ordering::Release);
if self.read_watermark.load(Ordering::Acquire) + 1 == ts.raw() {
self.advance();
}
}
fn advance(&self) {
loop {
let current = self.read_watermark.load(Ordering::Acquire);
let candidate = current + 1;
if self.completed[(candidate & RING_MASK) as usize].load(Ordering::Acquire) != candidate
{
return; }
if self
.read_watermark
.compare_exchange_weak(current, candidate, Ordering::AcqRel, Ordering::Acquire)
.is_err()
{
std::hint::spin_loop();
}
}
}
pub(crate) fn gc_watermark(&self) -> Timestamp {
let oldest = self
.active
.iter()
.filter_map(|s| s.snapshots.lock().iter().copied().min())
.min();
match oldest {
Some(ts) => Timestamp(ts),
None => Timestamp(self.read_watermark.load(Ordering::Acquire)),
}
}
pub(crate) fn active_count(&self) -> usize {
self.active.iter().map(|s| s.snapshots.lock().len()).sum()
}
fn shard(&self, id: TxnId) -> &ActiveShard {
&self.active[(id.0 as usize) % SHARDS]
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn watermark_waits_for_out_of_order_installs() {
let o = Oracle::new(OracleConfig::Centralised);
let first = o.begin_commit();
let second = o.begin_commit();
assert!(first < second);
o.publish(second);
assert!(
o.statement_snapshot() < first,
"watermark passed a commit that is still installing"
);
o.publish(first);
assert!(
o.statement_snapshot() >= second,
"watermark should now cover both"
);
}
#[test]
fn watermark_crosses_a_long_run_in_one_go() {
let o = Oracle::new(OracleConfig::Centralised);
let stamps: Vec<_> = (0..100).map(|_| o.begin_commit()).collect();
for ts in stamps.iter().rev() {
o.publish(*ts);
}
assert_eq!(o.statement_snapshot(), *stamps.last().unwrap());
}
#[test]
fn gc_watermark_is_pinned_by_the_oldest_reader() {
let o = Oracle::new(OracleConfig::Centralised);
let ts = o.begin_commit();
o.publish(ts);
let old = o.next_txn_id();
let old_reader = o.begin_snapshot(old);
let commit = o.begin_commit();
o.publish(commit);
let new = o.next_txn_id();
let _new_reader = o.begin_snapshot(new);
assert_eq!(o.gc_watermark(), old_reader, "the oldest reader pins GC");
o.release_snapshot(old, old_reader);
assert!(
o.gc_watermark() > old_reader,
"releasing it lets GC advance"
);
}
#[test]
fn transactions_sharing_a_snapshot_are_counted_separately() {
let o = Oracle::new(OracleConfig::Centralised);
let a = TxnId(1);
let b = TxnId(1 + SHARDS as u64);
let sa = o.begin_snapshot(a);
let sb = o.begin_snapshot(b);
assert_eq!(sa, sb, "both began before anything committed");
assert_eq!(o.active_count(), 2);
let ts = o.begin_commit();
o.publish(ts);
o.release_snapshot(a, sa);
assert_eq!(o.active_count(), 1, "b is still running");
assert_eq!(o.gc_watermark(), sb, "b's snapshot must still pin GC");
o.release_snapshot(b, sb);
assert_eq!(o.active_count(), 0);
assert!(o.gc_watermark() > sb, "now GC may advance");
}
#[test]
fn concurrent_commits_leave_the_watermark_consistent() {
use std::sync::Arc;
use std::thread;
let o = Arc::new(Oracle::new(OracleConfig::Centralised));
let threads: Vec<_> = (0..8)
.map(|_| {
let o = Arc::clone(&o);
thread::spawn(move || {
for _ in 0..2_000 {
let ts = o.begin_commit();
o.publish(ts);
}
})
})
.collect();
for t in threads {
t.join().expect("worker panicked");
}
let next = o.next_ts.load(Ordering::Acquire);
assert_eq!(
o.begin_snapshot(TxnId(1)),
Timestamp(next - 1),
"watermark stalled below a fully published sequence"
);
}
}