hooks: overhaul task hooks

This change overhauls our task hooks system, allowing users to attach
user data to individual tasks, which are passed into hooks. Spawn hooks
can see both parent and child metadata.
This commit is contained in:
noah
2026-05-29 16:25:29 -05:00
parent 32312ae0d6
commit a3860a9f37
24 changed files with 1431 additions and 231 deletions
+11
View File
@@ -307,6 +307,8 @@ impl Spawner {
Mandatory::NonMandatory,
SpawnMeta::new_unnamed(fn_size),
rt,
#[cfg(tokio_unstable)]
None,
)
} else {
self.spawn_blocking_inner(
@@ -314,6 +316,8 @@ impl Spawner {
Mandatory::NonMandatory,
SpawnMeta::new_unnamed(fn_size),
rt,
#[cfg(tokio_unstable)]
None,
)
};
@@ -345,6 +349,8 @@ impl Spawner {
Mandatory::Mandatory,
SpawnMeta::new_unnamed(fn_size),
rt,
#[cfg(tokio_unstable)]
None,
)
} else {
self.spawn_blocking_inner(
@@ -352,6 +358,8 @@ impl Spawner {
Mandatory::Mandatory,
SpawnMeta::new_unnamed(fn_size),
rt,
#[cfg(tokio_unstable)]
None,
)
};
@@ -370,6 +378,7 @@ impl Spawner {
is_mandatory: Mandatory,
spawn_meta: SpawnMeta<'_>,
rt: &Handle,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> (JoinHandle<R>, Result<(), SpawnError>)
where
F: FnOnce() -> R + Send + 'static,
@@ -384,6 +393,8 @@ impl Spawner {
BlockingSchedule::new(rt),
id,
task::SpawnLocation::capture(),
#[cfg(tokio_unstable)]
user_data,
);
let spawned = self.spawn_task(Task::new(task, is_mandatory), rt);
+11 -8
View File
@@ -1,7 +1,9 @@
#[cfg(feature = "test-util")]
use crate::runtime::scheduler;
use crate::runtime::task::{self, Task, TaskHarnessScheduleHooks};
use crate::runtime::task::{self, Task};
use crate::runtime::Handle;
#[cfg(tokio_unstable)]
use crate::runtime::{TaskCallback, TaskMeta};
/// `task::Schedule` implementation that does nothing (except some bookkeeping
/// in test-util builds). This is unique to the blocking scheduler as tasks
@@ -12,7 +14,8 @@ use crate::runtime::Handle;
pub(crate) struct BlockingSchedule {
#[cfg(feature = "test-util")]
handle: Handle,
hooks: TaskHarnessScheduleHooks,
#[cfg(tokio_unstable)]
task_terminate_callback: Option<TaskCallback>,
}
impl BlockingSchedule {
@@ -31,9 +34,8 @@ impl BlockingSchedule {
BlockingSchedule {
#[cfg(feature = "test-util")]
handle: handle.clone(),
hooks: TaskHarnessScheduleHooks {
task_terminate_callback: handle.inner.hooks().task_terminate_callback.clone(),
},
#[cfg(tokio_unstable)]
task_terminate_callback: handle.inner.hooks().task_terminate_callback.clone(),
}
}
}
@@ -58,9 +60,10 @@ impl task::Schedule for BlockingSchedule {
unreachable!();
}
fn hooks(&self) -> TaskHarnessScheduleHooks {
TaskHarnessScheduleHooks {
task_terminate_callback: self.hooks.task_terminate_callback.clone(),
#[cfg(tokio_unstable)]
fn task_terminate_callback(&self, meta: &mut TaskMeta<'_>) {
if let Some(task_terminate_callback) = &self.task_terminate_callback {
task_terminate_callback(meta);
}
}
}
+8 -7
View File
@@ -2,7 +2,8 @@
use crate::runtime::handle::Handle;
use crate::runtime::{
blocking, driver, Callback, HistogramBuilder, Runtime, TaskCallback, TimerFlavor,
blocking, driver, Callback, HistogramBuilder, Runtime, TaskCallback, TaskSpawnCallback,
TimerFlavor,
};
#[cfg(tokio_unstable)]
use crate::runtime::{metrics::HistogramConfiguration, TaskMeta};
@@ -96,7 +97,7 @@ pub struct Builder {
pub(super) after_unpark: Option<Callback>,
/// To run before each task is spawned.
pub(super) before_spawn: Option<TaskCallback>,
pub(super) before_spawn: Option<TaskSpawnCallback>,
/// To run before each poll
#[cfg(tokio_unstable)]
@@ -873,7 +874,7 @@ impl Builder {
/// # use tokio::runtime;
/// # pub fn main() {
/// let runtime = runtime::Builder::new_current_thread()
/// .on_task_spawn(|_| {
/// .on_task_spawn(|_, _| {
/// println!("spawning task");
/// })
/// .build()
@@ -892,7 +893,7 @@ impl Builder {
#[cfg_attr(docsrs, doc(cfg(tokio_unstable)))]
pub fn on_task_spawn<F>(&mut self, f: F) -> &mut Self
where
F: Fn(&TaskMeta<'_>) + Send + Sync + 'static,
F: Fn(&mut TaskMeta<'_>, Option<crate::runtime::TaskMetaRef<'_>>) + Send + Sync + 'static,
{
self.before_spawn = Some(std::sync::Arc::new(f));
self
@@ -939,7 +940,7 @@ impl Builder {
#[cfg_attr(docsrs, doc(cfg(tokio_unstable)))]
pub fn on_before_task_poll<F>(&mut self, f: F) -> &mut Self
where
F: Fn(&TaskMeta<'_>) + Send + Sync + 'static,
F: Fn(&mut TaskMeta<'_>) + Send + Sync + 'static,
{
self.before_poll = Some(std::sync::Arc::new(f));
self
@@ -986,7 +987,7 @@ impl Builder {
#[cfg_attr(docsrs, doc(cfg(tokio_unstable)))]
pub fn on_after_task_poll<F>(&mut self, f: F) -> &mut Self
where
F: Fn(&TaskMeta<'_>) + Send + Sync + 'static,
F: Fn(&mut TaskMeta<'_>) + Send + Sync + 'static,
{
self.after_poll = Some(std::sync::Arc::new(f));
self
@@ -1035,7 +1036,7 @@ impl Builder {
#[cfg_attr(docsrs, doc(cfg(tokio_unstable)))]
pub fn on_task_terminate<F>(&mut self, f: F) -> &mut Self
where
F: Fn(&TaskMeta<'_>) + Send + Sync + 'static,
F: Fn(&mut TaskMeta<'_>) + Send + Sync + 'static,
{
self.after_termination = Some(std::sync::Arc::new(f));
self
+2 -2
View File
@@ -2,7 +2,7 @@
any(not(all(tokio_unstable, feature = "full")), target_family = "wasm"),
allow(dead_code)
)]
use crate::runtime::{Callback, TaskCallback};
use crate::runtime::{Callback, TaskCallback, TaskSpawnCallback};
use crate::util::RngSeedGenerator;
pub(crate) struct Config {
@@ -19,7 +19,7 @@ pub(crate) struct Config {
pub(crate) after_unpark: Option<Callback>,
/// To run before each task is spawned.
pub(crate) before_spawn: Option<TaskCallback>,
pub(crate) before_spawn: Option<TaskSpawnCallback>,
/// To run after each task is terminated.
pub(crate) after_termination: Option<TaskCallback>,
+22
View File
@@ -21,6 +21,8 @@ cfg_rt! {
use crate::runtime::{scheduler, task::Id};
#[cfg(tokio_unstable)]
use std::ptr::NonNull;
use std::task::Waker;
cfg_taskdump! {
@@ -49,6 +51,9 @@ struct Context {
#[cfg(feature = "rt")]
current_task_id: Cell<Option<Id>>,
#[cfg(all(feature = "rt", tokio_unstable))]
current_task: Cell<Option<NonNull<()>>>,
/// Tracks if the current thread is currently driving a runtime.
/// Note, that if this is set to "entered", the current scheduler
/// handle may not reference the runtime currently executing. This
@@ -92,6 +97,9 @@ tokio_thread_local! {
#[cfg(feature = "rt")]
current_task_id: Cell::new(None),
#[cfg(all(feature = "rt", tokio_unstable))]
current_task: Cell::new(None),
// Tracks if the current thread is currently driving a runtime.
// Note, that if this is set to "entered", the current scheduler
// handle may not reference the runtime currently executing. This
@@ -159,6 +167,20 @@ cfg_rt! {
CONTEXT.try_with(|ctx| ctx.current_task_id.get()).unwrap_or(None)
}
#[cfg(tokio_unstable)]
pub(crate) fn set_current_task(task: Option<NonNull<()>>) -> Option<NonNull<()>> {
CONTEXT
.try_with(|ctx| ctx.current_task.replace(task))
.unwrap_or(None)
}
#[cfg(tokio_unstable)]
pub(crate) fn current_task() -> Option<NonNull<()>> {
CONTEXT
.try_with(|ctx| ctx.current_task.get())
.unwrap_or(None)
}
#[cfg(tokio_unstable)]
pub(crate) fn worker_index() -> Option<usize> {
with_scheduler(|ctx| ctx.and_then(|c| c.worker_index()))
+91 -2
View File
@@ -372,6 +372,40 @@ impl Handle {
#[track_caller]
pub(crate) fn spawn_named<F>(&self, future: F, meta: SpawnMeta<'_>) -> JoinHandle<F::Output>
where
F: Future + Send + 'static,
F::Output: Send + 'static,
{
self.spawn_named_inner(
future,
meta,
#[cfg(tokio_unstable)]
None,
)
}
#[cfg(all(tokio_unstable, feature = "tracing"))]
#[track_caller]
pub(crate) fn spawn_named_with_data<F>(
&self,
future: F,
meta: SpawnMeta<'_>,
user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where
F: Future + Send + 'static,
F::Output: Send + 'static,
{
self.spawn_named_inner(future, meta, user_data)
}
#[track_caller]
fn spawn_named_inner<F>(
&self,
future: F,
meta: SpawnMeta<'_>,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where
F: Future + Send + 'static,
F::Output: Send + 'static,
@@ -387,7 +421,13 @@ impl Handle {
let future = super::task::trace::Trace::root(future);
#[cfg(all(tokio_unstable, feature = "tracing"))]
let future = crate::util::trace::task(future, "task", meta, id.as_u64());
self.inner.spawn(future, id, meta.spawned_at)
self.inner.spawn(
future,
id,
meta.spawned_at,
#[cfg(tokio_unstable)]
user_data,
)
}
#[track_caller]
@@ -401,6 +441,47 @@ impl Handle {
future: F,
meta: SpawnMeta<'_>,
) -> JoinHandle<F::Output>
where
F: Future + 'static,
F::Output: 'static,
{
unsafe {
self.spawn_local_named_inner(
future,
meta,
#[cfg(tokio_unstable)]
None,
)
}
}
/// # Safety
///
/// This must only be called in `LocalRuntime` if the runtime has been verified to be owned
/// by the current thread.
#[cfg(all(tokio_unstable, feature = "tracing"))]
#[track_caller]
#[allow(dead_code)]
pub(crate) unsafe fn spawn_local_named_with_data<F>(
&self,
future: F,
meta: SpawnMeta<'_>,
user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where
F: Future + 'static,
F::Output: 'static,
{
unsafe { self.spawn_local_named_inner(future, meta, user_data) }
}
#[track_caller]
unsafe fn spawn_local_named_inner<F>(
&self,
future: F,
meta: SpawnMeta<'_>,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where
F: Future + 'static,
F::Output: 'static,
@@ -416,7 +497,15 @@ impl Handle {
let future = super::task::trace::Trace::root(future);
#[cfg(all(tokio_unstable, feature = "tracing"))]
let future = crate::util::trace::task(future, "task", meta, id.as_u64());
unsafe { self.inner.spawn_local(future, id, meta.spawned_at) }
unsafe {
self.inner.spawn_local(
future,
id,
meta.spawned_at,
#[cfg(tokio_unstable)]
user_data,
)
}
}
/// Returns the flavor of the current `Runtime`.
+4 -2
View File
@@ -589,9 +589,11 @@ cfg_rt! {
}
mod task_hooks;
pub(crate) use task_hooks::{TaskHooks, TaskCallback};
pub(crate) use task_hooks::{TaskCallback, TaskHooks, TaskSpawnCallback};
#[cfg(tokio_unstable)]
pub(crate) use task_hooks::TaskData;
cfg_unstable! {
pub use task_hooks::TaskMeta;
pub use task_hooks::{TaskMeta, TaskMetaRef};
}
#[cfg(not(tokio_unstable))]
pub(crate) use task_hooks::TaskMeta;
@@ -3,7 +3,7 @@ use crate::loom::sync::Arc;
use crate::runtime::driver::{self, Driver};
use crate::runtime::scheduler::{self, Defer, Inject};
use crate::runtime::task::{
self, JoinHandle, OwnedTasks, Schedule, SpawnLocation, Task, TaskHarnessScheduleHooks,
self, JoinHandle, OwnedTasks, Schedule, SpawnLocation, Task,
};
use crate::runtime::{
blocking, context, Config, MetricsBatch, SchedulerMetrics, TaskHooks, TaskMeta, WorkerMetrics,
@@ -470,18 +470,35 @@ impl Handle {
future: F,
id: crate::runtime::task::Id,
spawned_at: SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where
F: crate::future::Future + Send + 'static,
F::Output: Send + 'static,
{
let (handle, notified) = me.shared.owned.bind(future, me.clone(), id, spawned_at);
me.task_hooks.spawn(&TaskMeta {
#[cfg(tokio_unstable)]
let parent = task::current_task_meta();
#[cfg(tokio_unstable)]
let (handle, notified) = me.shared.owned.bind_with_spawn_hook(
future,
me.clone(),
id,
spawned_at,
_phantom: Default::default(),
});
user_data,
|task| {
// Safety: the task is freshly allocated and not published yet.
let mut meta = unsafe { task.task_meta() };
me.task_hooks.spawn(&mut meta, parent);
},
);
#[cfg(not(tokio_unstable))]
let (handle, notified) = me.shared.owned.bind(future, me.clone(), id, spawned_at);
#[cfg(not(tokio_unstable))]
{
let mut meta = TaskMeta::new(id, spawned_at);
me.task_hooks.spawn(&mut meta, None);
}
if let Some(notified) = notified {
me.schedule(notified);
@@ -503,23 +520,44 @@ impl Handle {
future: F,
id: crate::runtime::task::Id,
spawned_at: SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where
F: crate::future::Future + 'static,
F::Output: 'static,
{
// Safety: the caller guarantees that this is only called on a `LocalRuntime`.
#[cfg(tokio_unstable)]
let parent = task::current_task_meta();
#[cfg(tokio_unstable)]
let before_bind = |task: &Task<Arc<Handle>>| {
// Safety: the task is freshly allocated and not published yet.
let mut meta = unsafe { task.task_meta() };
me.task_hooks.spawn(&mut meta, parent);
};
#[cfg(tokio_unstable)]
let (handle, notified) = unsafe {
me.shared.owned.bind_local_with_spawn_hook(
future,
me.clone(),
id,
spawned_at,
user_data,
before_bind,
)
};
#[cfg(not(tokio_unstable))]
let (handle, notified) = unsafe {
me.shared
.owned
.bind_local(future, me.clone(), id, spawned_at)
};
me.task_hooks.spawn(&TaskMeta {
id,
spawned_at,
_phantom: Default::default(),
});
#[cfg(not(tokio_unstable))]
{
let mut meta = TaskMeta::new(id, spawned_at);
me.task_hooks.spawn(&mut meta, None);
}
if let Some(notified) = notified {
me.schedule(notified);
@@ -687,13 +725,19 @@ impl Schedule for Arc<Handle> {
});
}
fn hooks(&self) -> TaskHarnessScheduleHooks {
TaskHarnessScheduleHooks {
task_terminate_callback: self.task_hooks.task_terminate_callback.clone(),
}
}
cfg_unstable! {
fn task_terminate_callback(&self, meta: &mut TaskMeta<'_>) {
self.task_hooks.task_terminate_callback(meta);
}
fn task_poll_start_callback(&self, meta: &mut TaskMeta<'_>) {
self.task_hooks.poll_start_callback(meta);
}
fn task_poll_stop_callback(&self, meta: &mut TaskMeta<'_>) {
self.task_hooks.poll_stop_callback(meta);
}
fn unhandled_panic(&self) {
use crate::runtime::UnhandledPanic;
@@ -815,17 +859,8 @@ impl CoreGuard<'_> {
let task = context.handle.shared.owned.assert_owner(task);
#[cfg(tokio_unstable)]
let task_meta = task.task_meta();
let (c, ()) = context.run_task(core, || {
#[cfg(tokio_unstable)]
context.handle.task_hooks.poll_start_callback(&task_meta);
task.run();
#[cfg(tokio_unstable)]
context.handle.task_hooks.poll_stop_callback(&task_meta);
});
core = c;
+42 -6
View File
@@ -8,8 +8,8 @@ cfg_rt! {
pub(crate) mod inject;
pub(crate) use inject::Inject;
#[cfg(tokio_unstable)]
use crate::runtime::TaskHooks;
use crate::runtime::WorkerMetrics;
}
@@ -149,16 +149,36 @@ cfg_rt! {
}
}
pub(crate) fn spawn<F>(&self, future: F, id: Id, spawned_at: SpawnLocation) -> JoinHandle<F::Output>
pub(crate) fn spawn<F>(
&self,
future: F,
id: Id,
spawned_at: SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where
F: Future + Send + 'static,
F::Output: Send + 'static,
{
match self {
Handle::CurrentThread(h) => current_thread::Handle::spawn(h, future, id, spawned_at),
Handle::CurrentThread(h) => current_thread::Handle::spawn(
h,
future,
id,
spawned_at,
#[cfg(tokio_unstable)]
user_data,
),
#[cfg(feature = "rt-multi-thread")]
Handle::MultiThread(h) => multi_thread::Handle::spawn(h, future, id, spawned_at),
Handle::MultiThread(h) => multi_thread::Handle::spawn(
h,
future,
id,
spawned_at,
#[cfg(tokio_unstable)]
user_data,
),
}
}
@@ -170,14 +190,29 @@ cfg_rt! {
/// by the current thread.
#[allow(irrefutable_let_patterns)]
#[track_caller]
pub(crate) unsafe fn spawn_local<F>(&self, future: F, id: Id, spawned_at: SpawnLocation) -> JoinHandle<F::Output>
pub(crate) unsafe fn spawn_local<F>(
&self,
future: F,
id: Id,
spawned_at: SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where
F: Future + 'static,
F::Output: 'static,
{
if let Handle::CurrentThread(h) = self {
// Safety: caller guarantees that this is a `LocalRuntime`.
unsafe { current_thread::Handle::spawn_local(h, future, id, spawned_at) }
unsafe {
current_thread::Handle::spawn_local(
h,
future,
id,
spawned_at,
#[cfg(tokio_unstable)]
user_data,
)
}
} else {
panic!("Only current_thread and LocalSet have spawn_local internals implemented")
}
@@ -204,6 +239,7 @@ cfg_rt! {
}
}
#[cfg(tokio_unstable)]
pub(crate) fn hooks(&self) -> &TaskHooks {
match self {
Handle::CurrentThread(h) => &h.task_hooks,
@@ -1,7 +1,7 @@
use crate::future::Future;
use crate::loom::sync::Arc;
use crate::runtime::scheduler::multi_thread::worker;
use crate::runtime::task::{Notified, Task, TaskHarnessScheduleHooks};
use crate::runtime::task::{Notified, Task};
use crate::runtime::{
blocking, driver,
task::{self, JoinHandle, SpawnLocation},
@@ -57,12 +57,20 @@ impl Handle {
future: F,
id: task::Id,
spawned_at: SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where
F: crate::future::Future + Send + 'static,
F::Output: Send + 'static,
{
Self::bind_new_task(me, future, id, spawned_at)
Self::bind_new_task(
me,
future,
id,
spawned_at,
#[cfg(tokio_unstable)]
user_data,
)
}
#[cfg(all(tokio_unstable, feature = "time"))]
@@ -83,18 +91,35 @@ impl Handle {
future: T,
id: task::Id,
spawned_at: SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<T::Output>
where
T: Future + Send + 'static,
T::Output: Send + 'static,
{
let (handle, notified) = me.shared.owned.bind(future, me.clone(), id, spawned_at);
me.task_hooks.spawn(&TaskMeta {
#[cfg(tokio_unstable)]
let parent = task::current_task_meta();
#[cfg(tokio_unstable)]
let (handle, notified) = me.shared.owned.bind_with_spawn_hook(
future,
me.clone(),
id,
spawned_at,
_phantom: Default::default(),
});
user_data,
|task| {
// Safety: the task is freshly allocated and not published yet.
let mut meta = unsafe { task.task_meta() };
me.task_hooks.spawn(&mut meta, parent);
},
);
#[cfg(not(tokio_unstable))]
let (handle, notified) = me.shared.owned.bind(future, me.clone(), id, spawned_at);
#[cfg(not(tokio_unstable))]
{
let mut meta = TaskMeta::new(id, spawned_at);
me.task_hooks.spawn(&mut meta, None);
}
me.schedule_option_task_without_yield(notified);
@@ -111,15 +136,24 @@ impl task::Schedule for Arc<Handle> {
self.schedule_task(task, false);
}
fn hooks(&self) -> TaskHarnessScheduleHooks {
TaskHarnessScheduleHooks {
task_terminate_callback: self.task_hooks.task_terminate_callback.clone(),
}
}
fn yield_now(&self, task: Notified<Self>) {
self.schedule_task(task, true);
}
#[cfg(tokio_unstable)]
fn task_terminate_callback(&self, meta: &mut TaskMeta<'_>) {
self.task_hooks.task_terminate_callback(meta);
}
#[cfg(tokio_unstable)]
fn task_poll_start_callback(&self, meta: &mut TaskMeta<'_>) {
self.task_hooks.poll_start_callback(meta);
}
#[cfg(tokio_unstable)]
fn task_poll_stop_callback(&self, meta: &mut TaskMeta<'_>) {
self.task_hooks.poll_stop_callback(meta);
}
}
impl Handle {
@@ -628,9 +628,6 @@ impl Context {
}
fn run_task(&self, task: Notified, mut core: Box<Core>) -> RunResult {
#[cfg(tokio_unstable)]
let task_meta = task.task_meta();
let task = self.worker.handle.shared.owned.assert_owner(task);
// Make sure the worker is not in the **searching** state. This enables
@@ -673,19 +670,8 @@ impl Context {
// Run the task
coop::budget(|| {
// Unlike the poll time above, poll start callback is attached to the task id,
// so it is tightly associated with the actual poll invocation.
#[cfg(tokio_unstable)]
self.worker
.handle
.task_hooks
.poll_start_callback(&task_meta);
task.run();
#[cfg(tokio_unstable)]
self.worker.handle.task_hooks.poll_stop_callback(&task_meta);
let mut lifo_polls = 0;
// As long as there is budget remaining and a task exists in the
@@ -749,19 +735,7 @@ impl Context {
*self.core.borrow_mut() = Some(core);
let task = self.worker.handle.shared.owned.assert_owner(task);
#[cfg(tokio_unstable)]
let task_meta = task.task_meta();
#[cfg(tokio_unstable)]
self.worker
.handle
.task_hooks
.poll_start_callback(&task_meta);
task.run();
#[cfg(tokio_unstable)]
self.worker.handle.task_hooks.poll_stop_callback(&task_meta);
}
})
}
+47 -8
View File
@@ -23,7 +23,9 @@ use crate::loom::cell::UnsafeCell;
use crate::runtime::context;
use crate::runtime::task::raw::{self, Vtable};
use crate::runtime::task::state::State;
use crate::runtime::task::{Id, Schedule, TaskHarnessScheduleHooks};
use crate::runtime::task::{Id, Schedule};
#[cfg(tokio_unstable)]
use crate::runtime::TaskData;
use crate::util::linked_list;
use std::num::NonZeroU64;
@@ -203,9 +205,8 @@ pub(super) struct Trailer {
pub(super) owned: linked_list::Pointers<Header>,
/// Consumer task waiting on completion of this task.
pub(super) waker: UnsafeCell<Option<Waker>>,
/// Optional hooks needed in the harness.
#[cfg_attr(not(tokio_unstable), allow(dead_code))] //TODO: remove when hooks are stabilized
pub(super) hooks: TaskHarnessScheduleHooks,
#[cfg(tokio_unstable)]
pub(super) user_data: UnsafeCell<Option<TaskData>>,
}
generate_addr_of_methods! {
@@ -233,6 +234,7 @@ impl<T: Future, S: Schedule> Cell<T, S> {
state: State,
task_id: Id,
#[cfg(tokio_unstable)] spawned_at: &'static Location<'static>,
#[cfg(tokio_unstable)] user_data: Option<TaskData>,
) -> Box<Cell<T, S>> {
// Separated into a non-generic function to reduce LLVM codegen
fn new_header(
@@ -254,7 +256,10 @@ impl<T: Future, S: Schedule> Cell<T, S> {
let tracing_id = future.id();
let vtable = raw::vtable::<T, S>();
let result = Box::new(Cell {
trailer: Trailer::new(scheduler.hooks()),
trailer: Trailer::new(
#[cfg(tokio_unstable)]
user_data,
),
header: new_header(
state,
vtable,
@@ -359,7 +364,11 @@ impl<T: Future, S: Schedule> Core<T, S> {
///
/// `self` must also be pinned. This is handled by storing the task on the
/// heap.
pub(super) fn poll(&self, mut cx: Context<'_>) -> Poll<T::Output> {
pub(super) fn poll(
&self,
#[cfg(tokio_unstable)] header: NonNull<Header>,
mut cx: Context<'_>,
) -> Poll<T::Output> {
let res = {
self.stage.stage.with_mut(|ptr| {
// Safety: The caller ensures mutual exclusion to the field.
@@ -372,6 +381,8 @@ impl<T: Future, S: Schedule> Core<T, S> {
let future = unsafe { Pin::new_unchecked(future) };
let _guard = TaskIdGuard::enter(self.task_id);
#[cfg(tokio_unstable)]
let _current_task = CurrentTaskGuard::enter(header);
future.poll(&mut cx)
})
};
@@ -430,6 +441,27 @@ impl<T: Future, S: Schedule> Core<T, S> {
}
}
#[cfg(tokio_unstable)]
pub(crate) struct CurrentTaskGuard {
parent_task: Option<NonNull<()>>,
}
#[cfg(tokio_unstable)]
impl CurrentTaskGuard {
fn enter(header: NonNull<Header>) -> Self {
CurrentTaskGuard {
parent_task: context::set_current_task(Some(header.cast())),
}
}
}
#[cfg(tokio_unstable)]
impl Drop for CurrentTaskGuard {
fn drop(&mut self) {
context::set_current_task(self.parent_task);
}
}
impl Header {
pub(super) unsafe fn set_next(&self, next: Option<NonNull<Header>>) {
self.queue_next.with_mut(|ptr| *ptr = next);
@@ -537,14 +569,21 @@ impl Header {
}
impl Trailer {
fn new(hooks: TaskHarnessScheduleHooks) -> Self {
fn new(#[cfg(tokio_unstable)] user_data: Option<TaskData>) -> Self {
Trailer {
waker: UnsafeCell::new(None),
owned: linked_list::Pointers::new(),
hooks,
#[cfg(tokio_unstable)]
user_data: UnsafeCell::new(user_data),
}
}
#[cfg(tokio_unstable)]
pub(super) fn user_data_ptr(&self) -> NonNull<Option<TaskData>> {
self.user_data
.with_mut(|ptr| unsafe { NonNull::new_unchecked(ptr) })
}
pub(super) unsafe fn set_waker(&self, waker: Option<Waker>) {
self.waker.with_mut(|ptr| {
*ptr = waker;
+101 -9
View File
@@ -48,6 +48,21 @@ where
fn core(&self) -> &Core<T, S> {
unsafe { &self.cell.as_ref().core }
}
/// # Safety
///
/// The returned metadata must only be used while the caller has exclusive
/// access to task hook data.
#[cfg(tokio_unstable)]
unsafe fn task_meta<'meta>(&self) -> TaskMeta<'meta> {
unsafe {
TaskMeta::new(
self.core().task_id,
self.core().spawned_at.into(),
Some(self.trailer().user_data_ptr()),
)
}
}
}
/// Task operations that can be implemented without being generic over the
@@ -204,10 +219,56 @@ where
TransitionToIdle::Cancelled => PollFuture::Complete,
}
}
#[cfg(tokio_unstable)]
{
// Safety: the task is in the RUNNING state, which excludes
// concurrent shutdown and termination metadata access.
let mut task_meta = unsafe { self.task_meta() };
let res = panic::catch_unwind(panic::AssertUnwindSafe(|| {
self.core()
.scheduler
.task_poll_start_callback(&mut task_meta);
}));
if let Err(panic) = res {
poll_hook_panic(self.core(), panic);
return PollFuture::Complete;
}
if self.state().load().is_cancelled() {
cancel_task(self.core());
return PollFuture::Complete;
}
}
let header_ptr = self.header_ptr();
let waker_ref = waker_ref::<S>(&header_ptr);
let cx = Context::from_waker(&waker_ref);
let res = poll_future(self.core(), cx);
let res = poll_future(
self.core(),
#[cfg(tokio_unstable)]
header_ptr,
cx,
);
#[cfg(tokio_unstable)]
{
// Safety: the task is still in the RUNNING state, which
// excludes concurrent shutdown and termination metadata
// access.
let mut task_meta = unsafe { self.task_meta() };
let hook_res = panic::catch_unwind(panic::AssertUnwindSafe(|| {
self.core()
.scheduler
.task_poll_stop_callback(&mut task_meta);
}));
if let Err(panic) = hook_res {
poll_hook_panic(self.core(), panic);
return PollFuture::Complete;
}
}
if res == Poll::Ready(()) {
// The future completed. Move on to complete the task.
@@ -255,6 +316,8 @@ where
// because we are going to drop them. This only matters when running
// under loom.
self.trailer().waker.with_mut(|_| ());
#[cfg(tokio_unstable)]
self.trailer().user_data.with_mut(|_| ());
self.core().stage.with_mut(|_| ());
// Safety: The caller of this method just transitioned our ref-count to
@@ -369,13 +432,12 @@ where
// We call this in a separate block so that it runs after the task appears to have
// completed and will still run if the destructor panics.
#[cfg(tokio_unstable)]
if let Some(f) = self.trailer().hooks.task_terminate_callback.as_ref() {
{
// Safety: completion owns the task lifecycle transition, and the
// terminate hook is invoked synchronously with exclusive metadata access.
let mut meta = unsafe { self.task_meta() };
let _ = panic::catch_unwind(panic::AssertUnwindSafe(|| {
f(&TaskMeta {
id: self.core().task_id,
spawned_at: self.core().spawned_at.into(),
_phantom: Default::default(),
})
self.core().scheduler.task_terminate_callback(&mut meta)
}));
}
@@ -516,9 +578,35 @@ fn panic_result_to_join_error(
}
}
#[cfg(tokio_unstable)]
fn poll_hook_panic<T: Future, S: Schedule>(
core: &Core<T, S>,
hook_panic: Box<dyn Any + Send + 'static>,
) {
let drop_res = panic::catch_unwind(panic::AssertUnwindSafe(|| {
core.drop_future_or_output();
}));
let join_error = match drop_res {
Ok(()) => panic_to_error(&core.scheduler, core.task_id, hook_panic),
Err(drop_panic) => panic_to_error(&core.scheduler, core.task_id, drop_panic),
};
let res = panic::catch_unwind(panic::AssertUnwindSafe(|| {
core.store_output(Err(join_error));
}));
if res.is_err() {
core.scheduler.unhandled_panic();
}
}
/// Polls the future. If the future completes, the output is written to the
/// stage field.
fn poll_future<T: Future, S: Schedule>(core: &Core<T, S>, cx: Context<'_>) -> Poll<()> {
fn poll_future<T: Future, S: Schedule>(
core: &Core<T, S>,
#[cfg(tokio_unstable)] header: NonNull<Header>,
cx: Context<'_>,
) -> Poll<()> {
// Poll the future.
let output = panic::catch_unwind(panic::AssertUnwindSafe(|| {
struct Guard<'a, T: Future, S: Schedule> {
@@ -532,7 +620,11 @@ fn poll_future<T: Future, S: Schedule>(core: &Core<T, S>, cx: Context<'_>) -> Po
}
}
let guard = Guard { core };
let res = guard.core.poll(cx);
let res = guard.core.poll(
#[cfg(tokio_unstable)]
header,
cx,
);
mem::forget(guard);
res
}));
+76 -3
View File
@@ -86,19 +86,50 @@ impl<S: 'static> OwnedTasks<S> {
/// Binds the provided task to this `OwnedTasks` instance. This fails if the
/// `OwnedTasks` has been closed.
#[cfg_attr(tokio_unstable, allow(dead_code))]
pub(crate) fn bind<T>(
&self,
task: T,
scheduler: S,
id: super::Id,
spawned_at: SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> (JoinHandle<T::Output>, Option<Notified<S>>)
where
S: Schedule,
T: Future + Send + 'static,
T::Output: Send + 'static,
{
let (task, notified, join) = super::new_task(task, scheduler, id, spawned_at);
let (task, notified, join) = super::new_task(
task,
scheduler,
id,
spawned_at,
#[cfg(tokio_unstable)]
user_data,
);
let notified = unsafe { self.bind_inner(task, notified) };
(join, notified)
}
#[cfg(tokio_unstable)]
pub(crate) fn bind_with_spawn_hook<T>(
&self,
task: T,
scheduler: S,
id: super::Id,
spawned_at: SpawnLocation,
user_data: Option<crate::runtime::TaskData>,
before_bind: impl FnOnce(&Task<S>),
) -> (JoinHandle<T::Output>, Option<Notified<S>>)
where
S: Schedule,
T: Future + Send + 'static,
T::Output: Send + 'static,
{
let (task, notified, join) =
super::new_task(task, scheduler, id, spawned_at, user_data);
before_bind(&task);
let notified = unsafe { self.bind_inner(task, notified) };
(join, notified)
}
@@ -108,19 +139,53 @@ impl<S: 'static> OwnedTasks<S> {
/// # Safety
///
/// Only use this in `LocalRuntime` where the task cannot move
#[cfg_attr(tokio_unstable, allow(dead_code))]
pub(crate) unsafe fn bind_local<T>(
&self,
task: T,
scheduler: S,
id: super::Id,
spawned_at: SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> (JoinHandle<T::Output>, Option<Notified<S>>)
where
S: Schedule,
T: Future + 'static,
T::Output: 'static,
{
let (task, notified, join) = super::new_task(task, scheduler, id, spawned_at);
let (task, notified, join) = super::new_task(
task,
scheduler,
id,
spawned_at,
#[cfg(tokio_unstable)]
user_data,
);
let notified = unsafe { self.bind_inner(task, notified) };
(join, notified)
}
/// # Safety
///
/// Only use this in `LocalRuntime` where the task cannot move.
#[cfg(tokio_unstable)]
pub(crate) unsafe fn bind_local_with_spawn_hook<T>(
&self,
task: T,
scheduler: S,
id: super::Id,
spawned_at: SpawnLocation,
user_data: Option<crate::runtime::TaskData>,
before_bind: impl FnOnce(&Task<S>),
) -> (JoinHandle<T::Output>, Option<Notified<S>>)
where
S: Schedule,
T: Future + 'static,
T::Output: 'static,
{
let (task, notified, join) =
super::new_task(task, scheduler, id, spawned_at, user_data);
before_bind(&task);
let notified = unsafe { self.bind_inner(task, notified) };
(join, notified)
}
@@ -264,13 +329,21 @@ impl<S: 'static> LocalOwnedTasks<S> {
scheduler: S,
id: super::Id,
spawned_at: SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> (JoinHandle<T::Output>, Option<Notified<S>>)
where
S: Schedule,
T: Future + 'static,
T::Output: 'static,
{
let (task, notified, join) = super::new_task(task, scheduler, id, spawned_at);
let (task, notified, join) = super::new_task(
task,
scheduler,
id,
spawned_at,
#[cfg(tokio_unstable)]
user_data,
);
unsafe {
// safety: We just created the task, so we have exclusive access
+39 -37
View File
@@ -119,6 +119,12 @@
//! the `JoinHandle` exclusive access again so that it is able to drop the waker
//! at a later point.
//!
//! * The `user_data` field is accessed only through scoped task hook metadata.
//! The spawn hook runs before the task is scheduled, poll hooks run while
//! the task holds the RUNNING lock but outside the actual future poll, and
//! the terminate hook runs after completion. Parent task metadata exposed to
//! spawn hooks is read-only.
//!
//! All other fields are immutable and can be accessed immutably without
//! synchronization by anyone.
//!
@@ -221,7 +227,6 @@ use crate::future::Future;
use crate::util::linked_list;
use crate::util::sharded_list;
use crate::runtime::TaskCallback;
use std::marker::PhantomData;
use std::panic::Location;
use std::ptr::NonNull;
@@ -241,14 +246,6 @@ unsafe impl<S> Sync for Task<S> {}
#[repr(transparent)]
pub(crate) struct Notified<S: 'static>(Task<S>);
impl<S> Notified<S> {
#[cfg(all(tokio_unstable, feature = "rt-multi-thread"))]
#[inline]
pub(crate) fn task_meta<'meta>(&self) -> crate::runtime::TaskMeta<'meta> {
self.0.task_meta()
}
}
// safety: This type cannot be used to touch the task without first verifying
// that the value is on a thread where it is safe to poll the task.
unsafe impl<S: Schedule> Send for Notified<S> {}
@@ -262,14 +259,6 @@ pub(crate) struct LocalNotified<S: 'static> {
_not_send: PhantomData<*const ()>,
}
impl<S> LocalNotified<S> {
#[cfg(tokio_unstable)]
#[inline]
pub(crate) fn task_meta<'meta>(&self) -> crate::runtime::TaskMeta<'meta> {
self.task.task_meta()
}
}
/// A task that is not owned by any `OwnedTasks`. Used for blocking tasks.
/// This type holds two ref-counts.
pub(crate) struct UnownedTask<S: 'static> {
@@ -284,12 +273,6 @@ unsafe impl<S> Sync for UnownedTask<S> {}
/// Task result sent back.
pub(crate) type Result<T> = std::result::Result<T, JoinError>;
/// Hooks for scheduling tasks which are needed in the task harness.
#[derive(Clone)]
pub(crate) struct TaskHarnessScheduleHooks {
pub(crate) task_terminate_callback: Option<TaskCallback>,
}
pub(crate) trait Schedule: Sync + Sized + 'static {
/// The task has completed work and is ready to be released. The scheduler
/// should release it immediately and return it. The task module will batch
@@ -301,8 +284,6 @@ pub(crate) trait Schedule: Sync + Sized + 'static {
/// Schedule the task
fn schedule(&self, task: Notified<Self>);
fn hooks(&self) -> TaskHarnessScheduleHooks;
/// Schedule the task to run in the near future, yielding the thread to
/// other tasks.
fn yield_now(&self, task: Notified<Self>) {
@@ -313,6 +294,15 @@ pub(crate) trait Schedule: Sync + Sized + 'static {
fn unhandled_panic(&self) {
// By default, do nothing. This maintains the 1.0 behavior.
}
#[cfg(tokio_unstable)]
fn task_terminate_callback(&self, _meta: &mut crate::runtime::TaskMeta<'_>) {}
#[cfg(tokio_unstable)]
fn task_poll_start_callback(&self, _meta: &mut crate::runtime::TaskMeta<'_>) {}
#[cfg(tokio_unstable)]
fn task_poll_stop_callback(&self, _meta: &mut crate::runtime::TaskMeta<'_>) {}
}
cfg_rt! {
@@ -325,6 +315,7 @@ cfg_rt! {
scheduler: S,
id: Id,
spawned_at: SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> (Task<S>, Notified<S>, JoinHandle<T::Output>)
where
S: Schedule,
@@ -336,6 +327,8 @@ cfg_rt! {
scheduler,
id,
spawned_at,
#[cfg(tokio_unstable)]
user_data,
);
let task = Task {
raw,
@@ -359,6 +352,7 @@ cfg_rt! {
scheduler: S,
id: Id,
spawned_at: SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> (UnownedTask<S>, JoinHandle<T::Output>)
where
S: Schedule,
@@ -370,6 +364,8 @@ cfg_rt! {
scheduler,
id,
spawned_at,
#[cfg(tokio_unstable)]
user_data,
);
// This transfers the ref-count of task and notified into an UnownedTask.
@@ -429,22 +425,16 @@ impl<S: 'static> Task<S> {
unsafe { Header::get_id(self.raw.header_ptr()) }
}
#[cfg(tokio_unstable)]
pub(crate) fn spawned_at(&self) -> &'static Location<'static> {
// Safety: The header pointer is valid.
unsafe { Header::get_spawn_location(self.raw.header_ptr()) }
}
// Explicit `'task` and `'meta` lifetimes are necessary here, as otherwise,
// the compiler infers the lifetimes to be the same, and considers the task
// to be borrowed for the lifetime of the returned `TaskMeta`.
#[cfg(tokio_unstable)]
pub(crate) fn task_meta<'meta>(&self) -> crate::runtime::TaskMeta<'meta> {
crate::runtime::TaskMeta {
id: self.id(),
spawned_at: self.spawned_at().into(),
_phantom: PhantomData,
}
/// # Safety
///
/// The returned metadata must have exclusive access to hook data for as long
/// as it can expose mutable references.
pub(crate) unsafe fn task_meta<'meta>(&self) -> crate::runtime::TaskMeta<'meta> {
unsafe { self.raw.task_meta() }
}
cfg_taskdump! {
@@ -677,3 +667,15 @@ impl SpawnLocation {
Self::from(Location::caller())
}
}
#[cfg(tokio_unstable)]
pub(crate) fn current_task_meta<'meta>() -> Option<crate::runtime::TaskMetaRef<'meta>> {
let ptr = crate::runtime::context::current_task()?;
// Safety: the context stores this pointer only while the referenced task is
// being polled, so the allocation is alive for the duration of this call.
let raw = unsafe { RawTask::from_raw(ptr.cast()) };
// Safety: parent metadata is exposed read-only during synchronous spawn
// hook invocation while no mutable parent hook metadata is live.
Some(unsafe { raw.task_meta_ref() })
}
+39
View File
@@ -11,6 +11,8 @@ use crate::future::Future;
use crate::runtime::task::core::{Core, Trailer};
use crate::runtime::task::{Cell, Harness, Header, Id, Schedule, State};
#[cfg(tokio_unstable)]
use crate::runtime::TaskData;
#[cfg(tokio_unstable)]
use std::panic::Location;
use std::ptr::NonNull;
use std::task::{Poll, Waker};
@@ -213,6 +215,7 @@ impl RawTask {
scheduler: S,
id: Id,
_spawned_at: super::SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<TaskData>,
) -> RawTask
where
T: Future,
@@ -225,6 +228,8 @@ impl RawTask {
id,
#[cfg(tokio_unstable)]
_spawned_at.0,
#[cfg(tokio_unstable)]
user_data,
));
let ptr = unsafe { NonNull::new_unchecked(ptr.cast()) };
@@ -267,6 +272,40 @@ impl RawTask {
unsafe { &*self.trailer_ptr().as_ptr() }
}
#[cfg(tokio_unstable)]
/// # Safety
///
/// The task allocation must be live, and the returned metadata must have
/// exclusive access to hook data for as long as it can expose mutable references.
pub(crate) unsafe fn task_meta<'meta>(&self) -> crate::runtime::TaskMeta<'meta> {
// Safety: `self` holds a live task reference, and callers use the
// metadata only for the current hook invocation.
unsafe {
crate::runtime::TaskMeta::new(
Header::get_id(self.ptr),
Header::get_spawn_location(self.ptr).into(),
Some(self.trailer().user_data_ptr()),
)
}
}
#[cfg(tokio_unstable)]
/// # Safety
///
/// The task allocation must be live, and hook data must not be mutated while
/// references exposed through the returned metadata are live.
pub(crate) unsafe fn task_meta_ref<'meta>(&self) -> crate::runtime::TaskMetaRef<'meta> {
// Safety: `self` holds a live task reference, and this only exposes
// shared access to task data.
unsafe {
crate::runtime::TaskMetaRef::new(
Header::get_id(self.ptr),
Header::get_spawn_location(self.ptr).into(),
Some(self.trailer().user_data_ptr()),
)
}
}
/// Returns a reference to the task's state.
pub(super) fn state(&self) -> &State {
&self.header().state
+186 -6
View File
@@ -1,10 +1,18 @@
use super::Config;
#[cfg(tokio_unstable)]
use std::any::Any;
use std::marker::PhantomData;
#[cfg(tokio_unstable)]
use std::ptr::NonNull;
use std::sync::Arc;
#[cfg(tokio_unstable)]
pub(crate) type TaskData = Box<dyn Any + Send + Sync + 'static>;
impl TaskHooks {
pub(crate) fn spawn(&self, meta: &TaskMeta<'_>) {
pub(crate) fn spawn(&self, meta: &mut TaskMeta<'_>, parent: Option<TaskMetaRef<'_>>) {
if let Some(f) = self.task_spawn_callback.as_ref() {
f(meta)
f(meta, parent)
}
}
@@ -22,7 +30,7 @@ impl TaskHooks {
#[cfg(tokio_unstable)]
#[inline]
pub(crate) fn poll_start_callback(&self, meta: &TaskMeta<'_>) {
pub(crate) fn poll_start_callback(&self, meta: &mut TaskMeta<'_>) {
if let Some(poll_start) = &self.before_poll_callback {
(poll_start)(meta);
}
@@ -30,16 +38,25 @@ impl TaskHooks {
#[cfg(tokio_unstable)]
#[inline]
pub(crate) fn poll_stop_callback(&self, meta: &TaskMeta<'_>) {
pub(crate) fn poll_stop_callback(&self, meta: &mut TaskMeta<'_>) {
if let Some(poll_stop) = &self.after_poll_callback {
(poll_stop)(meta);
}
}
#[cfg(tokio_unstable)]
#[inline]
pub(crate) fn task_terminate_callback(&self, meta: &mut TaskMeta<'_>) {
if let Some(task_terminate) = &self.task_terminate_callback {
(task_terminate)(meta);
}
}
}
#[derive(Clone)]
pub(crate) struct TaskHooks {
pub(crate) task_spawn_callback: Option<TaskCallback>,
pub(crate) task_spawn_callback: Option<TaskSpawnCallback>,
#[cfg_attr(not(tokio_unstable), allow(dead_code))]
pub(crate) task_terminate_callback: Option<TaskCallback>,
#[cfg(tokio_unstable)]
pub(crate) before_poll_callback: Option<TaskCallback>,
@@ -62,10 +79,44 @@ pub struct TaskMeta<'a> {
/// The location where the task was spawned.
#[cfg_attr(not(tokio_unstable), allow(unreachable_pub, dead_code))]
pub(crate) spawned_at: crate::runtime::task::SpawnLocation,
#[cfg(tokio_unstable)]
pub(crate) user_data: Option<NonNull<Option<TaskData>>>,
pub(crate) _phantom: PhantomData<&'a ()>,
}
impl<'a> TaskMeta<'a> {
#[cfg(not(tokio_unstable))]
pub(crate) fn new(
id: super::task::Id,
spawned_at: crate::runtime::task::SpawnLocation,
) -> Self {
Self {
id,
spawned_at,
_phantom: PhantomData,
}
}
/// # Safety
///
/// If `user_data` is present, it must point to live task storage, and this
/// metadata value must have exclusive access to that storage while it can
/// expose mutable references to it.
#[cfg(tokio_unstable)]
pub(crate) unsafe fn new(
id: super::task::Id,
spawned_at: crate::runtime::task::SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<NonNull<Option<TaskData>>>,
) -> Self {
Self {
id,
spawned_at,
#[cfg(tokio_unstable)]
user_data,
_phantom: PhantomData,
}
}
/// Return the opaque ID of the task.
#[cfg_attr(not(tokio_unstable), allow(unreachable_pub, dead_code))]
pub fn id(&self) -> super::task::Id {
@@ -77,7 +128,136 @@ impl<'a> TaskMeta<'a> {
pub fn spawned_at(&self) -> &'static std::panic::Location<'static> {
self.spawned_at.0
}
/// Returns a shared reference to this task's user data when the stored type
/// is `T`.
#[cfg(tokio_unstable)]
pub fn data<T: Any>(&self) -> Option<&T> {
let user_data = self.user_data?;
// Safety: `TaskMeta` is only constructed while the task allocation is
// known to be alive. Shared access is allowed for the duration of hook
// invocation.
unsafe { user_data.as_ref().as_ref()?.downcast_ref::<T>() }
}
/// Returns a mutable reference to this task's user data when the stored type
/// is `T`.
#[cfg(tokio_unstable)]
pub fn data_mut<T: Any>(&mut self) -> Option<&mut T> {
let mut user_data = self.user_data?;
// Safety: mutable `TaskMeta` is only handed to hooks for the exact task
// currently being initialized, polled under the RUNNING task state, or
// terminated. Tokio does not hold this borrow across polling the future.
unsafe { user_data.as_mut().as_mut()?.downcast_mut::<T>() }
}
/// Replaces this task's user data.
#[cfg(tokio_unstable)]
pub fn set_data<T: Any + Send + Sync + 'static>(&mut self, data: T) {
if let Some(mut user_data) = self.user_data {
// Safety: see `data_mut`.
unsafe {
*user_data.as_mut() = Some(Box::new(data));
}
}
}
/// Takes this task's user data when the stored type is `T`.
#[cfg(tokio_unstable)]
pub fn take_data<T: Any>(&mut self) -> Option<Box<T>> {
let mut user_data = self.user_data?;
// Safety: see `data_mut`.
unsafe {
let user_data = user_data.as_mut();
if !user_data.as_ref()?.is::<T>() {
return None;
}
user_data.take()?.downcast::<T>().ok()
}
}
/// Clears this task's user data, returning whether any data was present.
#[cfg(tokio_unstable)]
pub fn clear_data(&mut self) -> bool {
let Some(mut user_data) = self.user_data else {
return false;
};
// Safety: see `data_mut`.
unsafe { user_data.as_mut().take().is_some() }
}
}
/// Read-only task metadata supplied to task spawn hooks for parent tasks.
///
/// **Note**: This is an [unstable API][unstable]. The public API of this type
/// may break in 1.x releases. See [the documentation on unstable
/// features][unstable] for details.
///
/// [unstable]: crate#unstable-features
#[allow(missing_debug_implementations)]
#[cfg_attr(not(tokio_unstable), allow(unreachable_pub))]
pub struct TaskMetaRef<'a> {
/// The opaque ID of the task.
pub(crate) id: super::task::Id,
/// The location where the task was spawned.
#[cfg_attr(not(tokio_unstable), allow(unreachable_pub, dead_code))]
pub(crate) spawned_at: crate::runtime::task::SpawnLocation,
#[cfg(tokio_unstable)]
pub(crate) user_data: Option<NonNull<Option<TaskData>>>,
pub(crate) _phantom: PhantomData<&'a ()>,
}
impl<'a> TaskMetaRef<'a> {
/// # Safety
///
/// If `user_data` is present, it must point to live task storage for the
/// duration of any references exposed through this metadata value.
#[cfg(tokio_unstable)]
pub(crate) unsafe fn new(
id: super::task::Id,
spawned_at: crate::runtime::task::SpawnLocation,
#[cfg(tokio_unstable)] user_data: Option<NonNull<Option<TaskData>>>,
) -> Self {
Self {
id,
spawned_at,
#[cfg(tokio_unstable)]
user_data,
_phantom: PhantomData,
}
}
/// Return the opaque ID of the task.
#[cfg_attr(not(tokio_unstable), allow(unreachable_pub, dead_code))]
pub fn id(&self) -> super::task::Id {
self.id
}
/// Return the source code location where the task was spawned.
#[cfg(tokio_unstable)]
pub fn spawned_at(&self) -> &'static std::panic::Location<'static> {
self.spawned_at.0
}
/// Returns a shared reference to this task's user data when the stored type
/// is `T`.
#[cfg(tokio_unstable)]
pub fn data<T: Any>(&self) -> Option<&T> {
let user_data = self.user_data?;
// Safety: `TaskMetaRef` is only constructed while the task allocation is
// known to be alive, and it does not expose mutation.
unsafe { user_data.as_ref().as_ref()?.downcast_ref::<T>() }
}
}
/// Runs on specific task-related events
pub(crate) type TaskCallback = std::sync::Arc<dyn Fn(&TaskMeta<'_>) + Send + Sync>;
pub(crate) type TaskCallback = Arc<dyn Fn(&mut TaskMeta<'_>) + Send + Sync>;
pub(crate) type TaskSpawnCallback =
Arc<dyn Fn(&mut TaskMeta<'_>, Option<TaskMetaRef<'_>>) + Send + Sync>;
+17 -8
View File
@@ -6,7 +6,7 @@ use self::noop_scheduler::NoopSchedule;
use self::unowned_wrapper::unowned;
mod noop_scheduler {
use crate::runtime::task::{self, Task, TaskHarnessScheduleHooks};
use crate::runtime::task::{self, Task};
/// `task::Schedule` implementation that does nothing, for testing.
pub(crate) struct NoopSchedule;
@@ -20,11 +20,6 @@ mod noop_scheduler {
unreachable!();
}
fn hooks(&self) -> TaskHarnessScheduleHooks {
TaskHarnessScheduleHooks {
task_terminate_callback: None,
}
}
}
}
@@ -43,7 +38,14 @@ mod unowned_wrapper {
let span = tracing::trace_span!("test_span");
let task = task.instrument(span);
let (task, handle) =
crate::runtime::task::unowned(task, NoopSchedule, Id::next(), SpawnLocation::capture());
crate::runtime::task::unowned(
task,
NoopSchedule,
Id::next(),
SpawnLocation::capture(),
#[cfg(tokio_unstable)]
None,
);
(task.into_notified(), handle)
}
@@ -55,7 +57,14 @@ mod unowned_wrapper {
T::Output: Send + 'static,
{
let (task, handle) =
crate::runtime::task::unowned(task, NoopSchedule, Id::next(), SpawnLocation::capture());
crate::runtime::task::unowned(
task,
NoopSchedule,
Id::next(),
SpawnLocation::capture(),
#[cfg(tokio_unstable)]
None,
);
(task.into_notified(), handle)
}
}
+30 -11
View File
@@ -1,6 +1,5 @@
use crate::runtime::task::{
self, unowned, Id, JoinHandle, OwnedTasks, Schedule, SpawnLocation, Task,
TaskHarnessScheduleHooks,
};
use crate::runtime::tests::NoopSchedule;
@@ -61,6 +60,8 @@ fn create_drop1() {
NoopSchedule,
Id::next(),
SpawnLocation::capture(),
#[cfg(tokio_unstable)]
None,
);
drop(notified);
handle.assert_not_dropped();
@@ -79,6 +80,8 @@ fn create_drop2() {
NoopSchedule,
Id::next(),
SpawnLocation::capture(),
#[cfg(tokio_unstable)]
None,
);
drop(join);
handle.assert_not_dropped();
@@ -97,6 +100,8 @@ fn drop_abort_handle1() {
NoopSchedule,
Id::next(),
SpawnLocation::capture(),
#[cfg(tokio_unstable)]
None,
);
let abort = join.abort_handle();
drop(join);
@@ -118,6 +123,8 @@ fn drop_abort_handle2() {
NoopSchedule,
Id::next(),
SpawnLocation::capture(),
#[cfg(tokio_unstable)]
None,
);
let abort = join.abort_handle();
drop(notified);
@@ -139,6 +146,8 @@ fn drop_abort_handle_clone() {
NoopSchedule,
Id::next(),
SpawnLocation::capture(),
#[cfg(tokio_unstable)]
None,
);
let abort = join.abort_handle();
let abort_clone = abort.clone();
@@ -164,6 +173,8 @@ fn create_shutdown1() {
NoopSchedule,
Id::next(),
SpawnLocation::capture(),
#[cfg(tokio_unstable)]
None,
);
drop(join);
handle.assert_not_dropped();
@@ -182,6 +193,8 @@ fn create_shutdown2() {
NoopSchedule,
Id::next(),
SpawnLocation::capture(),
#[cfg(tokio_unstable)]
None,
);
handle.assert_not_dropped();
notified.shutdown();
@@ -191,7 +204,14 @@ fn create_shutdown2() {
#[test]
fn unowned_poll() {
let (task, _) = unowned(async {}, NoopSchedule, Id::next(), SpawnLocation::capture());
let (task, _) = unowned(
async {},
NoopSchedule,
Id::next(),
SpawnLocation::capture(),
#[cfg(tokio_unstable)]
None,
);
task.run();
}
@@ -402,9 +422,14 @@ impl Runtime {
T::Output: 'static + Send,
{
let (handle, notified) =
self.0
.owned
.bind(future, self.clone(), Id::next(), SpawnLocation::capture());
self.0.owned.bind(
future,
self.clone(),
Id::next(),
SpawnLocation::capture(),
#[cfg(tokio_unstable)]
None,
);
if let Some(notified) = notified {
self.schedule(notified);
@@ -460,10 +485,4 @@ impl Schedule for Runtime {
fn schedule(&self, task: task::Notified<Self>) {
self.0.core.try_lock().unwrap().queue.push_back(task);
}
fn hooks(&self) -> TaskHarnessScheduleHooks {
TaskHarnessScheduleHooks {
task_terminate_callback: None,
}
}
}
+48 -15
View File
@@ -4,7 +4,7 @@ use crate::{
task::{JoinHandle, LocalSet},
util::trace::SpawnMeta,
};
use std::{future::Future, io, mem};
use std::{any::Any, fmt, future::Future, io, mem};
/// Factory which is used to configure the properties of a new task.
///
@@ -14,10 +14,11 @@ use std::{future::Future, io, mem};
///
/// Methods can be chained in order to configure it.
///
/// Currently, there is only one configuration option:
/// Configuration options include:
///
/// - [`name`], which specifies an associated name for
/// the task
/// - [`data`], which stores user data for runtime task hooks
///
/// There are three types of task that can be spawned from a Builder:
/// - [`spawn_local`] for executing not [`Send`] futures
@@ -55,13 +56,15 @@ use std::{future::Future, io, mem};
/// ```
/// [unstable]: crate#unstable-features
/// [`name`]: Builder::name
/// [`data`]: Builder::data
/// [`spawn_local`]: Builder::spawn_local
/// [`spawn`]: Builder::spawn
/// [`spawn_blocking`]: Builder::spawn_blocking
#[derive(Default, Debug)]
#[derive(Default)]
#[cfg_attr(docsrs, doc(cfg(all(tokio_unstable, feature = "tracing"))))]
pub struct Builder<'a> {
name: Option<&'a str>,
data: Option<crate::runtime::TaskData>,
}
impl<'a> Builder<'a> {
@@ -71,8 +74,22 @@ impl<'a> Builder<'a> {
}
/// Assigns a name to the task which will be spawned.
pub fn name(&self, name: &'a str) -> Self {
Self { name: Some(name) }
pub fn name(self, name: &'a str) -> Self {
Self {
name: Some(name),
..self
}
}
/// Sets task data visible to runtime task hooks.
pub fn data<T>(self, data: T) -> Self
where
T: Any + Send + Sync + 'static,
{
Self {
data: Some(Box::new(data)),
..self
}
}
/// Spawns a task with this builder's settings on the current runtime.
@@ -89,11 +106,12 @@ impl<'a> Builder<'a> {
Fut: Future + Send + 'static,
Fut::Output: Send + 'static,
{
let Builder { name, data } = self;
let fut_size = mem::size_of::<Fut>();
Ok(if fut_size > BOX_FUTURE_THRESHOLD {
super::spawn::spawn_inner(Box::pin(future), SpawnMeta::new(self.name, fut_size))
super::spawn::spawn_inner(Box::pin(future), SpawnMeta::new(name, fut_size), data)
} else {
super::spawn::spawn_inner(future, SpawnMeta::new(self.name, fut_size))
super::spawn::spawn_inner(future, SpawnMeta::new(name, fut_size), data)
})
}
@@ -110,11 +128,12 @@ impl<'a> Builder<'a> {
Fut: Future + Send + 'static,
Fut::Output: Send + 'static,
{
let Builder { name, data } = self;
let fut_size = mem::size_of::<Fut>();
Ok(if fut_size > BOX_FUTURE_THRESHOLD {
handle.spawn_named(Box::pin(future), SpawnMeta::new(self.name, fut_size))
handle.spawn_named_with_data(Box::pin(future), SpawnMeta::new(name, fut_size), data)
} else {
handle.spawn_named(future, SpawnMeta::new(self.name, fut_size))
handle.spawn_named_with_data(future, SpawnMeta::new(name, fut_size), data)
})
}
@@ -141,11 +160,12 @@ impl<'a> Builder<'a> {
Fut: Future + 'static,
Fut::Output: 'static,
{
let Builder { name, data } = self;
let fut_size = mem::size_of::<Fut>();
Ok(if fut_size > BOX_FUTURE_THRESHOLD {
super::local::spawn_local_inner(Box::pin(future), SpawnMeta::new(self.name, fut_size))
super::local::spawn_local_inner(Box::pin(future), SpawnMeta::new(name, fut_size), data)
} else {
super::local::spawn_local_inner(future, SpawnMeta::new(self.name, fut_size))
super::local::spawn_local_inner(future, SpawnMeta::new(name, fut_size), data)
})
}
@@ -166,11 +186,12 @@ impl<'a> Builder<'a> {
Fut: Future + 'static,
Fut::Output: 'static,
{
let Builder { name, data } = self;
let fut_size = mem::size_of::<Fut>();
Ok(if fut_size > BOX_FUTURE_THRESHOLD {
local_set.spawn_named(Box::pin(future), SpawnMeta::new(self.name, fut_size))
local_set.spawn_named_with_data(Box::pin(future), SpawnMeta::new(name, fut_size), data)
} else {
local_set.spawn_named(future, SpawnMeta::new(self.name, fut_size))
local_set.spawn_named_with_data(future, SpawnMeta::new(name, fut_size), data)
})
}
@@ -212,20 +233,23 @@ impl<'a> Builder<'a> {
Output: Send + 'static,
{
use crate::runtime::Mandatory;
let Builder { name, data } = self;
let fn_size = mem::size_of::<Function>();
let (join_handle, spawn_result) = if fn_size > BOX_FUTURE_THRESHOLD {
handle.inner.blocking_spawner().spawn_blocking_inner(
Box::new(function),
Mandatory::NonMandatory,
SpawnMeta::new(self.name, fn_size),
SpawnMeta::new(name, fn_size),
handle,
data,
)
} else {
handle.inner.blocking_spawner().spawn_blocking_inner(
function,
Mandatory::NonMandatory,
SpawnMeta::new(self.name, fn_size),
SpawnMeta::new(name, fn_size),
handle,
data,
)
};
@@ -233,3 +257,12 @@ impl<'a> Builder<'a> {
Ok(join_handle)
}
}
impl fmt::Debug for Builder<'_> {
fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt.debug_struct("Builder")
.field("name", &self.name)
.field("data", &self.data.as_ref().map(|_| "<task data>"))
.finish()
}
}
+9
View File
@@ -705,6 +705,15 @@ impl<'a, T: 'static> Builder<'a, T> {
Self { builder, ..self }
}
/// Sets task data visible to runtime task hooks.
pub fn data<D>(self, data: D) -> Self
where
D: std::any::Any + Send + Sync + 'static,
{
let builder = self.builder.data(data);
Self { builder, ..self }
}
/// Spawn the provided task with this builder's settings and store it in the
/// [`JoinSet`], returning an [`AbortHandle`] that can be used to remotely
/// cancel the task.
+82 -20
View File
@@ -2,9 +2,7 @@
use crate::loom::cell::UnsafeCell;
use crate::loom::sync::{Arc, Mutex};
use crate::runtime;
use crate::runtime::task::{
self, JoinHandle, LocalOwnedTasks, SpawnLocation, Task, TaskHarnessScheduleHooks,
};
use crate::runtime::task::{self, JoinHandle, LocalOwnedTasks, SpawnLocation, Task};
use crate::runtime::{context, ThreadId, BOX_FUTURE_THRESHOLD};
use crate::sync::AtomicWaker;
use crate::util::trace::SpawnMeta;
@@ -399,21 +397,37 @@ cfg_rt! {
{
let fut_size = std::mem::size_of::<F>();
if fut_size > BOX_FUTURE_THRESHOLD {
spawn_local_inner(Box::pin(future), SpawnMeta::new_unnamed(fut_size))
spawn_local_inner(
Box::pin(future),
SpawnMeta::new_unnamed(fut_size),
#[cfg(tokio_unstable)]
None,
)
} else {
spawn_local_inner(future, SpawnMeta::new_unnamed(fut_size))
spawn_local_inner(
future,
SpawnMeta::new_unnamed(fut_size),
#[cfg(tokio_unstable)]
None,
)
}
}
#[track_caller]
pub(super) fn spawn_local_inner<F>(future: F, meta: SpawnMeta<'_>) -> JoinHandle<F::Output>
pub(super) fn spawn_local_inner<F>(
future: F,
meta: SpawnMeta<'_>,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where F: Future + 'static,
F::Output: 'static
{
use crate::runtime::{context, task};
let mut future = Some(future);
#[cfg(tokio_unstable)]
let mut user_data = Some(user_data);
let res = context::with_current(|handle| {
Some(if handle.is_local() {
@@ -439,11 +453,24 @@ cfg_rt! {
let task = crate::util::trace::task(future, "task", meta, id.as_u64());
// safety: we have verified that this is a `LocalRuntime` owned by the current thread
unsafe { handle.spawn_local(task, id, meta.spawned_at) }
unsafe {
handle.spawn_local(
task,
id,
meta.spawned_at,
#[cfg(tokio_unstable)]
user_data.take().unwrap(),
)
}
} else {
match CURRENT.with(|LocalData { ctx, .. }| ctx.get()) {
None => panic!("`spawn_local` called from outside of a `task::LocalSet` or `runtime::LocalRuntime`"),
Some(cx) => cx.spawn(future.take().unwrap(), meta)
Some(cx) => cx.spawn(
future.take().unwrap(),
meta,
#[cfg(tokio_unstable)]
user_data.take().unwrap(),
)
}
})
});
@@ -453,7 +480,12 @@ cfg_rt! {
Ok(Some(join_handle)) => join_handle,
Err(_) => match CURRENT.with(|LocalData { ctx, .. }| ctx.get()) {
None => panic!("`spawn_local` called from outside of a `task::LocalSet` or `runtime::LocalRuntime`"),
Some(cx) => cx.spawn(future.unwrap(), meta)
Some(cx) => cx.spawn(
future.unwrap(),
meta,
#[cfg(tokio_unstable)]
user_data.unwrap(),
)
}
}
}
@@ -729,16 +761,46 @@ impl LocalSet {
F: Future + 'static,
F::Output: 'static,
{
self.spawn_named_inner(future, meta)
self.spawn_named_inner(
future,
meta,
#[cfg(tokio_unstable)]
None,
)
}
#[cfg(all(tokio_unstable, feature = "tracing"))]
#[track_caller]
fn spawn_named_inner<F>(&self, future: F, meta: SpawnMeta<'_>) -> JoinHandle<F::Output>
pub(in crate::task) fn spawn_named_with_data<F>(
&self,
future: F,
meta: SpawnMeta<'_>,
user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where
F: Future + 'static,
F::Output: 'static,
{
let handle = self.context.spawn(future, meta);
self.spawn_named_inner(future, meta, user_data)
}
#[track_caller]
fn spawn_named_inner<F>(
&self,
future: F,
meta: SpawnMeta<'_>,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where
F: Future + 'static,
F::Output: 'static,
{
let handle = self.context.spawn(
future,
meta,
#[cfg(tokio_unstable)]
user_data,
);
// Because a task was spawned from *outside* the `LocalSet`, wake the
// `LocalSet` future to execute the new task, if it hasn't been woken.
@@ -1024,7 +1086,12 @@ impl Drop for LocalSet {
impl Context {
#[track_caller]
fn spawn<F>(&self, future: F, meta: SpawnMeta<'_>) -> JoinHandle<F::Output>
fn spawn<F>(
&self,
future: F,
meta: SpawnMeta<'_>,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<F::Output>
where
F: Future + 'static,
F::Output: 'static,
@@ -1040,6 +1107,8 @@ impl Context {
self.shared.clone(),
id,
SpawnLocation::capture(),
#[cfg(tokio_unstable)]
user_data,
)
};
@@ -1148,13 +1217,6 @@ impl task::Schedule for Arc<Shared> {
Shared::schedule(self, task);
}
// localset does not currently support task hooks
fn hooks(&self) -> TaskHarnessScheduleHooks {
TaskHarnessScheduleHooks {
task_terminate_callback: None,
}
}
cfg_unstable! {
fn unhandled_panic(&self) {
use crate::runtime::UnhandledPanic;
+26 -4
View File
@@ -178,14 +178,28 @@ cfg_rt! {
{
let fut_size = std::mem::size_of::<F>();
if fut_size > BOX_FUTURE_THRESHOLD {
spawn_inner(Box::pin(future), SpawnMeta::new_unnamed(fut_size))
spawn_inner(
Box::pin(future),
SpawnMeta::new_unnamed(fut_size),
#[cfg(tokio_unstable)]
None,
)
} else {
spawn_inner(future, SpawnMeta::new_unnamed(fut_size))
spawn_inner(
future,
SpawnMeta::new_unnamed(fut_size),
#[cfg(tokio_unstable)]
None,
)
}
}
#[track_caller]
pub(super) fn spawn_inner<T>(future: T, meta: SpawnMeta<'_>) -> JoinHandle<T::Output>
pub(super) fn spawn_inner<T>(
future: T,
meta: SpawnMeta<'_>,
#[cfg(tokio_unstable)] user_data: Option<crate::runtime::TaskData>,
) -> JoinHandle<T::Output>
where
T: Future + Send + 'static,
T::Output: Send + 'static,
@@ -207,7 +221,15 @@ cfg_rt! {
let id = task::Id::next();
let task = crate::util::trace::task(future, "task", meta, id.as_u64());
match context::with_current(|handle| handle.spawn(task, id, meta.spawned_at)) {
match context::with_current(|handle| {
handle.spawn(
task,
id,
meta.spawned_at,
#[cfg(tokio_unstable)]
user_data,
)
}) {
Ok(join_handle) => join_handle,
Err(e) => panic!("{}", e),
}
+432 -18
View File
@@ -2,8 +2,12 @@
#![cfg(all(feature = "full", tokio_unstable, target_has_atomic = "64"))]
use std::collections::HashSet;
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use std::sync::{Arc, Condvar, Mutex};
use std::task::{Context, Poll};
use std::time::Duration;
use tokio::runtime::Builder;
@@ -19,7 +23,8 @@ fn spawn_task_hook_fires() {
let ids2 = Arc::clone(&ids);
let runtime = Builder::new_current_thread()
.on_task_spawn(move |data| {
.on_task_spawn(move |data, _parent| {
assert!(data.data::<usize>().is_none());
ids2.lock().unwrap().insert(data.id());
count2.fetch_add(1, Ordering::SeqCst);
@@ -85,11 +90,11 @@ fn task_hook_spawn_location_current_thread() {
"(current_thread) on_task_spawn",
&spawns,
))
.on_before_task_poll(mk_spawn_location_hook(
.on_before_task_poll(mk_poll_location_hook(
"(current_thread) on_before_task_poll",
&poll_starts,
))
.on_after_task_poll(mk_spawn_location_hook(
.on_after_task_poll(mk_poll_location_hook(
"(current_thread) on_after_task_poll",
&poll_ends,
))
@@ -136,11 +141,11 @@ fn task_hook_spawn_location_multi_thread() {
"(multi_thread) on_task_spawn",
&spawns,
))
.on_before_task_poll(mk_spawn_location_hook(
.on_before_task_poll(mk_poll_location_hook(
"(multi_thread) on_before_task_poll",
&poll_starts,
))
.on_after_task_poll(mk_spawn_location_hook(
.on_after_task_poll(mk_poll_location_hook(
"(multi_thread) on_after_task_poll",
&poll_ends,
))
@@ -174,21 +179,430 @@ fn task_hook_spawn_location_multi_thread() {
assert_eq!(poll_starts, poll_ends.fetch_add(0, Ordering::SeqCst));
}
#[derive(Debug)]
struct PollState {
before: usize,
after: usize,
}
#[test]
fn task_data_mutates_across_current_thread_hooks() {
let terminated = Arc::new(Mutex::new(Vec::new()));
let terminated2 = Arc::clone(&terminated);
let runtime = Builder::new_current_thread()
.on_task_spawn(|meta, parent| {
assert!(parent.is_none());
assert!(!meta.clear_data());
meta.set_data(PollState {
before: 0,
after: 0,
});
})
.on_before_task_poll(|meta| {
meta.data_mut::<PollState>().unwrap().before += 1;
})
.on_after_task_poll(|meta| {
if let Some(data) = meta.data_mut::<PollState>() {
data.after += 1;
}
})
.on_task_terminate(move |meta| {
let data = meta.take_data::<PollState>().unwrap();
assert!(meta.data::<PollState>().is_none());
terminated2.lock().unwrap().push((data.before, data.after));
})
.build()
.unwrap();
runtime.block_on(async {
tokio::spawn(async {
for _ in 0..3 {
tokio::task::yield_now().await;
}
})
.await
.unwrap();
});
let terminated = terminated.lock().unwrap();
assert_eq!(terminated.len(), 1);
let (before, after) = terminated[0];
assert!(before > 1);
assert_eq!(before, after);
}
#[cfg_attr(
target_os = "wasi",
ignore = "WASI does not support multi-threaded runtime"
)]
#[test]
fn task_data_mutates_across_multi_thread_hooks() {
let terminated = Arc::new(Mutex::new(Vec::new()));
let terminated2 = Arc::clone(&terminated);
let runtime = Builder::new_multi_thread()
.worker_threads(2)
.on_task_spawn(|meta, _parent| {
meta.set_data(PollState {
before: 0,
after: 0,
});
})
.on_before_task_poll(|meta| {
meta.data_mut::<PollState>().unwrap().before += 1;
})
.on_after_task_poll(|meta| {
if let Some(data) = meta.data_mut::<PollState>() {
data.after += 1;
}
})
.on_task_terminate(move |meta| {
let data = meta.take_data::<PollState>().unwrap();
terminated2.lock().unwrap().push((data.before, data.after));
})
.build()
.unwrap();
runtime.block_on(async {
tokio::spawn(async {
for _ in 0..3 {
tokio::task::yield_now().await;
}
})
.await
.unwrap();
});
runtime.shutdown_timeout(std::time::Duration::from_secs(60));
let terminated = terminated.lock().unwrap();
assert_eq!(terminated.len(), 1);
let (before, after) = terminated[0];
assert!(before > 1);
assert_eq!(before, after);
}
#[cfg_attr(
target_os = "wasi",
ignore = "WASI does not support multi-threaded runtime"
)]
#[test]
fn poll_hook_data_is_not_terminated_during_multi_thread_shutdown() {
let entered = Arc::new((Mutex::new(false), Condvar::new()));
let entered2 = Arc::clone(&entered);
let release = Arc::new((Mutex::new(false), Condvar::new()));
let release2 = Arc::clone(&release);
let terminated = Arc::new(Mutex::new(Vec::new()));
let terminated2 = Arc::clone(&terminated);
let runtime = Builder::new_multi_thread()
.worker_threads(2)
.on_task_spawn(|meta, _parent| {
meta.set_data(Vec::<&'static str>::new());
})
.on_before_task_poll(move |meta| {
let (lock, cvar) = &*entered2;
*lock.lock().unwrap() = true;
cvar.notify_one();
let (lock, cvar) = &*release2;
let mut released = lock.lock().unwrap();
while !*released {
released = cvar.wait(released).unwrap();
}
meta.data_mut::<Vec<&'static str>>()
.unwrap()
.push("before_done");
})
.on_task_terminate(move |meta| {
let data = meta.take_data::<Vec<&'static str>>().unwrap();
terminated2.lock().unwrap().push(*data);
})
.build()
.unwrap();
drop(runtime.spawn(std::future::pending::<()>()));
let (lock, cvar) = &*entered;
let entered_guard = lock.lock().unwrap();
let (entered_guard, wait_result) = cvar
.wait_timeout_while(entered_guard, Duration::from_secs(5), |entered| !*entered)
.unwrap();
assert!(*entered_guard);
assert!(!wait_result.timed_out());
let shutdown = std::thread::spawn(move || {
runtime.shutdown_timeout(Duration::from_secs(5));
});
std::thread::sleep(Duration::from_millis(50));
let (lock, cvar) = &*release;
*lock.lock().unwrap() = true;
cvar.notify_one();
shutdown.join().unwrap();
assert_eq!(*terminated.lock().unwrap(), vec![vec!["before_done"]]);
}
#[cfg_attr(
target_os = "wasi",
ignore = "WASI does not support multi-threaded runtime"
)]
#[test]
fn abort_during_before_poll_hook_does_not_poll_future() {
struct CountPolls {
polls: Arc<AtomicUsize>,
}
impl Future for CountPolls {
type Output = ();
fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<()> {
self.polls.fetch_add(1, Ordering::SeqCst);
Poll::Pending
}
}
let entered = Arc::new((Mutex::new(false), Condvar::new()));
let entered2 = Arc::clone(&entered);
let release = Arc::new((Mutex::new(false), Condvar::new()));
let release2 = Arc::clone(&release);
let polls = Arc::new(AtomicUsize::new(0));
let runtime = Builder::new_multi_thread()
.worker_threads(2)
.on_before_task_poll(move |_meta| {
let (lock, cvar) = &*entered2;
*lock.lock().unwrap() = true;
cvar.notify_one();
let (lock, cvar) = &*release2;
let mut released = lock.lock().unwrap();
while !*released {
released = cvar.wait(released).unwrap();
}
})
.build()
.unwrap();
let task = runtime.spawn(CountPolls {
polls: Arc::clone(&polls),
});
let (lock, cvar) = &*entered;
let entered_guard = lock.lock().unwrap();
let (entered_guard, wait_result) = cvar
.wait_timeout_while(entered_guard, Duration::from_secs(5), |entered| !*entered)
.unwrap();
assert!(*entered_guard);
assert!(!wait_result.timed_out());
task.abort();
let (lock, cvar) = &*release;
*lock.lock().unwrap() = true;
cvar.notify_one();
let err = runtime.block_on(task).unwrap_err();
assert!(err.is_cancelled());
assert_eq!(polls.load(Ordering::SeqCst), 0);
}
#[derive(Clone, Debug, Eq, PartialEq)]
struct Lineage {
depth: usize,
}
#[test]
fn spawn_hook_can_inherit_parent_task_data() {
let terminated = Arc::new(Mutex::new(Vec::new()));
let terminated2 = Arc::clone(&terminated);
let runtime = Builder::new_current_thread()
.on_task_spawn(|meta, parent| {
let depth = match parent {
Some(parent) => parent
.data::<Lineage>()
.map_or(0, |parent| parent.depth + 1),
None => 0,
};
meta.set_data(Lineage { depth });
})
.on_task_terminate(move |meta| {
let data = meta.take_data::<Lineage>().unwrap();
terminated2.lock().unwrap().push(data.depth);
})
.build()
.unwrap();
runtime.block_on(async {
tokio::spawn(async {
tokio::spawn(async {}).await.unwrap();
})
.await
.unwrap();
});
let mut terminated = terminated.lock().unwrap().clone();
terminated.sort_unstable();
assert_eq!(terminated, vec![0, 1]);
}
#[test]
fn spawn_hook_runs_before_terminate_when_current_thread_runtime_is_closed() {
struct ClosedSpawnData;
let events = Arc::new(Mutex::new(Vec::new()));
let events2 = Arc::clone(&events);
let events3 = Arc::clone(&events);
let runtime = Builder::new_current_thread()
.on_task_spawn(move |meta, _parent| {
meta.set_data(ClosedSpawnData);
events2.lock().unwrap().push("spawn");
})
.on_task_terminate(move |meta| {
if meta.take_data::<ClosedSpawnData>().is_some() {
events3.lock().unwrap().push("terminate_with_data");
}
})
.build()
.unwrap();
let handle = runtime.handle().clone();
drop(runtime);
drop(handle.spawn(async {}));
assert_eq!(*events.lock().unwrap(), ["spawn", "terminate_with_data"]);
}
#[cfg_attr(
target_os = "wasi",
ignore = "WASI does not support multi-threaded runtime"
)]
#[test]
fn spawn_hook_runs_before_terminate_when_multi_thread_runtime_is_closed() {
struct ClosedSpawnData;
let events = Arc::new(Mutex::new(Vec::new()));
let events2 = Arc::clone(&events);
let events3 = Arc::clone(&events);
let runtime = Builder::new_multi_thread()
.worker_threads(1)
.on_task_spawn(move |meta, _parent| {
meta.set_data(ClosedSpawnData);
events2.lock().unwrap().push("spawn");
})
.on_task_terminate(move |meta| {
if meta.take_data::<ClosedSpawnData>().is_some() {
events3.lock().unwrap().push("terminate_with_data");
}
})
.build()
.unwrap();
let handle = runtime.handle().clone();
drop(runtime);
drop(handle.spawn(async {}));
assert_eq!(*events.lock().unwrap(), ["spawn", "terminate_with_data"]);
}
#[cfg(feature = "tracing")]
#[test]
fn task_builder_data_is_visible_to_hooks() {
let terminated = Arc::new(Mutex::new(Vec::new()));
let terminated2 = Arc::clone(&terminated);
let runtime = Builder::new_current_thread()
.on_task_spawn(|meta, _parent| {
let value = meta.data_mut::<usize>().unwrap();
*value += 1;
})
.on_task_terminate(move |meta| {
let value = meta.take_data::<usize>().unwrap();
terminated2.lock().unwrap().push(*value);
})
.build()
.unwrap();
runtime.block_on(async {
tokio::task::Builder::new()
.data(41usize)
.spawn(async {})
.unwrap()
.await
.unwrap();
});
assert_eq!(*terminated.lock().unwrap(), vec![42]);
}
#[cfg(feature = "tracing")]
#[test]
fn task_builder_data_is_not_dropped_for_spawn_blocking() {
let terminated = Arc::new(Mutex::new(Vec::new()));
let terminated2 = Arc::clone(&terminated);
let runtime = Builder::new_current_thread()
.on_task_terminate(move |meta| {
if let Some(value) = meta.take_data::<usize>() {
terminated2.lock().unwrap().push(*value);
}
})
.build()
.unwrap();
runtime.block_on(async {
tokio::task::Builder::new()
.data(7usize)
.spawn_blocking(|| {})
.unwrap()
.await
.unwrap();
});
assert_eq!(*terminated.lock().unwrap(), vec![7]);
}
fn mk_spawn_location_hook(
event: &'static str,
count: &Arc<AtomicUsize>,
) -> impl Fn(&tokio::runtime::TaskMeta<'_>) {
) -> impl Fn(&mut tokio::runtime::TaskMeta<'_>, Option<tokio::runtime::TaskMetaRef<'_>>) {
let count = Arc::clone(count);
move |data| {
eprintln!("{event} ({:?}): {:?}", data.id(), data.spawned_at());
// Assert that the spawn location is in this file.
// Don't make assertions about line number/column here, as these
// may change as new code is added to the test file...
assert_eq!(
data.spawned_at().file(),
file!(),
"incorrect spawn location in {event} hook",
);
count.fetch_add(1, Ordering::SeqCst);
move |data, _parent| {
assert_spawn_location(event, count.as_ref(), data);
}
}
fn mk_poll_location_hook(
event: &'static str,
count: &Arc<AtomicUsize>,
) -> impl Fn(&mut tokio::runtime::TaskMeta<'_>) {
let count = Arc::clone(count);
move |data| {
assert_spawn_location(event, count.as_ref(), data);
}
}
fn assert_spawn_location(
event: &'static str,
count: &AtomicUsize,
data: &tokio::runtime::TaskMeta<'_>,
) {
eprintln!("{event} ({:?}): {:?}", data.id(), data.spawned_at());
assert_eq!(
data.spawned_at().file(),
file!(),
"incorrect spawn location in {event} hook",
);
count.fetch_add(1, Ordering::SeqCst);
}