use super::wait_queue::WaitQueue;
use crate::lockfree::{BoundedMpmcQueue, MpmcStack};
use std::cell::RefCell;
use std::collections::VecDeque;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::task::{Context, Poll};
#[repr(align(64))]
enum Queue<T> {
Bounded(Box<BoundedMpmcQueue<T>>),
Unbounded(MpmcStack<T>),
}
impl<T> Queue<T> {
#[inline(always)]
fn try_push(&self, value: T) -> Result<(), T> {
match self {
Self::Bounded(q) => q.try_push(value),
Self::Unbounded(q) => {
q.push(value);
Ok(())
}
}
}
}
#[repr(align(64))]
struct Shared<T> {
queue: Queue<T>,
sender_count: AtomicUsize,
receiver_dropped: AtomicBool,
recv_wait: WaitQueue,
send_wait: WaitQueue,
}
impl<T> Shared<T> {
#[inline(always)]
fn is_closed_for_send(&self) -> bool {
self.receiver_dropped.load(Ordering::Acquire)
}
#[inline(always)]
fn is_closed_for_recv(&self) -> bool {
self.sender_count.load(Ordering::Acquire) == 0
}
}
#[must_use]
#[inline]
pub fn channel<T>(capacity: usize) -> (Sender<T>, Receiver<T>) {
let shared = Arc::new(Shared {
queue: Queue::Bounded(Box::new(BoundedMpmcQueue::new(capacity.max(1)))),
sender_count: AtomicUsize::new(1),
receiver_dropped: AtomicBool::new(false),
recv_wait: WaitQueue::new(),
send_wait: WaitQueue::new(),
});
(
Sender {
shared: shared.clone(),
},
Receiver {
shared,
local: RefCell::new(VecDeque::new()),
},
)
}
#[must_use]
#[inline]
pub fn unbounded_channel<T>() -> (UnboundedSender<T>, UnboundedReceiver<T>) {
let shared = Arc::new(Shared {
queue: Queue::Unbounded(MpmcStack::new()),
sender_count: AtomicUsize::new(1),
receiver_dropped: AtomicBool::new(false),
recv_wait: WaitQueue::new(),
send_wait: WaitQueue::new(),
});
(
UnboundedSender {
inner: Sender {
shared: shared.clone(),
},
},
UnboundedReceiver {
inner: Receiver {
shared,
local: RefCell::new(VecDeque::new()),
},
},
)
}
#[repr(align(64))]
pub struct Sender<T> {
shared: Arc<Shared<T>>,
}
impl<T> Clone for Sender<T> {
#[inline(always)]
fn clone(&self) -> Self {
self.shared.sender_count.fetch_add(1, Ordering::AcqRel);
Self {
shared: self.shared.clone(),
}
}
}
impl<T> Drop for Sender<T> {
#[inline(always)]
fn drop(&mut self) {
if self.shared.sender_count.fetch_sub(1, Ordering::AcqRel) == 1 {
self.shared.recv_wait.wake_all();
}
}
}
impl<T> Sender<T> {
#[inline(always)]
pub async fn send(&self, value: T) -> Result<(), SendError<T>> {
let mut value = Some(value);
std::future::poll_fn(|cx| self.poll_send(cx, &mut value)).await
}
fn poll_send(&self, cx: &Context<'_>, value: &mut Option<T>) -> Poll<Result<(), SendError<T>>> {
if self.shared.is_closed_for_send() {
return Poll::Ready(Err(SendError(value.take().expect("value present"))));
}
if !self.shared.send_wait.has_waiters() {
match self
.shared
.queue
.try_push(value.take().expect("value present"))
{
Ok(()) => {
if self.shared.recv_wait.has_waiters() {
self.shared.recv_wait.wake_one();
}
return Poll::Ready(Ok(()));
}
Err(v) => *value = Some(v),
}
}
let token = self.shared.send_wait.register(cx.waker());
if self.shared.is_closed_for_send() {
self.shared.send_wait.cancel(token);
return Poll::Ready(Err(SendError(value.take().expect("value present"))));
}
match self
.shared
.queue
.try_push(value.take().expect("value present"))
{
Ok(()) => {
self.shared.send_wait.cancel(token);
if self.shared.recv_wait.has_waiters() {
self.shared.recv_wait.wake_one();
}
Poll::Ready(Ok(()))
}
Err(v) => {
*value = Some(v);
Poll::Pending
}
}
}
#[must_use]
#[inline(always)]
pub fn is_closed(&self) -> bool {
self.shared.is_closed_for_send()
}
}
#[repr(align(64))]
pub struct Receiver<T> {
shared: Arc<Shared<T>>,
local: RefCell<VecDeque<T>>,
}
unsafe impl<T: Send> Sync for Receiver<T> {}
impl<T> Drop for Receiver<T> {
#[inline(always)]
fn drop(&mut self) {
self.shared.receiver_dropped.store(true, Ordering::Release);
self.shared.send_wait.wake_all();
}
}
impl<T> Receiver<T> {
#[inline(always)]
pub async fn recv(&mut self) -> Option<T> {
std::future::poll_fn(|cx| self.poll_recv(cx)).await
}
#[inline]
fn poll_recv(&self, cx: &Context<'_>) -> Poll<Option<T>> {
if let Some(v) = self.try_pop() {
return Poll::Ready(Some(v));
}
if self.shared.is_closed_for_recv() {
return Poll::Ready(self.try_pop());
}
let token = self.shared.recv_wait.register(cx.waker());
if let Some(v) = self.try_pop() {
self.shared.recv_wait.cancel(token);
return Poll::Ready(Some(v));
}
if self.shared.is_closed_for_recv() {
let v = self.try_pop();
self.shared.recv_wait.cancel(token);
return Poll::Ready(v);
}
Poll::Pending
}
#[inline]
fn try_pop(&self) -> Option<T> {
let v = match &self.shared.queue {
Queue::Bounded(q) => q.try_pop(),
Queue::Unbounded(q) => {
let mut local = self.local.borrow_mut();
local.pop_front().map_or_else(
|| {
q.drain_into_vec_deque(&mut local);
local.pop_front()
},
Some,
)
}
};
if v.is_some() && self.shared.send_wait.has_waiters() {
self.shared.send_wait.wake_one();
}
v
}
}
#[repr(align(64))]
pub struct UnboundedSender<T> {
inner: Sender<T>,
}
impl<T> Clone for UnboundedSender<T> {
#[inline(always)]
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
}
}
}
impl<T> UnboundedSender<T> {
#[inline(always)]
pub fn send(&self, value: T) -> Result<(), SendError<T>> {
if self.inner.shared.is_closed_for_send() {
return Err(SendError(value));
}
self.inner
.shared
.queue
.try_push(value)
.unwrap_or_else(|_| unreachable!("unbounded queue variant never rejects a push"));
self.inner.shared.recv_wait.wake_one();
Ok(())
}
#[must_use]
#[inline(always)]
pub fn is_closed(&self) -> bool {
self.inner.is_closed()
}
}
#[repr(align(64))]
pub struct UnboundedReceiver<T> {
inner: Receiver<T>,
}
impl<T> UnboundedReceiver<T> {
#[inline(always)]
pub async fn recv(&mut self) -> Option<T> {
self.inner.recv().await
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(align(64))]
pub struct SendError<T>(pub T);
impl<T> std::fmt::Display for SendError<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("channel closed: receiver dropped")
}
}
impl<T: std::fmt::Debug> std::error::Error for SendError<T> {}