From cc1adbd168ea8f52473cf617785ea4c835c89ef9 Mon Sep 17 00:00:00 2001 From: Orthur Date: Mon, 14 Sep 2026 05:08:35 -0400 Subject: [PATCH 01/11] perf(mpmc): coordinate queue values and waiters under one lock --- asyncband/src/internal/mod.rs | 7 +- asyncband/src/internal/semaphore.rs | 4 +- asyncband/src/mpmc/bounded.rs | 17 +- asyncband/src/mpmc/queue.rs | 348 +++++++++++++++++++------ asyncband/src/mpmc/unbounded.rs | 10 +- benchmarks/asyncband/mpmc/bounded.rs | 41 +++ benchmarks/asyncband/mpmc/mod.rs | 3 + benchmarks/asyncband/mpmc/unbounded.rs | 49 ++++ 8 files changed, 380 insertions(+), 99 deletions(-) diff --git a/asyncband/src/internal/mod.rs b/asyncband/src/internal/mod.rs index 7998e25a..634e2e95 100644 --- a/asyncband/src/internal/mod.rs +++ b/asyncband/src/internal/mod.rs @@ -99,14 +99,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; diff --git a/asyncband/src/internal/semaphore.rs b/asyncband/src/internal/semaphore.rs index ca4c3019..48829769 100644 --- a/asyncband/src/internal/semaphore.rs +++ b/asyncband/src/internal/semaphore.rs @@ -204,7 +204,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() { @@ -213,7 +213,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 = vec![]; diff --git a/asyncband/src/mpmc/bounded.rs b/asyncband/src/mpmc/bounded.rs index 3bcd73f9..b42f0b1c 100644 --- a/asyncband/src/mpmc/bounded.rs +++ b/asyncband/src/mpmc/bounded.rs @@ -29,9 +29,9 @@ 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. +/// 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. /// /// # Panics /// @@ -84,9 +84,10 @@ 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. + /// returned `Pending` has not sent the value. If this call was woken for capacity that is still + /// free, cancelling it wakes the next waiting sender before `value` is dropped. 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 +138,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. If this call was woken for a value that + /// is still queued, cancelling it wakes the next waiting receiver instead. 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..2f0b700b 100644 --- a/asyncband/src/mpmc/queue.rs +++ b/asyncband/src/mpmc/queue.rs @@ -16,31 +16,136 @@ // 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; +use crate::internal::waker_batch::WakerBatch; -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. Waker callbacks and 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); + notify_one(&mut self.recv_waiters) + } + + /// 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, notify_one(&mut self.send_waiters))) + } 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, +} + +fn notify_one(waiters: &mut WaitList) -> Option { + let (_, waiter) = waiters.unlink_first_waiter(|_| true)?; + let Waiter::Waiting(waker) = mem::replace(waiter, Waiter::Notified) else { + unreachable!("only waiting operations remain linked"); + }; + Some(waker) +} + +fn notify_all(waiters: &mut WaitList) -> WakerBatch { + let mut wakers = WakerBatch::new(); + while let Some(waker) = notify_one(waiters) { + wakers.push(waker); + } + wakers +} + +fn remove_waiter(waiters: &mut WaitList, id: WaiterId) -> Waiter { + // Unlinking is idempotent, so notified waiters are removed the same way as linked ones. + waiters.unlink_waiter(id, |_| true); + waiters.remove_unlinked_waiter(id) +} + +fn wake(waker: Option) { + if let Some(waker) = waker { + waker.wake(); + } +} + +enum Registration { + /// The operation is queued; drop the replaced waker after releasing the lock. + Registered(Option), + /// Clone the current waker outside the lock, then retry. + NeedsWaker, +} + +/// 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. +fn register( + waiters: &mut WaitList, + id: &mut Option, + current: &Waker, + cloned: &mut Option, +) -> Registration { + if let Some(queued) = *id { + if let Waiter::Waiting(waker) = waiters.waiter_mut(queued) { + if waker.will_wake(current) { + return Registration::Registered(None); + } + return match cloned.take() { + Some(new) => Registration::Registered(Some(mem::replace(waker, new))), + None => Registration::NeedsWaker, + }; + } + } + let Some(waker) = cloned.take() else { + return Registration::NeedsWaker; + }; + if let Some(notified) = id.take() { + // The notification already took this node's waker, so nothing is retired. + remove_waiter(waiters, notified); + } + *id = Some(waiters.push_back(Waiter::Waiting(waker))); + Registration::Registered(None) } impl Shared { @@ -56,69 +161,66 @@ 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 wakers = { let mut state = self.state.lock(); state.senders -= 1; - state.senders == 0 + if state.senders != 0 { + return; + } + // Woken receivers drain buffered values before they observe disconnection. + notify_all(&mut state.recv_waiters) }; - if is_last { - self.recv_waiters.notify_all(); - } + wake_all(wakers.into_iter()); } 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, wakers) = { 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), + notify_all(&mut state.send_waiters), + ) }; - 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(wakers.into_iter()); 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); - } - self.recv_waiters.release_if_nonempty(1); + state.push(value) + }; + wake(waker); Ok(()) } @@ -130,24 +232,15 @@ 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()?; + wake(waker); Ok(value) } @@ -159,7 +252,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 +260,148 @@ 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"); + // The first poll follows a full `try_send` and most likely registers, so it clones before + // locking. A later poll clones only when it must store a new waker: its task changed, or a + // notification took the stored one. + let mut cloned = self.waiter.is_none().then(|| cx.waker().clone()); 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 { + Err(self.take_value()) + } else if state.has_capacity() { + Ok(state.push(self.take_value())) + } else { + match register( + &mut state.send_waiters, + &mut self.waiter, + cx.waker(), + &mut cloned, + ) { + Registration::Registered(replaced) => { + drop(state); + drop((replaced, cloned)); + return Poll::Pending; + } + Registration::NeedsWaker => { + drop(state); + cloned = Some(cx.waker().clone()); + continue; + } } - 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; - } + let retired = self + .waiter + .take() + .map(|id| remove_waiter(&mut state.send_waiters, id)); + drop(state); + // Deliver the notification before running waker destructors, which may panic. + let result = outcome.map(wake).map_err(SendError::new); + drop((retired, cloned)); + return 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(); + let retired = remove_waiter(&mut state.send_waiters, 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() { + notify_one(&mut state.send_waiters) + } else { + None + }; + (retired, waker) + }; + wake(waker); + 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> { + // The first poll follows an empty `try_recv` and most likely registers, so it clones before + // locking. A later poll clones only when it must store a new waker: its task changed, or a + // notification took the stored one. + let mut cloned = self.waiter.is_none().then(|| cx.waker().clone()); 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(); + let outcome = match state.pop() { + Ok(popped) => Ok(popped), + Err(TryRecvError::Disconnected) => Err(RecvError::Disconnected), + Err(TryRecvError::Empty) => match register( + &mut state.recv_waiters, + &mut self.waiter, + cx.waker(), + &mut cloned, + ) { + Registration::Registered(replaced) => { + drop(state); + drop((replaced, cloned)); + return Poll::Pending; + } + Registration::NeedsWaker => { + drop(state); + cloned = Some(cx.waker().clone()); + continue; + } + }, + }; + let retired = self + .waiter + .take() + .map(|id| remove_waiter(&mut state.recv_waiters, id)); + drop(state); + // Deliver the notification before running waker destructors, which may panic. + let result = outcome.map(|(value, waker)| { + wake(waker); + value + }); + drop((retired, cloned)); + return 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(); + let retired = remove_waiter(&mut state.recv_waiters, 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() { + notify_one(&mut state.recv_waiters) + } else { + None + }; + (retired, waker) + }; + wake(waker); + drop(retired); + } +} diff --git a/asyncband/src/mpmc/unbounded.rs b/asyncband/src/mpmc/unbounded.rs index 7e1db12a..048121c1 100644 --- a/asyncband/src/mpmc/unbounded.rs +++ b/asyncband/src/mpmc/unbounded.rs @@ -28,9 +28,9 @@ 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. +/// 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. pub fn unbounded() -> (UnboundedSender, UnboundedReceiver) { let shared = Arc::new(Shared::unbounded()); ( @@ -119,8 +119,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. If this call was woken for a value that + /// is still queued, cancelling it wakes the next waiting receiver instead. pub async fn recv(&self) -> Result { self.shared.recv().await } 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)); +} From e3ffa4402bb3a336d0afe477b34219bf1d9ec0ba Mon Sep 17 00:00:00 2001 From: Orthur Date: Mon, 14 Sep 2026 05:08:35 -0400 Subject: [PATCH 02/11] test(mpmc): cover disconnection wake-ups and callback reentrancy --- tests-integration/src/lib.rs | 30 ++ tests-integration/tests/mpmc_test.rs | 441 ------------------ .../tests/mpmc_test/callbacks.rs | 124 +++++ .../tests/mpmc_test/concurrency.rs | 107 +++++ tests-integration/tests/mpmc_test/main.rs | 143 ++++++ .../tests/mpmc_test/notification.rs | 216 +++++++++ .../tests/mpsc_test/callbacks.rs | 3 +- tests-integration/tests/mpsc_test/main.rs | 1 - tests-integration/tests/mpsc_test/support.rs | 47 -- xtask/src/main.rs | 1 + 10 files changed, 622 insertions(+), 491 deletions(-) delete mode 100644 tests-integration/tests/mpmc_test.rs create mode 100644 tests-integration/tests/mpmc_test/callbacks.rs create mode 100644 tests-integration/tests/mpmc_test/concurrency.rs create mode 100644 tests-integration/tests/mpmc_test/main.rs create mode 100644 tests-integration/tests/mpmc_test/notification.rs delete mode 100644 tests-integration/tests/mpsc_test/support.rs diff --git a/tests-integration/src/lib.rs b/tests-integration/src/lib.rs index 00f2045e..4862e5f6 100644 --- a/tests-integration/src/lib.rs +++ b/tests-integration/src/lib.rs @@ -24,6 +24,8 @@ use std::sync::atomic::AtomicUsize; use std::sync::atomic::Ordering; use std::task::Context; use std::task::Poll; +use std::task::RawWaker; +use std::task::RawWakerVTable; use std::task::Wake; use std::task::Waker; @@ -122,6 +124,34 @@ pub fn waker_on_drop(callback: impl FnOnce() + Send + 'static) -> Waker { Waker::from(Arc::new(OnDrop(Mutex::new(Some(Box::new(callback)))))) } +// 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: `Waker::clone` borrows the input waker, whose reference keeps the Arc alive while + // the callback runs. Taking the new reference last leaves the count unchanged if the + // callback panics; the returned waker owns that 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)) } +} + pub fn assert_completes_without_deadlock(test: impl FnOnce() + Send + 'static) { let (finished_tx, finished_rx) = std::sync::mpsc::channel(); let worker = std::thread::spawn(move || { 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..1ddc4b3f --- /dev/null +++ b/tests-integration/tests/mpmc_test/callbacks.rs @@ -0,0 +1,124 @@ +// 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_clone; +use tests_integration::waker_on_drop; + +// A registering poll clones its waker before taking the queue lock, so a value that the clone +// callback sends is already queued when the same poll looks for one. +#[test] +fn receive_registration_allows_a_waker_clone_to_send() { + assert_completes_without_deadlock(|| { + let (sender, receiver) = mpmc::unbounded(); + let reentrant = sender.clone(); + let waker = waker_on_clone(move || reentrant.send(1).unwrap()); + let mut recv = Box::pin(receiver.recv()); + + assert_eq!(expect_ready(poll_with(recv.as_mut(), &waker)), Ok(1)); + drop(sender); + }); +} + +#[test] +fn send_registration_allows_a_waker_clone_to_receive() { + assert_completes_without_deadlock(|| { + let (sender, receiver) = mpmc::bounded(1); + sender.try_send(0).unwrap(); + let reentrant = receiver.clone(); + let waker = waker_on_clone(move || assert_eq!(reentrant.try_recv(), Ok(0))); + let mut send = Box::pin(sender.send(1)); + + expect_ready(poll_with(send.as_mut(), &waker)).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..5d008504 --- /dev/null +++ b/tests-integration/tests/mpmc_test/main.rs @@ -0,0 +1,143 @@ +// 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; +} + +impl Receiver for mpmc::BoundedReceiver { + fn recv(&self) -> impl Future> + Send { + self.recv() + } +} + +impl Receiver for mpmc::UnboundedReceiver { + fn recv(&self) -> impl Future> + Send { + self.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..fd8db23a --- /dev/null +++ b/tests-integration/tests/mpmc_test/notification.rs @@ -0,0 +1,216 @@ +// 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; + +fn send_wakes_only_the_first_receiver(receiver: impl Receiver, send: impl FnOnce()) { + 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(); + + assert_eq!(first_wakes.count(), 1); + assert_eq!(second_wakes.count(), 0); + 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); +} + +#[test] +fn bounded_send_wakes_only_the_first_receiver() { + let (sender, receiver) = mpmc::bounded(2); + send_wakes_only_the_first_receiver(receiver, || sender.try_send(1).unwrap()); +} + +#[test] +fn unbounded_send_wakes_only_the_first_receiver() { + let (sender, receiver) = mpmc::unbounded(); + send_wakes_only_the_first_receiver(receiver, || sender.send(1).unwrap()); +} + +fn cancelled_notified_receiver_wakes_next_receiver( + receiver: impl Receiver, + send: impl FnOnce(), +) { + let competing = receiver.clone(); + let mut cancelled = Box::pin(receiver.recv()); + let mut waiting = Box::pin(competing.recv()); + 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()); + send(); + 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(waiting.as_mut(), &waiting_waker)), + Ok(1) + ); +} + +#[test] +fn bounded_cancelled_notified_receiver_wakes_next_receiver() { + let (sender, receiver) = mpmc::bounded(1); + cancelled_notified_receiver_wakes_next_receiver(receiver, || sender.try_send(1).unwrap()); +} + +#[test] +fn unbounded_cancelled_notified_receiver_wakes_next_receiver() { + let (sender, receiver) = mpmc::unbounded(); + cancelled_notified_receiver_wakes_next_receiver(receiver, || sender.send(1).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 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 (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()); + + drop(sender); + assert_eq!(first_wakes.count(), 1); + assert_eq!(second_wakes.count(), 1); + assert_eq!( + expect_ready(poll_with(first.as_mut(), &first_waker)), + Err(RecvError::Disconnected) + ); + 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 (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()); + + drop(receiver); + assert_eq!(first_wakes.count(), 1); + assert_eq!(second_wakes.count(), 1); + 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..e5452b2a 100644 --- a/tests-integration/tests/mpsc_test/callbacks.rs +++ b/tests-integration/tests/mpsc_test/callbacks.rs @@ -32,11 +32,10 @@ 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_clone; use tests_integration::waker_on_drop; use tests_integration::waker_on_wake; -use super::support::waker_on_clone; - struct HoldSender { _sender: S, } 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/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/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"], From b48cad4bbef6ad8987505f4a73436d760f61d3cc Mon Sep 17 00:00:00 2001 From: tison Date: Sun, 20 Sep 2026 10:07:21 +0800 Subject: [PATCH 03/11] refactor(mpmc): simplify waker registration and batch ownership --- asyncband/src/mpmc/bounded.rs | 4 +- asyncband/src/mpmc/queue.rs | 172 +++++++----------- asyncband/src/mpmc/unbounded.rs | 2 +- .../tests/mpmc_test/callbacks.rs | 28 +-- 4 files changed, 73 insertions(+), 133 deletions(-) diff --git a/asyncband/src/mpmc/bounded.rs b/asyncband/src/mpmc/bounded.rs index b42f0b1c..6700d6f3 100644 --- a/asyncband/src/mpmc/bounded.rs +++ b/asyncband/src/mpmc/bounded.rs @@ -30,8 +30,8 @@ use super::queue::Shared; /// the queue is full. /// /// 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. +/// waking tasks, dropping wakers, or dropping messages. The `try_*` methods do not wait for +/// capacity or messages, but may wait to acquire this mutex. /// /// # Panics /// diff --git a/asyncband/src/mpmc/queue.rs b/asyncband/src/mpmc/queue.rs index 2f0b700b..2012bc34 100644 --- a/asyncband/src/mpmc/queue.rs +++ b/asyncband/src/mpmc/queue.rs @@ -37,8 +37,8 @@ pub struct Shared { } /// Values, endpoint counts, and both waiter queues share one lock, so each transition and the -/// waiter it selects are decided together. Waker callbacks and value destructors run outside the -/// lock, because they may reenter the queue. +/// 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, @@ -90,12 +90,10 @@ fn notify_one(waiters: &mut WaitList) -> Option { Some(waker) } -fn notify_all(waiters: &mut WaitList) -> WakerBatch { - let mut wakers = WakerBatch::new(); +fn notify_all(waiters: &mut WaitList, wakers: &mut WakerBatch) { while let Some(waker) = notify_one(waiters) { wakers.push(waker); } - wakers } fn remove_waiter(waiters: &mut WaitList, id: WaiterId) -> Waiter { @@ -110,42 +108,30 @@ fn wake(waker: Option) { } } -enum Registration { - /// The operation is queued; drop the replaced waker after releasing the lock. - Registered(Option), - /// Clone the current waker outside the lock, then retry. - NeedsWaker, -} - /// 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( waiters: &mut WaitList, id: &mut Option, current: &Waker, - cloned: &mut Option, -) -> Registration { +) -> Option { if let Some(queued) = *id { if let Waiter::Waiting(waker) = waiters.waiter_mut(queued) { if waker.will_wake(current) { - return Registration::Registered(None); + return None; } - return match cloned.take() { - Some(new) => Registration::Registered(Some(mem::replace(waker, new))), - None => Registration::NeedsWaker, - }; + return Some(mem::replace(waker, current.clone())); } } - let Some(waker) = cloned.take() else { - return Registration::NeedsWaker; - }; + let waker = current.clone(); if let Some(notified) = id.take() { // The notification already took this node's waker, so nothing is retired. remove_waiter(waiters, notified); } *id = Some(waiters.push_back(Waiter::Waiting(waker))); - Registration::Registered(None) + None } impl Shared { @@ -175,16 +161,17 @@ impl Shared { } pub fn drop_sender(&self) { - let wakers = { + let mut wakers = WakerBatch::new(); + { let mut state = self.state.lock(); state.senders -= 1; if state.senders != 0 { return; } // Woken receivers drain buffered values before they observe disconnection. - notify_all(&mut state.recv_waiters) - }; - wake_all(wakers.into_iter()); + notify_all(&mut state.recv_waiters, &mut wakers); + } + wake_all(&mut wakers); } pub fn clone_receiver(&self) { @@ -192,20 +179,19 @@ impl Shared { } pub fn drop_receiver(&self) { - let (discarded, wakers) = { + let mut wakers = WakerBatch::new(); + let discarded = { let mut state = self.state.lock(); state.receivers -= 1; if state.receivers != 0 { return; } - ( - mem::take(&mut state.values), - notify_all(&mut state.send_waiters), - ) + notify_all(&mut state.send_waiters, &mut wakers); + mem::take(&mut state.values) }; // Release blocked senders before destroying buffered values. Local ownership still drops // the values if a wake callback unwinds. - wake_all(wakers.into_iter()); + wake_all(&mut wakers); drop(discarded); } @@ -272,45 +258,26 @@ impl Send<'_, T> { } fn poll(&mut self, cx: &mut Context<'_>) -> Poll>> { - // The first poll follows a full `try_send` and most likely registers, so it clones before - // locking. A later poll clones only when it must store a new waker: its task changed, or a - // notification took the stored one. - let mut cloned = self.waiter.is_none().then(|| cx.waker().clone()); - loop { - let mut state = self.shared.state.lock(); - let outcome = if state.receivers == 0 { - Err(self.take_value()) - } else if state.has_capacity() { - Ok(state.push(self.take_value())) - } else { - match register( - &mut state.send_waiters, - &mut self.waiter, - cx.waker(), - &mut cloned, - ) { - Registration::Registered(replaced) => { - drop(state); - drop((replaced, cloned)); - return Poll::Pending; - } - Registration::NeedsWaker => { - drop(state); - cloned = Some(cx.waker().clone()); - continue; - } - } - }; - let retired = self - .waiter - .take() - .map(|id| remove_waiter(&mut state.send_waiters, id)); + let mut state = self.shared.state.lock(); + let outcome = if state.receivers == 0 { + Err(self.take_value()) + } else if state.has_capacity() { + Ok(state.push(self.take_value())) + } else { + let retired = register(&mut state.send_waiters, &mut self.waiter, cx.waker()); drop(state); - // Deliver the notification before running waker destructors, which may panic. - let result = outcome.map(wake).map_err(SendError::new); - drop((retired, cloned)); - return Poll::Ready(result); - } + drop(retired); + return Poll::Pending; + }; + let retired = self + .waiter + .take() + .map(|id| remove_waiter(&mut state.send_waiters, id)); + drop(state); + // Deliver the notification before running waker destructors, which may panic. + let result = outcome.map(wake).map_err(SendError::new); + drop(retired); + Poll::Ready(result) } } @@ -342,46 +309,29 @@ struct Recv<'a, T> { impl Recv<'_, T> { fn poll(&mut self, cx: &mut Context<'_>) -> Poll> { - // The first poll follows an empty `try_recv` and most likely registers, so it clones before - // locking. A later poll clones only when it must store a new waker: its task changed, or a - // notification took the stored one. - let mut cloned = self.waiter.is_none().then(|| cx.waker().clone()); - loop { - let mut state = self.shared.state.lock(); - let outcome = match state.pop() { - Ok(popped) => Ok(popped), - Err(TryRecvError::Disconnected) => Err(RecvError::Disconnected), - Err(TryRecvError::Empty) => match register( - &mut state.recv_waiters, - &mut self.waiter, - cx.waker(), - &mut cloned, - ) { - Registration::Registered(replaced) => { - drop(state); - drop((replaced, cloned)); - return Poll::Pending; - } - Registration::NeedsWaker => { - drop(state); - cloned = Some(cx.waker().clone()); - continue; - } - }, - }; - let retired = self - .waiter - .take() - .map(|id| remove_waiter(&mut state.recv_waiters, id)); - drop(state); - // Deliver the notification before running waker destructors, which may panic. - let result = outcome.map(|(value, waker)| { - wake(waker); - value - }); - drop((retired, cloned)); - return Poll::Ready(result); - } + let mut state = self.shared.state.lock(); + let outcome = match state.pop() { + Ok(popped) => Ok(popped), + Err(TryRecvError::Disconnected) => Err(RecvError::Disconnected), + Err(TryRecvError::Empty) => { + let retired = register(&mut state.recv_waiters, &mut self.waiter, cx.waker()); + drop(state); + drop(retired); + return Poll::Pending; + } + }; + let retired = self + .waiter + .take() + .map(|id| remove_waiter(&mut state.recv_waiters, id)); + drop(state); + // Deliver the notification before running waker destructors, which may panic. + let result = outcome.map(|(value, waker)| { + wake(waker); + value + }); + drop(retired); + Poll::Ready(result) } } diff --git a/asyncband/src/mpmc/unbounded.rs b/asyncband/src/mpmc/unbounded.rs index 048121c1..850c0b4e 100644 --- a/asyncband/src/mpmc/unbounded.rs +++ b/asyncband/src/mpmc/unbounded.rs @@ -29,7 +29,7 @@ use super::queue::Shared; /// Sends are synchronous and values may be buffered until available memory is exhausted. /// /// 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 +/// waking tasks, dropping wakers, or dropping messages. Sending and trying to receive may wait to /// acquire this mutex, but never wait for capacity or new messages. pub fn unbounded() -> (UnboundedSender, UnboundedReceiver) { let shared = Arc::new(Shared::unbounded()); diff --git a/tests-integration/tests/mpmc_test/callbacks.rs b/tests-integration/tests/mpmc_test/callbacks.rs index 1ddc4b3f..cf851391 100644 --- a/tests-integration/tests/mpmc_test/callbacks.rs +++ b/tests-integration/tests/mpmc_test/callbacks.rs @@ -28,34 +28,24 @@ 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_clone; use tests_integration::waker_on_drop; -// A registering poll clones its waker before taking the queue lock, so a value that the clone -// callback sends is already queued when the same poll looks for one. #[test] -fn receive_registration_allows_a_waker_clone_to_send() { - assert_completes_without_deadlock(|| { - let (sender, receiver) = mpmc::unbounded(); - let reentrant = sender.clone(); - let waker = waker_on_clone(move || reentrant.send(1).unwrap()); - let mut recv = Box::pin(receiver.recv()); - - assert_eq!(expect_ready(poll_with(recv.as_mut(), &waker)), Ok(1)); - drop(sender); - }); -} - -#[test] -fn send_registration_allows_a_waker_clone_to_receive() { +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 waker = waker_on_clone(move || assert_eq!(reentrant.try_recv(), Ok(0))); + 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)); - expect_ready(poll_with(send.as_mut(), &waker)).unwrap(); + 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)); }); } From b1d5db502184d3cf02aecade5598ac0aae818b04 Mon Sep 17 00:00:00 2001 From: tison Date: Sun, 20 Sep 2026 10:07:26 +0800 Subject: [PATCH 04/11] test(mpmc): cover notification handoff after resource contention --- .../tests/mpmc_test/notification.rs | 102 ++++++++++++++++++ 1 file changed, 102 insertions(+) diff --git a/tests-integration/tests/mpmc_test/notification.rs b/tests-integration/tests/mpmc_test/notification.rs index fd8db23a..739d51ca 100644 --- a/tests-integration/tests/mpmc_test/notification.rs +++ b/tests-integration/tests/mpmc_test/notification.rs @@ -94,6 +94,55 @@ fn unbounded_cancelled_notified_receiver_wakes_next_receiver() { cancelled_notified_receiver_wakes_next_receiver(receiver, || sender.send(1).unwrap()); } +fn cancelling_a_receiver_after_its_value_is_taken_does_not_wake_next( + receiver: impl Receiver, + send: impl Fn(usize), + take_value: impl FnOnce() -> Result, +) { + let competing = receiver.clone(); + let mut cancelled = Box::pin(receiver.recv()); + let mut waiting = Box::pin(competing.recv()); + 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()); + send(1); + assert_eq!(cancelled_wakes.count(), 1); + assert_eq!(waiting_wakes.count(), 0); + assert_eq!(take_value(), Ok(1)); + drop(cancelled); + assert_eq!(waiting_wakes.count(), 0); + + // The next value must still notify the waiter whose predecessor was cancelled. + send(2); + assert_eq!(waiting_wakes.count(), 1); + assert_eq!( + expect_ready(poll_with(waiting.as_mut(), &waiting_waker)), + Ok(2) + ); +} + +#[test] +fn bounded_cancelling_a_receiver_after_its_value_is_taken_does_not_wake_next() { + let (sender, receiver) = mpmc::bounded(1); + cancelling_a_receiver_after_its_value_is_taken_does_not_wake_next( + receiver.clone(), + |value| sender.try_send(value).unwrap(), + || receiver.try_recv(), + ); +} + +#[test] +fn unbounded_cancelling_a_receiver_after_its_value_is_taken_does_not_wake_next() { + let (sender, receiver) = mpmc::unbounded(); + cancelling_a_receiver_after_its_value_is_taken_does_not_wake_next( + receiver.clone(), + |value| sender.send(value).unwrap(), + || receiver.try_recv(), + ); +} + #[test] fn notified_receiver_that_loses_the_value_queues_behind_waiting_receivers() { let (sender, receiver) = mpmc::unbounded(); @@ -120,6 +169,59 @@ fn notified_receiver_that_loses_the_value_queues_behind_waiting_receivers() { ); } +#[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. From e0f8f6126a864e3cddff2e93909d9d97e102dde9 Mon Sep 17 00:00:00 2001 From: tison Date: Sun, 20 Sep 2026 10:39:30 +0800 Subject: [PATCH 05/11] perf(mpmc): detach waiter storage on disconnection --- asyncband/src/internal/mod.rs | 1 - asyncband/src/mpmc/queue.rs | 41 +++++++++++-------- .../tests/mpmc_test/notification.rs | 16 +++++++- 3 files changed, 38 insertions(+), 20 deletions(-) diff --git a/asyncband/src/internal/mod.rs b/asyncband/src/internal/mod.rs index 1ff6ab01..c736edf9 100644 --- a/asyncband/src/internal/mod.rs +++ b/asyncband/src/internal/mod.rs @@ -129,7 +129,6 @@ pub(crate) mod waitlist; feature = "event", feature = "completion", feature = "latch", - feature = "mpmc", feature = "mpsc", feature = "mutex", feature = "once", diff --git a/asyncband/src/mpmc/queue.rs b/asyncband/src/mpmc/queue.rs index 2012bc34..11e87670 100644 --- a/asyncband/src/mpmc/queue.rs +++ b/asyncband/src/mpmc/queue.rs @@ -30,7 +30,6 @@ use crate::internal::mutex::Mutex; use crate::internal::waitlist::WaitList; use crate::internal::waitlist::WaiterId; use crate::internal::wake_all; -use crate::internal::waker_batch::WakerBatch; pub struct Shared { state: Mutex>, @@ -90,12 +89,6 @@ fn notify_one(waiters: &mut WaitList) -> Option { Some(waker) } -fn notify_all(waiters: &mut WaitList, wakers: &mut WakerBatch) { - while let Some(waker) = notify_one(waiters) { - wakers.push(waker); - } -} - fn remove_waiter(waiters: &mut WaitList, id: WaiterId) -> Waiter { // Unlinking is idempotent, so notified waiters are removed the same way as linked ones. waiters.unlink_waiter(id, |_| true); @@ -161,17 +154,17 @@ impl Shared { } pub fn drop_sender(&self) { - let mut wakers = WakerBatch::new(); - { + let mut waiters = { let mut state = self.state.lock(); state.senders -= 1; if state.senders != 0 { return; } - // Woken receivers drain buffered values before they observe disconnection. - notify_all(&mut state.recv_waiters, &mut wakers); - } - wake_all(&mut wakers); + // 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()) + }; + wake_all(std::iter::from_fn(|| notify_one(&mut waiters))); } pub fn clone_receiver(&self) { @@ -179,19 +172,20 @@ impl Shared { } pub fn drop_receiver(&self) { - let mut wakers = WakerBatch::new(); - let discarded = { + let (discarded, mut waiters) = { let mut state = self.state.lock(); state.receivers -= 1; if state.receivers != 0 { return; } - notify_all(&mut state.send_waiters, &mut wakers); - mem::take(&mut state.values) + ( + mem::take(&mut state.values), + mem::replace(&mut state.send_waiters, WaitList::new()), + ) }; // Release blocked senders before destroying buffered values. Local ownership still drops // the values if a wake callback unwinds. - wake_all(&mut wakers); + wake_all(std::iter::from_fn(|| notify_one(&mut waiters))); drop(discarded); } @@ -260,6 +254,7 @@ impl Send<'_, T> { fn poll(&mut self, cx: &mut Context<'_>) -> Poll>> { 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())) @@ -288,6 +283,9 @@ impl Drop for Send<'_, T> { }; let (retired, waker) = { let mut state = self.shared.state.lock(); + if state.receivers == 0 { + return; + } let retired = remove_waiter(&mut state.send_waiters, 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() { @@ -310,6 +308,10 @@ struct Recv<'a, T> { impl Recv<'_, T> { fn poll(&mut self, cx: &mut Context<'_>) -> Poll> { 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), @@ -342,6 +344,9 @@ impl Drop for Recv<'_, T> { }; let (retired, waker) = { let mut state = self.shared.state.lock(); + if state.senders == 0 { + return; + } let retired = remove_waiter(&mut state.recv_waiters, 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() { diff --git a/tests-integration/tests/mpmc_test/notification.rs b/tests-integration/tests/mpmc_test/notification.rs index 739d51ca..d79b8bea 100644 --- a/tests-integration/tests/mpmc_test/notification.rs +++ b/tests-integration/tests/mpmc_test/notification.rs @@ -270,17 +270,25 @@ fn last_sender_wakes_every_pending_receiver() { 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)), - Err(RecvError::Disconnected) + Ok(1) ); assert_eq!( expect_ready(poll_with(second.as_mut(), &second_waker)), @@ -295,14 +303,20 @@ fn last_receiver_wakes_every_pending_sender() { 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() From 6656599f62eaa4f1e999b942ff6ff534359a98d3 Mon Sep 17 00:00:00 2001 From: tison Date: Sun, 20 Sep 2026 10:42:01 +0800 Subject: [PATCH 06/11] refactor(mpsc): register task wakers only when needed --- CHANGELOG.md | 1 - asyncband/src/mpmc/bounded.rs | 15 +-- asyncband/src/mpmc/unbounded.rs | 8 +- asyncband/src/mpsc/bounded/mod.rs | 7 +- asyncband/src/mpsc/bounded/receiver.rs | 4 +- asyncband/src/mpsc/bounded/sender.rs | 8 +- asyncband/src/mpsc/mod.rs | 13 +++ asyncband/src/mpsc/unbounded/mod.rs | 5 +- asyncband/src/mpsc/unbounded/receiver.rs | 5 +- tests-integration/src/lib.rs | 30 ------ .../tests/mpsc_test/callbacks.rs | 102 ------------------ 11 files changed, 33 insertions(+), 165 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index e1bdd901..a2c8751d 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 diff --git a/asyncband/src/mpmc/bounded.rs b/asyncband/src/mpmc/bounded.rs index 6700d6f3..c103fbb0 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 an internal mutex; no lock is held across an await point or while -/// waking tasks, dropping wakers, or dropping messages. 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 /// @@ -83,11 +82,8 @@ 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. If this call was woken for capacity that is still - /// free, cancelling it wakes 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 } @@ -138,8 +134,7 @@ impl BoundedReceiver { /// /// # Cancel safety /// - /// Dropping a pending `recv` does not consume a value. If this call was woken for a value that - /// is still queued, cancelling it wakes the next waiting receiver instead. + /// 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/unbounded.rs b/asyncband/src/mpmc/unbounded.rs index 850c0b4e..c3927e30 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 an internal mutex; no lock is held across an await point or while -/// waking tasks, dropping wakers, or dropping messages. 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(Shared::unbounded()); ( @@ -119,8 +118,7 @@ impl UnboundedReceiver { /// /// # Cancel safety /// - /// Dropping a pending `recv` does not consume a value. If this call was woken for a value that - /// is still queued, cancelling it wakes the next waiting receiver instead. + /// 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..20c2fa66 100644 --- a/asyncband/src/mpsc/bounded/receiver.rs +++ b/asyncband/src/mpsc/bounded/receiver.rs @@ -28,6 +28,7 @@ use crate::internal::wake_all; use crate::internal::waker_batch::WakerBatch; use crate::mpsc::RecvError; use crate::mpsc::TryRecvError; +use crate::mpsc::register_waker; /// The receiving endpoint of a bounded mpsc channel. /// @@ -134,7 +135,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 +151,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..a4c5f921 100644 --- a/asyncband/src/mpsc/bounded/sender.rs +++ b/asyncband/src/mpsc/bounded/sender.rs @@ -28,6 +28,7 @@ use crate::internal::mutex::Mutex; use crate::internal::waitlist::WaiterId; use crate::mpsc::SendError; use crate::mpsc::TrySendError; +use crate::mpsc::register_waker; /// The sending endpoint of a bounded mpsc channel. /// @@ -235,7 +236,6 @@ 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 { return Poll::Ready(Err(SendError::new(()))); @@ -250,10 +250,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 +263,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 } diff --git a/asyncband/src/mpsc/mod.rs b/asyncband/src/mpsc/mod.rs index bad23f06..cd3268f0 100644 --- a/asyncband/src/mpsc/mod.rs +++ b/asyncband/src/mpsc/mod.rs @@ -21,6 +21,8 @@ //! to send while the receiver is alive, but a slow receiver can cause memory use to grow without a //! configured limit. +use std::task::Waker; + mod bounded; mod error; mod unbounded; @@ -36,3 +38,14 @@ pub use self::error::TrySendError; pub use self::unbounded::UnboundedReceiver; pub use self::unbounded::UnboundedSender; pub use self::unbounded::unbounded; + +/// Retains the current task waker, returning any replaced registration for unlocked destruction. +#[inline] +#[must_use = "drop the replaced waker after releasing the channel lock"] +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()) + } +} 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..381adfef 100644 --- a/asyncband/src/mpsc/unbounded/receiver.rs +++ b/asyncband/src/mpsc/unbounded/receiver.rs @@ -29,6 +29,7 @@ use super::buffer::pop_batch; use crate::internal::mutex::Mutex; use crate::mpsc::RecvError; use crate::mpsc::TryRecvError; +use crate::mpsc::register_waker; /// The receiving endpoint of an unbounded mpsc channel. /// @@ -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/tests-integration/src/lib.rs b/tests-integration/src/lib.rs index 4862e5f6..00f2045e 100644 --- a/tests-integration/src/lib.rs +++ b/tests-integration/src/lib.rs @@ -24,8 +24,6 @@ use std::sync::atomic::AtomicUsize; use std::sync::atomic::Ordering; use std::task::Context; use std::task::Poll; -use std::task::RawWaker; -use std::task::RawWakerVTable; use std::task::Wake; use std::task::Waker; @@ -124,34 +122,6 @@ pub fn waker_on_drop(callback: impl FnOnce() + Send + 'static) -> Waker { Waker::from(Arc::new(OnDrop(Mutex::new(Some(Box::new(callback)))))) } -// 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: `Waker::clone` borrows the input waker, whose reference keeps the Arc alive while - // the callback runs. Taking the new reference last leaves the count unchanged if the - // callback panics; the returned waker owns that 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)) } -} - pub fn assert_completes_without_deadlock(test: impl FnOnce() + Send + 'static) { let (finished_tx, finished_rx) = std::sync::mpsc::channel(); let worker = std::thread::spawn(move || { diff --git a/tests-integration/tests/mpsc_test/callbacks.rs b/tests-integration/tests/mpsc_test/callbacks.rs index e5452b2a..b2aae937 100644 --- a/tests-integration/tests/mpsc_test/callbacks.rs +++ b/tests-integration/tests/mpsc_test/callbacks.rs @@ -32,7 +32,6 @@ 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_clone; use tests_integration::waker_on_drop; use tests_integration::waker_on_wake; @@ -80,107 +79,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(|| { From 7e8b3543c86f1e19065f5035cc787d03c6f19bd9 Mon Sep 17 00:00:00 2001 From: tison Date: Sun, 20 Sep 2026 10:44:09 +0800 Subject: [PATCH 07/11] perf(mpsc): release disconnected waiter storage outside the lock --- CHANGELOG.md | 2 +- asyncband/src/internal/mod.rs | 1 - asyncband/src/mpsc/bounded/receiver.rs | 21 +++++++++++-------- asyncband/src/mpsc/bounded/sender.rs | 4 ++++ .../tests/mpsc_test/reservation.rs | 3 +++ 5 files changed, 20 insertions(+), 11 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index a2c8751d..f5e01ff8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -44,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/internal/mod.rs b/asyncband/src/internal/mod.rs index c736edf9..82d9f78b 100644 --- a/asyncband/src/internal/mod.rs +++ b/asyncband/src/internal/mod.rs @@ -129,7 +129,6 @@ pub(crate) mod waitlist; feature = "event", feature = "completion", feature = "latch", - feature = "mpsc", feature = "mutex", feature = "once", feature = "phaser", diff --git a/asyncband/src/mpsc/bounded/receiver.rs b/asyncband/src/mpsc/bounded/receiver.rs index 20c2fa66..2aaffdd4 100644 --- a/asyncband/src/mpsc/bounded/receiver.rs +++ b/asyncband/src/mpsc/bounded/receiver.rs @@ -24,8 +24,8 @@ use std::task::Poll; use super::State; use crate::internal::mutex::Mutex; +use crate::internal::waitlist::WaitList; use crate::internal::wake_all; -use crate::internal::waker_batch::WakerBatch; use crate::mpsc::RecvError; use crate::mpsc::TryRecvError; use crate::mpsc::register_waker; @@ -46,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); } diff --git a/asyncband/src/mpsc/bounded/sender.rs b/asyncband/src/mpsc/bounded/sender.rs index a4c5f921..7f548456 100644 --- a/asyncband/src/mpsc/bounded/sender.rs +++ b/asyncband/src/mpsc/bounded/sender.rs @@ -238,6 +238,7 @@ impl<'a, T> Reserve<'a, T> { fn poll(&mut self, cx: &mut Context<'_>) -> Poll, SendError<()>>> { let mut state = self.shared.lock(); if !state.receiver { + self.waiter = None; return Poll::Ready(Err(SendError::new(()))); } if let Some(index) = self.waiter { @@ -278,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/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()); From 303e8ed342067ff69aefa4a69d9d4cfbc29a9f14 Mon Sep 17 00:00:00 2001 From: tison Date: Sun, 20 Sep 2026 10:51:35 +0800 Subject: [PATCH 08/11] test: consolidate channel and phaser waker coverage --- asyncband/src/mpmc/bounded.rs | 6 +- asyncband/src/mpmc/unbounded.rs | 3 +- asyncband/src/mpsc/mod.rs | 5 +- tests-integration/tests/mpmc_test/main.rs | 9 + .../tests/mpmc_test/notification.rs | 166 ++++++--------- .../tests/mpsc_test/callbacks.rs | 1 - tests-integration/tests/phaser_test.rs | 196 +++++------------- 7 files changed, 133 insertions(+), 253 deletions(-) diff --git a/asyncband/src/mpmc/bounded.rs b/asyncband/src/mpmc/bounded.rs index c103fbb0..3122f839 100644 --- a/asyncband/src/mpmc/bounded.rs +++ b/asyncband/src/mpmc/bounded.rs @@ -83,7 +83,8 @@ impl BoundedSender { /// # Cancel safety /// /// 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. + /// [`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 } @@ -134,7 +135,8 @@ impl BoundedReceiver { /// /// # Cancel safety /// - /// Dropping a pending `recv` does not consume a value or prevent other receivers from receiving it. + /// 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/unbounded.rs b/asyncband/src/mpmc/unbounded.rs index c3927e30..7fc71ac6 100644 --- a/asyncband/src/mpmc/unbounded.rs +++ b/asyncband/src/mpmc/unbounded.rs @@ -118,7 +118,8 @@ impl UnboundedReceiver { /// /// # Cancel safety /// - /// Dropping a pending `recv` does not consume a value or prevent other receivers from receiving it. + /// 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/mod.rs b/asyncband/src/mpsc/mod.rs index cd3268f0..c86c6009 100644 --- a/asyncband/src/mpsc/mod.rs +++ b/asyncband/src/mpsc/mod.rs @@ -43,7 +43,10 @@ pub use self::unbounded::unbounded; #[inline] #[must_use = "drop the replaced waker after releasing the channel lock"] fn register_waker(slot: &mut Option, waker: &Waker) -> Option { - if slot.as_ref().is_some_and(|current| current.will_wake(waker)) { + if slot + .as_ref() + .is_some_and(|current| current.will_wake(waker)) + { None } else { slot.replace(waker.clone()) diff --git a/tests-integration/tests/mpmc_test/main.rs b/tests-integration/tests/mpmc_test/main.rs index 5d008504..8795d7ad 100644 --- a/tests-integration/tests/mpmc_test/main.rs +++ b/tests-integration/tests/mpmc_test/main.rs @@ -33,18 +33,27 @@ 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] diff --git a/tests-integration/tests/mpmc_test/notification.rs b/tests-integration/tests/mpmc_test/notification.rs index d79b8bea..1b2a944d 100644 --- a/tests-integration/tests/mpmc_test/notification.rs +++ b/tests-integration/tests/mpmc_test/notification.rs @@ -28,119 +28,74 @@ use tests_integration::poll_with; use super::Receiver; -fn send_wakes_only_the_first_receiver(receiver: impl Receiver, send: impl FnOnce()) { - 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(); - - assert_eq!(first_wakes.count(), 1); - assert_eq!(second_wakes.count(), 0); - 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); -} - -#[test] -fn bounded_send_wakes_only_the_first_receiver() { - let (sender, receiver) = mpmc::bounded(2); - send_wakes_only_the_first_receiver(receiver, || sender.try_send(1).unwrap()); -} - -#[test] -fn unbounded_send_wakes_only_the_first_receiver() { - let (sender, receiver) = mpmc::unbounded(); - send_wakes_only_the_first_receiver(receiver, || sender.send(1).unwrap()); +#[derive(Debug)] +enum NotifiedReceiver { + Receive, + Cancel, + TakeValueThenCancel, } -fn cancelled_notified_receiver_wakes_next_receiver( - receiver: impl Receiver, - send: impl FnOnce(), +fn receiver_notifications>( + channel: impl Fn() -> (S, R), + send: impl Fn(&S, usize), ) { - let competing = receiver.clone(); - let mut cancelled = Box::pin(receiver.recv()); - let mut waiting = Box::pin(competing.recv()); - 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()); - send(); - 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(waiting.as_mut(), &waiting_waker)), - Ok(1) - ); -} - -#[test] -fn bounded_cancelled_notified_receiver_wakes_next_receiver() { - let (sender, receiver) = mpmc::bounded(1); - cancelled_notified_receiver_wakes_next_receiver(receiver, || sender.try_send(1).unwrap()); -} - -#[test] -fn unbounded_cancelled_notified_receiver_wakes_next_receiver() { - let (sender, receiver) = mpmc::unbounded(); - cancelled_notified_receiver_wakes_next_receiver(receiver, || sender.send(1).unwrap()); -} - -fn cancelling_a_receiver_after_its_value_is_taken_does_not_wake_next( - receiver: impl Receiver, - send: impl Fn(usize), - take_value: impl FnOnce() -> Result, -) { - let competing = receiver.clone(); - let mut cancelled = Box::pin(receiver.recv()); - let mut waiting = Box::pin(competing.recv()); - 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()); - send(1); - assert_eq!(cancelled_wakes.count(), 1); - assert_eq!(waiting_wakes.count(), 0); - assert_eq!(take_value(), Ok(1)); - drop(cancelled); - assert_eq!(waiting_wakes.count(), 0); - - // The next value must still notify the waiter whose predecessor was cancelled. - send(2); - assert_eq!(waiting_wakes.count(), 1); - assert_eq!( - expect_ready(poll_with(waiting.as_mut(), &waiting_waker)), - Ok(2) - ); + 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_cancelling_a_receiver_after_its_value_is_taken_does_not_wake_next() { - let (sender, receiver) = mpmc::bounded(1); - cancelling_a_receiver_after_its_value_is_taken_does_not_wake_next( - receiver.clone(), - |value| sender.try_send(value).unwrap(), - || receiver.try_recv(), +fn bounded_receiver_notification_is_consumed_or_handed_off() { + receiver_notifications( + || mpmc::bounded(2), + |sender, value| sender.try_send(value).unwrap(), ); } #[test] -fn unbounded_cancelling_a_receiver_after_its_value_is_taken_does_not_wake_next() { - let (sender, receiver) = mpmc::unbounded(); - cancelling_a_receiver_after_its_value_is_taken_does_not_wake_next( - receiver.clone(), - |value| sender.send(value).unwrap(), - || receiver.try_recv(), - ); +fn unbounded_receiver_notification_is_consumed_or_handed_off() { + receiver_notifications(mpmc::unbounded, |sender, value| sender.send(value).unwrap()); } #[test] @@ -286,10 +241,7 @@ fn last_sender_wakes_every_pending_receiver() { 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(first.as_mut(), &first_waker)), Ok(1)); assert_eq!( expect_ready(poll_with(second.as_mut(), &second_waker)), Err(RecvError::Disconnected) diff --git a/tests-integration/tests/mpsc_test/callbacks.rs b/tests-integration/tests/mpsc_test/callbacks.rs index b2aae937..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; 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()); From 5fe3d07a51f7a35a843203723cf669a31140be8b Mon Sep 17 00:00:00 2001 From: tison Date: Sun, 20 Sep 2026 10:55:08 +0800 Subject: [PATCH 09/11] refactor(mpmc): group waiter operations into methods --- asyncband/src/mpmc/queue.rs | 120 +++++++++++++++++++----------------- 1 file changed, 64 insertions(+), 56 deletions(-) diff --git a/asyncband/src/mpmc/queue.rs b/asyncband/src/mpmc/queue.rs index 11e87670..a279470f 100644 --- a/asyncband/src/mpmc/queue.rs +++ b/asyncband/src/mpmc/queue.rs @@ -56,14 +56,14 @@ impl State { /// Queues a value and selects the receiver to wake. fn push(&mut self, value: T) -> Option { self.values.push_back(value); - notify_one(&mut self.recv_waiters) + 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, notify_one(&mut self.send_waiters))) + Ok((value, self.send_waiters.notify_one())) } else if self.senders == 0 { Err(TryRecvError::Disconnected) } else { @@ -81,50 +81,42 @@ enum Waiter { Notified, } -fn notify_one(waiters: &mut WaitList) -> Option { - let (_, waiter) = waiters.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(waiters: &mut WaitList, id: WaiterId) -> Waiter { - // Unlinking is idempotent, so notified waiters are removed the same way as linked ones. - waiters.unlink_waiter(id, |_| true); - waiters.remove_unlinked_waiter(id) -} +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 wake(waker: Option) { - if let Some(waker) = waker { - waker.wake(); + 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( - waiters: &mut WaitList, - id: &mut Option, - current: &Waker, -) -> Option { - if let Some(queued) = *id { - if let Waiter::Waiting(waker) = waiters.waiter_mut(queued) { - if waker.will_wake(current) { - return None; + /// 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())); } - 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 } - let waker = current.clone(); - if let Some(notified) = id.take() { - // The notification already took this node's waker, so nothing is retired. - remove_waiter(waiters, notified); - } - *id = Some(waiters.push_back(Waiter::Waiting(waker))); - None } impl Shared { @@ -164,7 +156,7 @@ impl Shared { // notification and reclamation happen without holding the queue lock. mem::replace(&mut state.recv_waiters, WaitList::new()) }; - wake_all(std::iter::from_fn(|| notify_one(&mut waiters))); + wake_all(std::iter::from_fn(|| waiters.notify_one())); } pub fn clone_receiver(&self) { @@ -185,7 +177,7 @@ impl Shared { }; // Release blocked senders before destroying buffered values. Local ownership still drops // the values if a wake callback unwinds. - wake_all(std::iter::from_fn(|| notify_one(&mut waiters))); + wake_all(std::iter::from_fn(|| waiters.notify_one())); drop(discarded); } @@ -200,7 +192,9 @@ impl Shared { } state.push(value) }; - wake(waker); + if let Some(waker) = waker { + waker.wake(); + } Ok(()) } @@ -220,7 +214,9 @@ impl Shared { pub fn try_recv(&self) -> Result { let (value, waker) = self.state.lock().pop()?; - wake(waker); + if let Some(waker) = waker { + waker.wake(); + } Ok(value) } @@ -259,7 +255,7 @@ impl Send<'_, T> { } else if state.has_capacity() { Ok(state.push(self.take_value())) } else { - let retired = register(&mut state.send_waiters, &mut self.waiter, cx.waker()); + let retired = state.send_waiters.register(&mut self.waiter, cx.waker()); drop(state); drop(retired); return Poll::Pending; @@ -267,10 +263,16 @@ impl Send<'_, T> { let retired = self .waiter .take() - .map(|id| remove_waiter(&mut state.send_waiters, id)); + .map(|id| state.send_waiters.remove_waiter(id)); drop(state); // Deliver the notification before running waker destructors, which may panic. - let result = outcome.map(wake).map_err(SendError::new); + let result = outcome + .map(|waker| { + if let Some(waker) = waker { + waker.wake(); + } + }) + .map_err(SendError::new); drop(retired); Poll::Ready(result) } @@ -286,16 +288,18 @@ impl Drop for Send<'_, T> { if state.receivers == 0 { return; } - let retired = remove_waiter(&mut state.send_waiters, id); + 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() { - notify_one(&mut state.send_waiters) + state.send_waiters.notify_one() } else { None }; (retired, waker) }; - wake(waker); + if let Some(waker) = waker { + waker.wake(); + } drop(retired); } } @@ -316,7 +320,7 @@ impl Recv<'_, T> { Ok(popped) => Ok(popped), Err(TryRecvError::Disconnected) => Err(RecvError::Disconnected), Err(TryRecvError::Empty) => { - let retired = register(&mut state.recv_waiters, &mut self.waiter, cx.waker()); + let retired = state.recv_waiters.register(&mut self.waiter, cx.waker()); drop(state); drop(retired); return Poll::Pending; @@ -325,11 +329,13 @@ impl Recv<'_, T> { let retired = self .waiter .take() - .map(|id| remove_waiter(&mut state.recv_waiters, id)); + .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)| { - wake(waker); + if let Some(waker) = waker { + waker.wake(); + } value }); drop(retired); @@ -347,16 +353,18 @@ impl Drop for Recv<'_, T> { if state.senders == 0 { return; } - let retired = remove_waiter(&mut state.recv_waiters, id); + 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() { - notify_one(&mut state.recv_waiters) + state.recv_waiters.notify_one() } else { None }; (retired, waker) }; - wake(waker); + if let Some(waker) = waker { + waker.wake(); + } drop(retired); } } From 35cca7f36d87728b500d2ab838a5fcfe8de5cec5 Mon Sep 17 00:00:00 2001 From: tison Date: Sun, 20 Sep 2026 11:11:21 +0800 Subject: [PATCH 10/11] docs: align the crate API map table --- asyncband/src/lib.rs | 52 ++++++++++++++++++++++---------------------- 1 file changed, 26 insertions(+), 26 deletions(-) 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 //! From 1a83536f67ab190de94d98b776aca75e876f1478 Mon Sep 17 00:00:00 2001 From: tison Date: Sun, 20 Sep 2026 11:11:22 +0800 Subject: [PATCH 11/11] refactor: share optional waker registration across primitives --- asyncband/src/event/manual_reset.rs | 26 ++++++------------------ asyncband/src/internal/mod.rs | 16 +++++++++++++++ asyncband/src/internal/semaphore.rs | 9 ++------ asyncband/src/mpsc/bounded/receiver.rs | 2 +- asyncband/src/mpsc/bounded/sender.rs | 2 +- asyncband/src/mpsc/mod.rs | 16 --------------- asyncband/src/mpsc/unbounded/receiver.rs | 2 +- 7 files changed, 27 insertions(+), 46 deletions(-) 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 82d9f78b..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 diff --git a/asyncband/src/internal/semaphore.rs b/asyncband/src/internal/semaphore.rs index 9ebf0eb9..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; @@ -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/mpsc/bounded/receiver.rs b/asyncband/src/mpsc/bounded/receiver.rs index 2aaffdd4..11f6f7f3 100644 --- a/asyncband/src/mpsc/bounded/receiver.rs +++ b/asyncband/src/mpsc/bounded/receiver.rs @@ -24,11 +24,11 @@ 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::mpsc::RecvError; use crate::mpsc::TryRecvError; -use crate::mpsc::register_waker; /// The receiving endpoint of a bounded mpsc channel. /// diff --git a/asyncband/src/mpsc/bounded/sender.rs b/asyncband/src/mpsc/bounded/sender.rs index 7f548456..2f286f04 100644 --- a/asyncband/src/mpsc/bounded/sender.rs +++ b/asyncband/src/mpsc/bounded/sender.rs @@ -25,10 +25,10 @@ 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; -use crate::mpsc::register_waker; /// The sending endpoint of a bounded mpsc channel. /// diff --git a/asyncband/src/mpsc/mod.rs b/asyncband/src/mpsc/mod.rs index c86c6009..bad23f06 100644 --- a/asyncband/src/mpsc/mod.rs +++ b/asyncband/src/mpsc/mod.rs @@ -21,8 +21,6 @@ //! to send while the receiver is alive, but a slow receiver can cause memory use to grow without a //! configured limit. -use std::task::Waker; - mod bounded; mod error; mod unbounded; @@ -38,17 +36,3 @@ pub use self::error::TrySendError; pub use self::unbounded::UnboundedReceiver; pub use self::unbounded::UnboundedSender; pub use self::unbounded::unbounded; - -/// Retains the current task waker, returning any replaced registration for unlocked destruction. -#[inline] -#[must_use = "drop the replaced waker after releasing the channel lock"] -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()) - } -} diff --git a/asyncband/src/mpsc/unbounded/receiver.rs b/asyncband/src/mpsc/unbounded/receiver.rs index 381adfef..90617be4 100644 --- a/asyncband/src/mpsc/unbounded/receiver.rs +++ b/asyncband/src/mpsc/unbounded/receiver.rs @@ -27,9 +27,9 @@ 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; -use crate::mpsc::register_waker; /// The receiving endpoint of an unbounded mpsc channel. ///