use std::{
cell::UnsafeCell,
mem::ManuallyDrop,
ops::Not,
sync::{
Arc, Weak,
atomic::{AtomicBool, AtomicPtr, Ordering},
},
task::Waker,
};
use futures::task::AtomicWaker;
pub unsafe trait StoredInQueue: Sized {
fn queue_link(&self) -> &QueueLink<Self>;
}
pub struct QueueLink<T> {
next: AtomicPtr<T>,
in_queue: AtomicBool,
}
impl<T> Default for QueueLink<T> {
fn default() -> Self {
Self {
next: Default::default(),
in_queue: Default::default(),
}
}
}
pub struct Queue<T: StoredInQueue> {
waker: AtomicWaker,
head: AtomicPtr<T>,
tail: UnsafeCell<*mut T>,
stub_predecessor: UnsafeCell<*mut T>,
stub: T,
}
impl<T: StoredInQueue> Queue<T> {
fn enqueue(&self, node: Arc<T>) {
if node.queue_link().in_queue.swap(true, Ordering::AcqRel) {
return;
}
node.queue_link()
.next
.store(std::ptr::null_mut(), Ordering::Relaxed);
let node = Arc::into_raw(node);
unsafe {
self.head
.swap(node.cast_mut(), Ordering::AcqRel)
.as_ref_unchecked()
.queue_link()
.next
.store(node.cast_mut(), Ordering::Release);
}
self.waker.wake();
}
unsafe fn dequeue(&self) -> Option<Arc<T>> {
let mut tail = unsafe {
*self.tail.get().as_ref_unchecked()
};
let mut next = unsafe {
tail.as_ref_unchecked()
.queue_link()
.next
.load(Ordering::Acquire)
};
if std::ptr::eq(tail, &self.stub) {
if next.is_null() {
return None;
}
unsafe {
*self.tail.get().as_mut_unchecked() = next;
}
tail = next;
next = unsafe {
next.as_mut_unchecked()
.queue_link()
.next
.load(Ordering::Acquire)
};
}
if next.is_null().not() {
unsafe {
*self.tail.get().as_mut_unchecked() = next;
}
debug_assert_ne!(
tail.cast_const(),
std::ptr::from_ref(&self.stub),
"popping stub, this should never happen",
);
let elem = unsafe {
Arc::from_raw(tail)
};
elem.queue_link().in_queue.store(false, Ordering::Release);
return Some(elem);
}
if std::ptr::eq(tail, self.head.load(Ordering::Acquire)).not() {
return None;
}
self.stub
.queue_link()
.next
.store(std::ptr::null_mut(), Ordering::Relaxed);
let stub = std::ptr::from_ref(&self.stub).cast_mut();
let stub_predecessor = self.head.swap(stub, Ordering::AcqRel);
unsafe {
stub_predecessor
.as_ref_unchecked()
.queue_link()
.next
.store(stub, Ordering::Release);
}
unsafe {
*self.stub_predecessor.get().as_mut_unchecked() = stub_predecessor;
}
next = unsafe {
tail.as_ref_unchecked()
.queue_link()
.next
.load(Ordering::Acquire)
};
if next.is_null() {
None
} else {
unsafe {
*self.tail.get().as_mut_unchecked() = next;
}
debug_assert_ne!(
tail.cast_const(),
std::ptr::from_ref(&self.stub),
"popping stub, this should never happen",
);
let elem = unsafe { Arc::from_raw(tail) };
elem.queue_link().in_queue.store(false, Ordering::Release);
Some(elem)
}
}
}
impl<T: StoredInQueue> Drop for Queue<T> {
fn drop(&mut self) {
struct DestroyOnDrop<'a, T: StoredInQueue>(&'a Queue<T>);
impl<T: StoredInQueue> Drop for DestroyOnDrop<'_, T> {
fn drop(&mut self) {
let inner_guard = DestroyOnDrop(self.0);
while unsafe { self.0.dequeue().is_some() } {}
let _ = ManuallyDrop::new(inner_guard);
}
}
DestroyOnDrop(self);
}
}
unsafe impl<T: StoredInQueue + Send> Send for Queue<T> {}
unsafe impl<T: StoredInQueue + Send> Sync for Queue<T> {}
pub struct Receiver<T: StoredInQueue>(Arc<Queue<T>>);
impl<T: StoredInQueue> Receiver<T> {
pub fn new<F>(make_stub: F) -> Self
where
F: FnOnce(WeakSender<T>) -> T,
{
let queue = Arc::<Queue<T>>::new_cyclic(|weak| {
let stub = make_stub(WeakSender(weak.clone()));
Queue {
waker: Default::default(),
head: AtomicPtr::new(std::ptr::null_mut()),
tail: UnsafeCell::new(std::ptr::null_mut()),
stub_predecessor: UnsafeCell::new(std::ptr::null_mut()),
stub,
}
});
let stub = std::ptr::from_ref(&queue.stub).cast_mut();
unsafe {
*queue.head.as_ptr().as_mut_unchecked() = stub;
*queue.tail.get().as_mut_unchecked() = stub;
}
Self(queue)
}
pub fn register(&mut self, waker: &Waker) -> *const T {
self.0.waker.register(waker);
let last = self.0.head.load(Ordering::Acquire);
if std::ptr::eq(last, &self.0.stub).not() {
return last;
}
let tail = unsafe {
*self.0.tail.get().as_ref_unchecked()
};
if std::ptr::eq(tail, &self.0.stub) {
return std::ptr::null();
}
unsafe {
*self.0.stub_predecessor.get().as_ref_unchecked()
}
}
pub fn recv(&mut self) -> Option<Arc<T>> {
unsafe {
self.0.dequeue()
}
}
pub fn send(&self, value: Arc<T>) {
self.0.enqueue(value);
}
pub fn weak_sender(&self) -> WeakSender<T> {
WeakSender(Arc::downgrade(&self.0))
}
pub fn is_parent(&self, weak: &WeakSender<T>) -> bool {
std::ptr::eq(Arc::as_ptr(&self.0), weak.0.as_ptr())
}
}
pub struct Sender<T: StoredInQueue>(Arc<Queue<T>>);
impl<T: StoredInQueue> Sender<T> {
pub fn send(&self, value: Arc<T>) {
self.0.enqueue(value);
}
}
pub struct WeakSender<T: StoredInQueue>(Weak<Queue<T>>);
impl<T: StoredInQueue> WeakSender<T> {
pub fn upgrade(&self) -> Option<Sender<T>> {
self.0.upgrade().map(Sender)
}
}
impl<T: StoredInQueue> Clone for WeakSender<T> {
fn clone(&self) -> Self {
Self(self.0.clone())
}
}
#[cfg(test)]
mod test {
use std::{
sync::{Arc, Barrier, atomic::Ordering},
task::Waker,
time::{Duration, Instant},
};
use rstest::rstest;
use crate::mpsc::{QueueLink, Receiver, StoredInQueue};
#[derive(Default)]
struct TestNode {
sender: usize,
link: QueueLink<Self>,
}
unsafe impl StoredInQueue for TestNode {
fn queue_link(&self) -> &QueueLink<Self> {
&self.link
}
}
#[rstest]
#[test]
fn concurrent_senders(#[values(1, 2, 4)] senders: usize) {
const SENT_BY_EACH: usize = 64 * 1024;
let barrier = Arc::new(Barrier::new(senders + 1));
let mut receiver = Receiver::new(|_| TestNode::default());
for i in 0..senders {
let barrier = barrier.clone();
let sender = receiver.weak_sender().upgrade().unwrap();
std::thread::spawn(move || {
let node = Arc::new(TestNode {
sender: i,
link: Default::default(),
});
let mut remaining = SENT_BY_EACH;
barrier.wait();
while remaining > 0 {
if node.queue_link().in_queue.load(Ordering::Acquire) {
std::hint::spin_loop();
continue;
}
sender.send(node.clone());
remaining -= 1;
}
});
}
let mut expected = vec![SENT_BY_EACH; senders];
let mut total_expected = SENT_BY_EACH * senders;
barrier.wait();
let start = Instant::now();
while total_expected > 0 {
let Some(node) = receiver.recv() else {
if start.elapsed() > Duration::from_secs(10) {
panic!("{expected:?}");
}
std::hint::spin_loop();
continue;
};
expected[node.sender] = expected[node.sender].checked_sub(1).unwrap();
total_expected = total_expected.checked_sub(1).unwrap();
}
assert_eq!(expected, vec![0; senders],);
}
#[test]
fn queue_cleanup() {
let receiver = Receiver::new(|_| TestNode {
sender: 0,
link: Default::default(),
});
let nodes = (0..10)
.map(|_| TestNode::default())
.map(Arc::new)
.map(|node| {
let weak = Arc::downgrade(&node);
receiver.send(node);
weak
})
.collect::<Vec<_>>();
drop(receiver);
nodes.iter().for_each(|node| {
assert!(node.upgrade().is_none());
});
}
#[test]
fn marker() {
let mut receiver = Receiver::new(|_| TestNode::default());
assert_eq!(receiver.register(Waker::noop()), std::ptr::null());
let node = Arc::new(TestNode::default());
receiver.send(node.clone());
assert_eq!(receiver.register(Waker::noop()), Arc::as_ptr(&node));
}
}