use std::{
num::NonZeroU32,
sync::atomic::{AtomicU32, Ordering},
};
use crossbeam_queue::ArrayQueue;
use diskann::utils::IntoUsize;
const SCAN_SIZE: u32 = 256;
#[derive(Debug)]
pub(crate) struct Freelist {
recycled: ArrayQueue<u32>,
max: u32,
next: AtomicU32,
scan_bucket: AtomicU32,
}
impl Freelist {
pub(crate) fn new(max: u32, recycled: NonZeroU32) -> Self {
Self {
recycled: ArrayQueue::new(recycled.get().into_usize()),
max,
next: AtomicU32::new(0),
scan_bucket: AtomicU32::new(0),
}
}
pub(crate) fn pop(&self) -> Id {
if let Some(id) = self.recycled.pop() {
return Id::Found(id);
}
let mut next = self.next.load(Ordering::Relaxed);
while next < self.max {
match self
.next
.compare_exchange(next, next + 1, Ordering::Relaxed, Ordering::Relaxed)
{
Ok(next) => return Id::Found(next),
Err(actual) => {
next = actual;
}
}
}
Id::Scan
}
pub(crate) fn pop_recycled(&self) -> Option<u32> {
self.recycled.pop()
}
pub(crate) fn scan(&self) -> Scan {
if self.max == 0 {
return Scan { start: 0, stop: 0 };
}
let num_buckets = self.max.div_ceil(SCAN_SIZE);
let bucket = self.scan_bucket.fetch_add(1, Ordering::Relaxed) % num_buckets;
let start = bucket * SCAN_SIZE;
let stop = match start.checked_add(SCAN_SIZE) {
Some(stop) => stop.min(self.max),
None => self.max,
};
Scan { start, stop }
}
pub(crate) fn push(&self, id: u32) -> bool {
if id < self.max {
self.recycled.push(id).is_ok()
} else {
false
}
}
}
#[derive(Debug, Clone, Copy)]
#[must_use]
pub(crate) enum Id {
Found(u32),
Scan,
}
#[cfg(test)]
impl Id {
fn unwrap(self) -> u32 {
match self {
Self::Found(i) => i,
Self::Scan => panic!("expected Id::Found, got Id::Scan"),
}
}
fn is_scan(self) -> bool {
matches!(self, Self::Scan)
}
}
#[derive(Debug)]
pub(crate) struct Scan {
start: u32,
stop: u32,
}
impl Scan {
#[cfg(test)]
fn as_range(&self) -> std::ops::Range<u32> {
self.start..self.stop
}
}
impl Iterator for Scan {
type Item = u32;
fn next(&mut self) -> Option<u32> {
if self.start >= self.stop {
None
} else {
let i = self.start;
self.start += 1;
Some(i)
}
}
fn size_hint(&self) -> (usize, Option<usize>) {
let len = (self.stop - self.start).into_usize();
(len, Some(len))
}
}
impl ExactSizeIterator for Scan {}
#[cfg(test)]
mod tests {
use super::*;
use std::{collections::HashSet, sync::Barrier, thread};
fn freelist(max: u32, recycled: u32) -> Freelist {
Freelist::new(max, NonZeroU32::new(recycled).unwrap())
}
#[test]
fn pop_mints_sequentially_until_exhausted() {
let fl = freelist(4, 8);
let mut got = Vec::new();
for _ in 0..4 {
got.push(fl.pop().unwrap());
}
assert_eq!(got, vec![0, 1, 2, 3]);
assert!(fl.pop().is_scan());
assert!(fl.pop().is_scan());
}
#[test]
fn pop_returns_scan_when_max_zero() {
let fl = freelist(0, 1);
assert!(fl.pop().is_scan());
}
#[test]
fn recycled_ids_take_precedence_over_minting() {
let fl = freelist(4, 8);
assert!(fl.push(2));
assert_eq!(fl.pop().unwrap(), 2);
assert_eq!(fl.pop().unwrap(), 0);
}
#[test]
fn push_rejects_ids_at_or_above_max() {
let fl = freelist(4, 8);
assert!(!fl.push(4));
assert!(!fl.push(u32::MAX));
assert!(fl.push(3));
assert_eq!(fl.pop_recycled().unwrap(), 3);
}
#[test]
fn push_returns_false_when_recycled_full() {
let fl = freelist(16, 2);
assert!(fl.push(2));
assert!(fl.push(3));
assert!(!fl.push(5));
assert_eq!(fl.pop().unwrap(), 2);
assert_eq!(fl.pop().unwrap(), 3);
}
#[test]
fn pop_recycled_empty_returns_none() {
let fl = freelist(4, 4);
assert!(fl.pop_recycled().is_none());
}
#[test]
fn pop_recycled_does_not_mint() {
let fl = freelist(4, 4);
assert!(fl.pop_recycled().is_none());
assert_eq!(fl.pop().unwrap(), 0);
}
fn as_vec<I>(itr: I) -> Vec<I::Item>
where
I: Iterator,
{
itr.collect()
}
#[test]
fn scan_on_empty_freelist_yields_nothing() {
let fl = freelist(0, 1);
let mut scan = fl.scan();
assert_eq!(scan.len(), 0);
assert!(scan.next().is_none());
}
#[test]
fn scan_covers_full_range_in_one_pass() {
let max = 2 * SCAN_SIZE + 50;
let fl = freelist(max, 4);
let scan = fl.scan();
assert_eq!(scan.as_range(), 0..SCAN_SIZE);
assert_eq!(scan.len(), SCAN_SIZE.into_usize());
assert_eq!(as_vec(scan), as_vec(0..SCAN_SIZE));
let scan = fl.scan();
assert_eq!(scan.as_range(), SCAN_SIZE..2 * SCAN_SIZE);
assert_eq!(scan.len(), SCAN_SIZE.into_usize());
assert_eq!(as_vec(scan), as_vec(SCAN_SIZE..2 * SCAN_SIZE));
let scan = fl.scan();
assert_eq!(scan.as_range(), 2 * SCAN_SIZE..(2 * SCAN_SIZE + 50));
assert_eq!(scan.len(), 50);
assert_eq!(as_vec(scan), as_vec((2 * SCAN_SIZE)..(2 * SCAN_SIZE + 50)));
let scan = fl.scan();
assert_eq!(scan.as_range(), 0..SCAN_SIZE);
assert_eq!(scan.len(), SCAN_SIZE.into_usize());
assert_eq!(as_vec(scan), as_vec(0..SCAN_SIZE));
}
#[test]
fn concurrent_pop_yields_unique_ids() {
let max = 4096u32;
let fl = Freelist::new(max, NonZeroU32::new(8).unwrap());
let nthreads = 8;
let barrier = Barrier::new(nthreads);
let results: Vec<Vec<u32>> = thread::scope(|s| {
let handles: Vec<_> = (0..nthreads)
.map(|_| {
s.spawn(|| {
let mut out = Vec::new();
barrier.wait();
while let Id::Found(id) = fl.pop() {
out.push(id);
}
out
})
})
.collect();
handles.into_iter().map(|h| h.join().unwrap()).collect()
});
let mut all: Vec<u32> = results.into_iter().flatten().collect();
all.sort();
let expected: Vec<u32> = (0..max).collect();
assert_eq!(all, expected, "all ids in [0, max) minted exactly once");
}
#[test]
fn concurrent_scan_partitions_one_pass() {
let max = SCAN_SIZE * 4;
let fl = Freelist::new(max, NonZeroU32::new(4).unwrap());
let num_buckets = max.div_ceil(SCAN_SIZE) as usize;
let nthreads = num_buckets;
let barrier = Barrier::new(nthreads);
let ids: Vec<u32> = thread::scope(|s| {
let handles: Vec<_> = (0..nthreads)
.map(|_| {
s.spawn(|| {
barrier.wait();
fl.scan().collect::<Vec<u32>>()
})
})
.collect();
handles
.into_iter()
.flat_map(|h| h.join().unwrap())
.collect()
});
let unique: HashSet<u32> = ids.iter().copied().collect();
assert_eq!(
unique.len(),
ids.len(),
"no id appeared twice across threads"
);
assert_eq!(
unique.len() as u32,
max,
"scans covered every id in [0, max)"
);
}
}