use dashmap::DashMap;
use parking_lot::{Condvar, Mutex};
use std::sync::atomic::{fence, AtomicBool, AtomicU64, AtomicUsize, Ordering};
use super::Key;
const PRUNE_INTERVAL_COMMITS: u64 = 4096;
const INFLIGHT_RING_SIZE: usize = 4096;
const INFLIGHT_RING_MASK: u64 = (INFLIGHT_RING_SIZE as u64) - 1;
const BARRIER_SPINS: u32 = 48;
const SLOT_EMPTY: u64 = 0;
#[inline]
fn slot_pending(ts: u64) -> u64 {
(ts << 1) | 1
}
#[inline]
fn slot_applied(ts: u64) -> u64 {
ts << 1
}
#[inline]
fn slot_is_pending(v: u64) -> bool {
v & 1 == 1
}
#[inline]
fn slot_ts(v: u64) -> u64 {
v >> 1
}
thread_local! {
static LAST_BEGIN_HINT: std::cell::Cell<(usize, u64, u64)> =
const { std::cell::Cell::new((0, 0, 0)) };
}
pub struct WriteConflictRegistry {
recent_writes: DashMap<Key, u64>,
active_snapshots: DashMap<u64, u64>,
ring: Box<[AtomicU64]>,
head: AtomicU64,
tail: AtomicU64,
reclaim_flag: AtomicBool,
pending_count: AtomicUsize,
applied_watermark: AtomicU64,
barrier_waiters: AtomicUsize,
barrier_lock: Mutex<()>,
barrier_cv: Condvar,
commit_count: AtomicU64,
gc_pins: DashMap<u64, u64>,
next_gc_pin: AtomicU64,
}
impl Default for WriteConflictRegistry {
fn default() -> Self {
Self::new()
}
}
impl WriteConflictRegistry {
pub fn new() -> Self {
let ring: Vec<AtomicU64> = (0..INFLIGHT_RING_SIZE).map(|_| AtomicU64::new(SLOT_EMPTY)).collect();
Self {
recent_writes: DashMap::new(),
active_snapshots: DashMap::new(),
ring: ring.into_boxed_slice(),
head: AtomicU64::new(0),
tail: AtomicU64::new(0),
reclaim_flag: AtomicBool::new(false),
pending_count: AtomicUsize::new(0),
applied_watermark: AtomicU64::new(0),
barrier_waiters: AtomicUsize::new(0),
barrier_lock: Mutex::new(()),
barrier_cv: Condvar::new(),
commit_count: AtomicU64::new(0),
gc_pins: DashMap::new(),
next_gc_pin: AtomicU64::new(0),
}
}
#[inline]
#[allow(clippy::indexing_slicing)]
fn slot(&self, seq: u64) -> &AtomicU64 {
&self.ring[(seq & INFLIGHT_RING_MASK) as usize]
}
#[inline]
fn instance_id(&self) -> usize {
std::ptr::from_ref(self) as usize
}
pub fn register_txn(&self, txn_id: u64, snapshot_ts: u64) {
self.active_snapshots.insert(txn_id, snapshot_ts);
}
pub fn refresh_txn(&self, txn_id: u64, snapshot_ts: u64) {
self.active_snapshots.insert(txn_id, snapshot_ts);
}
pub fn deregister_txn(&self, txn_id: u64) {
self.active_snapshots.remove(&txn_id);
}
pub fn snapshot_barrier(&self, snapshot_ts: u64) {
if self.pending_count.load(Ordering::Acquire) == 0 {
return;
}
if self.applied_watermark.load(Ordering::Acquire) >= snapshot_ts {
return;
}
if !self.has_pending_at_or_below(snapshot_ts) {
return;
}
self.snapshot_barrier_slow(snapshot_ts);
}
#[cold]
fn snapshot_barrier_slow(&self, snapshot_ts: u64) {
for _ in 0..BARRIER_SPINS {
std::thread::yield_now();
if !self.has_pending_at_or_below(snapshot_ts) {
return;
}
}
let mut guard = self.barrier_lock.lock();
self.barrier_waiters.fetch_add(1, Ordering::Relaxed);
loop {
fence(Ordering::SeqCst);
if !self.has_pending_at_or_below(snapshot_ts) {
break;
}
self.barrier_cv.wait(&mut guard);
}
self.barrier_waiters.fetch_sub(1, Ordering::Relaxed);
}
fn has_pending_at_or_below(&self, snapshot_ts: u64) -> bool {
loop {
let t = self.tail.load(Ordering::Acquire);
let h = self.head.load(Ordering::Acquire);
let mut s = t;
let mut early_false = false;
while s != h {
let v = self.slot(s).load(Ordering::Acquire);
if slot_is_pending(v) {
if slot_ts(v) <= snapshot_ts {
return true;
}
early_false = true;
break;
}
s += 1;
}
if !early_false || self.tail.load(Ordering::Acquire) == t {
return false;
}
}
}
pub fn applied_watermark(&self) -> u64 {
self.applied_watermark.load(Ordering::Acquire)
}
pub fn begin_commit(&self, commit_ts: u64) {
let h = self.head.load(Ordering::Relaxed);
if h.wrapping_sub(self.tail.load(Ordering::Acquire)) >= INFLIGHT_RING_SIZE as u64 {
self.wait_for_ring_space(h);
}
self.slot(h).store(slot_pending(commit_ts), Ordering::Release);
self.pending_count.fetch_add(1, Ordering::AcqRel);
self.head.store(h + 1, Ordering::Release);
LAST_BEGIN_HINT.with(|c| c.set((self.instance_id(), commit_ts, h)));
}
#[cold]
fn wait_for_ring_space(&self, h: u64) {
let mut spins: u32 = 0;
loop {
self.try_reclaim();
if h.wrapping_sub(self.tail.load(Ordering::Acquire)) < INFLIGHT_RING_SIZE as u64 {
return;
}
spins += 1;
if spins < 64 {
std::thread::yield_now();
} else {
std::thread::sleep(std::time::Duration::from_micros(50));
}
}
}
pub fn end_commit(&self, commit_ts: u64) {
let pending = slot_pending(commit_ts);
let applied = slot_applied(commit_ts);
let mut marked = false;
let hint = LAST_BEGIN_HINT.with(|c| c.get());
if hint.0 == self.instance_id() && hint.1 == commit_ts {
marked = self
.slot(hint.2)
.compare_exchange(pending, applied, Ordering::AcqRel, Ordering::Relaxed)
.is_ok();
}
if !marked {
let t = self.tail.load(Ordering::Acquire);
let h = self.head.load(Ordering::Acquire);
let mut s = t;
while s != h {
let slot = self.slot(s);
if slot.load(Ordering::Acquire) == pending {
marked = slot
.compare_exchange(pending, applied, Ordering::AcqRel, Ordering::Relaxed)
.is_ok();
break;
}
s += 1;
}
if !marked {
return; }
}
self.pending_count.fetch_sub(1, Ordering::AcqRel);
self.try_reclaim();
fence(Ordering::SeqCst);
if self.barrier_waiters.load(Ordering::Relaxed) > 0 {
let _guard = self.barrier_lock.lock();
self.barrier_cv.notify_all();
}
}
fn try_reclaim(&self) {
while self
.reclaim_flag
.compare_exchange(false, true, Ordering::Acquire, Ordering::Relaxed)
.is_ok()
{
loop {
let t = self.tail.load(Ordering::Relaxed); if t == self.head.load(Ordering::Acquire) {
break;
}
let slot = self.slot(t);
let v = slot.load(Ordering::Acquire);
if v == SLOT_EMPTY || slot_is_pending(v) {
break;
}
self.applied_watermark.fetch_max(slot_ts(v), Ordering::AcqRel);
slot.store(SLOT_EMPTY, Ordering::Release);
self.tail.store(t + 1, Ordering::Release);
}
self.reclaim_flag.store(false, Ordering::Release);
let t = self.tail.load(Ordering::Acquire);
if t == self.head.load(Ordering::Acquire) {
return;
}
let v = self.slot(t).load(Ordering::Acquire);
if v == SLOT_EMPTY || slot_is_pending(v) {
return;
}
}
}
pub fn validate_and_record(
&self,
write_set: &DashMap<Key, Option<Vec<u8>>>,
validate: bool,
snapshot_ts: u64,
commit_ts: u64,
) -> std::result::Result<(), (Key, u64)> {
if write_set.is_empty() {
return Ok(());
}
let mut recorded: Vec<(Key, Option<u64>)> = Vec::with_capacity(write_set.len());
let mut conflict: Option<(Key, u64)> = None;
for item in write_set.iter() {
let key = item.key();
let entry = self.recent_writes.entry(key.clone());
match entry {
dashmap::mapref::entry::Entry::Occupied(mut occupied) => {
let prior = *occupied.get();
if validate && prior > snapshot_ts {
conflict = Some((key.clone(), prior));
break;
}
occupied.insert(commit_ts);
recorded.push((key.clone(), Some(prior)));
}
dashmap::mapref::entry::Entry::Vacant(vacant) => {
vacant.insert(commit_ts);
recorded.push((key.clone(), None));
}
}
}
if let Some((key, ts)) = conflict {
for (k, prior) in recorded {
match prior {
Some(ts) => {
self.recent_writes.insert(k, ts);
}
None => {
self.recent_writes.remove(&k);
}
}
}
return Err((key, ts));
}
let n = self.commit_count.fetch_add(1, Ordering::Relaxed);
if n % PRUNE_INTERVAL_COMMITS == PRUNE_INTERVAL_COMMITS - 1 {
self.prune();
}
Ok(())
}
fn prune(&self) {
let min_active = self
.active_snapshots
.iter()
.map(|entry| *entry.value())
.min()
.unwrap_or(u64::MAX);
self.recent_writes.retain(|_, ts| *ts >= min_active);
}
pub fn tracked_keys(&self) -> usize {
self.recent_writes.len()
}
pub fn pin_snapshot(&self, snapshot_ts: u64) -> u64 {
let id = self.next_gc_pin.fetch_add(1, Ordering::Relaxed);
self.gc_pins.insert(id, snapshot_ts);
id
}
pub fn unpin_snapshot(&self, pin_id: u64) {
self.gc_pins.remove(&pin_id);
}
pub fn min_pinned_snapshot(&self) -> Option<u64> {
let pin_min = self.gc_pins.iter().map(|e| *e.value()).min();
let active_min = self.active_snapshots.iter().map(|e| *e.value()).min();
match (pin_min, active_min) {
(Some(a), Some(b)) => Some(a.min(b)),
(Some(a), None) => Some(a),
(None, Some(b)) => Some(b),
(None, None) => None,
}
}
pub fn gc_pin_count(&self) -> usize {
self.gc_pins.len()
}
}
pub struct GcPinGuard {
registry: std::sync::Arc<WriteConflictRegistry>,
pin_id: u64,
}
impl GcPinGuard {
pub fn new(registry: std::sync::Arc<WriteConflictRegistry>, snapshot_ts: u64) -> Self {
let pin_id = registry.pin_snapshot(snapshot_ts);
Self { registry, pin_id }
}
}
impl Drop for GcPinGuard {
fn drop(&mut self) {
self.registry.unpin_snapshot(self.pin_id);
}
}
#[cfg(test)]
mod tests {
use super::*;
fn ws(keys: &[&str]) -> DashMap<Key, Option<Vec<u8>>> {
let m = DashMap::new();
for k in keys {
m.insert(k.as_bytes().to_vec(), Some(vec![1]));
}
m
}
#[test]
fn first_committer_wins() {
let reg = WriteConflictRegistry::new();
assert!(reg.validate_and_record(&ws(&["k1"]), true, 10, 20).is_ok());
let err = reg.validate_and_record(&ws(&["k1"]), true, 15, 25).unwrap_err();
assert_eq!(err.1, 20);
assert!(reg.validate_and_record(&ws(&["k1"]), true, 30, 35).is_ok());
}
#[test]
fn losing_commit_undoes_partial_records() {
let reg = WriteConflictRegistry::new();
assert!(reg.validate_and_record(&ws(&["b"]), true, 10, 20).is_ok());
let _ = reg.validate_and_record(&ws(&["a", "b"]), true, 15, 25).unwrap_err();
assert!(reg.validate_and_record(&ws(&["a"]), true, 5, 30).is_ok());
}
#[test]
fn non_validating_commits_record() {
let reg = WriteConflictRegistry::new();
assert!(reg.validate_and_record(&ws(&["k"]), false, 0, 20).is_ok());
let err = reg.validate_and_record(&ws(&["k"]), true, 10, 25).unwrap_err();
assert_eq!(err.1, 20);
}
#[test]
fn prune_respects_active_snapshots() {
let reg = WriteConflictRegistry::new();
reg.register_txn(1, 50);
for i in 0..(PRUNE_INTERVAL_COMMITS + 1) {
let m = ws(&[format!("k{i}").as_str()]);
let _ = reg.validate_and_record(&m, false, 0, 40 + i);
}
assert!(reg.tracked_keys() > 0);
reg.deregister_txn(1);
}
#[test]
fn gc_pins_track_min_snapshot() {
let reg = std::sync::Arc::new(WriteConflictRegistry::new());
assert_eq!(reg.min_pinned_snapshot(), None);
let p1 = reg.pin_snapshot(100);
let _p2 = reg.pin_snapshot(50);
reg.register_txn(7, 80);
assert_eq!(reg.min_pinned_snapshot(), Some(50));
reg.unpin_snapshot(p1);
assert_eq!(reg.min_pinned_snapshot(), Some(50));
reg.deregister_txn(7);
{
let _g = GcPinGuard::new(reg.clone(), 10);
assert_eq!(reg.min_pinned_snapshot(), Some(10));
}
assert_eq!(reg.min_pinned_snapshot(), Some(50));
}
#[test]
fn snapshot_barrier_waits_for_inflight() {
let reg = std::sync::Arc::new(WriteConflictRegistry::new());
reg.begin_commit(10);
let r2 = reg.clone();
let h = std::thread::spawn(move || {
r2.snapshot_barrier(15);
});
std::thread::sleep(std::time::Duration::from_millis(20));
assert!(!h.is_finished(), "barrier returned while commit in flight");
reg.end_commit(10);
h.join().unwrap();
reg.begin_commit(100);
reg.snapshot_barrier(50);
reg.end_commit(100);
}
#[test]
fn watermark_advances_past_contiguous_applied_prefix() {
let reg = WriteConflictRegistry::new();
reg.begin_commit(10);
reg.begin_commit(20);
reg.begin_commit(30);
assert_eq!(reg.applied_watermark(), 0);
reg.end_commit(10);
assert_eq!(reg.applied_watermark(), 10);
reg.end_commit(20);
assert_eq!(reg.applied_watermark(), 20);
reg.end_commit(30);
assert_eq!(reg.applied_watermark(), 30);
}
#[test]
fn out_of_order_end_commit_holds_watermark_until_gap_fills() {
let reg = std::sync::Arc::new(WriteConflictRegistry::new());
reg.begin_commit(10);
reg.begin_commit(20);
reg.begin_commit(30);
reg.end_commit(30);
reg.end_commit(20);
assert_eq!(reg.applied_watermark(), 0);
let r2 = reg.clone();
let h = std::thread::spawn(move || r2.snapshot_barrier(25));
std::thread::sleep(std::time::Duration::from_millis(20));
assert!(!h.is_finished(), "barrier(25) returned while commit 10 pending");
reg.end_commit(10);
h.join().unwrap();
assert_eq!(reg.applied_watermark(), 30);
}
#[test]
fn barrier_ignores_gaps_from_never_begun_timestamps() {
let reg = std::sync::Arc::new(WriteConflictRegistry::new());
reg.begin_commit(10);
reg.begin_commit(40); let r2 = reg.clone();
let h = std::thread::spawn(move || r2.snapshot_barrier(25));
std::thread::sleep(std::time::Duration::from_millis(20));
assert!(!h.is_finished(), "barrier(25) returned while commit 10 pending");
reg.end_commit(10);
h.join().unwrap();
assert_eq!(reg.applied_watermark(), 10);
let r3 = reg.clone();
let h2 = std::thread::spawn(move || r3.snapshot_barrier(45));
std::thread::sleep(std::time::Duration::from_millis(20));
assert!(!h2.is_finished(), "barrier(45) returned while commit 40 pending");
reg.end_commit(40);
h2.join().unwrap();
assert_eq!(reg.applied_watermark(), 40);
}
#[test]
fn end_commit_is_idempotent_and_ignores_unknown_ts() {
let reg = WriteConflictRegistry::new();
reg.begin_commit(10);
reg.begin_commit(20);
reg.end_commit(999); reg.end_commit(20);
reg.end_commit(20); assert_eq!(reg.applied_watermark(), 0);
reg.end_commit(10);
assert_eq!(reg.applied_watermark(), 20);
reg.snapshot_barrier(u64::MAX);
assert_eq!(reg.pending_count.load(Ordering::Acquire), 0);
}
#[test]
fn cross_thread_end_commit_marks_via_scan() {
let reg = std::sync::Arc::new(WriteConflictRegistry::new());
reg.begin_commit(10);
let r2 = reg.clone();
std::thread::spawn(move || r2.end_commit(10)).join().unwrap();
assert_eq!(reg.applied_watermark(), 10);
reg.snapshot_barrier(u64::MAX);
}
#[test]
fn ring_overflow_blocks_begin_until_straggler_applies() {
let reg = std::sync::Arc::new(WriteConflictRegistry::new());
reg.begin_commit(1); for i in 0..(INFLIGHT_RING_SIZE as u64 - 1) {
let ts = 2 + i;
reg.begin_commit(ts);
reg.end_commit(ts); }
let r2 = reg.clone();
let h = std::thread::spawn(move || {
let ts = 2 + INFLIGHT_RING_SIZE as u64;
r2.begin_commit(ts); r2.end_commit(ts);
});
std::thread::sleep(std::time::Duration::from_millis(30));
assert!(!h.is_finished(), "begin_commit proceeded on a full ring");
reg.end_commit(1);
h.join().unwrap();
reg.snapshot_barrier(u64::MAX);
assert_eq!(reg.pending_count.load(Ordering::Acquire), 0);
}
#[test]
fn concurrent_begin_end_barrier_stress() {
let reg = std::sync::Arc::new(WriteConflictRegistry::new());
let alloc = std::sync::Arc::new(Mutex::new(0u64));
let committers = 8;
let per_thread = 500;
std::thread::scope(|s| {
for _ in 0..committers {
let reg = std::sync::Arc::clone(®);
let alloc = std::sync::Arc::clone(&alloc);
s.spawn(move || {
for i in 0..per_thread {
let ts = {
let mut next = alloc.lock();
*next += 1;
reg.begin_commit(*next);
*next
};
if i % 3 == 0 {
std::thread::yield_now();
}
reg.end_commit(ts);
}
});
}
for _ in 0..4 {
let reg = std::sync::Arc::clone(®);
let alloc = std::sync::Arc::clone(&alloc);
s.spawn(move || {
for _ in 0..per_thread {
let snap = {
let mut next = alloc.lock();
*next += 1;
*next
};
reg.snapshot_barrier(snap);
assert!(
!reg.has_pending_at_or_below(snap),
"barrier passed a pending commit <= its snapshot"
);
}
});
}
});
reg.snapshot_barrier(u64::MAX);
assert_eq!(reg.pending_count.load(Ordering::Acquire), 0, "ledger must drain");
reg.try_reclaim();
assert_eq!(
reg.tail.load(Ordering::Acquire),
reg.head.load(Ordering::Acquire),
"ring must drain"
);
}
}