diff --git a/codex-rs/core/src/client.rs b/codex-rs/core/src/client.rs index 9c133d85e4..58e8d520c0 100644 --- a/codex-rs/core/src/client.rs +++ b/codex-rs/core/src/client.rs @@ -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 { @@ -918,7 +920,7 @@ impl ModelClient { turn_state: Option<&Arc>>, 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) -> Option { + // Responses request bodies carry the report in client metadata. + let mut metadata = + serde_json::from_str::>(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 diff --git a/codex-rs/core/src/client_tests.rs b/codex-rs/core/src/client_tests.rs index f8d904036a..fac01a47dc 100644 --- a/codex-rs/core/src/client_tests.rs +++ b/codex-rs/core/src/client_tests.rs @@ -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); diff --git a/codex-rs/core/src/lib.rs b/codex-rs/core/src/lib.rs index b77f96cdb1..d3cd6e3fae 100644 --- a/codex-rs/core/src/lib.rs +++ b/codex-rs/core/src/lib.rs @@ -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; diff --git a/codex-rs/core/src/tool_timing.rs b/codex-rs/core/src/tool_timing.rs new file mode 100644 index 0000000000..43a407cdc8 --- /dev/null +++ b/codex-rs/core/src/tool_timing.rs @@ -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>, +} + +#[derive(Debug)] +struct ToolTimingStateInner { + origin: Instant, + next_entry_id: u64, + next_report_id: u64, + calls: Vec, +} + +#[derive(Debug)] +struct ToolTimingCallState { + entry_id: u64, + call_id: String, + tool_name: String, + source: ToolTimingSource, + started_us: u64, + execution_started_us: Option, + completed_us: Option, +} + +#[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, +} + +#[derive(Debug, Serialize)] +struct ToolCallTimingReport { + call_id: String, + tool_name: String, + source: &'static str, + #[serde(skip_serializing_if = "Option::is_none")] + cell_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + runtime_tool_call_id: Option, + started_s: f64, + #[serde(skip_serializing_if = "Option::is_none")] + execution_started_s: Option, + 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::>(); + 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::>(); + 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; diff --git a/codex-rs/core/src/tool_timing_tests.rs b/codex-rs/core/src/tool_timing_tests.rs new file mode 100644 index 0000000000..7efde38398 --- /dev/null +++ b/codex-rs/core/src/tool_timing_tests.rs @@ -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)); +} diff --git a/codex-rs/core/src/tools/parallel.rs b/codex-rs/core/src/tools/parallel.rs index 7db1ec9643..03ca90a832 100644 --- a/codex-rs/core/src/tools/parallel.rs +++ b/codex-rs/core/src/tools/parallel.rs @@ -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() => { diff --git a/codex-rs/core/src/turn_metadata.rs b/codex-rs/core/src/turn_metadata.rs index e66a1fbb50..a161b8e89b 100644 --- a/codex-rs/core/src/turn_metadata.rs +++ b/codex-rs/core/src/turn_metadata.rs @@ -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>>, turn_started_at_unix_ms: Arc>>, responsesapi_client_metadata: Arc>>>, + tool_timing_state: ToolTimingState, user_input_requested_during_turn: Arc, enrichment_task: Arc>>>, } @@ -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 { - 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::>(&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 { @@ -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() } diff --git a/codex-rs/core/src/turn_metadata_tests.rs b/codex-rs/core/src/turn_metadata_tests.rs index 32a559848a..bb9636a41c 100644 --- a/codex-rs/core/src/turn_metadata_tests.rs +++ b/codex-rs/core/src/turn_metadata_tests.rs @@ -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(®ular_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()); }