diff --git a/asyncband/src/condvar/mod.rs b/asyncband/src/condvar/mod.rs index ffae538..d270776 100644 --- a/asyncband/src/condvar/mod.rs +++ b/asyncband/src/condvar/mod.rs @@ -77,12 +77,7 @@ use crate::mutex::OwnedMutexGuard; /// /// See the [module level documentation](self) for more. pub struct Condvar { - waiters: Mutex>, -} - -#[derive(Debug)] -struct WaitNode { - state: WaitState, + waiters: Mutex>, } #[derive(Debug)] @@ -92,17 +87,12 @@ enum WaitState { NotifiedAll, } -fn notify_one_locked(waiters: &mut WaitList) -> Option { - let mut waker = None; - waiters.unlink_first_waiter(|node| { - let WaitState::Waiting(waiting) = mem::replace(&mut node.state, WaitState::NotifiedOne) - else { - unreachable!("only waiting tasks remain linked") - }; - waker = Some(waiting); - true - }); - waker +fn notify_one_locked(waiters: &mut WaitList) -> Option { + let (_, state) = waiters.unlink_first_waiter(|_| true)?; + let WaitState::Waiting(waker) = mem::replace(state, WaitState::NotifiedOne) else { + unreachable!("only waiting tasks remain linked") + }; + Some(waker) } impl fmt::Debug for Condvar { @@ -160,10 +150,8 @@ impl Condvar { { let mut waiters = self.waiters.lock(); - while let Some((_, node)) = waiters.unlink_first_waiter(|_| true) { - let WaitState::Waiting(waker) = - mem::replace(&mut node.state, WaitState::NotifiedAll) - else { + while let Some((_, state)) = waiters.unlink_first_waiter(|_| true) { + let WaitState::Waiting(waker) = mem::replace(state, WaitState::NotifiedAll) else { unreachable!("only waiting tasks remain linked") }; wakers.push(waker); @@ -341,9 +329,7 @@ where if this.guard.is_some() { let mut waiters = this.condvar.waiters.lock(); - this.index = Some(waiters.push_back(WaitNode { - state: WaitState::Waiting(cx.waker().clone()), - })); + this.index = Some(waiters.push_back(WaitState::Waiting(cx.waker().clone()))); let guard = this.guard.take().unwrap(); // Registration must happen before unlocking the associated mutex. A notifier that @@ -356,7 +342,7 @@ where let index = this.index.expect("wait future polled after completion"); let mut waiters = this.condvar.waiters.lock(); let mut old_waker = None; - let notify_one_baton = match &mut waiters.waiter_mut(index).state { + let notify_one_baton = match waiters.waiter_mut(index) { WaitState::Waiting(waker) => { if !waker.will_wake(cx.waker()) { old_waker = Some(mem::replace(waker, cx.waker().clone())); @@ -385,21 +371,12 @@ impl Drop for Wait<'_, G> { let (waiter, waker) = { let mut waiters = self.condvar.waiters.lock(); - let mut pass_notification = false; - waiters.unlink_waiter(index, |node| match &node.state { - WaitState::Waiting(_) => true, - WaitState::NotifiedOne => { - pass_notification = true; - false - } - WaitState::NotifiedAll => false, - }); + waiters.unlink_waiter(index, |_| true); let waiter = waiters.remove_unlinked_waiter(index); - let waker = if pass_notification { - notify_one_locked(&mut waiters) - } else { - None + let waker = match &waiter { + WaitState::NotifiedOne => notify_one_locked(&mut waiters), + _ => None, }; (waiter, waker) }; diff --git a/tests-integration/tests/condvar_test.rs b/tests-integration/tests/condvar_test.rs index fb7128f..4c2174e 100644 --- a/tests-integration/tests/condvar_test.rs +++ b/tests-integration/tests/condvar_test.rs @@ -178,45 +178,53 @@ fn cancelling_waiter_drops_its_waker_outside_the_waiter_lock() { #[test] fn cancelling_notified_waiter_passes_notify_one_to_next_waiter() { test_runtime().block_on(async { - let mutex = Mutex::new(()); - let condvar = Condvar::new(); + for reacquire in [false, true] { + let mutex = Mutex::new(()); + let condvar = Condvar::new(); - let mut first = Box::pin(condvar.wait(mutex.lock().await)); - assert!(poll_once(first.as_mut()).is_pending()); + let mut first = Box::pin(condvar.wait(mutex.lock().await)); + assert!(poll_once(first.as_mut()).is_pending()); - let mut second = Box::pin(condvar.wait(mutex.lock().await)); - assert!(poll_once(second.as_mut()).is_pending()); + let mut second = Box::pin(condvar.wait(mutex.lock().await)); + assert!(poll_once(second.as_mut()).is_pending()); - let held = mutex.lock().await; - condvar.notify_one(); + let held = mutex.lock().await; + condvar.notify_one(); - // The first waiter consumes the notification, then blocks while reacquiring the mutex. - assert!(poll_once(first.as_mut()).is_pending()); - drop(first); - drop(held); + // Cancel either before the selected waiter is polled again or while it reacquires + // the mutex. Both paths must pass the notification to the next waiter. + if reacquire { + assert!(poll_once(first.as_mut()).is_pending()); + } + drop(first); + drop(held); - // Cancelling the selected waiter passes the notification to an existing waiter. - drop(expect_ready(poll_once(second.as_mut()))); + drop(expect_ready(poll_once(second.as_mut()))); + } }); } #[test] fn cancelling_only_notified_waiter_does_not_buffer_notify_one() { test_runtime().block_on(async { - let mutex = Mutex::new(()); - let condvar = Condvar::new(); + for reacquire in [false, true] { + let mutex = Mutex::new(()); + let condvar = Condvar::new(); - let mut first = Box::pin(condvar.wait(mutex.lock().await)); - assert!(poll_once(first.as_mut()).is_pending()); + let mut first = Box::pin(condvar.wait(mutex.lock().await)); + assert!(poll_once(first.as_mut()).is_pending()); - let held = mutex.lock().await; - condvar.notify_one(); - assert!(poll_once(first.as_mut()).is_pending()); - drop(first); - drop(held); + let held = mutex.lock().await; + condvar.notify_one(); + if reacquire { + assert!(poll_once(first.as_mut()).is_pending()); + } + drop(first); + drop(held); - let mut late = Box::pin(condvar.wait(mutex.lock().await)); - assert!(poll_once(late.as_mut()).is_pending()); + let mut late = Box::pin(condvar.wait(mutex.lock().await)); + assert!(poll_once(late.as_mut()).is_pending()); + } }); }