diff --git a/tokio/src/runtime/context.rs b/tokio/src/runtime/context.rs index e9011b66f..880b67205 100644 --- a/tokio/src/runtime/context.rs +++ b/tokio/src/runtime/context.rs @@ -10,10 +10,11 @@ cfg_rt! { mod scoped; use scoped::Scoped; - use crate::runtime::{scheduler, task::Id, Defer}; + use crate::runtime::{scheduler, task::Id}; use std::cell::RefCell; use std::marker::PhantomData; + use std::task::Waker; use std::time::Duration; cfg_taskdump! { @@ -45,11 +46,6 @@ struct Context { #[cfg(feature = "rt")] runtime: Cell, - /// Yielded task wakers are stored here and notified after resource drivers - /// are polled. - #[cfg(feature = "rt")] - defer: RefCell>, - #[cfg(any(feature = "rt", feature = "macros"))] rng: FastRand, @@ -93,9 +89,6 @@ tokio_thread_local! { #[cfg(feature = "rt")] runtime: Cell::new(EnterRuntime::NotEntered), - #[cfg(feature = "rt")] - defer: RefCell::new(None), - #[cfg(any(feature = "rt", feature = "macros"))] rng: FastRand::new(RngSeed::new()), @@ -170,12 +163,6 @@ cfg_rt! { #[allow(dead_code)] // Only tracking the guard. pub(crate) handle: SetCurrentGuard, - - /// If true, then this is the root runtime guard. It is possible to nest - /// runtime guards by using `block_in_place` between the calls. We need - /// to track the root guard as this is the guard responsible for freeing - /// the deferred task queue. - is_root: bool, } /// Guard tracking that a caller has entered a blocking region. @@ -240,20 +227,9 @@ cfg_rt! { // Set the entered flag c.runtime.set(EnterRuntime::Entered { allow_block_in_place }); - // Initialize queue to track yielded tasks - let mut defer = c.defer.borrow_mut(); - - let is_root = if defer.is_none() { - *defer = Some(Defer::new()); - true - } else { - false - }; - Some(EnterRuntimeGuard { blocking: BlockingRegionGuard::new(), handle: c.set_current(handle), - is_root, }) } }) @@ -292,11 +268,17 @@ cfg_rt! { DisallowBlockInPlaceGuard(reset) } - pub(crate) fn with_defer(f: impl FnOnce(&mut Defer) -> R) -> Option { - CONTEXT.with(|c| { - let mut defer = c.defer.borrow_mut(); - defer.as_mut().map(f) - }) + #[track_caller] + pub(crate) fn defer(waker: &Waker) { + with_scheduler(|maybe_scheduler| { + if let Some(scheduler) = maybe_scheduler { + scheduler.defer(waker); + } else { + // Called from outside of the runtime, immediately wake the + // task. + waker.wake_by_ref(); + } + }); } pub(super) fn set_scheduler(v: &scheduler::Context, f: impl FnOnce() -> R) -> R { @@ -342,10 +324,6 @@ cfg_rt! { CONTEXT.with(|c| { assert!(c.runtime.get().is_entered()); c.runtime.set(EnterRuntime::NotEntered); - - if self.is_root { - *c.defer.borrow_mut() = None; - } }); } } @@ -354,6 +332,7 @@ cfg_rt! { fn new() -> BlockingRegionGuard { BlockingRegionGuard { _p: PhantomData } } + /// Blocks the thread on the specified future, returning the value with /// which that future completes. pub(crate) fn block_on(&mut self, f: F) -> Result @@ -397,10 +376,6 @@ cfg_rt! { return Err(()); } - // Wake any yielded tasks before parking in order to avoid - // blocking. - with_defer(|defer| defer.wake()); - park.park_timeout(when - now); } } diff --git a/tokio/src/runtime/defer.rs b/tokio/src/runtime/defer.rs deleted file mode 100644 index 559f9acbe..000000000 --- a/tokio/src/runtime/defer.rs +++ /dev/null @@ -1,38 +0,0 @@ -use std::task::Waker; - -pub(crate) struct Defer { - deferred: Vec, -} - -impl Defer { - pub(crate) fn new() -> Defer { - Defer { - deferred: Default::default(), - } - } - - pub(crate) fn defer(&mut self, waker: &Waker) { - // If the same task adds itself a bunch of times, then only add it once. - if let Some(last) = self.deferred.last() { - if last.will_wake(waker) { - return; - } - } - self.deferred.push(waker.clone()); - } - - pub(crate) fn is_empty(&self) -> bool { - self.deferred.is_empty() - } - - pub(crate) fn wake(&mut self) { - for waker in self.deferred.drain(..) { - waker.wake(); - } - } - - #[cfg(tokio_taskdump)] - pub(crate) fn take_deferred(&mut self) -> Vec { - std::mem::take(&mut self.deferred) - } -} diff --git a/tokio/src/runtime/mod.rs b/tokio/src/runtime/mod.rs index a2efba479..cb198f51f 100644 --- a/tokio/src/runtime/mod.rs +++ b/tokio/src/runtime/mod.rs @@ -230,9 +230,6 @@ cfg_rt! { pub use crate::util::rand::RngSeed; } - mod defer; - pub(crate) use defer::Defer; - cfg_taskdump! { pub mod dump; pub use dump::Dump; diff --git a/tokio/src/runtime/park.rs b/tokio/src/runtime/park.rs index dc86f42a9..2392846ab 100644 --- a/tokio/src/runtime/park.rs +++ b/tokio/src/runtime/park.rs @@ -284,11 +284,6 @@ impl CachedParkThread { return Ok(v); } - // Wake any yielded tasks before parking in order to avoid - // blocking. - #[cfg(feature = "rt")] - crate::runtime::context::with_defer(|defer| defer.wake()); - self.park(); } } diff --git a/tokio/src/runtime/scheduler/current_thread.rs b/tokio/src/runtime/scheduler/current_thread.rs index 9a0578c45..2cbe29739 100644 --- a/tokio/src/runtime/scheduler/current_thread.rs +++ b/tokio/src/runtime/scheduler/current_thread.rs @@ -2,9 +2,9 @@ use crate::future::poll_fn; use crate::loom::sync::atomic::AtomicBool; use crate::loom::sync::Arc; use crate::runtime::driver::{self, Driver}; +use crate::runtime::scheduler::{self, Defer}; use crate::runtime::task::{self, Inject, JoinHandle, OwnedTasks, Schedule, Task}; -use crate::runtime::{blocking, context, scheduler, Config}; -use crate::runtime::{MetricsBatch, SchedulerMetrics, WorkerMetrics}; +use crate::runtime::{blocking, context, Config, MetricsBatch, SchedulerMetrics, WorkerMetrics}; use crate::sync::notify::Notify; use crate::util::atomic_cell::AtomicCell; use crate::util::{waker_ref, RngSeedGenerator, Wake, WakerRef}; @@ -15,6 +15,7 @@ use std::fmt; use std::future::Future; use std::sync::atomic::Ordering::{AcqRel, Release}; use std::task::Poll::{Pending, Ready}; +use std::task::Waker; use std::time::Duration; /// Executes tasks on the current thread @@ -98,6 +99,9 @@ pub(crate) struct Context { /// Scheduler core, enabling the holder of `Context` to execute the /// scheduler. core: RefCell>>, + + /// Deferred tasks, usually ones that called `task::yield_now()`. + pub(crate) defer: Defer, } type Notified = task::Notified>; @@ -201,6 +205,7 @@ impl CurrentThread { context: scheduler::Context::CurrentThread(Context { handle: handle.clone(), core: RefCell::new(Some(core)), + defer: Defer::new(), }), scheduler: self, }) @@ -320,21 +325,11 @@ impl Core { } } -fn did_defer_tasks() -> bool { - context::with_defer(|deferred| !deferred.is_empty()).unwrap() -} - -fn wake_deferred_tasks() { - context::with_defer(|deferred| deferred.wake()); -} - #[cfg(tokio_taskdump)] -fn wake_deferred_tasks_and_free() { - let wakers = context::with_defer(|deferred| deferred.take_deferred()); - if let Some(wakers) = wakers { - for waker in wakers { - waker.wake(); - } +fn wake_deferred_tasks_and_free(context: &Context) { + let wakers = context.defer.take_deferred(); + for waker in wakers { + waker.wake(); } } @@ -372,7 +367,7 @@ impl Context { let (c, _) = self.enter(core, || { driver.park(&handle.driver); - wake_deferred_tasks(); + self.defer.wake(); }); core = c; @@ -398,7 +393,7 @@ impl Context { let (mut core, _) = self.enter(core, || { driver.park_timeout(&handle.driver, Duration::from_millis(0)); - wake_deferred_tasks(); + self.defer.wake(); }); core.driver = Some(driver); @@ -418,6 +413,10 @@ impl Context { let core = self.core.borrow_mut().take().expect("core missing"); (core, ret) } + + pub(crate) fn defer(&self, waker: &Waker) { + self.defer.defer(waker); + } } // ===== impl Handle ===== @@ -479,12 +478,15 @@ impl Handle { .into_iter() .map(dump::Task::new) .collect(); - }); - // Taking a taskdump could wakes every task, but we probably don't want - // the `yield_now` vector to be that large under normal circumstances. - // Therefore, we free its allocation. - wake_deferred_tasks_and_free(); + // Avoid double borrow panic + drop(maybe_core); + + // Taking a taskdump could wakes every task, but we probably don't want + // the `yield_now` vector to be that large under normal circumstances. + // Therefore, we free its allocation. + wake_deferred_tasks_and_free(context); + }); dump::Dump::new(traces) } @@ -671,7 +673,7 @@ impl CoreGuard<'_> { None => { core.metrics.end_processing_scheduled_tasks(); - core = if did_defer_tasks() { + core = if !context.defer.is_empty() { context.park_yield(core, handle) } else { context.park(core, handle) diff --git a/tokio/src/runtime/scheduler/defer.rs b/tokio/src/runtime/scheduler/defer.rs new file mode 100644 index 000000000..a4be8ef2e --- /dev/null +++ b/tokio/src/runtime/scheduler/defer.rs @@ -0,0 +1,43 @@ +use std::cell::RefCell; +use std::task::Waker; + +pub(crate) struct Defer { + deferred: RefCell>, +} + +impl Defer { + pub(crate) fn new() -> Defer { + Defer { + deferred: Default::default(), + } + } + + pub(crate) fn defer(&self, waker: &Waker) { + let mut deferred = self.deferred.borrow_mut(); + + // If the same task adds itself a bunch of times, then only add it once. + if let Some(last) = deferred.last() { + if last.will_wake(waker) { + return; + } + } + + deferred.push(waker.clone()); + } + + pub(crate) fn is_empty(&self) -> bool { + self.deferred.borrow().is_empty() + } + + pub(crate) fn wake(&self) { + while let Some(waker) = self.deferred.borrow_mut().pop() { + waker.wake(); + } + } + + #[cfg(tokio_taskdump)] + pub(crate) fn take_deferred(&self) -> Vec { + let mut deferred = self.deferred.borrow_mut(); + std::mem::take(&mut *deferred) + } +} diff --git a/tokio/src/runtime/scheduler/mod.rs b/tokio/src/runtime/scheduler/mod.rs index 4d423f6e4..118e555d2 100644 --- a/tokio/src/runtime/scheduler/mod.rs +++ b/tokio/src/runtime/scheduler/mod.rs @@ -1,6 +1,9 @@ cfg_rt! { pub(crate) mod current_thread; pub(crate) use current_thread::CurrentThread; + + mod defer; + use defer::Defer; } cfg_rt_multi_thread! { @@ -56,6 +59,7 @@ cfg_rt! { use crate::runtime::context; use crate::task::JoinHandle; use crate::util::RngSeedGenerator; + use std::task::Waker; impl Handle { #[track_caller] @@ -203,6 +207,14 @@ cfg_rt! { } } + pub(crate) fn defer(&self, waker: &Waker) { + match self { + Context::CurrentThread(context) => context.defer(waker), + #[cfg(all(feature = "rt-multi-thread", not(tokio_wasi)))] + Context::MultiThread(context) => context.defer(waker), + } + } + cfg_rt_multi_thread! { #[track_caller] pub(crate) fn expect_multi_thread(&self) -> &multi_thread::Context { diff --git a/tokio/src/runtime/scheduler/multi_thread/worker.rs b/tokio/src/runtime/scheduler/multi_thread/worker.rs index 5cc5bd1dc..7b87a9c0d 100644 --- a/tokio/src/runtime/scheduler/multi_thread/worker.rs +++ b/tokio/src/runtime/scheduler/multi_thread/worker.rs @@ -62,6 +62,7 @@ use crate::runtime::context; use crate::runtime::scheduler::multi_thread::{ idle, queue, Counters, Handle, Idle, Parker, Stats, Unparker, }; +use crate::runtime::scheduler::Defer; use crate::runtime::task::{Inject, OwnedTasks}; use crate::runtime::{ blocking, coop, driver, scheduler, task, Config, SchedulerMetrics, WorkerMetrics, @@ -70,6 +71,7 @@ use crate::util::atomic_cell::AtomicCell; use crate::util::rand::{FastRand, RngSeedGenerator}; use std::cell::RefCell; +use std::task::Waker; use std::time::Duration; /// A scheduler worker @@ -189,6 +191,10 @@ pub(crate) struct Context { /// Core data core: RefCell>>, + + /// Tasks to wake after resource drivers are polled. This is mostly to + /// handle yielded tasks. + pub(crate) defer: Defer, } /// Starts the workers @@ -432,6 +438,7 @@ fn run(worker: Arc) { let cx = scheduler::Context::MultiThread(Context { worker, core: RefCell::new(None), + defer: Defer::new(), }); context::set_scheduler(&cx, || { @@ -444,7 +451,7 @@ fn run(worker: Arc) { // Check if there are any deferred tasks to notify. This can happen when // the worker core is lost due to `block_in_place()` being called from // within the task. - wake_deferred_tasks(); + cx.defer.wake(); }); } @@ -484,7 +491,7 @@ impl Context { core = self.run_task(task, core)?; } else { // Wait for work - core = if did_defer_tasks() { + core = if !self.defer.is_empty() { self.park_timeout(core, Some(Duration::from_millis(0))) } else { self.park(core) @@ -669,7 +676,7 @@ impl Context { park.park(&self.worker.handle.driver); } - wake_deferred_tasks(); + self.defer.wake(); // Remove `core` from context core = self.core.borrow_mut().take().expect("core missing"); @@ -685,6 +692,10 @@ impl Context { core } + + pub(crate) fn defer(&self, waker: &Waker) { + self.defer.defer(waker); + } } impl Core { @@ -1054,14 +1065,6 @@ impl Handle { } } -fn did_defer_tasks() -> bool { - context::with_defer(|deferred| !deferred.is_empty()).unwrap() -} - -fn wake_deferred_tasks() { - context::with_defer(|deferred| deferred.wake()); -} - #[track_caller] fn with_current(f: impl FnOnce(Option<&Context>) -> R) -> R { use scheduler::Context::MultiThread; diff --git a/tokio/src/runtime/task/trace/mod.rs b/tokio/src/runtime/task/trace/mod.rs index 2f1d7dba4..99abc7548 100644 --- a/tokio/src/runtime/task/trace/mod.rs +++ b/tokio/src/runtime/task/trace/mod.rs @@ -1,6 +1,8 @@ use crate::loom::sync::Arc; -use crate::runtime::scheduler::current_thread; +use crate::runtime::context; +use crate::runtime::scheduler::{self, current_thread}; use crate::runtime::task::Inject; + use backtrace::BacktraceFrame; use std::cell::Cell; use std::collections::VecDeque; @@ -179,10 +181,15 @@ pub(crate) fn trace_leaf(cx: &mut task::Context<'_>) -> Poll<()> { if did_trace { // Use the same logic that `yield_now` uses to send out wakeups after // the task yields. - let defer = crate::runtime::context::with_defer(|rt| { - rt.defer(cx.waker()); + context::with_scheduler(|scheduler| { + if let Some(scheduler) = scheduler { + match scheduler { + scheduler::Context::CurrentThread(s) => s.defer.defer(cx.waker()), + #[cfg(all(feature = "rt-multi-thread", not(tokio_wasi)))] + scheduler::Context::MultiThread(s) => s.defer.defer(cx.waker()), + } + } }); - debug_assert!(defer.is_some()); Poll::Pending } else { diff --git a/tokio/src/task/yield_now.rs b/tokio/src/task/yield_now.rs index 8f8341242..428d124c3 100644 --- a/tokio/src/task/yield_now.rs +++ b/tokio/src/task/yield_now.rs @@ -54,15 +54,7 @@ pub async fn yield_now() { self.yielded = true; - let defer = context::with_defer(|rt| { - rt.defer(cx.waker()); - }); - - if defer.is_none() { - // Not currently in a runtime, just notify ourselves - // immediately. - cx.waker().wake_by_ref(); - } + context::defer(cx.waker()); Poll::Pending } diff --git a/tokio/tests/task_yield_now.rs b/tokio/tests/task_yield_now.rs new file mode 100644 index 000000000..b16bca528 --- /dev/null +++ b/tokio/tests/task_yield_now.rs @@ -0,0 +1,16 @@ +#![cfg(all(feature = "full", tokio_unstable))] + +use tokio::task; +use tokio_test::task::spawn; + +// `yield_now` is tested within the runtime in `rt_common`. +#[test] +fn yield_now_outside_of_runtime() { + let mut task = spawn(async { + task::yield_now().await; + }); + + assert!(task.poll().is_pending()); + assert!(task.is_woken()); + assert!(task.poll().is_ready()); +}