1use std::{
4 future::Future,
5 sync::{
6 atomic::{AtomicBool, Ordering},
7 Arc,
8 },
9};
10
11use futures::FutureExt;
12use tokio::task::JoinHandle;
13use tokio_util::{sync::CancellationToken, task::TaskTracker};
14use tracing::{debug, error, info};
15
16pub type FatalError = tower::BoxError;
18
19#[must_use]
21pub enum ShutdownReason {
22 Requested,
24
25 TaskFailed,
27}
28
29#[derive(Clone, Default)]
31pub struct TaskExecutor {
32 token: CancellationToken,
33 tracker: TaskTracker,
34 failed: Arc<AtomicBool>,
35}
36
37impl TaskExecutor {
38 pub fn new() -> Self {
40 Self::default()
41 }
42
43 pub fn spawn<F>(&self, future: F) -> JoinHandle<F::Output>
45 where
46 F: Future + Send + 'static,
47 F::Output: Send + 'static,
48 {
49 self.tracker.spawn(future)
50 }
51
52 pub fn spawn_critical<F>(&self, name: &'static str, future: F) -> JoinHandle<()>
54 where
55 F: Future<Output = Result<(), FatalError>> + Send + 'static,
56 {
57 let executor = self.clone();
58 self.tracker.spawn(future.map(move |result| {
59 if executor.token.is_cancelled() {
60 if let Err(e) = result {
62 debug!(subsystem = name, "{:#}", anyhow::Error::from_boxed(e));
63 }
64 return;
65 }
66 match result {
67 Ok(()) => error!(
68 subsystem = name,
69 "critical task exited before shutdown was requested"
70 ),
71 Err(e) => {
72 error!(subsystem = name, "{:#}", anyhow::Error::from_boxed(e));
73 }
74 }
75
76 executor.failed.store(true, Ordering::Relaxed);
77 executor.trigger_shutdown();
78 }))
79 }
80
81 pub fn cancellation_token(&self) -> CancellationToken {
83 self.token.clone()
84 }
85
86 pub fn trigger_shutdown(&self) {
88 if !self.token.is_cancelled() {
89 info!("Shutting down...");
90 }
91 self.token.cancel();
92 }
93
94 pub async fn wait_for_shutdown(&self) -> ShutdownReason {
96 self.token.cancelled().await;
97 self.tracker.close();
98 self.tracker.wait().await;
99 if self.failed.load(Ordering::Relaxed) {
100 ShutdownReason::TaskFailed
101 } else {
102 ShutdownReason::Requested
103 }
104 }
105}