From d5ad1eba7a3bd13f2c0049001fa35c30baf9ceb8 Mon Sep 17 00:00:00 2001 From: noah Date: Sun, 24 Apr 2022 14:58:20 -0500 Subject: [PATCH] runtime: add on_task_poll_start hook This hook will allow users to have a callback run before tasks are polled. Eventually this should be accompanied on_task_poll_stop, and maybe a few other callbacks. This was implemented in order to allow tokio-uring to run a periodic maintenance loop. --- tokio/src/runtime/basic_scheduler.rs | 11 ++++ tokio/src/runtime/builder.rs | 69 ++++++++++++++++++++++++- tokio/src/runtime/thread_pool/mod.rs | 11 +++- tokio/src/runtime/thread_pool/worker.rs | 15 ++++++ 4 files changed, 103 insertions(+), 3 deletions(-) diff --git a/tokio/src/runtime/basic_scheduler.rs b/tokio/src/runtime/basic_scheduler.rs index acebd0ab4..94df2256a 100644 --- a/tokio/src/runtime/basic_scheduler.rs +++ b/tokio/src/runtime/basic_scheduler.rs @@ -87,6 +87,9 @@ struct Shared { /// Callback for a worker unparking itself after_unpark: Option, + /// Callback for a task being polled + before_task_poll: Option, + /// Keeps track of various runtime metrics. scheduler_metrics: SchedulerMetrics, @@ -125,6 +128,7 @@ impl BasicScheduler { handle_inner: HandleInner, before_park: Option, after_unpark: Option, + before_task_enter: Option, ) -> BasicScheduler { let unpark = driver.unpark(); @@ -137,6 +141,7 @@ impl BasicScheduler { handle_inner, before_park, after_unpark, + before_task_poll: before_task_enter, scheduler_metrics: SchedulerMetrics::new(), worker_metrics: WorkerMetrics::new(), }), @@ -293,6 +298,12 @@ impl Context { /// thread-local context. fn run_task(&self, mut core: Box, f: impl FnOnce() -> R) -> (Box, R) { core.metrics.incr_poll_count(); + + // run before polling the task + if let Some(f) = &self.spawner.shared.before_task_poll { + f() + } + self.enter(core, || crate::coop::budget(f)) } diff --git a/tokio/src/runtime/builder.rs b/tokio/src/runtime/builder.rs index 618474c05..a470bacfa 100644 --- a/tokio/src/runtime/builder.rs +++ b/tokio/src/runtime/builder.rs @@ -76,6 +76,9 @@ pub struct Builder { /// To run after each thread is unparked. pub(super) after_unpark: Option, + /// To run before each task is entered. + pub(super) before_task_poll: Option, + /// Customizable keep alive timeout for BlockingPool pub(super) keep_alive: Option, } @@ -144,6 +147,9 @@ impl Builder { before_park: None, after_unpark: None, + // No task callbacks + before_task_poll: None, + keep_alive: None, } } @@ -496,6 +502,59 @@ impl Builder { self } + /// Executes a function `f` just before a task is polled. + /// + /// This is intended for bookkeeping and monitoring use cases; note that work + /// in this callback will increase latencies when the runtime polls tasks. + /// + /// Note: There can only be one pre-poll callback for a runtime; calling this function + /// more than once replaces the last callback defined, rather than adding to it. + /// + /// # Examples + /// ``` + /// # use tokio::runtime; + /// # use std::sync::{Arc, atomic::{AtomicBool, Ordering}}; + /// + /// // test that our task was entered + /// let state = Arc::new(AtomicBool::new(false)); + /// let callback_state = state.clone(); + /// + /// let runtime = runtime::Builder::new_multi_thread() + /// .on_task_poll_start(move || callback_state.store(true, Ordering::SeqCst)) + /// .build(); + /// + /// runtime.unwrap().block_on(async { + /// let state = state.clone(); + /// tokio::spawn(async move { assert!(state.load(Ordering::SeqCst)) }).await.unwrap() + /// }); + /// ``` + /// + /// ``` + /// # use tokio::runtime; + /// # use std::sync::{Arc, atomic::{AtomicBool, Ordering}}; + /// + /// // test that our task was entered + /// let state = Arc::new(AtomicBool::new(false)); + /// let callback_state = state.clone(); + /// + /// let runtime = runtime::Builder::new_current_thread() + /// .on_task_poll_start(move || callback_state.store(true, Ordering::SeqCst)) + /// .build(); + /// + /// runtime.unwrap().block_on(async { + /// let state = state.clone(); + /// tokio::spawn(async move { assert!(state.load(Ordering::SeqCst)) }).await.unwrap() + /// }); + /// ``` + #[cfg(not(loom))] + pub fn on_task_poll_start(&mut self, f: F) -> &mut Self + where + F: Fn() + Send + Sync + 'static, + { + self.before_task_poll = Some(std::sync::Arc::new(f)); + self + } + /// Creates the configured `Runtime`. /// /// The returned `Runtime` instance is ready to spawn tasks. @@ -580,6 +639,7 @@ impl Builder { handle_inner, self.before_park.clone(), self.after_unpark.clone(), + self.before_task_poll.clone(), ); let spawner = Spawner::Basic(scheduler.spawner().clone()); @@ -685,7 +745,14 @@ cfg_rt_multi_thread! { blocking_spawner, }; - let (scheduler, launch) = ThreadPool::new(core_threads, driver, handle_inner, self.before_park.clone(), self.after_unpark.clone()); + let (scheduler, launch) = ThreadPool::new( + core_threads, + driver, + handle_inner, + self.before_park.clone(), + self.after_unpark.clone(), + self.before_task_poll.clone(), + ); let spawner = Spawner::ThreadPool(scheduler.spawner().clone()); // Create the runtime handle diff --git a/tokio/src/runtime/thread_pool/mod.rs b/tokio/src/runtime/thread_pool/mod.rs index 76346c686..f9161ed9c 100644 --- a/tokio/src/runtime/thread_pool/mod.rs +++ b/tokio/src/runtime/thread_pool/mod.rs @@ -51,10 +51,17 @@ impl ThreadPool { handle_inner: HandleInner, before_park: Option, after_unpark: Option, + before_task_poll: Option, ) -> (ThreadPool, Launch) { let parker = Parker::new(driver); - let (shared, launch) = - worker::create(size, parker, handle_inner, before_park, after_unpark); + let (shared, launch) = worker::create( + size, + parker, + handle_inner, + before_park, + after_unpark, + before_task_poll, + ); let spawner = Spawner { shared }; let thread_pool = ThreadPool { spawner }; diff --git a/tokio/src/runtime/thread_pool/worker.rs b/tokio/src/runtime/thread_pool/worker.rs index 9b456570d..1b559e045 100644 --- a/tokio/src/runtime/thread_pool/worker.rs +++ b/tokio/src/runtime/thread_pool/worker.rs @@ -151,6 +151,9 @@ pub(super) struct Shared { /// Callback for a worker unparking itself after_unpark: Option, + /// Callback for a task being polled + before_task_poll: Option, + /// Collects metrics from the runtime. pub(super) scheduler_metrics: SchedulerMetrics, @@ -198,6 +201,7 @@ pub(super) fn create( handle_inner: HandleInner, before_park: Option, after_unpark: Option, + before_task_poll: Option, ) -> (Arc, Launch) { let mut cores = Vec::with_capacity(size); let mut remotes = Vec::with_capacity(size); @@ -234,6 +238,7 @@ pub(super) fn create( shutdown_cores: Mutex::new(vec![]), before_park, after_unpark, + before_task_poll, scheduler_metrics: SchedulerMetrics::new(), worker_metrics: worker_metrics.into_boxed_slice(), }); @@ -424,6 +429,11 @@ impl Context { // Make the core available to the runtime context core.metrics.incr_poll_count(); + + if let Some(f) = &self.worker.shared.before_task_poll { + f(); + } + *self.core.borrow_mut() = Some(core); // Run the task @@ -449,6 +459,11 @@ impl Context { if coop::has_budget_remaining() { // Run the LIFO task, then loop core.metrics.incr_poll_count(); + + if let Some(f) = &self.worker.shared.before_task_poll { + f(); + } + *self.core.borrow_mut() = Some(core); let task = self.worker.shared.owned.assert_owner(task); task.run();