Skip to main content

miden_node_utils/
tasks.rs

1use 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
10/// A tagged task set for supervising concurrently-running Tokio tasks.
11///
12/// Dropping a task set aborts all tasks that are still running.
13pub 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    /// Creates an empty task set.
29    pub fn new() -> Self {
30        Self::default()
31    }
32
33    /// Spawns a tagged task into the set.
34    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    /// Waits for the next task to complete and returns it with its tag.
45    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        // Every task is tagged when it is spawned, and a task is joined once.
52        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    /// Returns `true` if no tasks are currently in the set.
59    pub fn is_empty(&self) -> bool {
60        self.handles.is_empty()
61    }
62
63    /// Returns the number of tasks currently in the set.
64    pub fn len(&self) -> usize {
65        self.handles.len()
66    }
67
68    /// Returns the tag of every task in the set.
69    pub fn tags(&self) -> impl Iterator<Item = &M> {
70        self.tags.values()
71    }
72
73    /// Aborts every task in the set and waits for the tasks to finish.
74    pub async fn shutdown(&mut self) {
75        self.handles.shutdown().await;
76        self.tags.clear();
77    }
78}
79
80impl Tasks {
81    /// Spawns a named task that does not return an error.
82    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    /// Waits for the next task to complete, treating that completion as an error.
94    ///
95    /// This is intended for supervised task sets where every task is expected to run indefinitely.
96    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    /// Waits for either an unexpected task completion or a shutdown request.
105    ///
106    /// Before shutdown, any task completion is treated as fatal because this type supervises
107    /// long-running tasks. Such a completion triggers the shutdown itself: the token is cancelled
108    /// and the remaining tasks are drained before the error is returned. Returning without
109    /// draining would drop the set and abort the surviving tasks mid-work — e.g. the store's
110    /// block writer between its database commit and tree update, tearing persistent state.
111    ///
112    /// Once `token` is cancelled (whether externally or by a failure here), clean task exits are
113    /// accepted and this method waits for all tracked tasks to finish. The first failure observed
114    /// is returned as the root cause; subsequent failures are logged, since they are often
115    /// knock-on effects of the first.
116    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                    // Shut the remaining tasks down and fall through to the drain below.
128                    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                // No failure so far: this task's result (clean or failed) becomes the outcome.
136                (Ok(()), result) => outcome = result,
137                // A failure is already recorded as the root cause; later failures are often
138                // knock-on effects of it, so log them rather than mask it.
139                (Err(_), Err(err)) => {
140                    warn!(&err, "task failed during shutdown", task.name = task);
141                },
142                // A failure is already recorded and this task exited cleanly: nothing to add.
143                (Err(_), Ok(())) => {},
144            }
145        }
146
147        outcome
148    }
149
150    /// Interprets a task completion observed *before* shutdown was requested.
151    ///
152    /// Supervised tasks are expected to run until shutdown, so every completion — even a clean
153    /// exit — is an error here; the variants only differ in how much context the error carries
154    /// (task failure, or a panicked/aborted task surfacing as a [`JoinError`]).
155    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    /// Interprets a task completion observed *after* shutdown was requested.
167    ///
168    /// During shutdown a clean exit is the expected outcome, and a cancelled task is also fine —
169    /// abort is how a dropped set winds tasks down. A task error or a panic (a non-cancellation
170    /// [`JoinError`]) is still a failure worth reporting.
171    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                // Work past the cancellation point: an aborted task would never get here.
240                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}