taskdump: allow impl FnMut() in taskdumps instead of just fn() (#8040)

This commit is contained in:
Alice Ryhl
2026-04-14 19:18:56 +02:00
committed by GitHub
parent f943312865
commit 36d12d2686
3 changed files with 141 additions and 136 deletions
+86 -43
View File
@@ -34,7 +34,8 @@ pub(crate) struct Context {
/// The function that is invoked at each leaf future inside of Tokio
///
/// For example, within tokio::time:sleep, sockets. etc.
trace_leaf_fn: Cell<Option<fn(&TraceMeta)>>,
#[allow(clippy::type_complexity)]
trace_leaf_fn: Cell<Option<NonNull<dyn FnMut(&TraceMeta)>>>,
}
/// A [`Frame`] in an intrusive, doubly-linked tree of [`Frame`]s.
@@ -114,19 +115,39 @@ impl Context {
}
}
/// Calls the provided closure if we are being traced.
fn try_with_current_trace_leaf_fn<F, R>(f: F) -> Option<R>
where
F: FnOnce(&Cell<Option<fn(&TraceMeta)>>) -> R,
F: for<'a> FnOnce(&'a mut dyn FnMut(&TraceMeta)) -> R,
{
// SAFETY: This call can only access the trace_leaf_fn field, so it cannot
// break the trace frame linked list.
unsafe { Self::try_with_current(|context| f(&context.trace_leaf_fn)) }
let mut ret = None;
let inner = |context: &Context| {
if let Some(mut trace_leaf_fn) = context.trace_leaf_fn.replace(None) {
let _restore = defer(move || {
context.trace_leaf_fn.set(Some(trace_leaf_fn));
});
// SAFETY: The trace leaf fn is valid for the duration in which it's stored in the
// context. Furthermore, re-entrant calls are not possible because we store `None` for
// the duration in which we hold a mutable reference, so access is exclusive for that
// duration.
ret = Some(f(unsafe { trace_leaf_fn.as_mut() }));
}
};
// SAFETY: This call can only access the trace_leaf_fn field, so it cannot break the trace
// frame linked list.
unsafe { Self::try_with_current(inner) };
ret
}
/// Produces `true` if the current task is being traced; otherwise false.
pub(crate) fn is_tracing() -> bool {
Self::try_with_current_trace_leaf_fn(|maybe_trace_leaf| maybe_trace_leaf.get().is_some())
.unwrap_or(false)
// SAFETY: This call can only access the trace_leaf_fn field, so it cannot break the trace
// frame linked list.
unsafe { Self::try_with_current(|ctx| ctx.trace_leaf_fn.get().is_some()).unwrap_or(false) }
}
}
@@ -149,13 +170,9 @@ pub struct TraceMeta {
/// Runs `f`. If `f` hits a Tokio yield point `trace_leaf` will be invoked.
///
/// This allows taking a task dump with caller-provided task dump machinery. If `f` is the poll function of a future
/// and that future returns `Poll::Pending`, then `trace_leaf` will be invoked. `trace_leaf` can then take a backtrace
/// to determine exactly where the yield occurred.
///
/// `trace_leaf` is a function pointer (`fn`) rather than a closure (`Fn`) because it must be stored
/// in thread-local state via a `Cell`. Use thread-locals to communicate between the callback and
/// calling code (see example below).
/// This allows taking a task dump with caller-provided task dump machinery. If `f` is the poll
/// function of a future and that future returns `Poll::Pending`, then `trace_leaf` will be
/// invoked. `trace_leaf` can then take a backtrace to determine exactly where the yield occurred.
///
/// # Example
///
@@ -164,45 +181,69 @@ pub struct TraceMeta {
/// use std::task::Poll;
/// use tokio::runtime::dump::{trace_with, Trace, TraceMeta};
///
/// // Thread-local storage for the custom trace function.
/// std::thread_local! {
/// static LEAF_COUNT: std::cell::Cell<u32> = const { std::cell::Cell::new(0) };
/// fn my_trace_leaf(_meta: &TraceMeta, count: &mut u32) {
/// *count += 1;
/// }
///
/// fn my_trace_leaf(_meta: &TraceMeta) {
/// LEAF_COUNT.with(|c| c.set(c.get() + 1));
/// }
///
/// # async fn example() {
/// # #[tokio::main(flavor = "current_thread")]
/// # async fn main() {
/// let mut fut = std::pin::pin!(async {
/// tokio::task::yield_now().await;
/// });
///
/// LEAF_COUNT.with(|c| c.set(0));
/// let mut leaf_count = 0;
///
/// Trace::root(std::future::poll_fn(|cx| {
/// trace_with(|| { let _ = fut.as_mut().poll(cx); }, my_trace_leaf);
/// trace_with(
/// || { let _ = fut.as_mut().poll(cx); },
/// |meta| my_trace_leaf(meta, &mut leaf_count),
/// );
/// Poll::Ready(())
/// })).await;
///
/// let count = LEAF_COUNT.with(|c| c.get());
/// assert!(count > 0);
/// assert!(leaf_count > 0);
/// # }
/// ```
pub fn trace_with<F, R>(f: F, trace_leaf: fn(&TraceMeta)) -> R
pub fn trace_with<FN, FT, R>(f: FN, mut trace_leaf: FT) -> R
where
F: FnOnce() -> R,
FN: FnOnce() -> R,
FT: FnMut(&TraceMeta),
{
// store our new trace_leaf function
let previous =
Context::try_with_current_trace_leaf_fn(|current| current.replace(Some(trace_leaf)));
let trace_leaf_dyn = (&mut trace_leaf) as &mut (dyn FnMut(&TraceMeta) + '_);
// SAFETY: The raw pointer is removed from the thread local before `trace_leaf` is dropped, so
// this transmute cannot lead to the violation of any lifetime requirements.
let trace_leaf_dyn = unsafe {
std::mem::transmute::<
*mut (dyn FnMut(&TraceMeta) + '_),
*mut (dyn FnMut(&TraceMeta) + 'static),
>(trace_leaf_dyn)
};
// SAFETY: Pointer comes from reference, so not null.
let trace_leaf_dyn = unsafe { NonNull::new_unchecked(trace_leaf_dyn) };
let mut old_trace_leaf_fn = None;
// Even if this access fails, that's okay. In that case, we still call the closure without
// actually performing any tracing.
//
// SAFETY: This call can only access the trace_leaf_fn field, so it cannot break the trace
// frame linked list.
unsafe {
Context::try_with_current(|ctx| {
old_trace_leaf_fn = ctx.trace_leaf_fn.replace(Some(trace_leaf_dyn));
})
};
// restore previous on drop. This is ensures state remains consistent
// even if the trace_leaf function panics
let _restore = defer(move || {
if let Some(previous) = previous {
Context::try_with_current_trace_leaf_fn(|current| current.set(previous));
}
// This ensures that `trace_leaf_fn` cannot be accessed after this call returns.
//
// SAFETY: This call can only access the trace_leaf_fn field, so it cannot
// break the trace frame linked list.
unsafe {
Context::try_with_current(|ctx| {
ctx.trace_leaf_fn.set(old_trace_leaf_fn);
})
};
});
f()
@@ -249,13 +290,14 @@ impl Trace {
// internal implementation details of this crate).
#[inline(never)]
pub(crate) fn trace_leaf(cx: &mut task::Context<'_>) -> Poll<()> {
let trace_leaf_fn = Context::try_with_current_trace_leaf_fn(|cell| cell.get()).flatten();
if let Some(trace_leaf_fn) = trace_leaf_fn {
let root_addr = Context::current_frame_addr();
let ret = Context::try_with_current_trace_leaf_fn(|leaf_fn| {
let meta = TraceMeta {
root_addr: Context::current_frame_addr(),
root_addr,
trace_leaf_addr: trace_leaf as *const c_void,
};
trace_leaf_fn(&meta);
leaf_fn(&meta);
// Use the same logic that `yield_now` uses to send out wakeups after
// the task yields.
@@ -268,10 +310,11 @@ pub(crate) fn trace_leaf(cx: &mut task::Context<'_>) -> Poll<()> {
}
}
});
});
Poll::Pending
} else {
Poll::Ready(())
match ret {
Some(()) => Poll::Pending,
None => Poll::Ready(()),
}
}
+22 -53
View File
@@ -2,73 +2,42 @@
//!
//! This implementation may eventually be extracted into a separate `tokio-taskdump` crate.
use std::{cell::Cell, ptr};
use std::ptr;
use crate::runtime::task::trace::{trace_with, Trace, TraceMeta};
use super::defer;
/// Thread local state used to communicate between calling the trace and the interior `trace_leaf` function
struct TraceContext {
collector: Cell<Option<Trace>>,
}
thread_local! {
static TRACE_CONTEXT: TraceContext = const {
TraceContext {
collector: Cell::new(None),
}
};
}
/// Capture using the default `backtrace::trace`-based implementation.
#[inline(never)]
pub(super) fn capture<F, R>(f: F) -> (R, Trace)
where
F: FnOnce() -> R,
{
let collector = Trace::empty();
let mut trace = Trace::empty();
let previous = TRACE_CONTEXT.with(|state| state.collector.replace(Some(collector)));
let result = trace_with(f, |meta| trace_leaf(meta, &mut trace));
// restore previous collector on drop even if the callback panics
let _restore = defer(move || {
TRACE_CONTEXT.with(|state| state.collector.set(previous));
});
let result = trace_with(f, trace_leaf);
// take the collector before _restore runs
let collector = TRACE_CONTEXT.with(|state| state.collector.take()).unwrap();
(result, collector)
(result, trace)
}
/// Capture a backtrace via `backtrace::trace` and collect it into `STATE`
#[inline(never)]
pub(crate) fn trace_leaf(meta: &TraceMeta) {
TRACE_CONTEXT.with(|state| {
if let Some(mut collector) = state.collector.take() {
let mut frames: Vec<backtrace::BacktraceFrame> = vec![];
let mut above_leaf = false;
/// Capture a backtrace via `backtrace::trace` and collect it into `trace`.
pub(crate) fn trace_leaf(meta: &TraceMeta, trace: &mut Trace) {
let mut frames: Vec<backtrace::BacktraceFrame> = vec![];
let mut above_leaf = false;
if let Some(root_addr) = meta.root_addr {
backtrace::trace(|frame| {
let below_root = !ptr::eq(frame.symbol_address(), root_addr);
if let Some(root_addr) = meta.root_addr {
backtrace::trace(|frame| {
let below_root = !ptr::eq(frame.symbol_address(), root_addr);
if above_leaf && below_root {
frames.push(frame.to_owned().into());
}
if ptr::eq(frame.symbol_address(), meta.trace_leaf_addr) {
above_leaf = true;
}
below_root
});
if above_leaf && below_root {
frames.push(frame.to_owned().into());
}
collector.push_backtrace(frames);
state.collector.set(Some(collector));
}
});
if ptr::eq(frame.symbol_address(), meta.trace_leaf_addr) {
above_leaf = true;
}
below_root
});
}
trace.push_backtrace(frames);
}
+33 -40
View File
@@ -109,48 +109,41 @@ async fn task_trace_self() {
/// Collect frames between `trace_leaf_for_test` and `root_addr` using
/// `backtrace::trace`, resolve them, and store pretty-printed symbol names
/// (with compiler hashes stripped) into `TRACE_WITH_LOG`.
/// (with compiler hashes stripped) into `logs`.
#[inline(never)]
fn trace_leaf_for_test(meta: &TraceMeta) {
TRACE_WITH_LOG.with(|log| {
let mut frames: Vec<backtrace::BacktraceFrame> = vec![];
let mut above_leaf = false;
fn trace_leaf_for_test(meta: &TraceMeta, log: &mut Vec<Vec<String>>) {
let mut frames: Vec<backtrace::BacktraceFrame> = vec![];
let mut above_leaf = false;
if let Some(root_addr) = meta.root_addr {
backtrace::trace(|frame| {
let below_root = !ptr::eq(frame.symbol_address(), root_addr);
if let Some(root_addr) = meta.root_addr {
backtrace::trace(|frame| {
let below_root = !ptr::eq(frame.symbol_address(), root_addr);
if above_leaf && below_root {
frames.push(frame.to_owned().into());
}
if above_leaf && below_root {
frames.push(frame.to_owned().into());
}
if ptr::eq(frame.symbol_address(), meta.trace_leaf_addr) {
above_leaf = true;
}
if ptr::eq(frame.symbol_address(), meta.trace_leaf_addr) {
above_leaf = true;
}
below_root
});
}
below_root
});
}
// Resolve frames into human-readable symbol names with hashes stripped.
let mut bt = backtrace::Backtrace::from(frames);
bt.resolve();
let mut names = vec![];
for frame in bt.frames() {
for symbol in frame.symbols() {
if let Some(name) = symbol.name() {
names.push(strip_symbol_hash(&format!("{name}")).to_owned());
}
// Resolve frames into human-readable symbol names with hashes stripped.
let mut bt = backtrace::Backtrace::from(frames);
bt.resolve();
let mut names = vec![];
for frame in bt.frames() {
for symbol in frame.symbols() {
if let Some(name) = symbol.name() {
names.push(strip_symbol_hash(&format!("{name}")).to_owned());
}
}
}
log.borrow_mut().push(names);
});
}
thread_local! {
static TRACE_WITH_LOG: std::cell::RefCell<Vec<Vec<String>>> =
const { std::cell::RefCell::new(vec![]) };
log.push(names);
}
/// Strip the trailing `::h<hex>` hash that rustc appends to symbol names.
@@ -197,13 +190,17 @@ impl<F: Future> Future for TaskDump<F> {
Poll::Pending => {}
};
let mut logs = Vec::new();
// Tracing poll with a noop waker. If the future is at a yield
// point, trace_leaf fires our callback and returns Pending. We discard
// the result — this poll is purely for capturing the backtrace.
let noop = futures::task::noop_waker();
let mut noop_cx = Context::from_waker(&noop);
let logs = this.logs.clone();
let trace_poll = trace_with(|| this.f.as_mut().poll(&mut noop_cx), trace_leaf_for_test);
let trace_poll = trace_with(
|| this.f.as_mut().poll(&mut noop_cx),
|meta| trace_leaf_for_test(meta, &mut logs),
);
// trace should always produce poll pending
assert!(
matches!(trace_poll, Poll::Pending),
@@ -211,11 +208,7 @@ impl<F: Future> Future for TaskDump<F> {
);
// Drain any frames captured by trace_leaf_for_test into our log.
TRACE_WITH_LOG.with(|tl| {
let mut tl = tl.borrow_mut();
let mut dest = logs.lock().unwrap();
dest.append(&mut tl);
});
this.logs.lock().unwrap().extend(logs);
Poll::Pending
}
}