use std::ptr;
use std::sync::atomic::Ordering::{AcqRel, Acquire, Relaxed, Release};
use std::sync::atomic::{AtomicBool, AtomicPtr, AtomicUsize};
use crate::error::MesoError;
#[derive(Debug)]
pub struct BufferWheel<const N: usize, T> {
buffers: [AtomicPtr<T>; N],
read: AtomicUsize,
write: AtomicUsize,
full: AtomicBool,
}
impl<const N: usize, T> Default for BufferWheel<N, T> {
fn default() -> Self {
BufferWheel::new()
}
}
impl<const N: usize, T> BufferWheel<N, T> {
pub fn new() -> Self {
let buffers = array_init::array_init(|_| AtomicPtr::new(ptr::null_mut()));
Self {
buffers,
read: AtomicUsize::new(0),
write: AtomicUsize::new(0),
full: AtomicBool::new(false),
}
}
pub fn write(&self, data: T) -> Result<(), MesoError> {
let write = self.write.load(Relaxed);
if self.full.load(Acquire) {
return Err(MesoError::BuffersFull);
}
let new_ptr = Box::into_raw(Box::new(data));
let old_ptr = self.buffers[write].swap(new_ptr, AcqRel);
if !old_ptr.is_null() {
unsafe {
drop(Box::from_raw(old_ptr));
}
}
let next = (write + 1) % N;
self.write.store(next, Release);
if next == self.read.load(Acquire) {
self.full.store(true, Release);
}
Ok(())
}
pub fn read(&self) -> Result<T, MesoError> {
let read = self.read.load(Relaxed);
if read == self.write.load(Acquire) && !self.full.load(Acquire) {
return Err(MesoError::NoPendingUpdates);
}
let null = ptr::null_mut();
let old_ptr = self.buffers[read].swap(null, AcqRel);
if old_ptr.is_null() {
return Err(MesoError::ExpectedUpdate);
}
let edge = unsafe { Box::from_raw(old_ptr) };
let info = *edge;
let next = (read + 1) % N;
self.read.store(next, Release);
self.full.store(false, Release);
Ok(info)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::thread;
#[test]
fn sequential_write_read() {
let buf = BufferWheel::<3, i32>::default();
assert_eq!(buf.read().unwrap_err(), MesoError::NoPendingUpdates);
buf.write(42).expect("first write okay");
buf.write(1337).expect("second write okay");
assert_eq!(buf.read().unwrap(), 42);
assert_eq!(buf.read().unwrap(), 1337);
assert_eq!(buf.read().unwrap_err(), MesoError::NoPendingUpdates);
}
#[test]
fn capacity_full_and_recover() {
let buf = BufferWheel::<2, u8>::default();
assert!(buf.write(10).is_ok());
assert!(buf.write(20).is_ok());
let e = buf.write(30).unwrap_err();
assert_eq!(e, MesoError::BuffersFull);
assert_eq!(buf.read().unwrap(), 10);
buf.write(30).expect("recovered after read");
assert_eq!(buf.read().unwrap(), 20);
assert_eq!(buf.read().unwrap(), 30);
assert_eq!(buf.read().unwrap_err(), MesoError::NoPendingUpdates);
}
#[test]
fn spsc_concurrent_spinning() {
let buf = Arc::new(BufferWheel::<1, usize>::default());
let prod = Arc::clone(&buf);
let cons = Arc::clone(&buf);
let writer = thread::spawn(move || {
for i in 0..100 {
loop {
match prod.write(i) {
Ok(_) => break,
Err(MesoError::BuffersFull) => continue,
Err(e) => panic!("unexpected write error: {e:?}"),
}
}
}
});
let reader = thread::spawn(move || {
for expected in 0..100 {
loop {
match cons.read() {
Ok(v) => {
assert_eq!(v, expected);
break;
}
Err(MesoError::NoPendingUpdates) => continue,
Err(e) => panic!("unexpected read error: {e:?}"),
}
}
}
});
writer.join().unwrap();
reader.join().unwrap();
}
}