From 59096d9c155c8059a421b20b63e22b0d3d13a209 Mon Sep 17 00:00:00 2001 From: YoVinchen Date: Mon, 22 Dec 2025 00:33:34 +0800 Subject: [PATCH] refactor(failover): make auto-failover toggle per-app independent - Change setting key from 'auto_failover_enabled' to 'auto_failover_enabled_{app_type}' - Update provider_router to check per-app failover setting - When failover disabled, use current provider only; when enabled, use queue order - Add unit tests for failover enabled/disabled behavior --- src-tauri/src/commands/failover.rs | 31 +++-- src-tauri/src/proxy/provider_router.rs | 174 +++++++++++++++++-------- src/lib/api/failover.ts | 15 ++- src/lib/query/failover.ts | 48 ++++--- 4 files changed, 177 insertions(+), 91 deletions(-) diff --git a/src-tauri/src/commands/failover.rs b/src-tauri/src/commands/failover.rs index bc8058ce2..cd8b9a8f2 100644 --- a/src-tauri/src/commands/failover.rs +++ b/src-tauri/src/commands/failover.rs @@ -82,28 +82,35 @@ pub async fn set_failover_item_enabled( .set_failover_item_enabled(&app_type, &provider_id, enabled) .map_err(|e| e.to_string()) } - -/// 获取自动故障转移总开关状态 +/// 获取指定应用的自动故障转移开关状态 #[tauri::command] -pub async fn get_auto_failover_enabled(state: tauri::State<'_, AppState>) -> Result { +pub async fn get_auto_failover_enabled( + state: tauri::State<'_, AppState>, + app_type: String, +) -> Result { + let key = format!("auto_failover_enabled_{app_type}"); state .db - .get_setting("auto_failover_enabled") + .get_setting(&key) .map(|v| v.map(|s| s == "true").unwrap_or(false)) // 默认关闭 .map_err(|e| e.to_string()) } -/// 设置自动故障转移总开关状态 +/// 设置指定应用的自动故障转移开关状态 +/// +/// 注意:关闭故障转移时不会清除队列,队列内容会保留供下次开启时使用 #[tauri::command] pub async fn set_auto_failover_enabled( state: tauri::State<'_, AppState>, + app_type: String, enabled: bool, ) -> Result<(), String> { - state - .db - .set_setting( - "auto_failover_enabled", - if enabled { "true" } else { "false" }, - ) - .map_err(|e| e.to_string()) + let key = format!("auto_failover_enabled_{app_type}"); + let value = if enabled { "true" } else { "false" }; + + log::info!( + "[Failover] Setting auto_failover_enabled: key='{key}', value='{value}', app_type='{app_type}'" + ); + + state.db.set_setting(&key, value).map_err(|e| e.to_string()) } diff --git a/src-tauri/src/proxy/provider_router.rs b/src-tauri/src/proxy/provider_router.rs index 3a5466f36..6fba778ba 100644 --- a/src-tauri/src/proxy/provider_router.rs +++ b/src-tauri/src/proxy/provider_router.rs @@ -30,54 +30,46 @@ impl ProviderRouter { /// 选择可用的供应商(支持故障转移) /// /// 返回按优先级排序的可用供应商列表: - /// 1. 当前供应商(is_current=true)始终第一位 - /// 2. 故障转移队列中的其他供应商(按 queue_order 排序)- 仅当自动故障转移开关开启时 - /// 3. 只返回熔断器未打开的供应商 + /// - 故障转移关闭时:仅返回当前供应商 + /// - 故障转移开启时:完全按照故障转移队列顺序返回,忽略当前供应商设置 pub async fn select_providers(&self, app_type: &str) -> Result, AppError> { let mut result = Vec::new(); let all_providers = self.db.get_all_providers(app_type)?; - // 检查自动故障转移总开关是否开启 - let auto_failover_enabled = self - .db - .get_setting("auto_failover_enabled") - .map(|v| v.map(|s| s == "true").unwrap_or(false)) // 默认关闭 - .unwrap_or(false); - - // 1. 当前供应商始终第一位 - if let Some(current_id) = self.db.get_current_provider(app_type)? { - if let Some(current) = all_providers.get(¤t_id) { - let circuit_key = format!("{}:{}", app_type, current.id); - let breaker = self.get_or_create_circuit_breaker(&circuit_key).await; - - if breaker.is_available().await { - log::info!( - "[{}] Current provider available: {} ({})", - app_type, - current.name, - current.id - ); - result.push(current.clone()); - } else { - log::warn!( - "[{}] Current provider {} circuit breaker open, checking failover queue", - app_type, - current.name - ); - } + // 检查该应用的自动故障转移开关是否开启 + let failover_key = format!("auto_failover_enabled_{app_type}"); + let auto_failover_enabled = match self.db.get_setting(&failover_key) { + Ok(Some(value)) => { + let enabled = value == "true"; + log::info!( + "[{app_type}] Failover setting '{failover_key}' = '{value}', enabled: {enabled}" + ); + enabled } - } + Ok(None) => { + log::warn!( + "[{app_type}] Failover setting '{failover_key}' not found in database, defaulting to disabled" + ); + false + } + Err(e) => { + log::error!( + "[{app_type}] Failed to read failover setting '{failover_key}': {e}, defaulting to disabled" + ); + false + } + }; - // 2. 获取故障转移队列中的供应商(仅当自动故障转移开关开启时) if auto_failover_enabled { + // 故障转移开启:完全按照队列顺序,忽略当前供应商 let queue = self.db.get_failover_queue(app_type)?; + log::info!( + "[{}] Failover enabled, using queue order ({} items)", + app_type, + queue.len() + ); for item in queue { - // 跳过已添加的当前供应商 - if result.iter().any(|p| p.id == item.provider_id) { - continue; - } - // 跳过禁用的队列项 if !item.enabled { continue; @@ -91,7 +83,7 @@ impl ProviderRouter { if breaker.is_available().await { log::info!( - "[{}] Failover provider available: {} ({}) at queue position {}", + "[{}] Queue provider available: {} ({}) at position {}", app_type, provider.name, provider.id, @@ -100,7 +92,7 @@ impl ProviderRouter { result.push(provider.clone()); } else { log::debug!( - "[{}] Failover provider {} circuit breaker open, skipping", + "[{}] Queue provider {} circuit breaker open, skipping", app_type, provider.name ); @@ -108,7 +100,31 @@ impl ProviderRouter { } } } else { - log::info!("[{app_type}] Auto-failover disabled, using only current provider"); + // 故障转移关闭:仅使用当前供应商 + log::info!("[{app_type}] Failover disabled, using current provider only"); + + if let Some(current_id) = self.db.get_current_provider(app_type)? { + if let Some(current) = all_providers.get(¤t_id) { + let circuit_key = format!("{}:{}", app_type, current.id); + let breaker = self.get_or_create_circuit_breaker(&circuit_key).await; + + if breaker.is_available().await { + log::info!( + "[{}] Current provider available: {} ({})", + app_type, + current.name, + current.id + ); + result.push(current.clone()); + } else { + log::warn!( + "[{}] Current provider {} circuit breaker open", + app_type, + current.name + ); + } + } + } } if result.is_empty() { @@ -118,7 +134,7 @@ impl ProviderRouter { } log::info!( - "[{}] Failover chain: {} provider(s) available", + "[{}] Provider chain: {} provider(s) available", app_type, result.len() ); @@ -275,25 +291,14 @@ mod tests { let db = Arc::new(Database::memory().unwrap()); let router = ProviderRouter::new(db); - // 测试创建熔断器 let breaker = router.get_or_create_circuit_breaker("claude:test").await; assert!(breaker.allow_request().await.allowed); } #[tokio::test] - async fn select_providers_does_not_consume_half_open_permit() { + async fn test_failover_disabled_uses_current_provider() { let db = Arc::new(Database::memory().unwrap()); - // 配置:让熔断器 Open 后立刻进入 HalfOpen(timeout_seconds=0),并用 1 次失败就打开熔断器 - db.update_circuit_breaker_config(&CircuitBreakerConfig { - failure_threshold: 1, - timeout_seconds: 0, - ..Default::default() - }) - .await - .unwrap(); - - // 准备 2 个 Provider:A(当前)+ B(队列) let provider_a = Provider::with_id("a".to_string(), "Provider A".to_string(), json!({}), None); let provider_b = @@ -304,19 +309,74 @@ mod tests { db.set_current_provider("claude", "a").unwrap(); db.add_to_failover_queue("claude", "b").unwrap(); + let router = ProviderRouter::new(db.clone()); + let providers = router.select_providers("claude").await.unwrap(); + + assert_eq!(providers.len(), 1); + assert_eq!(providers[0].id, "a"); + } + + #[tokio::test] + async fn test_failover_enabled_uses_queue_order() { + let db = Arc::new(Database::memory().unwrap()); + + let provider_a = + Provider::with_id("a".to_string(), "Provider A".to_string(), json!({}), None); + let provider_b = + Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None); + + db.save_provider("claude", &provider_a).unwrap(); + db.save_provider("claude", &provider_b).unwrap(); + db.set_current_provider("claude", "a").unwrap(); + + db.add_to_failover_queue("claude", "b").unwrap(); + db.add_to_failover_queue("claude", "a").unwrap(); + db.set_setting("auto_failover_enabled_claude", "true") + .unwrap(); + + let router = ProviderRouter::new(db.clone()); + let providers = router.select_providers("claude").await.unwrap(); + + assert_eq!(providers.len(), 2); + assert_eq!(providers[0].id, "b"); + assert_eq!(providers[1].id, "a"); + } + + #[tokio::test] + async fn test_select_providers_does_not_consume_half_open_permit() { + let db = Arc::new(Database::memory().unwrap()); + + db.update_circuit_breaker_config(&CircuitBreakerConfig { + failure_threshold: 1, + timeout_seconds: 0, + ..Default::default() + }) + .await + .unwrap(); + + let provider_a = + Provider::with_id("a".to_string(), "Provider A".to_string(), json!({}), None); + let provider_b = + Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None); + + db.save_provider("claude", &provider_a).unwrap(); + db.save_provider("claude", &provider_b).unwrap(); + + db.add_to_failover_queue("claude", "a").unwrap(); + db.add_to_failover_queue("claude", "b").unwrap(); + db.set_setting("auto_failover_enabled_claude", "true") + .unwrap(); + let router = ProviderRouter::new(db.clone()); - // 让 B 进入 Open 状态(failure_threshold=1) router .record_result("b", "claude", false, false, Some("fail".to_string())) .await .unwrap(); - // select_providers 只做“可用性判断”,不应占用 HalfOpen 探测名额 let providers = router.select_providers("claude").await.unwrap(); assert_eq!(providers.len(), 2); - // 如果 select_providers 错误地消耗了 HalfOpen 名额,这里会返回 false(被限流拒绝) assert!(router.allow_provider_request("b", "claude").await.allowed); } } diff --git a/src/lib/api/failover.ts b/src/lib/api/failover.ts index 85c00c852..fffa7a54f 100644 --- a/src/lib/api/failover.ts +++ b/src/lib/api/failover.ts @@ -105,13 +105,16 @@ export const failoverApi = { }); }, - // 获取自动故障转移总开关状态 - async getAutoFailoverEnabled(): Promise { - return invoke("get_auto_failover_enabled"); + // 获取指定应用的自动故障转移开关状态 + async getAutoFailoverEnabled(appType: string): Promise { + return invoke("get_auto_failover_enabled", { appType }); }, - // 设置自动故障转移总开关状态 - async setAutoFailoverEnabled(enabled: boolean): Promise { - return invoke("set_auto_failover_enabled", { enabled }); + // 设置指定应用的自动故障转移开关状态 + async setAutoFailoverEnabled( + appType: string, + enabled: boolean, + ): Promise { + return invoke("set_auto_failover_enabled", { appType, enabled }); }, }; diff --git a/src/lib/query/failover.ts b/src/lib/query/failover.ts index bfba5b52a..29e55b657 100644 --- a/src/lib/query/failover.ts +++ b/src/lib/query/failover.ts @@ -18,6 +18,9 @@ export function useProviderHealth(providerId: string, appType: string) { /** * 重置熔断器 + * + * 重置后后端会检查是否应该切回优先级更高的供应商, + * 因此需要同时刷新供应商列表和代理状态。 */ export function useResetCircuitBreaker() { const queryClient = useQueryClient(); @@ -35,6 +38,14 @@ export function useResetCircuitBreaker() { queryClient.invalidateQueries({ queryKey: ["providerHealth", variables.providerId, variables.appType], }); + // 刷新供应商列表(因为可能发生了自动恢复切换) + queryClient.invalidateQueries({ + queryKey: ["providers", variables.appType], + }); + // 刷新代理状态(更新 active_targets) + queryClient.invalidateQueries({ + queryKey: ["proxyStatus"], + }); }, }); } @@ -236,55 +247,60 @@ export function useSetFailoverItemEnabled() { }); } -// ========== 自动故障转移总开关 Hooks ========== +// ========== 自动故障转移开关 Hooks ========== /** - * 获取自动故障转移总开关状态 + * 获取指定应用的自动故障转移开关状态 */ -export function useAutoFailoverEnabled() { +export function useAutoFailoverEnabled(appType: string) { return useQuery({ - queryKey: ["autoFailoverEnabled"], - queryFn: () => failoverApi.getAutoFailoverEnabled(), + queryKey: ["autoFailoverEnabled", appType], + queryFn: () => failoverApi.getAutoFailoverEnabled(appType), // 默认值为 false(与后端保持一致) placeholderData: false, }); } /** - * 设置自动故障转移总开关状态 + * 设置指定应用的自动故障转移开关状态 */ export function useSetAutoFailoverEnabled() { const queryClient = useQueryClient(); return useMutation({ - mutationFn: (enabled: boolean) => - failoverApi.setAutoFailoverEnabled(enabled), + mutationFn: ({ appType, enabled }: { appType: string; enabled: boolean }) => + failoverApi.setAutoFailoverEnabled(appType, enabled), // 乐观更新 - onMutate: async (enabled) => { - await queryClient.cancelQueries({ queryKey: ["autoFailoverEnabled"] }); + onMutate: async ({ appType, enabled }) => { + await queryClient.cancelQueries({ + queryKey: ["autoFailoverEnabled", appType], + }); const previousValue = queryClient.getQueryData([ "autoFailoverEnabled", + appType, ]); - queryClient.setQueryData(["autoFailoverEnabled"], enabled); + queryClient.setQueryData(["autoFailoverEnabled", appType], enabled); - return { previousValue }; + return { previousValue, appType }; }, // 错误时回滚 - onError: (_error, _enabled, context) => { + onError: (_error, _variables, context) => { if (context?.previousValue !== undefined) { queryClient.setQueryData( - ["autoFailoverEnabled"], + ["autoFailoverEnabled", context.appType], context.previousValue, ); } }, // 无论成功失败,都重新获取 - onSettled: () => { - queryClient.invalidateQueries({ queryKey: ["autoFailoverEnabled"] }); + onSettled: (_, __, variables) => { + queryClient.invalidateQueries({ + queryKey: ["autoFailoverEnabled", variables.appType], + }); }, }); }