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); + } }