#![allow(dead_code)]
use core::cell::UnsafeCell;
use core::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
#[derive(Debug)]
pub enum WriteError {
OversizedWrite,
QueueFull,
}
#[repr(align(64))]
struct CachePadded<T> {
value: T,
}
impl<T> CachePadded<T> {
const fn new(value: T) -> Self {
Self { value }
}
}
struct QueueInner {
buf: Box<[UnsafeCell<u8>]>,
capacity: usize,
mask: usize,
write_pos: CachePadded<AtomicUsize>,
read_pos: CachePadded<AtomicUsize>,
}
unsafe impl Sync for QueueInner {}
#[doc(hidden)]
pub struct Producer {
inner: Arc<QueueInner>,
cached_read: usize,
write: usize,
}
pub struct Consumer {
inner: Arc<QueueInner>,
cached_write: usize,
read: usize,
}
unsafe impl Send for Consumer {}
#[must_use]
pub fn new(capacity: usize) -> (Producer, Consumer) {
assert!(capacity > 0, "queue capacity must be > 0");
let capacity = capacity.next_power_of_two();
let buf: Box<[UnsafeCell<u8>]> = core::iter::repeat_with(|| UnsafeCell::new(0u8))
.take(2 * capacity)
.collect::<Vec<_>>()
.into_boxed_slice();
let inner = Arc::new(QueueInner {
buf,
capacity,
mask: capacity - 1,
write_pos: CachePadded::new(AtomicUsize::new(0)),
read_pos: CachePadded::new(AtomicUsize::new(0)),
});
let producer = Producer {
inner: Arc::clone(&inner),
cached_read: 0,
write: 0,
};
let consumer = Consumer {
inner,
cached_write: 0,
read: 0,
};
(producer, consumer)
}
impl Producer {
#[cfg_attr(feature = "rtsan", rtsan_standalone::nonblocking)]
pub fn write(&mut self, n: usize, f: impl FnOnce(&mut [u8])) -> Result<(), WriteError> {
let Some(ptr) = self.try_reserve(n) else {
return if n > self.inner.capacity {
Err(WriteError::OversizedWrite)
} else {
Err(WriteError::QueueFull)
};
};
let buf = unsafe { core::slice::from_raw_parts_mut(ptr, n) };
f(buf);
self.commit(n);
Ok(())
}
#[cfg_attr(feature = "rtsan", rtsan_standalone::nonblocking)]
fn try_reserve(&mut self, n: usize) -> Option<*mut u8> {
let capacity = self.inner.capacity;
if n > capacity {
return None;
}
let mut available = capacity - self.write.wrapping_sub(self.cached_read);
if available < n {
self.cached_read = self.inner.read_pos.value.load(Ordering::Acquire);
available = capacity - self.write.wrapping_sub(self.cached_read);
if available < n {
return None;
}
}
let offset = self.write & self.inner.mask;
Some(unsafe { UnsafeCell::raw_get(self.inner.buf.as_ptr().add(offset)) })
}
#[cfg_attr(feature = "rtsan", rtsan_standalone::nonblocking)]
fn commit(&mut self, n: usize) {
self.write = self.write.wrapping_add(n);
self.inner
.write_pos
.value
.store(self.write, Ordering::Release);
}
}
impl Consumer {
pub fn available(&mut self) -> usize {
if self.cached_write > self.read {
return self.cached_write.wrapping_sub(self.read);
}
self.cached_write = self.inner.write_pos.value.load(Ordering::Acquire);
self.cached_write.wrapping_sub(self.read)
}
pub fn peek(&mut self, len: usize) -> &[u8] {
let avail = self.available();
debug_assert!(
len <= avail,
"peek: requested {len} bytes but only {avail} available",
);
let offset = self.read & self.inner.mask;
unsafe {
let ptr = UnsafeCell::raw_get(self.inner.buf.as_ptr().add(offset)).cast_const();
core::slice::from_raw_parts(ptr, len)
}
}
pub fn read<R>(&mut self, len: usize, f: impl FnOnce(&[u8]) -> R) -> R {
let avail = self.available();
debug_assert!(
len <= avail,
"read: requested {len} bytes but only {avail} available",
);
let offset = self.read & self.inner.mask;
let result = unsafe {
let ptr = UnsafeCell::raw_get(self.inner.buf.as_ptr().add(offset)).cast_const();
f(core::slice::from_raw_parts(ptr, len))
};
self.advance(len);
result
}
fn advance(&mut self, n: usize) {
self.read = self.read.wrapping_add(n);
self.inner
.read_pos
.value
.store(self.read, Ordering::Release);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn oversized_write_error() {
let (mut prod, _cons) = new(64);
let result = prod.write(65, |_| {});
assert!(matches!(result, Err(WriteError::OversizedWrite)));
}
#[test]
fn queue_full_error() {
let (mut prod, _cons) = new(64);
prod.write(64, |_| {}).unwrap();
let result = prod.write(1, |_| {});
assert!(matches!(result, Err(WriteError::QueueFull)));
}
#[test]
#[cfg(debug_assertions)]
#[should_panic(expected = "peek: requested 5 bytes but only 4 available")]
fn peek_panics_on_insufficient_data() {
let (mut prod, mut cons) = new(64);
prod.write(4, |buf| buf.fill(0xAA)).unwrap();
cons.peek(5);
}
#[test]
#[cfg(debug_assertions)]
#[should_panic(expected = "read: requested 5 bytes but only 4 available")]
fn read_panics_on_insufficient_data() {
let (mut prod, mut cons) = new(64);
prod.write(4, |buf| buf.fill(0xAA)).unwrap();
cons.read(5, |_| {});
}
#[test]
#[cfg(debug_assertions)]
#[should_panic(expected = "peek: requested 1 bytes but only 0 available")]
fn peek_panics_on_empty_queue() {
let (_prod, mut cons) = new(64);
cons.peek(1);
}
#[test]
#[cfg(debug_assertions)]
#[should_panic(expected = "read: requested 1 bytes but only 0 available")]
fn read_panics_on_empty_queue() {
let (_prod, mut cons) = new(64);
cons.read(1, |_| {});
}
#[test]
fn basic_write_peek_read() {
let (mut prod, mut cons) = new(64);
prod.write(4, |buf| buf.copy_from_slice(&[1u8, 2, 3, 4]))
.unwrap();
assert_eq!(cons.available(), 4);
assert_eq!(cons.peek(4), &[1u8, 2, 3, 4]);
assert_eq!(cons.available(), 4);
assert_eq!(cons.read(4, <[u8]>::to_vec), &[1u8, 2, 3, 4]);
assert_eq!(cons.available(), 0);
}
#[test]
fn wrap_around_contiguous() {
let (mut prod, mut cons) = new(64);
prod.write(60, |buf| buf.fill(0xAA)).unwrap();
assert_eq!(cons.available(), 60);
cons.read(60, |_| {});
prod.write(16, |buf| buf.fill(0xBB)).unwrap();
assert_eq!(cons.available(), 16);
assert_eq!(cons.read(16, <[u8]>::to_vec), &[0xBBu8; 16]);
assert_eq!(cons.available(), 0);
}
#[test]
fn queue_reuse_after_cycle() {
let (mut prod, mut cons) = new(16);
for i in 0..8_u8 {
prod.write(4, |buf| buf.copy_from_slice(&[i, i + 1, i + 2, i + 3]))
.unwrap();
assert_eq!(cons.read(4, <[u8]>::to_vec), &[i, i + 1, i + 2, i + 3]);
}
assert_eq!(cons.available(), 0);
}
#[test]
fn write_closure_buffer_length() {
let (mut prod, _cons) = new(64);
prod.write(13, |buf| {
assert_eq!(buf.len(), 13);
})
.unwrap();
}
#[test]
fn successive_reads() {
let (mut prod, mut cons) = new(64);
prod.write(4, |buf| buf.copy_from_slice(&[1u8, 2, 3, 4]))
.unwrap();
prod.write(4, |buf| buf.copy_from_slice(&[5u8, 6, 7, 8]))
.unwrap();
assert_eq!(cons.read(4, <[u8]>::to_vec), &[1u8, 2, 3, 4]);
assert_eq!(cons.read(4, <[u8]>::to_vec), &[5u8, 6, 7, 8]);
assert_eq!(cons.available(), 0);
}
#[test]
fn successive_reads_across_capacity_boundary() {
let (mut prod, mut cons) = new(16);
prod.write(12, |buf| buf.fill(0x00)).unwrap();
cons.read(12, |_| {});
prod.write(8, |buf| buf.fill(0xAA)).unwrap();
prod.write(4, |buf| buf.fill(0xBB)).unwrap();
assert_eq!(cons.available(), 12);
assert_eq!(cons.peek(4), &[0xAAu8; 4]);
assert_eq!(cons.peek(8), &[0xAAu8; 8]);
assert_eq!(cons.available(), 12);
assert_eq!(cons.read(8, <[u8]>::to_vec), &[0xAAu8; 8]);
assert_eq!(cons.available(), 4);
assert_eq!(cons.read(4, <[u8]>::to_vec), &[0xBBu8; 4]);
assert_eq!(cons.available(), 0);
}
#[test]
#[should_panic(expected = "queue capacity must be > 0")]
fn new_panics_on_zero_capacity() {
let _ = new(0);
}
#[test]
fn capacity_rounds_to_next_power_of_two() {
let (mut prod, _cons) = new(3);
assert!(prod.write(4, |_| {}).is_ok());
let (mut prod, _cons) = new(3);
assert!(matches!(
prod.write(5, |_| {}),
Err(WriteError::OversizedWrite)
));
}
#[test]
fn full_capacity_write_after_partial_drain() {
let (mut prod, mut cons) = new(16);
prod.write(8, |buf| buf.fill(0x00)).unwrap();
cons.read(8, |_| {});
prod.write(16, |buf| buf.fill(0xCC)).unwrap();
assert_eq!(cons.available(), 16);
assert_eq!(cons.read(16, <[u8]>::to_vec), &[0xCCu8; 16]);
assert_eq!(cons.available(), 0);
}
#[test]
fn read_race_advance_before_data_used() {
let (mut prod, mut cons) = new(4);
prod.write(4, |buf| buf.copy_from_slice(&[1u8, 2, 3, 4]))
.unwrap();
let handle = std::thread::spawn(move || {
loop {
if prod.write(4, |buf| buf.fill(0xFF)).is_ok() {
break;
}
std::hint::spin_loop();
}
});
let got = cons.read(4, |slice| {
let copy: [u8; 4] = slice.try_into().unwrap();
copy
});
handle.join().unwrap();
assert_eq!(got, [1u8, 2, 3, 4]);
}
#[test]
fn threaded_producer_consumer() {
let (mut prod, mut cons) = new(64);
let consumer_thread = std::thread::spawn(move || {
loop {
if cons.available() >= 4 {
return cons.read(4, <[u8]>::to_vec);
}
std::hint::spin_loop();
}
});
prod.write(4, |buf| buf.copy_from_slice(&[1u8, 2, 3, 4]))
.unwrap();
let received = consumer_thread.join().unwrap();
assert_eq!(received, &[1u8, 2, 3, 4]);
}
}