From cc1adbd168ea8f52473cf617785ea4c835c89ef9 Mon Sep 17 00:00:00 2001 From: Orthur Date: Mon, 14 Sep 2026 05:08:35 -0400 Subject: [PATCH 1/2] 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 2/2] 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"],