Skip to main content

reth_tasks/
lib.rs

1//! Reth task management.
2//!
3//! # Feature Flags
4//!
5//! - `rayon`: Enable rayon thread pool for blocking tasks.
6
7#![doc(
8    html_logo_url = "https://raw.githubusercontent.com/paradigmxyz/reth/main/assets/reth-docs.png",
9    html_favicon_url = "https://avatars0.githubusercontent.com/u/97369466?s=256",
10    issue_tracker_base_url = "https://github.com/paradigmxyz/reth/issues/"
11)]
12#![cfg_attr(not(test), warn(unused_crate_dependencies))]
13#![cfg_attr(docsrs, feature(doc_cfg))]
14
15use crate::shutdown::{signal, Shutdown, Signal};
16use std::{
17    any::Any,
18    fmt::{Display, Formatter},
19    pin::Pin,
20    sync::{
21        atomic::{AtomicUsize, Ordering},
22        Arc,
23    },
24    task::{ready, Context, Poll},
25    thread,
26};
27use tokio::{
28    runtime::Handle,
29    sync::mpsc::{unbounded_channel, UnboundedReceiver, UnboundedSender},
30};
31use tracing::debug;
32
33pub mod cancel;
34pub mod lazy;
35pub mod metrics;
36pub mod runtime;
37pub mod shutdown;
38pub mod utils;
39pub(crate) mod worker_map;
40
41#[cfg(feature = "rayon")]
42pub mod pool;
43#[cfg(feature = "rayon")]
44pub use pool::{build_pool_with_panic_handler, Worker, WorkerPool};
45
46/// Lock-free ordered parallel iterator extension trait.
47#[cfg(feature = "rayon")]
48pub mod for_each_ordered;
49#[cfg(feature = "rayon")]
50pub use for_each_ordered::ForEachOrdered;
51
52pub use cancel::{CancelOnDrop, ManualCancel};
53pub use lazy::LazyHandle;
54#[cfg(feature = "rayon")]
55pub use runtime::RayonConfig;
56pub use runtime::{Runtime, RuntimeBuildError, RuntimeBuilder, RuntimeConfig, TokioConfig};
57
58/// A [`TaskExecutor`] is now an alias for [`Runtime`].
59pub type TaskExecutor = Runtime;
60
61/// Spawns an OS thread with the current tokio runtime context propagated.
62///
63/// This function captures the current tokio runtime handle (if available) and enters it
64/// in the newly spawned thread. This ensures that code running in the spawned thread can
65/// use [`Handle::current()`], [`Handle::spawn_blocking()`], and other tokio utilities that
66/// require a runtime context.
67#[track_caller]
68pub fn spawn_os_thread<F, T>(name: &str, f: F) -> thread::JoinHandle<T>
69where
70    F: FnOnce() -> T + Send + 'static,
71    T: Send + 'static,
72{
73    let handle = Handle::try_current().ok();
74    thread::Builder::new()
75        .name(name.to_string())
76        .spawn(move || {
77            let _guard = handle.as_ref().map(Handle::enter);
78            f()
79        })
80        .unwrap_or_else(|e| panic!("failed to spawn thread {name:?}: {e}"))
81}
82
83/// Spawns a scoped OS thread with the current tokio runtime context propagated.
84///
85/// This is the scoped thread version of [`spawn_os_thread`], for use with [`std::thread::scope`].
86#[track_caller]
87pub fn spawn_scoped_os_thread<'scope, 'env, F, T>(
88    scope: &'scope thread::Scope<'scope, 'env>,
89    name: &str,
90    f: F,
91) -> thread::ScopedJoinHandle<'scope, T>
92where
93    F: FnOnce() -> T + Send + 'scope,
94    T: Send + 'scope,
95{
96    let handle = Handle::try_current().ok();
97    thread::Builder::new()
98        .name(name.to_string())
99        .spawn_scoped(scope, move || {
100            let _guard = handle.as_ref().map(Handle::enter);
101            f()
102        })
103        .unwrap_or_else(|e| panic!("failed to spawn scoped thread {name:?}: {e}"))
104}
105
106/// Monitors critical tasks for panics and manages graceful shutdown.
107///
108/// The main purpose of this type is to be able to monitor if a critical task panicked, for
109/// diagnostic purposes, since tokio tasks essentially fail silently. Therefore, this type is a
110/// Future that resolves with the name of the panicked task. See [`Runtime::spawn_critical_task`].
111///
112/// Automatically spawned as a background task when building a [`Runtime`]. Use
113/// [`Runtime::take_task_manager_handle`] to extract the join handle if you need to poll for
114/// panic errors directly.
115#[derive(Debug)]
116#[must_use = "TaskManager must be polled to monitor critical tasks"]
117pub struct TaskManager {
118    /// Receiver for task events.
119    task_events_rx: UnboundedReceiver<TaskEvent>,
120    /// The [Signal] to fire when all tasks should be shutdown.
121    ///
122    /// This is fired when dropped.
123    signal: Option<Signal>,
124    /// How many [`GracefulShutdown`](crate::shutdown::GracefulShutdown) tasks are currently
125    /// active.
126    graceful_tasks: Arc<AtomicUsize>,
127}
128
129// === impl TaskManager ===
130
131impl TaskManager {
132    /// Create a new [`TaskManager`] without an associated [`Runtime`], returning
133    /// the shutdown/event primitives for [`RuntimeBuilder`] to wire up.
134    pub(crate) fn new_parts(
135        _handle: Handle,
136    ) -> (Self, Shutdown, UnboundedSender<TaskEvent>, Arc<AtomicUsize>) {
137        let (task_events_tx, task_events_rx) = unbounded_channel();
138        let (signal, on_shutdown) = signal();
139        let graceful_tasks = Arc::new(AtomicUsize::new(0));
140        let manager = Self {
141            task_events_rx,
142            signal: Some(signal),
143            graceful_tasks: Arc::clone(&graceful_tasks),
144        };
145        (manager, on_shutdown, task_events_tx, graceful_tasks)
146    }
147
148    /// Fires the shutdown signal and awaits until all tasks are shutdown.
149    pub fn graceful_shutdown(self) {
150        let _ = self.do_graceful_shutdown(None);
151    }
152
153    /// Fires the shutdown signal and awaits until all tasks are shutdown.
154    ///
155    /// Returns true if all tasks were shutdown before the timeout elapsed.
156    pub fn graceful_shutdown_with_timeout(self, timeout: std::time::Duration) -> bool {
157        self.do_graceful_shutdown(Some(timeout))
158    }
159
160    fn do_graceful_shutdown(self, timeout: Option<std::time::Duration>) -> bool {
161        drop(self.signal);
162        let deadline = timeout.map(|t| std::time::Instant::now() + t);
163        while self.graceful_tasks.load(Ordering::SeqCst) > 0 {
164            if deadline.is_some_and(|d| std::time::Instant::now() > d) {
165                debug!("graceful shutdown timed out");
166                return false;
167            }
168            thread::yield_now();
169        }
170        debug!("gracefully shut down");
171        true
172    }
173}
174
175/// An endless future that resolves if a critical task panicked.
176///
177/// See [`Runtime::spawn_critical_task`]
178impl std::future::Future for TaskManager {
179    type Output = Result<(), PanickedTaskError>;
180
181    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
182        match ready!(self.as_mut().get_mut().task_events_rx.poll_recv(cx)) {
183            Some(TaskEvent::Panic(err)) => Poll::Ready(Err(err)),
184            Some(TaskEvent::GracefulShutdown) | None => {
185                if let Some(signal) = self.get_mut().signal.take() {
186                    signal.fire();
187                }
188                Poll::Ready(Ok(()))
189            }
190        }
191    }
192}
193
194/// Error with the name of the task that panicked and an error downcasted to string, if possible.
195#[derive(Debug, thiserror::Error, PartialEq, Eq)]
196pub struct PanickedTaskError {
197    task_name: &'static str,
198    error: Option<String>,
199}
200
201impl Display for PanickedTaskError {
202    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
203        let task_name = self.task_name;
204        if let Some(error) = &self.error {
205            write!(f, "Critical task `{task_name}` panicked: `{error}`")
206        } else {
207            write!(f, "Critical task `{task_name}` panicked")
208        }
209    }
210}
211
212impl PanickedTaskError {
213    pub(crate) fn new(task_name: &'static str, error: Box<dyn Any>) -> Self {
214        let error = match error.downcast::<String>() {
215            Ok(value) => Some(*value),
216            Err(error) => match error.downcast::<&str>() {
217                Ok(value) => Some(value.to_string()),
218                Err(_) => None,
219            },
220        };
221
222        Self { task_name, error }
223    }
224}
225
226/// Represents the events that the `TaskManager`'s main future can receive.
227#[derive(Debug)]
228pub(crate) enum TaskEvent {
229    /// Indicates that a critical task has panicked.
230    Panic(PanickedTaskError),
231    /// A signal requesting a graceful shutdown of the `TaskManager`.
232    GracefulShutdown,
233}
234
235#[cfg(test)]
236mod tests {
237    use super::*;
238    use std::{
239        sync::atomic::{AtomicBool, AtomicUsize, Ordering},
240        time::Duration,
241    };
242
243    #[test]
244    fn test_critical() {
245        let rt = Runtime::test();
246        let handle = rt.take_task_manager_handle().unwrap();
247
248        rt.spawn_critical_task("this is a critical task", async { panic!("intentionally panic") });
249
250        rt.handle().block_on(async move {
251            let err_result = handle.await.unwrap();
252            assert!(err_result.is_err(), "Expected TaskManager to return an error due to panic");
253            let panicked_err = err_result.unwrap_err();
254
255            assert_eq!(panicked_err.task_name, "this is a critical task");
256            assert_eq!(panicked_err.error, Some("intentionally panic".to_string()));
257        })
258    }
259
260    #[test]
261    fn test_manager_shutdown_critical() {
262        let rt = Runtime::test();
263
264        let (signal, shutdown) = signal();
265
266        rt.spawn_critical_task("this is a critical task", async move {
267            tokio::time::sleep(Duration::from_millis(200)).await;
268            drop(signal);
269        });
270
271        rt.graceful_shutdown();
272
273        rt.handle().block_on(shutdown);
274    }
275
276    #[test]
277    fn test_manager_shutdown() {
278        let rt = Runtime::test();
279
280        let (signal, shutdown) = signal();
281
282        rt.spawn_task(async move {
283            tokio::time::sleep(Duration::from_millis(200)).await;
284            drop(signal);
285        });
286
287        rt.graceful_shutdown();
288
289        rt.handle().block_on(shutdown);
290    }
291
292    #[test]
293    fn test_manager_graceful_shutdown() {
294        let rt = Runtime::test();
295
296        let val = Arc::new(AtomicBool::new(false));
297        let c = val.clone();
298        rt.spawn_critical_with_graceful_shutdown_signal("grace", async move |shutdown| {
299            let _guard = shutdown.await;
300            tokio::time::sleep(Duration::from_millis(200)).await;
301            c.store(true, Ordering::Relaxed);
302        });
303
304        rt.graceful_shutdown();
305        assert!(val.load(Ordering::Relaxed));
306    }
307
308    #[test]
309    fn test_manager_graceful_shutdown_many() {
310        let rt = Runtime::test();
311
312        let counter = Arc::new(AtomicUsize::new(0));
313        let num = 10;
314        for _ in 0..num {
315            let c = counter.clone();
316            rt.spawn_critical_with_graceful_shutdown_signal("grace", async move |shutdown| {
317                let _guard = shutdown.await;
318                tokio::time::sleep(Duration::from_millis(200)).await;
319                c.fetch_add(1, Ordering::SeqCst);
320            });
321        }
322
323        rt.graceful_shutdown();
324        assert_eq!(counter.load(Ordering::Relaxed), num);
325    }
326
327    #[test]
328    fn test_manager_graceful_shutdown_timeout() {
329        let rt = Runtime::test();
330
331        let timeout = Duration::from_millis(500);
332        let val = Arc::new(AtomicBool::new(false));
333        let val2 = val.clone();
334        rt.spawn_critical_with_graceful_shutdown_signal("grace", async move |shutdown| {
335            let _guard = shutdown.await;
336            tokio::time::sleep(timeout * 3).await;
337            val2.store(true, Ordering::Relaxed);
338            unreachable!("should not be reached");
339        });
340
341        rt.graceful_shutdown_with_timeout(timeout);
342        assert!(!val.load(Ordering::Relaxed));
343    }
344
345    #[test]
346    fn can_build_runtime() {
347        let rt = Runtime::test();
348        let _handle = rt.handle();
349    }
350
351    #[test]
352    fn test_graceful_shutdown_triggered_by_executor() {
353        let rt = Runtime::test();
354        let task_manager_handle = rt.take_task_manager_handle().unwrap();
355
356        let task_did_shutdown_flag = Arc::new(AtomicBool::new(false));
357        let flag_clone = task_did_shutdown_flag.clone();
358
359        let spawned_task_handle = rt.spawn_with_signal(async move |shutdown_signal| {
360            shutdown_signal.await;
361            flag_clone.store(true, Ordering::SeqCst);
362        });
363
364        let send_result = rt.initiate_graceful_shutdown();
365        assert!(send_result.is_ok());
366
367        let manager_final_result = rt.handle().block_on(task_manager_handle);
368        assert!(manager_final_result.is_ok(), "TaskManager task should not panic");
369        assert_eq!(manager_final_result.unwrap(), Ok(()));
370
371        let task_join_result = rt.handle().block_on(spawned_task_handle);
372        assert!(task_join_result.is_ok());
373
374        assert!(task_did_shutdown_flag.load(Ordering::Relaxed));
375    }
376}