use std::sync::Arc;
use std::time::Instant;
use tracing::Subscriber;
use tracing_subscriber::Layer;
use tracing_subscriber::layer::Context;
use tracing_subscriber::registry::LookupSpan;
use crate::metrics::MetricsCollector;
const WATCHED_SPANS: &[(&str, TimingField)] = &[
("core.context.prepare_context", TimingField::PrepareContext),
("llm.chat_with_tools", TimingField::LlmChat),
("core.tool.native_loop", TimingField::ToolExec),
];
const MAIN_TURN_SPAN: &str = "llm.turn_call";
#[derive(Clone, Copy)]
pub(crate) enum TimingField {
PrepareContext,
LlmChat,
ToolExec,
}
impl TimingField {
pub(crate) const fn bridge_bit(self) -> u8 {
match self {
Self::PrepareContext => 1 << 0,
Self::LlmChat => 1 << 1,
Self::ToolExec => 1 << 2,
}
}
}
struct WatchedSpan;
struct SpanEntry(Instant);
struct SpanTiming(u64);
pub struct MetricsBridge {
collector: Arc<MetricsCollector>,
}
impl MetricsBridge {
#[must_use]
pub fn new(collector: Arc<MetricsCollector>) -> Self {
Self { collector }
}
}
impl<S> Layer<S> for MetricsBridge
where
S: Subscriber + for<'a> LookupSpan<'a>,
{
fn on_new_span(
&self,
attrs: &tracing::span::Attributes<'_>,
id: &tracing::span::Id,
ctx: Context<'_, S>,
) {
let name = attrs.metadata().name();
if WATCHED_SPANS.iter().any(|(n, _)| *n == name)
&& let Some(span) = ctx.span(id)
{
span.extensions_mut().insert(WatchedSpan);
}
}
fn on_enter(&self, id: &tracing::span::Id, ctx: Context<'_, S>) {
if let Some(span) = ctx.span(id) {
if span.extensions().get::<WatchedSpan>().is_some() {
span.extensions_mut().replace(SpanEntry(Instant::now()));
}
}
}
fn on_exit(&self, id: &tracing::span::Id, ctx: Context<'_, S>) {
if let Some(span) = ctx.span(id) {
let elapsed_ms = span
.extensions()
.get::<SpanEntry>()
.map(|e| u64::try_from(e.0.elapsed().as_millis()).unwrap_or(u64::MAX));
if let Some(elapsed_ms) = elapsed_ms {
let mut exts = span.extensions_mut();
if let Some(timing) = exts.get_mut::<SpanTiming>() {
timing.0 = timing.0.saturating_add(elapsed_ms);
} else {
exts.insert(SpanTiming(elapsed_ms));
}
}
}
}
fn on_close(&self, id: tracing::span::Id, ctx: Context<'_, S>) {
if let Some(span) = ctx.span(&id) {
let name = span.name();
if let Some((_, field)) = WATCHED_SPANS.iter().find(|(n, _)| *n == name) {
let field = *field;
if matches!(field, TimingField::LlmChat)
&& !span
.scope()
.any(|ancestor| ancestor.name() == MAIN_TURN_SPAN)
{
return;
}
let exts = span.extensions();
if let Some(timing) = exts.get::<SpanTiming>() {
let duration_ms = timing.0;
self.collector.update(|m| {
match field {
TimingField::PrepareContext => {
m.last_turn_timings.prepare_context_ms = duration_ms;
}
TimingField::LlmChat => {
m.last_turn_timings.llm_chat_ms =
m.last_turn_timings.llm_chat_ms.saturating_add(duration_ms);
}
TimingField::ToolExec => {
m.last_turn_timings.tool_exec_ms = duration_ms;
}
}
m.bridge_timings_written |= field.bridge_bit();
});
}
}
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use tracing_subscriber::Registry;
use tracing_subscriber::layer::SubscriberExt;
use super::MetricsBridge;
use crate::metrics::MetricsCollector;
fn make_bridge() -> (
MetricsBridge,
Arc<MetricsCollector>,
tokio::sync::watch::Receiver<crate::metrics::MetricsSnapshot>,
) {
let (collector, rx) = MetricsCollector::new();
let arc = Arc::new(collector);
(MetricsBridge::new(Arc::clone(&arc)), arc, rx)
}
#[test]
fn watched_span_updates_correct_field() {
let (bridge, _collector, rx) = make_bridge();
let subscriber = Registry::default().with(bridge);
tracing::subscriber::with_default(subscriber, || {
let turn_span = tracing::span!(tracing::Level::INFO, "llm.turn_call");
let _turn_guard = turn_span.enter();
let span = tracing::span!(tracing::Level::INFO, "llm.chat_with_tools");
let guard = span.enter();
drop(guard);
});
let snapshot = rx.borrow().clone();
assert_eq!(snapshot.last_turn_timings.prepare_context_ms, 0);
assert_eq!(snapshot.last_turn_timings.tool_exec_ms, 0);
assert_eq!(snapshot.last_turn_timings.persist_message_ms, 0);
let _ = snapshot.last_turn_timings.llm_chat_ms;
}
#[test]
fn bare_llm_chat_span_is_not_watched() {
let (bridge, _collector, rx) = make_bridge();
let subscriber = Registry::default().with(bridge);
tracing::subscriber::with_default(subscriber, || {
let span = tracing::span!(tracing::Level::INFO, "llm.chat");
let guard = span.enter();
drop(guard);
});
let snapshot = rx.borrow().clone();
assert_eq!(snapshot.last_turn_timings.llm_chat_ms, 0);
assert_eq!(snapshot.bridge_timings_written, 0);
}
#[test]
fn llm_chat_with_tools_span_accumulates_across_multiple_closes() {
let (bridge, _collector, rx) = make_bridge();
let subscriber = Registry::default().with(bridge);
tracing::subscriber::with_default(subscriber, || {
let turn_span = tracing::span!(tracing::Level::INFO, "llm.turn_call");
let _turn_guard = turn_span.enter();
for _ in 0..3 {
let span = tracing::span!(tracing::Level::INFO, "llm.chat_with_tools");
let guard = span.enter();
std::thread::sleep(std::time::Duration::from_millis(10));
drop(guard);
}
});
let snapshot = rx.borrow().clone();
assert!(
snapshot.last_turn_timings.llm_chat_ms >= 25,
"expected accumulated duration across 3 closes (~30ms+), got {}ms — \
on_close may be overwriting instead of accumulating",
snapshot.last_turn_timings.llm_chat_ms
);
}
#[test]
fn llm_chat_with_tools_outside_main_turn_span_is_ignored() {
let (bridge, _collector, rx) = make_bridge();
let subscriber = Registry::default().with(bridge);
tracing::subscriber::with_default(subscriber, || {
let span = tracing::span!(tracing::Level::INFO, "llm.chat_with_tools");
let guard = span.enter();
std::thread::sleep(std::time::Duration::from_millis(10));
drop(guard);
});
let snapshot = rx.borrow().clone();
assert_eq!(
snapshot.last_turn_timings.llm_chat_ms, 0,
"an llm.chat_with_tools span with no llm.turn_call ancestor must not be counted"
);
assert_eq!(
snapshot.bridge_timings_written, 0,
"an out-of-scope llm.chat_with_tools span must not set the LlmChat bridge bit"
);
}
#[test]
fn llm_chat_with_tools_inside_main_turn_span_counted_despite_concurrent_out_of_scope_span() {
let (bridge, _collector, rx) = make_bridge();
let subscriber = Registry::default().with(bridge);
tracing::subscriber::with_default(subscriber, || {
let outside = tracing::span!(tracing::Level::INFO, "llm.chat_with_tools");
let outside_guard = outside.enter();
drop(outside_guard);
let turn_span = tracing::span!(tracing::Level::INFO, "llm.turn_call");
let _turn_guard = turn_span.enter();
let inside = tracing::span!(tracing::Level::INFO, "llm.chat_with_tools");
let inside_guard = inside.enter();
std::thread::sleep(std::time::Duration::from_millis(10));
drop(inside_guard);
});
let snapshot = rx.borrow().clone();
assert!(
snapshot.last_turn_timings.llm_chat_ms >= 5,
"the in-scope span must still be counted, got {}ms",
snapshot.last_turn_timings.llm_chat_ms
);
}
#[test]
fn non_watched_span_produces_no_update() {
let (bridge, _collector, rx) = make_bridge();
let subscriber = Registry::default().with(bridge);
tracing::subscriber::with_default(subscriber, || {
let span = tracing::span!(tracing::Level::INFO, "some.other.span");
let guard = span.enter();
drop(guard);
});
let snapshot = rx.borrow().clone();
assert_eq!(snapshot.last_turn_timings.prepare_context_ms, 0);
assert_eq!(snapshot.last_turn_timings.llm_chat_ms, 0);
assert_eq!(snapshot.last_turn_timings.tool_exec_ms, 0);
assert_eq!(snapshot.last_turn_timings.persist_message_ms, 0);
}
#[test]
fn all_watched_span_names_registered() {
let expected = [
"core.context.prepare_context",
"llm.chat_with_tools",
"core.tool.native_loop",
];
for span_name in expected {
assert!(
super::WATCHED_SPANS.iter().any(|(n, _)| *n == span_name),
"span '{span_name}' not in WATCHED_SPANS",
);
}
assert_eq!(
super::WATCHED_SPANS.len(),
expected.len(),
"unexpected extra spans in WATCHED_SPANS"
);
}
#[test]
fn watched_spans_match_real_instrument_names_in_this_crate() {
let assembly_src = include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/src/agent/context/assembly.rs"
));
let tier_loop_src = include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/src/agent/tool_execution/tier_loop.rs"
));
assert!(
assembly_src.contains(r#"name = "core.context.prepare_context""#),
"core.context.prepare_context span not found in assembly.rs — WATCHED_SPANS has drifted"
);
assert!(
tier_loop_src.contains(r#"name = "core.tool.native_loop""#),
"core.tool.native_loop span not found in tier_loop.rs — WATCHED_SPANS has drifted"
);
}
}