use crate::{
QueueError, pack_entry,
traits::{QueueConsumer, QueueFactory, QueueProducer},
unpack_entry,
};
use crossbeam_utils::CachePadded;
use portable_atomic::AtomicU128;
use std::{
fmt,
marker::PhantomData,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
};
const MAX_ATTEMPTS: usize = if cfg!(test) { 1000 } else { u16::MAX as usize };
enum Storage<const N: usize> {
Static([CachePadded<AtomicU128>; N]),
Dynamic(Box<[CachePadded<AtomicU128>]>),
}
impl<const N: usize> Storage<N> {
#[inline]
fn load(&self, idx: usize, order: Ordering) -> u128 {
match self {
Self::Static(a) => a[idx].load(order),
Self::Dynamic(v) => v[idx].load(order),
}
}
#[inline]
fn compare_exchange(
&self,
idx: usize,
old: u128,
new: u128,
success: Ordering,
failure: Ordering,
) -> Result<u128, u128> {
match self {
Self::Static(a) => a[idx].compare_exchange(old, new, success, failure),
Self::Dynamic(v) => v[idx].compare_exchange(old, new, success, failure),
}
}
}
pub struct MpmcQueue<T, I = u32, const N: usize = 0>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
storage: Storage<N>,
capacity: usize,
mask: usize,
write_index: AtomicUsize,
pub(crate) read_index: AtomicUsize,
data_size: usize,
seq_shift: u32,
_phantom: PhantomData<(T, I)>,
}
impl<T, I, const N: usize> fmt::Debug for MpmcQueue<T, I, N>
where
T: Copy + Send + Sync + Default + fmt::Debug,
I: Copy + Into<u128> + fmt::Debug,
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MpmcQueue")
.field("capacity", &self.capacity)
.field("len", &self.len())
.field("is_empty", &self.is_empty())
.finish_non_exhaustive()
}
}
#[derive(Debug, Clone)]
pub struct QueueBuilder<T, I = u32>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
capacity: Option<usize>,
_phantom: PhantomData<(T, I)>,
}
impl<T, I> Default for QueueBuilder<T, I>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
fn default() -> Self {
Self::new()
}
}
impl<T, I> QueueBuilder<T, I>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
pub const fn new() -> Self {
Self {
capacity: None,
_phantom: PhantomData,
}
}
#[must_use]
pub const fn capacity(mut self, cap: usize) -> Self {
self.capacity = Some(cap);
self
}
pub fn build(self) -> Result<Arc<MpmcQueue<T, I>>, QueueError> {
let capacity = self.capacity.ok_or(QueueError::InvalidCapacity)?;
Ok(Arc::new(MpmcQueue::new(capacity)?))
}
pub fn build_static<const N: usize>(self) -> Result<Arc<MpmcQueue<T, I, N>>, QueueError> {
let capacity = self.capacity.unwrap_or(N);
Ok(Arc::new(MpmcQueue::new(capacity)?))
}
pub fn channels(self) -> Result<(Producer<T, I>, Consumer<T, I>), QueueError> {
let queue = self.build()?;
Ok((queue.producer(), queue.consumer()))
}
pub fn channels_static<const N: usize>(
self,
) -> Result<(Producer<T, I, N>, Consumer<T, I, N>), QueueError> {
let queue = self.build_static::<N>()?;
Ok((queue.producer(), queue.consumer()))
}
}
pub const fn queue<T>() -> QueueBuilder<T, u32>
where
T: Copy + Send + Sync + Default,
{
QueueBuilder::new()
}
pub const fn queue_with_index<T, I>() -> QueueBuilder<T, I>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
QueueBuilder::new()
}
impl<T, I, const N: usize> MpmcQueue<T, I, N>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
pub(crate) fn new(mut cap: usize) -> Result<Self, QueueError> {
if N > 0 && cap != N {
return Err(QueueError::CapacityMismatch);
}
cap = cap.max(2).next_power_of_two();
let data_size = size_of::<T>();
let index_size = size_of::<I>();
if data_size + index_size > 16 {
return Err(QueueError::TypeSizeExceeded { size: data_size });
}
let seq_shift = u32::try_from(index_size * 8).map_err(|_| QueueError::CapacityMismatch)?;
let seq_mask = (1u128 << seq_shift) - 1u128;
let storage = if N > 0 {
Storage::Static(std::array::from_fn(|i| {
let seq = ((i as u128) << 1) & seq_mask;
let packed = pack_entry::<T, I>(seq, T::default(), seq_shift);
CachePadded::new(AtomicU128::new(packed))
}))
} else {
Storage::Dynamic(
(0..cap)
.map(|i| {
let seq = ((i as u128) << 1) & seq_mask;
let packed = pack_entry::<T, I>(seq, T::default(), seq_shift);
CachePadded::new(AtomicU128::new(packed))
})
.collect(),
)
};
Ok(Self {
storage,
capacity: cap,
mask: cap - 1,
write_index: AtomicUsize::new(0),
read_index: AtomicUsize::new(0),
data_size,
seq_shift,
_phantom: PhantomData,
})
}
pub const fn capacity(&self) -> usize {
self.capacity
}
pub fn len(&self) -> usize {
let read = self.read_index.load(Ordering::Relaxed);
let write = self.write_index.load(Ordering::Relaxed);
write.wrapping_sub(read)
}
pub fn is_empty(&self) -> bool {
let rd = self.read_index.load(Ordering::Acquire);
let idx = rd & self.mask;
let packed = self.storage.load(idx, Ordering::Acquire);
let (seq, _) = unpack_entry::<T, I>(packed, self.seq_shift, self.data_size);
seq == ((rd as u128) << 1)
}
pub fn is_full(&self) -> bool {
self.len() >= self.capacity
}
pub fn exchange(&self, index: usize, old_value: T, new_value: T) -> bool {
let idx = index & self.mask;
let seq = ((index as u128) << 1) | 1u128;
let old_packed = pack_entry::<T, I>(seq, old_value, self.seq_shift);
let new_packed = pack_entry::<T, I>(seq, new_value, self.seq_shift);
self.storage
.compare_exchange(
idx,
old_packed,
new_packed,
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
}
pub fn try_push(&self, value: T) -> Result<(), (T, QueueError)> {
match self.push_impl(value, false).map(|_| ()) {
Ok(()) => Ok(()),
Err(QueueError::Full) => Err((value, QueueError::Full)),
Err(e) => Err((value, e)),
}
}
pub fn push(&self, value: T) -> Result<(), QueueError> {
self.push_impl(value, true).map(|_| ())
}
pub(crate) fn push_impl(&self, value: T, retry: bool) -> Result<usize, QueueError> {
let mut attempts = 0;
loop {
if !retry && attempts > 0 {
return Err(QueueError::Full);
}
attempts += 1;
if attempts > MAX_ATTEMPTS {
return Err(QueueError::Full);
}
let wr = self.write_index.load(Ordering::Acquire);
let idx = wr & self.mask;
let packed = self.storage.load(idx, Ordering::Acquire);
let (seq, _) = unpack_entry::<T, I>(packed, self.seq_shift, self.data_size);
let expected_seq = (wr as u128) << 1;
if seq == expected_seq {
let new_seq = expected_seq | 1u128;
let new_packed = pack_entry::<T, I>(new_seq, value, self.seq_shift);
if self
.storage
.compare_exchange(idx, packed, new_packed, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
let _ = self.write_index.compare_exchange_weak(
wr,
wr + 1,
Ordering::AcqRel,
Ordering::Acquire,
);
return Ok(wr); }
} else if seq == (expected_seq | 1u128) {
let _ = self.write_index.compare_exchange_weak(
wr,
wr + 1,
Ordering::AcqRel,
Ordering::Acquire,
);
} else if seq.wrapping_add((self.capacity as u128) << 1) == (expected_seq | 1u128) {
if !retry {
return Err(QueueError::Full);
}
std::hint::spin_loop();
}
}
}
pub fn try_pop(&self) -> Result<T, QueueError> {
self.pop_impl(false).map(|(data, _)| data)
}
pub fn pop(&self) -> Result<T, QueueError> {
self.pop_impl(true).map(|(data, _)| data)
}
pub(crate) fn pop_impl(&self, retry: bool) -> Result<(T, usize), QueueError> {
let mut attempts = 0;
loop {
if !retry && attempts > 0 {
return Err(QueueError::Empty);
}
attempts += 1;
if attempts > MAX_ATTEMPTS {
return Err(QueueError::Empty);
}
let rd = self.read_index.load(Ordering::Acquire);
let idx = rd & self.mask;
let packed = self.storage.load(idx, Ordering::Acquire);
let (seq, data) = unpack_entry::<T, I>(packed, self.seq_shift, self.data_size);
let expected_full = ((rd as u128) << 1) | 1u128;
if seq == expected_full {
let new_seq =
(((rd + self.capacity) as u128) << 1) & ((1u128 << self.seq_shift) - 1u128);
let new_packed = pack_entry::<T, I>(new_seq, T::default(), self.seq_shift);
if self
.storage
.compare_exchange(idx, packed, new_packed, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
let _ = self.read_index.compare_exchange_weak(
rd,
rd + 1,
Ordering::AcqRel,
Ordering::Acquire,
);
return Ok((data, rd)); }
} else if seq == ((rd as u128) << 1) {
if !retry {
return Err(QueueError::Empty);
}
std::hint::spin_loop();
} else {
let _ = self.read_index.compare_exchange_weak(
rd,
rd + 1,
Ordering::AcqRel,
Ordering::Acquire,
);
}
}
}
pub fn peek(&self) -> Result<T, QueueError> {
let rd = self.read_index.load(Ordering::Acquire);
let idx = rd & self.mask;
let packed = self.storage.load(idx, Ordering::Acquire);
let (seq, data) = unpack_entry::<T, I>(packed, self.seq_shift, self.data_size);
if seq == (((rd as u128) << 1) | 1u128) {
Ok(data)
} else {
Err(QueueError::Empty)
}
}
}
pub type Producer<T, I = u32, const N: usize = 0> = QueueProducerHandle<T, I, N>;
pub type Consumer<T, I = u32, const N: usize = 0> = QueueConsumerHandle<T, I, N>;
#[derive(Debug)]
pub struct QueueProducerHandle<T, I = u32, const N: usize = 0>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
queue: Arc<MpmcQueue<T, I, N>>,
}
impl<T, I, const N: usize> Clone for QueueProducerHandle<T, I, N>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
fn clone(&self) -> Self {
Self {
queue: self.queue.clone(),
}
}
}
impl<T, I, const N: usize> QueueProducer<T> for QueueProducerHandle<T, I, N>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
fn try_push(&self, value: T) -> Result<(), (T, QueueError)> {
self.queue.try_push(value)
}
fn push(&self, value: T) -> Result<(), QueueError> {
self.queue.push(value)
}
fn push_with_seq(&self, value: T) -> Result<usize, QueueError> {
self.queue.push_impl(value, true)
}
}
#[derive(Debug)]
pub struct QueueConsumerHandle<T, I = u32, const N: usize = 0>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
queue: Arc<MpmcQueue<T, I, N>>,
}
impl<T, I, const N: usize> Clone for QueueConsumerHandle<T, I, N>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
fn clone(&self) -> Self {
Self {
queue: self.queue.clone(),
}
}
}
impl<T, I, const N: usize> QueueConsumer<T> for QueueConsumerHandle<T, I, N>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
fn try_pop(&self) -> Result<T, QueueError> {
self.queue.try_pop()
}
fn pop(&self) -> Result<T, QueueError> {
self.queue.pop()
}
fn pop_with_seq(&self) -> Result<(T, usize), QueueError> {
self.queue.pop_impl(true)
}
fn peek(&self) -> Result<T, QueueError> {
self.queue.peek()
}
fn peek_with_seq(&self) -> Result<(T, usize), QueueError> {
let rd = self.queue.read_index.load(Ordering::Acquire);
self.queue.peek().map(|value| (value, rd))
}
fn pop_if<F>(&self, mut predicate: F) -> Result<T, QueueError>
where
F: FnMut(&T, usize) -> bool,
{
loop {
let rd = self.queue.read_index.load(Ordering::SeqCst);
let idx = rd & self.queue.mask;
let packed = self.queue.storage.load(idx, Ordering::SeqCst);
let (seq, data) =
unpack_entry::<T, I>(packed, self.queue.seq_shift, self.queue.data_size);
if seq == ((rd as u128) << 1) {
return Err(QueueError::Empty);
}
if seq == (((rd as u128) << 1) | 1u128) {
if !predicate(&data, rd) {
return Err(QueueError::Empty); }
let new_seq = (((rd + self.queue.capacity) as u128) << 1)
& ((1u128 << self.queue.seq_shift) - 1u128);
let new_packed = pack_entry::<T, I>(new_seq, T::default(), self.queue.seq_shift);
if self
.queue
.storage
.compare_exchange(idx, packed, new_packed, Ordering::SeqCst, Ordering::SeqCst)
.is_ok()
{
let _ = self.queue.read_index.compare_exchange(
rd,
rd + 1,
Ordering::SeqCst,
Ordering::SeqCst,
);
return Ok(data);
}
} else if (seq >> 1) == ((rd + self.queue.capacity) as u128) {
let _ = self.queue.read_index.compare_exchange(
rd,
rd + 1,
Ordering::SeqCst,
Ordering::SeqCst,
);
}
}
}
fn consume<F>(&self, mut consumer: F) -> usize
where
F: FnMut(T, usize) -> bool,
{
let mut count = 0;
while let Ok((value, seq)) = self.pop_with_seq() {
count += 1;
if consumer(value, seq) {
break;
}
}
count
}
fn is_empty(&self) -> bool {
self.queue.is_empty()
}
fn size(&self) -> usize {
self.queue.len()
}
}
impl<T, I, const N: usize> QueueFactory<T> for Arc<MpmcQueue<T, I, N>>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
type Producer = QueueProducerHandle<T, I, N>;
type Consumer = QueueConsumerHandle<T, I, N>;
fn producer(&self) -> Self::Producer {
QueueProducerHandle {
queue: self.clone(),
}
}
fn consumer(&self) -> Self::Consumer {
QueueConsumerHandle {
queue: self.clone(),
}
}
}
unsafe impl<T, I, const N: usize> Send for MpmcQueue<T, I, N>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
}
unsafe impl<T, I, const N: usize> Sync for MpmcQueue<T, I, N>
where
T: Copy + Send + Sync + Default,
I: Copy + Into<u128>,
{
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn runtime_basic() {
let q = queue::<u32>().capacity(8).build().unwrap();
assert_eq!(q.capacity(), 8);
assert_eq!(q.len(), 0);
let (producer, consumer) = q.channel();
producer.push(10).unwrap();
assert_eq!(consumer.pop().unwrap(), 10);
}
#[test]
fn static_basic() {
let q = queue::<u32>().capacity(4).build().unwrap();
assert_eq!(q.capacity(), 4);
assert_eq!(q.len(), 0);
let (producer, consumer) = q.channel();
producer.push(7).unwrap();
assert_eq!(consumer.pop().unwrap(), 7);
}
#[test]
fn push_pop_wrap() {
let (producer, consumer) = queue::<u32>().capacity(8).channels().unwrap();
for i in 0..8 {
producer.push(i).unwrap();
}
assert!(matches!(producer.push(99), Err(QueueError::Full)));
for i in 0..8 {
let v = consumer.pop().unwrap();
assert_eq!(v, i);
}
assert!(matches!(consumer.pop(), Err(QueueError::Empty)));
}
#[test]
fn test_with_seq_operations() {
let (producer, consumer) = queue::<u32>().capacity(8).channels().unwrap();
let seq1 = producer.push_with_seq(100).unwrap();
let seq2 = producer.push_with_seq(200).unwrap();
assert_eq!(seq1, 0);
assert_eq!(seq2, 1);
let (val1, pop_seq1) = consumer.pop_with_seq().unwrap();
let (val2, pop_seq2) = consumer.pop_with_seq().unwrap();
assert_eq!(val1, 100);
assert_eq!(val2, 200);
assert_eq!(pop_seq1, 0);
assert_eq!(pop_seq2, 1);
}
#[test]
fn test_peek_operations() {
let (producer, consumer) = queue::<u32>().capacity(8).channels().unwrap();
assert!(consumer.peek().is_err());
producer.push(42).unwrap();
assert_eq!(consumer.peek().unwrap(), 42);
let (val, seq) = consumer.peek_with_seq().unwrap();
assert_eq!(val, 42);
assert_eq!(seq, 0);
assert_eq!(consumer.pop().unwrap(), 42);
assert!(consumer.peek().is_err());
}
#[test]
fn test_pop_if() {
let (producer, consumer) = queue::<u32>().capacity(8).channels().unwrap();
producer.push(10).unwrap();
producer.push(20).unwrap();
producer.push(30).unwrap();
assert!(consumer.pop_if(|&v, _seq| v > 15).is_err());
let val = consumer.pop_if(|&v, _seq| v > 5).unwrap();
assert_eq!(val, 10);
let val = consumer.pop_if(|&v, _seq| v > 15).unwrap();
assert_eq!(val, 20);
}
#[test]
fn test_consume() {
let (producer, consumer) = queue::<u32>().capacity(8).channels().unwrap();
for i in 0..5 {
producer.push(i).unwrap();
}
let mut consumed = Vec::new();
let count = consumer.consume(|val, _seq| {
consumed.push(val);
val == 2 });
assert_eq!(count, 3); assert_eq!(consumed, vec![0, 1, 2]);
assert_eq!(consumer.pop().unwrap(), 3);
assert_eq!(consumer.pop().unwrap(), 4);
assert!(consumer.is_empty());
}
#[test]
fn test_exchange() {
let (producer, consumer) = queue::<u32>().capacity(4).channels().unwrap();
producer.push(100).unwrap();
producer.push(200).unwrap();
assert!(producer.queue.exchange(0, 100, 150));
assert!(!producer.queue.exchange(0, 100, 160));
assert_eq!(consumer.pop().unwrap(), 150);
assert_eq!(consumer.pop().unwrap(), 200);
}
use crate::traits::{QueueConsumer, QueueFactory, QueueProducer};
use std::{
collections::HashSet,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
time::Instant,
};
use tokio::{
task,
time::{Duration, sleep},
};
#[tokio::test(flavor = "multi_thread", worker_threads = 8)]
async fn mpmc_stress_dynamic() {
let producers = 4usize;
let consumers = 4usize;
let items_per_producer = 100_000usize;
let capacity = 1024usize;
let total = producers * items_per_producer;
let (producer, consumer) = queue::<u64>().capacity(capacity).channels().unwrap();
let seen = Arc::new(tokio::sync::Mutex::new(HashSet::<u64>::with_capacity(
total,
)));
let consumed = Arc::new(AtomicUsize::new(0));
let mut consumer_handles = Vec::with_capacity(consumers);
for _ in 0..consumers {
let seen_cl = seen.clone();
let consumed_cl = consumed.clone();
let total_cl = total;
let consumer = consumer.clone();
let h = task::spawn(async move {
loop {
if consumed_cl.load(Ordering::SeqCst) >= total_cl {
break;
}
match consumer.pop() {
Ok(val) => {
let inserted = seen_cl.lock().await.insert(val);
assert!(inserted, "duplicate value observed: {val}");
consumed_cl.fetch_add(1, Ordering::SeqCst);
},
Err(QueueError::Empty) => {
task::yield_now().await;
},
Err(e) => {
panic!("unexpected queue error in consumer: {e:?}");
},
}
}
});
consumer_handles.push(h);
}
let mut producer_handles = Vec::with_capacity(producers);
let start = Instant::now();
for pid in 0..producers {
let producer = producer.clone();
let h = task::spawn(async move {
for i in 0..items_per_producer {
let val = ((pid as u64) << 32) | (i as u64);
loop {
match producer.push(val) {
Ok(()) => break,
Err(QueueError::Full) => {
task::yield_now().await;
},
Err(e) => {
panic!("unexpected queue error in producer: {e:?}");
},
}
}
}
});
producer_handles.push(h);
}
for h in producer_handles {
h.await.expect("producer join");
}
while consumed.load(Ordering::SeqCst) < total {
sleep(Duration::from_millis(1)).await;
}
for h in consumer_handles {
h.await.expect("consumer join");
}
let elapsed = start.elapsed();
let throughput = (total as f64) / elapsed.as_secs_f64();
let seen_len = { seen.lock().await.len() };
assert_eq!(seen_len, total, "expected all items consumed once");
println!(
"DYNAMIC test: producers={producers} consumers={consumers} items/producer={items_per_producer} capacity={capacity} => total={total} elapsed={elapsed:?} throughput={throughput:.0} ops/sec"
);
}
const CAP: usize = 1024usize;
#[tokio::test(flavor = "multi_thread", worker_threads = 8)]
async fn mpmc_stress_static() {
let producers = 4usize;
let consumers = 4usize;
let items_per_producer = 100_000usize;
let total = producers * items_per_producer;
let q = queue::<u64>().capacity(CAP).build().unwrap();
let seen = Arc::new(tokio::sync::Mutex::new(HashSet::<u64>::with_capacity(
total,
)));
let consumed = Arc::new(AtomicUsize::new(0));
let mut consumer_handles = Vec::with_capacity(consumers);
for _ in 0..consumers {
let q_cl = q.clone();
let seen_cl = seen.clone();
let consumed_cl = consumed.clone();
let total_cl = total;
let h = task::spawn(async move {
loop {
if consumed_cl.load(Ordering::SeqCst) >= total_cl {
break;
}
let (_, consumer) = q_cl.channel();
match consumer.pop() {
Ok(val) => {
let inserted = seen_cl.lock().await.insert(val);
assert!(inserted, "duplicate value observed: {val}");
consumed_cl.fetch_add(1, Ordering::SeqCst);
},
Err(QueueError::Empty) => {
task::yield_now().await;
},
Err(e) => {
panic!("unexpected queue error in consumer: {e:?}");
},
}
}
});
consumer_handles.push(h);
}
let mut producer_handles = Vec::with_capacity(producers);
let start = Instant::now();
for pid in 0..producers {
let q_cl = q.clone();
let h = task::spawn(async move {
for i in 0..items_per_producer {
let val = ((pid as u64) << 32) | (i as u64);
loop {
let (producer, _) = q_cl.channel();
match producer.push(val) {
Ok(()) => break,
Err(QueueError::Full) => {
task::yield_now().await;
},
Err(e) => {
panic!("unexpected queue error in producer: {e:?}");
},
}
}
}
});
producer_handles.push(h);
}
for h in producer_handles {
h.await.expect("producer join");
}
while consumed.load(Ordering::SeqCst) < total {
sleep(Duration::from_millis(1)).await;
}
for h in consumer_handles {
h.await.expect("consumer join");
}
let elapsed = start.elapsed();
let throughput = (total as f64) / elapsed.as_secs_f64();
let seen_len = { seen.lock().await.len() };
assert_eq!(seen_len, total, "expected all items consumed once");
println!(
"STATIC test: producers={producers} consumers={consumers} items/producer={items_per_producer} capacity={CAP} => total={total} elapsed={elapsed:?} throughput={throughput:.0} ops/sec"
);
}
}