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
16pub type TaskID = usize;
18
19pub struct TaskGroup {
40 inner: Arc<Inner>,
41}
42
43struct Inner {
47 tasks: Mutex<HashMap<TaskID, TaskHandler>>,
48 next_id: AtomicUsize,
49 executor: Executor,
50}
51
52impl Inner {
53 fn remove(&self, id: TaskID) -> Option<TaskHandler> {
55 self.tasks.lock().remove(&id)
56 }
57}
58
59impl TaskGroup {
60 pub fn new() -> Self {
65 Self::with_inner(global_executor())
66 }
67
68 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 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 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 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 pub fn remove(&self, id: TaskID) -> Option<TaskHandler> {
129 self.inner.remove(id)
130 }
131
132 pub fn is_empty(&self) -> bool {
134 self.inner.tasks.lock().is_empty()
135 }
136
137 pub fn len(&self) -> usize {
139 self.inner.tasks.lock().len()
140 }
141
142 pub async fn cancel(&self) {
144 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#[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
174pub struct TaskHandler {
176 task: Task<()>,
177 stop_signal: Arc<CondWait>,
180 cancel_flag: Arc<CondWait>,
182}
183
184impl TaskHandler {
185 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 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 callback(result).await;
214
215 cancel_flag_c.signal().await;
216
217 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 fn detach(self) {
236 self.task.detach();
237 }
238
239 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 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 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 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 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}