From 653730639a2b206603e48a387e1694f86081831e Mon Sep 17 00:00:00 2001 From: mulfyx Date: Wed, 29 Jul 2026 19:14:41 +0500 Subject: [PATCH] fix(codex): report streaming input usage --- src/providers/codex/mod.rs | 9 ++- src/providers/codex/translate/live_stream.rs | 66 +++++++++++++++++++- src/providers/codex/translate/stream.rs | 49 +++++++++------ tests/smoke_cutover.rs | 8 +++ 4 files changed, 111 insertions(+), 21 deletions(-) diff --git a/src/providers/codex/mod.rs b/src/providers/codex/mod.rs index 22d7a51a..7d8a9455 100644 --- a/src/providers/codex/mod.rs +++ b/src/providers/codex/mod.rs @@ -353,10 +353,12 @@ impl Provider for CodexProvider { }; if want_stream { + let estimated_input_tokens = count_translated_tokens(&translated); let sse_bytes = match translate_stream_bytes_with_traffic( &upstream.body, &message_id, model, + estimated_input_tokens, ctx.traffic.as_deref(), ) { Ok(b) => b, @@ -647,7 +649,12 @@ async fn live_stream_response_once( request_body: translate::request::ResponsesRequest, compact_boundary: bool, ) -> LiveStreamStart { - let mut translator = LiveStreamTranslator::new(message_id, model.to_string()); + let estimated_input_tokens = count_translated_tokens(&request_body); + let mut translator = LiveStreamTranslator::with_estimated_input_tokens( + message_id, + model.to_string(), + estimated_input_tokens, + ); let mut upstream_sse_body = Vec::new(); // Keep protocol framing private until real output makes a transparent retry unsafe. // Every branch that consumes pending_chunk returns, so it is never flushed twice. diff --git a/src/providers/codex/translate/live_stream.rs b/src/providers/codex/translate/live_stream.rs index de79e163..badbaf35 100644 --- a/src/providers/codex/translate/live_stream.rs +++ b/src/providers/codex/translate/live_stream.rs @@ -64,11 +64,22 @@ pub struct LiveStreamTranslator { web_search_results: Vec, deferred_text: Vec<(usize, String)>, semantic_output_started: bool, + // Seeds Claude Code's live subagent counter until the provider returns + // authoritative usage in the terminal message_delta. + estimated_input_tokens: u64, finished: bool, } impl LiveStreamTranslator { pub fn new(message_id: impl Into, model: impl Into) -> Self { + Self::with_estimated_input_tokens(message_id, model, 0) + } + + pub fn with_estimated_input_tokens( + message_id: impl Into, + model: impl Into, + estimated_input_tokens: u64, + ) -> Self { Self { message_id: message_id.into(), model: model.into(), @@ -84,6 +95,7 @@ impl LiveStreamTranslator { web_search_results: Vec::new(), deferred_text: Vec::new(), semantic_output_started: false, + estimated_input_tokens, finished: false, } } @@ -231,7 +243,7 @@ impl LiveStreamTranslator { "stop_reason": null, "stop_sequence": null, "usage": { - "input_tokens": 0, + "input_tokens": self.estimated_input_tokens, "output_tokens": 0 } } @@ -1169,6 +1181,7 @@ fn error_message(payload: &serde_json::Value) -> String { #[cfg(test)] mod tests { use super::*; + use crate::anthropic::sse::parse_sse_events; use serde_json::json; fn render(events: Vec) -> String { @@ -1202,6 +1215,57 @@ mod tests { assert!(translator.has_semantic_output()); } + #[test] + fn estimated_input_is_visible_at_start_and_provider_usage_is_exact_at_finish() { + let mut translator = + LiveStreamTranslator::with_estimated_input_tokens("msg_1", "gpt-5.5", 321); + + let started = translator + .accept( + &json!({ + "type": "response.output_text.delta", + "output_index": 0, + "delta": "abcdefgh" + }), + None, + ) + .unwrap(); + let started = parse_sse_events(&started) + .into_iter() + .filter_map(|event| serde_json::from_str::(&event.data).ok()) + .find(|value| { + value.get("type").and_then(serde_json::Value::as_str) == Some("message_start") + }) + .unwrap(); + assert_eq!( + started.pointer("/message/usage/input_tokens"), + Some(&json!(321)) + ); + + let finished = translator + .accept( + &json!({ + "type": "response.completed", + "response": { + "id": "resp_1", + "status": "completed", + "usage": {"input_tokens": 300, "output_tokens": 9} + } + }), + None, + ) + .unwrap(); + let finished = parse_sse_events(&finished) + .into_iter() + .filter_map(|event| serde_json::from_str::(&event.data).ok()) + .find(|value| { + value.get("type").and_then(serde_json::Value::as_str) == Some("message_delta") + }) + .unwrap(); + assert_eq!(finished.pointer("/usage/input_tokens"), Some(&json!(300))); + assert_eq!(finished.pointer("/usage/output_tokens"), Some(&json!(9))); + } + #[test] fn finishes_text_stream() { let out = render(vec![ diff --git a/src/providers/codex/translate/stream.rs b/src/providers/codex/translate/stream.rs index 73596bb3..804b3846 100644 --- a/src/providers/codex/translate/stream.rs +++ b/src/providers/codex/translate/stream.rs @@ -15,6 +15,12 @@ enum OpenBlock { Tool { id: String, name: String }, } +struct MessageMetadata<'a> { + id: &'a str, + model: &'a str, + estimated_input_tokens: u64, +} + fn emit( out: &mut Vec, traffic: Option<&TrafficCapture>, @@ -37,23 +43,22 @@ fn ensure_message_start( out: &mut Vec, traffic: Option<&TrafficCapture>, message_started: &mut bool, - message_id: &str, - model: &str, + message: &MessageMetadata<'_>, ) { if !*message_started { *message_started = true; let data = serde_json::json!({ "type": "message_start", "message": { - "id": message_id, + "id": message.id, "type": "message", "role": "assistant", - "model": model, + "model": message.model, "content": [], "stop_reason": null, "stop_sequence": null, "usage": { - "input_tokens": 0, + "input_tokens": message.estimated_input_tokens, "output_tokens": 0, } } @@ -67,13 +72,12 @@ fn emit_content_event( traffic: Option<&TrafficCapture>, message_started: &mut bool, open_blocks: &mut BTreeMap, - message_id: &str, - model: &str, + message: &MessageMetadata<'_>, event: &ReducerEvent, ) -> bool { match event { ReducerEvent::ThinkingStart { index } => { - ensure_message_start(out, traffic, message_started, message_id, model); + ensure_message_start(out, traffic, message_started, message); open_blocks.insert(*index, OpenBlock::Thinking); emit( out, @@ -127,7 +131,7 @@ fn emit_content_event( true } ReducerEvent::TextStart { index } => { - ensure_message_start(out, traffic, message_started, message_id, model); + ensure_message_start(out, traffic, message_started, message); open_blocks.insert(*index, OpenBlock::Text); emit( out, @@ -168,7 +172,7 @@ fn emit_content_event( true } ReducerEvent::ToolStart { index, id, name } => { - ensure_message_start(out, traffic, message_started, message_id, model); + ensure_message_start(out, traffic, message_started, message); open_blocks.insert( *index, OpenBlock::Tool { @@ -250,13 +254,14 @@ pub fn translate_stream_bytes( message_id: &str, model: &str, ) -> Result, anyhow::Error> { - translate_stream_bytes_with_traffic(upstream, message_id, model, None) + translate_stream_bytes_with_traffic(upstream, message_id, model, 0, None) } pub fn translate_stream_bytes_with_traffic( upstream: &[u8], message_id: &str, model: &str, + estimated_input_tokens: u64, traffic: Option<&TrafficCapture>, ) -> Result, anyhow::Error> { let events = match reduce_upstream_bytes(upstream) { @@ -276,6 +281,11 @@ pub fn translate_stream_bytes_with_traffic( let mut open_blocks: BTreeMap = BTreeMap::new(); let mut web_search_events: Vec = Vec::new(); let mut deferred_content_events: Vec = Vec::new(); + let message = MessageMetadata { + id: message_id, + model, + estimated_input_tokens, + }; for event in &events { if matches!(event, ReducerEvent::WebSearch { .. }) { @@ -292,8 +302,7 @@ pub fn translate_stream_bytes_with_traffic( traffic, &mut message_started, &mut open_blocks, - message_id, - model, + &message, event, ) { continue; @@ -326,8 +335,7 @@ pub fn translate_stream_bytes_with_traffic( &mut out, traffic, &mut message_started, - message_id, - model, + &message, ); emit( &mut out, @@ -416,13 +424,12 @@ pub fn translate_stream_bytes_with_traffic( traffic, &mut message_started, &mut open_blocks, - message_id, - model, + &message, deferred, ); } - ensure_message_start(&mut out, traffic, &mut message_started, message_id, model); + ensure_message_start(&mut out, traffic, &mut message_started, &message); let mapped = map_codex_usage_to_anthropic(usage, Some(*web_search_requests)); emit( @@ -515,11 +522,15 @@ mod tests { ), ); let out = String::from_utf8( - translate_stream_bytes(upstream.as_bytes(), "msg_1", "gpt-5.5").unwrap(), + translate_stream_bytes_with_traffic(upstream.as_bytes(), "msg_1", "gpt-5.5", 321, None) + .unwrap(), ) .unwrap(); assert!(out.contains("message_start")); + assert!(out.contains(r#""input_tokens":321"#)); assert!(out.contains("text_delta")); + assert!(out.contains(r#""input_tokens":5"#)); + assert!(out.contains(r#""output_tokens":1"#)); assert!(out.contains("message_stop")); } diff --git a/tests/smoke_cutover.rs b/tests/smoke_cutover.rs index 955a3f11..717c8b0d 100644 --- a/tests/smoke_cutover.rs +++ b/tests/smoke_cutover.rs @@ -1012,6 +1012,10 @@ async fn smoke_codex_http_stream_retries_empty_completion() { body_text.contains("buffered stream retry ok"), "expected retried text in SSE body: {body_text}" ); + assert!( + !body_text.contains(r#""input_tokens":0"#), + "message_start should expose the request token estimate: {body_text}" + ); assert_eq!(attempts.load(std::sync::atomic::Ordering::SeqCst), 2); } @@ -1635,6 +1639,10 @@ async fn smoke_codex_websocket_stream_returns_delta_before_terminal() { assert!(read.is_ok(), "stream did not yield an early text delta"); let text = String::from_utf8_lossy(&collected); assert!(text.contains("early chunk"), "stream body: {text}"); + assert!( + !text.contains(r#""input_tokens":0"#), + "message_start should expose the request token estimate: {text}" + ); assert!( !text.contains("message_stop"), "stream finished too early: {text}"