Add exact tool timing metadata

This commit is contained in:
Suraj Srinivasan
2026-06-19 00:04:22 -07:00
parent 2ec687715d
commit cdf4f6eef1
8 changed files with 487 additions and 9 deletions

View File

@@ -85,6 +85,7 @@ use codex_rollout_trace::CompactionTraceContext;
use codex_rollout_trace::InferenceTraceAttempt;
use codex_rollout_trace::InferenceTraceContext;
use codex_tools::create_tools_json_for_responses_api;
use codex_utils_string::to_ascii_json_string;
use eventsource_stream::Event;
use eventsource_stream::EventStreamError;
use futures::StreamExt;
@@ -111,6 +112,7 @@ use crate::client_common::Prompt;
use crate::client_common::ResponseEvent;
use crate::client_common::ResponseStream;
use crate::feedback_tags;
use crate::tool_timing::TOOL_TIMING_KEY;
use crate::util::emit_feedback_auth_recovery_tags;
use codex_api::map_api_error;
use codex_feedback::FeedbackRequestTags;
@@ -650,7 +652,7 @@ impl ModelClient {
extra_headers
}
fn build_ws_client_metadata(
fn build_responses_client_metadata(
&self,
turn_metadata_header: Option<&str>,
) -> HashMap<String, String> {
@@ -918,7 +920,7 @@ impl ModelClient {
turn_state: Option<&Arc<OnceLock<String>>>,
turn_metadata_header: Option<&str>,
) -> ApiHeaderMap {
let turn_metadata_header = parse_turn_metadata_header(turn_metadata_header);
let turn_metadata_header = parse_transport_turn_metadata_header(turn_metadata_header);
let session_id = self.state.session_id.to_string();
let thread_id = self.state.thread_id.to_string();
let mut headers = build_responses_headers(
@@ -976,7 +978,7 @@ impl ModelClientSession {
turn_metadata_header: Option<&str>,
compression: Compression,
) -> ApiResponsesOptions {
let turn_metadata_header = parse_turn_metadata_header(turn_metadata_header);
let turn_metadata_header = parse_transport_turn_metadata_header(turn_metadata_header);
let session_id = self.client.state.session_id.to_string();
let thread_id = self.client.state.thread_id.to_string();
ApiResponsesOptions {
@@ -1263,7 +1265,7 @@ impl ModelClientSession {
.build_responses_options(turn_metadata_header, compression)
.await;
let request = self.client.build_responses_request(
let mut request = self.client.build_responses_request(
&client_setup.api_provider,
prompt,
model_info,
@@ -1271,6 +1273,10 @@ impl ModelClientSession {
summary,
service_tier.clone(),
)?;
request.client_metadata = Some(
self.client
.build_responses_client_metadata(turn_metadata_header),
);
let inference_trace_attempt = inference_trace.start_attempt();
inference_trace_attempt.add_request_headers(&mut options.extra_headers);
inference_trace_attempt.record_started(&request);
@@ -1372,7 +1378,7 @@ impl ModelClientSession {
let options = self
.build_responses_options(turn_metadata_header, compression)
.await;
let request = self.client.build_responses_request(
let mut request = self.client.build_responses_request(
&client_setup.api_provider,
prompt,
model_info,
@@ -1380,9 +1386,13 @@ impl ModelClientSession {
summary,
service_tier.clone(),
)?;
request.client_metadata = Some(
self.client
.build_responses_client_metadata(turn_metadata_header),
);
let mut ws_payload = ResponseCreateWsRequest {
client_metadata: response_create_client_metadata(
Some(self.client.build_ws_client_metadata(turn_metadata_header)),
request.client_metadata.clone(),
request_trace.as_ref(),
),
..ResponseCreateWsRequest::from(&request)
@@ -1654,6 +1664,15 @@ fn parse_turn_metadata_header(turn_metadata_header: Option<&str>) -> Option<Head
turn_metadata_header.and_then(|value| HeaderValue::from_str(value).ok())
}
fn parse_transport_turn_metadata_header(turn_metadata_header: Option<&str>) -> Option<HeaderValue> {
// Responses request bodies carry the report in client metadata.
let mut metadata =
serde_json::from_str::<serde_json::Map<String, serde_json::Value>>(turn_metadata_header?)
.ok()?;
metadata.remove(TOOL_TIMING_KEY);
HeaderValue::from_str(&to_ascii_json_string(&metadata).ok()?).ok()
}
/// Stamp a ResponsesWsRequest with the current time.
///
/// Meant to be called just before sending the request over the socket, to capture realistic

View File

@@ -7,6 +7,7 @@ use super::X_CODEX_PARENT_THREAD_ID_HEADER;
use super::X_CODEX_TURN_METADATA_HEADER;
use super::X_CODEX_WINDOW_ID_HEADER;
use super::X_OPENAI_SUBAGENT_HEADER;
use super::parse_transport_turn_metadata_header;
use crate::AttestationContext;
use crate::AttestationProvider;
use crate::GenerateAttestationFuture;
@@ -278,7 +279,7 @@ fn build_subagent_headers_sets_internal_memory_consolidation_label() {
}
#[test]
fn build_ws_client_metadata_includes_window_lineage_and_turn_metadata() {
fn build_responses_client_metadata_includes_window_lineage_and_turn_metadata() {
let parent_thread_id = ThreadId::new();
let client = test_model_client_with_parent(
SessionSource::SubAgent(SubAgentSource::ThreadSpawn {
@@ -293,7 +294,7 @@ fn build_ws_client_metadata_includes_window_lineage_and_turn_metadata() {
client.advance_window_generation();
let client_metadata = client.build_ws_client_metadata(Some(r#"{"turn_id":"turn-123"}"#));
let client_metadata = client.build_responses_client_metadata(Some(r#"{"turn_id":"turn-123"}"#));
let thread_id = client.state.thread_id;
assert_eq!(
client_metadata,
@@ -322,6 +323,18 @@ fn build_ws_client_metadata_includes_window_lineage_and_turn_metadata() {
);
}
#[test]
fn transport_turn_metadata_header_omits_tool_timing() {
let header = parse_transport_turn_metadata_header(Some(
r#"{"turn_id":"turn-123","tool_timing":{"version":1}}"#,
))
.expect("valid metadata header");
let metadata: serde_json::Value =
serde_json::from_str(header.to_str().expect("ASCII metadata header")).expect("JSON");
assert_eq!(metadata, json!({"turn_id": "turn-123"}));
}
#[tokio::test]
async fn summarize_memories_returns_empty_for_empty_input() {
let client = test_model_client(SessionSource::Cli);

View File

@@ -142,6 +142,7 @@ pub(crate) mod state_db_bridge;
pub use state_db_bridge::StateDbHandle;
pub use state_db_bridge::init_state_db;
mod thread_rollout_truncation;
mod tool_timing;
mod tools;
pub(crate) mod turn_diff_tracker;
mod turn_metadata;

View File

@@ -0,0 +1,268 @@
//! Turn-local, model-visible tool latency exported through Responses client metadata.
use std::mem;
use std::sync::Arc;
use std::sync::Mutex;
use std::time::Instant;
use serde::Serialize;
const TOOL_TIMING_VERSION: u8 = 1;
pub(crate) const TOOL_TIMING_KEY: &str = "tool_timing";
#[derive(Clone, Debug)]
pub(crate) enum ToolTimingSource {
Direct,
CodeMode {
cell_id: String,
runtime_tool_call_id: String,
},
}
#[derive(Clone, Debug)]
pub(crate) struct ToolTimingCall {
pub(crate) call_id: String,
pub(crate) tool_name: String,
pub(crate) source: ToolTimingSource,
}
#[derive(Clone, Debug)]
pub(crate) struct ToolTimingState {
inner: Arc<Mutex<ToolTimingStateInner>>,
}
#[derive(Debug)]
struct ToolTimingStateInner {
origin: Instant,
next_entry_id: u64,
next_report_id: u64,
calls: Vec<ToolTimingCallState>,
}
#[derive(Debug)]
struct ToolTimingCallState {
entry_id: u64,
call_id: String,
tool_name: String,
source: ToolTimingSource,
started_us: u64,
execution_started_us: Option<u64>,
completed_us: Option<u64>,
}
#[derive(Clone, Debug)]
pub(crate) struct ToolTimingMarker {
state: ToolTimingState,
entry_id: u64,
}
#[derive(Debug)]
pub(crate) struct ToolTimingGuard {
marker: ToolTimingMarker,
}
#[derive(Debug, Serialize)]
pub(crate) struct ToolTimingReport {
version: u8,
report_id: u64,
tool_active_s: f64,
calls: Vec<ToolCallTimingReport>,
}
#[derive(Debug, Serialize)]
struct ToolCallTimingReport {
call_id: String,
tool_name: String,
source: &'static str,
#[serde(skip_serializing_if = "Option::is_none")]
cell_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
runtime_tool_call_id: Option<String>,
started_s: f64,
#[serde(skip_serializing_if = "Option::is_none")]
execution_started_s: Option<f64>,
completed_s: f64,
dispatch_s: f64,
handler_s: f64,
total_s: f64,
}
impl Default for ToolTimingState {
fn default() -> Self {
Self {
inner: Arc::new(Mutex::new(ToolTimingStateInner {
origin: Instant::now(),
next_entry_id: 0,
next_report_id: 0,
calls: Vec::new(),
})),
}
}
}
impl ToolTimingState {
pub(crate) fn start_call(&self, call: ToolTimingCall) -> ToolTimingGuard {
let mut state = self
.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let entry_id = state.next_entry_id;
state.next_entry_id += 1;
let started_us = state.elapsed_us();
state.calls.push(ToolTimingCallState {
entry_id,
call_id: call.call_id,
tool_name: call.tool_name,
source: call.source,
started_us,
execution_started_us: None,
completed_us: None,
});
ToolTimingGuard {
marker: ToolTimingMarker {
state: self.clone(),
entry_id,
},
}
}
pub(crate) fn take_report(&self) -> ToolTimingReport {
let mut state = self
.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let report_id = state.next_report_id;
state.next_report_id += 1;
let mut completed_calls = Vec::new();
let mut pending_calls = Vec::new();
for call in mem::take(&mut state.calls) {
match call.completed_us {
Some(completed_us) => completed_calls.push((call, completed_us)),
None => pending_calls.push(call),
}
}
state.calls = pending_calls;
// Nested code-mode calls are diagnostic; their enclosing direct call owns the latency.
let mut intervals = completed_calls
.iter()
.filter(|(call, _)| matches!(&call.source, ToolTimingSource::Direct))
.map(|(call, completed_us)| (call.started_us, *completed_us))
.collect::<Vec<_>>();
intervals.sort_unstable_by_key(|interval| interval.0);
let mut tool_active_us = 0_u64;
if let Some((mut current_start_us, mut current_end_us)) = intervals.first().copied() {
for (started_us, completed_us) in intervals.into_iter().skip(1) {
if started_us <= current_end_us {
current_end_us = current_end_us.max(completed_us);
} else {
tool_active_us += current_end_us - current_start_us;
current_start_us = started_us;
current_end_us = completed_us;
}
}
tool_active_us += current_end_us - current_start_us;
}
let mut calls = completed_calls
.into_iter()
.map(|(call, completed_us)| {
let execution_started_us = call.execution_started_us;
let dispatch_completed_us = execution_started_us.unwrap_or(completed_us);
let (source, cell_id, runtime_tool_call_id) = match call.source {
ToolTimingSource::Direct => ("direct", None, None),
ToolTimingSource::CodeMode {
cell_id,
runtime_tool_call_id,
} => ("code_mode", Some(cell_id), Some(runtime_tool_call_id)),
};
ToolCallTimingReport {
call_id: call.call_id,
tool_name: call.tool_name,
source,
cell_id,
runtime_tool_call_id,
started_s: micros_to_seconds(call.started_us),
execution_started_s: execution_started_us.map(micros_to_seconds),
completed_s: micros_to_seconds(completed_us),
dispatch_s: micros_to_seconds(dispatch_completed_us - call.started_us),
handler_s: micros_to_seconds(completed_us - dispatch_completed_us),
total_s: micros_to_seconds(completed_us - call.started_us),
}
})
.collect::<Vec<_>>();
calls.sort_by(|left, right| left.started_s.total_cmp(&right.started_s));
ToolTimingReport {
version: TOOL_TIMING_VERSION,
report_id,
tool_active_s: micros_to_seconds(tool_active_us),
calls,
}
}
fn mark_execution_started(&self, entry_id: u64) {
let mut state = self
.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let execution_started_us = state.elapsed_us();
if let Some(call) = state
.calls
.iter_mut()
.find(|call| call.entry_id == entry_id)
&& call.completed_us.is_none()
&& call.execution_started_us.is_none()
{
call.execution_started_us = Some(execution_started_us);
}
}
fn complete_call(&self, entry_id: u64) {
let mut state = self
.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let completed_us = state.elapsed_us();
if let Some(call) = state
.calls
.iter_mut()
.find(|call| call.entry_id == entry_id)
{
call.completed_us = Some(completed_us);
}
}
}
impl ToolTimingStateInner {
fn elapsed_us(&self) -> u64 {
u64::try_from(self.origin.elapsed().as_micros()).unwrap_or(u64::MAX)
}
}
impl ToolTimingGuard {
pub(crate) fn marker(&self) -> ToolTimingMarker {
self.marker.clone()
}
}
impl Drop for ToolTimingGuard {
fn drop(&mut self) {
self.marker.state.complete_call(self.marker.entry_id);
}
}
impl ToolTimingMarker {
pub(crate) fn mark_execution_started(&self) {
self.state.mark_execution_started(self.entry_id);
}
}
fn micros_to_seconds(micros: u64) -> f64 {
micros as f64 / 1_000_000.0
}
#[cfg(test)]
#[path = "tool_timing_tests.rs"]
mod tests;

View File

@@ -0,0 +1,115 @@
use std::sync::Arc;
use std::sync::Mutex;
use std::time::Instant;
use pretty_assertions::assert_eq;
use serde_json::json;
use super::*;
#[test]
fn direct_calls_are_unioned_and_nested_calls_are_diagnostic() {
let state = ToolTimingState {
inner: Arc::new(Mutex::new(ToolTimingStateInner {
origin: Instant::now(),
next_entry_id: 3,
next_report_id: 4,
calls: vec![
ToolTimingCallState {
entry_id: 0,
call_id: "direct-a".to_string(),
tool_name: "functions.exec".to_string(),
source: ToolTimingSource::Direct,
started_us: 0,
execution_started_us: Some(500_000),
completed_us: Some(2_000_000),
},
ToolTimingCallState {
entry_id: 1,
call_id: "nested".to_string(),
tool_name: "mcp__example__read".to_string(),
source: ToolTimingSource::CodeMode {
cell_id: "cell-1".to_string(),
runtime_tool_call_id: "runtime-call-1".to_string(),
},
started_us: 500_000,
execution_started_us: Some(1_000_000),
completed_us: Some(2_500_000),
},
ToolTimingCallState {
entry_id: 2,
call_id: "direct-b".to_string(),
tool_name: "functions.wait".to_string(),
source: ToolTimingSource::Direct,
started_us: 1_000_000,
execution_started_us: Some(1_250_000),
completed_us: Some(3_000_000),
},
],
})),
};
assert_eq!(
serde_json::to_value(state.take_report()).expect("serialize report"),
json!({
"version": 1,
"report_id": 4,
"tool_active_s": 3.0,
"calls": [
{
"call_id": "direct-a",
"tool_name": "functions.exec",
"source": "direct",
"started_s": 0.0,
"execution_started_s": 0.5,
"completed_s": 2.0,
"dispatch_s": 0.5,
"handler_s": 1.5,
"total_s": 2.0,
},
{
"call_id": "nested",
"tool_name": "mcp__example__read",
"source": "code_mode",
"cell_id": "cell-1",
"runtime_tool_call_id": "runtime-call-1",
"started_s": 0.5,
"execution_started_s": 1.0,
"completed_s": 2.5,
"dispatch_s": 0.5,
"handler_s": 1.5,
"total_s": 2.0,
},
{
"call_id": "direct-b",
"tool_name": "functions.wait",
"source": "direct",
"started_s": 1.0,
"execution_started_s": 1.25,
"completed_s": 3.0,
"dispatch_s": 0.25,
"handler_s": 1.75,
"total_s": 2.0,
},
],
})
);
}
#[test]
fn execution_start_after_completion_is_ignored() {
let state = ToolTimingState::default();
let guard = state.start_call(ToolTimingCall {
call_id: "call-1".to_string(),
tool_name: "functions.exec".to_string(),
source: ToolTimingSource::Direct,
});
let marker = guard.marker();
drop(guard);
marker.mark_execution_started();
let report = serde_json::to_value(state.take_report()).expect("serialize report");
assert_eq!(report["calls"][0]["execution_started_s"], json!(null));
assert_eq!(report["calls"][0]["handler_s"], json!(0.0));
}

View File

@@ -15,6 +15,8 @@ use tracing::trace_span;
use crate::function_tool::FunctionCallError;
use crate::session::session::Session;
use crate::session::turn_context::TurnContext;
use crate::tool_timing::ToolTimingCall;
use crate::tool_timing::ToolTimingSource;
use crate::tools::context::AbortedToolOutput;
use crate::tools::context::SharedTurnDiffTracker;
use crate::tools::context::ToolPayload;
@@ -94,6 +96,22 @@ impl ToolCallRuntime {
let invocation_cancellation_token = cancellation_token.clone();
let wait_for_runtime_cancellation = self.router.tool_waits_for_runtime_cancellation(&call);
let started = Instant::now();
let tool_timing_source = match &source {
ToolCallSource::Direct => ToolTimingSource::Direct,
ToolCallSource::CodeMode {
cell_id,
runtime_tool_call_id,
} => ToolTimingSource::CodeMode {
cell_id: cell_id.clone(),
runtime_tool_call_id: runtime_tool_call_id.clone(),
},
};
let tool_timing_guard = turn.turn_metadata_state.start_tool_timing(ToolTimingCall {
call_id: call.call_id.clone(),
tool_name: call.tool_name.to_string(),
source: tool_timing_source,
});
let tool_timing_marker = tool_timing_guard.marker();
let abort_session = Arc::clone(&session);
let abort_source = source.clone();
let abort_turn = Arc::clone(&turn);
@@ -117,6 +135,7 @@ impl ToolCallRuntime {
} else {
Either::Right(lock.write().await)
};
tool_timing_marker.mark_execution_started();
router
.dispatch_tool_call_with_terminal_outcome(
@@ -133,6 +152,7 @@ impl ToolCallRuntime {
}));
async move {
let _tool_timing_guard = tool_timing_guard;
tokio::select! {
res = &mut handle => res.map_err(Self::tool_task_join_error)?,
_ = cancellation_token.cancelled() => {

View File

@@ -17,6 +17,10 @@ use serde_json::Value;
use tokio::task::JoinHandle;
use crate::sandbox_tags::permission_profile_sandbox_tag;
use crate::tool_timing::TOOL_TIMING_KEY;
use crate::tool_timing::ToolTimingCall;
use crate::tool_timing::ToolTimingGuard;
use crate::tool_timing::ToolTimingState;
use codex_git_utils::get_git_remote_urls_assume_git_repo;
use codex_git_utils::get_git_repo_root;
use codex_git_utils::get_has_changes;
@@ -197,6 +201,7 @@ fn merge_turn_metadata(
| REQUEST_KIND_KEY
| COMPACTION_KEY
| WINDOW_ID_KEY
| TOOL_TIMING_KEY
) {
continue;
}
@@ -252,6 +257,7 @@ pub(crate) struct TurnMetadataState {
enriched_header: Arc<RwLock<Option<String>>>,
turn_started_at_unix_ms: Arc<RwLock<Option<i64>>>,
responsesapi_client_metadata: Arc<RwLock<Option<HashMap<String, String>>>>,
tool_timing_state: ToolTimingState,
user_input_requested_during_turn: Arc<AtomicBool>,
enrichment_task: Arc<Mutex<Option<JoinHandle<()>>>>,
}
@@ -312,6 +318,7 @@ impl TurnMetadataState {
enriched_header: Arc::new(RwLock::new(None)),
turn_started_at_unix_ms: Arc::new(RwLock::new(None)),
responsesapi_client_metadata: Arc::new(RwLock::new(None)),
tool_timing_state: ToolTimingState::default(),
user_input_requested_during_turn: Arc::new(AtomicBool::new(false)),
enrichment_task: Arc::new(Mutex::new(None)),
}
@@ -401,7 +408,20 @@ impl TurnMetadataState {
}
pub(crate) fn current_header_value_for_model_request(&self, window_id: &str) -> Option<String> {
self.current_header_value_for_model_request_kind(window_id, TurnMetadataRequestKind::Turn)
let header = self.current_header_value_for_model_request_kind(
window_id,
TurnMetadataRequestKind::Turn,
)?;
let mut metadata = serde_json::from_str::<serde_json::Map<String, Value>>(&header).ok()?;
metadata.insert(
TOOL_TIMING_KEY.to_string(),
serde_json::to_value(self.tool_timing_state.take_report()).ok()?,
);
to_ascii_json_string(&metadata).ok()
}
pub(crate) fn start_tool_timing(&self, call: ToolTimingCall) -> ToolTimingGuard {
self.tool_timing_state.start_call(call)
}
pub(crate) fn current_header_value_for_prewarm(&self, window_id: &str) -> Option<String> {
@@ -425,6 +445,10 @@ impl TurnMetadataState {
COMPACTION_KEY.to_string(),
serde_json::to_value(compaction).ok()?,
);
metadata.insert(
TOOL_TIMING_KEY.to_string(),
serde_json::to_value(self.tool_timing_state.take_report()).ok()?,
);
to_ascii_json_string(&metadata).ok()
}

View File

@@ -662,6 +662,15 @@ fn turn_metadata_state_overlays_compaction_only_on_compaction_requests() {
assert_eq!(compact_json["request_kind"].as_str(), Some("compaction"));
assert_eq!(compact_json["turn_id"].as_str(), Some("turn-a"));
assert_eq!(compact_json[WINDOW_ID_KEY].as_str(), Some("thread-a:2"));
assert_eq!(
compact_json["tool_timing"],
serde_json::json!({
"version": 1,
"report_id": 0,
"tool_active_s": 0.0,
"calls": [],
})
);
assert_eq!(
compact_json["compaction"],
serde_json::json!({
@@ -679,6 +688,15 @@ fn turn_metadata_state_overlays_compaction_only_on_compaction_requests() {
let regular_json: Value = serde_json::from_str(&regular_header).expect("json");
assert_eq!(regular_json["request_kind"].as_str(), Some("turn"));
assert_eq!(regular_json[WINDOW_ID_KEY].as_str(), Some("thread-a:3"));
assert_eq!(
regular_json["tool_timing"],
serde_json::json!({
"version": 1,
"report_id": 1,
"tool_active_s": 0.0,
"calls": [],
})
);
assert!(regular_json.get("compaction").is_none());
}