From fef4c239c41e56184b1a3f7dc66fe26b729256b9 Mon Sep 17 00:00:00 2001 From: YoVinchen Date: Mon, 1 Dec 2025 17:11:23 +0800 Subject: [PATCH] feat(proxy): add streaming SSE transform and thinking parameter support MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit New features: - Add OpenAI → Anthropic SSE streaming response transformation - Support thinking parameter detection for reasoning model selection - Add ANTHROPIC_REASONING_MODEL config option for extended thinking Changes: - streaming.rs: Implement SSE event parsing and Anthropic format conversion - transform.rs: Add thinking detection logic and reasoning model mapping - handlers.rs: Integrate streaming transform for OpenRouter compatibility - Cargo.toml: Add async-stream and bytes dependencies --- src-tauri/Cargo.lock | 24 ++ src-tauri/Cargo.toml | 2 + src-tauri/src/proxy/handlers.rs | 136 +++++---- src-tauri/src/proxy/providers/mod.rs | 1 + src-tauri/src/proxy/providers/streaming.rs | 317 +++++++++++++++++++++ src-tauri/src/proxy/providers/transform.rs | 78 ++++- 6 files changed, 489 insertions(+), 69 deletions(-) create mode 100644 src-tauri/src/proxy/providers/streaming.rs diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 46b1b4896..6173ece95 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -257,6 +257,28 @@ dependencies = [ "windows-sys 0.61.1", ] +[[package]] +name = "async-stream" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b5a71a6f37880a80d1d7f19efd781e4b5de42c88f0722cc13bcb6cc2cfe8476" +dependencies = [ + "async-stream-impl", + "futures-core", + "pin-project-lite", +] + +[[package]] +name = "async-stream-impl" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7c24de15d275a1ecfd47a380fb4d5ec9bfe0933f309ed5e705b775596a3574d" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.106", +] + [[package]] name = "async-task" version = "4.7.1" @@ -676,9 +698,11 @@ name = "cc-switch" version = "3.8.1" dependencies = [ "anyhow", + "async-stream", "auto-launch", "axum", "base64 0.22.1", + "bytes", "chrono", "dirs 5.0.1", "futures", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 089a94ea6..9c0e8db48 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -40,6 +40,8 @@ toml_edit = "0.22" reqwest = { version = "0.12", features = ["rustls-tls", "json", "stream"] } tokio = { version = "1", features = ["macros", "rt-multi-thread", "time", "sync"] } futures = "0.3" +async-stream = "0.3" +bytes = "1.5" axum = "0.7" tower = "0.4" tower-http = { version = "0.5", features = ["cors"] } diff --git a/src-tauri/src/proxy/handlers.rs b/src-tauri/src/proxy/handlers.rs index f677eecb0..6d24449c4 100644 --- a/src-tauri/src/proxy/handlers.rs +++ b/src-tauri/src/proxy/handlers.rs @@ -10,7 +10,7 @@ use super::{ ProxyError, }; use crate::app_config::AppType; -use axum::{extract::State, http::StatusCode, Json}; +use axum::{extract::State, http::StatusCode, response::IntoResponse, Json}; use serde_json::{json, Value}; /// 健康检查 @@ -76,75 +76,97 @@ pub async fn handle_messages( let status = response.status(); log::info!("[Claude] 上游响应状态: {status}"); - // 如果需要转换且是非流式请求,转换响应 - if needs_transform && !is_stream { - log::info!("[Claude] 开始转换响应 (OpenAI → Anthropic)"); + // 如果需要转换 + if needs_transform { + if is_stream { + // 流式响应转换 + log::info!("[Claude] 开始流式响应转换 (OpenAI SSE → Anthropic SSE)"); - let response_headers = response.headers().clone(); + let stream = response.bytes_stream(); + let sse_stream = super::providers::streaming::create_anthropic_sse_stream(stream); - // 读取响应体 - let body_bytes = response.bytes().await.map_err(|e| { - log::error!("[Claude] 读取响应体失败: {e}"); - ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) - })?; + let mut headers = axum::http::HeaderMap::new(); + headers.insert( + "Content-Type", + axum::http::HeaderValue::from_static("text/event-stream"), + ); + headers.insert( + "Cache-Control", + axum::http::HeaderValue::from_static("no-cache"), + ); + headers.insert( + "Connection", + axum::http::HeaderValue::from_static("keep-alive"), + ); - let body_str = String::from_utf8_lossy(&body_bytes); - log::info!("[Claude] OpenAI 响应长度: {} bytes", body_bytes.len()); - log::debug!("[Claude] OpenAI 原始响应: {body_str}"); + let body = axum::body::Body::from_stream(sse_stream); + return Ok((headers, body).into_response()); + } else { + // 非流式响应转换 + log::info!("[Claude] 开始转换响应 (OpenAI → Anthropic)"); - // 解析并转换 - let openai_response: Value = serde_json::from_slice(&body_bytes).map_err(|e| { - log::error!("[Claude] 解析 OpenAI 响应失败: {e}, body: {body_str}"); - ProxyError::TransformError(format!("Failed to parse OpenAI response: {e}")) - })?; + let response_headers = response.headers().clone(); - log::info!("[Claude] 解析 OpenAI 响应成功"); + // 读取响应体 + let body_bytes = response.bytes().await.map_err(|e| { + log::error!("[Claude] 读取响应体失败: {e}"); + ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) + })?; - let anthropic_response = transform::openai_to_anthropic(openai_response).map_err(|e| { - log::error!("[Claude] 转换响应失败: {e}"); - e - })?; + let body_str = String::from_utf8_lossy(&body_bytes); + log::info!("[Claude] OpenAI 响应长度: {} bytes", body_bytes.len()); + log::debug!("[Claude] OpenAI 原始响应: {body_str}"); - log::info!("[Claude] 转换响应成功"); - log::info!( - "[Claude] Anthropic 响应: {}", - serde_json::to_string(&anthropic_response).unwrap_or_default() - ); + // 解析并转换 + let openai_response: Value = serde_json::from_slice(&body_bytes).map_err(|e| { + log::error!("[Claude] 解析 OpenAI 响应失败: {e}, body: {body_str}"); + ProxyError::TransformError(format!("Failed to parse OpenAI response: {e}")) + })?; - // 构建响应 - let mut builder = axum::response::Response::builder().status(status); + log::info!("[Claude] 解析 OpenAI 响应成功"); - // 复制响应头(排除 content-length,因为内容已改变) - for (key, value) in response_headers.iter() { - if key.as_str().to_lowercase() != "content-length" - && key.as_str().to_lowercase() != "transfer-encoding" - { - builder = builder.header(key, value); + let anthropic_response = + transform::openai_to_anthropic(openai_response).map_err(|e| { + log::error!("[Claude] 转换响应失败: {e}"); + e + })?; + + log::info!("[Claude] 转换响应成功"); + log::debug!( + "[Claude] Anthropic 响应: {}", + serde_json::to_string(&anthropic_response).unwrap_or_default() + ); + + // 构建响应 + let mut builder = axum::response::Response::builder().status(status); + + // 复制响应头(排除 content-length,因为内容已改变) + for (key, value) in response_headers.iter() { + if key.as_str().to_lowercase() != "content-length" + && key.as_str().to_lowercase() != "transfer-encoding" + { + builder = builder.header(key, value); + } } + + builder = builder.header("content-type", "application/json"); + + let response_body = serde_json::to_vec(&anthropic_response).map_err(|e| { + log::error!("[Claude] 序列化响应失败: {e}"); + ProxyError::TransformError(format!("Failed to serialize response: {e}")) + })?; + + log::info!( + "[Claude] 返回转换后的响应, 长度: {} bytes", + response_body.len() + ); + + let body = axum::body::Body::from(response_body); + return Ok(builder.body(body).unwrap()); } - - builder = builder.header("content-type", "application/json"); - - let response_body = serde_json::to_vec(&anthropic_response).map_err(|e| { - log::error!("[Claude] 序列化响应失败: {e}"); - ProxyError::TransformError(format!("Failed to serialize response: {e}")) - })?; - - log::info!( - "[Claude] 返回转换后的响应, 长度: {} bytes", - response_body.len() - ); - - let body = axum::body::Body::from(response_body); - return Ok(builder.body(body).unwrap()); } - // 流式请求需要特殊处理 - if needs_transform && is_stream { - log::warn!("[Claude] OpenRouter 流式请求暂不支持完整转换,透传响应"); - } - - // 透传响应(直连 Anthropic 或流式请求) + // 透传响应(直连 Anthropic) log::info!("[Claude] 透传响应"); let mut builder = axum::response::Response::builder().status(response.status()); diff --git a/src-tauri/src/proxy/providers/mod.rs b/src-tauri/src/proxy/providers/mod.rs index a0633ae65..4b9720e33 100644 --- a/src-tauri/src/proxy/providers/mod.rs +++ b/src-tauri/src/proxy/providers/mod.rs @@ -17,6 +17,7 @@ mod claude; mod codex; mod gemini; pub mod models; +pub mod streaming; pub mod transform; use crate::app_config::AppType; diff --git a/src-tauri/src/proxy/providers/streaming.rs b/src-tauri/src/proxy/providers/streaming.rs new file mode 100644 index 000000000..5e6f2bf52 --- /dev/null +++ b/src-tauri/src/proxy/providers/streaming.rs @@ -0,0 +1,317 @@ +//! 流式响应转换模块 +//! +//! 实现 OpenAI SSE → Anthropic SSE 格式转换 + +use bytes::Bytes; +use futures::stream::{Stream, StreamExt}; +use serde::{Deserialize, Serialize}; +use serde_json::json; + +/// OpenAI 流式响应数据结构 +#[derive(Debug, Deserialize)] +struct OpenAIStreamChunk { + id: String, + model: String, + choices: Vec, + #[serde(default)] + usage: Option, +} + +#[derive(Debug, Deserialize)] +struct StreamChoice { + delta: Delta, + #[serde(default)] + finish_reason: Option, +} + +#[derive(Debug, Deserialize)] +struct Delta { + #[serde(default)] + content: Option, + #[serde(default)] + reasoning: Option, // OpenRouter 的推理内容 + #[serde(default)] + tool_calls: Option>, +} + +#[derive(Debug, Deserialize, Serialize)] +struct DeltaToolCall { + index: usize, + #[serde(default)] + id: Option, + #[serde(rename = "type", default)] + call_type: Option, + #[serde(default)] + function: Option, +} + +#[derive(Debug, Deserialize, Serialize)] +struct DeltaFunction { + #[serde(default)] + name: Option, + #[serde(default)] + arguments: Option, +} + +#[derive(Debug, Deserialize)] +struct Usage { + completion_tokens: u32, +} + +/// 创建 Anthropic SSE 流 +pub fn create_anthropic_sse_stream( + stream: impl Stream> + Send + 'static, +) -> impl Stream> + Send { + async_stream::stream! { + let mut buffer = String::new(); + let mut message_id = None; + let mut current_model = None; + let mut content_index = 0; + let mut has_sent_message_start = false; + let mut current_block_type: Option = None; + let mut tool_call_id = None; + + tokio::pin!(stream); + + while let Some(chunk) = stream.next().await { + match chunk { + Ok(bytes) => { + let text = String::from_utf8_lossy(&bytes); + buffer.push_str(&text); + + while let Some(pos) = buffer.find("\n\n") { + let line = buffer[..pos].to_string(); + buffer = buffer[pos + 2..].to_string(); + + if line.trim().is_empty() { + continue; + } + + for l in line.lines() { + if let Some(data) = l.strip_prefix("data: ") { + if data.trim() == "[DONE]" { + let event = json!({"type": "message_stop"}); + let sse_data = format!("event: message_stop\ndata: {}\n\n", + serde_json::to_string(&event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + continue; + } + + if let Ok(chunk) = serde_json::from_str::(data) { + if message_id.is_none() { + message_id = Some(chunk.id.clone()); + } + if current_model.is_none() { + current_model = Some(chunk.model.clone()); + } + + if let Some(choice) = chunk.choices.first() { + if !has_sent_message_start { + let event = json!({ + "type": "message_start", + "message": { + "id": message_id.clone().unwrap_or_default(), + "type": "message", + "role": "assistant", + "model": current_model.clone().unwrap_or_default(), + "usage": { + "input_tokens": 0, + "output_tokens": 0 + } + } + }); + let sse_data = format!("event: message_start\ndata: {}\n\n", + serde_json::to_string(&event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + has_sent_message_start = true; + } + + // 处理 reasoning(thinking) + if let Some(reasoning) = &choice.delta.reasoning { + if current_block_type.is_none() { + let event = json!({ + "type": "content_block_start", + "index": content_index, + "content_block": { + "type": "thinking", + "thinking": "" + } + }); + let sse_data = format!("event: content_block_start\ndata: {}\n\n", + serde_json::to_string(&event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + current_block_type = Some("thinking".to_string()); + } + + let event = json!({ + "type": "content_block_delta", + "index": content_index, + "delta": { + "type": "thinking_delta", + "thinking": reasoning + } + }); + let sse_data = format!("event: content_block_delta\ndata: {}\n\n", + serde_json::to_string(&event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + } + + // 处理文本内容 + if let Some(content) = &choice.delta.content { + if !content.is_empty() { + if current_block_type.as_deref() != Some("text") { + if current_block_type.is_some() { + let event = json!({ + "type": "content_block_stop", + "index": content_index + }); + let sse_data = format!("event: content_block_stop\ndata: {}\n\n", + serde_json::to_string(&event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + content_index += 1; + } + + let event = json!({ + "type": "content_block_start", + "index": content_index, + "content_block": { + "type": "text", + "text": "" + } + }); + let sse_data = format!("event: content_block_start\ndata: {}\n\n", + serde_json::to_string(&event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + current_block_type = Some("text".to_string()); + } + + let event = json!({ + "type": "content_block_delta", + "index": content_index, + "delta": { + "type": "text_delta", + "text": content + } + }); + let sse_data = format!("event: content_block_delta\ndata: {}\n\n", + serde_json::to_string(&event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + } + } + + // 处理工具调用 + if let Some(tool_calls) = &choice.delta.tool_calls { + for tool_call in tool_calls { + if let Some(id) = &tool_call.id { + if current_block_type.is_some() { + let event = json!({ + "type": "content_block_stop", + "index": content_index + }); + let sse_data = format!("event: content_block_stop\ndata: {}\n\n", + serde_json::to_string(&event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + content_index += 1; + } + + tool_call_id = Some(id.clone()); + } + + if let Some(function) = &tool_call.function { + if let Some(name) = &function.name { + let event = json!({ + "type": "content_block_start", + "index": content_index, + "content_block": { + "type": "tool_use", + "id": tool_call_id.clone().unwrap_or_default(), + "name": name + } + }); + let sse_data = format!("event: content_block_start\ndata: {}\n\n", + serde_json::to_string(&event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + current_block_type = Some("tool_use".to_string()); + } + + if let Some(args) = &function.arguments { + let event = json!({ + "type": "content_block_delta", + "index": content_index, + "delta": { + "type": "input_json_delta", + "partial_json": args + } + }); + let sse_data = format!("event: content_block_delta\ndata: {}\n\n", + serde_json::to_string(&event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + } + } + } + } + + // 处理 finish_reason + if let Some(finish_reason) = &choice.finish_reason { + if current_block_type.is_some() { + let event = json!({ + "type": "content_block_stop", + "index": content_index + }); + let sse_data = format!("event: content_block_stop\ndata: {}\n\n", + serde_json::to_string(&event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + } + + let stop_reason = map_stop_reason(Some(finish_reason)); + 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 + })) + }); + let sse_data = format!("event: message_delta\ndata: {}\n\n", + serde_json::to_string(&event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + } + } + } + } + } + } + } + Err(e) => { + log::error!("Stream error: {e}"); + let error_event = json!({ + "type": "error", + "error": { + "type": "stream_error", + "message": format!("Stream error: {e}") + } + }); + let sse_data = format!("event: error\ndata: {}\n\n", + serde_json::to_string(&error_event).unwrap_or_default()); + yield Ok(Bytes::from(sse_data)); + break; + } + } + } + } +} + +/// 映射停止原因 +fn map_stop_reason(finish_reason: Option<&str>) -> Option { + finish_reason.map(|r| { + match r { + "tool_calls" => "tool_use", + "stop" => "end_turn", + "length" => "max_tokens", + _ => "end_turn", + } + .to_string() + }) +} diff --git a/src-tauri/src/proxy/providers/transform.rs b/src-tauri/src/proxy/providers/transform.rs index eaabb3e57..a841750ed 100644 --- a/src-tauri/src/proxy/providers/transform.rs +++ b/src-tauri/src/proxy/providers/transform.rs @@ -8,14 +8,31 @@ use crate::proxy::error::ProxyError; use serde_json::{json, Value}; /// 从 Provider 配置中获取模型映射 -fn get_model_from_provider(model: &str, provider: &Provider) -> String { +fn get_model_from_provider(model: &str, provider: &Provider, body: &Value) -> String { let env = provider.settings_config.get("env"); - - // 根据请求的模型类型选择对应的配置模型 let model_lower = model.to_lowercase(); + // 检测 thinking 参数 + let has_thinking = body + .get("thinking") + .and_then(|v| v.as_object()) + .and_then(|o| o.get("type")) + .and_then(|t| t.as_str()) + == Some("enabled"); + if let Some(env) = env { - // 检查是否是 haiku 模型 + // 如果启用 thinking,优先使用推理模型 + if has_thinking { + if let Some(m) = env + .get("ANTHROPIC_REASONING_MODEL") + .and_then(|v| v.as_str()) + { + log::debug!("[Transform] 使用推理模型: {m}"); + return m.to_string(); + } + } + + // 根据模型类型选择配置模型 if model_lower.contains("haiku") { if let Some(m) = env .get("ANTHROPIC_DEFAULT_HAIKU_MODEL") @@ -24,7 +41,6 @@ fn get_model_from_provider(model: &str, provider: &Provider) -> String { return m.to_string(); } } - // 检查是否是 opus 模型 if model_lower.contains("opus") { if let Some(m) = env .get("ANTHROPIC_DEFAULT_OPUS_MODEL") @@ -33,7 +49,6 @@ fn get_model_from_provider(model: &str, provider: &Provider) -> String { return m.to_string(); } } - // 检查是否是 sonnet 模型 if model_lower.contains("sonnet") { if let Some(m) = env .get("ANTHROPIC_DEFAULT_SONNET_MODEL") @@ -48,7 +63,6 @@ fn get_model_from_provider(model: &str, provider: &Provider) -> String { } } - // 如果没有配置,返回原始模型名 model.to_string() } @@ -56,9 +70,9 @@ fn get_model_from_provider(model: &str, provider: &Provider) -> String { pub fn anthropic_to_openai(body: Value, provider: &Provider) -> Result { let mut result = json!({}); - // 模型映射:使用 Provider 配置中的模型 + // 模型映射:使用 Provider 配置中的模型(支持 thinking 参数) if let Some(model) = body.get("model").and_then(|m| m.as_str()) { - let mapped_model = get_model_from_provider(model, provider); + let mapped_model = get_model_from_provider(model, provider, &body); result["model"] = json!(mapped_model); } @@ -551,22 +565,23 @@ mod tests { #[test] fn test_model_mapping_from_provider() { let provider = create_openrouter_provider(); + let body = json!({"model": "test"}); // sonnet 模型 assert_eq!( - get_model_from_provider("claude-sonnet-4-5-20250929", &provider), + get_model_from_provider("claude-sonnet-4-5-20250929", &provider, &body), "anthropic/claude-sonnet-4.5" ); // haiku 模型 assert_eq!( - get_model_from_provider("claude-haiku-4-5-20250929", &provider), + get_model_from_provider("claude-haiku-4-5-20250929", &provider, &body), "anthropic/claude-haiku-4.5" ); // opus 模型 assert_eq!( - get_model_from_provider("claude-opus-4-5", &provider), + get_model_from_provider("claude-opus-4-5", &provider, &body), "anthropic/claude-opus-4.5" ); } @@ -583,4 +598,43 @@ mod tests { let result = anthropic_to_openai(input, &provider).unwrap(); assert_eq!(result["model"], "anthropic/claude-sonnet-4.5"); } + + #[test] + fn test_thinking_parameter_detection() { + let mut provider = create_openrouter_provider(); + // 添加推理模型配置 + if let Some(env) = provider.settings_config.get_mut("env") { + env["ANTHROPIC_REASONING_MODEL"] = json!("anthropic/claude-sonnet-4.5:extended"); + } + + let input = json!({ + "model": "claude-sonnet-4-5", + "max_tokens": 1024, + "thinking": {"type": "enabled"}, + "messages": [{"role": "user", "content": "Solve this problem"}] + }); + + let result = anthropic_to_openai(input, &provider).unwrap(); + // 应该使用推理模型 + assert_eq!(result["model"], "anthropic/claude-sonnet-4.5:extended"); + } + + #[test] + fn test_thinking_parameter_disabled() { + let mut provider = create_openrouter_provider(); + if let Some(env) = provider.settings_config.get_mut("env") { + env["ANTHROPIC_REASONING_MODEL"] = json!("anthropic/claude-sonnet-4.5:extended"); + } + + let input = json!({ + "model": "claude-sonnet-4-5", + "max_tokens": 1024, + "thinking": {"type": "disabled"}, + "messages": [{"role": "user", "content": "Hello"}] + }); + + let result = anthropic_to_openai(input, &provider).unwrap(); + // 应该使用普通模型 + assert_eq!(result["model"], "anthropic/claude-sonnet-4.5"); + } }