diff --git a/tokio-util/src/task/join_queue.rs b/tokio-util/src/task/join_queue.rs index e327b93c7..9cdd0cb1e 100644 --- a/tokio-util/src/task/join_queue.rs +++ b/tokio-util/src/task/join_queue.rs @@ -188,6 +188,10 @@ impl JoinQueue { /// Note that on success the handle will panic on subsequent polls /// since it becomes consumed. fn try_poll_handle(jh: &mut AbortOnDropHandle) -> Option> { + if !jh.is_finished() { + return None; + } + let waker = futures_util::task::noop_waker(); let mut cx = Context::from_waker(&waker); diff --git a/tokio-util/tests/task_join_queue.rs b/tokio-util/tests/task_join_queue.rs index a61a968f2..7128f1a34 100644 --- a/tokio-util/tests/task_join_queue.rs +++ b/tokio-util/tests/task_join_queue.rs @@ -1,5 +1,8 @@ #![warn(rust_2018_idioms)] +use std::task::Context; + +use futures_test::task::new_count_waker; use tokio::sync::oneshot; use tokio::task::yield_now; use tokio::time::Duration; @@ -276,6 +279,60 @@ async fn test_join_queue_try_join_next() { check_try_join_next_is_noop(&mut queue); } +#[tokio::test] +async fn test_join_queue_try_join_next_does_not_replace_waker() { + let (send, recv) = oneshot::channel(); + let mut queue = JoinQueue::new(); + queue.spawn(async move { + recv.await.unwrap(); + 42 + }); + + let (waker, wake_count) = new_count_waker(); + let mut cx = Context::from_waker(&waker); + + assert_pending!(queue.poll_join_next(&mut cx)); + assert_eq!(wake_count, 0); + assert!(queue.try_join_next().is_none()); + + send.send(()).unwrap(); + yield_now().await; + + assert_eq!(wake_count, 1); + assert_eq!( + assert_ready!(queue.poll_join_next(&mut cx)) + .unwrap() + .unwrap(), + 42 + ); +} + +#[tokio::test] +async fn test_join_queue_try_join_next_with_id_does_not_replace_waker() { + let (send, recv) = oneshot::channel(); + let mut queue = JoinQueue::new(); + queue.spawn(async move { + recv.await.unwrap(); + 42 + }); + + let (waker, wake_count) = new_count_waker(); + let mut cx = Context::from_waker(&waker); + + assert_pending!(queue.poll_join_next_with_id(&mut cx)); + assert_eq!(wake_count, 0); + assert!(queue.try_join_next_with_id().is_none()); + + send.send(()).unwrap(); + yield_now().await; + + assert_eq!(wake_count, 1); + let (_, output) = assert_ready!(queue.poll_join_next_with_id(&mut cx)) + .unwrap() + .unwrap(); + assert_eq!(output, 42); +} + #[tokio::test] async fn test_join_queue_try_join_next_disabled_coop() { // This number is large enough to trigger coop. Without using `tokio::task::coop::unconstrained`