diff --git a/src-tauri/src/proxy/forwarder.rs b/src-tauri/src/proxy/forwarder.rs index b9b8fe8e5..f5b27d482 100644 --- a/src-tauri/src/proxy/forwarder.rs +++ b/src-tauri/src/proxy/forwarder.rs @@ -17,6 +17,16 @@ use std::sync::Arc; use std::time::{Duration, Instant}; use tokio::sync::RwLock; +pub struct ForwardResult { + pub response: Response, + pub provider: Provider, +} + +pub struct ForwardError { + pub error: ProxyError, + pub provider: Option, +} + pub struct RequestForwarder { client: Client, /// 共享的 ProviderRouter(持有熔断器状态) @@ -134,13 +144,16 @@ impl RequestForwarder { body: Value, headers: axum::http::HeaderMap, providers: Vec, - ) -> Result { + ) -> Result { // 获取适配器 let adapter = get_adapter(app_type); let app_type_str = app_type.as_str(); if providers.is_empty() { - return Err(ProxyError::NoAvailableProvider); + return Err(ForwardError { + error: ProxyError::NoAvailableProvider, + provider: None, + }); } log::info!( @@ -150,6 +163,7 @@ impl RequestForwarder { ); let mut last_error = None; + let mut last_provider = None; let mut attempted_providers = 0usize; // 依次尝试每个供应商 @@ -269,7 +283,10 @@ impl RequestForwarder { latency ); - return Ok(response); + return Ok(ForwardResult { + response, + provider: provider.clone(), + }); } Err(e) => { let latency = start.elapsed().as_millis() as u64; @@ -310,6 +327,7 @@ impl RequestForwarder { ); last_error = Some(e); + last_provider = Some(provider.clone()); // 继续尝试下一个供应商 continue; } @@ -331,7 +349,10 @@ impl RequestForwarder { provider.name, e ); - return Err(e); + return Err(ForwardError { + error: e, + provider: Some(provider.clone()), + }); } } } @@ -349,7 +370,10 @@ impl RequestForwarder { (status.success_requests as f32 / status.total_requests as f32) * 100.0; } } - return Err(ProxyError::NoAvailableProvider); + return Err(ForwardError { + error: ProxyError::NoAvailableProvider, + provider: None, + }); } // 所有供应商都失败了 @@ -369,7 +393,10 @@ impl RequestForwarder { providers.len() ); - Err(last_error.unwrap_or(ProxyError::MaxRetriesExceeded)) + Err(ForwardError { + error: last_error.unwrap_or(ProxyError::MaxRetriesExceeded), + provider: last_provider, + }) } /// 转发单个请求(使用适配器) diff --git a/src-tauri/src/proxy/handlers.rs b/src-tauri/src/proxy/handlers.rs index faeefec01..4df94ea15 100644 --- a/src-tauri/src/proxy/handlers.rs +++ b/src-tauri/src/proxy/handlers.rs @@ -61,27 +61,16 @@ pub async fn handle_messages( headers: axum::http::HeaderMap, Json(body): Json, ) -> Result { - let ctx = RequestContext::new(&state, &body, AppType::Claude, "Claude", "claude").await?; - - // 检查是否需要格式转换(OpenRouter 等中转服务) - let adapter = get_adapter(&AppType::Claude); - let needs_transform = adapter.needs_transform(&ctx.provider); + let mut ctx = RequestContext::new(&state, &body, AppType::Claude, "Claude", "claude").await?; let is_stream = body .get("stream") .and_then(|s| s.as_bool()) .unwrap_or(false); - log::info!( - "[Claude] Provider: {}, needs_transform: {}, is_stream: {}", - ctx.provider.name, - needs_transform, - is_stream - ); - // 转发请求 let forwarder = ctx.create_forwarder(&state); - let response = match forwarder + let result = match forwarder .forward_with_retry( &AppType::Claude, "/v1/messages", @@ -91,13 +80,30 @@ pub async fn handle_messages( ) .await { - Ok(resp) => resp, - Err(e) => { - log_forward_error(&state, &ctx, is_stream, &e); - return Err(e); + Ok(result) => result, + Err(mut err) => { + if let Some(provider) = err.provider.take() { + ctx.provider = provider; + } + log_forward_error(&state, &ctx, is_stream, &err.error); + return Err(err.error); } }; + ctx.provider = result.provider; + let response = result.response; + + // 检查是否需要格式转换(OpenRouter 等中转服务) + let adapter = get_adapter(&AppType::Claude); + let needs_transform = adapter.needs_transform(&ctx.provider); + + log::info!( + "[Claude] Provider: {}, needs_transform: {}, is_stream: {}", + ctx.provider.name, + needs_transform, + is_stream + ); + let status = response.status(); log::info!("[Claude] 上游响应状态: {status}"); @@ -295,7 +301,7 @@ pub async fn handle_chat_completions( ) -> Result { log::info!("[Codex] ====== /v1/chat/completions 请求开始 ======"); - let ctx = RequestContext::new(&state, &body, AppType::Codex, "Codex", "codex").await?; + let mut ctx = RequestContext::new(&state, &body, AppType::Codex, "Codex", "codex").await?; let is_stream = body .get("stream") @@ -309,7 +315,7 @@ pub async fn handle_chat_completions( ); let forwarder = ctx.create_forwarder(&state); - let response = match forwarder + let result = match forwarder .forward_with_retry( &AppType::Codex, "/v1/chat/completions", @@ -319,13 +325,19 @@ pub async fn handle_chat_completions( ) .await { - Ok(resp) => resp, - Err(e) => { - log_forward_error(&state, &ctx, is_stream, &e); - return Err(e); + Ok(result) => result, + Err(mut err) => { + if let Some(provider) = err.provider.take() { + ctx.provider = provider; + } + log_forward_error(&state, &ctx, is_stream, &err.error); + return Err(err.error); } }; + ctx.provider = result.provider; + let response = result.response; + log::info!("[Codex] 上游响应状态: {}", response.status()); process_response(response, &ctx, &state, &OPENAI_PARSER_CONFIG).await @@ -337,7 +349,7 @@ pub async fn handle_responses( headers: axum::http::HeaderMap, Json(body): Json, ) -> Result { - let ctx = RequestContext::new(&state, &body, AppType::Codex, "Codex", "codex").await?; + let mut ctx = RequestContext::new(&state, &body, AppType::Codex, "Codex", "codex").await?; let is_stream = body .get("stream") @@ -345,7 +357,7 @@ pub async fn handle_responses( .unwrap_or(false); let forwarder = ctx.create_forwarder(&state); - let response = match forwarder + let result = match forwarder .forward_with_retry( &AppType::Codex, "/v1/responses", @@ -355,13 +367,19 @@ pub async fn handle_responses( ) .await { - Ok(resp) => resp, - Err(e) => { - log_forward_error(&state, &ctx, is_stream, &e); - return Err(e); + Ok(result) => result, + Err(mut err) => { + if let Some(provider) = err.provider.take() { + ctx.provider = provider; + } + log_forward_error(&state, &ctx, is_stream, &err.error); + return Err(err.error); } }; + ctx.provider = result.provider; + let response = result.response; + log::info!("[Codex] 上游响应状态: {}", response.status()); process_response(response, &ctx, &state, &CODEX_PARSER_CONFIG).await @@ -379,7 +397,7 @@ pub async fn handle_gemini( Json(body): Json, ) -> Result { // Gemini 的模型名称在 URI 中 - let ctx = RequestContext::new(&state, &body, AppType::Gemini, "Gemini", "gemini") + let mut ctx = RequestContext::new(&state, &body, AppType::Gemini, "Gemini", "gemini") .await? .with_model_from_uri(&uri); @@ -397,7 +415,7 @@ pub async fn handle_gemini( .unwrap_or(false); let forwarder = ctx.create_forwarder(&state); - let response = match forwarder + let result = match forwarder .forward_with_retry( &AppType::Gemini, endpoint, @@ -407,13 +425,19 @@ pub async fn handle_gemini( ) .await { - Ok(resp) => resp, - Err(e) => { - log_forward_error(&state, &ctx, is_stream, &e); - return Err(e); + Ok(result) => result, + Err(mut err) => { + if let Some(provider) = err.provider.take() { + ctx.provider = provider; + } + log_forward_error(&state, &ctx, is_stream, &err.error); + return Err(err.error); } }; + ctx.provider = result.provider; + let response = result.response; + log::info!("[Gemini] 上游响应状态: {}", response.status()); process_response(response, &ctx, &state, &GEMINI_PARSER_CONFIG).await