Skip to main content

karyon_core/async_util/
task_group.rs

1use std::{
2    collections::HashMap,
3    future::Future,
4    sync::{
5        atomic::{AtomicUsize, Ordering},
6        Arc, Weak,
7    },
8};
9
10use parking_lot::Mutex;
11
12use crate::async_runtime::{global_executor, Executor, Task};
13
14use super::{select, CondWait, Either};
15
16/// Identifies a spawned task within a [`TaskGroup`].
17pub type TaskID = usize;
18
19/// TaskGroup A group that contains spawned tasks.
20///
21/// # Example
22///
23/// ```
24///
25/// use std::sync::Arc;
26///
27/// use karyon_core::async_util::{TaskGroup, sleep};
28///
29/// async {
30///     let group = TaskGroup::new();
31///
32///     group.spawn(sleep(std::time::Duration::MAX));
33///
34///     group.cancel().await;
35///
36/// };
37///
38/// ```
39pub struct TaskGroup {
40    inner: Arc<Inner>,
41}
42
43/// Shared state of a [`TaskGroup`]. Held behind an `Arc` so each task can
44/// keep a `Weak` reference back to the group and remove itself on
45/// completion, without forcing callers to wrap the group in an `Arc`.
46struct Inner {
47    tasks: Mutex<HashMap<TaskID, TaskHandler>>,
48    next_id: AtomicUsize,
49    executor: Executor,
50}
51
52impl Inner {
53    /// Removes a task by id, returning its handler if still present.
54    fn remove(&self, id: TaskID) -> Option<TaskHandler> {
55        self.tasks.lock().remove(&id)
56    }
57}
58
59impl TaskGroup {
60    /// Creates a new TaskGroup without providing an executor
61    ///
62    /// This will spawn tasks onto the process-wide multi-threaded
63    /// global executor.
64    pub fn new() -> Self {
65        Self::with_inner(global_executor())
66    }
67
68    /// Creates a new TaskGroup by providing an executor
69    pub fn with_executor(executor: Executor) -> Self {
70        Self::with_inner(executor)
71    }
72
73    fn with_inner(executor: Executor) -> Self {
74        Self {
75            inner: Arc::new(Inner {
76                tasks: Mutex::new(HashMap::new()),
77                next_id: AtomicUsize::new(0),
78                executor,
79            }),
80        }
81    }
82
83    /// Spawns a new task and ignores its result.
84    ///
85    /// Returns the task's [`TaskID`]. The task removes itself from the
86    /// group when it finishes, so the group does not grow without bound.
87    pub fn spawn<T, Fut>(&self, fut: Fut) -> TaskID
88    where
89        T: Send + Sync + 'static,
90        Fut: Future<Output = T> + Send + 'static,
91    {
92        self.spawn_then(fut, |_| async {})
93    }
94
95    /// Spawns a new task and calls the callback after it has completed
96    /// or been canceled. The callback will have the `TaskResult` as a
97    /// parameter, indicating whether the task completed or was canceled.
98    ///
99    /// Returns the task's [`TaskID`].
100    pub fn spawn_then<T, Fut, CallbackF, CallbackFut>(
101        &self,
102        fut: Fut,
103        callback: CallbackF,
104    ) -> TaskID
105    where
106        T: Send + Sync + 'static,
107        Fut: Future<Output = T> + Send + 'static,
108        CallbackF: FnOnce(TaskResult<T>) -> CallbackFut + Send + 'static,
109        CallbackFut: Future<Output = ()> + Send + 'static,
110    {
111        let id = self.inner.next_id.fetch_add(1, Ordering::Relaxed);
112        // Hold the lock across spawn and insert so the task cannot try to
113        // remove itself before it has been inserted.
114        let mut tasks = self.inner.tasks.lock();
115        let task = TaskHandler::new(
116            self.inner.executor.clone(),
117            fut,
118            callback,
119            Arc::downgrade(&self.inner),
120            id,
121        );
122        tasks.insert(id, task);
123        id
124    }
125
126    /// Removes a task by id, returning its handler if still present. The
127    /// caller takes ownership and may cancel it.
128    pub fn remove(&self, id: TaskID) -> Option<TaskHandler> {
129        self.inner.remove(id)
130    }
131
132    /// Checks if the TaskGroup is empty.
133    pub fn is_empty(&self) -> bool {
134        self.inner.tasks.lock().is_empty()
135    }
136
137    /// Get the number of the tasks in the group.
138    pub fn len(&self) -> usize {
139        self.inner.tasks.lock().len()
140    }
141
142    /// Cancels all tasks in the group.
143    pub async fn cancel(&self) {
144        // Take all handlers out, then cancel them without holding the lock.
145        let handlers: Vec<TaskHandler> = self.inner.tasks.lock().drain().map(|(_, h)| h).collect();
146        for handler in handlers {
147            handler.cancel().await;
148        }
149    }
150}
151
152impl Default for TaskGroup {
153    fn default() -> Self {
154        Self::new()
155    }
156}
157
158/// The result of a spawned task.
159#[derive(Debug)]
160pub enum TaskResult<T> {
161    Completed(T),
162    Cancelled,
163}
164
165impl<T: std::fmt::Debug> std::fmt::Display for TaskResult<T> {
166    fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
167        match self {
168            TaskResult::Cancelled => write!(f, "Task cancelled"),
169            TaskResult::Completed(res) => write!(f, "Task completed: {res:?}"),
170        }
171    }
172}
173
174/// TaskHandler
175pub struct TaskHandler {
176    task: Task<()>,
177    /// Per-task stop signal. Signaling it makes the task stop and run its
178    /// callback with `Cancelled`.
179    stop_signal: Arc<CondWait>,
180    /// Set once the task has finished running its callback.
181    cancel_flag: Arc<CondWait>,
182}
183
184impl TaskHandler {
185    /// Creates a new task handler
186    fn new<T, Fut, CallbackF, CallbackFut>(
187        ex: Executor,
188        fut: Fut,
189        callback: CallbackF,
190        group: Weak<Inner>,
191        id: TaskID,
192    ) -> TaskHandler
193    where
194        T: Send + Sync + 'static,
195        Fut: Future<Output = T> + Send + 'static,
196        CallbackF: FnOnce(TaskResult<T>) -> CallbackFut + Send + 'static,
197        CallbackFut: Future<Output = ()> + Send + 'static,
198    {
199        let stop_signal = Arc::new(CondWait::new());
200        let stop_signal_c = stop_signal.clone();
201        let cancel_flag = Arc::new(CondWait::new());
202        let cancel_flag_c = cancel_flag.clone();
203        let task = ex.spawn(async move {
204            // Waits for either the stop signal or the task to complete.
205            let result = select(stop_signal_c.wait(), fut).await;
206
207            let result = match result {
208                Either::Left(_) => TaskResult::Cancelled,
209                Either::Right(res) => TaskResult::Completed(res),
210            };
211
212            // Call the callback
213            callback(result).await;
214
215            cancel_flag_c.signal().await;
216
217            // Remove ourselves from the group. Detach instead of dropping
218            // the handler, so we are not cancelled from within our own
219            // task. If `cancel` already took us out, this is a no-op.
220            if let Some(group) = group.upgrade() {
221                if let Some(handler) = group.remove(id) {
222                    handler.detach();
223                }
224            }
225        });
226
227        TaskHandler {
228            task,
229            stop_signal,
230            cancel_flag,
231        }
232    }
233
234    /// Detaches the task, so dropping the handler does not cancel it.
235    fn detach(self) {
236        self.task.detach();
237    }
238
239    /// Cancels the task: tells it to stop, waits for its callback to run,
240    /// then aborts whatever is left.
241    async fn cancel(self) {
242        self.stop_signal.signal().await;
243        self.cancel_flag.wait().await;
244        self.task.cancel().await;
245    }
246}
247
248#[cfg(test)]
249mod tests {
250    use std::{future, sync::Arc};
251
252    use crate::async_runtime::block_on;
253    use crate::async_util::sleep;
254
255    use super::*;
256
257    #[cfg(feature = "tokio")]
258    #[test]
259    fn test_task_group_with_tokio_executor() {
260        let ex = Arc::new(tokio::runtime::Runtime::new().unwrap());
261        ex.clone().block_on(async move {
262            let group = Arc::new(TaskGroup::with_executor(ex.into()));
263
264            group.spawn_then(future::ready(0), |res| async move {
265                assert!(matches!(res, TaskResult::Completed(0)));
266            });
267
268            group.spawn_then(future::pending::<()>(), |res| async move {
269                assert!(matches!(res, TaskResult::Cancelled));
270            });
271
272            let groupc = group.clone();
273            group.spawn_then(
274                async move {
275                    groupc.spawn_then(future::pending::<()>(), |res| async move {
276                        assert!(matches!(res, TaskResult::Cancelled));
277                    });
278                },
279                |res| async move {
280                    assert!(matches!(res, TaskResult::Completed(_)));
281                },
282            );
283
284            // Do something
285            tokio::time::sleep(std::time::Duration::from_millis(50)).await;
286            group.cancel().await;
287        });
288    }
289
290    #[cfg(feature = "smol")]
291    #[test]
292    fn test_task_group_with_smol_executor() {
293        let ex = Arc::new(smol::Executor::new());
294        smol::block_on(ex.clone().run(async move {
295            let group = Arc::new(TaskGroup::with_executor(ex.into()));
296
297            group.spawn_then(future::ready(0), |res| async move {
298                assert!(matches!(res, TaskResult::Completed(0)));
299            });
300
301            group.spawn_then(future::pending::<()>(), |res| async move {
302                assert!(matches!(res, TaskResult::Cancelled));
303            });
304
305            let groupc = group.clone();
306            group.spawn_then(
307                async move {
308                    groupc.spawn_then(future::pending::<()>(), |res| async move {
309                        assert!(matches!(res, TaskResult::Cancelled));
310                    });
311                },
312                |res| async move {
313                    assert!(matches!(res, TaskResult::Completed(_)));
314                },
315            );
316
317            // Do something
318            smol::Timer::after(std::time::Duration::from_millis(50)).await;
319            group.cancel().await;
320        }));
321    }
322
323    #[test]
324    fn test_task_group() {
325        block_on(async {
326            let group = Arc::new(TaskGroup::new());
327
328            group.spawn_then(future::ready(0), |res| async move {
329                assert!(matches!(res, TaskResult::Completed(0)));
330            });
331
332            group.spawn_then(future::pending::<()>(), |res| async move {
333                assert!(matches!(res, TaskResult::Cancelled));
334            });
335
336            let groupc = group.clone();
337            group.spawn_then(
338                async move {
339                    groupc.spawn_then(future::pending::<()>(), |res| async move {
340                        assert!(matches!(res, TaskResult::Cancelled));
341                    });
342                },
343                |res| async move {
344                    assert!(matches!(res, TaskResult::Completed(_)));
345                },
346            );
347
348            // Do something
349            sleep(std::time::Duration::from_millis(50)).await;
350            group.cancel().await;
351        });
352    }
353
354    #[test]
355    fn test_task_group_removes_finished_tasks() {
356        block_on(async {
357            let group = Arc::new(TaskGroup::new());
358
359            // A finished task removes itself; a pending one stays.
360            group.spawn(future::ready(0));
361            group.spawn(future::pending::<()>());
362
363            sleep(std::time::Duration::from_millis(50)).await;
364            assert_eq!(group.len(), 1);
365
366            group.cancel().await;
367            assert!(group.is_empty());
368        });
369    }
370
371    #[test]
372    fn test_task_group_remove_by_id() {
373        block_on(async {
374            let group = Arc::new(TaskGroup::new());
375
376            let id = group.spawn(future::pending::<()>());
377            assert_eq!(group.len(), 1);
378
379            let handler = group.remove(id).expect("task is present");
380            assert!(group.is_empty());
381            handler.cancel().await;
382        });
383    }
384}