use std::{
collections::VecDeque,
sync::{
Mutex, PoisonError,
atomic::{AtomicU32, Ordering},
},
};
pub(super) struct SampleCursor {
next: AtomicU32,
end: u32,
local_retry: Mutex<VecDeque<(u32, u32)>>,
}
impl SampleCursor {
pub(super) const fn new(start: u32, end: u32) -> Self {
Self {
next: AtomicU32::new(start),
end,
local_retry: Mutex::new(VecDeque::new()),
}
}
pub(super) fn claim(&self, want: u32) -> Option<(u32, u32)> {
if want == 0 {
return None;
}
let start = self.next.fetch_add(want, Ordering::Relaxed);
if start >= self.end {
return None;
}
Some((start, want.min(self.end - start)))
}
pub(super) fn claim_local(&self, want: u32) -> Option<(u32, u32)> {
let retried = self
.local_retry
.lock()
.unwrap_or_else(PoisonError::into_inner)
.pop_front();
if retried.is_some() {
return retried;
}
self.claim(want)
}
pub(super) fn shared_pool_exhausted(&self) -> bool {
self.next.load(Ordering::Relaxed) >= self.end
}
pub(super) fn return_to_local(&self, start: u32, count: u32) {
if count == 0 {
return;
}
self.local_retry
.lock()
.unwrap_or_else(PoisonError::into_inner)
.push_back((start, count));
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::{collections::HashSet, thread};
#[test]
fn claim_hands_out_the_full_range_in_one_call_when_it_all_fits() {
let cursor = SampleCursor::new(0, 100);
assert_eq!(cursor.claim(1000), Some((0, 100)));
assert_eq!(cursor.claim(1), None, "the budget is now exhausted");
}
#[test]
fn claim_never_exceeds_the_budget_even_when_a_smaller_amount_remains() {
let cursor = SampleCursor::new(90, 100);
assert_eq!(cursor.claim(30), Some((90, 10)));
assert_eq!(cursor.claim(30), None);
}
#[test]
fn claim_starts_from_a_nonzero_offset() {
let cursor = SampleCursor::new(500, 600);
assert_eq!(cursor.claim(50), Some((500, 50)));
assert_eq!(cursor.claim(50), Some((550, 50)));
assert_eq!(cursor.claim(50), None);
}
#[test]
fn claim_of_zero_samples_is_always_none_and_never_advances_the_cursor() {
let cursor = SampleCursor::new(0, 100);
assert_eq!(cursor.claim(0), None);
assert_eq!(cursor.claim(100), Some((0, 100)));
}
#[test]
fn claim_local_prefers_the_retry_pile_over_fresh_work() {
let cursor = SampleCursor::new(0, 100);
cursor.return_to_local(40, 10);
assert_eq!(cursor.claim_local(5), Some((40, 10)));
assert_eq!(cursor.claim_local(5), Some((0, 5)));
}
#[test]
fn returned_ranges_are_invisible_to_claim() {
let cursor = SampleCursor::new(0, 0);
cursor.return_to_local(10, 5);
assert_eq!(cursor.claim(100), None);
assert_eq!(cursor.claim_local(100), Some((10, 5)));
}
#[test]
fn a_returned_remainder_is_reclaimable_exactly_once() {
let cursor = SampleCursor::new(0, 0);
cursor.return_to_local(200, 7);
assert_eq!(cursor.claim_local(100), Some((200, 7)));
assert_eq!(cursor.claim_local(100), None);
assert_eq!(cursor.claim(100), None);
}
#[test]
fn zero_count_returns_are_a_no_op() {
let cursor = SampleCursor::new(0, 0);
cursor.return_to_local(10, 0);
assert_eq!(cursor.claim_local(100), None);
}
#[test]
fn concurrent_claims_never_overlap_never_gap_and_never_exceed_the_budget() {
let total = 100_003u32; let cursor = SampleCursor::new(0, total);
let thread_count = 8;
let strides = [17u32, 40, 101, 256, 33, 512, 7, 999];
let claimed: Vec<Vec<(u32, u32)>> = thread::scope(|s| {
let handles: Vec<_> = (0..thread_count)
.map(|i| {
let cursor = &cursor;
let want = strides[i % strides.len()];
s.spawn(move || {
let mut mine = Vec::new();
while let Some(range) = cursor.claim(want) {
mine.push(range);
}
mine
})
})
.collect();
handles.into_iter().map(|h| h.join().unwrap()).collect()
});
let mut all: Vec<(u32, u32)> = claimed.into_iter().flatten().collect();
all.sort_unstable();
let mut cursor_pos = 0u32;
for (start, count) in &all {
assert_eq!(
*start, cursor_pos,
"expected the next range to start exactly at {cursor_pos}, found a gap \
or overlap at {start}"
);
assert!(*count > 0, "a claimed range must never be empty");
cursor_pos += *count;
}
assert_eq!(
cursor_pos, total,
"the union of every claimed range must cover the whole budget exactly once"
);
let mut seen = HashSet::with_capacity(total as usize);
for (start, count) in &all {
for idx in *start..(*start + *count) {
assert!(seen.insert(idx), "sample index {idx} was claimed twice");
}
}
assert_eq!(seen.len(), total as usize);
}
#[test]
fn concurrent_claim_local_still_partitions_the_budget_exactly() {
let total = 5_000u32;
let cursor = SampleCursor::new(0, total);
thread::scope(|s| {
let mut handles = Vec::new();
for i in 0..4 {
let cursor = &cursor;
let want = 30 + i as u32 * 11;
handles.push(s.spawn(move || {
let mut mine = Vec::new();
while let Some(range) = cursor.claim_local(want) {
mine.push(range);
}
mine
}));
}
let claimed: Vec<Vec<(u32, u32)>> =
handles.into_iter().map(|h| h.join().unwrap()).collect();
let mut all: Vec<(u32, u32)> = claimed.into_iter().flatten().collect();
all.sort_unstable();
let mut cursor_pos = 0u32;
for (start, count) in &all {
assert_eq!(*start, cursor_pos);
cursor_pos += *count;
}
assert_eq!(cursor_pos, total);
});
}
}