use std::sync::atomic::Ordering;
use diskann::utils::IntoUsize;
use thiserror::Error;
use crate::{
buffer::{Buffer, BufferError, RawSlice},
epoch,
num::{Align, Bytes, IdLimit},
store::{Lifecycle, slots},
tag::{self, AtomicTag, Tag},
};
#[derive(Debug, Clone)]
pub(crate) struct Config {
bytes: Bytes,
}
impl Config {
pub(crate) fn new(bytes: Bytes) -> Self {
Self { bytes }
}
}
impl slots::SlotsConfig for Config {
type Slots = Intrusive;
type Error = IntrusiveError;
unsafe fn build(
self,
handle: epoch::RegistryHandle,
tags: &tag::Authoritative,
) -> Result<Intrusive, IntrusiveError> {
let Self { bytes } = self;
unsafe { Intrusive::new(bytes, handle, tags.id_limit()) }
}
}
#[derive(Debug)]
pub(crate) struct Intrusive {
buffer: Buffer,
unpadded: Bytes,
handle: epoch::RegistryHandle,
}
impl Intrusive {
pub(crate) fn config(bytes: Bytes) -> Config {
Config::new(bytes)
}
pub(crate) unsafe fn new(
bytes: Bytes,
handle: epoch::RegistryHandle,
id_limit: IdLimit,
) -> Result<Self, IntrusiveError> {
let Some(unpadded) = bytes.checked_add(AtomicTag::SIZE) else {
return Err(IntrusiveError::bytes_overflowed());
};
let Some(padded_bytes) = unpadded.checked_next_multiple_of(Bytes::CACHELINE) else {
return Err(IntrusiveError::bytes_overflowed());
};
let buffer = match Buffer::new(id_limit.as_usize(), padded_bytes, Align::_128) {
Ok(buffer) => buffer,
Err(err) => return Err(IntrusiveError::buffer_error(err)),
};
Ok(Self {
buffer,
unpadded,
handle,
})
}
pub(crate) fn id_limit(&self) -> IdLimit {
IdLimit::new(self.buffer.len() as u32)
}
pub(crate) fn bytes(&self) -> Bytes {
self.bytes_plus_tag().unchecked_sub(AtomicTag::SIZE)
}
pub(crate) fn bytes_plus_tag(&self) -> Bytes {
self.unpadded
}
pub(crate) fn reader<'a>(&'a self, guard: epoch::Guard<'a>) -> Reader<'a> {
self.handle.assert_guard_belongs(&guard);
Reader {
buffer: &self.buffer,
unpadded: self.unpadded,
guard,
}
}
unsafe fn data_unchecked(&self, i: usize) -> (&AtomicTag, RawSlice<'_>) {
let (data, mirror) = unsafe { self.buffer.get_unchecked(i) }
.truncate(self.unpadded)
.split(self.unpadded.unchecked_sub(AtomicTag::SIZE));
(
unsafe { AtomicTag::from_ptr(mirror.as_mut_ptr().cast()) },
data,
)
}
fn data(&self, i: usize) -> Option<(&AtomicTag, RawSlice<'_>)> {
if i >= self.buffer.len() {
None
} else {
Some(unsafe { self.data_unchecked(i) })
}
}
}
#[derive(Debug, Error)]
#[error(transparent)]
pub(crate) struct IntrusiveError(IntrusiveErrorInner);
impl IntrusiveError {
fn bytes_overflowed() -> Self {
Self(IntrusiveErrorInner::BytesOverflowed)
}
fn buffer_error(err: BufferError) -> Self {
Self(IntrusiveErrorInner::BufferError(err))
}
}
#[derive(Debug, Error)]
enum IntrusiveErrorInner {
#[error("computation of the bytes per slot overflowed")]
BytesOverflowed,
#[error(transparent)]
BufferError(BufferError),
}
impl slots::Slots for Intrusive {
type Exclusive<'a> = Exclusive<'a>;
fn id_limit(&self) -> IdLimit {
<Intrusive>::id_limit(self)
}
#[expect(clippy::panic, reason = "out-of-bounds is a hard program bug")]
unsafe fn acquire(&self, i: u32, _: Lifecycle) -> Self::Exclusive<'_> {
let Some((tag, data)) = self.data(i.into_usize()) else {
panic!("index {i} is out-of-bounds");
};
debug_assert_eq!(
tag.load(Ordering::Relaxed),
Tag::AVAILABLE,
"concurrency violation",
);
tag.store(Tag::OWNED, Ordering::Relaxed);
Exclusive { tag, data }
}
#[expect(clippy::panic, reason = "out-of-bounds is a hard program bug")]
unsafe fn reclaim(&self, i: u32, _: Lifecycle) {
let Some((tag, _)) = self.data(i.into_usize()) else {
panic!("index {i} is out-of-bounds");
};
tag.store(Tag::AVAILABLE, Ordering::Release);
}
#[expect(clippy::panic, reason = "out-of-bounds is a hard program bug")]
unsafe fn retire(&self, i: u32, _: Lifecycle) {
let Some((tag, _)) = self.data(i.into_usize()) else {
panic!("index {i} is out-of-bounds");
};
tag.store(Tag::RETIRING, Ordering::Relaxed);
}
}
#[derive(Debug)]
pub(crate) struct Reader<'a> {
buffer: &'a Buffer,
unpadded: Bytes,
#[cfg_attr(
not(feature = "quantization"),
expect(unused, reason = "quantization uses this to share the guard")
)]
guard: epoch::Guard<'a>,
}
impl<'a> Reader<'a> {
#[inline]
pub(crate) fn read(&self, i: usize) -> Option<&[u8]> {
if self.is_in_bounds(i) {
unsafe { self.read_in_bounds(i) }
} else {
None
}
}
#[inline]
#[must_use = "this function has no side-effects"]
pub(crate) fn is_in_bounds(&self, i: usize) -> bool {
i < self.buffer.len()
}
#[inline]
#[must_use = "this function has no side-effects"]
pub(crate) fn id_limit(&self) -> IdLimit {
IdLimit::new(self.buffer.len() as u32)
}
#[cfg_attr(
not(test),
expect(
dead_code,
reason = "this is non-trivial method that is likely to be used in the future"
)
)]
pub(crate) fn can_read(&self, i: usize) -> Option<bool> {
if !self.is_in_bounds(i) {
return None;
}
let tag_ptr = unsafe {
self.buffer
.get_unchecked(i)
.as_mut_ptr()
.add(self.unpadded.unchecked_sub(AtomicTag::SIZE).value())
};
let can_read = unsafe { AtomicTag::from_ptr(tag_ptr.cast()) }
.load(Ordering::Acquire)
.can_read();
Some(can_read)
}
#[inline]
pub(crate) unsafe fn read_in_bounds(&self, i: usize) -> Option<&[u8]> {
debug_assert!(self.is_in_bounds(i));
let (data, tag_ptr) = unsafe {
self.buffer
.get_unchecked(i)
.truncate_unchecked(self.unpadded)
.split_unchecked(self.unpadded.unchecked_sub(AtomicTag::SIZE))
};
let can_read = unsafe { AtomicTag::from_ptr(tag_ptr.as_mut_ptr().cast()) }
.load(Ordering::Acquire)
.can_read();
if can_read {
Some(unsafe { data.as_slice() })
} else {
None
}
}
#[inline]
pub(crate) unsafe fn read_raw_unchecked(&self, i: usize) -> RawSlice<'_> {
unsafe { self.buffer.get_unchecked(i) }.truncate(self.unpadded)
}
pub(crate) fn bytes(&self) -> Bytes {
self.bytes_plus_tag().unchecked_sub(AtomicTag::SIZE)
}
pub(crate) fn bytes_plus_tag(&self) -> Bytes {
self.unpadded
}
#[cfg(feature = "quantization")]
pub(crate) fn guard(&self) -> &epoch::Guard<'a> {
&self.guard
}
}
#[derive(Debug)]
pub(crate) struct Exclusive<'a> {
tag: &'a AtomicTag,
data: RawSlice<'a>,
}
impl<'a> Exclusive<'a> {
pub(crate) fn as_mut_slice(&mut self) -> &mut [u8] {
unsafe { self.data.as_mut_slice() }
}
}
impl slots::Exclusive for Exclusive<'_> {
fn publish(self, _: Lifecycle) {
self.tag.store(Tag::PUBLISHED, Ordering::Release);
}
fn freeze(self, _: Lifecycle) {
self.tag.store(Tag::FROZEN, Ordering::Release);
}
fn abort(self, _: Lifecycle) {
self.tag.store(Tag::AVAILABLE, Ordering::Release);
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::{
assert_matches,
num::{NonZeroU32, NonZeroUsize},
};
use crate::{
num::{Capacity, MaxDegree},
store::{self, Store},
};
fn store(
entries: usize,
entry_bytes: usize,
frozen: usize,
) -> Result<Store<Intrusive>, store::StoreError> {
let store = Store::new(
store::Layout::new(
Capacity::new(entries),
MaxDegree::new(0),
frozen.try_into().unwrap(),
),
store::Config::__exhaustive(
NonZeroUsize::new(10).unwrap(),
NonZeroU32::new(16).unwrap(),
),
Config::new(Bytes::new(entry_bytes)),
)?;
for (base, id) in store.frozen().enumerate() {
let mut slot = store.slot(id).unwrap();
slot.data().as_mut_slice().fill(base as u8);
slot.freeze();
}
Ok(store)
}
#[test]
fn frozen_range_follows_writable_slots() {
let s = store(4, 8, 2).unwrap();
assert_eq!(s.frozen(), 4..6);
let reader = s.guard(|intrusive, guard| intrusive.reader(guard)).unwrap();
for i in 0..4 {
assert!(!s.can_read_approximate(i).unwrap());
assert!(!reader.can_read(i).unwrap());
assert!(reader.read(i).is_none());
}
assert!(s.can_read_approximate(4).unwrap());
assert!(reader.can_read(4).unwrap());
assert_eq!(reader.read(4).unwrap(), &[0, 0, 0, 0, 0, 0, 0, 0]);
assert!(s.can_read_approximate(5).unwrap());
assert!(reader.can_read(5).unwrap());
assert_eq!(reader.read(5).unwrap(), &[1, 1, 1, 1, 1, 1, 1, 1]);
assert!(s.can_read_approximate(6).is_none());
assert!(reader.can_read(6).is_none());
assert!(reader.read(6).is_none());
}
#[test]
fn acquire_write_publish_read_roundtrip() {
let s = store(4, 8, 1).unwrap();
let reader = s
.guard(|intrusive, guard| intrusive.reader(guard))
.expect("reader guard available");
let idx = {
let mut slot = s.acquire().expect("a fresh store has free slots");
let idx = slot.slot() as usize;
slot.data()
.as_mut_slice()
.copy_from_slice(&[1, 2, 3, 4, 5, 6, 7, 8]);
assert!(reader.read(idx).is_none());
assert!(!s.can_read_approximate(idx).unwrap());
slot.publish();
idx
};
assert_eq!(reader.read(idx), Some([1, 2, 3, 4, 5, 6, 7, 8].as_slice()));
assert!(s.can_read_approximate(idx).unwrap());
}
#[test]
fn unpublished_slots_are_immediately_available() {
let s = store(4, 8, 1).unwrap();
let reader = s
.guard(|intrusive, guard| intrusive.reader(guard))
.expect("reader guard available");
let idx = {
let mut slot = s.acquire().expect("a fresh store has free slots");
let idx = slot.slot() as usize;
slot.data()
.as_mut_slice()
.copy_from_slice(&[1, 2, 3, 4, 5, 6, 7, 8]);
assert!(reader.read(idx).is_none());
assert!(!s.can_read_approximate(idx).unwrap());
idx
};
assert!(reader.read(idx).is_none());
assert!(!s.can_read_approximate(idx).unwrap());
}
#[test]
fn acquire_exhausts_then_reports_none() {
let s = store(2, 8, 1).unwrap();
let _a = s.acquire().expect("first writable slot");
let _b = s.acquire().expect("second writable slot");
assert!(
s.acquire().is_none(),
"all writable slots are owned, so acquire must fail"
);
}
#[test]
fn retire_out_of_bounds() {
let s = store(4, 8, 1).unwrap();
assert_matches!(s.retire(999), Err(store::RetireError::OutOfBounds));
}
#[test]
fn retire_rejects_reserved_slots() {
let s = store(4, 8, 1).unwrap();
assert_matches!(s.retire(0), Err(store::RetireError::SlotIsReserved { .. }));
let frozen = s.frozen().start as usize;
assert_matches!(
s.retire(frozen),
Err(store::RetireError::SlotIsReserved { .. })
);
let slot = s.acquire().unwrap();
assert_matches!(
s.retire(slot.slot() as usize),
Err(store::RetireError::SlotIsReserved { .. })
);
}
#[test]
fn retire_published_slot_then_unreadable() {
let s = store(4, 8, 1).unwrap();
let idx = {
let slot = s.acquire().unwrap();
slot.publish() as usize
};
assert!(s.retire(idx).is_ok());
let reader = s
.guard(|intrusive, guard| intrusive.reader(guard))
.expect("reader guard available");
assert_eq!(reader.read(idx), None);
assert_eq!(reader.can_read(idx), Some(false));
assert_matches!(
s.retire(idx),
Err(store::RetireError::SlotIsReserved { .. })
);
}
#[test]
fn test_recycling() {
let entries = if cfg!(miri) { 16 } else { 2048 };
let s = store(entries, 4, 2).unwrap();
let mut count = 0;
while let Some(slot) = s.acquire() {
slot.publish();
count += 1;
}
assert_eq!(count, s.writable().len());
for i in s.writable() {
s.retire(i.into_usize()).unwrap();
}
let mut count = 0;
while let Some(slot) = s.acquire() {
slot.publish();
count += 1;
}
assert_eq!(count, s.writable().len());
}
}