feat(proxy): add streaming SSE transform and thinking parameter support

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
This commit is contained in:
YoVinchen
2025-12-01 17:11:23 +08:00
parent a35d112cd4
commit fef4c239c4
6 changed files with 489 additions and 69 deletions
+24
View File
@@ -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",
+2
View File
@@ -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"] }
+79 -57
View File
@@ -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());
+1
View File
@@ -17,6 +17,7 @@ mod claude;
mod codex;
mod gemini;
pub mod models;
pub mod streaming;
pub mod transform;
use crate::app_config::AppType;
+317
View File
@@ -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<StreamChoice>,
#[serde(default)]
usage: Option<Usage>,
}
#[derive(Debug, Deserialize)]
struct StreamChoice {
delta: Delta,
#[serde(default)]
finish_reason: Option<String>,
}
#[derive(Debug, Deserialize)]
struct Delta {
#[serde(default)]
content: Option<String>,
#[serde(default)]
reasoning: Option<String>, // OpenRouter 的推理内容
#[serde(default)]
tool_calls: Option<Vec<DeltaToolCall>>,
}
#[derive(Debug, Deserialize, Serialize)]
struct DeltaToolCall {
index: usize,
#[serde(default)]
id: Option<String>,
#[serde(rename = "type", default)]
call_type: Option<String>,
#[serde(default)]
function: Option<DeltaFunction>,
}
#[derive(Debug, Deserialize, Serialize)]
struct DeltaFunction {
#[serde(default)]
name: Option<String>,
#[serde(default)]
arguments: Option<String>,
}
#[derive(Debug, Deserialize)]
struct Usage {
completion_tokens: u32,
}
/// 创建 Anthropic SSE 流
pub fn create_anthropic_sse_stream(
stream: impl Stream<Item = Result<Bytes, reqwest::Error>> + Send + 'static,
) -> impl Stream<Item = Result<Bytes, std::io::Error>> + 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<String> = 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::<OpenAIStreamChunk>(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;
}
// 处理 reasoningthinking
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<String> {
finish_reason.map(|r| {
match r {
"tool_calls" => "tool_use",
"stop" => "end_turn",
"length" => "max_tokens",
_ => "end_turn",
}
.to_string()
})
}
+66 -12
View File
@@ -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<Value, ProxyError> {
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");
}
}