use std::sync::atomic::{AtomicU32, Ordering};
use std::mem::MaybeUninit;
use std::io::ErrorKind;
use std::fmt::{Display, Formatter};
pub struct RingBuffer<Slot, const RING_BUFFER_SIZE: usize> {
reserved_tail: AtomicU32,
published_tail: AtomicU32,
buffer: MaybeUninit<[Slot; RING_BUFFER_SIZE]>,
}
impl<Slot, const RING_BUFFER_SIZE: usize>
Default
for RingBuffer<Slot, RING_BUFFER_SIZE> {
fn default() -> Self {
Self::new()
}
}
impl<Slot, const RING_BUFFER_SIZE: usize> RingBuffer<Slot, RING_BUFFER_SIZE> {
pub const fn new() -> Self {
Self {
reserved_tail: AtomicU32::new(0),
published_tail: AtomicU32::new(0),
buffer: MaybeUninit::uninit(),
}
}
pub fn consumer(&self) -> RingBufferConsumer<'_, Slot, RING_BUFFER_SIZE> {
RingBufferConsumer {
head: AtomicU32::new(self.published_tail.load(Ordering::Relaxed)),
ring_buffer: self,
}
}
pub fn enqueue(&self, element: Slot) {
let reserved_tail = self.reserved_tail.fetch_add(1, Ordering::Relaxed);
let mutable_buffer = unsafe {
let const_ptr = self.buffer.as_ptr();
let mut_ptr = const_ptr as *mut [Slot; RING_BUFFER_SIZE];
&mut *mut_ptr
};
mutable_buffer[reserved_tail as usize % RING_BUFFER_SIZE] = element;
loop {
match self.published_tail.compare_exchange_weak(reserved_tail, reserved_tail+1, Ordering::Release, Ordering::Relaxed) {
Ok(_) => return,
Err(reloaded_val) => if reloaded_val > reserved_tail {
panic!("BUG: Infinite loop detected in Ring-Buffer. Please fix.");
},
}
}
}
pub fn get_buffer_size(&self) -> usize {
RING_BUFFER_SIZE
}
}
pub struct RingBufferConsumer<'a, Slot, const RING_BUFFER_SIZE: usize> {
head: AtomicU32,
ring_buffer: &'a RingBuffer<Slot, RING_BUFFER_SIZE>,
}
impl<Slot, const RING_BUFFER_SIZE: usize> RingBufferConsumer<'_, Slot, RING_BUFFER_SIZE> {
pub fn dequeue(&self) -> Result<Option<&Slot>, RingBufferOverflowError> {
let mut head = self.head.load(Ordering::Relaxed);
loop {
let published_tail = self.ring_buffer.published_tail.load(Ordering::Relaxed);
if head > published_tail {
head = self.head.load(Ordering::Relaxed);
continue;
}
if head == published_tail {
return Ok(None);
}
match self.head.compare_exchange_weak(head, head + 1, Ordering::Acquire, Ordering::Relaxed) {
Ok(_) => unsafe {
let ptr = self.ring_buffer.buffer.as_ptr();
let array = &*ptr;
if self.ring_buffer.reserved_tail.load(Ordering::Relaxed) - head > RING_BUFFER_SIZE as u32 {
return Err(RingBufferOverflowError { msg: format!("Ring-Buffer overflow: published_tail={}, head={} -- tail could not be farther from head than the ring buffer size of {}", published_tail, head, RING_BUFFER_SIZE) });
}
return Ok(Some(&array[head as usize % RING_BUFFER_SIZE]))
},
Err(reloaded_head) => head = reloaded_head,
}
}
}
pub fn peek_all(&self) -> Result<[&[Slot];2], RingBufferOverflowError> {
let head = self.head.load(Ordering::Relaxed);
let published_tail = self.ring_buffer.published_tail.load(Ordering::Relaxed);
let head_index = head as usize % RING_BUFFER_SIZE;
let published_tail_index = published_tail as usize % RING_BUFFER_SIZE;
if head == published_tail {
Ok([&[],&[]])
} else if published_tail - head > RING_BUFFER_SIZE as u32 {
Err(RingBufferOverflowError { msg: format!("Ring-Buffer overflow: published_tail={}, head={} -- tail could not be farther from head than the ring buffer size of {}", published_tail, head, RING_BUFFER_SIZE) })
} else if head_index < published_tail_index {
unsafe {
let ptr = self.ring_buffer.buffer.as_ptr();
let array = &*ptr;
Ok([&array[head_index .. published_tail_index], &[]])
}
} else {
unsafe {
let ptr = self.ring_buffer.buffer.as_ptr();
let array = &*ptr;
Ok([&array[head_index..RING_BUFFER_SIZE], &array[0..published_tail_index]])
}
}
}
}
#[derive(Debug)]
pub struct RingBufferOverflowError {
msg: String,
}
impl Display for RingBufferOverflowError {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "RingBufferOverflowError: {}", self.msg)
}
}
impl std::error::Error for RingBufferOverflowError {}
impl From<RingBufferOverflowError> for std::io::Error {
fn from(custom_error: RingBufferOverflowError) -> Self {
std::io::Error::new(ErrorKind::InvalidInput, custom_error)
}
}
#[cfg(test)]
mod tests {
use super::*;
use serial_test::serial;
use std::fmt::Debug;
#[test]
fn simple_enqueue_dequeue_use_cases() {
let ring_buffer = RingBuffer::<i32, 16>::new();
let consumer = ring_buffer.consumer();
match consumer.dequeue() {
Ok(None) => (), Ok(Some(existing_element)) => panic!("Something was dequeued when noting should have been: {:?}", existing_element),
Err(error) => panic!("RingBufferOverflowError while dequeueing : {:?}", error),
}
let expected = 123;
ring_buffer.enqueue(expected);
match consumer.dequeue() {
Ok(None) => panic!("No element was dequeued"),
Ok(Some(existing_element)) => assert_eq!(existing_element, &expected, "Wrong element dequeued"),
Err(error) => panic!("RingBufferOverflowError while dequeueing : {:?}", error),
}
for i in 0..2*ring_buffer.get_buffer_size() as i32 {
ring_buffer.enqueue(i);
match consumer.dequeue() {
Ok(None) => panic!("No element was dequeued"),
Ok(Some(existing_element)) => assert_eq!(existing_element, &i, "Wrong element dequeued"),
Err(error) => panic!("RingBufferOverflowError while dequeueing : {:?}", error),
}
}
for i in 0..ring_buffer.get_buffer_size() as i32 {
ring_buffer.enqueue(i);
}
for i in 0..ring_buffer.get_buffer_size() as i32 {
match consumer.dequeue() {
Ok(None) => panic!("No element was dequeued"),
Ok(Some(existing_element)) => assert_eq!(existing_element, &i, "Wrong element dequeued"),
Err(error) => panic!("RingBufferOverflowError while dequeueing : {:?}", error),
}
}
match consumer.dequeue() {
Ok(None) => (), Ok(Some(existing_element)) => panic!("No element should have been left behind, yet {} was dequeued", existing_element),
Err(error) => panic!("RingBufferOverflowError while dequeueing : {:?}", error),
}
}
#[test]
fn peek() -> Result<(), RingBufferOverflowError> {
let ring_buffer = RingBuffer::<u32, 16>::new();
let consumer = ring_buffer.consumer();
let check_name = "empty peek";
let expected_elements = &[];
assert_eq!(consumer.peek_all()?.concat(), expected_elements, "{} failed", check_name);
let check_name = "peek for a single element";
let expected_elements = &[1];
ring_buffer.enqueue(1);
assert_eq!(consumer.peek_all()?.concat(), expected_elements, "{} failed", check_name);
let check_name = "peek also an additional element";
let expected_elements = &[1, 2];
ring_buffer.enqueue(2);
assert_eq!(consumer.peek_all()?.concat(), expected_elements, "{} failed", check_name);
let check_name = "peek the whole ring-buffer";
for e in 3..1+ring_buffer.get_buffer_size() as u32 {
ring_buffer.enqueue(e);
}
let expected_elements: Vec<u32> = (1..1+ring_buffer.get_buffer_size() as u32).into_iter().collect();
assert_eq!(consumer.peek_all()?.concat(), expected_elements, "{} failed", check_name);
let check_name = "ring goes round";
let expected_elements = &[16,17];
for _ in 1..ring_buffer.get_buffer_size() as u32 {
consumer.dequeue().unwrap();
}
ring_buffer.enqueue(17);
assert_eq!(consumer.peek_all()?.concat(), expected_elements, "{} failed", check_name);
let check_name = "EXTRA: demonstration on how to iterate over peeked objects without a vector (or any other) allocation";
let mut observed_elements = Vec::<u32>::new();
for peeked_chunk in consumer.peek_all()? {
for peeked_element in peeked_chunk {
observed_elements.push(*peeked_element);
}
}
assert_eq!(observed_elements, expected_elements, "{} failed", check_name);
Ok(())
}
#[test]
#[serial] fn buffer_overflowing() {
let ring_buffer = RingBuffer::<i32, 16>::new();
let consumer = ring_buffer.consumer();
for i in 0..1+ring_buffer.get_buffer_size() as i32 {
ring_buffer.enqueue(i);
}
let peeked_chunks = consumer.peek_all();
assert_buffer_overflow("Peeking", peeked_chunks, "Ring-Buffer overflow: published_tail=17, head=0 -- tail could not be farther from head than the ring buffer size of 16");
let element = consumer.dequeue();
assert_buffer_overflow("Dequeueing", element, "Ring-Buffer overflow: published_tail=17, head=0 -- tail could not be farther from head than the ring buffer size of 16");
fn assert_buffer_overflow<E: Debug>(operation: &str, result: Result<E, RingBufferOverflowError>, expected_error_message: &str) {
if result.is_ok() {
panic!("{} from an overflowed ring buffer was allowed, when it shouldn't. Returned element was {:?} -- if overflow didn't happen, it would be 0", operation, result);
} else {
let observed_error_message = result.unwrap_err().msg;
assert_eq!(observed_error_message, expected_error_message, "Wrong error message received");
}
}
}
#[test]
#[serial]
fn concurrency() {
let ring_buffer = RingBuffer::<u32, 40960>::new();
let consumer = ring_buffer.consumer();
for threads in 1..16 {
let start = 0;
let finish = 40960/10;
multi_threaded_iterate(start, finish, threads, |i| ring_buffer.enqueue(i));
let expected_sum = (finish - 1) * (finish - start) / 2;
let observed_sum = AtomicU32::new(0);
multi_threaded_iterate(start, finish, threads, |_| match consumer.dequeue() {
Ok(None) => panic!("Ran out of elements prematurely"),
Ok(Some(existing_element)) => { observed_sum.fetch_add(*existing_element, Ordering::Relaxed); },
Err(error) => panic!("RingBufferOverflowError while dequeueing : {:?}", error),
});
assert_eq!(observed_sum.load(Ordering::Relaxed), expected_sum, "Error in all-in / all-out multi-threaded test (with {} threads)", threads);
}
let start = 0;
let finish = 92600/10;
let threads = 16;
let expected_sum = (start + (finish-1)) * ( (finish - start) / 2 );
let expected_callback_calls = finish - start;
let observed_callback_calls = AtomicU32::new(0);
let observed_sum = AtomicU32::new(0);
multi_threaded_iterate(start, finish, threads, |i| {
observed_callback_calls.fetch_add(1, Ordering::Relaxed);
ring_buffer.enqueue(i);
match consumer.dequeue() {
Ok(Some(existing_element)) => observed_sum.fetch_add(*existing_element, Ordering::Relaxed),
Ok(None) => panic!("Ran out of elements prematurely"),
Err(error) => panic!("RingBufferOverflowError while dequeueing : {:?}", error),
};
});
assert_eq!(observed_callback_calls.load(Ordering::Relaxed), expected_callback_calls, "¿Wrong number of callback calls?");
assert_eq!(observed_sum.load(Ordering::Relaxed), expected_sum, "Error in single-in / single-out multi-threaded test (with {} threads)", threads);
fn multi_threaded_iterate(start: u32, finish: u32, threads: u32, callback: impl Fn(u32) -> () + std::marker::Sync) {
crossbeam::scope(|scope| {
let cb = &callback;
let join_handlers: Vec<crossbeam::thread::ScopedJoinHandle<()>> = (start..start+threads).into_iter()
.map(|thread_number| scope.spawn(move |_| iterate(thread_number, finish, threads, &cb)))
.collect();
for join_handler in join_handlers {
join_handler.join().unwrap();
}
}).unwrap();
}
fn iterate(start: u32, finish: u32, step: u32, callback: impl Fn(u32) -> () + std::marker::Sync) {
for i in (start..finish).step_by(step as usize) {
callback(i);
}
}
}
}