use crate::doorbell::Doorbell;
use crate::waiters::{Ticket, Waiters};
use grommet_core::ring;
use std::sync::Arc;
use std::task::{Context, Poll};
#[cfg(loom)]
use loom::sync::atomic::{AtomicBool, AtomicUsize, Ordering, fence};
#[cfg(not(loom))]
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering, fence};
struct Shared<W> {
ring: ring::Producer<W>,
bell: Doorbell,
waiters: Waiters,
parked: AtomicBool,
producers: AtomicUsize,
departed: AtomicBool,
}
pub fn channel<W>(capacity: usize) -> (Mailbox<W>, Inbox<W>) {
assert!(capacity > 0, "a mailbox needs capacity");
let (producer, consumer) = ring::bounded(capacity);
let shared = Arc::new(Shared {
ring: producer,
bell: Doorbell::new(),
waiters: Waiters::new(),
parked: AtomicBool::new(false),
producers: AtomicUsize::new(1),
departed: AtomicBool::new(false),
});
(Mailbox { shared: Arc::clone(&shared) }, Inbox { ring: consumer, shared, freed: 0 })
}
pub struct Mailbox<W> {
shared: Arc<Shared<W>>,
}
impl<W> Clone for Mailbox<W> {
fn clone(&self) -> Self {
self.shared.producers.fetch_add(1, Ordering::Relaxed);
Self { shared: Arc::clone(&self.shared) }
}
}
impl<W> Drop for Mailbox<W> {
fn drop(&mut self) {
if self.shared.producers.fetch_sub(1, Ordering::Release) == 1 {
self.shared.bell.ring();
}
}
}
impl<W> std::fmt::Debug for Mailbox<W> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Mailbox").field("capacity", &self.shared.ring.capacity()).finish()
}
}
impl<W> Mailbox<W> {
pub async fn send(&self, item: W) -> Result<(), Closed<W>> {
let full = match self.try_send(item) {
Ok(()) => return Ok(()),
Err(TrySendError::Closed(item)) => return Err(Closed(item)),
Err(TrySendError::Full(item)) => item,
};
let shared = &*self.shared;
let mut item = Some(full);
let mut registration = Registration { shared, ticket: None };
std::future::poll_fn(move |cx| {
let value = item.take().expect("the item is put back on every pending path");
match registration.ticket {
Some(ticket) => shared.waiters.refresh(ticket, cx.waker()),
None => match shared.waiters.park(cx.waker()) {
Ok(ticket) => registration.ticket = Some(ticket),
Err(_) => return Poll::Ready(Err(Closed(value))),
},
}
fence(Ordering::SeqCst);
match self.try_send(value) {
Ok(()) => {
registration.finish();
Poll::Ready(Ok(()))
}
Err(TrySendError::Closed(value)) => {
registration.finish();
Poll::Ready(Err(Closed(value)))
}
Err(TrySendError::Full(value)) => {
item = Some(value);
Poll::Pending
}
}
})
.await
}
#[inline]
pub fn try_send(&self, item: W) -> Result<(), TrySendError<W>> {
let sent = self.try_send_deferred(item);
if sent.is_ok() {
self.announce();
}
sent
}
#[inline]
pub(crate) fn try_send_deferred(&self, item: W) -> Result<(), TrySendError<W>> {
let shared = &*self.shared;
if shared.departed.load(Ordering::Acquire) {
return Err(TrySendError::Closed(item));
}
match shared.ring.try_push(item) {
Ok(()) => Ok(()),
Err(item) if shared.departed.load(Ordering::Acquire) => Err(TrySendError::Closed(item)),
Err(item) => Err(TrySendError::Full(item)),
}
}
#[inline]
pub(crate) fn announce(&self) {
self.shared.wake_shard();
}
}
impl<W> Shared<W> {
#[inline]
fn wake_shard(&self) {
fence(Ordering::SeqCst);
if self.parked.load(Ordering::Relaxed) {
self.bell.ring();
}
}
}
struct Registration<'a, W> {
shared: &'a Shared<W>,
ticket: Option<Ticket>,
}
impl<W> Registration<'_, W> {
fn finish(&mut self) {
if let Some(ticket) = self.ticket.take() {
self.shared.waiters.cancel(ticket);
}
}
}
impl<W> Drop for Registration<'_, W> {
fn drop(&mut self) {
let Some(ticket) = self.ticket.take() else { return };
if !self.shared.waiters.cancel(ticket) {
self.shared.waiters.wake_one();
}
}
}
pub struct Inbox<W> {
ring: ring::Consumer<W>,
shared: Arc<Shared<W>>,
freed: usize,
}
impl<W> std::fmt::Debug for Inbox<W> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Inbox").finish_non_exhaustive()
}
}
impl<W> Drop for Inbox<W> {
fn drop(&mut self) {
self.shared.departed.store(true, Ordering::Release);
self.shared.waiters.close();
self.shared.bell.close();
}
}
impl<W> Inbox<W> {
#[inline]
pub async fn recv(&mut self) -> Option<W> {
std::future::poll_fn(|cx| self.poll_recv(cx)).await
}
pub(crate) fn poll_recv(&mut self, cx: &mut Context<'_>) -> Poll<Option<W>> {
if let Some(item) = self.take() {
return Poll::Ready(Some(item));
}
self.release_capacity();
let closed = self.shared.producers.load(Ordering::Acquire) == 0;
self.shared.bell.register(cx.waker());
self.shared.parked.store(true, Ordering::Relaxed);
fence(Ordering::SeqCst);
if let Some(item) = self.take() {
self.unpark();
self.release_capacity();
return Poll::Ready(Some(item));
}
if closed {
self.unpark();
return Poll::Ready(None);
}
Poll::Pending
}
pub fn try_recv(&mut self) -> Result<W, TryRecvError> {
let closed = self.shared.producers.load(Ordering::Acquire) == 0;
let item = self.take();
self.release_capacity();
match item {
Some(item) => Ok(item),
None if closed => Err(TryRecvError::Closed),
None => Err(TryRecvError::Empty),
}
}
#[inline]
fn take(&mut self) -> Option<W> {
let item = self.ring.pop()?;
self.unpark();
self.freed += 1;
Some(item)
}
pub(crate) fn release_capacity(&mut self) {
if self.freed == 0 {
return;
}
let freed = std::mem::take(&mut self.freed);
fence(Ordering::SeqCst);
if !self.shared.waiters.any() {
return;
}
for _ in 0..freed {
if !self.shared.waiters.wake_one() {
break;
}
}
}
#[inline]
fn unpark(&self) {
if self.shared.parked.load(Ordering::Relaxed) {
self.shared.parked.store(false, Ordering::Relaxed);
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Closed<W>(pub W);
impl<W> Closed<W> {
pub fn into_inner(self) -> W {
self.0
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum TrySendError<W> {
Full(W),
Closed(W),
}
impl<W> TrySendError<W> {
pub fn into_inner(self) -> W {
match self {
Self::Full(item) | Self::Closed(item) => item,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum TryRecvError {
Empty,
Closed,
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn items_arrive_in_submission_order() {
let (mailbox, mut inbox) = channel(4);
for item in 0..3 {
mailbox.send(item).await.unwrap();
}
assert_eq!(inbox.recv().await, Some(0));
assert_eq!(inbox.recv().await, Some(1));
assert_eq!(inbox.recv().await, Some(2));
}
#[tokio::test]
async fn a_full_mailbox_hands_the_item_back_rather_than_dropping_it() {
let (mailbox, mut inbox) = channel(1);
mailbox.try_send(1).unwrap();
assert_eq!(
mailbox.try_send(2),
Err(TrySendError::Full(2)),
"the caller gets its work back"
);
assert_eq!(inbox.recv().await, Some(1));
mailbox.try_send(2).unwrap();
assert_eq!(inbox.recv().await, Some(2));
}
#[tokio::test]
async fn a_departed_shard_is_reported_by_both_submission_paths() {
let (mailbox, inbox) = channel::<u8>(4);
drop(inbox);
assert_eq!(mailbox.try_send(1), Err(TrySendError::Closed(1)));
assert_eq!(mailbox.send(2).await, Err(Closed(2)));
assert_eq!(Closed(3).into_inner(), 3);
assert_eq!(TrySendError::Full(4).into_inner(), 4);
}
#[tokio::test]
async fn an_immediate_receive_separates_an_empty_queue_from_a_closed_one() {
let (mailbox, mut inbox) = channel(4);
assert_eq!(inbox.try_recv(), Err(TryRecvError::Empty), "senders remain, so this is a lull");
mailbox.send(9).await.unwrap();
assert_eq!(inbox.try_recv(), Ok(9));
drop(mailbox);
assert_eq!(inbox.try_recv(), Err(TryRecvError::Closed));
assert_eq!(inbox.recv().await, None);
}
#[tokio::test]
async fn a_drained_mailbox_still_delivers_what_was_already_queued() {
let (mailbox, mut inbox) = channel(4);
mailbox.send(1).await.unwrap();
mailbox.send(2).await.unwrap();
drop(mailbox);
assert_eq!(inbox.try_recv(), Ok(1));
assert_eq!(inbox.recv().await, Some(2));
assert_eq!(inbox.recv().await, None);
}
#[test]
#[should_panic(expected = "a mailbox needs capacity")]
fn a_zero_capacity_mailbox_is_refused() {
let _ = channel::<u8>(0);
}
#[test]
fn a_submitting_handle_is_one_pointer_wide() {
assert_eq!(std::mem::size_of::<Mailbox<u64>>(), std::mem::size_of::<usize>());
}
}
#[cfg(all(test, not(loom)))]
mod backpressure_tests {
use super::*;
use std::future::Future;
use std::pin::Pin;
use std::time::Duration;
async fn before_long<F: Future>(future: F) -> F::Output {
tokio::time::timeout(Duration::from_secs(5), future)
.await
.expect("a sender was never woken")
}
async fn poll_once<F: Future>(future: &mut Pin<Box<F>>) {
std::future::poll_fn(|cx| {
let _ = future.as_mut().poll(cx);
Poll::Ready(())
})
.await;
}
#[tokio::test]
async fn a_full_mailbox_suspends_the_sender_until_the_shard_drains() {
let (mailbox, mut inbox) = channel(1);
mailbox.send(1).await.unwrap();
let sender = mailbox.clone();
let parked = tokio::spawn(async move { sender.send(2).await });
tokio::task::yield_now().await;
assert_eq!(inbox.try_recv(), Ok(1));
before_long(parked).await.unwrap().unwrap();
assert_eq!(inbox.recv().await, Some(2));
}
#[tokio::test]
async fn waiting_senders_are_admitted_in_the_order_they_arrived() {
let (mailbox, mut inbox) = channel(1);
mailbox.send(0).await.unwrap();
let mut senders = Vec::new();
for item in 1..=3 {
let mailbox = mailbox.clone();
senders.push(tokio::spawn(async move { mailbox.send(item).await }));
tokio::task::yield_now().await;
}
let mut seen = Vec::new();
while seen.len() < 4 {
seen.push(before_long(inbox.recv()).await.expect("senders remain"));
}
for sender in senders {
before_long(sender).await.unwrap().unwrap();
}
assert_eq!(seen, [0, 1, 2, 3]);
}
#[tokio::test]
async fn abandoning_a_waiting_sender_leaves_the_queue_behind_it_intact() {
let (mailbox, mut inbox) = channel(1);
mailbox.send(0).await.unwrap();
let abandoned = mailbox.clone();
let mut abandoned = Box::pin(abandoned.send(1));
poll_once(&mut abandoned).await;
let next = mailbox.clone();
let waiting = tokio::spawn(async move { next.send(2).await });
tokio::task::yield_now().await;
drop(abandoned);
assert_eq!(inbox.try_recv(), Ok(0));
before_long(waiting).await.unwrap().unwrap();
assert_eq!(inbox.recv().await, Some(2));
}
#[tokio::test]
async fn a_sender_dropped_after_being_admitted_passes_its_slot_on() {
let (mailbox, mut inbox) = channel(1);
mailbox.send(0).await.unwrap();
let first = mailbox.clone();
let mut first = Box::pin(first.send(1));
poll_once(&mut first).await;
let second = mailbox.clone();
let waiting = tokio::spawn(async move { second.send(2).await });
tokio::task::yield_now().await;
assert_eq!(inbox.try_recv(), Ok(0), "this is the slot the first sender is offered");
drop(first);
before_long(waiting).await.unwrap().unwrap();
assert_eq!(inbox.recv().await, Some(2));
}
#[tokio::test]
async fn a_departing_shard_hands_every_waiting_sender_its_item_back() {
let (mailbox, inbox) = channel(1);
mailbox.send(0).await.unwrap();
let senders: Vec<_> = (1..=3)
.map(|item| {
let mailbox = mailbox.clone();
tokio::spawn(async move { mailbox.send(item).await })
})
.collect();
tokio::task::yield_now().await;
drop(inbox);
for (expected, sender) in (1..=3).zip(senders) {
let returned = before_long(sender).await.unwrap();
assert!(
matches!(returned, Err(Closed(item)) if item == expected) || returned == Ok(()),
"a parked sender neither sent nor got its item back: {returned:?}"
);
}
}
#[tokio::test]
async fn a_sender_that_arrives_after_the_shard_has_gone_is_told_so() {
let (mailbox, inbox) = channel::<u8>(1);
drop(inbox);
assert_eq!(mailbox.send(1).await, Err(Closed(1)));
}
}
#[cfg(all(test, loom))]
mod loom_tests {
use super::*;
use loom::sync::atomic::AtomicBool;
use std::sync::Arc as StdArc;
use std::task::{Wake, Waker};
struct Flag(AtomicBool);
impl Flag {
fn waker() -> (StdArc<Self>, Waker) {
let flag = StdArc::new(Self(AtomicBool::new(false)));
(StdArc::clone(&flag), Waker::from(flag))
}
fn woken(&self) -> bool {
self.0.load(Ordering::Acquire)
}
}
impl Wake for Flag {
fn wake(self: StdArc<Self>) {
self.0.store(true, Ordering::Release);
}
fn wake_by_ref(self: &StdArc<Self>) {
self.0.store(true, Ordering::Release);
}
}
#[test]
fn loom_a_shard_never_parks_on_an_item_already_pushed() {
loom::model(|| {
let (mailbox, mut inbox) = channel::<u32>(1);
let (flag, waker) = Flag::waker();
let sender = mailbox.clone();
let producer = loom::thread::spawn(move || sender.try_send(1).is_ok());
let polled = inbox.poll_recv(&mut Context::from_waker(&waker));
assert!(producer.join().unwrap(), "an empty mailbox refused a push");
match polled {
Poll::Ready(Some(item)) => assert_eq!(item, 1),
Poll::Pending => {
assert!(flag.woken(), "the shard parked on an item that was already there");
}
Poll::Ready(None) => panic!("a live sender was reported as gone"),
}
drop(mailbox);
});
}
#[test]
fn loom_a_sender_never_parks_on_room_already_freed() {
loom::model(|| {
let (mailbox, mut inbox) = channel::<u32>(1);
assert!(mailbox.try_send(1).is_ok(), "an empty mailbox refused a push");
let (flag, waker) = Flag::waker();
let drainer = loom::thread::spawn(move || {
let taken = inbox.try_recv();
(inbox, taken)
});
let mut send = Box::pin(mailbox.send(2));
let polled = send.as_mut().poll(&mut Context::from_waker(&waker));
let (inbox, taken) = drainer.join().unwrap();
assert_eq!(taken, Ok(1), "the queued item was not drained");
if polled.is_pending() {
assert!(flag.woken(), "a sender parked on room that was already free");
}
drop(send);
drop(inbox);
});
}
}
#[cfg(all(test, not(loom)))]
mod starvation_tests {
use super::*;
use std::future::Future;
use std::time::Duration;
async fn before_long<F: Future>(future: F) -> F::Output {
tokio::time::timeout(Duration::from_secs(5), future)
.await
.expect("a waiting sender was never admitted")
}
#[tokio::test]
async fn a_drain_that_stops_before_the_mailbox_empties_still_admits_a_waiting_sender() {
let (mailbox, mut inbox) = channel(2);
mailbox.try_send(1).unwrap();
mailbox.try_send(2).unwrap();
let sender = mailbox.clone();
let waiting = tokio::spawn(async move { sender.send(3).await });
tokio::task::yield_now().await;
let waker = futures::task::noop_waker();
let mut cx = Context::from_waker(&waker);
assert!(matches!(inbox.poll_recv(&mut cx), Poll::Ready(Some(1))));
inbox.release_capacity();
before_long(waiting).await.unwrap().unwrap();
assert_eq!(inbox.try_recv(), Ok(2));
assert_eq!(inbox.try_recv(), Ok(3));
}
}