mirror of
https://github.com/openai/codex.git
synced 2026-09-13 11:47:17 +00:00
Add exact tool timing metadata
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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;
|
||||
|
||||
268
codex-rs/core/src/tool_timing.rs
Normal file
268
codex-rs/core/src/tool_timing.rs
Normal 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;
|
||||
115
codex-rs/core/src/tool_timing_tests.rs
Normal file
115
codex-rs/core/src/tool_timing_tests.rs
Normal 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));
|
||||
}
|
||||
@@ -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() => {
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user