use std::ptr::NonNull;
use diskann::{graph::AdjacencyList, utils::IntoUsize};
use parking_lot::{RwLock, RwLockWriteGuard};
use thiserror::Error;
use crate::{
buffer::{Buffer, BufferError},
num::{Align, Bytes},
};
type Id = u32;
const LOCK_GRANULARITY: usize = 16;
fn lock_index(i: u32) -> usize {
i.into_usize() / LOCK_GRANULARITY
}
#[derive(Debug)]
pub(crate) struct Neighbors {
neighbors: Buffer,
locks: Vec<RwLock<()>>,
}
impl Neighbors {
pub(crate) fn new(entries: u32, max_length: u32) -> Result<Self, NeighborsError> {
let bytes = max_length
.into_usize()
.checked_add(1)
.and_then(|len| len.checked_mul(std::mem::size_of::<Id>()))
.map(Bytes::new)
.ok_or(NeighborsError::Overflow(max_length))?;
const ALIGN: Align = Align::_128;
const {
assert!(
ALIGN.value() >= Align::of::<Id>().value(),
"buffer alignment must be at least that of the ID"
);
}
let neighbors = Buffer::new(entries.into_usize(), bytes, ALIGN)?;
let locks = std::iter::repeat_with(|| RwLock::new(()))
.take(entries.into_usize().div_ceil(LOCK_GRANULARITY))
.collect();
Ok(Self { neighbors, locks })
}
pub(crate) fn max_length(&self) -> usize {
(self.neighbors.stride().value() - std::mem::size_of::<Id>()) / std::mem::size_of::<Id>()
}
pub(crate) fn max_length_u32(&self) -> u32 {
self.max_length() as u32
}
pub(crate) fn entries(&self) -> u32 {
self.neighbors.len() as u32
}
pub(crate) fn get(
&self,
i: u32,
neighbors: &mut AdjacencyList<u32>,
) -> Result<(), OutOfBounds> {
self.check(i)?;
let lock = unsafe { self.locks.get_unchecked(lock_index(i)) };
let _guard = lock.read();
let (prefix, rest) =
unsafe { self.neighbors.get_unchecked(i.into_usize()) }.split(Bytes::size_of::<Id>());
debug_assert_eq!(prefix.len(), Bytes::size_of::<Id>());
debug_assert!(prefix.as_ptr().cast::<Id>().is_aligned());
let len: usize = unsafe { prefix.as_ptr().cast::<Id>().read() }
.min(self.max_length_u32())
.into_usize();
let mut resizer = neighbors.resize(len);
unsafe {
std::ptr::copy_nonoverlapping(
rest.as_mut_ptr(),
resizer.as_mut_ptr().cast::<u8>(),
len * std::mem::size_of::<Id>(),
)
};
resizer.finish(len);
Ok(())
}
pub(crate) fn lock(&self, i: u32) -> Result<Lock<'_>, OutOfBounds> {
self.check(i)?;
Ok(unsafe { self.lock_unchecked(i) })
}
unsafe fn lock_unchecked(&self, i: u32) -> Lock<'_> {
let lock = unsafe { self.locks.get_unchecked(lock_index(i)) }.write();
let slice = unsafe { self.neighbors.get_unchecked(i.into_usize()) };
debug_assert!(slice.as_ptr().cast::<Id>().is_aligned());
Lock {
ptr: slice.as_non_null().cast::<Id>(),
capacity: self.max_length().into_usize(),
_lock: lock,
}
}
pub(crate) fn set(&self, i: u32, neighbors: &[u32]) -> Result<(), SetError> {
self.check(i).map_err(SetError::OutOfBounds)?;
if neighbors.len() > self.max_length().into_usize() {
return Err(SetError::TooLong(TooLong {
got: neighbors.len(),
max: self.max_length_u32(),
}));
}
let lock = unsafe { self.lock_unchecked(i) };
unsafe { lock.write_unchecked(neighbors) };
Ok(())
}
fn check(&self, i: u32) -> Result<(), OutOfBounds> {
if i >= self.entries() {
Err(OutOfBounds(i))
} else {
Ok(())
}
}
}
#[derive(Debug, Error)]
pub(crate) enum NeighborsError {
#[error("adjacency list length of {0} is too long")]
Overflow(u32),
#[error("neighbor buffer allocation failed")]
AllocationFailed(#[from] BufferError),
}
#[derive(Debug, Clone, Copy, Error)]
#[error("index {} is out-of-bounds", self.0)]
pub(crate) struct OutOfBounds(u32);
diskann::convert_error!(OutOfBounds);
#[derive(Debug, Clone, Copy, Error)]
#[error("length {} exceeds the max length {}", self.got, self.max)]
pub(crate) struct TooLong {
got: usize,
max: u32,
}
diskann::convert_error!(TooLong);
#[derive(Debug, Clone, Copy, Error)]
pub(crate) enum SetError {
#[error(transparent)]
OutOfBounds(OutOfBounds),
#[error(transparent)]
TooLong(TooLong),
}
diskann::convert_error!(SetError);
pub(crate) struct Lock<'a> {
ptr: NonNull<u32>,
capacity: usize,
_lock: RwLockWriteGuard<'a, ()>,
}
impl Lock<'_> {
pub(crate) fn capacity(&self) -> usize {
self.capacity
}
pub(crate) fn len(&self) -> usize {
unsafe { self.ptr.read() }.into_usize().min(self.capacity())
}
pub(crate) fn append(self, neighbors: &[u32]) -> Result<(), TooLong> {
let len = self.len();
let newlen = len.saturating_add(neighbors.len());
if newlen > self.capacity() {
return Err(TooLong {
got: newlen,
max: self.capacity as u32,
});
}
unsafe {
std::ptr::copy_nonoverlapping(
neighbors.as_ptr(),
self.ptr.add(len + 1).as_ptr(),
neighbors.len(),
)
}
unsafe { self.ptr.write(newlen as u32) };
Ok(())
}
unsafe fn write_unchecked(self, neighbors: &[u32]) {
let len = neighbors.len();
debug_assert!(len <= self.capacity());
unsafe { std::ptr::copy_nonoverlapping(neighbors.as_ptr(), self.ptr.as_ptr().add(1), len) }
unsafe { self.ptr.write(len as u32) };
}
#[cfg(test)]
fn as_slice(&self) -> &[u32] {
let len = self.len();
unsafe { std::slice::from_raw_parts(self.ptr.add(1).as_ptr().cast_const(), len) }
}
#[cfg(test)]
fn write(self, neighbors: &[u32]) -> Result<(), TooLong> {
if neighbors.len() > self.capacity() {
return Err(TooLong {
got: neighbors.len(),
max: self.capacity as u32,
});
}
unsafe { self.write_unchecked(neighbors) };
Ok(())
}
#[cfg(test)]
fn is_empty(&self) -> bool {
self.len() == 0
}
}
impl std::fmt::Debug for Lock<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Lock")
.field("ptr", &self.ptr)
.field("capacity", &self.capacity)
.field("lock", &())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test::Sequencer;
#[test]
fn out_of_bounds_rejects_indices_beyond_entries() {
let n = Neighbors::new(4, 4).unwrap();
let mut out = AdjacencyList::with_capacity(4);
for bad in [4u32, 5, 100, u32::MAX] {
assert!(matches!(n.get(bad, &mut out), Err(OutOfBounds(_))));
assert!(matches!(n.set(bad, &[]), Err(SetError::OutOfBounds(_))));
assert!(matches!(n.lock(bad), Err(OutOfBounds(_))));
}
}
#[test]
fn empty_neighbors_rejects_all_access() {
let n = Neighbors::new(0, 4).unwrap();
let mut out = AdjacencyList::with_capacity(4);
for i in [0u32, 1, u32::MAX] {
assert!(matches!(n.get(i, &mut out), Err(OutOfBounds(_))));
assert!(matches!(n.set(i, &[]), Err(SetError::OutOfBounds(_))));
assert!(matches!(n.lock(i), Err(OutOfBounds(_))));
}
}
#[test]
fn set_rejects_oversized_neighbors() {
let n = Neighbors::new(4, 3).unwrap();
let too_many = &[1, 2, 3, 4];
assert!(matches!(n.set(0, too_many), Err(SetError::TooLong(_))));
}
#[test]
fn lock_write_rejects_oversized_neighbors() {
let n = Neighbors::new(4, 3).unwrap();
let lock = n.lock(0).unwrap();
assert!(lock.write(&[1, 2, 3, 4]).is_err());
}
#[test]
fn lock_append_rejects_overflow() {
let n = Neighbors::new(4, 3).unwrap();
n.set(0, &[1, 2]).unwrap();
let lock = n.lock(0).unwrap();
assert!(lock.append(&[3, 4]).is_err());
}
#[test]
fn lock_implements_debug() {
let n = Neighbors::new(4, 3).unwrap();
let lock = n.lock(0).unwrap();
let _ = format!("{:?}", lock);
}
#[test]
fn append_preserves_existing_and_adds_new() {
let n = Neighbors::new(4, 6).unwrap();
n.set(0, &[10, 20]).unwrap();
let lock = n.lock(0).unwrap();
assert_eq!(lock.as_slice(), &[10, 20]);
lock.append(&[30, 40, 50]).unwrap();
let mut out = AdjacencyList::with_capacity(6);
n.get(0, &mut out).unwrap();
assert_eq!(&*out, &[10, 20, 30, 40, 50]);
}
#[test]
fn append_to_empty() {
let n = Neighbors::new(4, 4).unwrap();
let lock = n.lock(0).unwrap();
assert_eq!(lock.as_slice(), &[]);
lock.append(&[1, 2, 3]).unwrap();
let mut out = AdjacencyList::with_capacity(4);
n.get(0, &mut out).unwrap();
assert_eq!(&*out, &[1, 2, 3]);
}
#[test]
fn append_fills_to_capacity() {
let n = Neighbors::new(1, 3).unwrap();
n.set(0, &[1]).unwrap();
let lock = n.lock(0).unwrap();
lock.append(&[2, 3]).unwrap();
let mut out = AdjacencyList::with_capacity(3);
n.get(0, &mut out).unwrap();
assert_eq!(&*out, &[1, 2, 3]);
}
#[test]
fn append_empty_slice_is_noop() {
let n = Neighbors::new(1, 4).unwrap();
n.set(0, &[10, 20]).unwrap();
let lock = n.lock(0).unwrap();
lock.append(&[]).unwrap();
let mut out = AdjacencyList::with_capacity(4);
n.get(0, &mut out).unwrap();
assert_eq!(&*out, &[10, 20]);
}
#[test]
fn write_overwrites_longer_list() {
let n = Neighbors::new(1, 5).unwrap();
n.set(0, &[1, 2, 3, 4, 5]).unwrap();
let lock = n.lock(0).unwrap();
assert_eq!(lock.len(), 5);
lock.write(&[99]).unwrap();
let mut out = AdjacencyList::with_capacity(5);
n.get(0, &mut out).unwrap();
assert_eq!(&*out, &[99]);
}
fn clear(neighbors: &mut Neighbors) {
for i in 0..neighbors.entries() {
neighbors.set(i, &[]).unwrap();
}
assert_is_cleared(neighbors);
}
fn assert_is_cleared(neighbors: &mut Neighbors) {
for i in 0..neighbors.entries() {
assert!(neighbors.lock(i).unwrap().is_empty());
}
}
#[test]
fn basic_test() {
let mut neighbors = Neighbors::new(10, 4).unwrap();
assert_eq!(neighbors.entries(), 10);
assert_eq!(neighbors.max_length(), 4);
let mut list = AdjacencyList::new();
for i in 0..neighbors.entries() {
list.clear();
list.extend_from_slice(&[1, 2, 3, 4]);
neighbors.get(i, &mut list).unwrap();
assert!(list.is_empty());
let lock = neighbors.lock(i).unwrap();
assert_eq!(lock.capacity(), neighbors.max_length());
assert_eq!(lock.len(), 0);
assert!(lock.is_empty());
assert_eq!(lock.as_slice(), &[]);
}
let oob = neighbors.entries();
assert!(matches!(neighbors.get(oob, &mut list), Err(OutOfBounds(_))));
assert!(matches!(neighbors.lock(oob), Err(OutOfBounds(_))));
assert!(matches!(
neighbors.set(oob, &[1, 2, 3, 4, 5, 6]),
Err(SetError::OutOfBounds(_))
));
let generate =
|round: u32, entry: u32| -> Vec<u32> { (0..(round + 1)).map(|r| entry + r).collect() };
for round in 0..neighbors.max_length_u32() {
for i in 0..neighbors.entries() {
let v = generate(round, i);
neighbors.set(i, &v).unwrap();
}
for i in 0..neighbors.entries() {
let expected = generate(round, i);
neighbors.get(i, &mut list).unwrap();
assert_eq!(&*list, &*expected);
let lock = neighbors.lock(i).unwrap();
assert_eq!(lock.as_slice(), &*expected);
}
}
clear(&mut neighbors);
for round in 0..neighbors.max_length_u32() {
for i in 0..neighbors.entries() {
let v = generate(round, i);
neighbors.lock(i).unwrap().write(&v).unwrap();
}
for i in 0..neighbors.entries() {
let expected = generate(round, i);
neighbors.get(i, &mut list).unwrap();
assert_eq!(&*list, &*expected);
let lock = neighbors.lock(i).unwrap();
assert_eq!(lock.as_slice(), &*expected);
}
}
clear(&mut neighbors);
for round in 0..neighbors.max_length_u32() {
for i in 0..neighbors.entries() {
neighbors.lock(i).unwrap().append(&[round + i]).unwrap();
}
for i in 0..neighbors.entries() {
let expected = generate(round, i);
neighbors.get(i, &mut list).unwrap();
assert_eq!(&*list, &*expected);
let lock = neighbors.lock(i).unwrap();
assert_eq!(lock.as_slice(), &*expected);
}
}
clear(&mut neighbors);
}
#[test]
fn lock_blocks_get() {
for _ in 0..10 {
let neighbors = Neighbors::new(3, 4).unwrap();
let seq = Sequencer::new();
std::thread::scope(|s| {
let handle = s.spawn(|| {
seq.wait_for(0);
let mut list = AdjacencyList::new();
neighbors.get(0, &mut list).unwrap();
list
});
seq.until_waiting_for(0);
let lock = neighbors.lock(0).unwrap();
seq.advance_past(0);
lock.write(&[1, 2, 3, 4]).unwrap();
let list = handle.join().unwrap();
assert_eq!(&*list, &[1, 2, 3, 4]);
});
}
}
#[test]
fn many_appends() {
let max_length = if cfg!(miri) { 100 } else { 1000 };
let neighbors = Neighbors::new(1, max_length).unwrap();
let num_threads = 4;
let barrier = std::sync::Barrier::new(num_threads);
std::thread::scope(|s| {
let neighbors_ref = &neighbors;
let barrier_ref = &barrier;
for thread_id in 0..num_threads {
s.spawn(move || {
barrier_ref.wait();
let mut i = thread_id as u32;
let upper = neighbors_ref.max_length() as u32;
while i < upper {
neighbors_ref.lock(0).unwrap().append(&[i]).unwrap();
i += num_threads as u32;
}
});
}
});
let mut list = AdjacencyList::new();
let expected: Vec<_> = (0..neighbors.max_length()).map(|i| i as u32).collect();
neighbors.get(0, &mut list).unwrap();
list.sort();
assert_eq!(&*list, &*expected);
}
}