use moirai_utils::cache::CacheAligned;
use std::cell::{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_relaxed();
let consumer = self.consumer_acquire();
if current.wrapping_sub(consumer) >= self.buffer.len() {
return Err(value);
}
unsafe { self.produce_at(current, value) };
Ok(())
}
pub fn try_consume(&self) -> Option<T> {
let current = self.consumer_relaxed();
let producer = self.producer_acquire();
if current == producer {
return None;
}
Some(unsafe { self.consume_at(current) })
}
pub fn capacity(&self) -> usize {
self.buffer.len()
}
pub fn is_empty(&self) -> bool {
self.consumer_acquire() == self.producer_acquire()
}
pub fn is_full(&self) -> bool {
self.producer_acquire()
.wrapping_sub(self.consumer_acquire())
>= self.buffer.len()
}
pub fn len(&self) -> usize {
self.producer_acquire()
.wrapping_sub(self.consumer_acquire())
}
}
impl<T> RingBuffer<T> {
pub(crate) fn indices(&self) -> (usize, usize) {
(
self.producer_seq.0.load(Ordering::Relaxed),
self.consumer_seq.0.load(Ordering::Relaxed),
)
}
pub(crate) fn producer_relaxed(&self) -> usize {
self.producer_seq.0.load(Ordering::Relaxed)
}
pub(crate) fn producer_acquire(&self) -> usize {
self.producer_seq.0.load(Ordering::Acquire)
}
pub(crate) fn consumer_relaxed(&self) -> usize {
self.consumer_seq.0.load(Ordering::Relaxed)
}
pub(crate) fn consumer_acquire(&self) -> usize {
self.consumer_seq.0.load(Ordering::Acquire)
}
pub(crate) fn has_room(&self, producer: usize, cached_consumer: &Cell<usize>) -> bool {
if producer.wrapping_sub(cached_consumer.get()) < self.buffer.len() {
return true;
}
let consumer = self.consumer_seq.0.load(Ordering::Acquire);
cached_consumer.set(consumer);
producer.wrapping_sub(consumer) < self.buffer.len()
}
pub(crate) fn has_value(&self, consumer: usize, cached_producer: &Cell<usize>) -> bool {
if consumer != cached_producer.get() {
return true;
}
let producer = self.producer_seq.0.load(Ordering::Acquire);
cached_producer.set(producer);
consumer != producer
}
pub(crate) unsafe fn produce_at(&self, producer: usize, value: T) {
unsafe {
let slot = &mut *self.buffer[producer & self.mask].get();
slot.write(value);
}
self.producer_seq
.0
.store(producer.wrapping_add(1), Ordering::Release);
}
pub(crate) unsafe fn consume_at(&self, consumer: usize) -> T {
let value = unsafe {
let slot = &*self.buffer[consumer & self.mask].get();
slot.assume_init_read()
};
self.consumer_seq
.0
.store(consumer.wrapping_add(1), Ordering::Release);
value
}
}
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);
}
}