use core::cell::UnsafeCell;
#[derive(Debug)]
pub(crate) struct RingBuffer<E> {
buf: Box<[UnsafeCell<E>]>,
mask: usize,
}
unsafe impl<E> Sync for RingBuffer<E> where E: Sync {}
impl<E> RingBuffer<E> {
pub(crate) fn from_factory(size: usize, mut factory: impl FnMut() -> E) -> Self {
debug_assert!(size > 0 && size.is_power_of_two());
let buf = (0..size).map(|_| UnsafeCell::new(factory())).collect();
let mask = size - 1;
RingBuffer { buf, mask }
}
pub(crate) fn from_buffer(buf: Box<[E]>) -> Self {
let size = buf.len();
debug_assert!(size > 0 && size.is_power_of_two());
let buf: Box<[UnsafeCell<E>]> = unsafe { core::mem::transmute(buf) };
let mask = size - 1;
RingBuffer { buf, mask }
}
#[inline]
unsafe fn iter(&self, start: usize, end: usize) -> core::slice::Iter<'_, UnsafeCell<E>> {
let slice = unsafe { self.buf.get_unchecked(start..end) };
slice.iter()
}
#[inline]
pub(crate) unsafe fn apply<F>(&self, seq: i64, size: usize, mut func: F)
where
F: FnMut(*mut E, i64, bool),
{
debug_assert!(seq >= 0);
debug_assert!(size < self.size());
let index = (seq & self.mask as i64) as usize;
unsafe {
let end = seq + size as i64 - 1;
if index + size > self.size() {
let diff = self.size() - index;
self.iter(index, self.size())
.zip(seq..)
.for_each(|(elem, s)| func(elem.get(), s, false));
self.iter(0, size - diff)
.zip(seq + diff as i64..)
.for_each(|(elem, s)| func(elem.get(), s, s == end))
} else {
self.iter(index, index + size)
.zip(seq..)
.for_each(|(elem, s)| func(elem.get(), s, s == end))
}
}
}
#[inline]
pub(crate) unsafe fn try_apply<F, Err>(&self, seq: i64, size: usize, func: F) -> Result<(), Err>
where
F: FnMut(*mut E, i64, bool) -> Result<(), Err>,
{
let mut func = func;
debug_assert!(seq >= 0);
debug_assert!(size < self.size());
let index = (seq & self.mask as i64) as usize;
unsafe {
let end = seq + size as i64 - 1;
if index + size > self.size() {
let diff = self.size() - index;
self.iter(index, self.size())
.zip(seq..)
.try_for_each(|(elem, s)| func(elem.get(), s, false))?;
self.iter(0, size - diff)
.zip(seq + diff as i64..)
.try_for_each(|(elem, s)| func(elem.get(), s, s == end))
} else {
self.iter(index, index + size)
.zip(seq..)
.try_for_each(|(elem, s)| func(elem.get(), s, s == end))
}
}
}
#[inline]
pub(crate) const fn size(&self) -> usize {
self.buf.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_wraps_batch_end_count() {
let buffer = RingBuffer::from_factory(8, || 0u8);
let mut end_count = 0;
unsafe {
buffer.apply(6, 4, |_, _, end| {
if end {
end_count += 1;
}
})
};
assert_eq!(end_count, 1);
}
}