miden_node_utils/
tasks.rs1use std::collections::HashMap;
2use std::future::Future;
3
4use anyhow::Context;
5use miden_node_tracing::warn;
6use tokio::task::{Id, JoinError, JoinSet};
7
8use crate::shutdown::CancellationToken;
9
10pub struct Tasks<T = anyhow::Result<()>, M = String> {
14 handles: JoinSet<T>,
15 tags: HashMap<Id, M>,
16}
17
18impl<T, M> Default for Tasks<T, M> {
19 fn default() -> Self {
20 Self {
21 handles: JoinSet::new(),
22 tags: HashMap::new(),
23 }
24 }
25}
26
27impl<T: Send + 'static, M> Tasks<T, M> {
28 pub fn new() -> Self {
30 Self::default()
31 }
32
33 pub fn spawn(
35 &mut self,
36 tag: impl Into<M>,
37 task: impl Future<Output = T> + Send + 'static,
38 ) -> Id {
39 let id = self.handles.spawn(task).id();
40 self.tags.insert(id, tag.into());
41 id
42 }
43
44 pub async fn join_next(&mut self) -> Option<(M, Result<T, JoinError>)> {
46 let result = self.handles.join_next_with_id().await?;
47 let id = match &result {
48 Ok((id, _)) => *id,
49 Err(err) => err.id(),
50 };
51 let tag = self.tags.remove(&id).expect("a spawned task has a tag");
53 let result = result.map(|(_, output)| output);
54
55 Some((tag, result))
56 }
57
58 pub fn is_empty(&self) -> bool {
60 self.handles.is_empty()
61 }
62
63 pub fn len(&self) -> usize {
65 self.handles.len()
66 }
67
68 pub fn tags(&self) -> impl Iterator<Item = &M> {
70 self.tags.values()
71 }
72
73 pub async fn shutdown(&mut self) {
75 self.handles.shutdown().await;
76 self.tags.clear();
77 }
78}
79
80impl Tasks {
81 pub fn spawn_infallible(
83 &mut self,
84 name: impl Into<String>,
85 task: impl Future<Output = ()> + Send + 'static,
86 ) -> Id {
87 self.spawn(name, async move {
88 task.await;
89 Ok(())
90 })
91 }
92
93 pub async fn join_next_as_error(&mut self) -> anyhow::Result<()> {
97 let Some((task, result)) = self.join_next().await else {
98 anyhow::bail!("task set is empty");
99 };
100
101 Self::unexpected_completion(&task, result)
102 }
103
104 pub async fn join_next_or_cancelled(&mut self, token: CancellationToken) -> anyhow::Result<()> {
117 let mut outcome = Ok(());
118 while !token.is_cancelled() {
119 tokio::select! {
120 biased;
121 () = token.cancelled() => break,
122 result = self.join_next() => {
123 let Some((task, result)) = result else {
124 anyhow::bail!("task set is empty");
125 };
126 outcome = Self::unexpected_completion(&task, result);
127 token.cancel();
129 },
130 }
131 }
132
133 while let Some((task, result)) = self.join_next().await {
134 match (&outcome, Self::shutdown_completion(&task, result)) {
135 (Ok(()), result) => outcome = result,
137 (Err(_), Err(err)) => {
140 warn!(&err, "task failed during shutdown", task.name = task);
141 },
142 (Err(_), Ok(())) => {},
144 }
145 }
146
147 outcome
148 }
149
150 fn unexpected_completion(
156 task: &str,
157 result: Result<anyhow::Result<()>, JoinError>,
158 ) -> anyhow::Result<()> {
159 match result {
160 Ok(Ok(())) => anyhow::bail!("task {task} completed unexpectedly"),
161 Ok(Err(err)) => Err(err).with_context(|| format!("task {task} failed")),
162 Err(err) => Err(err).with_context(|| format!("task {task} failed to join")),
163 }
164 }
165
166 fn shutdown_completion(
172 task: &str,
173 result: Result<anyhow::Result<()>, JoinError>,
174 ) -> anyhow::Result<()> {
175 match result {
176 Ok(Ok(())) => Ok(()),
177 Ok(Err(err)) => Err(err).with_context(|| format!("task {task} failed during shutdown")),
178 Err(err) if err.is_cancelled() => Ok(()),
179 Err(err) => Err(err).with_context(|| format!("task {task} failed to join")),
180 }
181 }
182}
183
184#[cfg(test)]
185mod tests {
186 use std::time::Duration;
187
188 use super::*;
189
190 #[tokio::test]
191 async fn join_next_or_cancelled_accepts_clean_task_completion_after_cancellation() {
192 let token = crate::shutdown::CancellationToken::new();
193 let mut tasks = Tasks::new();
194 tasks.spawn("worker", {
195 let token = token.clone();
196 async move {
197 token.cancelled().await;
198 Ok(())
199 }
200 });
201
202 token.cancel();
203
204 tasks
205 .join_next_or_cancelled(token)
206 .await
207 .expect("clean shutdown should not be treated as an error");
208 }
209
210 #[tokio::test]
211 async fn join_next_or_cancelled_treats_task_completion_before_cancellation_as_error() {
212 let token = crate::shutdown::CancellationToken::new();
213 let mut tasks = Tasks::new();
214 tasks.spawn("worker", async { Ok(()) });
215
216 let err = tasks
217 .join_next_or_cancelled(token)
218 .await
219 .expect_err("unexpected task completion should fail before shutdown");
220
221 assert_eq!(err.to_string(), "task worker completed unexpectedly");
222 }
223
224 #[tokio::test]
225 async fn join_next_or_cancelled_drains_remaining_tasks_after_a_failure() {
226 use std::sync::Arc;
227 use std::sync::atomic::{AtomicBool, Ordering};
228
229 let token = crate::shutdown::CancellationToken::new();
230 let mut tasks = Tasks::new();
231 let survivor_finished = Arc::new(AtomicBool::new(false));
232
233 tasks.spawn("failing", async { anyhow::bail!("boom") });
234 tasks.spawn("survivor", {
235 let token = token.clone();
236 let finished = Arc::clone(&survivor_finished);
237 async move {
238 token.cancelled().await;
239 tokio::time::sleep(Duration::from_millis(10)).await;
241 finished.store(true, Ordering::Relaxed);
242 Ok(())
243 }
244 });
245
246 let err = tasks
247 .join_next_or_cancelled(token.clone())
248 .await
249 .expect_err("the failing task's error should be returned");
250
251 assert_eq!(err.to_string(), "task failing failed");
252 assert!(token.is_cancelled(), "a task failure should trigger shutdown");
253 assert!(tasks.is_empty(), "all tasks should be drained before returning");
254 assert!(
255 survivor_finished.load(Ordering::Relaxed),
256 "surviving tasks should shut down gracefully, not be aborted",
257 );
258 }
259
260 #[tokio::test]
261 async fn join_next_or_cancelled_drains_past_failures_during_shutdown() {
262 use std::sync::Arc;
263 use std::sync::atomic::{AtomicBool, Ordering};
264
265 let token = crate::shutdown::CancellationToken::new();
266 let mut tasks = Tasks::new();
267 let survivor_finished = Arc::new(AtomicBool::new(false));
268
269 tasks.spawn("failing", {
270 let token = token.clone();
271 async move {
272 token.cancelled().await;
273 anyhow::bail!("boom")
274 }
275 });
276 tasks.spawn("survivor", {
277 let token = token.clone();
278 let finished = Arc::clone(&survivor_finished);
279 async move {
280 token.cancelled().await;
281 tokio::time::sleep(Duration::from_millis(10)).await;
282 finished.store(true, Ordering::Relaxed);
283 Ok(())
284 }
285 });
286
287 token.cancel();
288
289 let err = tasks
290 .join_next_or_cancelled(token)
291 .await
292 .expect_err("a failure during shutdown should be reported");
293
294 assert_eq!(err.to_string(), "task failing failed during shutdown");
295 assert!(tasks.is_empty(), "draining should continue past the failed task");
296 assert!(
297 survivor_finished.load(Ordering::Relaxed),
298 "surviving tasks should shut down gracefully, not be aborted",
299 );
300 }
301
302 #[tokio::test]
303 async fn join_next_or_cancelled_waits_for_all_tasks_to_complete_after_cancellation() {
304 let token = crate::shutdown::CancellationToken::new();
305 let mut tasks = Tasks::new();
306 tasks.spawn("worker-a", {
307 let token = token.clone();
308 async move {
309 token.cancelled().await;
310 Ok(())
311 }
312 });
313 tasks.spawn("worker-b", {
314 let token = token.clone();
315 async move {
316 token.cancelled().await;
317 tokio::time::sleep(Duration::from_millis(10)).await;
318 Ok(())
319 }
320 });
321
322 token.cancel();
323
324 tasks
325 .join_next_or_cancelled(token)
326 .await
327 .expect("shutdown should wait for all clean task exits");
328 assert!(tasks.is_empty());
329 }
330}