diff --git a/CHANGELOG.md b/CHANGELOG.md index 66b5f57..65c424a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -32,3 +32,4 @@ All notable changes to this project will be documented in this file. ### Improvements * Remove the `slab` dependency in favor of a focused internal waiter arena. +* Reduce unbounded MPSC contention by avoiding receiver wake operations while no receive call is parked. diff --git a/asyncband/src/channel/mpsc/unbounded.rs b/asyncband/src/channel/mpsc/unbounded.rs index 3ee892a..bb5a37d 100644 --- a/asyncband/src/channel/mpsc/unbounded.rs +++ b/asyncband/src/channel/mpsc/unbounded.rs @@ -21,6 +21,7 @@ use std::fmt; use std::future::poll_fn; use std::sync::Arc; +use std::sync::atomic::AtomicBool; use std::sync::atomic::AtomicUsize; use std::sync::atomic::Ordering; use std::task::Context; @@ -43,6 +44,7 @@ use crate::internal::atomic_waker::AtomicWaker; pub fn unbounded() -> (UnboundedSender, UnboundedReceiver) { let state = Arc::new(UnboundedState { senders: AtomicUsize::new(1), + rx_waiting: AtomicBool::new(false), rx_waker: AtomicWaker::new(), }); let (sender, receiver) = std::sync::mpsc::channel(); @@ -59,9 +61,23 @@ pub fn unbounded() -> (UnboundedSender, UnboundedReceiver) { struct UnboundedState { senders: AtomicUsize, + rx_waiting: AtomicBool, rx_waker: AtomicWaker, } +impl UnboundedState { + fn wake_receiver(&self) { + if self.rx_waiting.load(Ordering::SeqCst) + && self + .rx_waiting + .compare_exchange(true, false, Ordering::SeqCst, Ordering::SeqCst) + .is_ok() + { + self.rx_waker.wake(); + } + } +} + /// Send values to the associated [`UnboundedReceiver`]. /// /// Instances are created by the [`unbounded`] function. @@ -95,7 +111,7 @@ impl Drop for UnboundedSender { 1 => { // If this is the last sender, we need to wake up the receiver so it can // observe the disconnected state. - self.state.rx_waker.wake(); + self.state.wake_receiver(); } _ => { // there are still other senders left, do nothing @@ -118,7 +134,7 @@ impl UnboundedSender { let sender = self.sender.as_ref().unwrap(); sender.send(value).map_err(|err| SendError::new(err.0))?; - self.state.rx_waker.wake(); + self.state.wake_receiver(); Ok(()) } @@ -243,9 +259,19 @@ impl UnboundedReceiver { Err(TryRecvError::Empty) => { self.state.rx_waker.register(cx.waker()); + // Publishing the wait before rechecking the queue prevents a sender from being + // missed: it either precedes the recheck or claims this flag and wakes us. + self.state.rx_waiting.store(true, Ordering::SeqCst); + match self.try_recv() { - Ok(v) => Poll::Ready(Ok(v)), - Err(TryRecvError::Disconnected) => Poll::Ready(Err(RecvError::Disconnected)), + Ok(v) => { + self.state.rx_waiting.store(false, Ordering::SeqCst); + Poll::Ready(Ok(v)) + } + Err(TryRecvError::Disconnected) => { + self.state.rx_waiting.store(false, Ordering::SeqCst); + Poll::Ready(Err(RecvError::Disconnected)) + } Err(TryRecvError::Empty) => Poll::Pending, } } diff --git a/tests-integration/tests/mpsc_test.rs b/tests-integration/tests/mpsc_test.rs index 72256a7..6a057fc 100644 --- a/tests-integration/tests/mpsc_test.rs +++ b/tests-integration/tests/mpsc_test.rs @@ -15,7 +15,14 @@ // specific language governing permissions and limitations // under the License. +use std::future::Future; +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::Instant; use asyncband::mpsc; @@ -33,6 +40,14 @@ fn expect_ready(poll: Poll) -> T { } } +struct WakeProbe(AtomicUsize); + +impl Wake for WakeProbe { + fn wake(self: Arc) { + self.0.fetch_add(1, Ordering::Relaxed); + } +} + #[test] fn test_unbounded_pressure() { let n = 1024 * 1024; @@ -209,6 +224,46 @@ fn try_recv_close_while_empty_unbounded() { assert_eq!(Err(TryRecvError::Disconnected), rx.try_recv()); } +#[test] +fn unbounded_burst_wakes_parked_receiver_once() { + let (tx, mut rx) = mpsc::unbounded(); + let probe = Arc::new(WakeProbe(AtomicUsize::new(0))); + let waker = Waker::from(probe.clone()); + let mut context = Context::from_waker(&waker); + let mut recv = Box::pin(rx.recv()); + + assert!(recv.as_mut().poll(&mut context).is_pending()); + for value in 0..100 { + tx.send(value).unwrap(); + } + + assert_eq!(probe.0.load(Ordering::Relaxed), 1); + assert_eq!(expect_ready(recv.as_mut().poll(&mut context)), Ok(0)); + drop(recv); + + for value in 1..100 { + assert_eq!(rx.try_recv(), Ok(value)); + } +} + +#[test] +fn unbounded_last_sender_drop_wakes_parked_receiver() { + let (tx, mut rx) = mpsc::unbounded::<()>(); + let probe = Arc::new(WakeProbe(AtomicUsize::new(0))); + let waker = Waker::from(probe.clone()); + let mut context = Context::from_waker(&waker); + let mut recv = Box::pin(rx.recv()); + + assert!(recv.as_mut().poll(&mut context).is_pending()); + drop(tx); + + assert_eq!(probe.0.load(Ordering::Relaxed), 1); + assert_eq!( + expect_ready(recv.as_mut().poll(&mut context)), + Err(RecvError::Disconnected) + ); +} + #[tokio::test] async fn send_recv_bounded() { let (tx, mut rx) = mpsc::bounded(1);