use moirai_utils::cache::CacheAligned;
use std::cell::UnsafeCell;
use std::mem::MaybeUninit;
use std::sync::atomic::{AtomicUsize, Ordering};
pub struct RingBuffer<T> {
buffer: Box<[UnsafeCell<MaybeUninit<T>>]>,
mask: usize,
producer_seq: CacheAligned<AtomicUsize>,
consumer_seq: CacheAligned<AtomicUsize>,
}
unsafe impl<T: Send> Send for RingBuffer<T> {}
impl<T> RingBuffer<T> {
pub fn new(capacity: usize) -> Self {
let capacity = capacity.next_power_of_two();
let buffer = (0..capacity)
.map(|_| UnsafeCell::new(MaybeUninit::uninit()))
.collect::<Vec<_>>()
.into_boxed_slice();
Self {
buffer,
mask: capacity - 1,
producer_seq: CacheAligned::new(AtomicUsize::new(0)),
consumer_seq: CacheAligned::new(AtomicUsize::new(0)),
}
}
pub fn try_produce(&self, value: T) -> Result<(), T> {
let current = self.producer_seq.0.load(Ordering::Relaxed);
let consumer = self.consumer_seq.0.load(Ordering::Acquire);
if current.wrapping_sub(consumer) >= self.buffer.len() {
return Err(value);
}
unsafe {
let slot = &mut *self.buffer[current & self.mask].get();
slot.write(value);
}
self.producer_seq
.0
.store(current.wrapping_add(1), Ordering::Release);
Ok(())
}
pub fn try_consume(&self) -> Option<T> {
let current = self.consumer_seq.0.load(Ordering::Relaxed);
let producer = self.producer_seq.0.load(Ordering::Acquire);
if current == producer {
return None;
}
let value = unsafe {
let slot = &*self.buffer[current & self.mask].get();
slot.assume_init_read()
};
self.consumer_seq
.0
.store(current.wrapping_add(1), Ordering::Release);
Some(value)
}
pub fn capacity(&self) -> usize {
self.buffer.len()
}
pub fn is_empty(&self) -> bool {
let consumer = self.consumer_seq.0.load(Ordering::Acquire);
let producer = self.producer_seq.0.load(Ordering::Acquire);
consumer == producer
}
pub fn is_full(&self) -> bool {
let consumer = self.consumer_seq.0.load(Ordering::Acquire);
let producer = self.producer_seq.0.load(Ordering::Acquire);
producer.wrapping_sub(consumer) >= self.buffer.len()
}
pub fn len(&self) -> usize {
let consumer = self.consumer_seq.0.load(Ordering::Acquire);
let producer = self.producer_seq.0.load(Ordering::Acquire);
producer.wrapping_sub(consumer)
}
}
impl<T> Drop for RingBuffer<T> {
fn drop(&mut self) {
let consumer = *self.consumer_seq.0.get_mut();
let producer = *self.producer_seq.0.get_mut();
let len = producer.wrapping_sub(consumer);
for i in 0..len {
let idx = (consumer.wrapping_add(i)) & self.mask;
unsafe {
let slot = &mut *self.buffer[idx].get();
slot.assume_init_drop();
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_wrapping_drop_correctness() {
use std::sync::atomic::{AtomicUsize, Ordering};
static DROP_COUNT: AtomicUsize = AtomicUsize::new(0);
struct TrackDrop;
impl Drop for TrackDrop {
fn drop(&mut self) {
DROP_COUNT.fetch_add(1, Ordering::SeqCst);
}
}
{
let mut rb = RingBuffer::<TrackDrop>::new(4);
let mask = rb.mask;
unsafe {
let slot1 = &mut *rb.buffer[(usize::MAX - 1) & mask].get();
slot1.write(TrackDrop);
let slot2 = &mut *rb.buffer[usize::MAX & mask].get();
slot2.write(TrackDrop);
let slot3 = &mut *rb.buffer[0].get();
slot3.write(TrackDrop);
}
*rb.consumer_seq.0.get_mut() = usize::MAX - 1;
*rb.producer_seq.0.get_mut() = 1;
}
assert_eq!(DROP_COUNT.load(Ordering::SeqCst), 3);
}
}