diff --git a/asyncband/src/barrier/mod.rs b/asyncband/src/barrier/mod.rs index b3a0637e..e52db7cd 100644 --- a/asyncband/src/barrier/mod.rs +++ b/asyncband/src/barrier/mod.rs @@ -52,10 +52,9 @@ use std::future::Future; use std::pin::Pin; use std::task::Context; use std::task::Poll; +use std::task::Waker; use crate::internal::mutex::Mutex; -use crate::internal::wake_all; -use crate::internal::waker_batch::WakerBatch; use crate::internal::wakerset::WakerSet; use crate::internal::wakerset::WakerToken; @@ -185,14 +184,13 @@ impl Barrier { state.arrived += 1; // The final arrival completes this generation. Advance the generation while holding - // the state lock, then wake the drained followers after releasing it. + // the state lock, then wake the detached followers after releasing it. if state.arrived == self.n { state.arrived = 0; state.generation += 1; - let mut wakers = WakerBatch::new(); - state.waiters.drain_into(&mut wakers); + let mut wakers = state.waiters.take_all(); drop(state); - wake_all(&mut wakers); + wakers.by_ref().for_each(Waker::wake); return BarrierWaitResult(true); } @@ -239,7 +237,7 @@ impl Future for BarrierWait<'_> { let mut state = barrier.state.lock(); if *generation < state.generation { - // Completion advances the generation and drains its old waiters under this same lock, + // Completion advances the generation and detaches its old waiters under this same lock, // so no registration represented by this token remains in the waker set. *token = None; return Poll::Ready(()); diff --git a/asyncband/src/broadcast/mpmc/bounded/mod.rs b/asyncband/src/broadcast/mpmc/bounded/mod.rs index 5cc64edd..48ec26c6 100644 --- a/asyncband/src/broadcast/mpmc/bounded/mod.rs +++ b/asyncband/src/broadcast/mpmc/bounded/mod.rs @@ -108,6 +108,7 @@ use std::sync::atomic::AtomicUsize; use std::sync::atomic::Ordering; use std::task::Context; use std::task::Poll; +use std::task::Waker; use super::common; use super::common::Backlog; @@ -119,8 +120,6 @@ use crate::internal::arena::SlotId; use crate::internal::mutex::Mutex; use crate::internal::semaphore::Acquire; use crate::internal::semaphore::Semaphore; -use crate::internal::wake_all; -use crate::internal::waker_batch::WakerBatch; use crate::internal::wakerset::WakerToken; #[cfg(test)] @@ -181,64 +180,39 @@ struct Shared { /// and parks again if another producer took the space first. The semaphore starts empty and /// only ever grows when a reclaim finds someone waiting, so an idle channel accumulates none. tx_permits: Semaphore, - /// How many producers are somewhere inside the waiting path of [`BoundedSender::send`]. - /// - /// An upper bound on the number of parked producers, and the only thing either release path - /// consults. It answers both questions a reclaim has — whether to wake anyone, and how many - /// permits are worth handing out — without taking the semaphore's lock. Reclaiming is far more - /// frequent than blocking — under fan-out every message is reclaimed, while a channel with - /// headroom never blocks at all — so paying an atomic load there instead of a lock acquisition - /// is what keeps an uncontended receive off the semaphore entirely. + /// Number of producers in the waiting path of [`BoundedSender::send`]. + /// + /// This upper bound lets receives skip the semaphore lock when no producer is waiting. blocked_senders: AtomicUsize, } impl Shared { - /// Hands `freed` released slots back to producers parked in `send`. - /// - /// Capacity is `retained()`, which is `buffer.len()`. The buffer grows only in - /// `Backlog::publish_retained` and shrinks only in `Backlog::reclaim_vacated`, which is - /// reachable from exactly two places: a receive that vacates the last cursor at the backlog - /// head, and removing a subscription. Those are the only callers of this method, so no path - /// can free capacity without waking a producer. Subscribing cannot: a new cursor starts at the - /// tail and never lowers `retained()`. + /// Notifies blocked producers after a receive or subscription removal frees capacity. /// - /// Callers must invoke this with the channel unlocked, and — on the receive path — before - /// touching the payload, since `common::take_msg` runs user code that may panic. + /// Call with the channel unlocked, before payload cloning or destruction can panic. fn release_reclaimed(&self, freed: usize) { - // Release no more permits than there are producers to wake. A permit the semaphore cannot - // hand to a waiter is kept as slack, and the next producer to block has to burn it off one - // futile publish attempt — a channel lock apiece — at a time before it can park. Freeing a - // large prefix at once is not exotic: dropping a lagging subscription reclaims the whole - // backlog, which would otherwise leave nearly `capacity` permits behind. - // - // Capping cannot lose a wake-up, by the same argument that lets this read the count at all: - // a producer this load observes is one the release covers, and one it misses incremented - // after the load, which it does before taking the channel lock to recheck — so its recheck - // runs after the reclaim and finds the capacity itself. + // Cap permits at the number of blocked producers. Surplus permits make later sends retry + // a full channel instead of parking, especially after dropping a lagging subscription. let waiting = self.waiting_senders(); if freed > 0 && waiting > 0 { self.tx_permits.release_if_nonempty(freed.min(waiting)); } } - /// Wakes every parked producer, however many slots came back. + /// Wakes every blocked producer when the last subscription leaves. /// - /// The last subscription leaving is not a reclaim of some number of slots — it removes the - /// limit itself, because a channel with no receivers discards instead of retaining. Releasing - /// only as many permits as that final reclaim freed would strand every producer beyond that - /// count, so this is the one release that must be unbounded. + /// Sends now discard payloads without consuming capacity, so every producer can proceed. fn release_all(&self) { if self.waiting_senders() > 0 { self.tx_permits.notify_all(); } } - /// How many producers might be waiting, answered without touching the semaphore's lock. + /// Returns an upper bound on parked producers without locking the semaphore. /// - /// This cannot miss a wake-up. A producer increments the count before it ever takes the - /// channel lock to recheck capacity, and every caller here loads it after releasing that same - /// lock, so the mutex orders the two: either this load observes the producer, or the - /// producer's recheck runs after the change and finds the capacity itself. + /// Producers increment before rechecking capacity under the channel lock; reclaim paths load + /// after releasing it. A producer is either covered by this count or rechecks capacity after + /// the reclaim, so skipping or capping notifications cannot strand it. fn waiting_senders(&self) -> usize { self.blocked_senders.load(Ordering::Acquire) } @@ -429,36 +403,26 @@ impl BoundedSender { self.publish(msg, |msg| msg) } - /// The publish step both send paths share. - /// - /// `into_msg` is called only once this decides the message will actually be retained, which is - /// what lets `try_send` defer its allocation past the capacity check while `try_publish` hands - /// over an `Arc` it allocated with the channel unlocked. + /// Publishes a message for both send paths. /// - /// Publishing and draining the wait set share one critical section, so a receiver can never - /// observe an empty buffer and park after this message became visible. + /// Calls `into_msg` only when retaining the payload, so `try_send` can defer its allocation + /// until capacity is available while `try_publish` passes through its existing `Arc`. fn publish

(&self, payload: P, into_msg: impl FnOnce(P) -> Arc) -> Result<(), P> { let mut discarded = None; - let mut wakers = WakerBatch::new(); - { - let mut inner = self.shared.inner.lock(); - - if !inner.log.has_receivers() { - // Nothing can read this message. The payload leaves the critical section with us - // and is dropped below, so `T::drop` never runs under the lock. - inner.log.publish_discarded(); - discarded = Some(payload); - } else if inner.log.retained() == self.shared.capacity { - // Nothing was published, so there is no wait set to drain. - return Err(payload); - } else { - inner.log.publish_retained(into_msg(payload)); - } - - inner.waiters.drain_into(&mut wakers); + let mut inner = self.shared.inner.lock(); + if !inner.log.has_receivers() { + // Drop the discarded payload after unlocking: its destructor may reenter the channel. + inner.log.publish_discarded(); + discarded = Some(payload); + } else if inner.log.retained() == self.shared.capacity { + // Leave waiters registered because no message was published. + return Err(payload); + } else { + inner.log.publish_retained(into_msg(payload)); } - - wake_all(&mut wakers); + let mut wakers = inner.waiters.take_all(); + drop(inner); + wakers.by_ref().for_each(Waker::wake); drop(discarded); Ok(()) } diff --git a/asyncband/src/broadcast/mpmc/common.rs b/asyncband/src/broadcast/mpmc/common.rs index 4b2d321b..73d96e0f 100644 --- a/asyncband/src/broadcast/mpmc/common.rs +++ b/asyncband/src/broadcast/mpmc/common.rs @@ -37,7 +37,6 @@ use super::error::TryRecvError; use crate::internal::arena::Arena; use crate::internal::arena::SlotId; use crate::internal::mutex::Mutex; -use crate::internal::wake_all; use crate::internal::wakerset::WakerSet; use crate::internal::wakerset::WakerToken; @@ -320,7 +319,7 @@ impl Backlog { /// subscription for the slowest cursor, so advancing the head costs the messages released /// instead of the receivers subscribed. /// - /// A receive releases exactly one message: its cursor is counted at the next version before + /// A `receive` releases exactly one message: its cursor is counted at the next version before /// it leaves `head`, so the zero-count prefix ends there. Only removing a lagging subscription /// can release more. /// @@ -406,10 +405,8 @@ impl Backlog { /// Buffer, receiver cursors, and parked receivers, all under one lock. /// -/// The wait set lives beside the backlog so that publishing a message and draining the waiters -/// happen in one critical section. That is what makes the park path race-free: a receiver that -/// finds no message and then registers still holds this lock, so a concurrent send cannot slip -/// between the two steps and skip the wake-up. +/// Checking the backlog and registering a waker share this lock with publishing and detaching +/// registrations, so a send cannot slip between an empty check and registration. pub struct Inner { pub log: Backlog, pub waiters: WakerSet, @@ -432,14 +429,11 @@ impl Inner { /// /// Both families call this from the last sender's `Drop`. pub fn disconnect(inner: &Mutex>) { - let wakers = { - let mut inner = inner.lock(); - inner.waiters.take_all() - }; - wake_all(wakers); + let wakers = mem::take(&mut inner.lock().waiters); + wakers.wake_all(); } -/// Releases a cancelled receive's waker registration, dropping the waker unlocked. +/// Removes the waker registration for a cancelled `receive`, dropping the waker unlocked. pub fn unregister( inner: &Mutex>, senders: &AtomicUsize, @@ -480,9 +474,8 @@ pub fn try_receive( /// The one poll step behind `recv` on both channels. /// -/// Checking the backlog and registering a waker under the same lock prevents a publication from -/// landing between those steps. Publication and disconnection detach all registrations, so their -/// ready paths clear the token without unregistering it. +/// Publication and disconnection detach all registrations, so their ready paths clear the token +/// without unregistering it. pub fn poll_receive( inner: &Mutex>, senders: &AtomicUsize, @@ -582,7 +575,7 @@ mod tests { assert_eq!(log.remove_receiver(b).len(), 3); log.assert_cursor_accounting(); - // A receive that catches up to the tail releases exactly one message and moves the cursor + // A `receive` that catches up to the tail releases exactly one message and moves the cursor // back to `at_tail`. assert!(log.publish(Arc::new(4)).is_none()); let (msg, reclaimed) = log.receive(a).unwrap(); diff --git a/asyncband/src/broadcast/mpmc/unbounded/mod.rs b/asyncband/src/broadcast/mpmc/unbounded/mod.rs index 663e5a6b..17f9d111 100644 --- a/asyncband/src/broadcast/mpmc/unbounded/mod.rs +++ b/asyncband/src/broadcast/mpmc/unbounded/mod.rs @@ -62,6 +62,7 @@ use std::sync::atomic::AtomicUsize; use std::sync::atomic::Ordering; use std::task::Context; use std::task::Poll; +use std::task::Waker; use super::common; use super::common::Backlog; @@ -70,8 +71,6 @@ use super::error::RecvError; use super::error::TryRecvError; use crate::internal::arena::SlotId; use crate::internal::mutex::Mutex; -use crate::internal::wake_all; -use crate::internal::waker_batch::WakerBatch; use crate::internal::wakerset::WakerToken; #[cfg(test)] @@ -175,19 +174,13 @@ impl UnboundedSender { pub fn send(&self, msg: T) { let msg = Arc::new(msg); - // Publishing and draining the wait set share one critical section, so a receiver can never - // observe an empty buffer and park after this message became visible. - let mut wakers = WakerBatch::new(); - let unretained = { - let mut inner = self.shared.inner.lock(); - let unretained = inner.log.publish(msg); - inner.waiters.drain_into(&mut wakers); - unretained - }; + let mut inner = self.shared.inner.lock(); + let unretained = inner.log.publish(msg); + let mut wakers = inner.waiters.take_all(); + drop(inner); - // Notify all waiting receivers. An unsent message is dropped here too, once the lock is - // released. - wake_all(&mut wakers); + // Wake callbacks and payload destruction may reenter the channel. + wakers.by_ref().for_each(Waker::wake); drop(unretained); } diff --git a/asyncband/src/completion/mod.rs b/asyncband/src/completion/mod.rs index 2055ec1c..95ce3c15 100644 --- a/asyncband/src/completion/mod.rs +++ b/asyncband/src/completion/mod.rs @@ -50,6 +50,7 @@ use std::fmt; use std::future::Future; +use std::mem; use std::pin::Pin; use std::sync::Arc; use std::sync::OnceLock; @@ -58,7 +59,6 @@ use std::task::Context; use std::task::Poll; use crate::internal::mutex::Mutex; -use crate::internal::wake_all; use crate::internal::wakerset::WakerSet; use crate::internal::wakerset::WakerToken; @@ -125,17 +125,14 @@ impl Completer { let Some(shared) = self.shared.upgrade() else { return Err(value); }; - let wakers = { - let mut waiters = shared.waiters.lock(); - let wakers = waiters.take_all(); - // The single completer publishes only after every waiter token has been invalidated. - assert!(shared.result.set(Some(value)).is_ok()); - wakers - }; - // `complete` consumes the only completer. Disarm its destructor before invoking arbitrary - // wake callbacks; the completed state no longer needs abandonment handling. + let mut waiters = shared.waiters.lock(); + // Detach registrations before publishing completion to lock-free observers. + let wakers = mem::take(&mut *waiters); + assert!(shared.result.set(Some(value)).is_ok()); + drop(waiters); + // Disarm abandonment handling before invoking wake callbacks. self.shared = Weak::new(); - wake_all(wakers); + wakers.wake_all(); Ok(()) } } @@ -145,13 +142,11 @@ impl Drop for Completer { let Some(shared) = self.shared.upgrade() else { return; }; - let wakers = { - let mut waiters = shared.waiters.lock(); - let wakers = waiters.take_all(); - assert!(shared.result.set(None).is_ok()); - wakers - }; - wake_all(wakers); + let mut waiters = shared.waiters.lock(); + let wakers = mem::take(&mut *waiters); + assert!(shared.result.set(None).is_ok()); + drop(waiters); + wakers.wake_all(); } } diff --git a/asyncband/src/condvar/mod.rs b/asyncband/src/condvar/mod.rs index f2b06524..9cc2e8db 100644 --- a/asyncband/src/condvar/mod.rs +++ b/asyncband/src/condvar/mod.rs @@ -68,7 +68,6 @@ use std::task::Waker; 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; use crate::mutex; use crate::mutex::MutexGuard; @@ -161,21 +160,17 @@ impl Condvar { { let mut waiters = self.waiters.lock(); - while waiters - .unlink_first_waiter(|node| { - let WaitState::Waiting(waker) = - mem::replace(&mut node.state, WaitState::NotifiedAll) - else { - unreachable!("only waiting tasks remain linked") - }; - wakers.push(waker); - true - }) - .is_some() - {} + while let Some((_, node)) = waiters.unlink_first_waiter(|_| true) { + let WaitState::Waiting(waker) = + mem::replace(&mut node.state, WaitState::NotifiedAll) + else { + unreachable!("only waiting tasks remain linked") + }; + wakers.push(waker); + } } - wake_all(&mut wakers); + wakers.by_ref().for_each(Waker::wake); } /// Waits for a notification, atomically releasing and then reacquiring the mutex. diff --git a/asyncband/src/event/manual_reset.rs b/asyncband/src/event/manual_reset.rs index 4e0e2ce9..39d5bcf9 100644 --- a/asyncband/src/event/manual_reset.rs +++ b/asyncband/src/event/manual_reset.rs @@ -27,7 +27,6 @@ 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; use crate::internal::waker_batch::WakerBatch; /// A reusable signal that releases all waiters and remains set until explicitly reset. @@ -117,8 +116,7 @@ impl ManualResetEvent { } } } - - wake_all(&mut wakers); + wakers.by_ref().for_each(Waker::wake); } /// Clears the set state. diff --git a/asyncband/src/internal/arena.rs b/asyncband/src/internal/arena.rs index ce0bc7bd..0cd8647f 100644 --- a/asyncband/src/internal/arena.rs +++ b/asyncband/src/internal/arena.rs @@ -60,6 +60,12 @@ pub struct Arena { len: usize, } +impl Default for Arena { + fn default() -> Self { + Self::new() + } +} + #[derive(Debug)] enum Slot { Occupied(T), @@ -166,32 +172,31 @@ impl Arena { value } - /// Drains every occupied value in slot order while retaining the allocation for reuse. + /// Collects occupied values in slot order, retaining the allocation for reuse. /// - /// After a non-empty drain, every previously issued slot ID becomes invalid, including IDs for - /// slots that were already vacant. Consumers that retain IDs across this operation must supply - /// their own epoch check. + /// All previous slot IDs become invalid and may be reused after collection. #[inline] - pub fn drain(&mut self) -> impl Iterator + '_ { + pub fn take_all>(&mut self) -> C { self.vacant_head = None; self.len = 0; - self.slots.drain(..).filter_map(|slot| match slot { - Slot::Occupied(value) => Some(value), - Slot::Vacant { .. } => None, - }) - } - - /// Takes every occupied value and the backing allocation in slot order. - #[inline] - pub fn take_all(&mut self) -> impl Iterator + use { - self.vacant_head = None; - self.len = 0; - mem::take(&mut self.slots) - .into_iter() + self.slots + .drain(..) .filter_map(|slot| match slot { Slot::Occupied(value) => Some(value), Slot::Vacant { .. } => None, }) + .collect() + } + + /// Consumes the arena, yielding occupied values in slot order. + /// + /// The returned iterator owns the backing allocation and releases it when dropped. + #[inline] + pub fn into_iter(self) -> impl Iterator { + self.slots.into_iter().filter_map(|slot| match slot { + Slot::Occupied(value) => Some(value), + Slot::Vacant { .. } => None, + }) } } @@ -219,7 +224,7 @@ mod tests { } #[test] - fn drain_restarts_slot_id_allocation() { + fn take_all_retains_capacity_and_restarts_slot_ids() { let mut arena = Arena::with_capacity(3); let first = arena.insert(1); let second = arena.insert(2); @@ -227,24 +232,23 @@ mod tests { let capacity = arena.slots.capacity(); arena.remove(second); - assert_eq!(arena.drain().collect::>(), vec![1, 3]); + assert_eq!(arena.take_all::>(), vec![1, 3]); assert_eq!(arena.len(), 0); assert_eq!(arena.slots.capacity(), capacity); - - let slot_ids = [arena.insert(4), arena.insert(5), arena.insert(6)]; - assert_eq!(slot_ids, [first, second, third]); + assert_eq!( + [arena.insert(4), arena.insert(5), arena.insert(6)], + [first, second, third] + ); } #[test] - fn take_all_releases_the_backing_allocation() { + fn into_iter_skips_vacant_slots() { let mut arena = Arena::new(); arena.insert(1); let removed = arena.insert(2); arena.insert(3); arena.remove(removed); - let values = arena.take_all(); - assert_eq!(arena.slots.capacity(), 0); - assert_eq!(values.collect::>(), vec![1, 3]); + assert_eq!(arena.into_iter().collect::>(), vec![1, 3]); } } diff --git a/asyncband/src/internal/countdown.rs b/asyncband/src/internal/countdown.rs index 53ff5e35..eb37cb93 100644 --- a/asyncband/src/internal/countdown.rs +++ b/asyncband/src/internal/countdown.rs @@ -15,13 +15,13 @@ // specific language governing permissions and limitations // under the License. +use std::mem; use std::sync::atomic::AtomicU32; use std::sync::atomic::Ordering; use std::task::Context; use std::task::Poll; use crate::internal::mutex::Mutex; -use crate::internal::wake_all; use crate::internal::wakerset::WakerSet; use crate::internal::wakerset::WakerToken; @@ -53,14 +53,10 @@ impl CountdownState { .map(|_| ()) } - /// Drains the waiter set under its lock, then wakes every waiter after releasing the lock. + /// Detaches the waiter set under its lock, then wakes every waiter after releasing the lock. pub fn wake_all(&self) { - let wakers = { - let mut waiters = self.waiters.lock(); - waiters.take_all() - }; - - wake_all(wakers); + let wakers = mem::take(&mut *self.waiters.lock()); + wakers.wake_all(); } /// Polls for zero, registering the current waker if the countdown is still active. @@ -74,7 +70,7 @@ impl CountdownState { let mut waiters = self.waiters.lock(); if self.state() == 0 { - // A concurrent zero transition will drain after this lock is released. + // A concurrent zero transition detaches the registrations under this same lock. *token = None; return Poll::Ready(()); } diff --git a/asyncband/src/internal/mod.rs b/asyncband/src/internal/mod.rs index 76d3ce80..f9dceb04 100644 --- a/asyncband/src/internal/mod.rs +++ b/asyncband/src/internal/mod.rs @@ -33,16 +33,6 @@ pub fn register_waker(slot: &mut Option, waker: &Waker) -> Option } } -/// Wakes every waker. -#[inline] -// A no-feature or blocking-only build has no primitive that fans notifications out. -#[allow(dead_code)] -pub(crate) fn wake_all(wakers: impl Iterator) { - for waker in wakers { - waker.wake(); - } -} - #[cfg(any( feature = "barrier", feature = "broadcast", @@ -123,8 +113,8 @@ pub(crate) mod waitlist; #[cfg(any( feature = "barrier", feature = "broadcast", - feature = "event", feature = "completion", + feature = "event", feature = "latch", feature = "mutex", feature = "once", @@ -134,9 +124,6 @@ pub(crate) mod waitlist; feature = "waitgroup", feature = "watch", ))] -// Only the semaphore refills a batch and asks whether it will spill, so other feature subsets -// leave that method unused. -#[allow(dead_code)] pub(crate) mod waker_batch; #[cfg(any( @@ -149,7 +136,6 @@ pub(crate) mod waker_batch; feature = "waitgroup", feature = "watch", ))] -// Reusable waker sets and terminal primitives use different lifecycle policies, so some feature -// subsets leave one constructor or detach operation unused. +// Some feature subsets use only one constructor or one of the two notification paths. #[allow(dead_code)] pub(crate) mod wakerset; diff --git a/asyncband/src/internal/semaphore.rs b/asyncband/src/internal/semaphore.rs index 51c7baed..4ce4ad71 100644 --- a/asyncband/src/internal/semaphore.rs +++ b/asyncband/src/internal/semaphore.rs @@ -37,7 +37,6 @@ 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; use crate::internal::waker_batch::WakerBatch; /// The internal semaphore that provides low-level async primitives. @@ -173,7 +172,7 @@ impl Semaphore { } } drop(waiters); - wake_all(&mut wakers); + wakers.by_ref().for_each(Waker::wake); } fn insert_permits_with_lock<'a>( @@ -212,7 +211,7 @@ impl Semaphore { } drop(waiters); - wake_all(&mut batch); + batch.by_ref().for_each(Waker::wake); if rem == 0 { return; } diff --git a/asyncband/src/internal/waker_batch.rs b/asyncband/src/internal/waker_batch.rs index a41c0f37..f1ec29e0 100644 --- a/asyncband/src/internal/waker_batch.rs +++ b/asyncband/src/internal/waker_batch.rs @@ -22,11 +22,8 @@ use std::task::Waker; /// An owning FIFO of wakers that stores the first [`Self::STACK_SIZE`] entries without allocating. /// -/// The batch is filled through [`WakerBatch::push`] or [`Extend`] and consumed as its own -/// iterator. Entries are written only as they are pushed, so constructing an empty or small batch -/// touches nothing beyond the two indices. Once every inline entry has been yielded the batch -/// reuses that storage, so a caller that alternates between filling and draining, as the -/// semaphore does, keeps running on the stack. +/// Only pushed entries are initialized. Iterate by reference to avoid moving the inline storage +/// into iterator adapters and to reuse the batch, as the semaphore does between notifications. pub struct WakerBatch { /// The initialized entries are exactly `start..end`. inline: [MaybeUninit; Self::STACK_SIZE], @@ -42,10 +39,9 @@ pub struct WakerBatch { } impl WakerBatch { - /// Wakers kept on the stack before the batch spills to the heap. + /// Number of wakers stored inline before using overflow storage. /// - /// This is also the most wakers the semaphore collects per lock acquisition, so a drain that - /// wakes a typical waiter set never allocates; larger sets pay one allocation for the overflow. + /// The semaphore's permit-release loop uses this as its batch limit before unlocking. pub const STACK_SIZE: usize = 32; pub const fn new() -> Self { @@ -61,31 +57,37 @@ impl WakerBatch { /// /// The semaphore stops filling a batch here so it can release its lock and wake what it has /// before collecting more. + #[inline] pub fn will_spill(&self) -> bool { self.end == Self::STACK_SIZE || !self.spilled.is_empty() } + #[inline] pub fn push(&mut self, waker: Waker) { - if self.end < Self::STACK_SIZE && self.spilled.is_empty() { + if self.will_spill() { + self.spilled.push_back(waker); + } else { self.inline[self.end].write(waker); self.end += 1; - } else { - self.spilled.push_back(waker); } } } -impl Extend for WakerBatch { - fn extend>(&mut self, iter: I) { +impl FromIterator for WakerBatch { + #[inline] + fn from_iter>(iter: T) -> Self { + let mut batch = Self::new(); for waker in iter { - self.push(waker); + batch.push(waker); } + batch } } impl Iterator for WakerBatch { type Item = Waker; + #[inline] fn next(&mut self) -> Option { if self.start < self.end { let index = self.start; @@ -102,6 +104,7 @@ impl Iterator for WakerBatch { } impl Drop for WakerBatch { + #[inline] fn drop(&mut self) { let initialized = ptr::slice_from_raw_parts_mut( self.inline[self.start..self.end] @@ -153,8 +156,10 @@ mod tests { })) } - fn wakers(log: &Arc, count: usize) -> impl Iterator + '_ { - (0..count).map(move |id| waker(log, id)) + fn push_wakers(batch: &mut WakerBatch, log: &Arc, count: usize) { + for id in 0..count { + batch.push(waker(log, id)); + } } fn alive(log: &Arc) -> usize { @@ -169,8 +174,7 @@ mod tests { fn yields_in_push_order_across_the_spill() { let log = log(); let count = STACK_SIZE + 8; - let mut batch = WakerBatch::new(); - batch.extend(wakers(&log, count)); + let mut batch: WakerBatch = (0..count).map(|id| waker(&log, id)).collect(); assert!(batch.will_spill()); for waker in &mut batch { @@ -188,7 +192,7 @@ mod tests { let count = STACK_SIZE + 8; for consumed in [0, 5, STACK_SIZE, STACK_SIZE + 3, count] { let mut batch = WakerBatch::new(); - batch.extend(wakers(&log, count)); + push_wakers(&mut batch, &log, count); for _ in 0..consumed { drop(batch.next().unwrap()); } @@ -204,7 +208,7 @@ mod tests { let log = log(); let mut batch = WakerBatch::new(); for _ in 0..3 { - batch.extend(wakers(&log, STACK_SIZE)); + push_wakers(&mut batch, &log, STACK_SIZE); assert!(batch.will_spill()); assert_eq!(batch.by_ref().count(), STACK_SIZE); assert!(!batch.will_spill()); @@ -216,7 +220,7 @@ mod tests { fn keeps_push_order_while_spilled() { let log = log(); let mut batch = WakerBatch::new(); - batch.extend(wakers(&log, STACK_SIZE + 1)); + push_wakers(&mut batch, &log, STACK_SIZE + 1); // Free inline room; the spilled entry must still come out before anything pushed now. for _ in 0..4 { batch.next().unwrap().wake(); diff --git a/asyncband/src/internal/wakerset.rs b/asyncband/src/internal/wakerset.rs index c74a89b5..75649cba 100644 --- a/asyncband/src/internal/wakerset.rs +++ b/asyncband/src/internal/wakerset.rs @@ -15,11 +15,11 @@ // specific language governing permissions and limitations // under the License. -//! Cancellable storage for task wakers whose lifecycle is owned by the caller. +//! Cancellable task wakers protected by the owning primitive's state lock. //! -//! A `WakerSet` is protected by the state lock of its owning primitive. The owner must clear a -//! token instead of unregistering it after an operation that detached the set. This lets each -//! primitive use its existing terminal state or generation to recognize stale registrations. +//! The primitive uses its generation or terminal state to recognize detached registrations; +//! their tokens must be cleared rather than passed back to the set. Returned wakers must be woken +//! or dropped after releasing the lock. use std::mem; use std::task::Waker; @@ -30,14 +30,12 @@ use crate::internal::waker_batch::WakerBatch; /// An exclusive handle to one waker slot in a [`WakerSet`]. /// -/// This token deliberately does not implement `Clone` or `Copy`. Its owner must not pass it back -/// to the set after the registration has been detached by [`WakerSet::drain_into`] or -/// [`WakerSet::take_all`]. +/// Removing the registration, calling [`WakerSet::take_all`], or replacing the set invalidates it. #[derive(Debug)] pub struct WakerToken(SlotId); /// Cancellable waker storage without an implicit lifecycle or generation. -#[derive(Debug)] +#[derive(Default, Debug)] pub struct WakerSet { wakers: Arena, } @@ -57,38 +55,36 @@ impl WakerSet { } } - /// Drains all registered wakers into `batch` while retaining slot capacity. + /// Collects registered wakers into an owned batch, retaining slot capacity for reuse. /// - /// The batch is filled in place because its inline storage is too large to move for free: - /// returning it by value costs every publish about 6ns even when nothing is registered. The - /// caller must invalidate every outstanding token and consume or drop the batch after - /// releasing the lock that protects this set. + /// Moves each waker into the batch; up to [`WakerBatch::STACK_SIZE`] fit without allocating. #[inline] - pub fn drain_into(&mut self, batch: &mut WakerBatch) { - batch.extend(self.wakers.drain()); + pub fn take_all(&mut self) -> WakerBatch { + if self.wakers.is_empty() { + return WakerBatch::new(); + } + self.wakers.take_all() } - /// Takes all registered wakers together with the set's backing allocation. + /// Consumes the set, waking every registered waker and releasing its allocation. /// - /// The caller must invalidate every outstanding token and consume or drop the iterator after - /// releasing the lock that protects this set. + /// Call after releasing the owning primitive's state lock. #[inline] - pub fn take_all(&mut self) -> impl Iterator + 'static { - self.wakers.take_all() + pub fn wake_all(self) { + self.wakers.into_iter().for_each(Waker::wake); } /// Registers or updates a waker. /// - /// If an existing waker is replaced, it is returned so the caller can drop it after releasing - /// the lock that protects this set. + /// Returns the previous waker only if it was replaced. #[inline] #[must_use = "drop the returned waker after releasing the waker set's state lock"] pub fn register(&mut self, token: &mut Option, waker: &Waker) -> Option { - if let Some(current) = token.as_ref().map(|token| { - self.wakers + if let Some(token) = token { + let current = self + .wakers .get_mut(token.0) - .expect("waker token must refer to an occupied slot") - }) { + .expect("waker token must refer to an occupied slot"); if current.will_wake(waker) { return None; } @@ -99,10 +95,7 @@ impl WakerSet { None } - /// Removes the waker identified by `token`. - /// - /// The owner must clear stale tokens without calling this method after detaching the set. The - /// returned waker must be dropped after releasing the lock that protects this set. + /// Removes and returns the waker identified by `token`, clearing the token. #[inline] #[must_use = "drop the returned waker after releasing the waker set's state lock"] pub fn unregister(&mut self, token: &mut Option) -> Option { diff --git a/asyncband/src/phaser/mod.rs b/asyncband/src/phaser/mod.rs index 9b00e7c5..5d97a51f 100644 --- a/asyncband/src/phaser/mod.rs +++ b/asyncband/src/phaser/mod.rs @@ -124,13 +124,14 @@ use std::fmt; use std::future::Future; use std::iter::FusedIterator; +use std::mem; use std::pin::Pin; use std::sync::Arc; use std::task::Context; use std::task::Poll; +use std::task::Waker; use crate::internal::mutex::Mutex; -use crate::internal::wake_all; use crate::internal::waker_batch::WakerBatch; use crate::internal::wakerset::WakerSet; use crate::internal::wakerset::WakerToken; @@ -172,14 +173,14 @@ struct State { } impl State { - /// Completes the phase once every participant has arrived, moving its waiters into `wakers`. - fn advance_if_ready(&mut self, wakers: &mut WakerBatch) { + /// Advances a completed phase and returns its waiters for waking outside the lock. + fn advance_if_ready(&mut self) -> WakerBatch { if self.closed || self.unarrived != 0 { - return; + return WakerBatch::new(); } self.phase = self.phase.wrapping_add(1); self.unarrived = self.registered; - self.waiters.drain_into(wakers); + self.waiters.take_all() } fn completion(&self, observed: u64) -> Poll> { @@ -247,9 +248,9 @@ impl Phaser { return; } state.closed = true; - state.waiters.take_all() + mem::take(&mut state.waiters) }; - wake_all(wakers); + wakers.wake_all(); } /// Returns an instantaneous count of registered participants, including those already arrived. @@ -385,16 +386,14 @@ impl Drop for PhaserParticipants { if self.remaining == 0 { return; } - let mut wakers = WakerBatch::new(); - { - let mut state = self.phaser.state.lock(); - // Unyielded participants have never arrived and prevent their phase from advancing. - state.registered -= self.remaining; - state.unarrived -= self.remaining; - self.remaining = 0; - state.advance_if_ready(&mut wakers); - } - wake_all(&mut wakers); + let mut state = self.phaser.state.lock(); + // Unyielded participants have never arrived and prevent their phase from advancing. + state.registered -= self.remaining; + state.unarrived -= self.remaining; + self.remaining = 0; + let mut wakers = state.advance_if_ready(); + drop(state); + wakers.by_ref().for_each(Waker::wake); } } @@ -433,21 +432,18 @@ impl PhaserParticipant { /// Repeated calls within one phase count only once. After advancement, an explicit new call /// arrives in the new phase and replaces any previous pending observation. pub fn arrive(&mut self) -> Result { - let mut wakers = WakerBatch::new(); - let phase = { - let mut state = self.phaser.state.lock(); - if state.closed { - return Err(Closed(())); - } - let phase = state.phase; - if self.pending != Some(phase) { - state.unarrived -= 1; - } - self.pending = Some(phase); - state.advance_if_ready(&mut wakers); - phase - }; - wake_all(&mut wakers); + let mut state = self.phaser.state.lock(); + if state.closed { + return Err(Closed(())); + } + let phase = state.phase; + if self.pending != Some(phase) { + state.unarrived -= 1; + } + self.pending = Some(phase); + let mut wakers = state.advance_if_ready(); + drop(state); + wakers.by_ref().for_each(Waker::wake); Ok(phase) } @@ -482,23 +478,20 @@ impl PhaserParticipant { } fn do_deregister(&mut self) -> Result { - let mut wakers = WakerBatch::new(); - let result = { - let mut state = self.phaser.state.lock(); - self.registered = false; - state.registered -= 1; - if self.pending != Some(state.phase) { - state.unarrived -= 1; - } - let result = if state.closed { - Err(Closed(())) - } else { - Ok(state.phase) - }; - state.advance_if_ready(&mut wakers); - result + let mut state = self.phaser.state.lock(); + self.registered = false; + state.registered -= 1; + if self.pending != Some(state.phase) { + state.unarrived -= 1; + } + let result = if state.closed { + Err(Closed(())) + } else { + Ok(state.phase) }; - wake_all(&mut wakers); + let mut wakers = state.advance_if_ready(); + drop(state); + wakers.by_ref().for_each(Waker::wake); result } } diff --git a/asyncband/src/waitgroup/mod.rs b/asyncband/src/waitgroup/mod.rs index b7e445a4..4fb7c403 100644 --- a/asyncband/src/waitgroup/mod.rs +++ b/asyncband/src/waitgroup/mod.rs @@ -59,6 +59,7 @@ use std::fmt; use std::future::Future; use std::future::IntoFuture; +use std::mem; use std::pin::Pin; use std::sync::Arc; use std::sync::atomic::AtomicUsize; @@ -67,7 +68,6 @@ use std::task::Context; use std::task::Poll; use crate::internal::mutex::Mutex; -use crate::internal::wake_all; use crate::internal::wakerset::WakerSet; use crate::internal::wakerset::WakerToken; @@ -102,11 +102,8 @@ impl State { return; } - let wakers = { - let mut waiters = self.waiters.lock(); - waiters.take_all() - }; - wake_all(wakers); + let wakers = mem::take(&mut *self.waiters.lock()); + wakers.wake_all(); } fn poll_wait(&self, token: &mut Option, cx: &mut Context<'_>) -> Poll<()> { diff --git a/asyncband/src/watch/mod.rs b/asyncband/src/watch/mod.rs index 90b71923..c2f20721 100644 --- a/asyncband/src/watch/mod.rs +++ b/asyncband/src/watch/mod.rs @@ -70,12 +70,11 @@ use std::pin::Pin; use std::sync::Arc; use std::task::Context; use std::task::Poll; +use std::task::Waker; pub use self::error::RecvError; pub use self::error::SendError; use crate::internal::mutex::Mutex; -use crate::internal::wake_all; -use crate::internal::waker_batch::WakerBatch; use crate::internal::wakerset::WakerSet; use crate::internal::wakerset::WakerToken; @@ -147,16 +146,15 @@ impl fmt::Debug for Sender { impl Drop for Sender { fn drop(&mut self) { - // Only the final sender detaches the parked receivers; their wake callbacks run unlocked. let wakers = { let mut state = self.shared.state.lock(); state.senders -= 1; if state.senders != 0 { return; } - state.waiters.take_all() + mem::take(&mut state.waiters) }; - wake_all(wakers); + wakers.wake_all(); } } @@ -169,23 +167,21 @@ impl Sender { /// /// Panics if the channel has already published `u64::MAX` updates. pub fn send(&self, value: T) -> Result<(), SendError> { - let mut wakers = WakerBatch::new(); - let replaced = { - let mut state = self.shared.state.lock(); - if state.receivers == 0 { - return Err(SendError::new(value)); - } - let version = state - .version - .checked_add(1) - .expect("watch channel version counter overflowed"); - let replaced = mem::replace(&mut state.value, value); - state.version = version; - state.waiters.drain_into(&mut wakers); - replaced - }; + let mut state = self.shared.state.lock(); + if state.receivers == 0 { + return Err(SendError::new(value)); + } + let version = state + .version + .checked_add(1) + .expect("watch channel version counter overflowed"); + let replaced = mem::replace(&mut state.value, value); + state.version = version; + let mut wakers = state.waiters.take_all(); + drop(state); + // Waker callbacks and the replaced value's destructor may reenter this channel. - wake_all(&mut wakers); + wakers.by_ref().for_each(Waker::wake); drop(replaced); Ok(()) } @@ -199,19 +195,16 @@ impl Sender { /// /// Panics if the channel has already published `u64::MAX` updates. pub fn send_replace(&self, value: T) -> T { - let mut wakers = WakerBatch::new(); - let replaced = { - let mut state = self.shared.state.lock(); - let version = state - .version - .checked_add(1) - .expect("watch channel version counter overflowed"); - let replaced = mem::replace(&mut state.value, value); - state.version = version; - state.waiters.drain_into(&mut wakers); - replaced - }; - wake_all(&mut wakers); + let mut state = self.shared.state.lock(); + let version = state + .version + .checked_add(1) + .expect("watch channel version counter overflowed"); + let replaced = mem::replace(&mut state.value, value); + state.version = version; + let mut wakers = state.waiters.take_all(); + drop(state); + wakers.by_ref().for_each(Waker::wake); replaced } diff --git a/tests-integration/tests/phaser_test.rs b/tests-integration/tests/phaser_test.rs index 28c225a1..af81580d 100644 --- a/tests-integration/tests/phaser_test.rs +++ b/tests-integration/tests/phaser_test.rs @@ -343,6 +343,53 @@ fn advancing_a_phase_wakes_every_registered_waiter_once() { )); } +#[test] +fn partial_arrivals_and_withdrawals_preserve_pending_waits() { + let phaser = Phaser::new(); + let observed = phaser.phase(); + let mut participants = phaser.register(4).unwrap(); + let mut first = participants.next().unwrap(); + let withdrawing = participants.next().unwrap(); + let mut last = participants.next().unwrap(); + let (waker, counter) = WakeCounter::new(); + let mut wait = Box::pin(phaser.wait(observed)); + assert!(poll_with(wait.as_mut(), &waker).is_pending()); + + first.arrive().unwrap(); + first.arrive().unwrap(); + withdrawing.deregister().unwrap(); + drop(participants); + let early_wakes = counter.count(); + + last.arrive().unwrap(); + assert_eq!(early_wakes, 0); + assert_eq!(counter.count(), 1); + assert_eq!(poll_once(wait.as_mut()), Poll::Ready(Ok(phaser.phase()))); +} + +#[test] +fn cancelling_after_partial_arrival_preserves_other_waiters() { + let phaser = Phaser::new(); + let observed = phaser.phase(); + let mut first = phaser.register_one().unwrap(); + let mut second = phaser.register_one().unwrap(); + let mut cancelled = Box::pin(phaser.wait(observed)); + assert!(poll_once(cancelled.as_mut()).is_pending()); + + first.arrive().unwrap(); + let (waker, counter) = WakeCounter::new(); + let mut remaining = Box::pin(phaser.wait(observed)); + assert!(poll_with(remaining.as_mut(), &waker).is_pending()); + drop(cancelled); + + second.arrive().unwrap(); + assert_eq!(counter.count(), 1); + assert_eq!( + poll_once(remaining.as_mut()), + Poll::Ready(Ok(phaser.phase())) + ); +} + #[test] fn repolling_updates_the_task_that_will_be_notified() { let phaser = Phaser::new();