From 1121347a4508875893507660c91d4df391160a73 Mon Sep 17 00:00:00 2001 From: YoVinchen Date: Thu, 4 Dec 2025 10:59:35 +0800 Subject: [PATCH] fix(proxy): resolve token parsing for OpenRouter streaming responses Problem: - OpenRouter and similar third-party services return streaming responses where input_tokens appear in message_delta instead of message_start - The previous implementation only extracted input_tokens from message_start, causing input_tokens to be recorded as 0 for these providers Changes: - streaming.rs: Add prompt_tokens field to Usage struct and include input_tokens in the transformed message_delta event when converting OpenAI format to Anthropic format - parser.rs: Update from_claude_stream_events() to handle input_tokens from both message_start (native Claude API) and message_delta (OpenRouter) - Use if-let pattern instead of direct unwrap for safer parsing - Only update input_tokens from message_delta if not already set - logger.rs: Adjust test parameters to match updated function signature Tests: - Add test_openrouter_stream_parsing() for OpenRouter format validation - Add test_native_claude_stream_parsing() for native Claude API validation --- src-tauri/src/proxy/providers/streaming.rs | 13 +++- src-tauri/src/proxy/usage/logger.rs | 1 + src-tauri/src/proxy/usage/parser.rs | 77 +++++++++++++++++++++- 3 files changed, 85 insertions(+), 6 deletions(-) diff --git a/src-tauri/src/proxy/providers/streaming.rs b/src-tauri/src/proxy/providers/streaming.rs index 48ecefa6c..ab310a4b5 100644 --- a/src-tauri/src/proxy/providers/streaming.rs +++ b/src-tauri/src/proxy/providers/streaming.rs @@ -53,8 +53,12 @@ struct DeltaFunction { arguments: Option, } +/// OpenAI 流式响应的 usage 信息(完整版) #[derive(Debug, Deserialize)] struct Usage { + #[serde(default)] + prompt_tokens: u32, + #[serde(default)] completion_tokens: u32, } @@ -278,15 +282,18 @@ pub fn create_anthropic_sse_stream( } let stop_reason = map_stop_reason(Some(finish_reason)); + // 构建 usage 信息,包含 input_tokens 和 output_tokens + let usage_json = chunk.usage.as_ref().map(|u| json!({ + "input_tokens": u.prompt_tokens, + "output_tokens": u.completion_tokens + })); let event = json!({ "type": "message_delta", "delta": { "stop_reason": stop_reason, "stop_sequence": null }, - "usage": chunk.usage.as_ref().map(|u| json!({ - "output_tokens": u.completion_tokens - })) + "usage": usage_json }); let sse_data = format!("event: message_delta\ndata: {}\n\n", serde_json::to_string(&event).unwrap_or_default()); diff --git a/src-tauri/src/proxy/usage/logger.rs b/src-tauri/src/proxy/usage/logger.rs index 6dccb8dcf..ce99b9bac 100644 --- a/src-tauri/src/proxy/usage/logger.rs +++ b/src-tauri/src/proxy/usage/logger.rs @@ -235,6 +235,7 @@ mod tests { usage, Decimal::from(1), 100, + None, 200, None, Some("claude".to_string()), diff --git a/src-tauri/src/proxy/usage/parser.rs b/src-tauri/src/proxy/usage/parser.rs index 65fe61a69..1aa99ae1e 100644 --- a/src-tauri/src/proxy/usage/parser.rs +++ b/src-tauri/src/proxy/usage/parser.rs @@ -59,7 +59,10 @@ impl TokenUsage { match event_type { "message_start" => { if let Some(msg_usage) = event.get("message").and_then(|m| m.get("usage")) { - usage.input_tokens = msg_usage.get("input_tokens")?.as_u64()? as u32; + // 从 message_start 获取 input_tokens(原生 Claude API) + if let Some(input) = msg_usage.get("input_tokens").and_then(|v| v.as_u64()) { + usage.input_tokens = input as u32; + } usage.cache_read_tokens = msg_usage .get("cache_read_input_tokens") .and_then(|v| v.as_u64()) @@ -74,8 +77,17 @@ impl TokenUsage { } "message_delta" => { if let Some(delta_usage) = event.get("usage") { - usage.output_tokens = - delta_usage.get("output_tokens")?.as_u64()? as u32; + // 从 message_delta 获取 output_tokens + if let Some(output) = delta_usage.get("output_tokens").and_then(|v| v.as_u64()) { + usage.output_tokens = output as u32; + } + // OpenRouter 转换后的流式响应:input_tokens 也在 message_delta 中 + // 如果 message_start 中没有 input_tokens,则从 message_delta 获取 + if usage.input_tokens == 0 { + if let Some(input) = delta_usage.get("input_tokens").and_then(|v| v.as_u64()) { + usage.input_tokens = input as u32; + } + } } } _ => {} @@ -454,4 +466,63 @@ mod tests { assert_eq!(usage.input_tokens, 0); assert_eq!(usage.cache_read_tokens, 200); } + + #[test] + fn test_openrouter_stream_parsing() { + // 测试 OpenRouter 转换后的流式响应解析 + // OpenRouter 流式响应经过转换后,input_tokens 在 message_delta 中 + let events = vec![ + json!({ + "type": "message_start", + "message": { + "usage": { + "input_tokens": 0, + "output_tokens": 0 + } + } + }), + json!({ + "type": "message_delta", + "delta": { + "stop_reason": "end_turn" + }, + "usage": { + "input_tokens": 150, + "output_tokens": 75 + } + }), + ]; + + let usage = TokenUsage::from_claude_stream_events(&events).unwrap(); + assert_eq!(usage.input_tokens, 150); + assert_eq!(usage.output_tokens, 75); + } + + #[test] + fn test_native_claude_stream_parsing() { + // 测试原生 Claude API 流式响应解析 + // 原生 Claude API 的 input_tokens 在 message_start 中 + let events = vec![ + json!({ + "type": "message_start", + "message": { + "usage": { + "input_tokens": 200, + "cache_read_input_tokens": 50 + } + } + }), + json!({ + "type": "message_delta", + "usage": { + "output_tokens": 100 + } + }), + ]; + + let usage = TokenUsage::from_claude_stream_events(&events).unwrap(); + assert_eq!(usage.input_tokens, 200); + assert_eq!(usage.output_tokens, 100); + assert_eq!(usage.cache_read_tokens, 50); + } }