From d57da0b700c48b11dedf2f31c295c227088e868e Mon Sep 17 00:00:00 2001 From: YoVinchen Date: Mon, 8 Dec 2025 11:02:47 +0800 Subject: [PATCH] feat(proxy): add failover Tauri commands and integrate with forwarder Expose failover functionality to frontend: - Add Tauri commands: get_proxy_targets, set_proxy_target, get_provider_health, reset_circuit_breaker, get/update_circuit_breaker_config, get_circuit_breaker_stats - Register all new commands in lib.rs invoke handler - Update forwarder with improved error handling and logging - Integrate ProviderRouter with proxy server startup - Add provider health tracking in request handlers --- src-tauri/src/commands/proxy.rs | 114 ++++++++++++++++++++++++++ src-tauri/src/lib.rs | 8 ++ src-tauri/src/proxy/forwarder.rs | 132 +++++++++++++++++++++++-------- src-tauri/src/proxy/handlers.rs | 4 + src-tauri/src/proxy/router.rs | 1 + src-tauri/src/proxy/server.rs | 24 +++--- 6 files changed, 239 insertions(+), 44 deletions(-) diff --git a/src-tauri/src/commands/proxy.rs b/src-tauri/src/commands/proxy.rs index 7f30eea0c..79a8e49db 100644 --- a/src-tauri/src/commands/proxy.rs +++ b/src-tauri/src/commands/proxy.rs @@ -2,7 +2,9 @@ //! //! 提供前端调用的 API 接口 +use crate::provider::Provider; use crate::proxy::types::*; +use crate::proxy::{CircuitBreakerConfig, CircuitBreakerStats}; use crate::store::AppState; /// 启动代理服务器 @@ -45,3 +47,115 @@ pub async fn update_proxy_config( pub async fn is_proxy_running(state: tauri::State<'_, AppState>) -> Result { Ok(state.proxy_service.is_running().await) } + +// ==================== 故障转移相关命令 ==================== + +/// 获取代理目标列表 +#[tauri::command] +pub async fn get_proxy_targets( + state: tauri::State<'_, AppState>, + app_type: String, +) -> Result, String> { + let db = &state.db; + db.get_proxy_targets(&app_type) + .await + .map_err(|e| e.to_string()) + .map(|providers| providers.into_values().collect()) +} + +/// 设置代理目标 +#[tauri::command] +pub async fn set_proxy_target( + state: tauri::State<'_, AppState>, + provider_id: String, + app_type: String, + enabled: bool, +) -> Result<(), String> { + let db = &state.db; + + // 设置代理目标状态 + db.set_proxy_target(&provider_id, &app_type, enabled) + .await + .map_err(|e| e.to_string())?; + + // 如果是禁用代理目标,重置健康状态 + if !enabled { + log::info!( + "Resetting health status for provider {provider_id} (app: {app_type}) after disabling proxy target" + ); + if let Err(e) = db.reset_provider_health(&provider_id, &app_type).await { + log::warn!("Failed to reset provider health: {e}"); + } + } + + Ok(()) +} + +/// 获取供应商健康状态 +#[tauri::command] +pub async fn get_provider_health( + state: tauri::State<'_, AppState>, + provider_id: String, + app_type: String, +) -> Result { + let db = &state.db; + db.get_provider_health(&provider_id, &app_type) + .await + .map_err(|e| e.to_string()) +} + +/// 重置熔断器 +#[tauri::command] +pub async fn reset_circuit_breaker( + state: tauri::State<'_, AppState>, + provider_id: String, + app_type: String, +) -> Result<(), String> { + // 重置数据库健康状态 + let db = &state.db; + db.update_provider_health(&provider_id, &app_type, true, None) + .await + .map_err(|e| e.to_string())?; + + // 注意:熔断器状态在内存中,重启代理服务器后会重置 + // 如果代理服务器正在运行,需要通知它重置熔断器 + // 目前先通过数据库重置健康状态,熔断器会在下次超时后自动尝试半开 + + Ok(()) +} + +/// 获取熔断器配置 +#[tauri::command] +pub async fn get_circuit_breaker_config( + state: tauri::State<'_, AppState>, +) -> Result { + let db = &state.db; + db.get_circuit_breaker_config() + .await + .map_err(|e| e.to_string()) +} + +/// 更新熔断器配置 +#[tauri::command] +pub async fn update_circuit_breaker_config( + state: tauri::State<'_, AppState>, + config: CircuitBreakerConfig, +) -> Result<(), String> { + let db = &state.db; + db.update_circuit_breaker_config(&config) + .await + .map_err(|e| e.to_string()) +} + +/// 获取熔断器统计信息(仅当代理服务器运行时) +#[tauri::command] +pub async fn get_circuit_breaker_stats( + state: tauri::State<'_, AppState>, + provider_id: String, + app_type: String, +) -> Result, String> { + // 这个功能需要访问运行中的代理服务器的内存状态 + // 目前先返回 None,后续可以通过 ProxyService 暴露接口来实现 + let _ = (state, provider_id, app_type); + Ok(None) +} diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index af735ecdb..4d6a36ab4 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -650,6 +650,14 @@ pub fn run() { commands::get_proxy_config, commands::update_proxy_config, commands::is_proxy_running, + // Proxy failover commands + commands::get_proxy_targets, + commands::set_proxy_target, + commands::get_provider_health, + commands::reset_circuit_breaker, + commands::get_circuit_breaker_config, + commands::update_circuit_breaker_config, + commands::get_circuit_breaker_stats, // Usage statistics commands::get_usage_summary, commands::get_usage_trends, diff --git a/src-tauri/src/proxy/forwarder.rs b/src-tauri/src/proxy/forwarder.rs index 245759d04..b4226ba00 100644 --- a/src-tauri/src/proxy/forwarder.rs +++ b/src-tauri/src/proxy/forwarder.rs @@ -4,8 +4,8 @@ use super::{ error::*, + provider_router::ProviderRouter as NewProviderRouter, providers::{get_adapter, ProviderAdapter}, - router::ProviderRouter, types::ProxyStatus, ProxyError, }; @@ -18,9 +18,11 @@ use tokio::sync::RwLock; pub struct RequestForwarder { client: Client, - router: ProviderRouter, + router: Arc, + #[allow(dead_code)] max_retries: u8, status: Arc>, + current_providers: Arc>>, } impl RequestForwarder { @@ -29,6 +31,7 @@ impl RequestForwarder { timeout_secs: u64, max_retries: u8, status: Arc>, + current_providers: Arc>>, ) -> Self { let mut client_builder = Client::builder(); if timeout_secs > 0 { @@ -41,13 +44,14 @@ impl RequestForwarder { Self { client, - router: ProviderRouter::new(db), + router: Arc::new(NewProviderRouter::new(db)), max_retries, status, + current_providers, } } - /// 转发请求(带重试和故障转移) + /// 转发请求(带故障转移) pub async fn forward_with_retry( &self, app_type: &AppType, @@ -55,21 +59,39 @@ impl RequestForwarder { body: Value, headers: axum::http::HeaderMap, ) -> Result { - let mut failed_ids = Vec::new(); - let mut failover_happened = false; - // 获取适配器 let adapter = get_adapter(app_type); + let app_type_str = app_type.as_str(); - for attempt in 0..self.max_retries { - // 选择Provider - let provider = self.router.select_provider(app_type, &failed_ids).await?; + // 使用新的 ProviderRouter 选择所有可用供应商 + let providers = self + .router + .select_providers(app_type_str) + .await + .map_err(|e| ProxyError::DatabaseError(e.to_string()))?; - log::debug!( - "尝试 {} - 使用Provider: {} ({})", + if providers.is_empty() { + return Err(ProxyError::NoAvailableProvider); + } + + log::info!( + "[{}] 故障转移链: {} 个可用供应商", + app_type_str, + providers.len() + ); + + let mut last_error = None; + let mut failover_happened = false; + + // 依次尝试每个供应商 + for (attempt, provider) in providers.iter().enumerate() { + log::info!( + "[{}] 尝试 {}/{} - 使用Provider: {} (sort_index: {})", + app_type_str, attempt + 1, + providers.len(), provider.name, - provider.id + provider.sort_index.unwrap_or(999999) ); // 更新状态中的当前Provider信息 @@ -88,16 +110,29 @@ impl RequestForwarder { // 转发请求 match self - .forward(&provider, endpoint, &body, &headers, adapter.as_ref()) + .forward(provider, endpoint, &body, &headers, adapter.as_ref()) .await { Ok(response) => { - let _latency = start.elapsed().as_millis() as u64; + let latency = start.elapsed().as_millis() as u64; - // 成功:更新健康状态 - self.router - .update_health(&provider, app_type, true, None) - .await; + // 成功:记录成功并更新熔断器 + if let Err(e) = self + .router + .record_result(&provider.id, app_type_str, true, None) + .await + { + log::warn!("Failed to record success: {e}"); + } + + // 更新当前应用类型使用的 provider + { + let mut current_providers = self.current_providers.write().await; + current_providers.insert( + app_type_str.to_string(), + (provider.id.clone(), provider.name.clone()), + ); + } // 更新成功统计 { @@ -106,6 +141,12 @@ impl RequestForwarder { status.last_error = None; if failover_happened { status.failover_count += 1; + log::info!( + "[{}] 故障转移成功!切换到 Provider: {} (耗时: {}ms)", + app_type_str, + provider.name, + latency + ); } // 重新计算成功率 if status.total_requests > 0 { @@ -115,23 +156,33 @@ impl RequestForwarder { } } + log::info!( + "[{}] 请求成功 - Provider: {} - {}ms", + app_type_str, + provider.name, + latency + ); + return Ok(response); } Err(e) => { let latency = start.elapsed().as_millis() as u64; - // 失败:分类错误 + // 失败:记录失败并更新熔断器 + if let Err(record_err) = self + .router + .record_result(&provider.id, app_type_str, false, Some(e.to_string())) + .await + { + log::warn!("Failed to record failure: {record_err}"); + } + + // 分类错误 let category = self.categorize_proxy_error(&e); match category { ErrorCategory::Retryable => { - // 可重试:更新健康状态,添加到失败列表 - self.router - .update_health(&provider, app_type, false, Some(e.to_string())) - .await; - failed_ids.push(provider.id.clone()); - - // 更新错误信息 + // 可重试:更新错误信息,继续尝试下一个供应商 { let mut status = self.status.write().await; status.last_error = @@ -139,15 +190,19 @@ impl RequestForwarder { } log::warn!( - "请求失败(可重试): Provider {} - {} - {}ms", + "[{}] Provider {} 失败(可重试): {} - {}ms", + app_type_str, provider.name, e, latency ); + + last_error = Some(e); + // 继续尝试下一个供应商 continue; } ErrorCategory::NonRetryable | ErrorCategory::ClientAbort => { - // 不可重试:更新失败统计并返回 + // 不可重试:直接返回错误 { let mut status = self.status.write().await; status.failed_requests += 1; @@ -158,7 +213,12 @@ impl RequestForwarder { * 100.0; } } - log::error!("请求失败(不可重试): {e}"); + log::error!( + "[{}] Provider {} 失败(不可重试): {}", + app_type_str, + provider.name, + e + ); return Err(e); } } @@ -166,18 +226,24 @@ impl RequestForwarder { } } - // 所有重试都失败 + // 所有供应商都失败了 { let mut status = self.status.write().await; status.failed_requests += 1; - status.last_error = Some("已达到最大重试次数".to_string()); + status.last_error = Some("所有供应商都失败".to_string()); if status.total_requests > 0 { status.success_rate = (status.success_requests as f32 / status.total_requests as f32) * 100.0; } } - Err(ProxyError::MaxRetriesExceeded) + log::error!( + "[{}] 所有 {} 个供应商都失败了", + app_type_str, + providers.len() + ); + + Err(last_error.unwrap_or(ProxyError::MaxRetriesExceeded)) } /// 转发单个请求(使用适配器) diff --git a/src-tauri/src/proxy/handlers.rs b/src-tauri/src/proxy/handlers.rs index 54ad0cce2..679fd2a73 100644 --- a/src-tauri/src/proxy/handlers.rs +++ b/src-tauri/src/proxy/handlers.rs @@ -322,6 +322,7 @@ pub async fn handle_messages( config.request_timeout, config.max_retries, state.status.clone(), + state.current_providers.clone(), ); let response = forwarder @@ -641,6 +642,7 @@ pub async fn handle_gemini( config.request_timeout, config.max_retries, state.status.clone(), + state.current_providers.clone(), ); // 提取完整的路径和查询参数 @@ -806,6 +808,7 @@ pub async fn handle_responses( config.request_timeout, config.max_retries, state.status.clone(), + state.current_providers.clone(), ); let response = forwarder @@ -985,6 +988,7 @@ pub async fn handle_chat_completions( config.request_timeout, config.max_retries, state.status.clone(), + state.current_providers.clone(), ); let response = forwarder diff --git a/src-tauri/src/proxy/router.rs b/src-tauri/src/proxy/router.rs index 47efa952b..c59c99bf8 100644 --- a/src-tauri/src/proxy/router.rs +++ b/src-tauri/src/proxy/router.rs @@ -57,6 +57,7 @@ impl ProviderRouter { } /// 更新Provider健康状态(保留接口但不影响选择) + #[allow(dead_code)] pub async fn update_health( &self, _provider: &Provider, diff --git a/src-tauri/src/proxy/server.rs b/src-tauri/src/proxy/server.rs index 10932c380..6da189c36 100644 --- a/src-tauri/src/proxy/server.rs +++ b/src-tauri/src/proxy/server.rs @@ -20,6 +20,8 @@ pub struct ProxyState { pub config: Arc>, pub status: Arc>, pub start_time: Arc>>, + /// 每个应用类型当前使用的 provider (app_type -> (provider_id, provider_name)) + pub current_providers: Arc>>, } /// 代理HTTP服务器 @@ -36,6 +38,7 @@ impl ProxyServer { config: Arc::new(RwLock::new(config.clone())), status: Arc::new(RwLock::new(ProxyStatus::default())), start_time: Arc::new(RwLock::new(None)), + current_providers: Arc::new(RwLock::new(std::collections::HashMap::new())), }; Self { @@ -121,17 +124,16 @@ impl ProxyServer { status.uptime_seconds = start.elapsed().as_secs(); } - // 获取所有活跃的代理目标 - if let Ok(targets) = self.state.db.get_all_proxy_targets() { - status.active_targets = targets - .into_iter() - .map(|(app_type, name, id)| ActiveTarget { - app_type, - provider_name: name, - provider_id: id, - }) - .collect(); - } + // 从 current_providers HashMap 获取每个应用类型当前正在使用的 provider + let current_providers = self.state.current_providers.read().await; + status.active_targets = current_providers + .iter() + .map(|(app_type, (provider_id, provider_name))| ActiveTarget { + app_type: app_type.clone(), + provider_id: provider_id.clone(), + provider_name: provider_name.clone(), + }) + .collect(); status }