diff --git a/src/proxy/translate/responses_stream.rs b/src/proxy/translate/responses_stream.rs index b443f3e..7d20009 100644 --- a/src/proxy/translate/responses_stream.rs +++ b/src/proxy/translate/responses_stream.rs @@ -64,6 +64,10 @@ where } } + if state.error_emitted { + return; + } + // Finalize: close any open block and send message_delta + message_stop if state.block_started { yield Ok(Bytes::from(format_sse("content_block_stop", &json!({ @@ -76,7 +80,12 @@ where yield Ok(Bytes::from(format_sse("message_delta", &json!({ "type": "message_delta", "delta": {"stop_reason": stop_reason, "stop_sequence": null}, - "usage": {"output_tokens": state.output_tokens} + "usage": { + "input_tokens": state.input_tokens, + "cache_creation_input_tokens": state.cache_creation_input_tokens, + "cache_read_input_tokens": state.cache_read_input_tokens, + "output_tokens": state.output_tokens, + } })))); yield Ok(Bytes::from(format_sse("message_stop", &json!({"type": "message_stop"})))); }; @@ -90,7 +99,11 @@ struct ResponsesStreamState { block_started: bool, has_tool_use: bool, stop_reason: String, + input_tokens: u64, + cache_creation_input_tokens: u64, + cache_read_input_tokens: u64, output_tokens: u64, + error_emitted: bool, } impl ResponsesStreamState { @@ -101,11 +114,19 @@ impl ResponsesStreamState { block_started: false, has_tool_use: false, stop_reason: "end_turn".to_string(), + input_tokens: 0, + cache_creation_input_tokens: 0, + cache_read_input_tokens: 0, output_tokens: 0, + error_emitted: false, } } fn process_line(&mut self, line: &str) -> Vec { + if self.error_emitted { + return vec![]; + } + // Responses API SSE format: "event: \ndata: " or just "data: " // We may receive "event:" and "data:" lines separately if line.starts_with("event:") { @@ -264,6 +285,19 @@ impl ResponsesStreamState { // Extract usage from the completed response if let Some(resp) = json.get("response") { if let Some(usage) = resp.get("usage") { + let total_input_tokens = usage + .get("input_tokens") + .and_then(Value::as_u64) + .unwrap_or(0); + let cached_input_tokens = usage + .get("input_tokens_details") + .and_then(|details| details.get("cached_tokens")) + .and_then(Value::as_u64) + .unwrap_or(0) + .min(total_input_tokens); + self.input_tokens = total_input_tokens.saturating_sub(cached_input_tokens); + self.cache_creation_input_tokens = 0; + self.cache_read_input_tokens = cached_input_tokens; self.output_tokens = usage .get("output_tokens") .and_then(|v| v.as_u64()) @@ -281,8 +315,31 @@ impl ResponsesStreamState { vec![] } "response.failed" => { - self.stop_reason = "end_turn".to_string(); - vec![] + let error = json + .get("response") + .and_then(|response| response.get("error")); + let code = error + .and_then(|error| error.get("code").or_else(|| error.get("type"))) + .and_then(Value::as_str) + .unwrap_or("api_error"); + let error_type = if code.eq_ignore_ascii_case("context_length_exceeded") { + "invalid_request_error" + } else { + "api_error" + }; + let message = error + .and_then(|error| error.get("message")) + .and_then(Value::as_str) + .filter(|message| !message.is_empty()) + .unwrap_or("Upstream response failed"); + self.error_emitted = true; + vec![format_sse( + "error", + &json!({ + "type": "error", + "error": {"type": error_type, "message": message}, + }), + )] } _ => vec![], } @@ -292,6 +349,37 @@ impl ResponsesStreamState { #[cfg(test)] mod tests { use super::*; + use futures::stream; + + async fn translate(lines: &[&str]) -> String { + let input = lines.join("\n\n") + "\n\n"; + let source = stream::iter(vec![Ok::(Bytes::from(input))]); + let mut translated = translate_responses_stream(source, ToolNameMap::new()); + let mut result = String::new(); + while let Some(chunk) = translated.next().await { + result.push_str(std::str::from_utf8(&chunk.unwrap()).unwrap()); + } + result + } + + fn accumulated_usage(output: &str) -> Value { + let mut usage = serde_json::Map::new(); + for frame in output.split("\n\n") { + let Some(data) = frame.lines().find_map(|line| line.strip_prefix("data: ")) else { + continue; + }; + let payload: Value = serde_json::from_str(data).unwrap(); + let event_usage = match payload.get("type").and_then(Value::as_str) { + Some("message_start") => payload.pointer("/message/usage"), + Some("message_delta") => payload.get("usage"), + _ => None, + }; + if let Some(event_usage) = event_usage.and_then(Value::as_object) { + usage.extend(event_usage.clone()); + } + } + Value::Object(usage) + } #[test] fn test_text_delta() { @@ -340,6 +428,91 @@ mod tests { state.process_line( r#"data: {"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":100,"output_tokens":50,"total_tokens":150}}}"#, ); + assert_eq!(state.input_tokens, 100); + assert_eq!(state.cache_creation_input_tokens, 0); + assert_eq!(state.cache_read_input_tokens, 0); assert_eq!(state.output_tokens, 50); } + + #[tokio::test] + async fn test_completed_usage_is_accumulated_from_final_delta() { + let output = translate(&[ + r#"data: {"type":"response.output_text.delta","delta":"Hello"}"#, + r#"data: {"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":12,"output_tokens":3}}}"#, + ]) + .await; + assert_eq!( + accumulated_usage(&output), + json!({ + "input_tokens": 12, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "output_tokens": 3, + }) + ); + } + + #[tokio::test] + async fn test_cached_input_tokens_are_not_double_counted() { + let output = translate(&[ + r#"data: {"type":"response.output_text.delta","delta":"Hello"}"#, + r#"data: {"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":12,"input_tokens_details":{"cached_tokens":5},"output_tokens":3}}}"#, + ]) + .await; + assert_eq!( + accumulated_usage(&output), + json!({ + "input_tokens": 7, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 5, + "output_tokens": 3, + }) + ); + } + + #[tokio::test] + async fn test_malformed_usage_fields_degrade_to_zero() { + let output = translate(&[ + r#"data: {"type":"response.output_text.delta","delta":"Hello"}"#, + r#"data: {"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":"unknown","input_tokens_details":{"cached_tokens":5},"output_tokens":null}}}"#, + ]) + .await; + assert_eq!( + accumulated_usage(&output), + json!({ + "input_tokens": 0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "output_tokens": 0, + }) + ); + } + + #[tokio::test] + async fn test_context_overflow_emits_only_invalid_request_error() { + let output = translate(&[ + r#"data: {"type":"response.output_text.delta","delta":"partial"}"#, + r#"data: {"type":"response.failed","response":{"status":"failed","error":{"code":"context_length_exceeded","message":"Context window exceeded"}}}"#, + ]) + .await; + assert_eq!(output.matches("event: error").count(), 1); + assert!(output.contains("invalid_request_error")); + assert!(output.contains("Context window exceeded")); + assert!(!output.contains("content_block_stop")); + assert!(!output.contains("message_delta")); + assert!(!output.contains("message_stop")); + } + + #[tokio::test] + async fn test_unrelated_failure_remains_api_error_without_normal_stop() { + let output = translate(&[ + r#"data: {"type":"response.failed","response":{"status":"failed","error":{"code":"server_error","message":"Upstream failed"}}}"#, + ]) + .await; + assert_eq!(output.matches("event: error").count(), 1); + assert!(output.contains("api_error")); + assert!(output.contains("Upstream failed")); + assert!(!output.contains("message_delta")); + assert!(!output.contains("message_stop")); + } }