Skip to content
Draft
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
179 changes: 176 additions & 3 deletions src/proxy/translate/responses_stream.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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!({
Expand All @@ -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"}))));
};
Expand All @@ -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 {
Expand All @@ -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<String> {
if self.error_emitted {
return vec![];
}

// Responses API SSE format: "event: <type>\ndata: <json>" or just "data: <json>"
// We may receive "event:" and "data:" lines separately
if line.starts_with("event:") {
Expand Down Expand Up @@ -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())
Expand All @@ -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![],
}
Expand All @@ -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, reqwest::Error>(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() {
Expand Down Expand Up @@ -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"));
}
}