diff --git a/CHANGELOG.md b/CHANGELOG.md index e1bdd901..f5e01ff8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -37,7 +37,6 @@ All notable changes to this project will be documented in this file. * Complete semaphore permit releases and notify all eligible waiters even if a wake callback panics. * Release MPSC receiver wakers when the receiver is dropped, avoiding retained tasks and ownership cycles when a waker holds a sender. * Notify all blocked bounded MPSC senders on receiver disconnection even when a buffered message destructor panics. -* Avoid deadlocks when a bounded MPSC sender's waker clone callback receives from the same channel. ### Improvements @@ -45,7 +44,7 @@ All notable changes to this project will be documented in this file. * Finish releasing buffered bounded MPSC messages even if one message destructor panics. * Improve unbounded MPSC throughput with batched receiving and incremental storage reclamation; empty-buffer retention is bounded independently of previous peak occupancy. * Make completed and abandoned `Completion` waits lock-free while preserving cancellable pending registration. -* Avoid heap allocation when waking up to 32 waiters in Barrier, broadcast, condvar, event, MPSC, phaser, and watch notifications; larger waiter sets spill to a single heap allocation. +* Reduce allocations when notifying waiters in Barrier, broadcast, condvar, event, MPSC, phaser, and watch. ## v0.7.2 (2026-09-11) diff --git a/asyncband/src/event/manual_reset.rs b/asyncband/src/event/manual_reset.rs index 074dea4f..47722aa2 100644 --- a/asyncband/src/event/manual_reset.rs +++ b/asyncband/src/event/manual_reset.rs @@ -17,7 +17,6 @@ use std::fmt; use std::future::Future; -use std::mem; use std::pin::Pin; use std::sync::Arc; use std::task::Context; @@ -25,6 +24,7 @@ use std::task::Poll; use std::task::Waker; use crate::internal::mutex::Mutex; +use crate::internal::register_waker; use crate::internal::waitlist::WaitList; use crate::internal::waitlist::WaiterId; use crate::internal::wake_all; @@ -232,8 +232,11 @@ impl ManualResetEvent { "a linked waiter must belong to an unset event" ); let waiter = state.waiters.waiter_mut(id); - let retired = (!waiter.will_wake(cx.waker())) - .then(|| waiter.replace_waker(cx.waker().clone())); + assert!( + waiter.waker.is_some(), + "an unnotified waiter must retain its waker" + ); + let retired = register_waker(&mut waiter.waker, cx.waker()); (Poll::Pending, retired) } None if state.is_set => (Poll::Ready(()), None), @@ -297,23 +300,6 @@ struct Waiter { waker: Option, } -impl Waiter { - fn will_wake(&self, waker: &Waker) -> bool { - self.waker - .as_ref() - .expect("an unnotified waiter must retain its waker") - .will_wake(waker) - } - - fn replace_waker(&mut self, waker: Waker) -> Waker { - let current = self - .waker - .as_mut() - .expect("an unnotified waiter must retain its waker"); - mem::replace(current, waker) - } -} - #[must_use = "futures do nothing unless you `.await` or poll them"] struct Wait<'a> { event: &'a ManualResetEvent, diff --git a/asyncband/src/internal/mod.rs b/asyncband/src/internal/mod.rs index 076646d8..0b8211bc 100644 --- a/asyncband/src/internal/mod.rs +++ b/asyncband/src/internal/mod.rs @@ -19,6 +19,22 @@ use std::panic; use std::panic::AssertUnwindSafe; use std::task::Waker; +/// Retains the current task waker, returning any replaced registration for unlocked destruction. +#[inline] +#[must_use = "drop the replaced waker after releasing the state lock"] +// Some feature subsets have no primitive that registers a single waker slot. +#[allow(dead_code)] +pub fn register_waker(slot: &mut Option, waker: &Waker) -> Option { + if slot + .as_ref() + .is_some_and(|current| current.will_wake(waker)) + { + None + } else { + slot.replace(waker.clone()) + } +} + /// Wakes every waker while preserving the first panic. /// /// If a wake callback panics, the remaining callbacks are still attempted during unwinding. Any @@ -99,14 +115,13 @@ pub(crate) mod mutex; #[cfg(any( feature = "broadcast", - feature = "mpmc", feature = "mutex", feature = "rwlock", feature = "semaphore", ))] -// Broadcast and MPMC use waiter notifications; mutexes and rwlocks use acquire/release operations; -// the public semaphore also exposes permit accounting. Single-primitive builds leave part of this -// API unused. +// Broadcast uses waiter notifications; mutexes and rwlocks use acquire/release operations; the +// public semaphore also exposes permit accounting. Single-primitive builds leave part of this API +// unused. #[allow(dead_code)] pub(crate) mod semaphore; @@ -130,8 +145,6 @@ pub(crate) mod waitlist; feature = "event", feature = "completion", feature = "latch", - feature = "mpmc", - feature = "mpsc", feature = "mutex", feature = "once", feature = "phaser", diff --git a/asyncband/src/internal/semaphore.rs b/asyncband/src/internal/semaphore.rs index 704affc5..6c7cc4ae 100644 --- a/asyncband/src/internal/semaphore.rs +++ b/asyncband/src/internal/semaphore.rs @@ -34,6 +34,7 @@ use std::task::Poll; use std::task::Waker; use crate::internal::mutex::Mutex; +use crate::internal::register_waker; use crate::internal::waitlist::WaitList; use crate::internal::waitlist::WaiterId; use crate::internal::wake_all; @@ -148,7 +149,7 @@ impl Semaphore { } /// Adds `n` permits to the semaphore if there is any waiter. - #[cfg(any(feature = "broadcast", feature = "mpmc"))] + #[cfg(feature = "broadcast")] pub fn release_if_nonempty(&self, n: usize) { let waiters = self.waiters.lock(); if !waiters.is_empty() { @@ -157,7 +158,7 @@ impl Semaphore { } /// Adds as many permits until there is no waiter. - #[cfg(any(feature = "broadcast", feature = "mpmc"))] + #[cfg(feature = "broadcast")] pub fn notify_all(&self) { let mut waiters = self.waiters.lock(); let mut wakers = WakerBatch::new(); @@ -295,13 +296,7 @@ impl Acquire<'_> { let ready = { let node = waiters.waiter_mut(*idx); if node.permits > 0 { - let update_waker = node - .waker - .as_ref() - .is_none_or(|current| !current.will_wake(waker)); - if update_waker { - old_waker = node.waker.replace(waker.clone()); - } + old_waker = register_waker(&mut node.waker, waker); false } else { true diff --git a/asyncband/src/lib.rs b/asyncband/src/lib.rs index a68e63db..a0472fa7 100644 --- a/asyncband/src/lib.rs +++ b/asyncband/src/lib.rs @@ -54,32 +54,32 @@ //! //! # API map //! -//! | Area | API | Feature | Use | -//! |----------------------------|-----------------------------------------------|----------------|---------------------------------------------------------------------------------------------------------------| -//! | Locks and conditions | [`Mutex`](mutex::Mutex) | `mutex` | Protect shared data with asynchronous mutual exclusion. | -//! | | [`RwLock`](rwlock::RwLock) | `rwlock` | Allow multiple readers or one writer. | -//! | | [`Condvar`](condvar::Condvar) | `condvar` | Wait for notifications while releasing a mutex. | -//! | Coordination | [`Semaphore`](semaphore::Semaphore) | `semaphore` | Limit concurrent work by acquiring permits. | -//! | | [`Barrier`](barrier::Barrier) | `barrier` | Synchronize a fixed number of participants at a reusable rendezvous. | -//! | | [`ManualResetEvent`](event::ManualResetEvent) | `event` | Signal current and future waits until explicitly reset. | -//! | | [`AutoResetEvent`](event::AutoResetEvent) | `event` | Retain one signal and release one waiter per consumed signal. | -//! | | [`Latch`](latch::Latch) | `latch` | Wait until a fixed one-way countdown reaches zero. | -//! | | [`Phaser`](phaser::Phaser) | `phaser` | Coordinate repeated phases with a dynamic participant set. | -//! | | [`WaitGroup`](waitgroup::WaitGroup) | `waitgroup` | Dynamically register participants and wait until all have completed. | -//! | | [`Shutdown`](shutdown::Shutdown) | `shutdown` | Request shutdown and wait until all completion guards are dropped. | -//! | Work coalescing | [`Once`](once::Once) | `once` | Complete one asynchronous initialization; cancelled or panicked attempts may be retried. | -//! | | [`OnceCell`](once::OnceCell) | `once-cell` | Store one value from an access-time initializer; failed, cancelled, or panicked attempts may be retried. | -//! | | [`LazyCell`](once::LazyCell) | `lazy-cell` | Initialize one value with a stored function and resume the same in-flight future after caller cancellation. | -//! | | [`OnceMap`](once::OnceMap) | `once-map` | Coalesce work per key and retain each successful value until explicitly removed. | -//! | | [`Group`](singleflight::Group) | `singleflight` | Coalesce overlapping work per key without retaining completed values. | -//! | Communication | [`Completion`](completion::Completion) | `completion` | Publish one shared result to any number of current and future observers. | -//! | | [`oneshot`] | `oneshot` | Send one value from one sender to one receiver. | -//! | | [`mpmc`] | `mpmc` | Distribute each value to exactly one of multiple competing receivers. | -//! | | [`mpsc`] | `mpsc` | Send each value from multiple producers to one receiver with bounded backpressure or an unbounded queue. | -//! | | [`broadcast`] | `broadcast` | Deliver every value to active receivers with bounded backpressure or unbounded retention. | -//! | | [`watch`] | `watch` | Publish cloneable latest state from one or more senders; receivers independently coalesce intermediate updates. | -//! | Object reuse | [`pool`] | `pool` | Reuse objects through bounded or unbounded pool variants. | -//! | Sync interop | [`FutureExt`](blocking::FutureExt) | `blocking` | Drive one runtime-agnostic future from a blocking thread. | +//! | Area | API | Feature | Use | +//! |----------------------|-----------------------------------------------|----------------|-----------------------------------------------------------------------------------------------------------------| +//! | Locks and conditions | [`Mutex`](mutex::Mutex) | `mutex` | Protect shared data with asynchronous mutual exclusion. | +//! | | [`RwLock`](rwlock::RwLock) | `rwlock` | Allow multiple readers or one writer. | +//! | | [`Condvar`](condvar::Condvar) | `condvar` | Wait for notifications while releasing a mutex. | +//! | Coordination | [`Semaphore`](semaphore::Semaphore) | `semaphore` | Limit concurrent work by acquiring permits. | +//! | | [`Barrier`](barrier::Barrier) | `barrier` | Synchronize a fixed number of participants at a reusable rendezvous. | +//! | | [`ManualResetEvent`](event::ManualResetEvent) | `event` | Signal current and future waits until explicitly reset. | +//! | | [`AutoResetEvent`](event::AutoResetEvent) | `event` | Retain one signal and release one waiter per consumed signal. | +//! | | [`Latch`](latch::Latch) | `latch` | Wait until a fixed one-way countdown reaches zero. | +//! | | [`Phaser`](phaser::Phaser) | `phaser` | Coordinate repeated phases with a dynamic participant set. | +//! | | [`WaitGroup`](waitgroup::WaitGroup) | `waitgroup` | Dynamically register participants and wait until all have completed. | +//! | | [`Shutdown`](shutdown::Shutdown) | `shutdown` | Request shutdown and wait until all completion guards are dropped. | +//! | Work coalescing | [`Once`](once::Once) | `once` | Complete one asynchronous initialization; cancelled or panicked attempts may be retried. | +//! | | [`OnceCell`](once::OnceCell) | `once-cell` | Store one value from an access-time initializer; failed, cancelled, or panicked attempts may be retried. | +//! | | [`LazyCell`](once::LazyCell) | `lazy-cell` | Initialize one value with a stored function and resume the same in-flight future after caller cancellation. | +//! | | [`OnceMap`](once::OnceMap) | `once-map` | Coalesce work per key and retain each successful value until explicitly removed. | +//! | | [`Group`](singleflight::Group) | `singleflight` | Coalesce overlapping work per key without retaining completed values. | +//! | Communication | [`Completion`](completion::Completion) | `completion` | Publish one shared result to any number of current and future observers. | +//! | | [`oneshot`] | `oneshot` | Send one value from one sender to one receiver. | +//! | | [`mpmc`] | `mpmc` | Distribute each value to exactly one of multiple competing receivers. | +//! | | [`mpsc`] | `mpsc` | Send each value from multiple producers to one receiver with bounded backpressure or an unbounded queue. | +//! | | [`broadcast`] | `broadcast` | Deliver every value to active receivers with bounded backpressure or unbounded retention. | +//! | | [`watch`] | `watch` | Publish cloneable latest state from one or more senders; receivers independently coalesce intermediate updates. | +//! | Object reuse | [`pool`] | `pool` | Reuse objects through bounded or unbounded pool variants. | +//! | Sync interop | [`FutureExt`](blocking::FutureExt) | `blocking` | Drive one runtime-agnostic future from a blocking thread. | //! //! # Scope and runtime model //! diff --git a/asyncband/src/mpmc/bounded.rs b/asyncband/src/mpmc/bounded.rs index 3bcd73f9..3122f839 100644 --- a/asyncband/src/mpmc/bounded.rs +++ b/asyncband/src/mpmc/bounded.rs @@ -29,9 +29,8 @@ use super::queue::Shared; /// The queue stores at most `capacity` values. Sending waits for a receiver to free capacity when /// the queue is full. /// -/// Operations briefly acquire internal mutexes. No lock is held across an await point, while -/// waking tasks, or while dropping messages. The `try_*` methods do not wait for capacity or -/// messages, but may wait to acquire a mutex. +/// The `try_*` methods do not wait for capacity or messages, but may briefly block on an internal +/// mutex. /// /// # Panics /// @@ -83,10 +82,9 @@ impl BoundedSender { /// /// # Cancel safety /// - /// Dropping a pending `send` removes it from the wait queue and drops `value`; a call that has - /// returned `Pending` has not sent the value. Any selected capacity notification is passed to - /// the next waiting sender before `value` is dropped. Use [`try_send`](Self::try_send) when - /// the caller must retain ownership if capacity is unavailable. + /// Dropping a pending `send` drops `value` without sending it or retaining capacity. Use + /// [`try_send`](Self::try_send) when the caller must retain ownership if capacity is + /// unavailable. pub async fn send(&self, value: T) -> Result<(), SendError> { self.shared.send(value).await } @@ -137,8 +135,8 @@ impl BoundedReceiver { /// /// # Cancel safety /// - /// Dropping a pending `recv` does not consume a value. Any selected value notification is - /// passed to another waiting receiver, so cancellation does not prevent it from receiving. + /// Dropping a pending `recv` does not consume a value or prevent other receivers from receiving + /// it. pub async fn recv(&self) -> Result { self.shared.recv().await } diff --git a/asyncband/src/mpmc/queue.rs b/asyncband/src/mpmc/queue.rs index fe51bce6..a279470f 100644 --- a/asyncband/src/mpmc/queue.rs +++ b/asyncband/src/mpmc/queue.rs @@ -16,31 +16,107 @@ // under the License. use std::collections::VecDeque; -use std::future::Future; use std::future::poll_fn; -use std::pin::Pin; +use std::mem; use std::task::Context; use std::task::Poll; +use std::task::Waker; use super::RecvError; use super::SendError; use super::TryRecvError; use super::TrySendError; use crate::internal::mutex::Mutex; -use crate::internal::semaphore::Acquire; -use crate::internal::semaphore::Semaphore; +use crate::internal::waitlist::WaitList; +use crate::internal::waitlist::WaiterId; +use crate::internal::wake_all; -pub(super) struct Shared { +pub struct Shared { state: Mutex>, - recv_waiters: Semaphore, - send_waiters: Semaphore, - capacity: Option, } +/// Values, endpoint counts, and both waiter queues share one lock, so each transition and the +/// waiter it selects are decided together. Wake callbacks and waker or value destructors run +/// outside the lock, because they may reenter the queue. struct State { values: VecDeque, + capacity: Option, senders: usize, receivers: usize, + recv_waiters: WaitList, + send_waiters: WaitList, +} + +impl State { + fn has_capacity(&self) -> bool { + self.capacity + .is_none_or(|capacity| self.values.len() < capacity) + } + + /// Queues a value and selects the receiver to wake. + fn push(&mut self, value: T) -> Option { + self.values.push_back(value); + self.recv_waiters.notify_one() + } + + /// Takes the next value and selects the sender to wake. + fn pop(&mut self) -> Result<(T, Option), TryRecvError> { + if let Some(value) = self.values.pop_front() { + // Unbounded queues never block senders, so their sender queue is always empty. + Ok((value, self.send_waiters.notify_one())) + } else if self.senders == 0 { + Err(TryRecvError::Disconnected) + } else { + Err(TryRecvError::Empty) + } + } +} + +/// A pending receive or bounded send. +/// +/// Notification makes a waiter runnable; it does not reserve a value or slot. The detached node +/// remains owned by its future until it retries or is dropped. +enum Waiter { + Waiting(Waker), + Notified, +} + +impl WaitList { + fn notify_one(&mut self) -> Option { + let (_, waiter) = self.unlink_first_waiter(|_| true)?; + let Waiter::Waiting(waker) = mem::replace(waiter, Waiter::Notified) else { + unreachable!("only waiting operations remain linked"); + }; + Some(waker) + } + + fn remove_waiter(&mut self, id: WaiterId) -> Waiter { + // Unlinking is idempotent, so notified waiters are removed the same way as linked ones. + self.unlink_waiter(id, |_| true); + self.remove_unlinked_waiter(id) + } + + /// Queues a blocked operation or refreshes the waker of a queued one. + /// + /// A notified operation that still found no value or slot queues again at the back. + #[must_use = "drop the replaced waker after releasing the queue lock"] + fn register(&mut self, id: &mut Option, current: &Waker) -> Option { + if let Some(queued) = *id { + if let Waiter::Waiting(waker) = self.waiter_mut(queued) { + if waker.will_wake(current) { + return None; + } + return Some(mem::replace(waker, current.clone())); + } + } + let waker = current.clone(); + if let Some(notified) = id.take() { + // The notification already took this node's waker, so nothing is retired. + self.remove_waiter(notified); + } + *id = Some(self.push_back(Waiter::Waiting(waker))); + None + } } impl Shared { @@ -56,69 +132,69 @@ impl Shared { Self { state: Mutex::new(State { values: VecDeque::new(), + capacity, senders: 1, receivers: 1, + recv_waiters: WaitList::new(), + send_waiters: WaitList::new(), }), - recv_waiters: Semaphore::new(0), - send_waiters: Semaphore::new(0), - capacity, } } pub fn clone_sender(&self) { - let mut state = self.state.lock(); - state.senders = state - .senders - .checked_add(1) - .expect("mpmc sender count overflow"); + self.state.lock().senders += 1; } pub fn drop_sender(&self) { - let is_last = { + let mut waiters = { let mut state = self.state.lock(); state.senders -= 1; - state.senders == 0 + if state.senders != 0 { + return; + } + // Disconnection invalidates every receiver waiter ID. Move the storage out so both + // notification and reclamation happen without holding the queue lock. + mem::replace(&mut state.recv_waiters, WaitList::new()) }; - if is_last { - self.recv_waiters.notify_all(); - } + wake_all(std::iter::from_fn(|| waiters.notify_one())); } pub fn clone_receiver(&self) { - let mut state = self.state.lock(); - state.receivers = state - .receivers - .checked_add(1) - .expect("mpmc receiver count overflow"); + self.state.lock().receivers += 1; } pub fn drop_receiver(&self) { - let discarded = { + let (discarded, mut waiters) = { let mut state = self.state.lock(); state.receivers -= 1; - (state.receivers == 0).then(|| std::mem::take(&mut state.values)) + if state.receivers != 0 { + return; + } + ( + mem::take(&mut state.values), + mem::replace(&mut state.send_waiters, WaitList::new()), + ) }; - if discarded.is_some() { - self.send_waiters.notify_all(); - } + // Release blocked senders before destroying buffered values. Local ownership still drops + // the values if a wake callback unwinds. + wake_all(std::iter::from_fn(|| waiters.notify_one())); drop(discarded); } pub fn try_send(&self, value: T) -> Result<(), TrySendError> { - { + let waker = { let mut state = self.state.lock(); if state.receivers == 0 { return Err(TrySendError::Disconnected(value)); } - if self - .capacity - .is_some_and(|capacity| state.values.len() >= capacity) - { + if !state.has_capacity() { return Err(TrySendError::Full(value)); } - state.values.push_back(value); + state.push(value) + }; + if let Some(waker) = waker { + waker.wake(); } - self.recv_waiters.release_if_nonempty(1); Ok(()) } @@ -130,23 +206,16 @@ impl Shared { }; let mut send = Send { shared: self, - acquire: self.send_waiters.poll_acquire(1), + waiter: None, value: Some(value), }; poll_fn(|cx| send.poll(cx)).await } pub fn try_recv(&self) -> Result { - let value = { - let mut state = self.state.lock(); - match state.values.pop_front() { - Some(value) => value, - None if state.senders == 0 => return Err(TryRecvError::Disconnected), - None => return Err(TryRecvError::Empty), - } - }; - if self.capacity.is_some() { - self.send_waiters.release_if_nonempty(1); + let (value, waker) = self.state.lock().pop()?; + if let Some(waker) = waker { + waker.wake(); } Ok(value) } @@ -159,7 +228,7 @@ impl Shared { } let mut recv = Recv { shared: self, - acquire: self.recv_waiters.poll_acquire(1), + waiter: None, }; poll_fn(|cx| recv.poll(cx)).await } @@ -167,53 +236,135 @@ impl Shared { struct Send<'a, T> { shared: &'a Shared, - // Cancel the wait and pass on its notification before dropping a value whose destructor - // may depend on another blocked sender making progress. Fields drop in declaration order. - acquire: Acquire<'a>, + waiter: Option, + // `Drop` passes an unconsumed notification on before this value is destroyed, because its + // destructor may depend on another blocked sender making progress. value: Option, } impl Send<'_, T> { + fn take_value(&mut self) -> T { + self.value.take().expect("pending send must own its value") + } + fn poll(&mut self, cx: &mut Context<'_>) -> Poll>> { - let mut value = self.value.take().expect("pending send must own its value"); - loop { - let notified = Pin::new(&mut self.acquire).poll(cx); - value = match self.shared.try_send(value) { - Ok(()) => return Poll::Ready(Ok(())), - Err(TrySendError::Disconnected(value)) => { - return Poll::Ready(Err(SendError::new(value))); + let mut state = self.shared.state.lock(); + let outcome = if state.receivers == 0 { + self.waiter = None; + Err(self.take_value()) + } else if state.has_capacity() { + Ok(state.push(self.take_value())) + } else { + let retired = state.send_waiters.register(&mut self.waiter, cx.waker()); + drop(state); + drop(retired); + return Poll::Pending; + }; + let retired = self + .waiter + .take() + .map(|id| state.send_waiters.remove_waiter(id)); + drop(state); + // Deliver the notification before running waker destructors, which may panic. + let result = outcome + .map(|waker| { + if let Some(waker) = waker { + waker.wake(); } - Err(TrySendError::Full(value)) => value, - }; - if notified.is_ready() { - self.acquire = self.shared.send_waiters.poll_acquire(1); - } else { - self.value = Some(value); - return Poll::Pending; + }) + .map_err(SendError::new); + drop(retired); + Poll::Ready(result) + } +} + +impl Drop for Send<'_, T> { + fn drop(&mut self) { + let Some(id) = self.waiter.take() else { + return; + }; + let (retired, waker) = { + let mut state = self.shared.state.lock(); + if state.receivers == 0 { + return; } + let retired = state.send_waiters.remove_waiter(id); + // Hand an unconsumed notification to the next sender while the slot is still free. + let waker = if matches!(retired, Waiter::Notified) && state.has_capacity() { + state.send_waiters.notify_one() + } else { + None + }; + (retired, waker) + }; + if let Some(waker) = waker { + waker.wake(); } + drop(retired); } } struct Recv<'a, T> { shared: &'a Shared, - acquire: Acquire<'a>, + waiter: Option, } impl Recv<'_, T> { fn poll(&mut self, cx: &mut Context<'_>) -> Poll> { - loop { - let notified = Pin::new(&mut self.acquire).poll(cx); - match self.shared.try_recv() { - Ok(value) => return Poll::Ready(Ok(value)), - Err(TryRecvError::Disconnected) => { - return Poll::Ready(Err(RecvError::Disconnected)); - } - Err(TryRecvError::Empty) if notified.is_ready() => { - self.acquire = self.shared.recv_waiters.poll_acquire(1); - } - Err(TryRecvError::Empty) => return Poll::Pending, + let mut state = self.shared.state.lock(); + if state.senders == 0 { + // Buffered values remain readable after the waiter storage has been detached. + self.waiter = None; + } + let outcome = match state.pop() { + Ok(popped) => Ok(popped), + Err(TryRecvError::Disconnected) => Err(RecvError::Disconnected), + Err(TryRecvError::Empty) => { + let retired = state.recv_waiters.register(&mut self.waiter, cx.waker()); + drop(state); + drop(retired); + return Poll::Pending; + } + }; + let retired = self + .waiter + .take() + .map(|id| state.recv_waiters.remove_waiter(id)); + drop(state); + // Deliver the notification before running waker destructors, which may panic. + let result = outcome.map(|(value, waker)| { + if let Some(waker) = waker { + waker.wake(); } + value + }); + drop(retired); + Poll::Ready(result) + } +} + +impl Drop for Recv<'_, T> { + fn drop(&mut self) { + let Some(id) = self.waiter.take() else { + return; + }; + let (retired, waker) = { + let mut state = self.shared.state.lock(); + if state.senders == 0 { + return; + } + let retired = state.recv_waiters.remove_waiter(id); + // Hand an unconsumed notification to the next receiver while a value still waits. + let waker = if matches!(retired, Waiter::Notified) && !state.values.is_empty() { + state.recv_waiters.notify_one() + } else { + None + }; + (retired, waker) + }; + if let Some(waker) = waker { + waker.wake(); } + drop(retired); } } diff --git a/asyncband/src/mpmc/unbounded.rs b/asyncband/src/mpmc/unbounded.rs index 7e1db12a..7fc71ac6 100644 --- a/asyncband/src/mpmc/unbounded.rs +++ b/asyncband/src/mpmc/unbounded.rs @@ -28,9 +28,8 @@ use super::queue::Shared; /// /// Sends are synchronous and values may be buffered until available memory is exhausted. /// -/// Operations briefly acquire internal mutexes. No lock is held across an await point, while -/// waking tasks, or while dropping messages. Sending and trying to receive may wait to acquire -/// a mutex, but never wait for capacity or new messages. +/// Sending and trying to receive never wait for capacity or new messages, but may briefly block +/// on an internal mutex. pub fn unbounded() -> (UnboundedSender, UnboundedReceiver) { let shared = Arc::new(Shared::unbounded()); ( @@ -119,8 +118,8 @@ impl UnboundedReceiver { /// /// # Cancel safety /// - /// Dropping a pending `recv` does not consume a value. Any selected value notification is - /// passed to another waiting receiver, so cancellation does not prevent it from receiving. + /// Dropping a pending `recv` does not consume a value or prevent other receivers from receiving + /// it. pub async fn recv(&self) -> Result { self.shared.recv().await } diff --git a/asyncband/src/mpsc/bounded/mod.rs b/asyncband/src/mpsc/bounded/mod.rs index a9ae677b..60557b77 100644 --- a/asyncband/src/mpsc/bounded/mod.rs +++ b/asyncband/src/mpsc/bounded/mod.rs @@ -43,9 +43,8 @@ pub use self::sender::Permit; /// Message storage is preallocated for `buffer` values. Queued messages and outstanding /// reservations together occupy at most `buffer` capacity units. /// -/// Operations briefly acquire an internal mutex; no lock is held across an await point or while -/// invoking waker callbacks or message destructors. The `try_*` methods do not wait for capacity -/// or messages, but may wait to acquire this mutex. +/// The `try_*` methods do not wait for capacity or messages, but may briefly block on an internal +/// mutex. /// /// # Panics /// @@ -72,7 +71,7 @@ pub fn bounded(buffer: usize) -> (BoundedSender, BoundedReceiver) { } // While open, capacity belongs to available, a queued message, a Permit, or a granted waiter. -// All transitions hold one mutex. Waker callbacks and message destruction run after unlocking. +// All transitions hold one mutex. Wake callbacks and waker or message destruction run unlocked. struct State { queue: VecDeque, available: usize, diff --git a/asyncband/src/mpsc/bounded/receiver.rs b/asyncband/src/mpsc/bounded/receiver.rs index cf7436c7..11f6f7f3 100644 --- a/asyncband/src/mpsc/bounded/receiver.rs +++ b/asyncband/src/mpsc/bounded/receiver.rs @@ -24,8 +24,9 @@ use std::task::Poll; use super::State; use crate::internal::mutex::Mutex; +use crate::internal::register_waker; +use crate::internal::waitlist::WaitList; use crate::internal::wake_all; -use crate::internal::waker_batch::WakerBatch; use crate::mpsc::RecvError; use crate::mpsc::TryRecvError; @@ -45,21 +46,24 @@ impl fmt::Debug for BoundedReceiver { impl Drop for BoundedReceiver { fn drop(&mut self) { - let mut wakers = WakerBatch::new(); - let (queue, recv_waker) = { + let (queue, recv_waker, mut waiters) = { let mut state = self.shared.lock(); state.receiver = false; let queue = mem::take(&mut state.queue); let recv_waker = state.recv_waker.take(); - while let Some((_, waiter)) = state.send_waiters.unlink_first_waiter(|_| true) { + // The disconnected state invalidates both queued and granted waiter IDs. + let waiters = mem::replace(&mut state.send_waiters, WaitList::new()); + (queue, recv_waker, waiters) + }; + // Local ownership also drains the queue if a wake or waker destructor unwinds. + wake_all(std::iter::from_fn(|| { + loop { + let (_, waiter) = waiters.unlink_first_waiter(|_| true)?; if let Some(waker) = waiter.waker.take() { - wakers.push(waker); + return Some(waker); } } - (queue, recv_waker) - }; - // Local ownership also drains the queue if a wake or waker destructor unwinds. - wake_all(&mut wakers); + })); drop(recv_waker); drop(queue); } @@ -134,7 +138,6 @@ impl BoundedReceiver { } fn poll_recv(&mut self, cx: &mut Context<'_>) -> Poll> { - let waker = cx.waker().clone(); let mut state = self.shared.lock(); match state.pop() { Ok((value, wake)) => { @@ -151,7 +154,7 @@ impl BoundedReceiver { Poll::Ready(Err(RecvError::Disconnected)) } Err(TryRecvError::Empty) => { - let old = state.recv_waker.replace(waker); + let old = register_waker(&mut state.recv_waker, cx.waker()); drop(state); drop(old); Poll::Pending diff --git a/asyncband/src/mpsc/bounded/sender.rs b/asyncband/src/mpsc/bounded/sender.rs index 392fa159..2f286f04 100644 --- a/asyncband/src/mpsc/bounded/sender.rs +++ b/asyncband/src/mpsc/bounded/sender.rs @@ -25,6 +25,7 @@ use std::task::Poll; use super::State; use super::Waiter; use crate::internal::mutex::Mutex; +use crate::internal::register_waker; use crate::internal::waitlist::WaiterId; use crate::mpsc::SendError; use crate::mpsc::TrySendError; @@ -235,9 +236,9 @@ struct Reserve<'a, T> { impl<'a, T> Reserve<'a, T> { fn poll(&mut self, cx: &mut Context<'_>) -> Poll, SendError<()>>> { - let waker = cx.waker().clone(); let mut state = self.shared.lock(); if !state.receiver { + self.waiter = None; return Poll::Ready(Err(SendError::new(()))); } if let Some(index) = self.waiter { @@ -250,10 +251,9 @@ impl<'a, T> Reserve<'a, T> { }; drop(state); drop(waiter); - drop(waker); return Poll::Ready(Ok(permit)); } - let old = waiter.waker.replace(waker); + let old = register_waker(&mut waiter.waker, cx.waker()); drop(state); drop(old); return Poll::Pending; @@ -264,12 +264,11 @@ impl<'a, T> Reserve<'a, T> { shared: self.shared, }; drop(state); - drop(waker); return Poll::Ready(Ok(permit)); } self.waiter = Some(state.send_waiters.push_back(Waiter { grant: false, - waker: Some(waker), + waker: Some(cx.waker().clone()), })); Poll::Pending } @@ -280,6 +279,9 @@ impl Drop for Reserve<'_, T> { let Some(index) = self.waiter else { return }; let (waiter, wake) = { let mut state = self.shared.lock(); + if !state.receiver { + return; + } state.send_waiters.unlink_waiter(index, |_| true); let waiter = state.send_waiters.remove_unlinked_waiter(index); let wake = if waiter.grant { state.release() } else { None }; diff --git a/asyncband/src/mpsc/unbounded/mod.rs b/asyncband/src/mpsc/unbounded/mod.rs index be2b087c..42802cfa 100644 --- a/asyncband/src/mpsc/unbounded/mod.rs +++ b/asyncband/src/mpsc/unbounded/mod.rs @@ -43,9 +43,8 @@ pub use self::sender::UnboundedSender; /// Storage is reclaimed incrementally as messages are received. A bounded amount of empty /// storage may be retained for reuse, independently of the channel's previous peak occupancy. /// -/// Operations briefly acquire an internal mutex; no lock is held across an await point or while -/// invoking waker callbacks or message destructors. Sending and trying to receive may wait to -/// acquire this mutex, but never wait for capacity or new messages. +/// Sending and trying to receive never wait for capacity or new messages, but may briefly block +/// on an internal mutex. pub fn unbounded() -> (UnboundedSender, UnboundedReceiver) { let shared = Arc::new(Mutex::new(State { buffer: Buffer::new(), diff --git a/asyncband/src/mpsc/unbounded/receiver.rs b/asyncband/src/mpsc/unbounded/receiver.rs index f521cf06..90617be4 100644 --- a/asyncband/src/mpsc/unbounded/receiver.rs +++ b/asyncband/src/mpsc/unbounded/receiver.rs @@ -27,6 +27,7 @@ use super::State; use super::buffer::Buffer; use super::buffer::pop_batch; use crate::internal::mutex::Mutex; +use crate::internal::register_waker; use crate::mpsc::RecvError; use crate::mpsc::TryRecvError; @@ -149,8 +150,6 @@ impl UnboundedReceiver { if !batch.is_empty() { return Poll::Ready(Ok(pop_batch(batch))); } - // Waker clone/drop callbacks can send into this channel, so run them outside the lock. - let waker = cx.waker().clone(); let mut state = self.shared.lock(); let retired = state.buffer.refill(batch); if !batch.is_empty() { @@ -165,7 +164,7 @@ impl UnboundedReceiver { drop(old); return Poll::Ready(Err(RecvError::Disconnected)); } - let old = state.recv_waker.replace(waker); + let old = register_waker(&mut state.recv_waker, cx.waker()); drop(state); drop(retired); drop(old); diff --git a/benchmarks/asyncband/mpmc/bounded.rs b/benchmarks/asyncband/mpmc/bounded.rs index d4af76fb..1bf7253f 100644 --- a/benchmarks/asyncband/mpmc/bounded.rs +++ b/benchmarks/asyncband/mpmc/bounded.rs @@ -15,9 +15,14 @@ // specific language governing permissions and limitations // under the License. +use std::pin::pin; + +use asyncband::mpmc; use divan::Bencher; +use divan::black_box; use divan::counter::ItemsCount; +use super::FAST_SAMPLE_SIZE; use crate::mpmc_support::adapters::Asyncband; use crate::mpmc_support::support::BATCH_MESSAGES; use crate::mpmc_support::support::BOUNDED_CAPACITY; @@ -27,6 +32,10 @@ use crate::mpmc_support::support::TaskBatch; use crate::mpmc_support::support::ThreadBatch; use crate::mpmc_support::support::Topology; use crate::mpmc_support::support::runtime; +use crate::support::bench_context; +use crate::support::poll_pending; +use crate::support::poll_pinned_ready; +use crate::support::poll_ready; #[divan::bench( args = TOPOLOGIES, @@ -53,3 +62,35 @@ fn tokio_tasks(bencher: Bencher, topology: Topology) { .with_inputs(|| TaskBatch::new::>(&runtime, topology)) .bench_local_refs(|batch| runtime.block_on(batch.run())); } + +#[divan::bench(sample_size = FAST_SAMPLE_SIZE)] +fn try_send_then_try_recv(bencher: Bencher) { + let (sender, receiver) = mpmc::bounded(1); + bencher.bench_local(|| { + sender.try_send(black_box(1usize)).unwrap(); + black_box(receiver.try_recv().unwrap()) + }); +} + +#[divan::bench(sample_size = FAST_SAMPLE_SIZE)] +fn send_then_recv(bencher: Bencher) { + let mut context = bench_context(); + let (sender, receiver) = mpmc::bounded(1); + bencher.bench_local(|| { + poll_ready(sender.send(black_box(1usize)), &mut context).unwrap(); + black_box(poll_ready(receiver.recv(), &mut context).unwrap()) + }); +} + +#[divan::bench(sample_size = FAST_SAMPLE_SIZE)] +fn wake_blocked_sender(bencher: Bencher) { + let mut context = bench_context(); + let (sender, receiver) = mpmc::bounded(1); + sender.try_send(0usize).unwrap(); + bencher.bench_local(|| { + let mut send = pin!(sender.send(black_box(1))); + poll_pending(send.as_mut(), &mut context); + black_box(receiver.try_recv().unwrap()); + poll_pinned_ready(send.as_mut(), &mut context).unwrap(); + }); +} diff --git a/benchmarks/asyncband/mpmc/mod.rs b/benchmarks/asyncband/mpmc/mod.rs index e0ac8347..a00276ab 100644 --- a/benchmarks/asyncband/mpmc/mod.rs +++ b/benchmarks/asyncband/mpmc/mod.rs @@ -17,3 +17,6 @@ mod bounded; mod unbounded; + +// Fixed sample sizes keep one-time warm-up work from changing Divan's iteration granularity. +const FAST_SAMPLE_SIZE: u32 = 256; diff --git a/benchmarks/asyncband/mpmc/unbounded.rs b/benchmarks/asyncband/mpmc/unbounded.rs index 965a12ed..34234c42 100644 --- a/benchmarks/asyncband/mpmc/unbounded.rs +++ b/benchmarks/asyncband/mpmc/unbounded.rs @@ -15,9 +15,14 @@ // specific language governing permissions and limitations // under the License. +use std::pin::pin; + +use asyncband::mpmc; use divan::Bencher; +use divan::black_box; use divan::counter::ItemsCount; +use super::FAST_SAMPLE_SIZE; use crate::mpmc_support::adapters::Asyncband; use crate::mpmc_support::support::BATCH_MESSAGES; use crate::mpmc_support::support::TOPOLOGIES; @@ -26,6 +31,10 @@ use crate::mpmc_support::support::ThreadBatch; use crate::mpmc_support::support::Topology; use crate::mpmc_support::support::Unbounded; use crate::mpmc_support::support::runtime; +use crate::support::bench_context; +use crate::support::poll_pending; +use crate::support::poll_pinned_ready; +use crate::support::poll_ready; #[divan::bench( args = TOPOLOGIES, @@ -52,3 +61,43 @@ fn tokio_tasks(bencher: Bencher, topology: Topology) { .with_inputs(|| TaskBatch::new::>(&runtime, topology)) .bench_local_refs(|batch| runtime.block_on(batch.run())); } + +#[divan::bench(sample_size = FAST_SAMPLE_SIZE)] +fn send_then_try_recv(bencher: Bencher) { + let (sender, receiver) = mpmc::unbounded(); + bencher.bench_local(|| { + sender.send(black_box(1usize)).unwrap(); + black_box(receiver.try_recv().unwrap()) + }); +} + +#[divan::bench(sample_size = FAST_SAMPLE_SIZE)] +fn send_then_recv(bencher: Bencher) { + let mut context = bench_context(); + let (sender, receiver) = mpmc::unbounded(); + bencher.bench_local(|| { + sender.send(black_box(1usize)).unwrap(); + black_box(poll_ready(receiver.recv(), &mut context).unwrap()) + }); +} + +#[divan::bench(sample_size = FAST_SAMPLE_SIZE)] +fn wake_pending_receiver(bencher: Bencher) { + let mut context = bench_context(); + let (sender, receiver) = mpmc::unbounded(); + bencher.bench_local(|| { + let mut recv = pin!(receiver.recv()); + poll_pending(recv.as_mut(), &mut context); + sender.send(black_box(usize::MAX)).unwrap(); + black_box(poll_pinned_ready(recv.as_mut(), &mut context).unwrap()) + }); +} + +#[divan::bench(sample_size = FAST_SAMPLE_SIZE)] +fn repoll_pending_receiver(bencher: Bencher) { + let mut context = bench_context(); + let (_sender, receiver) = mpmc::unbounded::(); + let mut recv = pin!(receiver.recv()); + poll_pending(recv.as_mut(), &mut context); + bencher.bench_local(|| poll_pending(recv.as_mut(), &mut context)); +} diff --git a/tests-integration/tests/mpmc_test.rs b/tests-integration/tests/mpmc_test.rs deleted file mode 100644 index f3f9e0a0..00000000 --- a/tests-integration/tests/mpmc_test.rs +++ /dev/null @@ -1,441 +0,0 @@ -// Licensed to the Apache Software Foundation (ASF) under one -// or more contributor license agreements. See the NOTICE file -// distributed with this work for additional information -// regarding copyright ownership. The ASF licenses this file -// to you under the Apache License, Version 2.0 (the -// "License"); you may not use this file except in compliance -// with the License. You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, -// software distributed under the License is distributed on an -// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -// KIND, either express or implied. See the License for the -// specific language governing permissions and limitations -// under the License. - -use std::future::Future; -use std::pin::Pin; -use std::sync::Arc; -use std::sync::atomic::AtomicUsize; -use std::sync::atomic::Ordering; -use std::task::Context; -use std::task::Poll; -use std::task::Wake; -use std::task::Waker; -use std::time::Duration; - -use asyncband::mpmc; -use asyncband::mpmc::RecvError; -use asyncband::mpmc::TryRecvError; -use asyncband::mpmc::TrySendError; -use tests_integration::poll_once; - -#[derive(Debug)] -struct WakeCounter(AtomicUsize); - -impl WakeCounter { - fn count(&self) -> usize { - self.0.load(Ordering::SeqCst) - } -} - -impl Wake for WakeCounter { - fn wake(self: Arc) { - self.0.fetch_add(1, Ordering::SeqCst); - } -} - -fn expect_ready(poll: Poll) -> T { - match poll { - Poll::Ready(value) => value, - Poll::Pending => panic!("future should be ready"), - } -} - -fn poll_with_waker(future: Pin<&mut F>, waker: &Waker) -> Poll { - future.poll(&mut Context::from_waker(waker)) -} - -#[test] -fn bounded_enforces_exact_capacity_and_fifo_order() { - let (sender, receiver) = mpmc::bounded(2); - let competing = receiver.clone(); - - sender.try_send(0).unwrap(); - sender.try_send(1).unwrap(); - assert_eq!(sender.try_send(2), Err(TrySendError::Full(2))); - - assert_eq!(receiver.try_recv(), Ok(0)); - assert_eq!(competing.try_recv(), Ok(1)); - assert_eq!(receiver.try_recv(), Err(TryRecvError::Empty)); -} - -#[test] -#[should_panic(expected = "mpmc bounded queue requires capacity > 0")] -fn bounded_rejects_zero_capacity() { - let _ = mpmc::bounded::<()>(0); -} - -#[test] -fn receiver_and_sender_clone_counts_control_disconnection() { - let (sender, receiver) = mpmc::unbounded(); - let sender_clone = sender.clone(); - let receiver_clone = receiver.clone(); - - drop(receiver); - sender.send(1).unwrap(); - assert_eq!(receiver_clone.try_recv(), Ok(1)); - - drop(sender); - assert_eq!(receiver_clone.try_recv(), Err(TryRecvError::Empty)); - drop(sender_clone); - assert_eq!(receiver_clone.try_recv(), Err(TryRecvError::Disconnected)); - - drop(receiver_clone); -} - -#[test] -fn last_receiver_returns_each_unsent_value_once() { - let (bounded_sender, bounded_receiver) = mpmc::bounded(1); - let bounded_receiver_clone = bounded_receiver.clone(); - drop(bounded_receiver); - drop(bounded_receiver_clone); - assert_eq!( - bounded_sender.try_send(1), - Err(TrySendError::Disconnected(1)) - ); - assert_eq!( - bounded_sender.try_send(2), - Err(TrySendError::Disconnected(2)) - ); - - let (unbounded_sender, unbounded_receiver) = mpmc::unbounded(); - drop(unbounded_receiver); - assert_eq!(unbounded_sender.send(3).unwrap_err().into_inner(), 3); -} - -#[tokio::test] -async fn buffered_values_drain_before_disconnection() { - let (sender, receiver) = mpmc::bounded(3); - sender.send(0).await.unwrap(); - sender.send(1).await.unwrap(); - sender.send(2).await.unwrap(); - drop(sender); - - assert_eq!(receiver.recv().await, Ok(0)); - assert_eq!(receiver.recv().await, Ok(1)); - assert_eq!(receiver.recv().await, Ok(2)); - assert_eq!(receiver.recv().await, Err(RecvError::Disconnected)); -} - -#[tokio::test] -async fn unbounded_preserves_fifo_order_and_drains_before_disconnection() { - let (sender, receiver) = mpmc::unbounded(); - sender.send(0).unwrap(); - sender.send(1).unwrap(); - sender.send(2).unwrap(); - drop(sender); - - assert_eq!(receiver.recv().await, Ok(0)); - assert_eq!(receiver.recv().await, Ok(1)); - assert_eq!(receiver.recv().await, Ok(2)); - assert_eq!(receiver.recv().await, Err(RecvError::Disconnected)); -} - -#[test] -fn bounded_send_wakes_only_the_first_receiver() { - let (sender, receiver) = mpmc::bounded(2); - let competing = receiver.clone(); - let mut first = Box::pin(receiver.recv()); - let mut second = Box::pin(competing.recv()); - let first_wakes = Arc::new(WakeCounter(AtomicUsize::new(0))); - let second_wakes = Arc::new(WakeCounter(AtomicUsize::new(0))); - let first_waker = Waker::from(first_wakes.clone()); - let second_waker = Waker::from(second_wakes.clone()); - - assert!(poll_with_waker(first.as_mut(), &first_waker).is_pending()); - assert!(poll_with_waker(second.as_mut(), &second_waker).is_pending()); - sender.try_send(1).unwrap(); - - assert_eq!(first_wakes.count(), 1); - assert_eq!(second_wakes.count(), 0); - assert_eq!( - expect_ready(poll_with_waker(first.as_mut(), &first_waker)), - Ok(1) - ); - assert!(poll_with_waker(second.as_mut(), &second_waker).is_pending()); - assert_eq!(second_wakes.count(), 0); -} - -#[test] -fn unbounded_send_wakes_only_the_first_receiver() { - let (sender, receiver) = mpmc::unbounded(); - let competing = receiver.clone(); - let mut first = Box::pin(receiver.recv()); - let mut second = Box::pin(competing.recv()); - let first_wakes = Arc::new(WakeCounter(AtomicUsize::new(0))); - let second_wakes = Arc::new(WakeCounter(AtomicUsize::new(0))); - let first_waker = Waker::from(first_wakes.clone()); - let second_waker = Waker::from(second_wakes.clone()); - - assert!(poll_with_waker(first.as_mut(), &first_waker).is_pending()); - assert!(poll_with_waker(second.as_mut(), &second_waker).is_pending()); - sender.send(1).unwrap(); - - assert_eq!(first_wakes.count(), 1); - assert_eq!(second_wakes.count(), 0); - assert_eq!( - expect_ready(poll_with_waker(first.as_mut(), &first_waker)), - Ok(1) - ); - assert!(poll_with_waker(second.as_mut(), &second_waker).is_pending()); - assert_eq!(second_wakes.count(), 0); -} - -#[test] -fn cancelled_notified_receiver_passes_value_to_next_receiver() { - let (sender, receiver) = mpmc::unbounded(); - let competing = receiver.clone(); - let mut cancelled = Box::pin(receiver.recv()); - let mut waiting = Box::pin(competing.recv()); - let cancelled_wakes = Arc::new(WakeCounter(AtomicUsize::new(0))); - let waiting_wakes = Arc::new(WakeCounter(AtomicUsize::new(0))); - let cancelled_waker = Waker::from(cancelled_wakes.clone()); - let waiting_waker = Waker::from(waiting_wakes.clone()); - - assert!(poll_with_waker(cancelled.as_mut(), &cancelled_waker).is_pending()); - assert!(poll_with_waker(waiting.as_mut(), &waiting_waker).is_pending()); - sender.send(1).unwrap(); - assert_eq!(cancelled_wakes.count(), 1); - assert_eq!(waiting_wakes.count(), 0); - drop(cancelled); - - assert_eq!(waiting_wakes.count(), 1); - assert_eq!( - expect_ready(poll_with_waker(waiting.as_mut(), &waiting_waker)), - Ok(1) - ); -} - -#[test] -fn bounded_cancelled_notified_receiver_passes_value_to_next_receiver() { - let (sender, receiver) = mpmc::bounded(1); - let competing = receiver.clone(); - let mut cancelled = Box::pin(receiver.recv()); - let mut waiting = Box::pin(competing.recv()); - let cancelled_wakes = Arc::new(WakeCounter(AtomicUsize::new(0))); - let waiting_wakes = Arc::new(WakeCounter(AtomicUsize::new(0))); - let cancelled_waker = Waker::from(cancelled_wakes.clone()); - let waiting_waker = Waker::from(waiting_wakes.clone()); - - assert!(poll_with_waker(cancelled.as_mut(), &cancelled_waker).is_pending()); - assert!(poll_with_waker(waiting.as_mut(), &waiting_waker).is_pending()); - sender.try_send(1).unwrap(); - assert_eq!(cancelled_wakes.count(), 1); - assert_eq!(waiting_wakes.count(), 0); - drop(cancelled); - - assert_eq!(waiting_wakes.count(), 1); - assert_eq!( - expect_ready(poll_with_waker(waiting.as_mut(), &waiting_waker)), - Ok(1) - ); -} - -#[test] -fn bounded_cancelled_sender_notifies_next_sender_before_dropping_value() { - #[derive(Debug)] - struct Value { - id: usize, - wake_observer: Option<(Arc, Arc)>, - } - - impl Drop for Value { - fn drop(&mut self) { - if let Some((wakes, observed)) = &self.wake_observer { - // A message destructor may depend on another blocked sender making progress. - observed.store(wakes.count(), Ordering::SeqCst); - } - } - } - - let (sender, receiver) = mpmc::bounded(1); - sender - .try_send(Value { - id: 0, - wake_observer: None, - }) - .unwrap(); - let first_sender = sender.clone(); - let second_sender = sender.clone(); - let cancelled_wakes = Arc::new(WakeCounter(AtomicUsize::new(0))); - let waiting_wakes = Arc::new(WakeCounter(AtomicUsize::new(0))); - let wakes_during_drop = Arc::new(AtomicUsize::new(usize::MAX)); - let mut cancelled = Box::pin(first_sender.send(Value { - id: 1, - wake_observer: Some((waiting_wakes.clone(), wakes_during_drop.clone())), - })); - let mut waiting = Box::pin(second_sender.send(Value { - id: 2, - wake_observer: None, - })); - let cancelled_waker = Waker::from(cancelled_wakes.clone()); - let waiting_waker = Waker::from(waiting_wakes.clone()); - - assert!(poll_with_waker(cancelled.as_mut(), &cancelled_waker).is_pending()); - assert!(poll_with_waker(waiting.as_mut(), &waiting_waker).is_pending()); - assert_eq!(receiver.try_recv().unwrap().id, 0); - assert_eq!(cancelled_wakes.count(), 1); - assert_eq!(waiting_wakes.count(), 0); - drop(cancelled); - - assert_eq!(wakes_during_drop.load(Ordering::SeqCst), 1); - assert_eq!(waiting_wakes.count(), 1); - expect_ready(poll_with_waker(waiting.as_mut(), &waiting_waker)).unwrap(); - assert_eq!(receiver.try_recv().unwrap().id, 2); - assert!(matches!(receiver.try_recv(), Err(TryRecvError::Empty))); -} - -#[test] -fn last_endpoint_wakes_all_opposite_waiters() { - let (sender, receiver) = mpmc::bounded(1); - sender.try_send(0).unwrap(); - let sender_clone = sender.clone(); - let mut first_send = Box::pin(sender.send(1)); - let mut second_send = Box::pin(sender_clone.send(2)); - assert!(poll_once(first_send.as_mut()).is_pending()); - assert!(poll_once(second_send.as_mut()).is_pending()); - drop(receiver); - assert_eq!( - expect_ready(poll_once(first_send.as_mut())) - .unwrap_err() - .into_inner(), - 1 - ); - assert_eq!( - expect_ready(poll_once(second_send.as_mut())) - .unwrap_err() - .into_inner(), - 2 - ); - - let (sender, receiver) = mpmc::unbounded::(); - let competing = receiver.clone(); - let mut first_recv = Box::pin(receiver.recv()); - let mut second_recv = Box::pin(competing.recv()); - assert!(poll_once(first_recv.as_mut()).is_pending()); - assert!(poll_once(second_recv.as_mut()).is_pending()); - drop(sender); - assert_eq!( - expect_ready(poll_once(first_recv.as_mut())), - Err(RecvError::Disconnected) - ); - assert_eq!( - expect_ready(poll_once(second_recv.as_mut())), - Err(RecvError::Disconnected) - ); -} - -#[tokio::test(flavor = "multi_thread", worker_threads = 4)] -async fn bounded_values_are_delivered_exactly_once_under_contention() { - const PRODUCERS: usize = 8; - const CONSUMERS: usize = 8; - const VALUES_PER_PRODUCER: usize = 512; - const TOTAL: usize = PRODUCERS * VALUES_PER_PRODUCER; - - let (sender, receiver) = mpmc::bounded(32); - let consumers = (0..CONSUMERS) - .map(|_| { - let receiver = receiver.clone(); - tokio::spawn(async move { - let mut values = Vec::new(); - while let Ok(value) = receiver.recv().await { - values.push(value); - } - values - }) - }) - .collect::>(); - drop(receiver); - - let producers = (0..PRODUCERS) - .map(|producer| { - let sender = sender.clone(); - tokio::spawn(async move { - let first = producer * VALUES_PER_PRODUCER; - for value in first..first + VALUES_PER_PRODUCER { - sender.send(value).await.unwrap(); - } - }) - }) - .collect::>(); - drop(sender); - - for producer in producers { - producer.await.unwrap(); - } - let mut received = Vec::with_capacity(TOTAL); - for consumer in consumers { - received.extend( - tokio::time::timeout(Duration::from_secs(10), consumer) - .await - .expect("bounded consumers must make progress") - .unwrap(), - ); - } - received.sort_unstable(); - assert_eq!(received, (0..TOTAL).collect::>()); -} - -#[tokio::test(flavor = "multi_thread", worker_threads = 4)] -async fn unbounded_values_are_delivered_exactly_once_under_contention() { - const PRODUCERS: usize = 8; - const CONSUMERS: usize = 8; - const VALUES_PER_PRODUCER: usize = 512; - const TOTAL: usize = PRODUCERS * VALUES_PER_PRODUCER; - - let (sender, receiver) = mpmc::unbounded(); - let consumers = (0..CONSUMERS) - .map(|_| { - let receiver = receiver.clone(); - tokio::spawn(async move { - let mut values = Vec::new(); - while let Ok(value) = receiver.recv().await { - values.push(value); - } - values - }) - }) - .collect::>(); - drop(receiver); - - let producers = (0..PRODUCERS) - .map(|producer| { - let sender = sender.clone(); - tokio::spawn(async move { - let first = producer * VALUES_PER_PRODUCER; - for value in first..first + VALUES_PER_PRODUCER { - sender.send(value).unwrap(); - } - }) - }) - .collect::>(); - drop(sender); - - for producer in producers { - producer.await.unwrap(); - } - let mut received = Vec::with_capacity(TOTAL); - for consumer in consumers { - received.extend( - tokio::time::timeout(Duration::from_secs(10), consumer) - .await - .expect("unbounded consumers must make progress") - .unwrap(), - ); - } - received.sort_unstable(); - assert_eq!(received, (0..TOTAL).collect::>()); -} diff --git a/tests-integration/tests/mpmc_test/callbacks.rs b/tests-integration/tests/mpmc_test/callbacks.rs new file mode 100644 index 00000000..cf851391 --- /dev/null +++ b/tests-integration/tests/mpmc_test/callbacks.rs @@ -0,0 +1,114 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use std::panic::AssertUnwindSafe; +use std::panic::catch_unwind; +use std::sync::Arc; +use std::task::Waker; + +use asyncband::mpmc; +use asyncband::mpmc::RecvError; +use tests_integration::PanicWake; +use tests_integration::WakeCounter; +use tests_integration::assert_completes_without_deadlock; +use tests_integration::expect_ready; +use tests_integration::poll_once; +use tests_integration::poll_with; +use tests_integration::waker_on_drop; + +#[test] +fn replacing_a_send_waker_allows_its_destructor_to_receive() { + assert_completes_without_deadlock(|| { + let (sender, receiver) = mpmc::bounded(1); + sender.try_send(0).unwrap(); + let reentrant = receiver.clone(); + let first = waker_on_drop(move || assert_eq!(reentrant.try_recv(), Ok(0))); + let (second, second_wakes) = WakeCounter::new(); + let mut send = Box::pin(sender.send(1)); + + assert!(poll_with(send.as_mut(), &first).is_pending()); + drop(first); + // Replacing the stored waker frees capacity and wakes the newly registered task. + assert!(poll_with(send.as_mut(), &second).is_pending()); + assert_eq!(second_wakes.count(), 1); + expect_ready(poll_with(send.as_mut(), &second)).unwrap(); + assert_eq!(receiver.try_recv(), Ok(1)); + }); +} + +#[test] +fn replacing_a_receive_waker_allows_its_destructor_to_send() { + assert_completes_without_deadlock(|| { + let (sender, receiver) = mpmc::unbounded(); + let reentrant = sender.clone(); + let first = waker_on_drop(move || reentrant.send(1).unwrap()); + let (second, second_wakes) = WakeCounter::new(); + let mut recv = Box::pin(receiver.recv()); + + assert!(poll_with(recv.as_mut(), &first).is_pending()); + drop(first); + // Replacing the stored waker releases its last reference, whose destructor sends. + assert!(poll_with(recv.as_mut(), &second).is_pending()); + assert_eq!(second_wakes.count(), 1); + assert_eq!(expect_ready(poll_with(recv.as_mut(), &second)), Ok(1)); + drop(sender); + }); +} + +#[test] +fn last_sender_attempts_every_wake_after_one_panics() { + let (sender, receiver) = mpmc::unbounded::(); + let competing = receiver.clone(); + let mut panicking = Box::pin(receiver.recv()); + let mut waiting = Box::pin(competing.recv()); + let panic_waker = Waker::from(Arc::new(PanicWake)); + let (waker, wakes) = WakeCounter::new(); + assert!(poll_with(panicking.as_mut(), &panic_waker).is_pending()); + assert!(poll_with(waiting.as_mut(), &waker).is_pending()); + + assert!(catch_unwind(AssertUnwindSafe(|| drop(sender))).is_err()); + assert_eq!(wakes.count(), 1); + assert_eq!( + expect_ready(poll_with(waiting.as_mut(), &waker)), + Err(RecvError::Disconnected) + ); + assert_eq!( + expect_ready(poll_once(panicking.as_mut())), + Err(RecvError::Disconnected) + ); +} + +#[test] +fn last_receiver_releases_senders_and_values_after_a_wake_panics() { + let buffered = Arc::new(()); + let released = Arc::downgrade(&buffered); + let (sender, receiver) = mpmc::bounded(1); + sender.try_send(buffered).unwrap(); + let competing = sender.clone(); + let mut panicking = Box::pin(sender.send(Arc::new(()))); + let mut waiting = Box::pin(competing.send(Arc::new(()))); + let panic_waker = Waker::from(Arc::new(PanicWake)); + let (waker, wakes) = WakeCounter::new(); + assert!(poll_with(panicking.as_mut(), &panic_waker).is_pending()); + assert!(poll_with(waiting.as_mut(), &waker).is_pending()); + + assert!(catch_unwind(AssertUnwindSafe(|| drop(receiver))).is_err()); + assert_eq!(wakes.count(), 1); + assert!(released.upgrade().is_none()); + assert!(expect_ready(poll_with(waiting.as_mut(), &waker)).is_err()); + assert!(expect_ready(poll_once(panicking.as_mut())).is_err()); +} diff --git a/tests-integration/tests/mpmc_test/concurrency.rs b/tests-integration/tests/mpmc_test/concurrency.rs new file mode 100644 index 00000000..2d9d2660 --- /dev/null +++ b/tests-integration/tests/mpmc_test/concurrency.rs @@ -0,0 +1,107 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use std::ops::Range; +use std::time::Duration; + +use asyncband::mpmc; +use tokio::task::JoinHandle; + +use super::Receiver; + +const PRODUCERS: usize = 8; +const CONSUMERS: usize = 8; +const VALUES_PER_PRODUCER: usize = 512; +const TOTAL: usize = PRODUCERS * VALUES_PER_PRODUCER; + +fn values_of(producer: usize) -> Range { + let first = producer * VALUES_PER_PRODUCER; + first..first + VALUES_PER_PRODUCER +} + +/// Consumes until disconnection and asserts that every produced value arrived exactly once. +async fn assert_delivered_exactly_once(receiver: R, producers: Vec>) +where + R: Receiver + Send + 'static, +{ + let consumers = (0..CONSUMERS) + .map(|_| { + let receiver = receiver.clone(); + tokio::spawn(async move { + let mut values = Vec::new(); + while let Ok(value) = receiver.recv().await { + values.push(value); + } + values + }) + }) + .collect::>(); + drop(receiver); + + let mut received = tokio::time::timeout(Duration::from_secs(10), async { + for producer in producers { + producer.await.unwrap(); + } + let mut received = Vec::with_capacity(TOTAL); + for consumer in consumers { + received.extend(consumer.await.unwrap()); + } + received + }) + .await + .expect("producers and consumers must make progress"); + received.sort_unstable(); + assert_eq!(received, (0..TOTAL).collect::>()); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +#[cfg_attr(miri, ignore = "requires an OS-backed Tokio runtime")] +async fn bounded_values_are_delivered_exactly_once_under_contention() { + let (sender, receiver) = mpmc::bounded(32); + let producers = (0..PRODUCERS) + .map(|producer| { + let sender = sender.clone(); + tokio::spawn(async move { + for value in values_of(producer) { + sender.send(value).await.unwrap(); + } + }) + }) + .collect(); + drop(sender); + + assert_delivered_exactly_once(receiver, producers).await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +#[cfg_attr(miri, ignore = "requires an OS-backed Tokio runtime")] +async fn unbounded_values_are_delivered_exactly_once_under_contention() { + let (sender, receiver) = mpmc::unbounded(); + let producers = (0..PRODUCERS) + .map(|producer| { + let sender = sender.clone(); + tokio::spawn(async move { + for value in values_of(producer) { + sender.send(value).unwrap(); + } + }) + }) + .collect(); + drop(sender); + + assert_delivered_exactly_once(receiver, producers).await; +} diff --git a/tests-integration/tests/mpmc_test/main.rs b/tests-integration/tests/mpmc_test/main.rs new file mode 100644 index 00000000..8795d7ad --- /dev/null +++ b/tests-integration/tests/mpmc_test/main.rs @@ -0,0 +1,152 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use std::future::Future; +use std::pin::pin; + +use asyncband::mpmc; +use asyncband::mpmc::RecvError; +use asyncband::mpmc::TryRecvError; +use asyncband::mpmc::TrySendError; +use tests_integration::expect_ready; +use tests_integration::poll_once; + +// Public queue contracts. The other suites cover notifications, callbacks, and concurrency. +mod callbacks; +mod concurrency; +mod notification; + +/// Either receiver flavor, so one case can cover both queues. +trait Receiver: Clone { + fn recv(&self) -> impl Future> + Send; + fn try_recv(&self) -> Result; +} + +impl Receiver for mpmc::BoundedReceiver { + fn recv(&self) -> impl Future> + Send { + self.recv() + } + + fn try_recv(&self) -> Result { + self.try_recv() + } +} + +impl Receiver for mpmc::UnboundedReceiver { + fn recv(&self) -> impl Future> + Send { + self.recv() + } + + fn try_recv(&self) -> Result { + self.try_recv() + } +} + +#[test] +fn bounded_enforces_exact_capacity_and_fifo_order() { + let (sender, receiver) = mpmc::bounded(2); + let competing = receiver.clone(); + + sender.try_send(0).unwrap(); + sender.try_send(1).unwrap(); + assert_eq!(sender.try_send(2), Err(TrySendError::Full(2))); + + assert_eq!(receiver.try_recv(), Ok(0)); + assert_eq!(competing.try_recv(), Ok(1)); + assert_eq!(receiver.try_recv(), Err(TryRecvError::Empty)); +} + +#[test] +#[should_panic(expected = "mpmc bounded queue requires capacity > 0")] +fn bounded_rejects_zero_capacity() { + let _ = mpmc::bounded::<()>(0); +} + +#[test] +fn receiver_and_sender_clone_counts_control_disconnection() { + let (sender, receiver) = mpmc::unbounded(); + let sender_clone = sender.clone(); + let receiver_clone = receiver.clone(); + + drop(receiver); + sender.send(1).unwrap(); + assert_eq!(receiver_clone.try_recv(), Ok(1)); + + drop(sender); + assert_eq!(receiver_clone.try_recv(), Err(TryRecvError::Empty)); + drop(sender_clone); + assert_eq!(receiver_clone.try_recv(), Err(TryRecvError::Disconnected)); +} + +#[test] +fn sends_after_the_last_receiver_return_the_value() { + let (bounded_sender, bounded_receiver) = mpmc::bounded(1); + let bounded_receiver_clone = bounded_receiver.clone(); + drop(bounded_receiver); + drop(bounded_receiver_clone); + assert_eq!( + bounded_sender.try_send(1), + Err(TrySendError::Disconnected(1)) + ); + assert_eq!( + bounded_sender.try_send(2), + Err(TrySendError::Disconnected(2)) + ); + + let (unbounded_sender, unbounded_receiver) = mpmc::unbounded(); + drop(unbounded_receiver); + assert_eq!(unbounded_sender.send(3).unwrap_err().into_inner(), 3); +} + +fn buffered_values_drain_in_order_before_disconnection( + sender: S, + send: impl Fn(&S, usize), + receiver: impl Receiver, +) { + for value in 0..3 { + send(&sender, value); + } + drop(sender); + + for expected in 0..3 { + assert_eq!(expect_ready(poll_once(pin!(receiver.recv()))), Ok(expected)); + } + assert_eq!( + expect_ready(poll_once(pin!(receiver.recv()))), + Err(RecvError::Disconnected) + ); +} + +#[test] +fn bounded_buffered_values_drain_in_order_before_disconnection() { + let (sender, receiver) = mpmc::bounded(3); + buffered_values_drain_in_order_before_disconnection( + sender, + |sender, value| sender.try_send(value).unwrap(), + receiver, + ); +} + +#[test] +fn unbounded_buffered_values_drain_in_order_before_disconnection() { + let (sender, receiver) = mpmc::unbounded(); + buffered_values_drain_in_order_before_disconnection( + sender, + |sender, value| sender.send(value).unwrap(), + receiver, + ); +} diff --git a/tests-integration/tests/mpmc_test/notification.rs b/tests-integration/tests/mpmc_test/notification.rs new file mode 100644 index 00000000..1b2a944d --- /dev/null +++ b/tests-integration/tests/mpmc_test/notification.rs @@ -0,0 +1,284 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; + +use asyncband::mpmc; +use asyncband::mpmc::RecvError; +use asyncband::mpmc::TryRecvError; +use tests_integration::WakeCounter; +use tests_integration::expect_ready; +use tests_integration::poll_with; + +use super::Receiver; + +#[derive(Debug)] +enum NotifiedReceiver { + Receive, + Cancel, + TakeValueThenCancel, +} + +fn receiver_notifications>( + channel: impl Fn() -> (S, R), + send: impl Fn(&S, usize), +) { + use NotifiedReceiver::*; + + for action in [Receive, Cancel, TakeValueThenCancel] { + let (sender, receiver) = channel(); + let competing = receiver.clone(); + let mut first = Box::pin(receiver.recv()); + let mut second = Box::pin(competing.recv()); + let (first_waker, first_wakes) = WakeCounter::new(); + let (second_waker, second_wakes) = WakeCounter::new(); + + assert!(poll_with(first.as_mut(), &first_waker).is_pending()); + assert!(poll_with(second.as_mut(), &second_waker).is_pending()); + send(&sender, 1); + assert_eq!(first_wakes.count(), 1); + assert_eq!(second_wakes.count(), 0); + + let expected = match action { + Receive => { + assert_eq!(expect_ready(poll_with(first.as_mut(), &first_waker)), Ok(1)); + assert!(poll_with(second.as_mut(), &second_waker).is_pending()); + assert_eq!(second_wakes.count(), 0); + send(&sender, 2); + 2 + } + Cancel => { + drop(first); + 1 + } + TakeValueThenCancel => { + assert_eq!(receiver.try_recv(), Ok(1)); + drop(first); + assert_eq!(second_wakes.count(), 0); + // Cancellation must leave the next send able to notify this waiter. + send(&sender, 2); + 2 + } + }; + assert_eq!(second_wakes.count(), 1, "{action:?}"); + assert_eq!( + expect_ready(poll_with(second.as_mut(), &second_waker)), + Ok(expected), + "{action:?}" + ); + } +} + +#[test] +fn bounded_receiver_notification_is_consumed_or_handed_off() { + receiver_notifications( + || mpmc::bounded(2), + |sender, value| sender.try_send(value).unwrap(), + ); +} + +#[test] +fn unbounded_receiver_notification_is_consumed_or_handed_off() { + receiver_notifications(mpmc::unbounded, |sender, value| sender.send(value).unwrap()); +} + +#[test] +fn notified_receiver_that_loses_the_value_queues_behind_waiting_receivers() { + let (sender, receiver) = mpmc::unbounded(); + let second_receiver = receiver.clone(); + let barging = receiver.clone(); + let mut first = Box::pin(receiver.recv()); + let mut second = Box::pin(second_receiver.recv()); + let (first_waker, first_wakes) = WakeCounter::new(); + let (second_waker, second_wakes) = WakeCounter::new(); + + assert!(poll_with(first.as_mut(), &first_waker).is_pending()); + assert!(poll_with(second.as_mut(), &second_waker).is_pending()); + sender.send(1).unwrap(); + assert_eq!(first_wakes.count(), 1); + assert_eq!(barging.try_recv(), Ok(1)); + assert!(poll_with(first.as_mut(), &first_waker).is_pending()); + + sender.send(2).unwrap(); + assert_eq!(first_wakes.count(), 1); + assert_eq!(second_wakes.count(), 1); + assert_eq!( + expect_ready(poll_with(second.as_mut(), &second_waker)), + Ok(2) + ); +} + +#[test] +fn notified_sender_that_loses_capacity_queues_behind_waiting_senders() { + let (sender, receiver) = mpmc::bounded(1); + sender.try_send(0).unwrap(); + let competing = sender.clone(); + let mut first = Box::pin(sender.send(1)); + let mut second = Box::pin(competing.send(2)); + let (first_waker, first_wakes) = WakeCounter::new(); + let (second_waker, second_wakes) = WakeCounter::new(); + + assert!(poll_with(first.as_mut(), &first_waker).is_pending()); + assert!(poll_with(second.as_mut(), &second_waker).is_pending()); + assert_eq!(receiver.try_recv(), Ok(0)); + assert_eq!(first_wakes.count(), 1); + sender.try_send(3).unwrap(); + assert!(poll_with(first.as_mut(), &first_waker).is_pending()); + + assert_eq!(receiver.try_recv(), Ok(3)); + assert_eq!(first_wakes.count(), 1); + assert_eq!(second_wakes.count(), 1); + expect_ready(poll_with(second.as_mut(), &second_waker)).unwrap(); + assert_eq!(receiver.try_recv(), Ok(2)); + assert_eq!(first_wakes.count(), 2); + expect_ready(poll_with(first.as_mut(), &first_waker)).unwrap(); + assert_eq!(receiver.try_recv(), Ok(1)); +} + +#[test] +fn cancelling_a_sender_after_capacity_is_taken_does_not_wake_next() { + let (sender, receiver) = mpmc::bounded(1); + sender.try_send(0).unwrap(); + let competing = sender.clone(); + let mut cancelled = Box::pin(sender.send(1)); + let mut waiting = Box::pin(competing.send(2)); + let (cancelled_waker, cancelled_wakes) = WakeCounter::new(); + let (waiting_waker, waiting_wakes) = WakeCounter::new(); + + assert!(poll_with(cancelled.as_mut(), &cancelled_waker).is_pending()); + assert!(poll_with(waiting.as_mut(), &waiting_waker).is_pending()); + assert_eq!(receiver.try_recv(), Ok(0)); + assert_eq!(cancelled_wakes.count(), 1); + assert_eq!(waiting_wakes.count(), 0); + sender.try_send(3).unwrap(); + drop(cancelled); + assert_eq!(waiting_wakes.count(), 0); + + assert_eq!(receiver.try_recv(), Ok(3)); + assert_eq!(waiting_wakes.count(), 1); + expect_ready(poll_with(waiting.as_mut(), &waiting_waker)).unwrap(); + assert_eq!(receiver.try_recv(), Ok(2)); + assert_eq!(receiver.try_recv(), Err(TryRecvError::Empty)); +} + +#[test] +fn bounded_cancelled_sender_notifies_next_sender_before_dropping_value() { + // A message destructor may depend on another blocked sender making progress. + struct WakesSeenOnDrop { + wakes: Arc, + seen: Arc, + } + + impl Drop for WakesSeenOnDrop { + fn drop(&mut self) { + self.seen.store(self.wakes.count(), Ordering::Relaxed); + } + } + + let (sender, receiver) = mpmc::bounded(1); + sender.try_send((0, None)).unwrap(); + let first_sender = sender.clone(); + let second_sender = sender.clone(); + let (cancelled_waker, cancelled_wakes) = WakeCounter::new(); + let (waiting_waker, waiting_wakes) = WakeCounter::new(); + let wakes_during_drop = Arc::new(AtomicUsize::new(usize::MAX)); + let observer = WakesSeenOnDrop { + wakes: waiting_wakes.clone(), + seen: wakes_during_drop.clone(), + }; + let mut cancelled = Box::pin(first_sender.send((1, Some(observer)))); + let mut waiting = Box::pin(second_sender.send((2, None))); + + assert!(poll_with(cancelled.as_mut(), &cancelled_waker).is_pending()); + assert!(poll_with(waiting.as_mut(), &waiting_waker).is_pending()); + assert_eq!(receiver.try_recv().unwrap().0, 0); + assert_eq!(cancelled_wakes.count(), 1); + assert_eq!(waiting_wakes.count(), 0); + drop(cancelled); + + assert_eq!(wakes_during_drop.load(Ordering::Relaxed), 1); + assert_eq!(waiting_wakes.count(), 1); + expect_ready(poll_with(waiting.as_mut(), &waiting_waker)).unwrap(); + assert_eq!(receiver.try_recv().unwrap().0, 2); + assert!(matches!(receiver.try_recv(), Err(TryRecvError::Empty))); +} + +#[test] +fn last_sender_wakes_every_pending_receiver() { + let (sender, receiver) = mpmc::unbounded::(); + let competing = receiver.clone(); + let mut first = Box::pin(receiver.recv()); + let mut second = Box::pin(competing.recv()); + let mut cancelled = Box::pin(receiver.recv()); + let (first_waker, first_wakes) = WakeCounter::new(); + let (second_waker, second_wakes) = WakeCounter::new(); + let (cancelled_waker, cancelled_wakes) = WakeCounter::new(); + assert!(poll_with(first.as_mut(), &first_waker).is_pending()); + assert!(poll_with(second.as_mut(), &second_waker).is_pending()); + assert!(poll_with(cancelled.as_mut(), &cancelled_waker).is_pending()); + + // Disconnection must handle both notified and still-linked registrations. + sender.send(1).unwrap(); + + drop(sender); + assert_eq!(first_wakes.count(), 1); + assert_eq!(second_wakes.count(), 1); + assert_eq!(cancelled_wakes.count(), 1); + drop(cancelled); + assert_eq!(expect_ready(poll_with(first.as_mut(), &first_waker)), Ok(1)); + assert_eq!( + expect_ready(poll_with(second.as_mut(), &second_waker)), + Err(RecvError::Disconnected) + ); +} + +#[test] +fn last_receiver_wakes_every_pending_sender() { + let (sender, receiver) = mpmc::bounded(1); + sender.try_send(0).unwrap(); + let competing = sender.clone(); + let mut first = Box::pin(sender.send(1)); + let mut second = Box::pin(competing.send(2)); + let mut cancelled = Box::pin(sender.send(3)); + let (first_waker, first_wakes) = WakeCounter::new(); + let (second_waker, second_wakes) = WakeCounter::new(); + let (cancelled_waker, cancelled_wakes) = WakeCounter::new(); + assert!(poll_with(first.as_mut(), &first_waker).is_pending()); + assert!(poll_with(second.as_mut(), &second_waker).is_pending()); + assert!(poll_with(cancelled.as_mut(), &cancelled_waker).is_pending()); + assert_eq!(receiver.try_recv(), Ok(0)); + + drop(receiver); + assert_eq!(first_wakes.count(), 1); + assert_eq!(second_wakes.count(), 1); + assert_eq!(cancelled_wakes.count(), 1); + drop(cancelled); + assert_eq!( + expect_ready(poll_with(first.as_mut(), &first_waker)) + .unwrap_err() + .into_inner(), + 1 + ); + assert_eq!( + expect_ready(poll_with(second.as_mut(), &second_waker)) + .unwrap_err() + .into_inner(), + 2 + ); +} diff --git a/tests-integration/tests/mpsc_test/callbacks.rs b/tests-integration/tests/mpsc_test/callbacks.rs index 772c3683..d3174f68 100644 --- a/tests-integration/tests/mpsc_test/callbacks.rs +++ b/tests-integration/tests/mpsc_test/callbacks.rs @@ -16,7 +16,6 @@ // under the License. use std::sync::Arc; -use std::sync::Mutex; use std::sync::atomic::AtomicBool; use std::sync::atomic::AtomicUsize; use std::sync::atomic::Ordering; @@ -35,8 +34,6 @@ use tests_integration::poll_with; use tests_integration::waker_on_drop; use tests_integration::waker_on_wake; -use super::support::waker_on_clone; - struct HoldSender { _sender: S, } @@ -81,107 +78,6 @@ fn unbounded_receiver_drop_releases_registered_waker() { assert!(retained.upgrade().is_none()); } -#[test] -fn bounded_send_rechecks_capacity_freed_by_waker_clone() { - assert_completes_without_deadlock(|| { - let (tx, rx) = mpsc::bounded(1); - tx.try_send(1).unwrap(); - let receiver = Arc::new(Mutex::new(rx)); - let received = AtomicBool::new(false); - let waker = waker_on_clone({ - let receiver = receiver.clone(); - move || { - if !received.swap(true, Ordering::Relaxed) { - assert_eq!(receiver.lock().unwrap().try_recv(), Ok(1)); - } - } - }); - assert_eq!( - poll_with(Box::pin(tx.send(2)).as_mut(), &waker), - Poll::Ready(Ok(())) - ); - assert_eq!(receiver.lock().unwrap().try_recv(), Ok(2)); - }); -} - -#[cfg(panic = "unwind")] -#[test] -fn bounded_send_returns_capacity_when_an_unused_waker_panics_on_drop() { - use std::task::RawWaker; - use std::task::RawWakerVTable; - - struct Callbacks { - receiver: Mutex>, - drop_panics: AtomicBool, - } - - unsafe fn clone(data: *const ()) -> RawWaker { - let pointer = data.cast::(); - // SAFETY: The input waker owns a live Arc. The returned clone gains its own reference. - unsafe { - assert_eq!((*pointer).receiver.lock().unwrap().try_recv(), Ok(1)); - Arc::increment_strong_count(pointer); - } - RawWaker::new(data, &VTABLE) - } - - unsafe fn release(data: *const ()) { - // SAFETY: Consumes this waker's Arc reference, including if the callback unwinds. - let callbacks = unsafe { Arc::from_raw(data.cast::()) }; - assert!( - !callbacks.drop_panics.swap(false, Ordering::Relaxed), - "unused cloned waker panicked on drop" - ); - } - - // A raw vtable is needed to run callbacks for cloning and dropping each waker reference. - static VTABLE: RawWakerVTable = RawWakerVTable::new(clone, release, |_| {}, release); - - let (tx, rx) = mpsc::bounded(1); - tx.try_send(1).unwrap(); - let callbacks = Arc::new(Callbacks { - receiver: Mutex::new(rx), - drop_panics: AtomicBool::new(true), - }); - let data = Arc::into_raw(callbacks.clone()).cast(); - // SAFETY: Every waker owns an Arc reference. All callbacks preserve ownership and use only - // synchronized state; wake_by_ref does not touch the reference count. - let waker = unsafe { Waker::from_raw(RawWaker::new(data, &VTABLE)) }; - - let mut send = Box::pin(tx.send(2)); - assert!( - std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { - poll_with(send.as_mut(), &waker) - })) - .is_err() - ); - drop(send); - let mut receiver = callbacks.receiver.lock().unwrap(); - assert_eq!(receiver.try_recv(), Err(TryRecvError::Empty)); - tx.try_send(3) - .expect("unwinding must return acquired capacity"); - assert_eq!(receiver.try_recv(), Ok(3)); -} - -#[test] -fn receive_rechecks_messages_sent_by_waker_clone() { - assert_completes_without_deadlock(|| { - let (tx, mut rx) = mpsc::bounded(1); - let waker = waker_on_clone(move || tx.try_send(7).unwrap()); - assert_eq!( - poll_with(Box::pin(rx.recv()).as_mut(), &waker), - Poll::Ready(Ok(7)) - ); - - let (tx, mut rx) = mpsc::unbounded(); - let waker = waker_on_clone(move || tx.send(7).unwrap()); - assert_eq!( - poll_with(Box::pin(rx.recv()).as_mut(), &waker), - Poll::Ready(Ok(7)) - ); - }); -} - #[test] fn wake_callbacks_can_send_into_the_same_channel() { assert_completes_without_deadlock(|| { diff --git a/tests-integration/tests/mpsc_test/main.rs b/tests-integration/tests/mpsc_test/main.rs index 057c792c..35b21b05 100644 --- a/tests-integration/tests/mpsc_test/main.rs +++ b/tests-integration/tests/mpsc_test/main.rs @@ -31,7 +31,6 @@ mod backpressure; mod callbacks; mod concurrency; mod reservation; -mod support; #[test] fn unbounded_try_recv_preserves_order_and_reports_state() { diff --git a/tests-integration/tests/mpsc_test/reservation.rs b/tests-integration/tests/mpsc_test/reservation.rs index 86ab54ed..b16f7efe 100644 --- a/tests-integration/tests/mpsc_test/reservation.rs +++ b/tests-integration/tests/mpsc_test/reservation.rs @@ -104,10 +104,13 @@ fn closing_after_a_grant_returns_the_unsent_message() { tx.try_send(String::from("queued")).unwrap(); let mut send = Box::pin(tx.send(String::from("unsent"))); let mut reservation = Box::pin(tx.reserve()); + let mut cancelled = Box::pin(tx.reserve()); assert!(poll_once(send.as_mut()).is_pending()); assert!(poll_once(reservation.as_mut()).is_pending()); + assert!(poll_once(cancelled.as_mut()).is_pending()); assert_eq!(rx.try_recv().unwrap(), "queued"); drop(rx); + drop(cancelled); let error = expect_ready(poll_once(send.as_mut())).unwrap_err(); assert_eq!(error.into_inner(), "unsent"); assert!(expect_ready(poll_once(reservation.as_mut())).is_err()); diff --git a/tests-integration/tests/mpsc_test/support.rs b/tests-integration/tests/mpsc_test/support.rs deleted file mode 100644 index 9b0340b0..00000000 --- a/tests-integration/tests/mpsc_test/support.rs +++ /dev/null @@ -1,47 +0,0 @@ -// Licensed to the Apache Software Foundation (ASF) under one -// or more contributor license agreements. See the NOTICE file -// distributed with this work for additional information -// regarding copyright ownership. The ASF licenses this file -// to you under the Apache License, Version 2.0 (the -// "License"); you may not use this file except in compliance -// with the License. You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, -// software distributed under the License is distributed on an -// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -// KIND, either express or implied. See the License for the -// specific language governing permissions and limitations -// under the License. - -use std::sync::Arc; -use std::task::RawWaker; -use std::task::RawWakerVTable; -use std::task::Waker; - -// RawWaker is needed only to exercise clone callbacks, which the safe Wake trait cannot override. -pub fn waker_on_clone(callback: impl Fn() + Send + Sync + 'static) -> Waker { - struct OnClone(Box); - - unsafe fn clone(data: *const ()) -> RawWaker { - let pointer = data.cast::(); - // SAFETY: The input waker owns a live Arc; the returned waker gains its own reference. - unsafe { - ((*pointer).0)(); - Arc::increment_strong_count(pointer); - } - RawWaker::new(data, &VTABLE) - } - - unsafe fn release(data: *const ()) { - // SAFETY: Consumes the one Arc reference owned by this waker. - drop(unsafe { Arc::from_raw(data.cast::()) }); - } - - static VTABLE: RawWakerVTable = RawWakerVTable::new(clone, release, |_| {}, release); - let pointer = Arc::into_raw(Arc::new(OnClone(Box::new(callback)))).cast(); - // SAFETY: Each waker owns one Arc; its callback is Send + Sync and all vtable operations - // preserve that ownership. wake_by_ref borrows the reference without changing it. - unsafe { Waker::from_raw(RawWaker::new(pointer, &VTABLE)) } -} diff --git a/tests-integration/tests/phaser_test.rs b/tests-integration/tests/phaser_test.rs index 09c5bf20..615672c3 100644 --- a/tests-integration/tests/phaser_test.rs +++ b/tests-integration/tests/phaser_test.rs @@ -15,52 +15,18 @@ // specific language governing permissions and limitations // under the License. -use std::future::Future; use std::panic; use std::sync::Arc; -use std::sync::atomic::AtomicUsize; -use std::sync::atomic::Ordering; -use std::task::Context; use std::task::Poll; -use std::task::Wake; use std::task::Waker; use asyncband::blocking::FutureExt; use asyncband::phaser::Phaser; - -fn poll_once(future: std::pin::Pin<&mut F>) -> Poll { - future.poll(&mut Context::from_waker(Waker::noop())) -} - -struct CountWake(AtomicUsize); - -impl Wake for CountWake { - fn wake(self: Arc) { - self.0.fetch_add(1, Ordering::Relaxed); - } -} - -struct PanicWake; - -impl Wake for PanicWake { - fn wake(self: Arc) { - panic!("wake failed"); - } -} - -struct CloseOnDrop(Phaser); - -impl Wake for CloseOnDrop { - fn wake(self: Arc) { - self.0.close(); - } -} - -impl Drop for CloseOnDrop { - fn drop(&mut self) { - self.0.close(); - } -} +use tests_integration::PanicWake; +use tests_integration::WakeCounter; +use tests_integration::poll_once; +use tests_integration::poll_with; +use tests_integration::waker_on_drop; #[test] fn batch_registration_joins_one_observed_phase() { @@ -78,16 +44,11 @@ fn batch_registration_joins_one_observed_phase() { assert_eq!(participants.len(), 1); assert_eq!(phaser.phase(), observed); - let counter = Arc::new(CountWake(AtomicUsize::new(0))); - let waker = Waker::from(counter.clone()); + let (waker, counter) = WakeCounter::new(); let mut wait = Box::pin(first.wait()); - assert!( - wait.as_mut() - .poll(&mut Context::from_waker(&waker)) - .is_pending() - ); + assert!(poll_with(wait.as_mut(), &waker).is_pending()); drop(participants); - assert_eq!(counter.0.load(Ordering::Relaxed), 1); + assert_eq!(counter.count(), 1); assert_eq!(poll_once(wait.as_mut()), Poll::Ready(Ok(phaser.phase()))); assert_ne!(phaser.phase(), observed); assert_eq!(phaser.registered_parties(), 2); @@ -342,20 +303,18 @@ fn registration_after_last_participant_drop_joins_the_advanced_phase() { fn observer_wait_is_cancel_safe_and_does_not_participate() { let phaser = Phaser::new(); let phase = phaser.phase(); - let counter = Arc::new(CountWake(AtomicUsize::new(0))); - let waker = Waker::from(Arc::clone(&counter)); - let mut context = Context::from_waker(&waker); + let (waker, counter) = WakeCounter::new(); { let mut wait = Box::pin(phaser.wait(phase)); - assert_eq!(Future::poll(wait.as_mut(), &mut context), Poll::Pending); + assert_eq!(poll_with(wait.as_mut(), &waker), Poll::Pending); assert_eq!(phaser.registered_parties(), 0); } assert_eq!(phaser.registered_parties(), 0); let participant = phaser.register_one().unwrap(); drop(participant); - assert_eq!(counter.0.load(Ordering::Relaxed), 0); + assert_eq!(counter.count(), 0); } #[test] @@ -363,33 +322,26 @@ fn advancing_a_phase_wakes_every_registered_waiter_once() { let phaser = Phaser::new(); let observed = phaser.phase(); let participant = phaser.register_one().unwrap(); - let first_counter = Arc::new(CountWake(AtomicUsize::new(0))); - let second_counter = Arc::new(CountWake(AtomicUsize::new(0))); - let first_waker = Waker::from(Arc::clone(&first_counter)); - let second_waker = Waker::from(Arc::clone(&second_counter)); - let mut first_context = Context::from_waker(&first_waker); - let mut second_context = Context::from_waker(&second_waker); + let (first_waker, first_counter) = WakeCounter::new(); + let (second_waker, second_counter) = WakeCounter::new(); let mut first_wait = Box::pin(phaser.wait(observed)); let mut second_wait = Box::pin(phaser.wait(observed)); + assert_eq!(poll_with(first_wait.as_mut(), &first_waker), Poll::Pending); assert_eq!( - Future::poll(first_wait.as_mut(), &mut first_context), - Poll::Pending - ); - assert_eq!( - Future::poll(second_wait.as_mut(), &mut second_context), + poll_with(second_wait.as_mut(), &second_waker), Poll::Pending ); drop(participant); - assert_eq!(first_counter.0.load(Ordering::Relaxed), 1); - assert_eq!(second_counter.0.load(Ordering::Relaxed), 1); + assert_eq!(first_counter.count(), 1); + assert_eq!(second_counter.count(), 1); assert!(matches!( - Future::poll(first_wait.as_mut(), &mut first_context), + poll_with(first_wait.as_mut(), &first_waker), Poll::Ready(_) )); assert!(matches!( - Future::poll(second_wait.as_mut(), &mut second_context), + poll_with(second_wait.as_mut(), &second_waker), Poll::Ready(_) )); } @@ -398,23 +350,17 @@ fn advancing_a_phase_wakes_every_registered_waiter_once() { fn repolling_updates_the_task_that_will_be_notified() { let phaser = Phaser::new(); let participant = phaser.register_one().unwrap(); - let first = Arc::new(CountWake(AtomicUsize::new(0))); - let second = Arc::new(CountWake(AtomicUsize::new(0))); - let first_waker = Waker::from(first.clone()); - let second_waker = Waker::from(second.clone()); + let (first_waker, first) = WakeCounter::new(); + let (second_waker, second) = WakeCounter::new(); let mut wait = Box::pin(phaser.wait(phaser.phase())); for waker in [&first_waker, &first_waker, &second_waker, &second_waker] { - assert!( - wait.as_mut() - .poll(&mut Context::from_waker(waker)) - .is_pending() - ); + assert!(poll_with(wait.as_mut(), waker).is_pending()); } drop(participant); - assert_eq!(first.0.load(Ordering::Relaxed), 0); - assert_eq!(second.0.load(Ordering::Relaxed), 1); + assert_eq!(first.count(), 0); + assert_eq!(second.count(), 1); assert_eq!(poll_once(wait.as_mut()), Poll::Ready(Ok(phaser.phase()))); } @@ -423,35 +369,28 @@ fn cancelling_a_woken_waiter_does_not_unregister_a_next_phase_waiter() { let phaser = Phaser::new(); let phase0 = phaser.phase(); let participant = phaser.register_one().unwrap(); - let stale_counter = Arc::new(CountWake(AtomicUsize::new(0))); - let stale_waker = Waker::from(Arc::clone(&stale_counter)); - let mut stale_context = Context::from_waker(&stale_waker); + let (stale_waker, stale_counter) = WakeCounter::new(); let mut stale_wait = Box::pin(phaser.wait(phase0)); - assert_eq!( - Future::poll(stale_wait.as_mut(), &mut stale_context), - Poll::Pending - ); + assert_eq!(poll_with(stale_wait.as_mut(), &stale_waker), Poll::Pending); drop(participant); let phase1 = phaser.phase(); assert_ne!(phase1, phase0); - assert_eq!(stale_counter.0.load(Ordering::Relaxed), 1); + assert_eq!(stale_counter.count(), 1); let participant = phaser.register_one().unwrap(); - let current_counter = Arc::new(CountWake(AtomicUsize::new(0))); - let current_waker = Waker::from(Arc::clone(¤t_counter)); - let mut current_context = Context::from_waker(¤t_waker); + let (current_waker, current_counter) = WakeCounter::new(); let mut current_wait = Box::pin(phaser.wait(phase1)); assert_eq!( - Future::poll(current_wait.as_mut(), &mut current_context), + poll_with(current_wait.as_mut(), ¤t_waker), Poll::Pending ); drop(stale_wait); drop(participant); - assert_eq!(current_counter.0.load(Ordering::Relaxed), 1); + assert_eq!(current_counter.count(), 1); assert!(matches!( - Future::poll(current_wait.as_mut(), &mut current_context), + poll_with(current_wait.as_mut(), ¤t_waker), Poll::Ready(_) )); } @@ -463,20 +402,15 @@ fn panicking_waker_does_not_lose_a_pending_phase() { let mut first = phaser.register_one().unwrap(); let mut second = phaser.register_one().unwrap(); let panic_waker = Waker::from(Arc::new(PanicWake)); - let mut panic_context = Context::from_waker(&panic_waker); let mut observer = Box::pin(phaser.wait(phase0)); - assert_eq!( - Future::poll(observer.as_mut(), &mut panic_context), - Poll::Pending - ); + assert_eq!(poll_with(observer.as_mut(), &panic_waker), Poll::Pending); assert_eq!(first.arrive().unwrap(), phase0); - let polling_waker = Waker::from(Arc::new(CountWake(AtomicUsize::new(0)))); - let mut polling_context = Context::from_waker(&polling_waker); + let (polling_waker, _) = WakeCounter::new(); let mut wait = Box::pin(second.wait()); let result = panic::catch_unwind(panic::AssertUnwindSafe(|| { - Future::poll(wait.as_mut(), &mut polling_context) + poll_with(wait.as_mut(), &polling_waker) })); assert!(result.is_err()); @@ -580,18 +514,16 @@ fn closing_wakes_all_waiters_once_and_rejects_new_obligations() { let mut first = phaser.register_one().unwrap(); let second = phaser.register_one().unwrap(); let observed = first.arrive().unwrap(); - let counter = Arc::new(CountWake(AtomicUsize::new(0))); - let waker = Waker::from(counter.clone()); - let mut context = Context::from_waker(&waker); + let (waker, counter) = WakeCounter::new(); let mut observer = Box::pin(phaser.wait(observed)); let mut wait = Box::pin(first.wait()); - assert!(observer.as_mut().poll(&mut context).is_pending()); - assert!(wait.as_mut().poll(&mut context).is_pending()); + assert!(poll_with(observer.as_mut(), &waker).is_pending()); + assert!(poll_with(wait.as_mut(), &waker).is_pending()); phaser.clone().close(); phaser.close(); assert!(phaser.is_closed()); - assert_eq!(counter.0.load(Ordering::Relaxed), 2); + assert_eq!(counter.count(), 2); assert!(matches!(poll_once(observer.as_mut()), Poll::Ready(Err(_)))); assert!(matches!(poll_once(wait.as_mut()), Poll::Ready(Err(_)))); drop(wait); @@ -637,26 +569,15 @@ fn close_survives_a_panicking_waker_and_notifies_other_waiters() { let phaser = Phaser::new(); let observed = phaser.phase(); let panic_waker = Waker::from(Arc::new(PanicWake)); - let counter = Arc::new(CountWake(AtomicUsize::new(0))); - let count_waker = Waker::from(counter.clone()); + let (count_waker, counter) = WakeCounter::new(); let mut first = Box::pin(phaser.wait(observed)); let mut second = Box::pin(phaser.wait(observed)); - assert!( - first - .as_mut() - .poll(&mut Context::from_waker(&panic_waker)) - .is_pending() - ); - assert!( - second - .as_mut() - .poll(&mut Context::from_waker(&count_waker)) - .is_pending() - ); + assert!(poll_with(first.as_mut(), &panic_waker).is_pending()); + assert!(poll_with(second.as_mut(), &count_waker).is_pending()); assert!(panic::catch_unwind(|| phaser.close()).is_err()); assert!(phaser.is_closed()); - assert_eq!(counter.0.load(Ordering::Relaxed), 1); + assert_eq!(counter.count(), 1); assert!(matches!(poll_once(first.as_mut()), Poll::Ready(Err(_)))); assert!(matches!(poll_once(second.as_mut()), Poll::Ready(Err(_)))); } @@ -664,38 +585,31 @@ fn close_survives_a_panicking_waker_and_notifies_other_waiters() { #[test] fn replacing_a_waiter_waker_can_close_the_phaser_from_its_destructor() { let phaser = Phaser::new(); - let waker = Waker::from(Arc::new(CloseOnDrop(phaser.clone()))); + let waker = waker_on_drop({ + let phaser = phaser.clone(); + move || phaser.close() + }); let mut wait = Box::pin(phaser.wait(phaser.phase())); - assert!( - wait.as_mut() - .poll(&mut Context::from_waker(&waker)) - .is_pending() - ); + assert!(poll_with(wait.as_mut(), &waker).is_pending()); drop(waker); assert!(!phaser.is_closed()); - let counter = Arc::new(CountWake(AtomicUsize::new(0))); - let replacement = Waker::from(counter.clone()); - assert!( - wait.as_mut() - .poll(&mut Context::from_waker(&replacement)) - .is_pending() - ); + let (replacement, counter) = WakeCounter::new(); + assert!(poll_with(wait.as_mut(), &replacement).is_pending()); assert!(phaser.is_closed()); - assert_eq!(counter.0.load(Ordering::Relaxed), 1); + assert_eq!(counter.count(), 1); assert!(matches!(poll_once(wait.as_mut()), Poll::Ready(Err(_)))); } #[test] fn cancelling_a_waiter_can_close_the_phaser_from_its_waker_destructor() { let phaser = Phaser::new(); - let waker = Waker::from(Arc::new(CloseOnDrop(phaser.clone()))); + let waker = waker_on_drop({ + let phaser = phaser.clone(); + move || phaser.close() + }); let mut wait = Box::pin(phaser.wait(phaser.phase())); - assert!( - wait.as_mut() - .poll(&mut Context::from_waker(&waker)) - .is_pending() - ); + assert!(poll_with(wait.as_mut(), &waker).is_pending()); drop(waker); assert!(!phaser.is_closed()); diff --git a/xtask/src/main.rs b/xtask/src/main.rs index 51fb9866..f2733663 100644 --- a/xtask/src/main.rs +++ b/xtask/src/main.rs @@ -121,6 +121,7 @@ impl CommandMiri { &["--test", "unsafe_paths_test"], )); run_command(make_miri_cmd("tests-integration", &["--test", "mpsc_test"])); + run_command(make_miri_cmd("tests-integration", &["--test", "mpmc_test"])); run_command(make_miri_cmd( "tests-integration", &["--test", "phaser_test"],