From 749322d3510a6828d9b118027fddedbc2d174d28 Mon Sep 17 00:00:00 2001 From: Elichai Turkel Date: Tue, 25 Nov 2025 12:31:21 +0200 Subject: [PATCH] task: implement `Extend` for `JoinSet` (#7195) --- tokio/src/task/join_set.rs | 45 ++++++++++++++++++++++++++++++++++++ tokio/tests/task_join_set.rs | 17 ++++++++++++++ 2 files changed, 62 insertions(+) diff --git a/tokio/src/task/join_set.rs b/tokio/src/task/join_set.rs index 544d6a234..f280a7461 100644 --- a/tokio/src/task/join_set.rs +++ b/tokio/src/task/join_set.rs @@ -649,6 +649,51 @@ where } } +/// Extend a [`JoinSet`] with futures from an iterator. +/// +/// This is equivalent to calling [`JoinSet::spawn`] on each element of the iterator. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(not(target_family = "wasm"))] +/// # { +/// use tokio::task::JoinSet; +/// +/// #[tokio::main] +/// async fn main() { +/// let mut set: JoinSet<_> = (0..5).map(|i| async move { i }).collect(); +/// +/// set.extend((5..10).map(|i| async move { i })); +/// +/// let mut seen = [false; 10]; +/// while let Some(res) = set.join_next().await { +/// let idx = res.unwrap(); +/// seen[idx] = true; +/// } +/// +/// for i in 0..10 { +/// assert!(seen[i]); +/// } +/// } +/// # } +/// ``` +impl std::iter::Extend for JoinSet +where + F: Future, + F: Send + 'static, + T: Send + 'static, +{ + fn extend(&mut self, iter: I) + where + I: IntoIterator, + { + iter.into_iter().for_each(|task| { + self.spawn(task); + }); + } +} + // === impl Builder === #[cfg(all(tokio_unstable, feature = "tracing"))] diff --git a/tokio/tests/task_join_set.rs b/tokio/tests/task_join_set.rs index 0c2de0969..a123a5e82 100644 --- a/tokio/tests/task_join_set.rs +++ b/tokio/tests/task_join_set.rs @@ -404,6 +404,23 @@ async fn try_join_next_with_id() { assert_eq!(joined, spawned); } +#[tokio::test] +async fn extend() { + let mut set: JoinSet<_> = (0..5).map(|i| async move { i }).collect(); + + set.extend((5..10).map(|i| async move { i })); + + let mut seen = [false; 10]; + while let Some(res) = set.join_next().await { + let idx = res.unwrap(); + seen[idx] = true; + } + + for s in &seen { + assert!(s); + } +} + mod spawn_local { use super::*;