diff --git a/src-tauri/src/proxy/circuit_breaker.rs b/src-tauri/src/proxy/circuit_breaker.rs new file mode 100644 index 000000000..acedef87d --- /dev/null +++ b/src-tauri/src/proxy/circuit_breaker.rs @@ -0,0 +1,334 @@ +//! 熔断器模块 +//! +//! 实现熔断器模式,用于防止向不健康的供应商发送请求 + +use serde::{Deserialize, Serialize}; +use std::sync::atomic::{AtomicU32, Ordering}; +use std::sync::Arc; +use std::time::Instant; +use tokio::sync::RwLock; + +/// 熔断器状态 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum CircuitState { + /// 关闭状态 - 正常工作 + Closed, + /// 打开状态 - 熔断激活,拒绝请求 + Open, + /// 半开状态 - 尝试恢复,允许部分请求通过 + HalfOpen, +} + +impl std::fmt::Display for CircuitState { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + CircuitState::Closed => write!(f, "closed"), + CircuitState::Open => write!(f, "open"), + CircuitState::HalfOpen => write!(f, "half_open"), + } + } +} + +/// 熔断器配置 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct CircuitBreakerConfig { + /// 失败阈值 - 连续失败多少次后打开熔断器 + pub failure_threshold: u32, + /// 成功阈值 - 半开状态下成功多少次后关闭熔断器 + pub success_threshold: u32, + /// 超时时间 - 熔断器打开后多久尝试半开(秒) + pub timeout_seconds: u64, + /// 错误率阈值 - 错误率超过此值时打开熔断器 (0.0-1.0) + pub error_rate_threshold: f64, + /// 最小请求数 - 计算错误率前的最小请求数 + pub min_requests: u32, +} + +impl Default for CircuitBreakerConfig { + fn default() -> Self { + Self { + failure_threshold: 5, + success_threshold: 2, + timeout_seconds: 60, + error_rate_threshold: 0.5, + min_requests: 10, + } + } +} + +/// 熔断器实例 +pub struct CircuitBreaker { + /// 当前状态 + state: Arc>, + /// 连续失败计数 + consecutive_failures: Arc, + /// 连续成功计数(半开状态) + consecutive_successes: Arc, + /// 总请求计数 + total_requests: Arc, + /// 失败请求计数 + failed_requests: Arc, + /// 上次打开时间 + last_opened_at: Arc>>, + /// 配置 + config: CircuitBreakerConfig, +} + +impl CircuitBreaker { + /// 创建新的熔断器 + pub fn new(config: CircuitBreakerConfig) -> Self { + Self { + state: Arc::new(RwLock::new(CircuitState::Closed)), + consecutive_failures: Arc::new(AtomicU32::new(0)), + consecutive_successes: Arc::new(AtomicU32::new(0)), + total_requests: Arc::new(AtomicU32::new(0)), + failed_requests: Arc::new(AtomicU32::new(0)), + last_opened_at: Arc::new(RwLock::new(None)), + config, + } + } + + /// 检查是否允许请求通过 + pub async fn allow_request(&self) -> bool { + let state = *self.state.read().await; + + match state { + CircuitState::Closed => true, + CircuitState::Open => { + // 检查是否应该尝试半开 + if let Some(opened_at) = *self.last_opened_at.read().await { + if opened_at.elapsed().as_secs() >= self.config.timeout_seconds { + log::info!( + "Circuit breaker transitioning from Open to HalfOpen (timeout reached)" + ); + self.transition_to_half_open().await; + return true; + } + } + false + } + CircuitState::HalfOpen => true, + } + } + + /// 记录成功 + pub async fn record_success(&self) { + let state = *self.state.read().await; + + // 重置失败计数 + self.consecutive_failures.store(0, Ordering::SeqCst); + self.total_requests.fetch_add(1, Ordering::SeqCst); + + match state { + CircuitState::HalfOpen => { + let successes = self.consecutive_successes.fetch_add(1, Ordering::SeqCst) + 1; + log::debug!( + "Circuit breaker HalfOpen: {} consecutive successes (threshold: {})", + successes, + self.config.success_threshold + ); + + if successes >= self.config.success_threshold { + log::info!("Circuit breaker transitioning from HalfOpen to Closed (success threshold reached)"); + self.transition_to_closed().await; + } + } + CircuitState::Closed => { + log::debug!("Circuit breaker Closed: request succeeded"); + } + _ => {} + } + } + + /// 记录失败 + pub async fn record_failure(&self) { + let state = *self.state.read().await; + + // 更新计数器 + let failures = self.consecutive_failures.fetch_add(1, Ordering::SeqCst) + 1; + self.total_requests.fetch_add(1, Ordering::SeqCst); + self.failed_requests.fetch_add(1, Ordering::SeqCst); + + // 重置成功计数 + self.consecutive_successes.store(0, Ordering::SeqCst); + + log::debug!( + "Circuit breaker {:?}: {} consecutive failures (threshold: {})", + state, + failures, + self.config.failure_threshold + ); + + // 检查是否应该打开熔断器 + match state { + CircuitState::Closed | CircuitState::HalfOpen => { + // 检查连续失败次数 + if failures >= self.config.failure_threshold { + log::warn!( + "Circuit breaker opening due to {} consecutive failures (threshold: {})", + failures, + self.config.failure_threshold + ); + self.transition_to_open().await; + } else { + // 检查错误率 + let total = self.total_requests.load(Ordering::SeqCst); + let failed = self.failed_requests.load(Ordering::SeqCst); + + if total >= self.config.min_requests { + let error_rate = failed as f64 / total as f64; + log::debug!( + "Circuit breaker error rate: {:.2}% ({}/{} requests)", + error_rate * 100.0, + failed, + total + ); + + if error_rate >= self.config.error_rate_threshold { + log::warn!( + "Circuit breaker opening due to high error rate: {:.2}% (threshold: {:.2}%)", + error_rate * 100.0, + self.config.error_rate_threshold * 100.0 + ); + self.transition_to_open().await; + } + } + } + } + _ => {} + } + } + + /// 获取当前状态 + pub async fn get_state(&self) -> CircuitState { + *self.state.read().await + } + + /// 获取统计信息 + #[allow(dead_code)] + pub async fn get_stats(&self) -> CircuitBreakerStats { + CircuitBreakerStats { + state: *self.state.read().await, + consecutive_failures: self.consecutive_failures.load(Ordering::SeqCst), + consecutive_successes: self.consecutive_successes.load(Ordering::SeqCst), + total_requests: self.total_requests.load(Ordering::SeqCst), + failed_requests: self.failed_requests.load(Ordering::SeqCst), + } + } + + /// 重置熔断器(手动恢复) + #[allow(dead_code)] + pub async fn reset(&self) { + log::info!("Circuit breaker manually reset to Closed state"); + self.transition_to_closed().await; + } + + /// 转换到打开状态 + async fn transition_to_open(&self) { + *self.state.write().await = CircuitState::Open; + *self.last_opened_at.write().await = Some(Instant::now()); + self.consecutive_failures.store(0, Ordering::SeqCst); + self.consecutive_successes.store(0, Ordering::SeqCst); + } + + /// 转换到半开状态 + async fn transition_to_half_open(&self) { + *self.state.write().await = CircuitState::HalfOpen; + self.consecutive_successes.store(0, Ordering::SeqCst); + } + + /// 转换到关闭状态 + async fn transition_to_closed(&self) { + *self.state.write().await = CircuitState::Closed; + self.consecutive_failures.store(0, Ordering::SeqCst); + self.consecutive_successes.store(0, Ordering::SeqCst); + // 重置计数器 + self.total_requests.store(0, Ordering::SeqCst); + self.failed_requests.store(0, Ordering::SeqCst); + } +} + +/// 熔断器统计信息 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct CircuitBreakerStats { + pub state: CircuitState, + pub consecutive_failures: u32, + pub consecutive_successes: u32, + pub total_requests: u32, + pub failed_requests: u32, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn test_circuit_breaker_closed_to_open() { + let config = CircuitBreakerConfig { + failure_threshold: 3, + ..Default::default() + }; + let breaker = CircuitBreaker::new(config); + + // 初始状态应该是关闭 + assert_eq!(breaker.get_state().await, CircuitState::Closed); + assert!(breaker.allow_request().await); + + // 记录 3 次失败 + for _ in 0..3 { + breaker.record_failure().await; + } + + // 应该转换到打开状态 + assert_eq!(breaker.get_state().await, CircuitState::Open); + assert!(!breaker.allow_request().await); + } + + #[tokio::test] + async fn test_circuit_breaker_half_open_to_closed() { + let config = CircuitBreakerConfig { + failure_threshold: 2, + success_threshold: 2, + ..Default::default() + }; + let breaker = CircuitBreaker::new(config); + + // 打开熔断器 + breaker.record_failure().await; + breaker.record_failure().await; + assert_eq!(breaker.get_state().await, CircuitState::Open); + + // 手动转换到半开状态 + breaker.transition_to_half_open().await; + assert_eq!(breaker.get_state().await, CircuitState::HalfOpen); + + // 记录 2 次成功 + breaker.record_success().await; + breaker.record_success().await; + + // 应该转换到关闭状态 + assert_eq!(breaker.get_state().await, CircuitState::Closed); + } + + #[tokio::test] + async fn test_circuit_breaker_reset() { + let config = CircuitBreakerConfig { + failure_threshold: 2, + ..Default::default() + }; + let breaker = CircuitBreaker::new(config); + + // 打开熔断器 + breaker.record_failure().await; + breaker.record_failure().await; + assert_eq!(breaker.get_state().await, CircuitState::Open); + + // 重置 + breaker.reset().await; + assert_eq!(breaker.get_state().await, CircuitState::Closed); + assert!(breaker.allow_request().await); + } +} diff --git a/src-tauri/src/proxy/mod.rs b/src-tauri/src/proxy/mod.rs index e8390252d..c23d277c1 100644 --- a/src-tauri/src/proxy/mod.rs +++ b/src-tauri/src/proxy/mod.rs @@ -2,10 +2,12 @@ //! //! 提供本地HTTP代理服务,支持多Provider故障转移和请求透传 +pub mod circuit_breaker; pub mod error; mod forwarder; mod handlers; mod health; +pub mod provider_router; pub mod providers; pub mod response_handler; mod router; @@ -16,8 +18,14 @@ pub mod usage; // 公开导出给外部使用(commands, services等模块需要) #[allow(unused_imports)] +pub use circuit_breaker::{ + CircuitBreaker, CircuitBreakerConfig, CircuitBreakerStats, CircuitState, +}; +#[allow(unused_imports)] pub use error::ProxyError; #[allow(unused_imports)] +pub use provider_router::ProviderRouter; +#[allow(unused_imports)] pub use response_handler::{NonStreamHandler, ResponseType, StreamHandler}; #[allow(unused_imports)] pub use session::{ClientFormat, ProxySession}; diff --git a/src-tauri/src/proxy/provider_router.rs b/src-tauri/src/proxy/provider_router.rs new file mode 100644 index 000000000..9849ecd65 --- /dev/null +++ b/src-tauri/src/proxy/provider_router.rs @@ -0,0 +1,216 @@ +//! 供应商路由器模块 +//! +//! 负责选择和管理代理目标供应商,实现智能故障转移 + +use crate::database::Database; +use crate::error::AppError; +use crate::provider::Provider; +use crate::proxy::circuit_breaker::CircuitBreaker; +use std::collections::HashMap; +use std::sync::Arc; +use tokio::sync::RwLock; + +/// 供应商路由器 +pub struct ProviderRouter { + /// 数据库连接 + db: Arc, + /// 熔断器管理器 - key 格式: "app_type:provider_id" + circuit_breakers: Arc>>>, +} + +impl ProviderRouter { + /// 创建新的供应商路由器 + pub fn new(db: Arc) -> Self { + Self { + db, + circuit_breakers: Arc::new(RwLock::new(HashMap::new())), + } + } + + /// 选择可用的供应商(支持故障转移) + /// 返回按优先级排序的可用供应商列表 + pub async fn select_providers(&self, app_type: &str) -> Result, AppError> { + // 1. 获取所有启用代理的供应商 + let providers = self.db.get_proxy_targets(app_type).await?; + + if providers.is_empty() { + return Err(AppError::Config( + "No proxy target providers configured".to_string(), + )); + } + + log::debug!( + "Found {} proxy target providers for app_type: {}", + providers.len(), + app_type + ); + + // 2. 按 sort_index 排序(已经在数据库查询中排序了) + let sorted_providers: Vec<_> = providers.into_values().collect(); + + // 3. 过滤可用的供应商(检查熔断器状态) + let mut available_providers = Vec::new(); + + for provider in sorted_providers { + let circuit_key = format!("{}:{}", app_type, provider.id); + let breaker = self.get_or_create_circuit_breaker(&circuit_key).await; + + if breaker.allow_request().await { + log::debug!( + "Provider {} is available (circuit state: {:?})", + provider.id, + breaker.get_state().await + ); + available_providers.push(provider); + } else { + log::warn!( + "Provider {} is unavailable (circuit breaker open)", + provider.id + ); + } + } + + if available_providers.is_empty() { + return Err(AppError::Config( + "All proxy target providers are unavailable (circuit breakers open)".to_string(), + )); + } + + log::info!( + "Selected {} available providers for failover chain", + available_providers.len() + ); + + Ok(available_providers) + } + + /// 记录供应商请求结果 + pub async fn record_result( + &self, + provider_id: &str, + app_type: &str, + success: bool, + error_msg: Option, + ) -> Result<(), AppError> { + // 1. 更新熔断器状态 + let circuit_key = format!("{app_type}:{provider_id}"); + let breaker = self.get_or_create_circuit_breaker(&circuit_key).await; + + if success { + breaker.record_success().await; + log::debug!("Provider {provider_id} request succeeded"); + } else { + breaker.record_failure().await; + log::warn!( + "Provider {} request failed: {}", + provider_id, + error_msg.as_deref().unwrap_or("Unknown error") + ); + } + + // 2. 更新数据库健康状态 + self.db + .update_provider_health(provider_id, app_type, success, error_msg.clone()) + .await?; + + // 3. 如果连续失败达到熔断阈值,自动禁用代理目标 + if !success { + let health = self.db.get_provider_health(provider_id, app_type).await?; + + // 获取熔断器配置 + let config = self.db.get_circuit_breaker_config().await.ok(); + let failure_threshold = config.map(|c| c.failure_threshold).unwrap_or(5); + + // 如果连续失败达到阈值,自动关闭该供应商的代理开关 + if health.consecutive_failures >= failure_threshold { + log::warn!( + "Provider {} has failed {} times (threshold: {}), auto-disabling proxy target", + provider_id, + health.consecutive_failures, + failure_threshold + ); + self.db + .set_proxy_target(provider_id, app_type, false) + .await?; + } + } + + Ok(()) + } + + /// 重置熔断器(手动恢复) + #[allow(dead_code)] + pub async fn reset_circuit_breaker(&self, circuit_key: &str) { + let breakers = self.circuit_breakers.read().await; + if let Some(breaker) = breakers.get(circuit_key) { + log::info!("Manually resetting circuit breaker for {circuit_key}"); + breaker.reset().await; + } + } + + /// 获取熔断器状态 + #[allow(dead_code)] + pub async fn get_circuit_breaker_stats( + &self, + provider_id: &str, + app_type: &str, + ) -> Option { + let circuit_key = format!("{app_type}:{provider_id}"); + let breakers = self.circuit_breakers.read().await; + + if let Some(breaker) = breakers.get(&circuit_key) { + Some(breaker.get_stats().await) + } else { + None + } + } + + /// 获取或创建熔断器 + async fn get_or_create_circuit_breaker(&self, key: &str) -> Arc { + // 先尝试读锁获取 + { + let breakers = self.circuit_breakers.read().await; + if let Some(breaker) = breakers.get(key) { + return breaker.clone(); + } + } + + // 如果不存在,获取写锁创建 + let mut breakers = self.circuit_breakers.write().await; + + // 双重检查,防止竞争条件 + if let Some(breaker) = breakers.get(key) { + return breaker.clone(); + } + + // 从数据库加载配置 + let config = self + .db + .get_circuit_breaker_config() + .await + .unwrap_or_default(); + + log::debug!("Creating new circuit breaker for {key} with config: {config:?}"); + + let breaker = Arc::new(CircuitBreaker::new(config)); + breakers.insert(key.to_string(), breaker.clone()); + + breaker + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::database::Database; + + #[tokio::test] + async fn test_provider_router_creation() { + let db = Arc::new(Database::new_in_memory().unwrap()); + let router = ProviderRouter::new(db); + + // 测试创建熔断器 + let breaker = router.get_or_create_circuit_breaker("claude:test").await; + assert!(breaker.allow_request().await); + } +}