#![warn(missing_docs)]
mod queue;
mod sync;
pub mod mpsc;
use crate::queue::{Backoff, Queue};
use crate::sync::Ordering::{AcqRel, Acquire, Relaxed, Release, SeqCst};
use crate::sync::{Arc, AtomicBool, AtomicU64, AtomicUsize, CachePadded, Mutex, MutexGuard};
use std::collections::VecDeque;
use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::task::{Context, Poll, Waker};
#[derive(PartialEq, Eq, Clone, Copy)]
pub struct SendError<T>(pub T);
impl<T> SendError<T> {
#[inline]
pub fn into_inner(self) -> T {
self.0
}
}
impl<T> fmt::Debug for SendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SendError").finish()
}
}
impl<T> fmt::Display for SendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "sending on a closed channel")
}
}
impl<T> std::error::Error for SendError<T> {}
#[derive(PartialEq, Eq, Clone, Copy, Debug)]
pub struct RecvError;
impl fmt::Display for RecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "receiving on an empty and closed channel")
}
}
impl std::error::Error for RecvError {}
#[derive(PartialEq, Eq, Clone, Copy)]
pub enum TrySendError<T> {
Full(T),
Closed(T),
}
impl<T> TrySendError<T> {
#[inline]
pub fn into_inner(self) -> T {
match self {
TrySendError::Full(val) => val,
TrySendError::Closed(val) => val,
}
}
#[inline]
pub fn is_full(&self) -> bool {
matches!(self, TrySendError::Full(_))
}
#[inline]
pub fn is_closed(&self) -> bool {
matches!(self, TrySendError::Closed(_))
}
}
impl<T> fmt::Debug for TrySendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
TrySendError::Full(_) => f.debug_tuple("TrySendError::Full").finish(),
TrySendError::Closed(_) => f.debug_tuple("TrySendError::Closed").finish(),
}
}
}
impl<T> fmt::Display for TrySendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
TrySendError::Full(_) => write!(f, "sending on a full channel"),
TrySendError::Closed(_) => write!(f, "sending on a closed channel"),
}
}
}
impl<T> std::error::Error for TrySendError<T> {}
#[derive(PartialEq, Eq, Clone, Copy, Debug)]
pub enum TryRecvError {
Empty,
Closed,
}
impl TryRecvError {
#[inline]
pub fn is_empty(&self) -> bool {
matches!(self, TryRecvError::Empty)
}
#[inline]
pub fn is_closed(&self) -> bool {
matches!(self, TryRecvError::Closed)
}
}
impl fmt::Display for TryRecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
TryRecvError::Empty => write!(f, "receiving on an empty channel"),
TryRecvError::Closed => write!(f, "receiving on an empty and closed channel"),
}
}
}
impl std::error::Error for TryRecvError {}
struct Flags {
closed: AtomicBool,
waiting_receivers: AtomicUsize,
waiting_senders: AtomicUsize,
senders: AtomicUsize,
receivers: AtomicUsize,
next_waiter_id: AtomicU64,
}
struct WaiterList {
list: Mutex<VecDeque<(u64, Waker)>>,
}
impl WaiterList {
fn new() -> Self {
WaiterList {
list: Mutex::new(VecDeque::with_capacity(4)),
}
}
#[inline]
fn lock(&self) -> MutexGuard<'_, VecDeque<(u64, Waker)>> {
self.list.lock().unwrap_or_else(|e| e.into_inner())
}
#[cold]
fn register(
&self,
count: &AtomicUsize,
next_id: &AtomicU64,
id: &mut Option<u64>,
waker: &Waker,
) {
let mut list = self.lock();
if let Some(my) = *id {
if let Some(entry) = list.iter_mut().find(|(entry_id, _)| *entry_id == my) {
if !entry.1.will_wake(waker) {
entry.1 = waker.clone();
}
return;
}
}
let my = match *id {
Some(my) => my,
None => {
let my = next_id.fetch_add(1, Relaxed);
*id = Some(my);
my
}
};
list.push_back((my, waker.clone()));
count.fetch_add(1, Release);
}
#[cold]
fn unregister(&self, count: &AtomicUsize, id: &mut Option<u64>, forward: bool) {
if let Some(my) = id.take() {
if !forward && count.load(Acquire) == 0 {
return;
}
let mut list = self.lock();
if let Some(pos) = list.iter().position(|(entry_id, _)| *entry_id == my) {
list.remove(pos);
count.fetch_sub(1, SeqCst);
return;
}
drop(list);
if forward {
self.notify_one(count);
}
}
}
#[cold]
fn notify_one(&self, count: &AtomicUsize) {
let waker = {
let mut list = self.lock();
match list.pop_front() {
Some((_, waker)) => {
count.fetch_sub(1, SeqCst);
Some(waker)
}
None => None,
}
};
if let Some(waker) = waker {
waker.wake();
}
}
#[cold]
fn notify_all(&self, count: &AtomicUsize) {
let wakers: Vec<Waker> = {
let mut list = self.lock();
count.store(0, SeqCst);
list.drain(..).map(|(_, waker)| waker).collect()
};
for waker in wakers {
waker.wake();
}
}
#[cold]
fn notify_many(&self, count: &AtomicUsize, limit: usize) {
if limit == 1 {
self.notify_one(count);
return;
}
let wakers: Vec<Waker> = {
let mut list = self.lock();
let n = limit.min(list.len());
let mut wakers = Vec::with_capacity(n);
for _ in 0..n {
wakers.push(list.pop_front().unwrap().1);
}
count.fetch_sub(n, SeqCst);
wakers
};
for waker in wakers {
waker.wake();
}
}
}
struct Inner<T> {
queue: Queue<T>,
flags: CachePadded<Flags>,
recv_waiters: CachePadded<WaiterList>,
send_waiters: CachePadded<WaiterList>,
}
impl<T> Inner<T> {
fn new(capacity: Option<usize>) -> Self {
Inner {
queue: Queue::new(capacity),
flags: CachePadded(Flags {
closed: AtomicBool::new(false),
waiting_receivers: AtomicUsize::new(0),
waiting_senders: AtomicUsize::new(0),
senders: AtomicUsize::new(1),
receivers: AtomicUsize::new(1),
next_waiter_id: AtomicU64::new(1),
}),
recv_waiters: CachePadded(WaiterList::new()),
send_waiters: CachePadded(WaiterList::new()),
}
}
#[inline(always)]
fn is_closed(&self) -> bool {
self.flags.closed.load(Acquire)
}
fn close(&self) -> bool {
if self.flags.closed.swap(true, SeqCst) {
return false;
}
self.recv_waiters.notify_all(&self.flags.waiting_receivers);
self.send_waiters.notify_all(&self.flags.waiting_senders);
true
}
#[inline(always)]
fn try_send(&self, msg: T) -> Result<(), TrySendError<T>> {
if self.flags.closed.load(Relaxed) {
return Err(TrySendError::Closed(msg));
}
match self.queue.push(msg, &self.flags.waiting_receivers) {
Ok(wake) => {
if wake {
self.recv_waiters.notify_one(&self.flags.waiting_receivers);
}
Ok(())
}
Err(msg) => Err(TrySendError::Full(msg)),
}
}
#[inline(always)]
fn try_recv(&self) -> Result<T, TryRecvError> {
let mut wait = Backoff::new();
loop {
match self.try_recv_nonblocking() {
Ok(msg) => return Ok(msg),
Err(RecvState::Empty) => return Err(TryRecvError::Empty),
Err(RecvState::Closed) => return Err(TryRecvError::Closed),
Err(RecvState::Busy) => wait.snooze(),
}
}
}
#[inline(always)]
fn try_recv_nonblocking(&self) -> Result<T, RecvState> {
if let Some((msg, wake)) = self.queue.pop(&self.flags.waiting_senders) {
if wake {
self.send_waiters.notify_one(&self.flags.waiting_senders);
}
return Ok(msg);
}
Err(self.recv_state_when_empty())
}
#[cold]
fn recv_state_when_empty(&self) -> RecvState {
if !self.flags.closed.load(SeqCst) {
RecvState::Empty
} else if self.queue.is_empty_seqcst() {
RecvState::Closed
} else {
RecvState::Busy
}
}
}
enum RecvState {
Empty,
Closed,
Busy,
}
pub struct Sender<T> {
inner: Arc<Inner<T>>,
}
pub struct Receiver<T> {
inner: Arc<Inner<T>>,
}
unsafe impl<T: std::marker::Send> std::marker::Send for Sender<T> {}
unsafe impl<T: std::marker::Send> Sync for Sender<T> {}
unsafe impl<T: std::marker::Send> std::marker::Send for Receiver<T> {}
unsafe impl<T: std::marker::Send> Sync for Receiver<T> {}
impl<T> Clone for Sender<T> {
fn clone(&self) -> Self {
self.inner.flags.senders.fetch_add(1, Relaxed);
Sender {
inner: self.inner.clone(),
}
}
}
impl<T> Clone for Receiver<T> {
fn clone(&self) -> Self {
self.inner.flags.receivers.fetch_add(1, Relaxed);
Receiver {
inner: self.inner.clone(),
}
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
if self.inner.flags.senders.fetch_sub(1, AcqRel) == 1 {
self.inner.close();
}
}
}
impl<T> Drop for Receiver<T> {
fn drop(&mut self) {
if self.inner.flags.receivers.fetch_sub(1, AcqRel) == 1 {
self.inner.close();
}
}
}
impl<T> fmt::Debug for Sender<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Sender")
.field("len", &self.len())
.field("capacity", &self.capacity())
.field("is_closed", &self.is_closed())
.field("senders", &self.sender_count())
.field("receivers", &self.receiver_count())
.finish()
}
}
impl<T> fmt::Debug for Receiver<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Receiver")
.field("len", &self.len())
.field("capacity", &self.capacity())
.field("is_closed", &self.is_closed())
.field("senders", &self.sender_count())
.field("receivers", &self.receiver_count())
.finish()
}
}
pub fn unbounded<T>() -> (Sender<T>, Receiver<T>) {
let inner = Arc::new(Inner::new(None));
(
Sender {
inner: inner.clone(),
},
Receiver { inner },
)
}
pub fn bounded<T>(capacity: usize) -> (Sender<T>, Receiver<T>) {
assert!(capacity > 0, "capacity must be greater than zero");
let inner = Arc::new(Inner::new(Some(capacity)));
(
Sender {
inner: inner.clone(),
},
Receiver { inner },
)
}
impl<T> Sender<T> {
#[inline(always)]
pub fn try_send(&self, msg: T) -> Result<(), TrySendError<T>> {
self.inner.try_send(msg)
}
#[inline]
pub fn send(&self, msg: T) -> Send<'_, T> {
Send {
sender: self,
msg: Some(msg),
waiter_id: None,
}
}
#[inline]
pub fn close(&self) -> bool {
self.inner.close()
}
#[inline]
pub fn is_closed(&self) -> bool {
self.inner.is_closed()
}
#[inline]
pub fn len(&self) -> usize {
self.inner.queue.len()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[inline]
pub fn capacity(&self) -> Option<usize> {
self.inner.queue.capacity()
}
#[inline]
pub fn sender_count(&self) -> usize {
self.inner.flags.senders.load(Relaxed)
}
#[inline]
pub fn receiver_count(&self) -> usize {
self.inner.flags.receivers.load(Relaxed)
}
}
impl<T> Receiver<T> {
#[doc(hidden)]
pub fn __debug_dump(&self) -> String {
self.inner.queue.debug_dump()
}
#[inline(always)]
pub fn try_recv(&self) -> Result<T, TryRecvError> {
self.inner.try_recv()
}
#[inline]
pub fn recv(&self) -> Recv<'_, T> {
Recv {
receiver: self,
waiter_id: None,
}
}
#[inline]
pub fn close(&self) -> bool {
self.inner.close()
}
#[inline]
pub fn is_closed(&self) -> bool {
self.inner.is_closed()
}
#[inline]
pub fn len(&self) -> usize {
self.inner.queue.len()
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[inline]
pub fn capacity(&self) -> Option<usize> {
self.inner.queue.capacity()
}
#[inline]
pub fn sender_count(&self) -> usize {
self.inner.flags.senders.load(Relaxed)
}
#[inline]
pub fn receiver_count(&self) -> usize {
self.inner.flags.receivers.load(Relaxed)
}
}
pub struct Send<'a, T> {
sender: &'a Sender<T>,
msg: Option<T>,
waiter_id: Option<u64>,
}
impl<T> Send<'_, T> {
#[inline(always)]
fn unregister(&mut self, forward: bool) {
if self.waiter_id.is_some() {
let inner = &*self.sender.inner;
inner.send_waiters.unregister(
&inner.flags.waiting_senders,
&mut self.waiter_id,
forward,
);
}
}
}
impl<T> Drop for Send<'_, T> {
fn drop(&mut self) {
self.unregister(true);
}
}
impl<T> Unpin for Send<'_, T> {}
impl<T> Unpin for Recv<'_, T> {}
impl<T> Future for Send<'_, T> {
type Output = Result<(), SendError<T>>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
this.unregister(false);
let inner = &*this.sender.inner;
let msg = this.msg.take().expect("Send polled after completion");
let msg = match inner.try_send(msg) {
Ok(()) => {
return Poll::Ready(Ok(()));
}
Err(TrySendError::Closed(msg)) => {
return Poll::Ready(Err(SendError(msg)));
}
Err(TrySendError::Full(msg)) => msg,
};
inner.send_waiters.register(
&inner.flags.waiting_senders,
&inner.flags.next_waiter_id,
&mut this.waiter_id,
cx.waker(),
);
inner.queue.sender_parking();
match inner.try_send(msg) {
Ok(()) => {
this.unregister(true);
Poll::Ready(Ok(()))
}
Err(TrySendError::Closed(msg)) => {
this.unregister(false);
Poll::Ready(Err(SendError(msg)))
}
Err(TrySendError::Full(msg)) => {
this.msg = Some(msg);
Poll::Pending
}
}
}
}
pub struct Recv<'a, T> {
receiver: &'a Receiver<T>,
waiter_id: Option<u64>,
}
impl<T> Recv<'_, T> {
#[inline(always)]
fn unregister(&mut self, forward: bool) {
if self.waiter_id.is_some() {
let inner = &*self.receiver.inner;
inner.recv_waiters.unregister(
&inner.flags.waiting_receivers,
&mut self.waiter_id,
forward,
);
}
}
}
impl<T> Drop for Recv<'_, T> {
fn drop(&mut self) {
self.unregister(true);
}
}
impl<T> Future for Recv<'_, T> {
type Output = Result<T, RecvError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
this.unregister(false);
let inner = &*this.receiver.inner;
match inner.try_recv_nonblocking() {
Ok(msg) => {
return Poll::Ready(Ok(msg));
}
Err(RecvState::Closed) => {
return Poll::Ready(Err(RecvError));
}
Err(RecvState::Empty) | Err(RecvState::Busy) => {}
}
inner.recv_waiters.register(
&inner.flags.waiting_receivers,
&inner.flags.next_waiter_id,
&mut this.waiter_id,
cx.waker(),
);
let tail = inner.queue.receiver_parking();
let mut wait = Backoff::new();
loop {
let busy = match inner.try_recv_nonblocking() {
Ok(msg) => {
this.unregister(true);
return Poll::Ready(Ok(msg));
}
Err(RecvState::Closed) => {
this.unregister(false);
return Poll::Ready(Err(RecvError));
}
Err(RecvState::Busy) => true,
Err(RecvState::Empty) => inner.queue.claims_below(tail),
};
if !busy {
return Poll::Pending;
}
if !wait.snooze_bounded() {
cx.waker().wake_by_ref();
return Poll::Pending;
}
}
}
}