use std::sync::atomic::{AtomicU64, Ordering};
use dashmap::DashMap;
use crate::branch::ShardId;
use crate::sync::ballot::{Ballot, Stamp};
#[derive(Debug)]
struct ShardOwnerStamp {
live_epoch: Ballot,
seq: AtomicU64,
}
impl ShardOwnerStamp {
fn bottom() -> Self {
Self {
live_epoch: Ballot::bottom(),
seq: AtomicU64::new(0),
}
}
}
#[derive(Debug, Default)]
pub struct OwnerStamps {
shards: DashMap<ShardId, ShardOwnerStamp>,
}
impl OwnerStamps {
pub fn record_won(&self, shard: ShardId, won_ballot: Ballot) {
self.shards.insert(
shard,
ShardOwnerStamp {
live_epoch: won_ballot,
seq: AtomicU64::new(0),
},
);
}
pub fn next_stamp(&self, shard: ShardId) -> Stamp {
let entry = self
.shards
.entry(shard)
.or_insert_with(ShardOwnerStamp::bottom);
let seq = entry.seq.fetch_add(1, Ordering::Relaxed);
Stamp::new(entry.live_epoch.clone(), seq)
}
pub fn live_epoch(&self, shard: ShardId) -> Ballot {
self.shards
.get(&shard)
.map_or_else(Ballot::bottom, |entry| entry.live_epoch.clone())
}
#[doc(hidden)]
pub fn peek_stamp(&self, shard: ShardId) -> Stamp {
self.shards.get(&shard).map_or_else(Stamp::bottom, |entry| {
Stamp::new(entry.live_epoch.clone(), entry.seq.load(Ordering::Relaxed))
})
}
}
impl super::Database {
#[must_use]
pub fn current_owner_epoch(&self, shard_id: usize) -> crate::sync::Ballot {
self.owner_stamps.live_epoch(shard_id)
}
#[must_use]
pub fn is_current_owner(&self, shard_id: usize) -> bool {
self.current_owner_epoch(shard_id) != crate::sync::Ballot::bottom()
}
pub(crate) fn next_stamp_for_key(&self, key: &[u8]) -> crate::sync::Stamp {
self.owner_stamps.next_stamp(self.shard_for(key))
}
}
#[cfg(test)]
mod tests {
use std::collections::HashSet;
use std::sync::Arc;
use crate::sync::ballot::Ballot;
use crate::sync::topology::SyncNodeId;
use super::OwnerStamps;
fn ballot(counter: u64, node: &str) -> Ballot {
Ballot::new(counter, SyncNodeId::new(node))
}
#[test]
fn default_stamp_is_bottom_epoch_with_monotonic_seq() {
let stamps = OwnerStamps::default();
assert_eq!(stamps.live_epoch(0), Ballot::bottom());
let first = stamps.next_stamp(0);
let second = stamps.next_stamp(0);
assert_eq!(first.epoch, Ballot::bottom());
assert_eq!(first.seq, 0);
assert_eq!(second.seq, 1);
}
#[test]
fn record_won_sets_live_epoch_and_resets_seq() {
let stamps = OwnerStamps::default();
assert_eq!(stamps.next_stamp(0).seq, 0);
assert_eq!(stamps.next_stamp(0).seq, 1);
stamps.record_won(0, ballot(5, "A"));
assert_eq!(stamps.live_epoch(0), ballot(5, "A"));
let s0 = stamps.next_stamp(0);
let s1 = stamps.next_stamp(0);
assert_eq!(s0.epoch, ballot(5, "A"));
assert_eq!(s0.seq, 0, "seq resets to 0 on a new live epoch");
assert_eq!(s1.seq, 1);
stamps.record_won(0, ballot(6, "A"));
let s = stamps.next_stamp(0);
assert_eq!(s.epoch, ballot(6, "A"));
assert_eq!(s.seq, 0);
}
#[test]
fn concurrent_draws_get_distinct_seq_no_toctou() {
const THREADS: usize = 8;
const PER_THREAD: u64 = 1000;
let stamps = Arc::new(OwnerStamps::default());
stamps.record_won(3, ballot(1, "owner"));
let mut handles = Vec::new();
for _ in 0..THREADS {
let stamps = Arc::clone(&stamps);
handles.push(std::thread::spawn(move || {
let mut seqs = Vec::with_capacity(PER_THREAD as usize);
for _ in 0..PER_THREAD {
let drawn = stamps.next_stamp(3);
assert_eq!(drawn.epoch, ballot(1, "owner"));
seqs.push(drawn.seq);
}
seqs
}));
}
let mut all = HashSet::new();
for handle in handles {
let seqs = handle.join().unwrap_or_default();
assert!(
!seqs.is_empty(),
"worker thread produced no seqs (panicked?)"
);
for seq in seqs {
assert!(
all.insert(seq),
"duplicate seq {seq} — TOCTOU / non-atomic draw"
);
}
}
let total = THREADS as u64 * PER_THREAD;
assert_eq!(all.len() as u64, total);
for expected in 0..total {
assert!(all.contains(&expected), "missing seq {expected}");
}
}
}