use std::{
num::NonZeroUsize,
sync::atomic::{AtomicU64, AtomicUsize, Ordering},
};
use crossbeam_queue::SegQueue;
use diskann::utils::IntoUsize;
use parking_lot::{Mutex, MutexGuard};
const DEFAULT_GUARD_SLOTS: NonZeroUsize = NonZeroUsize::new(256).unwrap();
#[derive(Debug)]
pub(crate) struct Registry {
guards: Box<[AtomicU64]>,
hint: AtomicUsize,
epoch: AtomicU64,
drain: Mutex<()>,
retiring: Box<[SegQueue<u32>; 4]>,
}
fn queue(epoch: u64) -> usize {
epoch.into_usize() % 4
}
fn last_queue(epoch: u64) -> usize {
queue(epoch.wrapping_sub(2))
}
impl Registry {
pub(crate) const fn default_guard_slots() -> NonZeroUsize {
DEFAULT_GUARD_SLOTS
}
#[cfg(test)]
pub(crate) fn new() -> Self {
Self::with_capacity(DEFAULT_GUARD_SLOTS)
}
pub(crate) fn with_capacity(capacity: NonZeroUsize) -> Self {
Self {
guards: std::iter::repeat_with(|| AtomicU64::new(0))
.take(capacity.get())
.collect(),
hint: AtomicUsize::new(0),
epoch: AtomicU64::new(1),
retiring: Box::new(core::array::from_fn(|_| SegQueue::new())),
drain: Mutex::new(()),
}
}
pub(crate) fn epoch(&self) -> u64 {
self.epoch.load(Ordering::Acquire)
}
pub(crate) fn guard(&self) -> Result<Guard<'_>, Unavailable> {
self.guard_inner(NoDelay)
}
#[inline]
fn guard_inner<T>(&self, mut delay: T) -> Result<Guard<'_>, Unavailable>
where
T: GuardDelay,
{
let mut epoch = self.epoch();
let hint = self.hint.fetch_add(1, Ordering::Relaxed);
delay.post_guard_check();
let nguards = self.guards.len();
for i in 0..nguards {
let slot = hint.wrapping_add(i) % nguards;
let guard_slot = &self.guards[slot];
delay.pre_cas();
if guard_slot.load(Ordering::Relaxed) == 0
&& guard_slot
.compare_exchange(0, epoch, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
delay.post_cas();
let mut reset = false;
loop {
delay.pre_fence();
std::sync::atomic::fence(Ordering::SeqCst);
delay.post_fence();
let current = self.epoch();
if current == epoch {
break;
}
reset = true;
epoch = current;
}
if reset {
guard_slot.store(epoch, Ordering::Relaxed);
}
return Ok(Guard {
slot: guard_slot,
retire: &self.retiring[queue(epoch)],
#[cfg(test)]
epoch,
#[cfg(test)]
slot_index: slot,
});
}
}
Err(Unavailable)
}
fn can_advance<T>(&self, delay: &mut T) -> (bool, u64)
where
T: CanAdvanceDelay,
{
delay.pre_fence();
std::sync::atomic::fence(Ordering::SeqCst);
delay.post_fence();
let current = self.epoch();
let mut min = current;
for s in self.guards.iter() {
let guarded = s.load(Ordering::Relaxed);
if guarded != 0 {
min = min.min(guarded);
}
}
std::sync::atomic::fence(Ordering::Acquire);
(min == current, min)
}
pub(crate) fn try_advance(&self) -> Option<Drain<'_>> {
self.try_advance_inner(NoDelay)
}
#[expect(
clippy::panic,
reason = "the panic is exceedingly unlikely to happen and if it does, we can't continue"
)]
fn try_advance_inner<T>(&self, mut delay: T) -> Option<Drain<'_>>
where
T: TryAdvanceDelay,
{
let drain = self.drain.try_lock()?;
let (can_advance, current) = self.can_advance(&mut delay);
if current == u64::MAX {
panic!(
"we've managed to go through nearly `u64::MAX` ids - this is unlikely in a real program"
);
}
if can_advance {
let _previous = self.epoch.fetch_add(1, Ordering::SeqCst);
debug_assert_eq!(_previous, current, "concurrency violation");
let queue = &self.retiring[last_queue(current)];
Some(Drain {
queue,
_drain: drain,
})
} else {
None
}
}
#[cfg(test)]
fn assert_no_workers(&self) {
for s in self.guards.iter() {
assert_eq!(s.load(Ordering::Relaxed), 0);
}
}
#[cfg(test)]
fn waiting(&self) -> u64 {
self.can_advance(&mut NoDelay).1
}
}
#[derive(Debug)]
pub(crate) struct Guard<'a> {
slot: &'a AtomicU64,
retire: &'a SegQueue<u32>,
#[cfg(test)]
pub(super) epoch: u64,
#[cfg(test)]
slot_index: usize,
}
impl Guard<'_> {
#[inline]
pub(crate) fn retire(&self, i: u32) {
self.retire.push(i)
}
}
impl Drop for Guard<'_> {
fn drop(&mut self) {
self.slot.store(0, Ordering::Release);
}
}
#[derive(Debug)]
pub(crate) struct Drain<'a> {
queue: &'a SegQueue<u32>,
_drain: MutexGuard<'a, ()>,
}
impl Drain<'_> {
#[must_use = "reclaimed ids must be reclaimed"]
pub(crate) fn pop(&self) -> Option<u32> {
self.queue.pop()
}
pub(crate) fn len(&self) -> usize {
self.queue.len()
}
#[cfg(test)]
fn is_empty(&self) -> bool {
self.len() == 0
}
}
impl Iterator for Drain<'_> {
type Item = u32;
fn next(&mut self) -> Option<u32> {
self.pop()
}
fn size_hint(&self) -> (usize, Option<usize>) {
(self.len(), Some(self.len()))
}
}
impl ExactSizeIterator for Drain<'_> {}
#[derive(Debug)]
#[non_exhaustive]
pub(crate) struct Unavailable;
impl std::fmt::Display for Unavailable {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("all available registry guard slots are occupied")
}
}
impl std::error::Error for Unavailable {}
diskann::convert_error!(Unavailable);
#[derive(Debug)]
struct NoDelay;
trait GuardDelay {
fn post_guard_check(&mut self) {}
fn pre_cas(&mut self) {}
fn post_cas(&mut self) {}
fn pre_fence(&mut self) {}
fn post_fence(&mut self) {}
}
impl GuardDelay for NoDelay {}
trait CanAdvanceDelay {
fn pre_fence(&mut self) {}
fn post_fence(&mut self) {}
}
impl CanAdvanceDelay for NoDelay {}
trait TryAdvanceDelay: CanAdvanceDelay {}
impl TryAdvanceDelay for NoDelay {}
#[cfg(test)]
mod tests {
use super::*;
use crate::test::Sequencer;
#[test]
fn test_cas_race() {
let seq = Sequencer::new();
let mut thread_a_loop_count = 0;
let mut thread_b_loop_count = 0;
let delay = TestGuardDelay::default()
.post_guard_check(|| seq.wait_for(0))
.with_post_fence(|| thread_a_loop_count += 1);
let registry = Registry::with_capacity(NonZeroUsize::new(2).unwrap());
std::thread::scope(|s| {
s.spawn(|| {
let g = registry.guard_inner(delay).unwrap();
assert_eq!(g.slot_index, 1);
seq.wait_for(1);
});
s.spawn(|| {
seq.until_waiting_for(0);
{
let delay =
TestGuardDelay::default().with_post_fence(|| thread_b_loop_count += 1);
let g = registry.guard_inner(delay).unwrap();
assert_eq!(g.slot_index, 1);
}
let g = registry.guard_inner(NoDelay).unwrap();
assert_eq!(g.slot_index, 0);
seq.advance_past(0);
seq.advance_past(1);
});
});
assert_eq!(thread_a_loop_count, 1);
assert_eq!(thread_b_loop_count, 1);
registry.assert_no_workers();
}
#[test]
fn test_register_wait() {
let seq = Sequencer::new();
let mut loop_count = 0;
let delay = TestGuardDelay::default()
.post_guard_check(|| seq.wait_for(0))
.with_post_cas(|| seq.wait_for(1))
.with_pre_fence(|| loop_count += 1);
let registry = Registry::with_capacity(NonZeroUsize::new(2).unwrap());
std::thread::scope(|s| {
let handle = s.spawn(|| {
let guard = registry.guard_inner(delay).unwrap();
guard.retire(10);
guard.retire(1);
guard.retire(2);
guard.retire(3);
guard
});
seq.until_waiting_for(0);
assert_eq!(registry.waiting(), 1);
{
let drain = registry.try_advance().unwrap();
assert!(drain.is_empty());
assert_eq!(registry.epoch(), 2);
}
{
let drain = registry.try_advance().unwrap();
assert!(drain.is_empty());
assert_eq!(registry.epoch(), 3);
}
seq.advance_past(0);
seq.until_waiting_for(1);
let (can_advance, waiter) = registry.can_advance(&mut NoDelay);
assert!(!can_advance);
assert_eq!(
waiter, 1,
"waiting thread registers an older generation before observing the change"
);
seq.advance_past(1);
let expected = 3;
let r = handle.join().unwrap();
assert_eq!(r.epoch, expected);
assert_eq!(registry.waiting(), expected);
});
assert_eq!(
loop_count, 2,
"the registering thread should have looped to update its generation"
);
registry.assert_no_workers();
{
let drain = registry.try_advance().unwrap();
assert!(drain.is_empty());
}
{
let drain = registry.try_advance().unwrap();
assert!(drain.is_empty());
}
{
let drain = registry.try_advance().unwrap();
let ids: Vec<_> = drain.collect();
assert_eq!(ids, &[10, 1, 2, 3]);
}
}
#[test]
fn test_slot_exhaustion() {
let registry = Registry::with_capacity(NonZeroUsize::new(2).unwrap());
let g0 = registry.guard().unwrap();
let g1 = registry.guard().unwrap();
assert!(matches!(registry.guard(), Err(Unavailable)));
assert!(matches!(registry.guard(), Err(Unavailable)));
let freed_slot = g0.slot_index;
drop(g0);
let g2 = registry.guard().unwrap();
assert_eq!(
g2.slot_index, freed_slot,
"newly freed slot should be reclaimed"
);
assert!(matches!(registry.guard(), Err(Unavailable)));
drop(g1);
drop(g2);
registry.assert_no_workers();
}
#[test]
fn test_slot_wrap_around() {
let registry = Registry::with_capacity(NonZeroUsize::new(4).unwrap());
let (g2, g3) = {
let _g0 = registry.guard().unwrap();
let _g1 = registry.guard().unwrap();
let g2 = registry.guard().unwrap();
let g3 = registry.guard().unwrap();
(g2, g3)
};
assert_eq!(g2.slot_index, 2);
assert_eq!(g3.slot_index, 3);
let f = || {
for _ in 0..10 {
let g0 = registry.guard().unwrap();
let g1 = registry.guard().unwrap();
let s0 = g0.slot_index;
let s1 = g1.slot_index;
if s0 < s1 {
assert_eq!((s0, s1), (0, 1));
} else {
assert_eq!((s0, s1), (1, 0));
};
assert!(matches!(registry.guard(), Err(Unavailable)));
}
};
f();
registry.hint.store(usize::MAX - 10, Ordering::Relaxed);
f();
drop((g2, g3));
registry.assert_no_workers();
}
#[test]
fn test_concurrent_try_advance() {
let registry = Registry::with_capacity(NonZeroUsize::new(2).unwrap());
let drain = registry
.try_advance()
.expect("first try_advance must succeed");
let gen_after_first = registry.epoch();
assert_eq!(gen_after_first, 2);
std::thread::scope(|s| {
s.spawn(|| {
assert!(
registry.try_advance().is_none(),
"try_advance must fail while another holds the drain mutex"
);
assert_eq!(
registry.epoch(),
gen_after_first,
"generation must not advance when drain is contended"
);
});
});
drop(drain);
let _drain2 = registry
.try_advance()
.expect("try_advance must succeed once drain is released");
assert_eq!(registry.epoch(), 3);
}
#[test]
fn test_drain_rotation() {
let registry = Registry::with_capacity(NonZeroUsize::new(1).unwrap());
let retire_at = |id: u32| {
let g = registry.guard().unwrap();
let epoch = g.epoch;
g.retire(id);
epoch
};
let gen_a = retire_at(100);
assert_eq!(gen_a, 1);
{
let drain = registry.try_advance().unwrap();
assert!(
drain.is_empty(),
"100 must not drain on 1st advance after A"
);
}
let gen_b = retire_at(200);
assert_eq!(gen_b, gen_a + 1);
{
let drain = registry.try_advance().unwrap();
assert!(
drain.is_empty(),
"100 must not drain on 2nd advance after A"
);
}
let _gen_c = retire_at(300);
{
let drained: Vec<_> = registry.try_advance().unwrap().collect();
assert_eq!(drained, &[100]);
}
{
let drained: Vec<_> = registry.try_advance().unwrap().collect();
assert_eq!(drained, &[200]);
}
{
let drained: Vec<_> = registry.try_advance().unwrap().collect();
assert_eq!(drained, &[300]);
}
{
let drain = registry.try_advance().unwrap();
assert!(
drain.is_empty(),
"rotation should leave queues empty after one cycle"
);
}
registry.assert_no_workers();
}
macro_rules! tester {
($struct:ident, $trait:ident, $($with:ident => $f:ident),* $(,)?) => {
#[derive(Default)]
struct $struct<'a> {
$($f: Option<Box<dyn FnMut() + Send + 'a>>,)*
}
impl<'a> $struct<'a> {
$(
fn $with<F>(mut self, f: F) -> Self
where
F: FnMut() + Send + 'a
{
self.$f = Some(Box::new(f));
self
}
)*
}
impl $trait for $struct<'_> {
$(
fn $f(&mut self) {
if let Some(f) = self.$f.as_mut() {
f()
}
}
)*
}
}
}
tester! {
TestGuardDelay,
GuardDelay,
post_guard_check => post_guard_check,
with_post_cas => post_cas,
with_pre_fence => pre_fence,
with_post_fence => post_fence,
}
}