use core::cell::UnsafeCell;
use core::sync::atomic::{AtomicUsize, Ordering};
use crate::cache::CacheAligned;
#[cfg(feature = "std")]
use std::boxed::Box;
#[cfg(not(feature = "std"))]
use alloc::boxed::Box;
const DEFAULT_QUEUE_CAPACITY: usize = 65536;
struct Slot<T> {
sequence: AtomicUsize,
data: UnsafeCell<Option<T>>,
}
#[repr(align(64))]
pub struct LockFreeQueue<T> {
buffer: Box<[Slot<T>]>,
mask: usize,
capacity: usize,
head: CacheAligned<AtomicUsize>,
tail: CacheAligned<AtomicUsize>,
}
unsafe impl<T: Send> Send for LockFreeQueue<T> {}
unsafe impl<T: Send> Sync for LockFreeQueue<T> {}
impl<T> LockFreeQueue<T> {
pub fn new() -> Self {
Self::with_capacity(DEFAULT_QUEUE_CAPACITY)
}
pub fn with_capacity(capacity: usize) -> Self {
assert!(capacity > 0, "Capacity must be greater than 0");
assert!(capacity.is_power_of_two(), "Capacity must be a power of 2");
#[cfg(feature = "std")]
let buffer: Box<[Slot<T>]> = (0..capacity)
.map(|i| Slot {
sequence: AtomicUsize::new(i),
data: UnsafeCell::new(None),
})
.collect::<std::vec::Vec<_>>()
.into_boxed_slice();
#[cfg(not(feature = "std"))]
let buffer: Box<[Slot<T>]> = (0..capacity)
.map(|i| Slot {
sequence: AtomicUsize::new(i),
data: UnsafeCell::new(None),
})
.collect::<alloc::vec::Vec<_>>()
.into_boxed_slice();
Self {
buffer,
mask: capacity - 1,
capacity,
head: CacheAligned::new(AtomicUsize::new(0)),
tail: CacheAligned::new(AtomicUsize::new(0)),
}
}
#[inline]
pub fn try_enqueue(&self, item: T) -> Result<(), T> {
let mut pos = self.tail.load(Ordering::Relaxed);
loop {
let slot = &self.buffer[pos & self.mask];
let seq = slot.sequence.load(Ordering::Acquire);
let diff = seq.wrapping_sub(pos) as isize;
if diff == 0 {
match self.tail.compare_exchange_weak(
pos,
pos.wrapping_add(1),
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => {
unsafe {
*slot.data.get() = Some(item);
}
slot.sequence.store(pos.wrapping_add(1), Ordering::Release);
return Ok(());
}
Err(actual) => pos = actual,
}
} else if diff < 0 {
return Err(item);
} else {
pos = self.tail.load(Ordering::Relaxed);
}
}
}
#[inline]
pub fn enqueue(&self, item: T) {
let mut backoff: usize = 1;
let mut item = Some(item);
loop {
match self.try_enqueue(item.take().expect("invariant: item present")) {
Ok(()) => return,
Err(returned) => {
item = Some(returned);
for _ in 0..backoff {
core::hint::spin_loop();
}
if backoff < 64 {
backoff = backoff.saturating_mul(2);
} else {
#[cfg(feature = "std")]
{
std::thread::yield_now();
}
backoff = 1;
}
}
}
}
}
#[inline]
pub fn try_dequeue(&self) -> Option<T> {
let mut pos = self.head.load(Ordering::Relaxed);
loop {
let slot = &self.buffer[pos & self.mask];
let seq = slot.sequence.load(Ordering::Acquire);
let diff = seq.wrapping_sub(pos.wrapping_add(1)) as isize;
if diff == 0 {
match self.head.compare_exchange_weak(
pos,
pos.wrapping_add(1),
Ordering::Relaxed,
Ordering::Relaxed,
) {
Ok(_) => {
let item = unsafe { (*slot.data.get()).take() };
slot.sequence
.store(pos.wrapping_add(self.capacity), Ordering::Release);
return item;
}
Err(actual) => pos = actual,
}
} else if diff < 0 {
return None;
} else {
pos = self.head.load(Ordering::Relaxed);
}
}
}
pub fn is_empty(&self) -> bool {
let head = self.head.load(Ordering::Relaxed);
let tail = self.tail.load(Ordering::Relaxed);
head == tail
}
pub const fn capacity(&self) -> usize {
self.capacity
}
}
impl<T> Default for LockFreeQueue<T> {
fn default() -> Self {
Self::new()
}
}
impl<T> Drop for LockFreeQueue<T> {
fn drop(&mut self) {
while self.try_dequeue().is_some() {}
}
}
#[cfg(test)]
mod tests {
use super::*;
use core::sync::atomic::AtomicUsize;
#[cfg(feature = "std")]
use std::sync::Arc;
#[cfg(not(feature = "std"))]
use alloc::sync::Arc;
#[test]
fn test_lock_free_queue_basic() {
let queue = LockFreeQueue::<i32>::with_capacity(4);
assert!(queue.is_empty());
queue.enqueue(1);
queue.enqueue(2);
assert!(!queue.is_empty());
assert_eq!(queue.try_dequeue(), Some(1));
assert_eq!(queue.try_dequeue(), Some(2));
assert_eq!(queue.try_dequeue(), None);
assert!(queue.is_empty());
}
#[test]
fn test_lock_free_queue_wrap_around() {
let queue = LockFreeQueue::<i32>::with_capacity(4);
for round in 0..16 {
for i in 0..3 {
queue.enqueue(round * 3 + i);
}
for i in 0..3 {
assert_eq!(
queue.try_dequeue(),
Some(round * 3 + i),
"round {round}, item {i}"
);
}
assert!(queue.try_dequeue().is_none(), "round {round} not empty");
}
}
#[test]
fn test_lock_free_queue_full_try_enqueue() {
let queue = LockFreeQueue::<i32>::with_capacity(4);
for i in 0..4 {
queue.try_enqueue(i).unwrap();
}
assert!(queue.try_enqueue(99).is_err());
assert_eq!(queue.try_dequeue(), Some(0));
queue.try_enqueue(99).unwrap();
}
#[test]
fn test_lock_free_queue_drop_runs_destructors() {
struct DropCounter {
counter: Arc<AtomicUsize>,
}
impl Drop for DropCounter {
fn drop(&mut self) {
self.counter.fetch_add(1, Ordering::Relaxed);
}
}
let counter = Arc::new(AtomicUsize::new(0));
{
let queue = LockFreeQueue::<DropCounter>::with_capacity(4);
for _ in 0..3 {
queue.enqueue(DropCounter {
counter: Arc::clone(&counter),
});
}
}
assert_eq!(counter.load(Ordering::Relaxed), 3);
}
#[cfg(feature = "std")]
#[test]
fn test_lock_free_queue_concurrent_mpmc() {
use std::thread;
let queue = Arc::new(LockFreeQueue::<i32>::with_capacity(1024));
let num_producers = 4;
let num_consumers = 4;
let items_per_producer = 1000;
let total_items = num_producers * items_per_producer;
let mut handles = Vec::new();
for p in 0..num_producers {
let q = Arc::clone(&queue);
handles.push(thread::spawn(move || {
for i in 0..items_per_producer {
q.enqueue((p * items_per_producer + i) as i32);
}
}));
}
let consumed = Arc::new(AtomicUsize::new(0));
for _ in 0..num_consumers {
let q = Arc::clone(&queue);
let c = Arc::clone(&consumed);
handles.push(thread::spawn(move || {
while c.load(Ordering::Relaxed) < total_items {
if q.try_dequeue().is_some() {
c.fetch_add(1, Ordering::Relaxed);
} else {
std::thread::yield_now();
}
}
}));
}
for h in handles {
h.join().unwrap();
}
assert_eq!(consumed.load(Ordering::Relaxed), total_items);
}
}