diff --git a/assets/partners/logos/aigocode.png b/assets/partners/logos/aigocode.png index c9a30c7d4..6dd5965ac 100644 Binary files a/assets/partners/logos/aigocode.png and b/assets/partners/logos/aigocode.png differ diff --git a/src-tauri/src/commands/failover.rs b/src-tauri/src/commands/failover.rs new file mode 100644 index 000000000..ddceeef13 --- /dev/null +++ b/src-tauri/src/commands/failover.rs @@ -0,0 +1,84 @@ +//! 故障转移队列命令 +//! +//! 管理代理模式下的故障转移队列 + +use crate::database::FailoverQueueItem; +use crate::provider::Provider; +use crate::store::AppState; + +/// 获取故障转移队列 +#[tauri::command] +pub async fn get_failover_queue( + state: tauri::State<'_, AppState>, + app_type: String, +) -> Result, String> { + state + .db + .get_failover_queue(&app_type) + .map_err(|e| e.to_string()) +} + +/// 获取可添加到故障转移队列的供应商(不在队列中的) +#[tauri::command] +pub async fn get_available_providers_for_failover( + state: tauri::State<'_, AppState>, + app_type: String, +) -> Result, String> { + state + .db + .get_available_providers_for_failover(&app_type) + .map_err(|e| e.to_string()) +} + +/// 添加供应商到故障转移队列 +#[tauri::command] +pub async fn add_to_failover_queue( + state: tauri::State<'_, AppState>, + app_type: String, + provider_id: String, +) -> Result<(), String> { + state + .db + .add_to_failover_queue(&app_type, &provider_id) + .map_err(|e| e.to_string()) +} + +/// 从故障转移队列移除供应商 +#[tauri::command] +pub async fn remove_from_failover_queue( + state: tauri::State<'_, AppState>, + app_type: String, + provider_id: String, +) -> Result<(), String> { + state + .db + .remove_from_failover_queue(&app_type, &provider_id) + .map_err(|e| e.to_string()) +} + +/// 重新排序故障转移队列 +#[tauri::command] +pub async fn reorder_failover_queue( + state: tauri::State<'_, AppState>, + app_type: String, + provider_ids: Vec, +) -> Result<(), String> { + state + .db + .reorder_failover_queue(&app_type, &provider_ids) + .map_err(|e| e.to_string()) +} + +/// 设置故障转移队列项的启用状态 +#[tauri::command] +pub async fn set_failover_item_enabled( + state: tauri::State<'_, AppState>, + app_type: String, + provider_id: String, + enabled: bool, +) -> Result<(), String> { + state + .db + .set_failover_item_enabled(&app_type, &provider_id, enabled) + .map_err(|e| e.to_string()) +} diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index 9b7e87b98..85dfb3adc 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -3,6 +3,7 @@ mod config; mod deeplink; mod env; +mod failover; mod import_export; mod mcp; mod misc; @@ -18,6 +19,7 @@ mod usage; pub use config::*; pub use deeplink::*; pub use env::*; +pub use failover::*; pub use import_export::*; pub use mcp::*; pub use misc::*; diff --git a/src-tauri/src/commands/provider.rs b/src-tauri/src/commands/provider.rs index b25b63edd..cce5aabd3 100644 --- a/src-tauri/src/commands/provider.rs +++ b/src-tauri/src/commands/provider.rs @@ -86,19 +86,6 @@ pub fn switch_provider( .map_err(|e| e.to_string()) } -/// 设置代理目标供应商 -#[tauri::command] -pub fn set_proxy_target_provider( - state: State<'_, AppState>, - app: String, - id: String, -) -> Result { - let app_type = AppType::from_str(&app).map_err(|e| e.to_string())?; - ProviderService::set_proxy_target(state.inner(), app_type, &id) - .map(|_| true) - .map_err(|e| e.to_string()) -} - fn import_default_config_internal(state: &AppState, app_type: AppType) -> Result { ProviderService::import_default_config(state, app_type) } diff --git a/src-tauri/src/commands/proxy.rs b/src-tauri/src/commands/proxy.rs index 9c9a34212..395ce307a 100644 --- a/src-tauri/src/commands/proxy.rs +++ b/src-tauri/src/commands/proxy.rs @@ -2,7 +2,6 @@ //! //! 提供前端调用的 API 接口 -use crate::provider::Provider; use crate::proxy::types::*; use crate::proxy::{CircuitBreakerConfig, CircuitBreakerStats}; use crate::store::AppState; @@ -69,47 +68,6 @@ pub async fn switch_proxy_provider( // ==================== 故障转移相关命令 ==================== -/// 获取代理目标列表 -#[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( @@ -130,15 +88,17 @@ pub async fn reset_circuit_breaker( provider_id: String, app_type: String, ) -> Result<(), String> { - // 重置数据库健康状态 + // 1. 重置数据库健康状态 let db = &state.db; db.update_provider_health(&provider_id, &app_type, true, None) .await .map_err(|e| e.to_string())?; - // 注意:熔断器状态在内存中,重启代理服务器后会重置 - // 如果代理服务器正在运行,需要通知它重置熔断器 - // 目前先通过数据库重置健康状态,熔断器会在下次超时后自动尝试半开 + // 2. 如果代理正在运行,重置内存中的熔断器状态 + state + .proxy_service + .reset_provider_circuit_breaker(&provider_id, &app_type) + .await?; Ok(()) } @@ -161,9 +121,19 @@ pub async fn update_circuit_breaker_config( config: CircuitBreakerConfig, ) -> Result<(), String> { let db = &state.db; + + // 1. 更新数据库配置 db.update_circuit_breaker_config(&config) .await - .map_err(|e| e.to_string()) + .map_err(|e| e.to_string())?; + + // 2. 如果代理正在运行,热更新内存中的熔断器配置 + state + .proxy_service + .update_circuit_breaker_configs(config) + .await?; + + Ok(()) } /// 获取熔断器统计信息(仅当代理服务器运行时) diff --git a/src-tauri/src/commands/stream_check.rs b/src-tauri/src/commands/stream_check.rs index d0fa13dbf..13ed33074 100644 --- a/src-tauri/src/commands/stream_check.rs +++ b/src-tauri/src/commands/stream_check.rs @@ -6,6 +6,7 @@ use crate::services::stream_check::{ HealthStatus, StreamCheckConfig, StreamCheckResult, StreamCheckService, }; use crate::store::AppState; +use std::collections::HashSet; use tauri::State; /// 流式健康检查(单个供应商) @@ -44,10 +45,28 @@ pub async fn stream_check_all_providers( let providers = state.db.get_all_providers(app_type.as_str())?; let mut results = Vec::new(); + let allowed_ids: Option> = if proxy_targets_only { + let mut ids = HashSet::new(); + if let Ok(Some(current_id)) = state.db.get_current_provider(app_type.as_str()) { + ids.insert(current_id); + } + if let Ok(queue) = state.db.get_failover_queue(app_type.as_str()) { + for item in queue { + if item.enabled { + ids.insert(item.provider_id); + } + } + } + Some(ids) + } else { + None + }; for (id, provider) in providers { - if proxy_targets_only && !provider.is_proxy_target.unwrap_or(false) { - continue; + if let Some(ids) = &allowed_ids { + if !ids.contains(&id) { + continue; + } } let result = StreamCheckService::check_with_retry(&app_type, &provider, &config) diff --git a/src-tauri/src/database/dao/failover.rs b/src-tauri/src/database/dao/failover.rs new file mode 100644 index 000000000..66e24f92d --- /dev/null +++ b/src-tauri/src/database/dao/failover.rs @@ -0,0 +1,248 @@ +//! 故障转移队列 DAO +//! +//! 管理代理模式下的故障转移队列 + +use crate::database::{lock_conn, Database}; +use crate::error::AppError; +use crate::provider::Provider; +use serde::{Deserialize, Serialize}; +use std::time::{SystemTime, UNIX_EPOCH}; + +/// 故障转移队列条目 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct FailoverQueueItem { + pub provider_id: String, + pub provider_name: String, + pub queue_order: i32, + pub enabled: bool, + pub created_at: i64, +} + +impl Database { + /// 获取故障转移队列(按 queue_order 排序) + pub fn get_failover_queue(&self, app_type: &str) -> Result, AppError> { + let conn = lock_conn!(self.conn); + + let mut stmt = conn + .prepare( + "SELECT fq.provider_id, p.name, fq.queue_order, fq.enabled, fq.created_at + FROM failover_queue fq + JOIN providers p ON fq.provider_id = p.id AND fq.app_type = p.app_type + WHERE fq.app_type = ?1 + ORDER BY fq.queue_order ASC", + ) + .map_err(|e| AppError::Database(e.to_string()))?; + + let items = stmt + .query_map([app_type], |row| { + Ok(FailoverQueueItem { + provider_id: row.get(0)?, + provider_name: row.get(1)?, + queue_order: row.get(2)?, + enabled: row.get(3)?, + created_at: row.get(4)?, + }) + }) + .map_err(|e| AppError::Database(e.to_string()))? + .collect::, _>>() + .map_err(|e| AppError::Database(e.to_string()))?; + + Ok(items) + } + + /// 获取故障转移队列中的供应商(完整 Provider 信息,按顺序) + pub fn get_failover_providers(&self, app_type: &str) -> Result, AppError> { + let queue = self.get_failover_queue(app_type)?; + let all_providers = self.get_all_providers(app_type)?; + + let mut result = Vec::new(); + for item in queue { + if item.enabled { + if let Some(provider) = all_providers.get(&item.provider_id) { + result.push(provider.clone()); + } + } + } + + Ok(result) + } + + /// 添加供应商到故障转移队列末尾 + pub fn add_to_failover_queue( + &self, + app_type: &str, + provider_id: &str, + ) -> Result<(), AppError> { + let conn = lock_conn!(self.conn); + + // 获取当前最大 queue_order + let max_order: i32 = conn + .query_row( + "SELECT COALESCE(MAX(queue_order), 0) FROM failover_queue WHERE app_type = ?1", + [app_type], + |row| row.get(0), + ) + .map_err(|e| AppError::Database(e.to_string()))?; + + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() as i64; + + conn.execute( + "INSERT OR IGNORE INTO failover_queue (app_type, provider_id, queue_order, enabled, created_at) + VALUES (?1, ?2, ?3, 1, ?4)", + rusqlite::params![app_type, provider_id, max_order + 1, now], + ) + .map_err(|e| AppError::Database(e.to_string()))?; + + Ok(()) + } + + /// 从故障转移队列中移除供应商 + pub fn remove_from_failover_queue( + &self, + app_type: &str, + provider_id: &str, + ) -> Result<(), AppError> { + let conn = lock_conn!(self.conn); + + // 获取被删除项的 queue_order + let removed_order: Option = conn + .query_row( + "SELECT queue_order FROM failover_queue WHERE app_type = ?1 AND provider_id = ?2", + [app_type, provider_id], + |row| row.get(0), + ) + .ok(); + + // 删除该项 + conn.execute( + "DELETE FROM failover_queue WHERE app_type = ?1 AND provider_id = ?2", + [app_type, provider_id], + ) + .map_err(|e| AppError::Database(e.to_string()))?; + + // 重新排序后面的项(填补空隙) + if let Some(order) = removed_order { + conn.execute( + "UPDATE failover_queue + SET queue_order = queue_order - 1 + WHERE app_type = ?1 AND queue_order > ?2", + rusqlite::params![app_type, order], + ) + .map_err(|e| AppError::Database(e.to_string()))?; + } + + Ok(()) + } + + /// 重新排序故障转移队列 + /// provider_ids: 按新顺序排列的 provider_id 列表 + pub fn reorder_failover_queue( + &self, + app_type: &str, + provider_ids: &[String], + ) -> Result<(), AppError> { + let conn = lock_conn!(self.conn); + + // 使用事务确保原子性 + conn.execute("BEGIN TRANSACTION", []) + .map_err(|e| AppError::Database(e.to_string()))?; + + let result = (|| { + for (index, provider_id) in provider_ids.iter().enumerate() { + conn.execute( + "UPDATE failover_queue + SET queue_order = ?3 + WHERE app_type = ?1 AND provider_id = ?2", + rusqlite::params![app_type, provider_id, (index + 1) as i32], + ) + .map_err(|e| AppError::Database(e.to_string()))?; + } + Ok(()) + })(); + + match result { + Ok(_) => { + conn.execute("COMMIT", []) + .map_err(|e| AppError::Database(e.to_string()))?; + Ok(()) + } + Err(e) => { + conn.execute("ROLLBACK", []).ok(); + Err(e) + } + } + } + + /// 设置故障转移队列中供应商的启用状态 + pub fn set_failover_item_enabled( + &self, + app_type: &str, + provider_id: &str, + enabled: bool, + ) -> Result<(), AppError> { + let conn = lock_conn!(self.conn); + + conn.execute( + "UPDATE failover_queue SET enabled = ?3 WHERE app_type = ?1 AND provider_id = ?2", + rusqlite::params![app_type, provider_id, enabled], + ) + .map_err(|e| AppError::Database(e.to_string()))?; + + Ok(()) + } + + /// 清空故障转移队列 + pub fn clear_failover_queue(&self, app_type: &str) -> Result<(), AppError> { + let conn = lock_conn!(self.conn); + + conn.execute( + "DELETE FROM failover_queue WHERE app_type = ?1", + [app_type], + ) + .map_err(|e| AppError::Database(e.to_string()))?; + + Ok(()) + } + + /// 检查供应商是否在故障转移队列中 + pub fn is_in_failover_queue( + &self, + app_type: &str, + provider_id: &str, + ) -> Result { + let conn = lock_conn!(self.conn); + + let count: i32 = conn + .query_row( + "SELECT COUNT(*) FROM failover_queue WHERE app_type = ?1 AND provider_id = ?2", + [app_type, provider_id], + |row| row.get(0), + ) + .map_err(|e| AppError::Database(e.to_string()))?; + + Ok(count > 0) + } + + /// 获取可添加到故障转移队列的供应商(不在队列中的) + pub fn get_available_providers_for_failover( + &self, + app_type: &str, + ) -> Result, AppError> { + let all_providers = self.get_all_providers(app_type)?; + let queue = self.get_failover_queue(app_type)?; + + let queue_ids: std::collections::HashSet<_> = + queue.iter().map(|item| &item.provider_id).collect(); + + let available: Vec = all_providers + .into_values() + .filter(|p| !queue_ids.contains(&p.id)) + .collect(); + + Ok(available) + } +} diff --git a/src-tauri/src/database/dao/mod.rs b/src-tauri/src/database/dao/mod.rs index a3b6ca471..c759c2a86 100644 --- a/src-tauri/src/database/dao/mod.rs +++ b/src-tauri/src/database/dao/mod.rs @@ -2,6 +2,7 @@ //! //! Database access operations for each domain +pub mod failover; pub mod mcp; pub mod prompts; pub mod providers; @@ -11,3 +12,5 @@ pub mod skills; pub mod stream_check; // 所有 DAO 方法都通过 Database impl 提供,无需单独导出 +// 导出 FailoverQueueItem 供外部使用 +pub use failover::FailoverQueueItem; diff --git a/src-tauri/src/database/dao/providers.rs b/src-tauri/src/database/dao/providers.rs index b9cf8de8a..1ba3683c5 100644 --- a/src-tauri/src/database/dao/providers.rs +++ b/src-tauri/src/database/dao/providers.rs @@ -17,7 +17,7 @@ impl Database { ) -> Result, AppError> { let conn = lock_conn!(self.conn); let mut stmt = conn.prepare( - "SELECT id, name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta, is_proxy_target + "SELECT id, name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta FROM providers WHERE app_type = ?1 ORDER BY COALESCE(sort_index, 999999), created_at ASC, id ASC" ).map_err(|e| AppError::Database(e.to_string()))?; @@ -35,7 +35,6 @@ impl Database { let icon: Option = row.get(8)?; let icon_color: Option = row.get(9)?; let meta_str: String = row.get(10)?; - let is_proxy_target: bool = row.get(11)?; let settings_config = serde_json::from_str(&settings_config_str).unwrap_or(serde_json::Value::Null); @@ -55,7 +54,6 @@ impl Database { meta: Some(meta), icon, icon_color, - is_proxy_target: Some(is_proxy_target), }, )) }) @@ -131,7 +129,7 @@ impl Database { ) -> Result, AppError> { let conn = lock_conn!(self.conn); let result = conn.query_row( - "SELECT name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta, is_proxy_target + "SELECT name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta FROM providers WHERE id = ?1 AND app_type = ?2", params![id, app_type], |row| { @@ -145,7 +143,6 @@ impl Database { let icon: Option = row.get(7)?; let icon_color: Option = row.get(8)?; let meta_str: String = row.get(9)?; - let is_proxy_target: bool = row.get(10)?; let settings_config = serde_json::from_str(&settings_config_str).unwrap_or(serde_json::Value::Null); let meta: ProviderMeta = serde_json::from_str(&meta_str).unwrap_or_default(); @@ -162,7 +159,6 @@ impl Database { meta: Some(meta), icon, icon_color, - is_proxy_target: Some(is_proxy_target), }) }, ); @@ -174,26 +170,6 @@ impl Database { } } - /// 获取代理目标供应商 ID - pub fn get_proxy_target_provider(&self, app_type: &str) -> Result, AppError> { - let conn = lock_conn!(self.conn); - let mut stmt = conn - .prepare("SELECT id FROM providers WHERE app_type = ?1 AND is_proxy_target = 1 LIMIT 1") - .map_err(|e| AppError::Database(e.to_string()))?; - - let mut rows = stmt - .query(params![app_type]) - .map_err(|e| AppError::Database(e.to_string()))?; - - if let Some(row) = rows.next().map_err(|e| AppError::Database(e.to_string()))? { - Ok(Some( - row.get(0).map_err(|e| AppError::Database(e.to_string()))?, - )) - } else { - Ok(None) - } - } - /// 保存供应商(新增或更新) /// /// 注意:更新模式下不同步 endpoints,因为编辑模式下端点通过单独的 API 管理 @@ -208,17 +184,17 @@ impl Database { let mut meta_clone = provider.meta.clone().unwrap_or_default(); let endpoints = std::mem::take(&mut meta_clone.custom_endpoints); - // 检查是否存在(用于判断新增/更新,以及保留 is_current 和 is_proxy_target) - let existing: Option<(bool, bool)> = tx + // 检查是否存在(用于判断新增/更新,以及保留 is_current) + let existing: Option = tx .query_row( - "SELECT is_current, is_proxy_target FROM providers WHERE id = ?1 AND app_type = ?2", + "SELECT is_current FROM providers WHERE id = ?1 AND app_type = ?2", params![provider.id, app_type], - |row| Ok((row.get(0)?, row.get(1)?)), + |row| row.get(0), ) .ok(); let is_update = existing.is_some(); - let (is_current, is_proxy_target) = existing.unwrap_or((false, false)); + let is_current = existing.unwrap_or(false); if is_update { // 更新模式:使用 UPDATE 避免触发 ON DELETE CASCADE @@ -234,9 +210,8 @@ impl Database { icon = ?8, icon_color = ?9, meta = ?10, - is_current = ?11, - is_proxy_target = ?12 - WHERE id = ?13 AND app_type = ?14", + is_current = ?11 + WHERE id = ?12 AND app_type = ?13", params![ provider.name, serde_json::to_string(&provider.settings_config).unwrap(), @@ -249,7 +224,6 @@ impl Database { provider.icon_color, serde_json::to_string(&meta_clone).unwrap(), is_current, - is_proxy_target, provider.id, app_type, ], @@ -260,8 +234,8 @@ impl Database { tx.execute( "INSERT INTO providers ( id, app_type, name, settings_config, website_url, category, - created_at, sort_index, notes, icon, icon_color, meta, is_current, is_proxy_target - ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)", + created_at, sort_index, notes, icon, icon_color, meta, is_current + ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13)", params![ provider.id, app_type, @@ -276,7 +250,6 @@ impl Database { provider.icon_color, serde_json::to_string(&meta_clone).unwrap(), is_current, - is_proxy_target, ], ) .map_err(|e| AppError::Database(e.to_string()))?; @@ -332,157 +305,6 @@ impl Database { Ok(()) } - /// 设置代理目标供应商 - pub fn set_proxy_target_provider(&self, app_type: &str, id: &str) -> Result<(), AppError> { - let mut conn = lock_conn!(self.conn); - let tx = conn - .transaction() - .map_err(|e| AppError::Database(e.to_string()))?; - - // 重置所有为 0 - tx.execute( - "UPDATE providers SET is_proxy_target = 0 WHERE app_type = ?1", - params![app_type], - ) - .map_err(|e| AppError::Database(e.to_string()))?; - - // 设置新的代理目标供应商 - tx.execute( - "UPDATE providers SET is_proxy_target = 1 WHERE id = ?1 AND app_type = ?2", - params![id, app_type], - ) - .map_err(|e| AppError::Database(e.to_string()))?; - - tx.commit().map_err(|e| AppError::Database(e.to_string()))?; - Ok(()) - } - - /// 设置单个供应商的代理目标状态(支持多个代理目标) - pub async fn set_proxy_target( - &self, - provider_id: &str, - app_type: &str, - enabled: bool, - ) -> Result<(), AppError> { - let conn = lock_conn!(self.conn); - conn.execute( - "UPDATE providers SET is_proxy_target = ?1 - WHERE id = ?2 AND app_type = ?3", - params![if enabled { 1 } else { 0 }, provider_id, app_type], - ) - .map_err(|e| AppError::Database(e.to_string()))?; - Ok(()) - } - - /// 获取指定应用类型的所有代理目标供应商(按 sort_index 排序) - pub async fn get_proxy_targets( - &self, - app_type: &str, - ) -> Result, AppError> { - let conn = lock_conn!(self.conn); - let mut stmt = conn.prepare( - "SELECT id, name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta, is_proxy_target - FROM providers WHERE app_type = ?1 AND is_proxy_target = 1 - ORDER BY COALESCE(sort_index, 999999), created_at ASC, id ASC" - ).map_err(|e| AppError::Database(e.to_string()))?; - - let provider_iter = stmt - .query_map(params![app_type], |row| { - let id: String = row.get(0)?; - let name: String = row.get(1)?; - let settings_config_str: String = row.get(2)?; - let website_url: Option = row.get(3)?; - let category: Option = row.get(4)?; - let created_at: Option = row.get(5)?; - let sort_index: Option = row.get(6)?; - let notes: Option = row.get(7)?; - let icon: Option = row.get(8)?; - let icon_color: Option = row.get(9)?; - let meta_str: String = row.get(10)?; - let is_proxy_target: bool = row.get(11)?; - - let settings_config = - serde_json::from_str(&settings_config_str).unwrap_or(serde_json::Value::Null); - let meta: ProviderMeta = serde_json::from_str(&meta_str).unwrap_or_default(); - - Ok(( - id, - Provider { - id: "".to_string(), - name, - settings_config, - website_url, - category, - created_at, - sort_index, - notes, - meta: Some(meta), - icon, - icon_color, - is_proxy_target: Some(is_proxy_target), - }, - )) - }) - .map_err(|e| AppError::Database(e.to_string()))?; - - let mut providers = IndexMap::new(); - for provider_res in provider_iter { - let (id, mut provider) = provider_res.map_err(|e| AppError::Database(e.to_string()))?; - provider.id = id.clone(); - - // 加载 endpoints - let mut stmt_endpoints = conn.prepare( - "SELECT url, added_at FROM provider_endpoints WHERE provider_id = ?1 AND app_type = ?2 ORDER BY added_at ASC, url ASC" - ).map_err(|e| AppError::Database(e.to_string()))?; - - let endpoints_iter = stmt_endpoints - .query_map(params![id, app_type], |row| { - let url: String = row.get(0)?; - let added_at: Option = row.get(1)?; - Ok(( - url, - crate::settings::CustomEndpoint { - url: "".to_string(), - added_at: added_at.unwrap_or(0), - last_used: None, - }, - )) - }) - .map_err(|e| AppError::Database(e.to_string()))?; - - let mut custom_endpoints = HashMap::new(); - for ep_res in endpoints_iter { - let (url, mut ep) = ep_res.map_err(|e| AppError::Database(e.to_string()))?; - ep.url = url.clone(); - custom_endpoints.insert(url, ep); - } - - if let Some(meta) = &mut provider.meta { - meta.custom_endpoints = custom_endpoints; - } - - providers.insert(id, provider); - } - - Ok(providers) - } - - /// 获取所有活跃的代理目标 - pub fn get_all_proxy_targets(&self) -> Result, AppError> { - let conn = lock_conn!(self.conn); - let mut stmt = conn - .prepare("SELECT app_type, name, id FROM providers WHERE is_proxy_target = 1") - .map_err(|e| AppError::Database(e.to_string()))?; - - let targets = stmt - .query_map([], |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?))) - .map_err(|e| AppError::Database(e.to_string()))? - .collect::, _>>() - .map_err(|e| AppError::Database(e.to_string()))?; - - Ok(targets) - } - /// 更新供应商的 settings_config(仅更新配置,不改变其他字段) pub fn update_provider_settings_config( &self, diff --git a/src-tauri/src/database/dao/proxy.rs b/src-tauri/src/database/dao/proxy.rs index 961760c42..1e6344b61 100644 --- a/src-tauri/src/database/dao/proxy.rs +++ b/src-tauri/src/database/dao/proxy.rs @@ -129,12 +129,31 @@ impl Database { } /// 更新Provider健康状态 + /// + /// 使用默认阈值(5)判断是否健康,建议使用 `update_provider_health_with_threshold` 传入配置的阈值 pub async fn update_provider_health( &self, provider_id: &str, app_type: &str, success: bool, error_msg: Option, + ) -> Result<(), AppError> { + // 默认阈值与 CircuitBreakerConfig::default() 保持一致 + self.update_provider_health_with_threshold(provider_id, app_type, success, error_msg, 5) + .await + } + + /// 更新Provider健康状态(带阈值参数) + /// + /// # Arguments + /// * `failure_threshold` - 连续失败多少次后标记为不健康 + pub async fn update_provider_health_with_threshold( + &self, + provider_id: &str, + app_type: &str, + success: bool, + error_msg: Option, + failure_threshold: u32, ) -> Result<(), AppError> { let conn = lock_conn!(self.conn); @@ -142,7 +161,7 @@ impl Database { // 先查询当前状态 let current = conn.query_row( - "SELECT consecutive_failures FROM provider_health + "SELECT consecutive_failures FROM provider_health WHERE provider_id = ?1 AND app_type = ?2", rusqlite::params![provider_id, app_type], |row| Ok(row.get::<_, i64>(0)? as u32), @@ -154,7 +173,8 @@ impl Database { } else { // 失败:增加失败计数 let failures = current.unwrap_or(0) + 1; - let healthy = if failures >= 3 { 0 } else { 1 }; + // 使用传入的阈值而非硬编码 + let healthy = if failures >= failure_threshold { 0 } else { 1 }; (healthy, failures) }; @@ -169,10 +189,10 @@ impl Database { "INSERT OR REPLACE INTO provider_health (provider_id, app_type, is_healthy, consecutive_failures, last_success_at, last_failure_at, last_error, updated_at) - VALUES (?1, ?2, ?3, ?4, - COALESCE(?5, (SELECT last_success_at FROM provider_health + VALUES (?1, ?2, ?3, ?4, + COALESCE(?5, (SELECT last_success_at FROM provider_health WHERE provider_id = ?1 AND app_type = ?2)), - COALESCE(?6, (SELECT last_failure_at FROM provider_health + COALESCE(?6, (SELECT last_failure_at FROM provider_health WHERE provider_id = ?1 AND app_type = ?2)), ?7, ?8)", rusqlite::params![ @@ -210,6 +230,17 @@ impl Database { Ok(()) } + /// 清空所有Provider健康状态(代理停止时调用) + pub async fn clear_all_provider_health(&self) -> Result<(), AppError> { + let conn = lock_conn!(self.conn); + + conn.execute("DELETE FROM provider_health", []) + .map_err(|e| AppError::Database(e.to_string()))?; + + log::debug!("Cleared all provider health records"); + Ok(()) + } + // ==================== Circuit Breaker Config ==================== /// 获取熔断器配置 diff --git a/src-tauri/src/database/mod.rs b/src-tauri/src/database/mod.rs index 2786a20b9..73a4877a4 100644 --- a/src-tauri/src/database/mod.rs +++ b/src-tauri/src/database/mod.rs @@ -31,6 +31,9 @@ mod schema; #[cfg(test)] mod tests; +// DAO 类型导出供外部使用 +pub use dao::FailoverQueueItem; + use crate::config::get_app_config_dir; use crate::error::AppError; use rusqlite::Connection; diff --git a/src-tauri/src/database/schema.rs b/src-tauri/src/database/schema.rs index 3dbdca557..134397a40 100644 --- a/src-tauri/src/database/schema.rs +++ b/src-tauri/src/database/schema.rs @@ -31,19 +31,12 @@ impl Database { icon_color TEXT, meta TEXT NOT NULL DEFAULT '{}', is_current BOOLEAN NOT NULL DEFAULT 0, - is_proxy_target BOOLEAN NOT NULL DEFAULT 0, PRIMARY KEY (id, app_type) )", [], ) .map_err(|e| AppError::Database(e.to_string()))?; - // 尝试添加 is_proxy_target 列(如果表已存在但缺少该列) - let _ = conn.execute( - "ALTER TABLE providers ADD COLUMN is_proxy_target BOOLEAN NOT NULL DEFAULT 0", - [], - ); - // 2. Provider Endpoints 表 conn.execute( "CREATE TABLE IF NOT EXISTS provider_endpoints ( @@ -319,6 +312,30 @@ impl Database { [], ); + // 14. Failover Queue 表 (故障转移队列) + conn.execute( + "CREATE TABLE IF NOT EXISTS failover_queue ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + app_type TEXT NOT NULL, + provider_id TEXT NOT NULL, + queue_order INTEGER NOT NULL, + enabled INTEGER NOT NULL DEFAULT 1, + created_at INTEGER NOT NULL, + UNIQUE (app_type, provider_id), + FOREIGN KEY (provider_id, app_type) REFERENCES providers(id, app_type) ON DELETE CASCADE + )", + [], + ) + .map_err(|e| AppError::Database(e.to_string()))?; + + // 为故障转移队列创建索引 + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_failover_queue_order + ON failover_queue(app_type, queue_order)", + [], + ) + .map_err(|e| AppError::Database(e.to_string()))?; + Ok(()) } diff --git a/src-tauri/src/database/tests.rs b/src-tauri/src/database/tests.rs index 2ce719039..44bf304ac 100644 --- a/src-tauri/src/database/tests.rs +++ b/src-tauri/src/database/tests.rs @@ -245,7 +245,6 @@ fn dry_run_validates_schema_compatibility() { meta: None, icon: None, icon_color: None, - is_proxy_target: Some(false), }, ); diff --git a/src-tauri/src/deeplink/mod.rs b/src-tauri/src/deeplink/mod.rs index 5c1e1a72e..4111089e2 100644 --- a/src-tauri/src/deeplink/mod.rs +++ b/src-tauri/src/deeplink/mod.rs @@ -113,4 +113,27 @@ pub struct DeepLinkImportRequest { /// Remote config URL #[serde(skip_serializing_if = "Option::is_none")] pub config_url: Option, + + // ============ Usage script fields (v3.9+) ============ + /// Whether to enable usage query (default: true if usage_script is provided) + #[serde(skip_serializing_if = "Option::is_none")] + pub usage_enabled: Option, + /// Base64 encoded usage query script code + #[serde(skip_serializing_if = "Option::is_none")] + pub usage_script: Option, + /// Usage query API key (if different from provider API key) + #[serde(skip_serializing_if = "Option::is_none")] + pub usage_api_key: Option, + /// Usage query base URL (if different from provider endpoint) + #[serde(skip_serializing_if = "Option::is_none")] + pub usage_base_url: Option, + /// Usage query access token (for NewAPI template) + #[serde(skip_serializing_if = "Option::is_none")] + pub usage_access_token: Option, + /// Usage query user ID (for NewAPI template) + #[serde(skip_serializing_if = "Option::is_none")] + pub usage_user_id: Option, + /// Auto query interval in minutes (0 to disable) + #[serde(skip_serializing_if = "Option::is_none")] + pub usage_auto_interval: Option, } diff --git a/src-tauri/src/deeplink/parser.rs b/src-tauri/src/deeplink/parser.rs index 316af0032..61cf7b68f 100644 --- a/src-tauri/src/deeplink/parser.rs +++ b/src-tauri/src/deeplink/parser.rs @@ -122,6 +122,19 @@ fn parse_provider_deeplink( let config_url = params.get("configUrl").cloned(); let enabled = params.get("enabled").and_then(|v| v.parse::().ok()); + // Extract usage script fields (v3.9+) + let usage_enabled = params + .get("usageEnabled") + .and_then(|v| v.parse::().ok()); + let usage_script = params.get("usageScript").cloned(); + let usage_api_key = params.get("usageApiKey").cloned(); + let usage_base_url = params.get("usageBaseUrl").cloned(); + let usage_access_token = params.get("usageAccessToken").cloned(); + let usage_user_id = params.get("usageUserId").cloned(); + let usage_auto_interval = params + .get("usageAutoInterval") + .and_then(|v| v.parse::().ok()); + Ok(DeepLinkImportRequest { version, resource, @@ -146,6 +159,13 @@ fn parse_provider_deeplink( config, config_format, config_url, + usage_enabled, + usage_script, + usage_api_key, + usage_base_url, + usage_access_token, + usage_user_id, + usage_auto_interval, }) } @@ -206,6 +226,13 @@ fn parse_prompt_deeplink( config: None, config_format: None, config_url: None, + usage_enabled: None, + usage_script: None, + usage_api_key: None, + usage_base_url: None, + usage_access_token: None, + usage_user_id: None, + usage_auto_interval: None, }) } @@ -261,6 +288,13 @@ fn parse_mcp_deeplink( directory: None, branch: None, config_url: None, + usage_enabled: None, + usage_script: None, + usage_api_key: None, + usage_base_url: None, + usage_access_token: None, + usage_user_id: None, + usage_auto_interval: None, }) } @@ -309,5 +343,12 @@ fn parse_skill_deeplink( config: None, config_format: None, config_url: None, + usage_enabled: None, + usage_script: None, + usage_api_key: None, + usage_base_url: None, + usage_access_token: None, + usage_user_id: None, + usage_auto_interval: None, }) } diff --git a/src-tauri/src/deeplink/provider.rs b/src-tauri/src/deeplink/provider.rs index 3b6cd8631..76fce86c1 100644 --- a/src-tauri/src/deeplink/provider.rs +++ b/src-tauri/src/deeplink/provider.rs @@ -5,7 +5,7 @@ use super::utils::{decode_base64_param, infer_homepage_from_endpoint}; use super::DeepLinkImportRequest; use crate::error::AppError; -use crate::provider::Provider; +use crate::provider::{Provider, ProviderMeta, UsageScript}; use crate::services::ProviderService; use crate::store::AppState; use crate::AppType; @@ -117,6 +117,9 @@ pub(crate) fn build_provider_from_request( AppType::Gemini => build_gemini_settings(request), }; + // Build usage script configuration if provided + let meta = build_provider_meta(request)?; + let provider = Provider { id: String::new(), // Will be generated by caller name: request.name.clone().unwrap_or_default(), @@ -126,15 +129,65 @@ pub(crate) fn build_provider_from_request( created_at: None, sort_index: None, notes: request.notes.clone(), - meta: None, + meta, icon: request.icon.clone(), icon_color: None, - is_proxy_target: None, }; Ok(provider) } +/// Build provider meta with usage script configuration +fn build_provider_meta(request: &DeepLinkImportRequest) -> Result, AppError> { + // Check if any usage script fields are provided + if request.usage_script.is_none() + && request.usage_enabled.is_none() + && request.usage_api_key.is_none() + && request.usage_base_url.is_none() + && request.usage_access_token.is_none() + && request.usage_user_id.is_none() + && request.usage_auto_interval.is_none() + { + return Ok(None); + } + + // Decode usage script code if provided + let code = if let Some(script_b64) = &request.usage_script { + let decoded = decode_base64_param("usage_script", script_b64)?; + String::from_utf8(decoded) + .map_err(|e| AppError::InvalidInput(format!("Invalid UTF-8 in usage_script: {e}")))? + } else { + String::new() + }; + + // Determine enabled state: explicit param > has code > false + let enabled = request.usage_enabled.unwrap_or(!code.is_empty()); + + // Build UsageScript - use provider's API key and endpoint as defaults + let usage_script = UsageScript { + enabled, + language: "javascript".to_string(), + code, + timeout: Some(10), + api_key: request + .usage_api_key + .clone() + .or_else(|| request.api_key.clone()), + base_url: request + .usage_base_url + .clone() + .or_else(|| request.endpoint.clone()), + access_token: request.usage_access_token.clone(), + user_id: request.usage_user_id.clone(), + auto_query_interval: request.usage_auto_interval, + }; + + Ok(Some(ProviderMeta { + usage_script: Some(usage_script), + ..Default::default() + })) +} + /// Build Claude settings configuration fn build_claude_settings(request: &DeepLinkImportRequest) -> serde_json::Value { let mut env = serde_json::Map::new(); diff --git a/src-tauri/src/deeplink/tests.rs b/src-tauri/src/deeplink/tests.rs index b20fa55c4..0f7fb5e45 100644 --- a/src-tauri/src/deeplink/tests.rs +++ b/src-tauri/src/deeplink/tests.rs @@ -145,6 +145,13 @@ fn test_build_gemini_provider_with_model() { content: None, description: None, enabled: None, + usage_enabled: None, + usage_script: None, + usage_api_key: None, + usage_base_url: None, + usage_access_token: None, + usage_user_id: None, + usage_auto_interval: None, }; let provider = build_provider_from_request(&AppType::Gemini, &request).unwrap(); @@ -191,6 +198,13 @@ fn test_build_gemini_provider_without_model() { content: None, description: None, enabled: None, + usage_enabled: None, + usage_script: None, + usage_api_key: None, + usage_base_url: None, + usage_access_token: None, + usage_user_id: None, + usage_auto_interval: None, }; let provider = build_provider_from_request(&AppType::Gemini, &request).unwrap(); @@ -232,6 +246,13 @@ fn test_parse_and_merge_config_claude() { content: None, description: None, enabled: None, + usage_enabled: None, + usage_script: None, + usage_api_key: None, + usage_base_url: None, + usage_access_token: None, + usage_user_id: None, + usage_auto_interval: None, }; let merged = parse_and_merge_config(&request).unwrap(); @@ -275,6 +296,13 @@ fn test_parse_and_merge_config_url_override() { content: None, description: None, enabled: None, + usage_enabled: None, + usage_script: None, + usage_api_key: None, + usage_base_url: None, + usage_access_token: None, + usage_user_id: None, + usage_auto_interval: None, }; let merged = parse_and_merge_config(&request).unwrap(); diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index c791a0a28..147bffd37 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -39,8 +39,8 @@ pub use mcp::{ }; pub use provider::{Provider, ProviderMeta}; pub use services::{ - ConfigService, EndpointLatency, McpService, PromptService, ProviderService, SkillService, - SpeedtestService, + ConfigService, EndpointLatency, McpService, PromptService, ProviderService, ProxyService, + SkillService, SpeedtestService, }; pub use settings::{update_settings, AppSettings}; pub use store::AppState; @@ -332,6 +332,11 @@ pub fn run() { let app_state = AppState::new(db); + // 设置 AppHandle 用于代理故障转移时的 UI 更新 + app_state + .proxy_service + .set_app_handle(app.handle().clone()); + // ============================================================ // 按表独立判断的导入逻辑(各类数据独立检查,互不影响) // ============================================================ @@ -522,10 +527,33 @@ pub fn run() { } } - // 自动启动代理服务器 + // 异常退出恢复 + 自动启动代理服务器 let app_handle = app.handle().clone(); tauri::async_runtime::spawn(async move { let state = app_handle.state::(); + + // 1. 检测异常退出并恢复 Live 配置 + match state.db.is_live_takeover_active().await { + Ok(true) => { + // 接管标志为 true 但代理未运行 → 上次异常退出 + if !state.proxy_service.is_running().await { + log::warn!("检测到上次异常退出,正在恢复 Live 配置..."); + if let Err(e) = state.proxy_service.recover_from_crash().await { + log::error!("恢复 Live 配置失败: {e}"); + } else { + log::info!("Live 配置已从异常退出中恢复"); + } + } + } + Ok(false) => { + // 正常状态,无需恢复 + } + Err(e) => { + log::error!("检查接管状态失败: {e}"); + } + } + + // 2. 自动启动代理服务器(如果配置为启用) match state.db.get_proxy_config().await { Ok(config) => { if config.enabled { @@ -553,7 +581,6 @@ pub fn run() { commands::update_provider, commands::delete_provider, commands::switch_provider, - commands::set_proxy_target_provider, commands::import_default_config, commands::get_claude_config_status, commands::get_config_status, @@ -656,13 +683,18 @@ pub fn run() { commands::is_live_takeover_active, commands::switch_proxy_provider, // 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, + // Failover queue management + commands::get_failover_queue, + commands::get_available_providers_for_failover, + commands::add_to_failover_queue, + commands::remove_from_failover_queue, + commands::reorder_failover_queue, + commands::set_failover_item_enabled, // Usage statistics commands::get_usage_summary, commands::get_usage_trends, @@ -697,6 +729,10 @@ pub fn run() { tauri::async_runtime::spawn(async move { cleanup_before_exit(&app_handle).await; log::info!("清理完成,退出应用"); + + // 短暂等待确保所有 I/O 操作(如数据库写入)刷新到磁盘 + tokio::time::sleep(std::time::Duration::from_millis(100)).await; + // 使用 std::process::exit 避免再次触发 ExitRequested std::process::exit(0); }); diff --git a/src-tauri/src/provider.rs b/src-tauri/src/provider.rs index 7de68b706..966b284c5 100644 --- a/src-tauri/src/provider.rs +++ b/src-tauri/src/provider.rs @@ -36,10 +36,6 @@ pub struct Provider { #[serde(skip_serializing_if = "Option::is_none")] #[serde(rename = "iconColor")] pub icon_color: Option, - /// 是否为代理目标(数据库专用字段,不写入配置文件) - #[serde(skip_serializing_if = "Option::is_none")] - #[serde(rename = "isProxyTarget")] - pub is_proxy_target: Option, } impl Provider { @@ -62,7 +58,6 @@ impl Provider { meta: None, icon: None, icon_color: None, - is_proxy_target: None, } } } diff --git a/src-tauri/src/proxy/circuit_breaker.rs b/src-tauri/src/proxy/circuit_breaker.rs index 8b932c75f..f2c5093ff 100644 --- a/src-tauri/src/proxy/circuit_breaker.rs +++ b/src-tauri/src/proxy/circuit_breaker.rs @@ -72,8 +72,10 @@ pub struct CircuitBreaker { failed_requests: Arc, /// 上次打开时间 last_opened_at: Arc>>, - /// 配置 - config: CircuitBreakerConfig, + /// 配置(支持热更新) + config: Arc>, + /// 半开状态已放行的请求数(用于限流) + half_open_requests: Arc, } impl CircuitBreaker { @@ -86,20 +88,35 @@ impl CircuitBreaker { total_requests: Arc::new(AtomicU32::new(0)), failed_requests: Arc::new(AtomicU32::new(0)), last_opened_at: Arc::new(RwLock::new(None)), - config, + config: Arc::new(RwLock::new(config)), + half_open_requests: Arc::new(AtomicU32::new(0)), } } - /// 检查是否允许请求通过 - pub async fn allow_request(&self) -> bool { + /// 更新熔断器配置(热更新,不重置状态) + pub async fn update_config(&self, new_config: CircuitBreakerConfig) { + *self.config.write().await = new_config; + log::debug!("Circuit breaker config updated"); + } + + /// 判断当前 Provider 是否“可被纳入候选链路” + /// + /// 这个方法不会占用 HalfOpen 探测名额,仅用于路由选择阶段的“可用性判断”: + /// - Closed / HalfOpen:可用(返回 true) + /// - Open:若超时到达则切到 HalfOpen 并返回 true,否则返回 false + /// + /// 注意:真正发起请求前仍需调用 `allow_request()` 来获取 HalfOpen 探测名额, + /// 并在请求结束后通过 `record_success()` / `record_failure()` 释放。 + pub async fn is_available(&self) -> bool { let state = *self.state.read().await; + let config = self.config.read().await; match state { - CircuitState::Closed => true, + CircuitState::Closed | CircuitState::HalfOpen => true, CircuitState::Open => { - // 检查是否应该尝试半开 if let Some(opened_at) = *self.last_opened_at.read().await { - if opened_at.elapsed().as_secs() >= self.config.timeout_seconds { + if opened_at.elapsed().as_secs() >= config.timeout_seconds { + drop(config); // 释放读锁再转换状态 log::info!( "Circuit breaker transitioning from Open to HalfOpen (timeout reached)" ); @@ -109,13 +126,62 @@ impl CircuitBreaker { } false } - CircuitState::HalfOpen => true, + } + } + + /// 检查是否允许请求通过 + pub async fn allow_request(&self) -> bool { + let state = *self.state.read().await; + let config = self.config.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() >= config.timeout_seconds { + drop(config); // 释放读锁再转换状态 + log::info!( + "Circuit breaker transitioning from Open to HalfOpen (timeout reached)" + ); + self.transition_to_half_open().await; + // 增加计数,确保 record_success/record_failure 减计数时不会下溢 + self.half_open_requests.fetch_add(1, Ordering::SeqCst); + return true; + } + } + false + } + CircuitState::HalfOpen => { + // 半开状态限流:只允许有限请求通过进行探测 + // 默认最多允许 1 个请求(可在配置中扩展) + let max_half_open_requests = 1u32; + let current = self.half_open_requests.fetch_add(1, Ordering::SeqCst); + + if current < max_half_open_requests { + log::debug!( + "Circuit breaker HalfOpen: allowing probe request ({}/{})", + current + 1, + max_half_open_requests + ); + true + } else { + // 超过限额,回退计数,拒绝请求 + self.half_open_requests.fetch_sub(1, Ordering::SeqCst); + log::debug!( + "Circuit breaker HalfOpen: rejecting request (limit reached: {})", + max_half_open_requests + ); + false + } + } } } /// 记录成功 pub async fn record_success(&self) { let state = *self.state.read().await; + let config = self.config.read().await; // 重置失败计数 self.consecutive_failures.store(0, Ordering::SeqCst); @@ -123,14 +189,18 @@ impl CircuitBreaker { match state { CircuitState::HalfOpen => { + // 释放 in-flight 名额(探测请求结束) + self.half_open_requests.fetch_sub(1, Ordering::SeqCst); + let successes = self.consecutive_successes.fetch_add(1, Ordering::SeqCst) + 1; log::debug!( "Circuit breaker HalfOpen: {} consecutive successes (threshold: {})", successes, - self.config.success_threshold + config.success_threshold ); - if successes >= self.config.success_threshold { + if successes >= config.success_threshold { + drop(config); // 释放读锁再转换状态 log::info!("Circuit breaker transitioning from HalfOpen to Closed (success threshold reached)"); self.transition_to_closed().await; } @@ -145,6 +215,7 @@ impl CircuitBreaker { /// 记录失败 pub async fn record_failure(&self) { let state = *self.state.read().await; + let config = self.config.read().await; // 更新计数器 let failures = self.consecutive_failures.fetch_add(1, Ordering::SeqCst) + 1; @@ -158,26 +229,38 @@ impl CircuitBreaker { "Circuit breaker {:?}: {} consecutive failures (threshold: {})", state, failures, - self.config.failure_threshold + config.failure_threshold ); // 检查是否应该打开熔断器 match state { - CircuitState::Closed | CircuitState::HalfOpen => { + CircuitState::HalfOpen => { + // 释放 in-flight 名额(探测请求结束) + self.half_open_requests.fetch_sub(1, Ordering::SeqCst); + + // HalfOpen 状态下失败,立即转为 Open + log::warn!( + "Circuit breaker HalfOpen probe failed, transitioning to Open" + ); + drop(config); + self.transition_to_open().await; + } + CircuitState::Closed => { // 检查连续失败次数 - if failures >= self.config.failure_threshold { + if failures >= config.failure_threshold { log::warn!( "Circuit breaker opening due to {} consecutive failures (threshold: {})", failures, - self.config.failure_threshold + config.failure_threshold ); + drop(config); // 释放读锁再转换状态 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 { + if total >= config.min_requests { let error_rate = failed as f64 / total as f64; log::debug!( "Circuit breaker error rate: {:.2}% ({}/{} requests)", @@ -186,12 +269,13 @@ impl CircuitBreaker { total ); - if error_rate >= self.config.error_rate_threshold { + if error_rate >= 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 + config.error_rate_threshold * 100.0 ); + drop(config); // 释放读锁再转换状态 self.transition_to_open().await; } } @@ -238,6 +322,8 @@ impl CircuitBreaker { async fn transition_to_half_open(&self) { *self.state.write().await = CircuitState::HalfOpen; self.consecutive_successes.store(0, Ordering::SeqCst); + // 重置半开状态的请求限流计数 + self.half_open_requests.store(0, Ordering::SeqCst); } /// 转换到关闭状态 diff --git a/src-tauri/src/proxy/failover_switch.rs b/src-tauri/src/proxy/failover_switch.rs new file mode 100644 index 000000000..22d636e34 --- /dev/null +++ b/src-tauri/src/proxy/failover_switch.rs @@ -0,0 +1,148 @@ +//! 故障转移切换模块 +//! +//! 处理故障转移成功后的供应商切换逻辑,包括: +//! - 去重控制(避免多个请求同时触发) +//! - 数据库更新 +//! - 托盘菜单更新 +//! - 前端事件发射 +//! - Live 备份更新 + +use crate::database::Database; +use crate::error::AppError; +use std::collections::HashSet; +use std::str::FromStr; +use std::sync::Arc; +use tauri::{Emitter, Manager}; +use tokio::sync::RwLock; + +/// 故障转移切换管理器 +/// +/// 负责处理故障转移成功后的供应商切换,确保 UI 能够直观反映当前使用的供应商。 +#[derive(Clone)] +pub struct FailoverSwitchManager { + /// 正在处理中的切换(key = "app_type:provider_id") + pending_switches: Arc>>, + db: Arc, +} + +impl FailoverSwitchManager { + pub fn new(db: Arc) -> Self { + Self { + pending_switches: Arc::new(RwLock::new(HashSet::new())), + db, + } + } + + /// 尝试执行故障转移切换 + /// + /// 如果相同的切换已在进行中,则跳过;否则执行切换逻辑。 + /// + /// # Returns + /// - `Ok(true)` - 切换成功执行 + /// - `Ok(false)` - 切换已在进行中,跳过 + /// - `Err(e)` - 切换过程中发生错误 + pub async fn try_switch( + &self, + app_handle: Option<&tauri::AppHandle>, + app_type: &str, + provider_id: &str, + provider_name: &str, + ) -> Result { + let switch_key = format!("{}:{}", app_type, provider_id); + + // 去重检查:如果相同切换已在进行中,跳过 + { + let mut pending = self.pending_switches.write().await; + if pending.contains(&switch_key) { + log::debug!( + "[Failover] 切换已在进行中,跳过: {} -> {}", + app_type, + provider_id + ); + return Ok(false); + } + pending.insert(switch_key.clone()); + } + + // 执行切换(确保最后清理 pending 标记) + let result = self + .do_switch(app_handle, app_type, provider_id, provider_name) + .await; + + // 清理 pending 标记 + { + let mut pending = self.pending_switches.write().await; + pending.remove(&switch_key); + } + + result + } + + async fn do_switch( + &self, + app_handle: Option<&tauri::AppHandle>, + app_type: &str, + provider_id: &str, + provider_name: &str, + ) -> Result { + log::info!( + "[Failover] 开始切换供应商: {} -> {} ({})", + app_type, + provider_name, + provider_id + ); + + // 1. 更新数据库 is_current + self.db.set_current_provider(app_type, provider_id)?; + + // 2. 更新本地 settings(设备级) + let app_type_enum = crate::app_config::AppType::from_str(app_type) + .map_err(|_| AppError::Message(format!("无效的应用类型: {}", app_type)))?; + crate::settings::set_current_provider(&app_type_enum, Some(provider_id))?; + + // 3. 更新托盘菜单和发射事件 + if let Some(app) = app_handle { + // 更新托盘菜单 + if let Some(app_state) = app.try_state::() { + // 更新 Live 备份(确保代理停止时恢复正确配置) + if let Ok(Some(provider)) = self.db.get_provider_by_id(provider_id, app_type) { + if let Err(e) = app_state + .proxy_service + .update_live_backup_from_provider(app_type, &provider) + .await + { + log::warn!("[Failover] 更新 Live 备份失败: {e}"); + } + } + + // 重建托盘菜单 + if let Ok(new_menu) = crate::tray::create_tray_menu(app, app_state.inner()) { + if let Some(tray) = app.tray_by_id("main") { + if let Err(e) = tray.set_menu(Some(new_menu)) { + log::error!("[Failover] 更新托盘菜单失败: {e}"); + } + } + } + } + + // 发射事件到前端 + let event_data = serde_json::json!({ + "appType": app_type, + "providerId": provider_id, + "source": "failover" // 标识来源是故障转移 + }); + if let Err(e) = app.emit("provider-switched", event_data) { + log::error!("[Failover] 发射供应商切换事件失败: {e}"); + } + } + + log::info!( + "[Failover] 供应商切换完成: {} -> {} ({})", + app_type, + provider_name, + provider_id + ); + + Ok(true) + } +} diff --git a/src-tauri/src/proxy/forwarder.rs b/src-tauri/src/proxy/forwarder.rs index b4226ba00..8d2300008 100644 --- a/src-tauri/src/proxy/forwarder.rs +++ b/src-tauri/src/proxy/forwarder.rs @@ -4,12 +4,13 @@ use super::{ error::*, - provider_router::ProviderRouter as NewProviderRouter, + failover_switch::FailoverSwitchManager, + provider_router::ProviderRouter, providers::{get_adapter, ProviderAdapter}, types::ProxyStatus, ProxyError, }; -use crate::{app_config::AppType, database::Database, provider::Provider}; +use crate::{app_config::AppType, provider::Provider}; use reqwest::{Client, Response}; use serde_json::Value; use std::sync::Arc; @@ -18,20 +19,27 @@ use tokio::sync::RwLock; pub struct RequestForwarder { client: Client, - router: Arc, - #[allow(dead_code)] + /// 共享的 ProviderRouter(持有熔断器状态) + router: Arc, + /// 单个 Provider 内的最大重试次数 max_retries: u8, status: Arc>, current_providers: Arc>>, + /// 故障转移切换管理器 + failover_manager: Arc, + /// AppHandle,用于发射事件和更新托盘 + app_handle: Option, } impl RequestForwarder { pub fn new( - db: Arc, + router: Arc, timeout_secs: u64, max_retries: u8, status: Arc>, current_providers: Arc>>, + failover_manager: Arc, + app_handle: Option, ) -> Self { let mut client_builder = Client::builder(); if timeout_secs > 0 { @@ -44,32 +52,90 @@ impl RequestForwarder { Self { client, - router: Arc::new(NewProviderRouter::new(db)), + router, max_retries, status, current_providers, + failover_manager, + app_handle, } } + /// 对单个 Provider 执行请求(带重试) + /// + /// 在同一个 Provider 上最多重试 max_retries 次,使用指数退避 + async fn forward_with_provider_retry( + &self, + provider: &Provider, + endpoint: &str, + body: &Value, + headers: &axum::http::HeaderMap, + adapter: &dyn ProviderAdapter, + ) -> Result { + let mut last_error = None; + + for attempt in 0..=self.max_retries { + if attempt > 0 { + // 指数退避:100ms, 200ms, 400ms, ... + let delay_ms = 100 * 2u64.pow(attempt as u32 - 1); + log::info!( + "[{}] 重试第 {}/{} 次(等待 {}ms)", + adapter.name(), + attempt, + self.max_retries, + delay_ms + ); + tokio::time::sleep(Duration::from_millis(delay_ms)).await; + } + + match self + .forward(provider, endpoint, body, headers, adapter) + .await + { + Ok(response) => return Ok(response), + Err(e) => { + let category = self.categorize_proxy_error(&e); + + // 只有可重试的错误才继续重试 + if category == ErrorCategory::NonRetryable { + return Err(e); + } + + log::debug!( + "[{}] Provider {} 第 {} 次请求失败: {}", + adapter.name(), + provider.name, + attempt + 1, + e + ); + last_error = Some(e); + } + } + } + + Err(last_error.unwrap_or(ProxyError::MaxRetriesExceeded)) + } + /// 转发请求(带故障转移) + /// + /// # Arguments + /// * `app_type` - 应用类型 + /// * `endpoint` - API 端点 + /// * `body` - 请求体 + /// * `headers` - 请求头 + /// * `providers` - 已选择的 Provider 列表(由 RequestContext 提供,避免重复调用 select_providers) pub async fn forward_with_retry( &self, app_type: &AppType, endpoint: &str, body: Value, headers: axum::http::HeaderMap, + providers: Vec, ) -> Result { // 获取适配器 let adapter = get_adapter(app_type); let app_type_str = app_type.as_str(); - // 使用新的 ProviderRouter 选择所有可用供应商 - let providers = self - .router - .select_providers(app_type_str) - .await - .map_err(|e| ProxyError::DatabaseError(e.to_string()))?; - if providers.is_empty() { return Err(ProxyError::NoAvailableProvider); } @@ -82,13 +148,33 @@ impl RequestForwarder { let mut last_error = None; let mut failover_happened = false; + let mut attempted_providers = 0usize; // 依次尝试每个供应商 - for (attempt, provider) in providers.iter().enumerate() { + for provider in providers.iter() { + // 发起请求前先获取熔断器放行许可(HalfOpen 会占用探测名额) + if !self + .router + .allow_provider_request(&provider.id, app_type_str) + .await + { + log::debug!( + "[{}] Provider {} 熔断器拒绝本次请求,跳过", + app_type_str, + provider.name + ); + continue; + } + + attempted_providers += 1; + if attempted_providers > 1 { + failover_happened = true; + } + log::info!( "[{}] 尝试 {}/{} - 使用Provider: {} (sort_index: {})", app_type_str, - attempt + 1, + attempted_providers, providers.len(), provider.name, provider.sort_index.unwrap_or(999999) @@ -101,16 +187,13 @@ impl RequestForwarder { status.current_provider_id = Some(provider.id.clone()); status.total_requests += 1; status.last_request_at = Some(chrono::Utc::now().to_rfc3339()); - if attempt > 0 { - failover_happened = true; - } } let start = Instant::now(); - // 转发请求 + // 转发请求(带单 Provider 内重试) match self - .forward(provider, endpoint, &body, &headers, adapter.as_ref()) + .forward_with_provider_retry(provider, endpoint, &body, &headers, adapter.as_ref()) .await { Ok(response) => { @@ -147,6 +230,20 @@ impl RequestForwarder { provider.name, latency ); + + // 异步触发供应商切换,更新 UI 和托盘菜单 + let fm = self.failover_manager.clone(); + let ah = self.app_handle.clone(); + let pid = provider.id.clone(); + let pname = provider.name.clone(); + let at = app_type_str.to_string(); + + tokio::spawn(async move { + if let Err(e) = fm.try_switch(ah.as_ref(), &at, &pid, &pname).await + { + log::error!("[Failover] 切换供应商失败: {e}"); + } + }); } // 重新计算成功率 if status.total_requests > 0 { @@ -226,6 +323,21 @@ impl RequestForwarder { } } + if attempted_providers == 0 { + // providers 列表非空,但全部被熔断器拒绝(典型:HalfOpen 探测名额被占用) + { + let mut status = self.status.write().await; + status.failed_requests += 1; + 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; + } + } + return Err(ProxyError::NoAvailableProvider); + } + // 所有供应商都失败了 { let mut status = self.status.write().await; @@ -373,21 +485,24 @@ impl RequestForwarder { } /// 分类ProxyError + /// + /// 决定哪些错误应该触发故障转移到下一个 Provider + /// + /// 设计原则:既然用户配置了多个供应商,就应该让所有供应商都尝试一遍。 + /// 只有明确是客户端中断的情况才不重试。 fn categorize_proxy_error(&self, error: &ProxyError) -> ErrorCategory { match error { + // 网络和上游错误:都应该尝试下一个供应商 ProxyError::Timeout(_) => ErrorCategory::Retryable, ProxyError::ForwardFailed(_) => ErrorCategory::Retryable, - ProxyError::UpstreamError { status, .. } => { - if *status >= 500 { - ErrorCategory::Retryable - } else if *status >= 400 && *status < 500 { - ErrorCategory::NonRetryable - } else { - ErrorCategory::Retryable - } - } ProxyError::ProviderUnhealthy(_) => ErrorCategory::Retryable, + // 上游 HTTP 错误:无论状态码如何,都尝试下一个供应商 + // 原因:不同供应商有不同的限制和认证,一个供应商的 4xx 错误 + // 不代表其他供应商也会失败 + ProxyError::UpstreamError { .. } => ErrorCategory::Retryable, + // 无可用供应商:所有供应商都试过了,无法重试 ProxyError::NoAvailableProvider => ErrorCategory::NonRetryable, + // 其他错误(配置错误、数据库错误等):不是供应商问题,无需重试 _ => ErrorCategory::NonRetryable, } } diff --git a/src-tauri/src/proxy/handler_config.rs b/src-tauri/src/proxy/handler_config.rs new file mode 100644 index 000000000..5710f0b48 --- /dev/null +++ b/src-tauri/src/proxy/handler_config.rs @@ -0,0 +1,164 @@ +//! Handler 配置模块 +//! +//! 定义各 API 处理器的配置结构和使用量解析器 + +use crate::app_config::AppType; +use crate::proxy::usage::parser::TokenUsage; +use serde_json::Value; + +/// 使用量解析器类型别名 +pub type StreamUsageParser = fn(&[Value]) -> Option; +pub type ResponseUsageParser = fn(&Value) -> Option; + +/// 模型提取器类型别名 +/// 参数: (流式事件列表, 请求中的模型名称) -> 最终使用的模型名称 +pub type StreamModelExtractor = fn(&[Value], &str) -> String; + +/// 各 API 的使用量解析配置 +#[derive(Clone, Copy)] +pub struct UsageParserConfig { + /// 流式响应解析器 + pub stream_parser: StreamUsageParser, + /// 非流式响应解析器 + pub response_parser: ResponseUsageParser, + /// 流式响应中的模型提取器 + pub model_extractor: StreamModelExtractor, + /// 应用类型字符串(用于日志记录) + pub app_type_str: &'static str, +} + +// ============================================================================ +// 模型提取器实现 +// ============================================================================ + +/// Claude 流式响应模型提取(直接使用请求模型) +fn claude_model_extractor(_events: &[Value], request_model: &str) -> String { + request_model.to_string() +} + +/// OpenAI Chat Completions 流式响应模型提取 +fn openai_model_extractor(events: &[Value], request_model: &str) -> String { + events + .iter() + .find_map(|e| e.get("model")?.as_str()) + .unwrap_or(request_model) + .to_string() +} + +/// Codex Responses API 流式响应模型提取 +fn codex_model_extractor(events: &[Value], request_model: &str) -> String { + events + .iter() + .find_map(|e| { + if e.get("type")?.as_str()? == "response.completed" { + e.get("response")?.get("model")?.as_str() + } else { + None + } + }) + .unwrap_or(request_model) + .to_string() +} + +/// Gemini 流式响应模型提取(优先使用 usage.model) +fn gemini_model_extractor(events: &[Value], request_model: &str) -> String { + // 首先尝试从解析的 usage 中获取模型 + if let Some(usage) = TokenUsage::from_gemini_stream_chunks(events) { + if let Some(model) = usage.model { + return model; + } + } + request_model.to_string() +} + +// ============================================================================ +// 预定义配置 +// ============================================================================ + +/// Claude API 解析配置 +pub const CLAUDE_PARSER_CONFIG: UsageParserConfig = UsageParserConfig { + stream_parser: TokenUsage::from_claude_stream_events, + response_parser: TokenUsage::from_claude_response, + model_extractor: claude_model_extractor, + app_type_str: "claude", +}; + +/// OpenAI Chat Completions API 解析配置(用于 Codex /v1/chat/completions) +pub const OPENAI_PARSER_CONFIG: UsageParserConfig = UsageParserConfig { + stream_parser: TokenUsage::from_openai_stream_events, + response_parser: TokenUsage::from_openai_response, + model_extractor: openai_model_extractor, + app_type_str: "codex", +}; + +/// Codex Responses API 解析配置(用于 /v1/responses) +pub const CODEX_PARSER_CONFIG: UsageParserConfig = UsageParserConfig { + stream_parser: TokenUsage::from_codex_stream_events, + response_parser: TokenUsage::from_codex_response, + model_extractor: codex_model_extractor, + app_type_str: "codex", +}; + +/// Gemini API 解析配置 +pub const GEMINI_PARSER_CONFIG: UsageParserConfig = UsageParserConfig { + stream_parser: TokenUsage::from_gemini_stream_chunks, + response_parser: TokenUsage::from_gemini_response, + model_extractor: gemini_model_extractor, + app_type_str: "gemini", +}; + +// ============================================================================ +// Handler 配置(预留,用于进一步简化) +// ============================================================================ + +/// Handler 基础配置 +/// +/// 预留结构,可用于进一步统一各 handler 的配置 +#[allow(dead_code)] +#[derive(Clone)] +pub struct HandlerConfig { + /// 应用类型 + pub app_type: AppType, + /// 日志标签 + pub tag: &'static str, + /// 应用类型字符串 + pub app_type_str: &'static str, + /// 使用量解析配置 + pub parser_config: &'static UsageParserConfig, +} + +/// Claude Handler 配置 +#[allow(dead_code)] +pub const CLAUDE_HANDLER_CONFIG: HandlerConfig = HandlerConfig { + app_type: AppType::Claude, + tag: "Claude", + app_type_str: "claude", + parser_config: &CLAUDE_PARSER_CONFIG, +}; + +/// Codex Chat Completions Handler 配置 +#[allow(dead_code)] +pub const CODEX_CHAT_HANDLER_CONFIG: HandlerConfig = HandlerConfig { + app_type: AppType::Codex, + tag: "Codex", + app_type_str: "codex", + parser_config: &OPENAI_PARSER_CONFIG, +}; + +/// Codex Responses Handler 配置 +#[allow(dead_code)] +pub const CODEX_RESPONSES_HANDLER_CONFIG: HandlerConfig = HandlerConfig { + app_type: AppType::Codex, + tag: "Codex", + app_type_str: "codex", + parser_config: &CODEX_PARSER_CONFIG, +}; + +/// Gemini Handler 配置 +#[allow(dead_code)] +pub const GEMINI_HANDLER_CONFIG: HandlerConfig = HandlerConfig { + app_type: AppType::Gemini, + tag: "Gemini", + app_type_str: "gemini", + parser_config: &GEMINI_PARSER_CONFIG, +}; diff --git a/src-tauri/src/proxy/handler_context.rs b/src-tauri/src/proxy/handler_context.rs new file mode 100644 index 000000000..e8ba7c214 --- /dev/null +++ b/src-tauri/src/proxy/handler_context.rs @@ -0,0 +1,151 @@ +//! 请求上下文模块 +//! +//! 提供请求生命周期的上下文管理,封装通用初始化逻辑 + +use crate::app_config::AppType; +use crate::provider::Provider; +use crate::proxy::{ + forwarder::RequestForwarder, server::ProxyState, types::ProxyConfig, ProxyError, +}; +use std::time::Instant; + +/// 请求上下文 +/// +/// 贯穿整个请求生命周期,包含: +/// - 计时信息 +/// - 代理配置 +/// - 选中的 Provider 列表(用于故障转移) +/// - 请求模型名称 +/// - 日志标签 +pub struct RequestContext { + /// 请求开始时间 + pub start_time: Instant, + /// 代理配置快照 + pub config: ProxyConfig, + /// 选中的 Provider(故障转移链的第一个) + pub provider: Provider, + /// 完整的 Provider 列表(用于故障转移) + providers: Vec, + /// 请求中的模型名称 + pub request_model: String, + /// 日志标签(如 "Claude"、"Codex"、"Gemini") + pub tag: &'static str, + /// 应用类型字符串(如 "claude"、"codex"、"gemini") + pub app_type_str: &'static str, + /// 应用类型(预留,目前通过 app_type_str 使用) + #[allow(dead_code)] + pub app_type: AppType, +} + +impl RequestContext { + /// 创建请求上下文 + /// + /// # Arguments + /// * `state` - 代理服务器状态 + /// * `body` - 请求体 JSON + /// * `app_type` - 应用类型 + /// * `tag` - 日志标签 + /// * `app_type_str` - 应用类型字符串 + /// + /// # Errors + /// 返回 `ProxyError` 如果 Provider 选择失败 + pub async fn new( + state: &ProxyState, + body: &serde_json::Value, + app_type: AppType, + tag: &'static str, + app_type_str: &'static str, + ) -> Result { + let start_time = Instant::now(); + let config = state.config.read().await.clone(); + + // 从请求体提取模型名称 + let request_model = body + .get("model") + .and_then(|m| m.as_str()) + .unwrap_or("unknown") + .to_string(); + + // 使用共享的 ProviderRouter 选择 Provider(熔断器状态跨请求保持) + // 注意:只在这里调用一次,结果传递给 forwarder,避免重复消耗 HalfOpen 名额 + let providers = state + .provider_router + .select_providers(app_type_str) + .await + .map_err(|e| ProxyError::DatabaseError(e.to_string()))?; + + let provider = providers + .first() + .cloned() + .ok_or(ProxyError::NoAvailableProvider)?; + + log::info!( + "[{}] Provider: {}, model: {}, failover chain: {} providers", + tag, + provider.name, + request_model, + providers.len() + ); + + Ok(Self { + start_time, + config, + provider, + providers, + request_model, + tag, + app_type_str, + app_type, + }) + } + + /// 从 URI 提取模型名称(Gemini 专用) + /// + /// Gemini API 的模型名称在 URI 中,格式如: + /// `/v1beta/models/gemini-pro:generateContent` + pub fn with_model_from_uri(mut self, uri: &axum::http::Uri) -> Self { + let endpoint = uri + .path_and_query() + .map(|pq| pq.as_str()) + .unwrap_or(uri.path()); + + self.request_model = endpoint + .split('/') + .find(|s| s.starts_with("models/")) + .and_then(|s| s.strip_prefix("models/")) + .map(|s| s.split(':').next().unwrap_or(s)) + .unwrap_or("unknown") + .to_string(); + + log::info!("[{}] 从 URI 提取模型: {}", self.tag, self.request_model); + self + } + + /// 创建 RequestForwarder + /// + /// 使用共享的 ProviderRouter,确保熔断器状态跨请求保持 + pub fn create_forwarder(&self, state: &ProxyState) -> RequestForwarder { + RequestForwarder::new( + state.provider_router.clone(), + self.config.request_timeout, + self.config.max_retries, + state.status.clone(), + state.current_providers.clone(), + state.failover_manager.clone(), + state.app_handle.clone(), + ) + } + + /// 获取 Provider 列表(用于故障转移) + /// + /// 返回在创建上下文时已选择的 providers,避免重复调用 select_providers() + pub fn get_providers(&self) -> Vec { + self.providers.clone() + } + + /// 计算请求延迟(毫秒) + #[inline] + pub fn latency_ms(&self) -> u64 { + self.start_time.elapsed().as_millis() as u64 + } +} diff --git a/src-tauri/src/proxy/handlers.rs b/src-tauri/src/proxy/handlers.rs index 36d95390e..b5949f3ad 100644 --- a/src-tauri/src/proxy/handlers.rs +++ b/src-tauri/src/proxy/handlers.rs @@ -1,88 +1,456 @@ //! 请求处理器 //! //! 处理各种API端点的HTTP请求 +//! +//! 重构后的结构: +//! - 通用逻辑提取到 `handler_context` 和 `response_processor` 模块 +//! - 各 handler 只保留独特的业务逻辑 +//! - Claude 的格式转换逻辑保留在此文件(独有功能) use super::{ error_mapper::{get_error_message, map_proxy_error_to_status}, - forwarder::RequestForwarder, - providers::{get_adapter, transform, ProviderType}, + handler_config::{ + CLAUDE_PARSER_CONFIG, CODEX_PARSER_CONFIG, GEMINI_PARSER_CONFIG, OPENAI_PARSER_CONFIG, + }, + handler_context::RequestContext, + providers::{get_adapter, streaming::create_anthropic_sse_stream, transform}, + response_processor::{ + create_logged_passthrough_stream, process_response, SseUsageCollector, + }, server::ProxyState, - session::ProxySession, types::*, - usage::{logger::UsageLogger, parser::TokenUsage}, + usage::parser::TokenUsage, ProxyError, }; use crate::app_config::AppType; use axum::{extract::State, http::StatusCode, response::IntoResponse, Json}; -use bytes::Bytes; -use futures::stream::{Stream, StreamExt}; use rust_decimal::Decimal; use serde_json::{json, Value}; -use std::{ - str::FromStr, - sync::{ - atomic::{AtomicBool, Ordering}, - Arc, - }, -}; -use tokio::sync::Mutex; +use std::str::FromStr; -/// 记录请求使用量(带 ProxySession 支持) -#[allow(dead_code, clippy::too_many_arguments)] -async fn log_usage_with_session( - state: &ProxyState, - session: &ProxySession, - provider_id: &str, - app_type: &str, - usage: TokenUsage, - latency_ms: u64, - first_token_ms: Option, - status_code: u16, - provider_type: Option<&ProviderType>, -) { - let logger = UsageLogger::new(&state.db); +// ============================================================================ +// 健康检查和状态查询(简单端点) +// ============================================================================ - // 获取 provider 的 cost_multiplier - let multiplier = match state.db.get_provider_by_id(provider_id, app_type) { - Ok(Some(p)) => { - if let Some(meta) = p.meta { - if let Some(cm) = meta.cost_multiplier { - Decimal::from_str(&cm).unwrap_or(Decimal::from(1)) - } else { - Decimal::from(1) - } - } else { - Decimal::from(1) - } +/// 健康检查 +pub async fn health_check() -> (StatusCode, Json) { + ( + StatusCode::OK, + Json(json!({ + "status": "healthy", + "timestamp": chrono::Utc::now().to_rfc3339(), + })), + ) +} + +/// 获取服务状态 +pub async fn get_status(State(state): State) -> Result, ProxyError> { + let status = state.status.read().await.clone(); + Ok(Json(status)) +} + +// ============================================================================ +// Claude API 处理器(包含格式转换逻辑) +// ============================================================================ + +/// 处理 /v1/messages 请求(Claude API) +/// +/// Claude 处理器包含独特的格式转换逻辑: +/// - 当使用 OpenRouter 等中转服务时,需要将 Anthropic 格式转换为 OpenAI 格式 +/// - 响应需要从 OpenAI 格式转回 Anthropic 格式 +pub async fn handle_messages( + State(state): State, + 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 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 + .forward_with_retry( + &AppType::Claude, + "/v1/messages", + body.clone(), + headers, + ctx.get_providers(), + ) + .await + { + Ok(resp) => resp, + Err(e) => { + log_forward_error(&state, &ctx, is_stream, &e); + return Err(e); } - _ => Decimal::from(1), }; - let model = session - .model - .clone() - .unwrap_or_else(|| "unknown".to_string()); - let provider_type_str = provider_type.map(|pt| pt.as_str().to_string()); + let status = response.status(); + log::info!("[Claude] 上游响应状态: {status}"); - if let Err(e) = logger.log_with_calculation( - session.session_id.clone(), - provider_id.to_string(), - app_type.to_string(), - model, - usage, - multiplier, - latency_ms, - first_token_ms, + // Claude 特有:格式转换处理 + if needs_transform { + return handle_claude_transform(response, &ctx, &state, &body, is_stream).await; + } + + // 通用响应处理(透传模式) + process_response(response, &ctx, &state, &CLAUDE_PARSER_CONFIG).await +} + +/// Claude 格式转换处理(独有逻辑) +/// +/// 处理 OpenRouter 等需要格式转换的中转服务 +async fn handle_claude_transform( + response: reqwest::Response, + ctx: &RequestContext, + state: &ProxyState, + _original_body: &Value, + is_stream: bool, +) -> Result { + let status = response.status(); + + if is_stream { + // 流式响应转换 (OpenAI SSE → Anthropic SSE) + log::info!("[Claude] 开始流式响应转换 (OpenAI SSE → Anthropic SSE)"); + + let stream = response.bytes_stream(); + let sse_stream = create_anthropic_sse_stream(stream); + + // 创建使用量收集器 + let usage_collector = { + let state = state.clone(); + let provider_id = ctx.provider.id.clone(); + let model = ctx.request_model.clone(); + let status_code = status.as_u16(); + let start_time = ctx.start_time; + + SseUsageCollector::new(start_time, move |events, first_token_ms| { + if let Some(usage) = TokenUsage::from_claude_stream_events(&events) { + let latency_ms = start_time.elapsed().as_millis() as u64; + let state = state.clone(); + let provider_id = provider_id.clone(); + let model = model.clone(); + + tokio::spawn(async move { + log_usage( + &state, + &provider_id, + "claude", + &model, + usage, + latency_ms, + first_token_ms, + true, + status_code, + ) + .await; + }); + } else { + log::debug!("[Claude] OpenRouter 流式响应缺少 usage 统计,跳过消费记录"); + } + }) + }; + + let logged_stream = create_logged_passthrough_stream( + sse_stream, + "Claude/OpenRouter", + Some(usage_collector), + ); + + 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 = axum::body::Body::from_stream(logged_stream); + log::info!("[Claude] ====== 请求结束 (流式转换) ======"); + return Ok((headers, body).into_response()); + } + + // 非流式响应转换 (OpenAI → Anthropic) + log::info!("[Claude] 开始转换响应 (OpenAI → Anthropic)"); + + let response_headers = response.headers().clone(); + + let body_bytes = response.bytes().await.map_err(|e| { + log::error!("[Claude] 读取响应体失败: {e}"); + ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) + })?; + + let body_str = String::from_utf8_lossy(&body_bytes); + log::info!("[Claude] OpenAI 响应长度: {} bytes", body_bytes.len()); + log::debug!("[Claude] OpenAI 原始响应: {body_str}"); + + 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}")) + })?; + + log::info!("[Claude] 解析 OpenAI 响应成功"); + log::info!( + "[Claude] <<< OpenAI 响应 JSON:\n{}", + serde_json::to_string_pretty(&openai_response).unwrap_or_default() + ); + + let anthropic_response = transform::openai_to_anthropic(openai_response).map_err(|e| { + log::error!("[Claude] 转换响应失败: {e}"); + e + })?; + + log::info!("[Claude] 转换响应成功"); + log::info!( + "[Claude] <<< Anthropic 响应 JSON:\n{}", + serde_json::to_string_pretty(&anthropic_response).unwrap_or_default() + ); + + // 记录使用量 + if let Some(usage) = TokenUsage::from_claude_response(&anthropic_response) { + let model = anthropic_response + .get("model") + .and_then(|m| m.as_str()) + .unwrap_or("unknown"); + let latency_ms = ctx.latency_ms(); + + tokio::spawn({ + let state = state.clone(); + let provider_id = ctx.provider.id.clone(); + let model = model.to_string(); + async move { + log_usage( + &state, + &provider_id, + "claude", + &model, + usage, + latency_ms, + None, + false, + status.as_u16(), + ) + .await; + } + }); + } + + log::info!("[Claude] ====== 请求结束 ======"); + + // 构建响应 + let mut builder = axum::response::Response::builder().status(status); + + 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); + Ok(builder.body(body).unwrap()) +} + +// ============================================================================ +// Codex API 处理器 +// ============================================================================ + +/// 处理 /v1/chat/completions 请求(OpenAI Chat Completions API - Codex CLI) +pub async fn handle_chat_completions( + State(state): State, + headers: axum::http::HeaderMap, + Json(body): Json, +) -> Result { + log::info!("[Codex] ====== /v1/chat/completions 请求开始 ======"); + + let ctx = RequestContext::new(&state, &body, AppType::Codex, "Codex", "codex").await?; + + let is_stream = body + .get("stream") + .and_then(|v| v.as_bool()) + .unwrap_or(false); + + log::info!( + "[Codex] 请求模型: {}, 流式: {}", + ctx.request_model, + is_stream + ); + + let forwarder = ctx.create_forwarder(&state); + let response = match forwarder + .forward_with_retry( + &AppType::Codex, + "/v1/chat/completions", + body, + headers, + ctx.get_providers(), + ) + .await + { + Ok(resp) => resp, + Err(e) => { + log_forward_error(&state, &ctx, is_stream, &e); + return Err(e); + } + }; + + log::info!("[Codex] 上游响应状态: {}", response.status()); + + process_response(response, &ctx, &state, &OPENAI_PARSER_CONFIG).await +} + +/// 处理 /v1/responses 请求(OpenAI Responses API - Codex CLI 透传) +pub async fn handle_responses( + State(state): State, + headers: axum::http::HeaderMap, + Json(body): Json, +) -> Result { + let ctx = RequestContext::new(&state, &body, AppType::Codex, "Codex", "codex").await?; + + let is_stream = body + .get("stream") + .and_then(|v| v.as_bool()) + .unwrap_or(false); + + let forwarder = ctx.create_forwarder(&state); + let response = match forwarder + .forward_with_retry( + &AppType::Codex, + "/v1/responses", + body, + headers, + ctx.get_providers(), + ) + .await + { + Ok(resp) => resp, + Err(e) => { + log_forward_error(&state, &ctx, is_stream, &e); + return Err(e); + } + }; + + log::info!("[Codex] 上游响应状态: {}", response.status()); + + process_response(response, &ctx, &state, &CODEX_PARSER_CONFIG).await +} + +// ============================================================================ +// Gemini API 处理器 +// ============================================================================ + +/// 处理 Gemini API 请求(透传,包括查询参数) +pub async fn handle_gemini( + State(state): State, + uri: axum::http::Uri, + headers: axum::http::HeaderMap, + Json(body): Json, +) -> Result { + // Gemini 的模型名称在 URI 中 + let ctx = RequestContext::new(&state, &body, AppType::Gemini, "Gemini", "gemini") + .await? + .with_model_from_uri(&uri); + + // 提取完整的路径和查询参数 + let endpoint = uri + .path_and_query() + .map(|pq| pq.as_str()) + .unwrap_or(uri.path()); + + log::info!("[Gemini] 请求端点: {}", endpoint); + + let is_stream = body + .get("stream") + .and_then(|v| v.as_bool()) + .unwrap_or(false); + + let forwarder = ctx.create_forwarder(&state); + let response = match forwarder + .forward_with_retry( + &AppType::Gemini, + endpoint, + body, + headers, + ctx.get_providers(), + ) + .await + { + Ok(resp) => resp, + Err(e) => { + log_forward_error(&state, &ctx, is_stream, &e); + return Err(e); + } + }; + + log::info!("[Gemini] 上游响应状态: {}", response.status()); + + process_response(response, &ctx, &state, &GEMINI_PARSER_CONFIG).await +} + +// ============================================================================ +// 使用量记录(保留用于 Claude 转换逻辑) +// ============================================================================ + +fn log_forward_error(state: &ProxyState, ctx: &RequestContext, is_streaming: bool, error: &ProxyError) { + use super::usage::logger::UsageLogger; + + let logger = UsageLogger::new(&state.db); + let status_code = map_proxy_error_to_status(error); + let error_message = get_error_message(error); + let request_id = uuid::Uuid::new_v4().to_string(); + + if let Err(e) = logger.log_error_with_context( + request_id.clone(), + ctx.provider.id.clone(), + ctx.app_type_str.to_string(), + ctx.request_model.clone(), status_code, - Some(session.session_id.clone()), - provider_type_str, - session.is_streaming, + error_message, + ctx.latency_ms(), + is_streaming, + Some(request_id), + None, ) { - log::warn!("记录使用量失败: {e}"); + log::warn!("记录失败请求日志失败: {e}"); } } -/// 记录请求使用量(兼容旧接口) +/// 记录请求使用量 #[allow(clippy::too_many_arguments)] async fn log_usage( state: &ProxyState, @@ -95,6 +463,8 @@ async fn log_usage( is_streaming: bool, status_code: u16, ) { + use super::usage::logger::UsageLogger; + let logger = UsageLogger::new(&state.db); // 获取 provider 的 cost_multiplier @@ -132,1121 +502,3 @@ async fn log_usage( log::warn!("记录使用量失败: {e}"); } } - -type UsageCallbackWithTiming = Arc, Option) + Send + Sync + 'static>; - -#[derive(Clone)] -struct SseUsageCollector { - inner: Arc, -} - -struct SseUsageCollectorInner { - events: Mutex>, - first_event_time: Mutex>, - start_time: std::time::Instant, - on_complete: UsageCallbackWithTiming, - finished: AtomicBool, -} - -impl SseUsageCollector { - fn new( - start_time: std::time::Instant, - callback: impl Fn(Vec, Option) + Send + Sync + 'static, - ) -> Self { - let on_complete: UsageCallbackWithTiming = Arc::new(callback); - Self { - inner: Arc::new(SseUsageCollectorInner { - events: Mutex::new(Vec::new()), - first_event_time: Mutex::new(None), - start_time, - on_complete, - finished: AtomicBool::new(false), - }), - } - } - - async fn push(&self, event: Value) { - // 记录首个事件时间 - { - let mut first_time = self.inner.first_event_time.lock().await; - if first_time.is_none() { - *first_time = Some(std::time::Instant::now()); - } - } - let mut events = self.inner.events.lock().await; - events.push(event); - } - - async fn finish(&self) { - if self.inner.finished.swap(true, Ordering::SeqCst) { - return; - } - - let events = { - let mut guard = self.inner.events.lock().await; - std::mem::take(&mut *guard) - }; - - let first_token_ms = { - let first_time = self.inner.first_event_time.lock().await; - first_time.map(|t| (t - self.inner.start_time).as_millis() as u64) - }; - - (self.inner.on_complete)(events, first_token_ms); - } -} - -/// 创建带日志记录的透传流 -fn create_logged_passthrough_stream( - stream: impl Stream> + Send + 'static, - tag: &'static str, - usage_collector: Option, -) -> impl Stream> + Send { - async_stream::stream! { - let mut buffer = String::new(); - let mut collector = usage_collector; - - 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); - - // 尝试解析并记录完整的 SSE 事件 - while let Some(pos) = buffer.find("\n\n") { - let event_text = buffer[..pos].to_string(); - buffer = buffer[pos + 2..].to_string(); - - if !event_text.trim().is_empty() { - // 提取 data 部分并尝试解析为 JSON - for line in event_text.lines() { - if let Some(data) = line.strip_prefix("data: ") { - if data.trim() != "[DONE]" { - if let Ok(json_value) = serde_json::from_str::(data) { - if let Some(c) = &collector { - c.push(json_value.clone()).await; - } - log::info!( - "[{}] <<< SSE 事件:\n{}", - tag, - serde_json::to_string_pretty(&json_value).unwrap_or_else(|_| data.to_string()) - ); - } else { - log::info!("[{tag}] <<< SSE 数据: {data}"); - } - } else { - log::info!("[{tag}] <<< SSE: [DONE]"); - } - } - } - } - } - - yield Ok(bytes); - } - Err(e) => { - log::error!("[{tag}] 流错误: {e}"); - yield Err(std::io::Error::other(e.to_string())); - break; - } - } - } - - log::info!("[{}] ====== 流结束 ======", tag); - - if let Some(c) = collector.take() { - c.finish().await; - } - } -} - -/// 健康检查 -pub async fn health_check() -> (StatusCode, Json) { - ( - StatusCode::OK, - Json(json!({ - "status": "healthy", - "timestamp": chrono::Utc::now().to_rfc3339(), - })), - ) -} - -/// 获取服务状态 -pub async fn get_status(State(state): State) -> Result, ProxyError> { - let status = state.status.read().await.clone(); - Ok(Json(status)) -} - -/// 处理 /v1/messages 请求(Claude API) -pub async fn handle_messages( - State(state): State, - headers: axum::http::HeaderMap, - Json(body): Json, -) -> Result { - let start_time = std::time::Instant::now(); - let session_id = uuid::Uuid::new_v4().to_string(); - - let config = state.config.read().await.clone(); - let request_model = body - .get("model") - .and_then(|m| m.as_str()) - .unwrap_or("unknown") - .to_string(); - - // 选择目标 Provider - let router = super::router::ProviderRouter::new(state.db.clone()); - let failed_ids = Vec::new(); - let provider = router - .select_provider(&AppType::Claude, &failed_ids) - .await?; - - // 检查是否需要转换(OpenRouter) - let adapter = get_adapter(&AppType::Claude); - let needs_transform = adapter.needs_transform(&provider); - - // 检查是否是流式请求 - let is_stream = body - .get("stream") - .and_then(|s| s.as_bool()) - .unwrap_or(false); - - log::info!( - "[Claude] Provider: {}, needs_transform: {}, is_stream: {}", - provider.name, - needs_transform, - is_stream - ); - - let forwarder = RequestForwarder::new( - state.db.clone(), - config.request_timeout, - config.max_retries, - state.status.clone(), - state.current_providers.clone(), - ); - - // 捕获错误并记录 - let response = match forwarder - .forward_with_retry(&AppType::Claude, "/v1/messages", body, headers) - .await - { - Ok(resp) => resp, - Err(e) => { - // 记录错误请求 - let logger = UsageLogger::new(&state.db); - let status_code = map_proxy_error_to_status(&e); - let error_message = get_error_message(&e); - - log::error!("[Claude] 请求失败: status={status_code}, error={error_message}"); - - let _ = logger.log_error_with_context( - session_id.clone(), - provider.id.clone(), - "claude".to_string(), - request_model.clone(), - status_code, - error_message, - start_time.elapsed().as_millis() as u64, - is_stream, - Some(session_id), - None, // provider_type 暂时设置为 None - ); - - return Err(e); - } - }; - - let status = response.status(); - log::info!("[Claude] 上游响应状态: {status}"); - - // 如果需要转换 - if needs_transform { - if is_stream { - // 流式响应转换 - log::info!("[Claude] 开始流式响应转换 (OpenAI SSE → Anthropic SSE)"); - - let stream = response.bytes_stream(); - let sse_stream = super::providers::streaming::create_anthropic_sse_stream(stream); - - let usage_collector = { - let state = state.clone(); - let provider_id = provider.id.clone(); - let model = request_model.clone(); - let status_code = status.as_u16(); - let start_time_clone = start_time; - SseUsageCollector::new(start_time, move |events, first_token_ms| { - if let Some(usage) = TokenUsage::from_claude_stream_events(&events) { - let latency_ms = start_time_clone.elapsed().as_millis() as u64; - let state = state.clone(); - let provider_id = provider_id.clone(); - let model = model.clone(); - tokio::spawn(async move { - log_usage( - &state, - &provider_id, - "claude", - &model, - usage, - latency_ms, - first_token_ms, - true, // is_streaming - status_code, - ) - .await; - }); - } else { - log::debug!("[Claude] OpenRouter 流式响应缺少 usage 统计,跳过消费记录"); - } - }) - }; - - let logged_stream = create_logged_passthrough_stream( - sse_stream, - "Claude/OpenRouter", - Some(usage_collector), - ); - - 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 = axum::body::Body::from_stream(logged_stream); - log::info!("[Claude] ====== 请求结束 (流式转换) ======"); - return Ok((headers, body).into_response()); - } else { - // 非流式响应转换 - log::info!("[Claude] 开始转换响应 (OpenAI → Anthropic)"); - - let response_headers = response.headers().clone(); - - // 读取响应体 - let body_bytes = response.bytes().await.map_err(|e| { - log::error!("[Claude] 读取响应体失败: {e}"); - ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) - })?; - - let body_str = String::from_utf8_lossy(&body_bytes); - log::info!("[Claude] OpenAI 响应长度: {} bytes", body_bytes.len()); - log::debug!("[Claude] OpenAI 原始响应: {body_str}"); - - // 解析并转换 - 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}")) - })?; - - log::info!("[Claude] 解析 OpenAI 响应成功"); - log::info!( - "[Claude] <<< OpenAI 响应 JSON:\n{}", - serde_json::to_string_pretty(&openai_response).unwrap_or_default() - ); - - let anthropic_response = - transform::openai_to_anthropic(openai_response).map_err(|e| { - log::error!("[Claude] 转换响应失败: {e}"); - e - })?; - - log::info!("[Claude] 转换响应成功"); - log::info!( - "[Claude] <<< Anthropic 响应 JSON:\n{}", - serde_json::to_string_pretty(&anthropic_response).unwrap_or_default() - ); - - // 记录使用量 - if let Some(usage) = TokenUsage::from_claude_response(&anthropic_response) { - let model = anthropic_response - .get("model") - .and_then(|m| m.as_str()) - .unwrap_or("unknown"); - let latency_ms = start_time.elapsed().as_millis() as u64; - - tokio::spawn({ - let state = state.clone(); - let provider_id = provider.id.clone(); - let model = model.to_string(); - async move { - log_usage( - &state, - &provider_id, - "claude", - &model, - usage, - latency_ms, - None, - false, - status.as_u16(), - ) - .await; - } - }); - } - - log::info!("[Claude] ====== 请求结束 ======"); - - // 构建响应 - 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()); - } - } - - // 透传响应(直连 Anthropic) - log::info!("[Claude] 透传响应模式"); - - // 检查是否流式响应 - let content_type = response - .headers() - .get("content-type") - .and_then(|v| v.to_str().ok()) - .unwrap_or(""); - let is_sse = content_type.contains("text/event-stream"); - - if is_sse { - // 流式透传:使用包装流记录 SSE 事件 - log::info!("[Claude] 流式透传响应 (SSE)"); - let mut builder = axum::response::Response::builder().status(status); - - for (key, value) in response.headers() { - builder = builder.header(key, value); - } - - let stream = response - .bytes_stream() - .map(|chunk| chunk.map_err(|e| std::io::Error::other(e.to_string()))); - let usage_collector = { - let state = state.clone(); - let provider_id = provider.id.clone(); - let model = request_model.clone(); - let status_code = status.as_u16(); - let start_time_clone = start_time; - SseUsageCollector::new(start_time, move |events, first_token_ms| { - if let Some(usage) = TokenUsage::from_claude_stream_events(&events) { - let latency_ms = start_time_clone.elapsed().as_millis() as u64; - let state = state.clone(); - let provider_id = provider_id.clone(); - let model = model.clone(); - tokio::spawn(async move { - log_usage( - &state, - &provider_id, - "claude", - &model, - usage, - latency_ms, - first_token_ms, - true, - status_code, - ) - .await; - }); - } else { - log::debug!("[Claude] 流式响应缺少 usage 统计,跳过消费记录"); - } - }) - }; - let logged_stream = - create_logged_passthrough_stream(stream, "Claude", Some(usage_collector)); - - let body = axum::body::Body::from_stream(logged_stream); - log::info!("[Claude] ====== 请求结束 (流式) ======"); - Ok(builder.body(body).unwrap()) - } else { - // 非流式透传:读取完整响应并记录 - let response_headers = response.headers().clone(); - let status = response.status(); - - let body_bytes = response.bytes().await.map_err(|e| { - log::error!("[Claude] 读取透传响应失败: {e}"); - ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) - })?; - - // 记录响应 JSON - if let Ok(json_value) = serde_json::from_slice::(&body_bytes) { - log::info!( - "[Claude] <<< Anthropic 透传响应 JSON:\n{}", - serde_json::to_string_pretty(&json_value).unwrap_or_default() - ); - - // 记录使用量 - if let Some(usage) = TokenUsage::from_claude_response(&json_value) { - let model = json_value - .get("model") - .and_then(|m| m.as_str()) - .unwrap_or("unknown"); - let latency_ms = start_time.elapsed().as_millis() as u64; - - tokio::spawn({ - let state = state.clone(); - let provider_id = provider.id.clone(); - let model = model.to_string(); - async move { - log_usage( - &state, - &provider_id, - "claude", - &model, - usage, - latency_ms, - None, - false, - status.as_u16(), - ) - .await; - } - }); - } - } else { - log::info!( - "[Claude] <<< 透传响应 (非 JSON): {} bytes", - body_bytes.len() - ); - } - log::info!("[Claude] ====== 请求结束 ======"); - - let mut builder = axum::response::Response::builder().status(status); - for (key, value) in response_headers.iter() { - builder = builder.header(key, value); - } - - let body = axum::body::Body::from(body_bytes); - Ok(builder.body(body).unwrap()) - } -} - -/// 处理 Gemini API 请求(透传,包括查询参数) -pub async fn handle_gemini( - State(state): State, - uri: axum::http::Uri, - headers: axum::http::HeaderMap, - Json(body): Json, -) -> Result { - let start_time = std::time::Instant::now(); - let session_id = uuid::Uuid::new_v4().to_string(); - - let config = state.config.read().await.clone(); - - // 选择目标 Provider - let router = super::router::ProviderRouter::new(state.db.clone()); - let failed_ids = Vec::new(); - let provider = router - .select_provider(&AppType::Gemini, &failed_ids) - .await?; - - let forwarder = RequestForwarder::new( - state.db.clone(), - config.request_timeout, - config.max_retries, - state.status.clone(), - state.current_providers.clone(), - ); - - // 提取完整的路径和查询参数 - let endpoint = uri - .path_and_query() - .map(|pq| pq.as_str()) - .unwrap_or(uri.path()); - let gemini_model = endpoint - .split('/') - .find(|s| s.starts_with("models/")) - .and_then(|s| s.strip_prefix("models/")) - .map(|s| s.split(':').next().unwrap_or(s)) - .unwrap_or("unknown") - .to_string(); - - // 检查是否是流式请求(从endpoint判断) - let is_stream = endpoint.contains("streamGenerateContent"); - - log::info!("[Gemini] 请求端点: {endpoint}"); - - // 捕获错误并记录 - let response = match forwarder - .forward_with_retry(&AppType::Gemini, endpoint, body, headers) - .await - { - Ok(resp) => resp, - Err(e) => { - // 记录错误请求 - let logger = UsageLogger::new(&state.db); - let status_code = map_proxy_error_to_status(&e); - let error_message = get_error_message(&e); - - log::error!("[Gemini] 请求失败: status={status_code}, error={error_message}"); - - let _ = logger.log_error_with_context( - session_id.clone(), - provider.id.clone(), - "gemini".to_string(), - gemini_model.clone(), - status_code, - error_message, - start_time.elapsed().as_millis() as u64, - is_stream, - Some(session_id), - None, // provider_type 暂时设置为 None - ); - - return Err(e); - } - }; - - let status = response.status(); - log::info!("[Gemini] 上游响应状态: {status}"); - - // 检查是否流式响应 - let content_type = response - .headers() - .get("content-type") - .and_then(|v| v.to_str().ok()) - .unwrap_or(""); - let is_sse = content_type.contains("text/event-stream"); - - if is_sse { - // 流式透传 - log::info!("[Gemini] 流式透传响应 (SSE)"); - let mut builder = axum::response::Response::builder().status(status); - - for (key, value) in response.headers() { - builder = builder.header(key, value); - } - - let stream = response - .bytes_stream() - .map(|chunk| chunk.map_err(|e| std::io::Error::other(e.to_string()))); - let usage_collector = { - let state = state.clone(); - let provider_id = provider.id.clone(); - let fallback_model = gemini_model.clone(); - let status_code = status.as_u16(); - let start_time_clone = start_time; - SseUsageCollector::new(start_time, move |events, first_token_ms| { - if let Some(usage) = TokenUsage::from_gemini_stream_chunks(&events) { - // 优先使用响应中的实际模型名称,否则使用从 URI 提取的模型名称 - let model = usage - .model - .clone() - .unwrap_or_else(|| fallback_model.clone()); - let latency_ms = start_time_clone.elapsed().as_millis() as u64; - let state = state.clone(); - let provider_id = provider_id.clone(); - tokio::spawn(async move { - log_usage( - &state, - &provider_id, - "gemini", - &model, - usage, - latency_ms, - first_token_ms, - true, - status_code, - ) - .await; - }); - } else { - log::debug!("[Gemini] 流式响应缺少 usage 统计,跳过消费记录"); - } - }) - }; - let logged_stream = - create_logged_passthrough_stream(stream, "Gemini", Some(usage_collector)); - - let body = axum::body::Body::from_stream(logged_stream); - Ok(builder.body(body).unwrap()) - } else { - // 非流式透传 - let response_headers = response.headers().clone(); - let status = response.status(); - - let body_bytes = response.bytes().await.map_err(|e| { - log::error!("[Gemini] 读取响应失败: {e}"); - ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) - })?; - - // 记录响应 JSON - if let Ok(json_value) = serde_json::from_slice::(&body_bytes) { - log::info!( - "[Gemini] <<< 响应 JSON:\n{}", - serde_json::to_string_pretty(&json_value).unwrap_or_default() - ); - - // 记录使用量 - if let Some(usage) = TokenUsage::from_gemini_response(&json_value) { - // 优先使用响应中的实际模型名称,否则使用从 URI 提取的模型名称 - let model = usage.model.clone().unwrap_or_else(|| gemini_model.clone()); - let latency_ms = start_time.elapsed().as_millis() as u64; - tokio::spawn({ - let state = state.clone(); - let provider_id = provider.id.clone(); - async move { - log_usage( - &state, - &provider_id, - "gemini", - &model, - usage, - latency_ms, - None, - false, - status.as_u16(), - ) - .await; - } - }); - } - } else { - log::info!("[Gemini] <<< 响应 (非 JSON): {} bytes", body_bytes.len()); - } - log::info!("[Gemini] ====== 请求结束 ======"); - - let mut builder = axum::response::Response::builder().status(status); - for (key, value) in response_headers.iter() { - builder = builder.header(key, value); - } - - let body = axum::body::Body::from(body_bytes); - Ok(builder.body(body).unwrap()) - } -} - -/// 处理 /v1/responses 请求(OpenAI Responses API - Codex CLI 透传) -pub async fn handle_responses( - State(state): State, - headers: axum::http::HeaderMap, - Json(body): Json, -) -> Result { - let start_time = std::time::Instant::now(); - let session_id = uuid::Uuid::new_v4().to_string(); - - let config = state.config.read().await.clone(); - let request_model = body - .get("model") - .and_then(|m| m.as_str()) - .unwrap_or("unknown") - .to_string(); - - // 选择目标 Provider - let router = super::router::ProviderRouter::new(state.db.clone()); - let failed_ids = Vec::new(); - let provider = router.select_provider(&AppType::Codex, &failed_ids).await?; - - let forwarder = RequestForwarder::new( - state.db.clone(), - config.request_timeout, - config.max_retries, - state.status.clone(), - state.current_providers.clone(), - ); - - // 检查是否是流式请求 - let is_stream = body - .get("stream") - .and_then(|s| s.as_bool()) - .unwrap_or(false); - - // 捕获错误并记录 - let response = match forwarder - .forward_with_retry(&AppType::Codex, "/v1/responses", body, headers) - .await - { - Ok(resp) => resp, - Err(e) => { - // 记录错误请求 - let logger = UsageLogger::new(&state.db); - let status_code = map_proxy_error_to_status(&e); - let error_message = get_error_message(&e); - - log::error!("[Codex] 请求失败: status={status_code}, error={error_message}"); - - let _ = logger.log_error_with_context( - session_id.clone(), - provider.id.clone(), - "codex".to_string(), - request_model.clone(), - status_code, - error_message, - start_time.elapsed().as_millis() as u64, - is_stream, - Some(session_id), - None, // provider_type 暂时设置为 None - ); - - return Err(e); - } - }; - - let status = response.status(); - log::info!("[Codex] 上游响应状态: {status}"); - - // 检查是否流式响应 - let content_type = response - .headers() - .get("content-type") - .and_then(|v| v.to_str().ok()) - .unwrap_or(""); - let is_sse = content_type.contains("text/event-stream"); - - if is_sse { - // 流式透传 - log::info!("[Codex] 流式透传响应 (SSE)"); - let mut builder = axum::response::Response::builder().status(status); - - for (key, value) in response.headers() { - builder = builder.header(key, value); - } - - let stream = response - .bytes_stream() - .map(|chunk| chunk.map_err(|e| std::io::Error::other(e.to_string()))); - let usage_collector = { - let state = state.clone(); - let provider_id = provider.id.clone(); - let request_model = request_model.clone(); - let status_code = status.as_u16(); - let start_time_clone = start_time; - SseUsageCollector::new(start_time, move |events, first_token_ms| { - if let Some(usage) = TokenUsage::from_codex_stream_events(&events) { - // 尝试从事件中提取模型,回退到请求模型 - let model = events - .iter() - .find_map(|e| { - if e.get("type")?.as_str()? == "response.completed" { - e.get("response")?.get("model")?.as_str() - } else { - None - } - }) - .unwrap_or(&request_model) - .to_string(); - let latency_ms = start_time_clone.elapsed().as_millis() as u64; - - let state = state.clone(); - let provider_id = provider_id.clone(); - tokio::spawn(async move { - log_usage( - &state, - &provider_id, - "codex", - &model, - usage, - latency_ms, - first_token_ms, - true, - status_code, - ) - .await; - }); - } else { - log::debug!("[Codex] 流式响应缺少 usage 统计,跳过消费记录"); - } - }) - }; - let logged_stream = - create_logged_passthrough_stream(stream, "Codex", Some(usage_collector)); - - let body = axum::body::Body::from_stream(logged_stream); - Ok(builder.body(body).unwrap()) - } else { - // 非流式透传 - let response_headers = response.headers().clone(); - let status = response.status(); - - let body_bytes = response.bytes().await.map_err(|e| { - log::error!("[Codex] 读取响应失败: {e}"); - ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) - })?; - - // 记录响应 JSON - if let Ok(json_value) = serde_json::from_slice::(&body_bytes) { - log::info!( - "[Codex] <<< 响应 JSON:\n{}", - serde_json::to_string_pretty(&json_value).unwrap_or_default() - ); - - // 记录使用量 - if let Some(usage) = TokenUsage::from_codex_response(&json_value) { - let model = json_value - .get("model") - .and_then(|m| m.as_str()) - .unwrap_or("unknown"); - let latency_ms = start_time.elapsed().as_millis() as u64; - - log::info!( - "[Codex] 解析到 usage: input={}, output={}", - usage.input_tokens, - usage.output_tokens - ); - - tokio::spawn({ - let state = state.clone(); - let provider_id = provider.id.clone(); - let model = model.to_string(); - async move { - log_usage( - &state, - &provider_id, - "codex", - &model, - usage, - latency_ms, - None, - false, - status.as_u16(), - ) - .await; - } - }); - } else { - log::warn!("[Codex] 未能解析 usage 信息,跳过记录"); - } - } else { - log::info!("[Codex] <<< 响应 (非 JSON): {} bytes", body_bytes.len()); - } - log::info!("[Codex] ====== 请求结束 ======"); - - let mut builder = axum::response::Response::builder().status(status); - for (key, value) in response_headers.iter() { - builder = builder.header(key, value); - } - - let body = axum::body::Body::from(body_bytes); - Ok(builder.body(body).unwrap()) - } -} - -/// 处理 /v1/chat/completions 请求(OpenAI Chat Completions API - Codex CLI) -pub async fn handle_chat_completions( - State(state): State, - headers: axum::http::HeaderMap, - Json(body): Json, -) -> Result { - let start_time = std::time::Instant::now(); - let session_id = uuid::Uuid::new_v4().to_string(); - - log::info!("[Codex] ====== /v1/chat/completions 请求开始 ======"); - - let config = state.config.read().await.clone(); - let request_model = body - .get("model") - .and_then(|m| m.as_str()) - .unwrap_or("unknown") - .to_string(); - let is_stream = body - .get("stream") - .and_then(|v| v.as_bool()) - .unwrap_or(false); - - log::info!("[Codex] 请求模型: {request_model}, 流式: {is_stream}"); - - // 选择目标 Provider - let router = super::router::ProviderRouter::new(state.db.clone()); - let failed_ids = Vec::new(); - let provider = router.select_provider(&AppType::Codex, &failed_ids).await?; - - log::info!("[Codex] 选择 Provider: {}", provider.id); - - let forwarder = RequestForwarder::new( - state.db.clone(), - config.request_timeout, - config.max_retries, - state.status.clone(), - state.current_providers.clone(), - ); - - // 捕获错误并记录 - let response = match forwarder - .forward_with_retry(&AppType::Codex, "/v1/chat/completions", body, headers) - .await - { - Ok(resp) => resp, - Err(e) => { - // 记录错误请求 - let logger = UsageLogger::new(&state.db); - let status_code = map_proxy_error_to_status(&e); - let error_message = get_error_message(&e); - - log::error!("[Codex] 请求失败: status={status_code}, error={error_message}"); - - let _ = logger.log_error_with_context( - session_id.clone(), - provider.id.clone(), - "codex".to_string(), - request_model.clone(), - status_code, - error_message, - start_time.elapsed().as_millis() as u64, - is_stream, - Some(session_id), - None, // provider_type 暂时设置为 None - ); - - return Err(e); - } - }; - - let status = response.status(); - log::info!("[Codex] 上游响应状态: {status}"); - - // 检查是否流式响应 - let content_type = response - .headers() - .get("content-type") - .and_then(|v| v.to_str().ok()) - .unwrap_or(""); - let is_sse = content_type.contains("text/event-stream"); - - if is_sse { - // 流式透传 - log::info!("[Codex] 流式透传响应 (SSE)"); - let mut builder = axum::response::Response::builder().status(status); - - for (key, value) in response.headers() { - builder = builder.header(key, value); - } - - let stream = response - .bytes_stream() - .map(|chunk| chunk.map_err(|e| std::io::Error::other(e.to_string()))); - - let usage_collector = { - let state = state.clone(); - let provider_id = provider.id.clone(); - let request_model = request_model.clone(); - let status_code = status.as_u16(); - let start_time_clone = start_time; - SseUsageCollector::new(start_time, move |events, first_token_ms| { - if let Some(usage) = TokenUsage::from_openai_stream_events(&events) { - let model = events - .iter() - .find_map(|e| e.get("model")?.as_str()) - .unwrap_or(&request_model) - .to_string(); - let latency_ms = start_time_clone.elapsed().as_millis() as u64; - - let state = state.clone(); - let provider_id = provider_id.clone(); - tokio::spawn(async move { - log_usage( - &state, - &provider_id, - "codex", - &model, - usage, - latency_ms, - first_token_ms, - true, - status_code, - ) - .await; - }); - } else { - log::debug!("[Codex] 流式响应缺少 usage 统计,跳过消费记录"); - } - }) - }; - let logged_stream = - create_logged_passthrough_stream(stream, "Codex", Some(usage_collector)); - - let body = axum::body::Body::from_stream(logged_stream); - Ok(builder.body(body).unwrap()) - } else { - // 非流式透传 - let response_headers = response.headers().clone(); - let status = response.status(); - - let body_bytes = response.bytes().await.map_err(|e| { - log::error!("[Codex] 读取响应失败: {e}"); - ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) - })?; - - // 记录响应 JSON - if let Ok(json_value) = serde_json::from_slice::(&body_bytes) { - log::info!( - "[Codex] <<< 响应 JSON:\n{}", - serde_json::to_string_pretty(&json_value).unwrap_or_default() - ); - - // 记录使用量 (OpenAI 格式: prompt_tokens, completion_tokens) - if let Some(usage) = TokenUsage::from_openai_response(&json_value) { - let model = json_value - .get("model") - .and_then(|m| m.as_str()) - .unwrap_or("unknown"); - let latency_ms = start_time.elapsed().as_millis() as u64; - - log::info!( - "[Codex] 解析到 usage: input={}, output={}", - usage.input_tokens, - usage.output_tokens - ); - - tokio::spawn({ - let state = state.clone(); - let provider_id = provider.id.clone(); - let model = model.to_string(); - async move { - log_usage( - &state, - &provider_id, - "codex", - &model, - usage, - latency_ms, - None, - false, - status.as_u16(), - ) - .await; - } - }); - } else { - log::warn!("[Codex] 未能解析 usage 信息,跳过记录"); - } - } else { - log::info!("[Codex] <<< 响应 (非 JSON): {} bytes", body_bytes.len()); - } - log::info!("[Codex] ====== 请求结束 ======"); - - let mut builder = axum::response::Response::builder().status(status); - for (key, value) in response_headers.iter() { - builder = builder.header(key, value); - } - - let body = axum::body::Body::from(body_bytes); - Ok(builder.body(body).unwrap()) - } -} diff --git a/src-tauri/src/proxy/mod.rs b/src-tauri/src/proxy/mod.rs index 9a9170868..4f65b853d 100644 --- a/src-tauri/src/proxy/mod.rs +++ b/src-tauri/src/proxy/mod.rs @@ -5,13 +5,16 @@ pub mod circuit_breaker; pub mod error; pub mod error_mapper; +pub(crate) mod failover_switch; mod forwarder; +pub mod handler_config; +pub mod handler_context; mod handlers; mod health; pub mod provider_router; pub mod providers; pub mod response_handler; -mod router; +pub mod response_processor; pub(crate) mod server; pub mod session; pub(crate) mod types; diff --git a/src-tauri/src/proxy/provider_router.rs b/src-tauri/src/proxy/provider_router.rs index 931a36ab1..dbac96e1c 100644 --- a/src-tauri/src/proxy/provider_router.rs +++ b/src-tauri/src/proxy/provider_router.rs @@ -5,7 +5,7 @@ use crate::database::Database; use crate::error::AppError; use crate::provider::Provider; -use crate::proxy::circuit_breaker::CircuitBreaker; +use crate::proxy::circuit_breaker::{CircuitBreaker, CircuitBreakerConfig}; use std::collections::HashMap; use std::sync::Arc; use tokio::sync::RwLock; @@ -28,44 +28,106 @@ impl ProviderRouter { } /// 选择可用的供应商(支持故障转移) - /// 返回按优先级排序的可用供应商列表 + /// + /// 返回按优先级排序的可用供应商列表: + /// 1. 当前供应商(is_current=true)始终第一位 + /// 2. 故障转移队列中的其他供应商(按 queue_order 排序) + /// 3. 只返回熔断器未打开的供应商 pub async fn select_providers(&self, app_type: &str) -> Result, AppError> { - // 直接获取当前选中的供应商(基于 is_current 字段) - let current_id = self - .db - .get_current_provider(app_type)? - .ok_or_else(|| AppError::Config(format!("No current provider for {app_type}")))?; + let mut result = Vec::new(); + let all_providers = self.db.get_all_providers(app_type)?; - let providers = self.db.get_all_providers(app_type)?; - let provider = providers - .get(¤t_id) - .ok_or_else(|| AppError::Config(format!("Current provider {current_id} not found")))? - .clone(); + // 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; - log::info!( - "[{}] Selected current provider: {} ({})", - app_type, - provider.name, - provider.id - ); + 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 circuit_key = format!("{}:{}", app_type, provider.id); - let breaker = self.get_or_create_circuit_breaker(&circuit_key).await; + // 2. 获取故障转移队列中的供应商 + let queue = self.db.get_failover_queue(app_type)?; - if !breaker.allow_request().await { - log::warn!( - "Provider {} is unavailable (circuit breaker open)", - provider.id - ); + for item in queue { + // 跳过已添加的当前供应商 + if result.iter().any(|p| p.id == item.provider_id) { + continue; + } + + // 跳过禁用的队列项 + if !item.enabled { + continue; + } + + // 获取供应商信息 + if let Some(provider) = all_providers.get(&item.provider_id) { + // 检查熔断器状态 + let circuit_key = format!("{}:{}", app_type, provider.id); + let breaker = self.get_or_create_circuit_breaker(&circuit_key).await; + + if breaker.is_available().await { + log::info!( + "[{}] Failover provider available: {} ({}) at queue position {}", + app_type, + provider.name, + provider.id, + item.queue_order + ); + result.push(provider.clone()); + } else { + log::debug!( + "[{}] Failover provider {} circuit breaker open, skipping", + app_type, + provider.name + ); + } + } + } + + if result.is_empty() { return Err(AppError::Config(format!( - "Current provider {} is unavailable (circuit breaker open)", - provider.name + "No available provider for {} (all circuit breakers open or no providers configured)", + app_type ))); } - // 返回单个供应商(保留 Vec 接口以兼容现有代码) - Ok(vec![provider]) + log::info!( + "[{}] Failover chain: {} provider(s) available", + app_type, + result.len() + ); + + Ok(result) + } + + /// 请求执行前获取熔断器“放行许可” + /// + /// - Closed:直接放行 + /// - Open:超时到达后切到 HalfOpen 并放行一次探测 + /// - HalfOpen:按限流规则放行探测 + /// + /// 注意:调用方必须在请求结束后通过 `record_result()` 释放 HalfOpen 名额, + /// 否则会导致该 Provider 长时间无法进入探测状态。 + pub async fn allow_provider_request(&self, provider_id: &str, app_type: &str) -> bool { + let circuit_key = format!("{app_type}:{provider_id}"); + let breaker = self.get_or_create_circuit_breaker(&circuit_key).await; + breaker.allow_request().await } /// 记录供应商请求结果 @@ -76,7 +138,11 @@ impl ProviderRouter { success: bool, error_msg: Option, ) -> Result<(), AppError> { - // 1. 更新熔断器状态 + // 1. 获取熔断器配置(用于更新健康状态和判断是否禁用) + let config = self.db.get_circuit_breaker_config().await.ok(); + let failure_threshold = config.map(|c| c.failure_threshold).unwrap_or(5); + + // 2. 更新熔断器状态 let circuit_key = format!("{app_type}:{provider_id}"); let breaker = self.get_or_create_circuit_breaker(&circuit_key).await; @@ -92,38 +158,21 @@ impl ProviderRouter { ); } - // 2. 更新数据库健康状态 + // 3. 更新数据库健康状态(使用配置的阈值) self.db - .update_provider_health(provider_id, app_type, success, error_msg.clone()) + .update_provider_health_with_threshold( + provider_id, + app_type, + success, + error_msg.clone(), + failure_threshold, + ) .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) { @@ -132,6 +181,27 @@ impl ProviderRouter { } } + /// 重置指定供应商的熔断器 + pub async fn reset_provider_breaker(&self, provider_id: &str, app_type: &str) { + let circuit_key = format!("{app_type}:{provider_id}"); + self.reset_circuit_breaker(&circuit_key).await; + } + + /// 更新所有熔断器的配置(热更新) + /// + /// 当用户在 UI 中修改熔断器配置后调用此方法, + /// 所有现有的熔断器会立即使用新配置 + pub async fn update_all_configs(&self, config: CircuitBreakerConfig) { + let breakers = self.circuit_breakers.read().await; + let count = breakers.len(); + + for breaker in breakers.values() { + breaker.update_config(config.clone()).await; + } + + log::info!("已更新 {} 个熔断器的配置", count); + } + /// 获取熔断器状态 #[allow(dead_code)] pub async fn get_circuit_breaker_stats( @@ -187,6 +257,7 @@ impl ProviderRouter { mod tests { use super::*; use crate::database::Database; + use serde_json::json; #[tokio::test] async fn test_provider_router_creation() { @@ -197,4 +268,44 @@ mod tests { let breaker = router.get_or_create_circuit_breaker("claude:test").await; assert!(breaker.allow_request().await); } + + #[tokio::test] + async fn select_providers_does_not_consume_half_open_permit() { + 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 = + 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(); + + let router = ProviderRouter::new(db.clone()); + + // 让 B 进入 Open 状态(failure_threshold=1) + router + .record_result("b", "claude", 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); + } } diff --git a/src-tauri/src/proxy/providers/claude.rs b/src-tauri/src/proxy/providers/claude.rs index f54321ca5..a9bc2e567 100644 --- a/src-tauri/src/proxy/providers/claude.rs +++ b/src-tauri/src/proxy/providers/claude.rs @@ -253,7 +253,6 @@ mod tests { meta: None, icon: None, icon_color: None, - is_proxy_target: None, } } diff --git a/src-tauri/src/proxy/providers/codex.rs b/src-tauri/src/proxy/providers/codex.rs index a7599b526..22c33c75a 100644 --- a/src-tauri/src/proxy/providers/codex.rs +++ b/src-tauri/src/proxy/providers/codex.rs @@ -174,7 +174,6 @@ mod tests { meta: None, icon: None, icon_color: None, - is_proxy_target: None, } } diff --git a/src-tauri/src/proxy/providers/gemini.rs b/src-tauri/src/proxy/providers/gemini.rs index b148f5a48..1b4ef3593 100644 --- a/src-tauri/src/proxy/providers/gemini.rs +++ b/src-tauri/src/proxy/providers/gemini.rs @@ -120,11 +120,7 @@ impl GeminiAdapter { /// 从 Provider 配置中提取原始 API Key fn extract_key_raw(&self, provider: &Provider) -> Option { if let Some(env) = provider.settings_config.get("env") { - // 优先使用 GOOGLE_GEMINI_API_KEY - if let Some(key) = env.get("GOOGLE_GEMINI_API_KEY").and_then(|v| v.as_str()) { - return Some(key.to_string()); - } - // 备选 GEMINI_API_KEY + // 使用 GEMINI_API_KEY if let Some(key) = env.get("GEMINI_API_KEY").and_then(|v| v.as_str()) { return Some(key.to_string()); } @@ -254,7 +250,6 @@ mod tests { meta: None, icon: None, icon_color: None, - is_proxy_target: None, } } @@ -276,7 +271,7 @@ mod tests { let adapter = GeminiAdapter::new(); let provider = create_provider(json!({ "env": { - "GOOGLE_GEMINI_API_KEY": "AIza-test-key-12345678" + "GEMINI_API_KEY": "AIza-test-key-12345678" } })); @@ -291,7 +286,7 @@ mod tests { let adapter = GeminiAdapter::new(); let provider = create_provider(json!({ "env": { - "GOOGLE_GEMINI_API_KEY": "ya29.test-access-token-12345" + "GEMINI_API_KEY": "ya29.test-access-token-12345" } })); @@ -308,7 +303,7 @@ mod tests { let adapter = GeminiAdapter::new(); let provider = create_provider(json!({ "env": { - "GOOGLE_GEMINI_API_KEY": "{\"access_token\":\"ya29.test-token\",\"refresh_token\":\"1//refresh\"}" + "GEMINI_API_KEY": "{\"access_token\":\"ya29.test-token\",\"refresh_token\":\"1//refresh\"}" } })); @@ -324,7 +319,7 @@ mod tests { // API Key let api_key_provider = create_provider(json!({ "env": { - "GOOGLE_GEMINI_API_KEY": "AIza-test-key" + "GEMINI_API_KEY": "AIza-test-key" } })); assert_eq!( @@ -335,7 +330,7 @@ mod tests { // OAuth access_token let oauth_provider = create_provider(json!({ "env": { - "GOOGLE_GEMINI_API_KEY": "ya29.test-token" + "GEMINI_API_KEY": "ya29.test-token" } })); assert_eq!( @@ -346,7 +341,7 @@ mod tests { // OAuth JSON let oauth_json_provider = create_provider(json!({ "env": { - "GOOGLE_GEMINI_API_KEY": "{\"access_token\":\"ya29.test\"}" + "GEMINI_API_KEY": "{\"access_token\":\"ya29.test\"}" } })); assert_eq!( diff --git a/src-tauri/src/proxy/providers/mod.rs b/src-tauri/src/proxy/providers/mod.rs index f41a03b88..f604f60a0 100644 --- a/src-tauri/src/proxy/providers/mod.rs +++ b/src-tauri/src/proxy/providers/mod.rs @@ -205,7 +205,6 @@ mod tests { meta: None, icon: None, icon_color: None, - is_proxy_target: None, } } @@ -369,7 +368,7 @@ mod tests { fn test_from_app_type_gemini_api_key() { let provider = create_provider(json!({ "env": { - "GOOGLE_GEMINI_API_KEY": "AIza-test-key" + "GEMINI_API_KEY": "AIza-test-key" } })); @@ -381,7 +380,7 @@ mod tests { fn test_from_app_type_gemini_cli_oauth() { let provider = create_provider(json!({ "env": { - "GOOGLE_GEMINI_API_KEY": "ya29.test-access-token" + "GEMINI_API_KEY": "ya29.test-access-token" } })); @@ -393,7 +392,7 @@ mod tests { fn test_from_app_type_gemini_cli_json() { let provider = create_provider(json!({ "env": { - "GOOGLE_GEMINI_API_KEY": "{\"access_token\":\"ya29.test\",\"refresh_token\":\"1//test\"}" + "GEMINI_API_KEY": "{\"access_token\":\"ya29.test\",\"refresh_token\":\"1//test\"}" } })); diff --git a/src-tauri/src/proxy/providers/transform.rs b/src-tauri/src/proxy/providers/transform.rs index a841750ed..141506fd5 100644 --- a/src-tauri/src/proxy/providers/transform.rs +++ b/src-tauri/src/proxy/providers/transform.rs @@ -394,7 +394,6 @@ mod tests { meta: None, icon: None, icon_color: None, - is_proxy_target: None, } } diff --git a/src-tauri/src/proxy/response_processor.rs b/src-tauri/src/proxy/response_processor.rs new file mode 100644 index 000000000..38b59faeb --- /dev/null +++ b/src-tauri/src/proxy/response_processor.rs @@ -0,0 +1,411 @@ +//! 响应处理器模块 +//! +//! 统一处理流式和非流式 API 响应 + +use super::{ + handler_config::UsageParserConfig, handler_context::RequestContext, server::ProxyState, + usage::parser::TokenUsage, ProxyError, +}; +use axum::response::Response; +use bytes::Bytes; +use futures::stream::{Stream, StreamExt}; +use rust_decimal::Decimal; +use serde_json::Value; +use std::{ + str::FromStr, + sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }, +}; +use tokio::sync::Mutex; + +// ============================================================================ +// 公共接口 +// ============================================================================ + +/// 检测响应是否为 SSE 流式响应 +#[inline] +pub fn is_sse_response(response: &reqwest::Response) -> bool { + response + .headers() + .get("content-type") + .and_then(|v| v.to_str().ok()) + .map(|ct| ct.contains("text/event-stream")) + .unwrap_or(false) +} + +/// 处理流式响应 +pub async fn handle_streaming( + response: reqwest::Response, + ctx: &RequestContext, + state: &ProxyState, + parser_config: &UsageParserConfig, +) -> Response { + log::info!("[{}] 流式透传响应 (SSE)", ctx.tag); + + let status = response.status(); + let mut builder = axum::response::Response::builder().status(status); + + // 复制响应头 + for (key, value) in response.headers() { + builder = builder.header(key, value); + } + + // 创建字节流 + let stream = response + .bytes_stream() + .map(|chunk| chunk.map_err(|e| std::io::Error::other(e.to_string()))); + + // 创建使用量收集器 + let usage_collector = create_usage_collector(ctx, state, status.as_u16(), parser_config); + + // 创建带日志的透传流 + let logged_stream = create_logged_passthrough_stream(stream, ctx.tag, Some(usage_collector)); + + let body = axum::body::Body::from_stream(logged_stream); + builder.body(body).unwrap() +} + +/// 处理非流式响应 +pub async fn handle_non_streaming( + response: reqwest::Response, + ctx: &RequestContext, + state: &ProxyState, + parser_config: &UsageParserConfig, +) -> Result { + let response_headers = response.headers().clone(); + let status = response.status(); + + // 读取响应体 + let body_bytes = response.bytes().await.map_err(|e| { + log::error!("[{}] 读取响应失败: {e}", ctx.tag); + ProxyError::ForwardFailed(format!("Failed to read response body: {e}")) + })?; + + // 解析并记录使用量 + if let Ok(json_value) = serde_json::from_slice::(&body_bytes) { + log::info!( + "[{}] <<< 响应 JSON:\n{}", + ctx.tag, + serde_json::to_string_pretty(&json_value).unwrap_or_default() + ); + + // 解析使用量 + if let Some(usage) = (parser_config.response_parser)(&json_value) { + let model = json_value + .get("model") + .and_then(|m| m.as_str()) + .unwrap_or(&ctx.request_model); + + spawn_log_usage(state, ctx, usage, model, status.as_u16(), false); + } else { + log::debug!( + "[{}] 未能解析 usage 信息,跳过记录", + parser_config.app_type_str + ); + } + } else { + log::info!( + "[{}] <<< 响应 (非 JSON): {} bytes", + ctx.tag, + body_bytes.len() + ); + } + + log::info!("[{}] ====== 请求结束 ======", ctx.tag); + + // 构建响应 + let mut builder = axum::response::Response::builder().status(status); + for (key, value) in response_headers.iter() { + builder = builder.header(key, value); + } + + let body = axum::body::Body::from(body_bytes); + Ok(builder.body(body).unwrap()) +} + +/// 通用响应处理入口 +/// +/// 根据响应类型自动选择流式或非流式处理 +pub async fn process_response( + response: reqwest::Response, + ctx: &RequestContext, + state: &ProxyState, + parser_config: &UsageParserConfig, +) -> Result { + if is_sse_response(&response) { + Ok(handle_streaming(response, ctx, state, parser_config).await) + } else { + handle_non_streaming(response, ctx, state, parser_config).await + } +} + +// ============================================================================ +// SSE 使用量收集器 +// ============================================================================ + +type UsageCallbackWithTiming = Arc, Option) + Send + Sync + 'static>; + +/// SSE 使用量收集器 +#[derive(Clone)] +pub struct SseUsageCollector { + inner: Arc, +} + +struct SseUsageCollectorInner { + events: Mutex>, + first_event_time: Mutex>, + start_time: std::time::Instant, + on_complete: UsageCallbackWithTiming, + finished: AtomicBool, +} + +impl SseUsageCollector { + /// 创建新的使用量收集器 + pub fn new( + start_time: std::time::Instant, + callback: impl Fn(Vec, Option) + Send + Sync + 'static, + ) -> Self { + let on_complete: UsageCallbackWithTiming = Arc::new(callback); + Self { + inner: Arc::new(SseUsageCollectorInner { + events: Mutex::new(Vec::new()), + first_event_time: Mutex::new(None), + start_time, + on_complete, + finished: AtomicBool::new(false), + }), + } + } + + /// 推送 SSE 事件 + pub async fn push(&self, event: Value) { + // 记录首个事件时间 + { + let mut first_time = self.inner.first_event_time.lock().await; + if first_time.is_none() { + *first_time = Some(std::time::Instant::now()); + } + } + let mut events = self.inner.events.lock().await; + events.push(event); + } + + /// 完成收集并触发回调 + pub async fn finish(&self) { + if self.inner.finished.swap(true, Ordering::SeqCst) { + return; + } + + let events = { + let mut guard = self.inner.events.lock().await; + std::mem::take(&mut *guard) + }; + + let first_token_ms = { + let first_time = self.inner.first_event_time.lock().await; + first_time.map(|t| (t - self.inner.start_time).as_millis() as u64) + }; + + (self.inner.on_complete)(events, first_token_ms); + } +} + +// ============================================================================ +// 内部辅助函数 +// ============================================================================ + +/// 创建使用量收集器 +fn create_usage_collector( + ctx: &RequestContext, + state: &ProxyState, + status_code: u16, + parser_config: &UsageParserConfig, +) -> SseUsageCollector { + let state = state.clone(); + let provider_id = ctx.provider.id.clone(); + let request_model = ctx.request_model.clone(); + let app_type_str = parser_config.app_type_str; + let tag = ctx.tag; + let start_time = ctx.start_time; + let stream_parser = parser_config.stream_parser; + let model_extractor = parser_config.model_extractor; + + SseUsageCollector::new(start_time, move |events, first_token_ms| { + if let Some(usage) = stream_parser(&events) { + let model = model_extractor(&events, &request_model); + let latency_ms = start_time.elapsed().as_millis() as u64; + + let state = state.clone(); + let provider_id = provider_id.clone(); + + tokio::spawn(async move { + log_usage_internal( + &state, + &provider_id, + app_type_str, + &model, + usage, + latency_ms, + first_token_ms, + true, // is_streaming + status_code, + ) + .await; + }); + } else { + log::debug!("[{tag}] 流式响应缺少 usage 统计,跳过消费记录"); + } + }) +} + +/// 异步记录使用量 +fn spawn_log_usage( + state: &ProxyState, + ctx: &RequestContext, + usage: TokenUsage, + model: &str, + status_code: u16, + is_streaming: bool, +) { + let state = state.clone(); + let provider_id = ctx.provider.id.clone(); + let app_type_str = ctx.app_type_str.to_string(); + let model = model.to_string(); + let latency_ms = ctx.latency_ms(); + + tokio::spawn(async move { + log_usage_internal( + &state, + &provider_id, + &app_type_str, + &model, + usage, + latency_ms, + None, + is_streaming, + status_code, + ) + .await; + }); +} + +/// 内部使用量记录函数 +#[allow(clippy::too_many_arguments)] +async fn log_usage_internal( + state: &ProxyState, + provider_id: &str, + app_type: &str, + model: &str, + usage: TokenUsage, + latency_ms: u64, + first_token_ms: Option, + is_streaming: bool, + status_code: u16, +) { + use super::usage::logger::UsageLogger; + + let logger = UsageLogger::new(&state.db); + + // 获取 provider 的 cost_multiplier + let multiplier = match state.db.get_provider_by_id(provider_id, app_type) { + Ok(Some(p)) => { + if let Some(meta) = p.meta { + if let Some(cm) = meta.cost_multiplier { + Decimal::from_str(&cm).unwrap_or(Decimal::from(1)) + } else { + Decimal::from(1) + } + } else { + Decimal::from(1) + } + } + _ => Decimal::from(1), + }; + + let request_id = uuid::Uuid::new_v4().to_string(); + + if let Err(e) = logger.log_with_calculation( + request_id, + provider_id.to_string(), + app_type.to_string(), + model.to_string(), + usage, + multiplier, + latency_ms, + first_token_ms, + status_code, + None, + None, // provider_type + is_streaming, + ) { + log::warn!("记录使用量失败: {e}"); + } +} + +/// 创建带日志记录的透传流 +pub fn create_logged_passthrough_stream( + stream: impl Stream> + Send + 'static, + tag: &'static str, + usage_collector: Option, +) -> impl Stream> + Send { + async_stream::stream! { + let mut buffer = String::new(); + let mut collector = usage_collector; + + 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); + + // 尝试解析并记录完整的 SSE 事件 + while let Some(pos) = buffer.find("\n\n") { + let event_text = buffer[..pos].to_string(); + buffer = buffer[pos + 2..].to_string(); + + if !event_text.trim().is_empty() { + // 提取 data 部分并尝试解析为 JSON + for line in event_text.lines() { + if let Some(data) = line.strip_prefix("data: ") { + if data.trim() != "[DONE]" { + if let Ok(json_value) = serde_json::from_str::(data) { + if let Some(c) = &collector { + c.push(json_value.clone()).await; + } + log::info!( + "[{}] <<< SSE 事件:\n{}", + tag, + serde_json::to_string_pretty(&json_value).unwrap_or_else(|_| data.to_string()) + ); + } else { + log::info!("[{tag}] <<< SSE 数据: {data}"); + } + } else { + log::info!("[{tag}] <<< SSE: [DONE]"); + } + } + } + } + } + + yield Ok(bytes); + } + Err(e) => { + log::error!("[{tag}] 流错误: {e}"); + yield Err(std::io::Error::other(e.to_string())); + break; + } + } + } + + log::info!("[{}] ====== 流结束 ======", tag); + + if let Some(c) = collector.take() { + c.finish().await; + } + } +} diff --git a/src-tauri/src/proxy/router.rs b/src-tauri/src/proxy/router.rs deleted file mode 100644 index c59c99bf8..000000000 --- a/src-tauri/src/proxy/router.rs +++ /dev/null @@ -1,70 +0,0 @@ -//! Provider路由器 -//! -//! 负责选择合适的Provider进行请求转发 - -use super::ProxyError; -use crate::{app_config::AppType, database::Database, provider::Provider}; -use std::sync::Arc; - -pub struct ProviderRouter { - db: Arc, -} - -impl ProviderRouter { - pub fn new(db: Arc) -> Self { - Self { db } - } - - /// 选择Provider(只使用标记为代理目标的 Provider) - pub async fn select_provider( - &self, - app_type: &AppType, - _failed_ids: &[String], - ) -> Result { - // 1. 获取 Proxy Target Provider ID - let proxy_target_id = self - .db - .get_proxy_target_provider(app_type.as_str()) - .map_err(|e| ProxyError::DatabaseError(e.to_string()))?; - - let target_id = proxy_target_id.ok_or_else(|| { - log::warn!("[{}] 未设置代理目标 Provider", app_type.as_str()); - ProxyError::NoAvailableProvider - })?; - - // 2. 获取所有 Provider - let providers = self - .db - .get_all_providers(app_type.as_str()) - .map_err(|e| ProxyError::DatabaseError(e.to_string()))?; - - // 3. 找到目标 Provider - let target = providers.get(&target_id).ok_or_else(|| { - log::warn!( - "[{}] 代理目标 Provider 不存在: {}", - app_type.as_str(), - target_id - ); - ProxyError::NoAvailableProvider - })?; - - log::info!( - "[{}] 使用代理目标 Provider: {}", - app_type.as_str(), - target.name - ); - Ok(target.clone()) - } - - /// 更新Provider健康状态(保留接口但不影响选择) - #[allow(dead_code)] - pub async fn update_health( - &self, - _provider: &Provider, - _app_type: &AppType, - _success: bool, - _error_msg: Option, - ) { - // 不再记录健康状态 - } -} diff --git a/src-tauri/src/proxy/server.rs b/src-tauri/src/proxy/server.rs index b66b0026b..11bfd8ed1 100644 --- a/src-tauri/src/proxy/server.rs +++ b/src-tauri/src/proxy/server.rs @@ -2,7 +2,10 @@ //! //! 基于Axum的HTTP服务器,处理代理请求 -use super::{handlers, types::*, ProxyError}; +use super::{ + failover_switch::FailoverSwitchManager, handlers, provider_router::ProviderRouter, types::*, + ProxyError, +}; use crate::database::Database; use axum::{ routing::{get, post}, @@ -11,6 +14,7 @@ use axum::{ use std::net::SocketAddr; use std::sync::Arc; use tokio::sync::{oneshot, RwLock}; +use tokio::task::JoinHandle; use tower_http::cors::{Any, CorsLayer}; /// 代理服务器状态(共享) @@ -22,6 +26,12 @@ pub struct ProxyState { pub start_time: Arc>>, /// 每个应用类型当前使用的 provider (app_type -> (provider_id, provider_name)) pub current_providers: Arc>>, + /// 共享的 ProviderRouter(持有熔断器状态,跨请求保持) + pub provider_router: Arc, + /// AppHandle,用于发射事件和更新托盘菜单 + pub app_handle: Option, + /// 故障转移切换管理器 + pub failover_manager: Arc, } /// 代理HTTP服务器 @@ -29,22 +39,37 @@ pub struct ProxyServer { config: ProxyConfig, state: ProxyState, shutdown_tx: Arc>>>, + /// 服务器任务句柄,用于等待服务器实际关闭 + server_handle: Arc>>>, } impl ProxyServer { - pub fn new(config: ProxyConfig, db: Arc) -> Self { + pub fn new( + config: ProxyConfig, + db: Arc, + app_handle: Option, + ) -> Self { + // 创建共享的 ProviderRouter(熔断器状态将跨所有请求保持) + let provider_router = Arc::new(ProviderRouter::new(db.clone())); + // 创建故障转移切换管理器 + let failover_manager = Arc::new(FailoverSwitchManager::new(db.clone())); + let state = ProxyState { db, 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())), + provider_router, + app_handle, + failover_manager, }; Self { config, state, shutdown_tx: Arc::new(RwLock::new(None)), + server_handle: Arc::new(RwLock::new(None)), } } @@ -87,7 +112,7 @@ impl ProxyServer { // 启动服务器 let state = self.state.clone(); - tokio::spawn(async move { + let handle = tokio::spawn(async move { axum::serve(listener, app) .with_graceful_shutdown(async { shutdown_rx.await.ok(); @@ -100,6 +125,9 @@ impl ProxyServer { *state.start_time.write().await = None; }); + // 保存服务器任务句柄 + *self.server_handle.write().await = Some(handle); + Ok(ProxyServerInfo { address: self.config.listen_address.clone(), port: self.config.listen_port, @@ -108,12 +136,23 @@ impl ProxyServer { } pub async fn stop(&self) -> Result<(), ProxyError> { + // 1. 发送关闭信号 if let Some(tx) = self.shutdown_tx.write().await.take() { let _ = tx.send(()); - Ok(()) } else { - Err(ProxyError::NotRunning) + return Err(ProxyError::NotRunning); } + + // 2. 等待服务器任务结束(带 5 秒超时保护) + if let Some(handle) = self.server_handle.write().await.take() { + match tokio::time::timeout(std::time::Duration::from_secs(5), handle).await { + Ok(Ok(())) => log::info!("代理服务器已完全停止"), + Ok(Err(e)) => log::warn!("代理服务器任务异常终止: {e}"), + Err(_) => log::warn!("代理服务器停止超时(5秒),强制继续"), + } + } + + Ok(()) } pub async fn get_status(&self) -> ProxyStatus { @@ -174,4 +213,22 @@ impl ProxyServer { pub async fn apply_runtime_config(&self, config: &ProxyConfig) { *self.state.config.write().await = config.clone(); } + + /// 热更新熔断器配置 + /// + /// 将新配置应用到所有已创建的熔断器实例 + pub async fn update_circuit_breaker_configs( + &self, + config: super::circuit_breaker::CircuitBreakerConfig, + ) { + self.state.provider_router.update_all_configs(config).await; + } + + /// 重置指定 Provider 的熔断器 + pub async fn reset_provider_circuit_breaker(&self, provider_id: &str, app_type: &str) { + self.state + .provider_router + .reset_provider_breaker(provider_id, app_type) + .await; + } } diff --git a/src-tauri/src/services/provider/mod.rs b/src-tauri/src/services/provider/mod.rs index bc6be2206..3446b356b 100644 --- a/src-tauri/src/services/provider/mod.rs +++ b/src-tauri/src/services/provider/mod.rs @@ -206,17 +206,27 @@ impl ProviderService { id ); + // 获取新供应商的完整配置(用于更新备份) + let provider = providers + .get(id) + .ok_or_else(|| AppError::Message(format!("供应商 {id} 不存在")))?; + // Update database is_current state.db.set_current_provider(app_type.as_str(), id)?; - // 同时更新 is_proxy_target(代理路由器使用此字段选择供应商) - state.db.set_proxy_target_provider(app_type.as_str(), id)?; - // Update local settings for consistency crate::settings::set_current_provider(&app_type, Some(id))?; + // 更新 Live 备份(确保代理关闭时恢复正确的供应商配置) + futures::executor::block_on( + state + .proxy_service + .update_live_backup_from_provider(app_type.as_str(), provider), + ) + .map_err(|e| AppError::Message(format!("更新 Live 备份失败: {e}")))?; + // Note: No Live config write, no MCP sync - // The proxy server will route requests to the new provider via is_proxy_target + // The proxy server will route requests to the new provider via is_current return Ok(()); } @@ -274,18 +284,6 @@ impl ProviderService { Ok(()) } - /// Set proxy target provider - pub fn set_proxy_target(state: &AppState, app_type: AppType, id: &str) -> Result<(), AppError> { - // Check if provider exists - let providers = state.db.get_all_providers(app_type.as_str())?; - if !providers.contains_key(id) { - return Err(AppError::Message(format!("供应商 {id} 不存在"))); - } - - state.db.set_proxy_target_provider(app_type.as_str(), id)?; - Ok(()) - } - /// Sync current provider to live configuration (re-export) pub fn sync_current_to_live(state: &AppState) -> Result<(), AppError> { sync_current_to_live(state) diff --git a/src-tauri/src/services/provider/usage.rs b/src-tauri/src/services/provider/usage.rs index ee95c00c7..6ff94ca93 100644 --- a/src-tauri/src/services/provider/usage.rs +++ b/src-tauri/src/services/provider/usage.rs @@ -79,6 +79,34 @@ pub(crate) async fn execute_and_format_usage_result( } } +/// Extract API key from provider configuration +fn extract_api_key_from_provider(provider: &crate::provider::Provider) -> Option { + if let Some(env) = provider.settings_config.get("env") { + // Try multiple possible API key fields + env.get("ANTHROPIC_AUTH_TOKEN") + .or_else(|| env.get("ANTHROPIC_API_KEY")) + .or_else(|| env.get("OPENROUTER_API_KEY")) + .or_else(|| env.get("GOOGLE_API_KEY")) + .and_then(|v| v.as_str()) + .map(|s| s.to_string()) + } else { + None + } +} + +/// Extract base URL from provider configuration +fn extract_base_url_from_provider(provider: &crate::provider::Provider) -> Option { + if let Some(env) = provider.settings_config.get("env") { + // Try multiple possible base URL fields + env.get("ANTHROPIC_BASE_URL") + .or_else(|| env.get("GOOGLE_GEMINI_BASE_URL")) + .and_then(|v| v.as_str()) + .map(|s| s.trim_end_matches('/').to_string()) + } else { + None + } +} + /// Query provider usage (using saved script configuration) pub async fn query_usage( state: &AppState, @@ -114,12 +142,26 @@ pub async fn query_usage( )); } - // Get credentials directly from UsageScript, no longer extract from provider config + // Get credentials: prioritize UsageScript values, fallback to provider config + let api_key = usage_script + .api_key + .clone() + .filter(|k| !k.is_empty()) + .or_else(|| extract_api_key_from_provider(provider)) + .unwrap_or_default(); + + let base_url = usage_script + .base_url + .clone() + .filter(|u| !u.is_empty()) + .or_else(|| extract_base_url_from_provider(provider)) + .unwrap_or_default(); + ( usage_script.code.clone(), usage_script.timeout.unwrap_or(10), - usage_script.api_key.clone().unwrap_or_default(), - usage_script.base_url.clone().unwrap_or_default(), + api_key, + base_url, usage_script.access_token.clone(), usage_script.user_id.clone(), ) diff --git a/src-tauri/src/services/proxy.rs b/src-tauri/src/services/proxy.rs index 4fd172ad2..a897af1a1 100644 --- a/src-tauri/src/services/proxy.rs +++ b/src-tauri/src/services/proxy.rs @@ -5,6 +5,7 @@ use crate::app_config::AppType; use crate::config::{get_claude_settings_path, read_json_file, write_json_file}; use crate::database::Database; +use crate::provider::Provider; use crate::proxy::server::ProxyServer; use crate::proxy::types::*; use serde_json::{json, Value}; @@ -16,6 +17,8 @@ use tokio::sync::RwLock; pub struct ProxyService { db: Arc, server: Arc>>, + /// AppHandle,用于传递给 ProxyServer 以支持故障转移时的 UI 更新 + app_handle: Arc>>, } impl ProxyService { @@ -23,9 +26,17 @@ impl ProxyService { Self { db, server: Arc::new(RwLock::new(None)), + app_handle: Arc::new(RwLock::new(None)), } } + /// 设置 AppHandle(在应用初始化时调用) + pub fn set_app_handle(&self, handle: tauri::AppHandle) { + futures::executor::block_on(async { + *self.app_handle.write().await = Some(handle); + }); + } + /// 启动代理服务器 pub async fn start(&self) -> Result { // 1. 获取配置 @@ -44,7 +55,8 @@ impl ProxyService { } // 4. 创建并启动服务器 - let server = ProxyServer::new(config.clone(), self.db.clone()); + let app_handle = self.app_handle.read().await.clone(); + let server = ProxyServer::new(config.clone(), self.db.clone(), app_handle); let info = server .start() .await @@ -65,25 +77,22 @@ impl ProxyService { /// 启动代理服务器(带 Live 配置接管) pub async fn start_with_takeover(&self) -> Result { - // 1. 自动将各应用当前选中的供应商设置为代理目标 - self.setup_proxy_targets().await?; - - // 2. 备份各应用的 Live 配置 + // 1. 备份各应用的 Live 配置 self.backup_live_configs().await?; - // 3. 同步 Live 配置中的 Token 到数据库(确保代理能读到最新的 Token) + // 2. 同步 Live 配置中的 Token 到数据库(确保代理能读到最新的 Token) self.sync_live_to_providers().await?; - // 4. 接管各应用的 Live 配置(写入代理地址,清空 Token) + // 3. 接管各应用的 Live 配置(写入代理地址,清空 Token) self.takeover_live_configs().await?; - // 5. 设置接管状态 + // 4. 设置接管状态 self.db .set_live_takeover_active(true) .await .map_err(|e| format!("设置接管状态失败: {e}"))?; - // 6. 启动代理服务器 + // 5. 启动代理服务器 match self.start().await { Ok(info) => Ok(info), Err(e) => { @@ -96,27 +105,6 @@ impl ProxyService { } } - /// 自动设置代理目标:将各应用当前选中的供应商设置为代理目标 - async fn setup_proxy_targets(&self) -> Result<(), String> { - let app_types = ["claude", "codex", "gemini"]; - - for app_type in app_types { - // 获取当前选中的供应商 - if let Ok(Some(provider_id)) = self.db.get_current_provider(app_type) { - // 设置为代理目标 - if let Err(e) = self.db.set_proxy_target(&provider_id, app_type, true).await { - log::warn!("设置 {app_type} 的代理目标 {provider_id} 失败: {e}"); - } else { - log::info!("已将 {app_type} 的当前供应商 {provider_id} 设置为代理目标"); - } - } else { - log::debug!("{app_type} 没有当前供应商,跳过代理目标设置"); - } - } - - Ok(()) - } - /// 同步 Live 配置中的 Token 到数据库 /// /// 在清空 Live Token 之前调用,确保数据库中的 Provider 配置有最新的 Token。 @@ -154,7 +142,8 @@ impl ProxyService { log::warn!("同步 Claude Token 到数据库失败: {e}"); } else { log::info!( - "已同步 Claude Token 到数据库 (provider: {provider_id})" + "已同步 Claude Token 到数据库 (provider: {})", + provider_id ); } } @@ -193,7 +182,8 @@ impl ProxyService { log::warn!("同步 Codex Token 到数据库失败: {e}"); } else { log::info!( - "已同步 Codex Token 到数据库 (provider: {provider_id})" + "已同步 Codex Token 到数据库 (provider: {})", + provider_id ); } } @@ -203,13 +193,13 @@ impl ProxyService { } } - // Gemini: 同步 GOOGLE_API_KEY + // Gemini: 同步 GEMINI_API_KEY if let Ok(live_config) = self.read_gemini_live() { if let Some(provider_id) = self.db.get_current_provider("gemini").ok().flatten() { if let Ok(Some(mut provider)) = self.db.get_provider_by_id(&provider_id, "gemini") { // 从 live 配置提取 token if let Some(env) = live_config.get("env") { - if let Some(token) = env.get("GOOGLE_API_KEY").and_then(|v| v.as_str()) { + if let Some(token) = env.get("GEMINI_API_KEY").and_then(|v| v.as_str()) { if !token.is_empty() { // 更新 provider 的 settings_config if let Some(env_obj) = provider @@ -217,10 +207,10 @@ impl ProxyService { .get_mut("env") .and_then(|v| v.as_object_mut()) { - env_obj.insert("GOOGLE_API_KEY".to_string(), json!(token)); + env_obj.insert("GEMINI_API_KEY".to_string(), json!(token)); } else { provider.settings_config["env"] = json!({ - "GOOGLE_API_KEY": token + "GEMINI_API_KEY": token }); } // 保存到数据库 @@ -232,7 +222,8 @@ impl ProxyService { log::warn!("同步 Gemini Token 到数据库失败: {e}"); } else { log::info!( - "已同步 Gemini Token 到数据库 (provider: {provider_id})" + "已同步 Gemini Token 到数据库 (provider: {})", + provider_id ); } } @@ -287,6 +278,12 @@ impl ProxyService { .await .map_err(|e| format!("删除备份失败: {e}"))?; + // 5. 重置健康状态(让健康徽章恢复为正常) + self.db + .clear_all_provider_health() + .await + .map_err(|e| format!("重置健康状态失败: {e}"))?; + log::info!("代理已停止,Live 配置已恢复"); Ok(()) } @@ -357,34 +354,42 @@ impl ProxyService { }); } self.write_claude_live(&live_config)?; - log::info!("Claude Live 配置已接管,代理地址: {proxy_url}"); + log::info!("Claude Live 配置已接管,代理地址: {}", proxy_url); } - // Codex: 修改 OPENAI_BASE_URL,使用占位符替代真实 Token(代理会注入真实 Token) + // Codex: 修改 config.toml 的 base_url,auth.json 的 OPENAI_API_KEY(代理会注入真实 Token) if let Ok(mut live_config) = self.read_codex_live() { + // 1. 修改 auth.json 中的 OPENAI_API_KEY(使用占位符) if let Some(auth) = live_config.get_mut("auth").and_then(|v| v.as_object_mut()) { - auth.insert("OPENAI_BASE_URL".to_string(), json!(&proxy_url)); - // 使用占位符,避免显示缺少 key 的警告 auth.insert("OPENAI_API_KEY".to_string(), json!("PROXY_MANAGED")); } + + // 2. 修改 config.toml 中的 base_url + let config_str = live_config + .get("config") + .and_then(|v| v.as_str()) + .unwrap_or(""); + let updated_config = Self::update_toml_base_url(config_str, &proxy_url); + live_config["config"] = json!(updated_config); + self.write_codex_live(&live_config)?; - log::info!("Codex Live 配置已接管,代理地址: {proxy_url}"); + log::info!("Codex Live 配置已接管,代理地址: {}", proxy_url); } - // Gemini: 修改 GEMINI_API_BASE,使用占位符替代真实 Token(代理会注入真实 Token) + // Gemini: 修改 GOOGLE_GEMINI_BASE_URL,使用占位符替代真实 Token(代理会注入真实 Token) if let Ok(mut live_config) = self.read_gemini_live() { if let Some(env) = live_config.get_mut("env").and_then(|v| v.as_object_mut()) { - env.insert("GEMINI_API_BASE".to_string(), json!(&proxy_url)); + env.insert("GOOGLE_GEMINI_BASE_URL".to_string(), json!(&proxy_url)); // 使用占位符,避免显示缺少 key 的警告 - env.insert("GOOGLE_API_KEY".to_string(), json!("PROXY_MANAGED")); + env.insert("GEMINI_API_KEY".to_string(), json!("PROXY_MANAGED")); } else { live_config["env"] = json!({ - "GEMINI_API_BASE": &proxy_url, - "GOOGLE_API_KEY": "PROXY_MANAGED" + "GOOGLE_GEMINI_BASE_URL": &proxy_url, + "GEMINI_API_KEY": "PROXY_MANAGED" }); } self.write_gemini_live(&live_config)?; - log::info!("Gemini Live 配置已接管,代理地址: {proxy_url}"); + log::info!("Gemini Live 配置已接管,代理地址: {}", proxy_url); } Ok(()) @@ -427,6 +432,73 @@ impl ProxyService { .map_err(|e| format!("检查接管状态失败: {e}")) } + /// 从异常退出中恢复(启动时调用) + /// + /// 检测到 live_takeover_active=true 但代理未运行时调用此方法。 + /// 会恢复 Live 配置、清除接管标志、删除备份。 + pub async fn recover_from_crash(&self) -> Result<(), String> { + // 1. 恢复 Live 配置 + self.restore_live_configs().await?; + + // 2. 清除接管标志 + self.db + .set_live_takeover_active(false) + .await + .map_err(|e| format!("清除接管状态失败: {e}"))?; + + // 3. 删除备份 + self.db + .delete_all_live_backups() + .await + .map_err(|e| format!("删除备份失败: {e}"))?; + + log::info!("已从异常退出中恢复 Live 配置"); + Ok(()) + } + + /// 从供应商配置更新 Live 备份(用于代理模式下的热切换) + /// + /// 与 backup_live_configs() 不同,此方法从供应商的 settings_config 生成备份, + /// 而不是从 Live 文件读取(因为 Live 文件已被代理接管)。 + pub async fn update_live_backup_from_provider( + &self, + app_type: &str, + provider: &Provider, + ) -> Result<(), String> { + let backup_json = match app_type { + "claude" => { + // Claude: settings_config 直接作为备份 + serde_json::to_string(&provider.settings_config) + .map_err(|e| format!("序列化 Claude 配置失败: {e}"))? + } + "codex" => { + // Codex: settings_config 包含 {"auth": ..., "config": ...},直接使用 + serde_json::to_string(&provider.settings_config) + .map_err(|e| format!("序列化 Codex 配置失败: {e}"))? + } + "gemini" => { + // Gemini: 只提取 env 字段(与原始备份格式一致) + // proxy.rs 的 read_gemini_live() 返回 {"env": {...}} + let env_backup = if let Some(env) = provider.settings_config.get("env") { + json!({ "env": env }) + } else { + json!({ "env": {} }) + }; + serde_json::to_string(&env_backup) + .map_err(|e| format!("序列化 Gemini 配置失败: {e}"))? + } + _ => return Err(format!("未知的应用类型: {app_type}")), + }; + + self.db + .save_live_backup(app_type, &backup_json) + .await + .map_err(|e| format!("更新 {app_type} 备份失败: {e}"))?; + + log::info!("已更新 {app_type} Live 备份(热切换)"); + Ok(()) + } + /// 代理模式下切换供应商(热切换,不写 Live) pub async fn switch_proxy_target( &self, @@ -441,12 +513,29 @@ impl ProxyService { .set_current_provider(app_type_enum.as_str(), provider_id) .map_err(|e| format!("更新当前供应商失败: {e}"))?; - log::info!("代理模式:已切换 {app_type} 的目标供应商为 {provider_id}"); + log::info!( + "代理模式:已切换 {} 的目标供应商为 {}", + app_type, + provider_id + ); Ok(()) } // ==================== Live 配置读写辅助方法 ==================== + /// 更新 TOML 字符串中的 base_url + fn update_toml_base_url(toml_str: &str, new_url: &str) -> String { + use toml_edit::DocumentMut; + + let mut doc = toml_str + .parse::() + .unwrap_or_else(|_| DocumentMut::new()); + + doc["base_url"] = toml_edit::value(new_url); + + doc.to_string() + } + fn read_claude_live(&self) -> Result { let path = get_claude_settings_path(); if !path.exists() { @@ -582,7 +671,8 @@ impl ProxyService { .map_err(|e| format!("重启前停止代理服务器失败: {e}"))?; } - let new_server = ProxyServer::new(new_config, self.db.clone()); + let app_handle = self.app_handle.read().await.clone(); + let new_server = ProxyServer::new(new_config, self.db.clone(), app_handle); new_server .start() .await @@ -602,4 +692,41 @@ impl ProxyService { pub async fn is_running(&self) -> bool { self.server.read().await.is_some() } + + /// 热更新熔断器配置 + /// + /// 如果代理服务器正在运行,将新配置应用到所有已创建的熔断器实例 + pub async fn update_circuit_breaker_configs( + &self, + config: crate::proxy::CircuitBreakerConfig, + ) -> Result<(), String> { + if let Some(server) = self.server.read().await.as_ref() { + server.update_circuit_breaker_configs(config).await; + log::info!("已热更新运行中的熔断器配置"); + } else { + log::debug!("代理服务器未运行,熔断器配置将在下次启动时生效"); + } + Ok(()) + } + + /// 重置指定 Provider 的熔断器 + /// + /// 如果代理服务器正在运行,立即重置内存中的熔断器状态 + pub async fn reset_provider_circuit_breaker( + &self, + provider_id: &str, + app_type: &str, + ) -> Result<(), String> { + if let Some(server) = self.server.read().await.as_ref() { + server + .reset_provider_circuit_breaker(provider_id, app_type) + .await; + log::info!( + "已重置 Provider {} (app: {}) 的熔断器", + provider_id, + app_type + ); + } + Ok(()) + } } diff --git a/src-tauri/tests/deeplink_import.rs b/src-tauri/tests/deeplink_import.rs index 2768f8953..5a0238c61 100644 --- a/src-tauri/tests/deeplink_import.rs +++ b/src-tauri/tests/deeplink_import.rs @@ -1,6 +1,8 @@ use std::sync::Arc; -use cc_switch_lib::{import_provider_from_deeplink, parse_deeplink_url, AppState, Database}; +use cc_switch_lib::{ + import_provider_from_deeplink, parse_deeplink_url, AppState, Database, ProxyService, +}; #[path = "support.rs"] mod support; @@ -16,8 +18,11 @@ fn deeplink_import_claude_provider_persists_to_db() { let request = parse_deeplink_url(url).expect("parse deeplink url"); let db = Arc::new(Database::memory().expect("create memory db")); - - let state = AppState { db: db.clone() }; + let proxy_service = ProxyService::new(db.clone()); + let state = AppState { + db: db.clone(), + proxy_service, + }; let provider_id = import_provider_from_deeplink(&state, request.clone()) .expect("import provider from deeplink"); @@ -53,8 +58,11 @@ fn deeplink_import_codex_provider_builds_auth_and_config() { let request = parse_deeplink_url(url).expect("parse deeplink url"); let db = Arc::new(Database::memory().expect("create memory db")); - - let state = AppState { db: db.clone() }; + let proxy_service = ProxyService::new(db.clone()); + let state = AppState { + db: db.clone(), + proxy_service, + }; let provider_id = import_provider_from_deeplink(&state, request.clone()) .expect("import provider from deeplink"); diff --git a/src-tauri/tests/support.rs b/src-tauri/tests/support.rs index b954bad79..bc5c8b716 100644 --- a/src-tauri/tests/support.rs +++ b/src-tauri/tests/support.rs @@ -1,7 +1,9 @@ use std::path::{Path, PathBuf}; use std::sync::{Arc, Mutex, OnceLock}; -use cc_switch_lib::{update_settings, AppSettings, AppState, Database, MultiAppConfig}; +use cc_switch_lib::{ + update_settings, AppSettings, AppState, Database, MultiAppConfig, ProxyService, +}; /// 为测试设置隔离的 HOME 目录,避免污染真实用户数据。 pub fn ensure_test_home() -> &'static Path { @@ -48,15 +50,17 @@ pub fn test_mutex() -> &'static Mutex<()> { /// 创建测试用的 AppState,包含一个空的数据库 pub fn create_test_state() -> Result> { - let db = Database::init()?; - Ok(AppState { db: Arc::new(db) }) + let db = Arc::new(Database::init()?); + let proxy_service = ProxyService::new(db.clone()); + Ok(AppState { db, proxy_service }) } /// 创建测试用的 AppState,并从 MultiAppConfig 迁移数据 pub fn create_test_state_with_config( config: &MultiAppConfig, ) -> Result> { - let db = Database::init()?; + let db = Arc::new(Database::init()?); db.migrate_from_json(config)?; - Ok(AppState { db: Arc::new(db) }) + let proxy_service = ProxyService::new(db.clone()); + Ok(AppState { db, proxy_service }) } diff --git a/src/App.tsx b/src/App.tsx index d37e5f4ac..41efe2640 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -142,6 +142,7 @@ function App() { if (migrated) { toast.success( t("migration.success", { defaultValue: "配置迁移成功" }), + { closeButton: true }, ); } } catch (error) { diff --git a/src/components/DeepLinkImportDialog.tsx b/src/components/DeepLinkImportDialog.tsx index 96035e38c..86e6b0e6f 100644 --- a/src/components/DeepLinkImportDialog.tsx +++ b/src/components/DeepLinkImportDialog.tsx @@ -128,6 +128,7 @@ export function DeepLinkImportDialog() { description: t("deeplink.mcpImportSuccessDescription", { count: summary.importedCount, }), + closeButton: true, }); } }; @@ -142,6 +143,7 @@ export function DeepLinkImportDialog() { description: t("deeplink.importSuccessDescription", { name: request.name, }), + closeButton: true, }); } else if (result.type === "prompt") { // Prompts don't use React Query, trigger a custom event for refresh @@ -154,6 +156,7 @@ export function DeepLinkImportDialog() { description: t("deeplink.promptImportSuccessDescription", { name: request.name, }), + closeButton: true, }); } else if (result.type === "mcp") { await refreshMcp(result); @@ -171,6 +174,7 @@ export function DeepLinkImportDialog() { description: t("deeplink.skillImportSuccessDescription", { repo: request.repo, }), + closeButton: true, }); } } else if (isMcpImportResult(result)) { @@ -185,6 +189,7 @@ export function DeepLinkImportDialog() { description: t("deeplink.importSuccessDescription", { name: request.name, }), + closeButton: true, }); } @@ -605,6 +610,86 @@ export function DeepLinkImportDialog() { )} + {/* Usage Script Configuration (v3.9+) */} + {request.usageScript && ( +
+
+
+ {t("deeplink.usageScript", { + defaultValue: "用量查询", + })} +
+
+ + {request.usageEnabled !== false + ? t("deeplink.usageScriptEnabled", { + defaultValue: "已启用", + }) + : t("deeplink.usageScriptDisabled", { + defaultValue: "未启用", + })} + +
+
+ + {/* Usage API Key (if different from provider) */} + {request.usageApiKey && + request.usageApiKey !== request.apiKey && ( +
+
+ {t("deeplink.usageApiKey", { + defaultValue: "用量 API Key", + })} +
+
+ {request.usageApiKey.length > 4 + ? `${request.usageApiKey.substring(0, 4)}${"*".repeat(12)}` + : "****"} +
+
+ )} + + {/* Usage Base URL (if different from provider) */} + {request.usageBaseUrl && + request.usageBaseUrl !== request.endpoint && ( +
+
+ {t("deeplink.usageBaseUrl", { + defaultValue: "用量查询地址", + })} +
+
+ {request.usageBaseUrl} +
+
+ )} + + {/* Auto Query Interval */} + {request.usageAutoInterval && + request.usageAutoInterval > 0 && ( +
+
+ {t("deeplink.usageAutoInterval", { + defaultValue: "自动查询", + })} +
+
+ {t("deeplink.usageAutoIntervalValue", { + defaultValue: "每 {{minutes}} 分钟", + minutes: request.usageAutoInterval, + })} +
+
+ )} +
+ )} + {/* Warning */}
{t("deeplink.warning")} diff --git a/src/components/JsonEditor.tsx b/src/components/JsonEditor.tsx index 72972e578..860660252 100644 --- a/src/components/JsonEditor.tsx +++ b/src/components/JsonEditor.tsx @@ -234,7 +234,9 @@ const JsonEditor: React.FC = ({ try { const formatted = formatJSON(currentValue); onChange(formatted); - toast.success(t("common.formatSuccess", { defaultValue: "格式化成功" })); + toast.success(t("common.formatSuccess", { defaultValue: "格式化成功" }), { + closeButton: true, + }); } catch (error) { const errorMessage = error instanceof Error ? error.message : String(error); diff --git a/src/components/UsageScriptModal.tsx b/src/components/UsageScriptModal.tsx index bbda3c05e..c23dbb552 100644 --- a/src/components/UsageScriptModal.tsx +++ b/src/components/UsageScriptModal.tsx @@ -227,6 +227,7 @@ const UsageScriptModal: React.FC = ({ .join(", "); toast.success(`${t("usageScript.testSuccess")}${summary}`, { duration: 3000, + closeButton: true, }); } else { toast.error( @@ -259,7 +260,10 @@ const UsageScriptModal: React.FC = ({ printWidth: 80, }); setScript({ ...script, code: formatted.trim() }); - toast.success(t("usageScript.formatSuccess"), { duration: 1000 }); + toast.success(t("usageScript.formatSuccess"), { + duration: 1000, + closeButton: true, + }); } catch (error: any) { toast.error( `${t("usageScript.formatFailed")}: ${error?.message || t("jsonEditor.invalidJson")}`, @@ -400,15 +404,25 @@ const UsageScriptModal: React.FC = ({ {/* 凭证配置 */} {shouldShowCredentialsConfig && (
-

- {t("usageScript.credentialsConfig")} -

+
+

+ {t("usageScript.credentialsConfig")} +

+

+ {t("usageScript.credentialsHint")} +

+
{selectedTemplate === TEMPLATE_KEYS.GENERAL && ( <>
- +
= ({ onChange={(e) => setScript({ ...script, apiKey: e.target.value }) } - placeholder="sk-xxxxx" + placeholder={t("usageScript.apiKeyPlaceholder")} autoComplete="off" className="border-white/10" /> @@ -444,7 +458,10 @@ const UsageScriptModal: React.FC = ({
= ({ onChange={(e) => setScript({ ...script, baseUrl: e.target.value }) } - placeholder="https://api.example.com" + placeholder={t("usageScript.baseUrlPlaceholder")} autoComplete="off" className="border-white/10" /> diff --git a/src/components/env/EnvWarningBanner.tsx b/src/components/env/EnvWarningBanner.tsx index 7a2ec5d06..55c5621fa 100644 --- a/src/components/env/EnvWarningBanner.tsx +++ b/src/components/env/EnvWarningBanner.tsx @@ -79,6 +79,7 @@ export function EnvWarningBanner({ path: backupInfo.backupPath, }), duration: 5000, + closeButton: true, }); // 清空选择并通知父组件 diff --git a/src/components/mcp/McpFormModal.tsx b/src/components/mcp/McpFormModal.tsx index d75e5b301..5c608a4ff 100644 --- a/src/components/mcp/McpFormModal.tsx +++ b/src/components/mcp/McpFormModal.tsx @@ -391,7 +391,7 @@ const McpFormModal: React.FC = ({ } await upsertMutation.mutateAsync(entry); - toast.success(t("common.success")); + toast.success(t("common.success"), { closeButton: true }); await onSave(); } catch (error: any) { const detail = extractErrorMessage(error); diff --git a/src/components/mcp/UnifiedMcpPanel.tsx b/src/components/mcp/UnifiedMcpPanel.tsx index 0e5448cea..08f6c340b 100644 --- a/src/components/mcp/UnifiedMcpPanel.tsx +++ b/src/components/mcp/UnifiedMcpPanel.tsx @@ -99,7 +99,7 @@ const UnifiedMcpPanel = React.forwardRef< try { await deleteServerMutation.mutateAsync(id); setConfirmDialog(null); - toast.success(t("common.success")); + toast.success(t("common.success"), { closeButton: true }); } catch (error) { toast.error(t("common.error"), { description: String(error), diff --git a/src/components/providers/ProviderActions.tsx b/src/components/providers/ProviderActions.tsx index 990b26913..361fa2b2c 100644 --- a/src/components/providers/ProviderActions.tsx +++ b/src/components/providers/ProviderActions.tsx @@ -7,7 +7,6 @@ import { Play, TestTube2, Trash2, - // RotateCcw, // TODO: 暂时注释,等待故障转移功能启用 } from "lucide-react"; import { useTranslation } from "react-i18next"; import { Button } from "@/components/ui/button"; @@ -23,9 +22,6 @@ interface ProviderActionsProps { onTest?: () => void; onConfigureUsage: () => void; onDelete: () => void; - onResetCircuitBreaker?: () => void; - isProxyTarget?: boolean; - consecutiveFailures?: number; } export function ProviderActions({ @@ -38,9 +34,6 @@ export function ProviderActions({ onTest, onConfigureUsage, onDelete, - onResetCircuitBreaker: _onResetCircuitBreaker, // 暂未使用,前缀 _ 避免 lint 警告 - isProxyTarget: _isProxyTarget, // 暂未使用,前缀 _ 避免 lint 警告 - consecutiveFailures: _consecutiveFailures = 0, // 暂未使用,前缀 _ 避免 lint 警告 }: ProviderActionsProps) { const { t } = useTranslation(); const iconButtonClass = "h-8 w-8 p-1"; @@ -123,33 +116,6 @@ export function ProviderActions({ - {/* 重置熔断器按钮 - 代理目标启用时显示 */} - {/* TODO: 暂时隐藏,后续根据故障转移功能启用 */} - {/* {onResetCircuitBreaker && isProxyTarget && ( - - )} */} -
-
+

{provider.name}

@@ -253,19 +248,56 @@ export function ProviderCard({
-
-
- +
+ {/* 用量信息区域 - hover 时向左移动,为操作按钮腾出空间 */} +
+
+ {/* 多套餐时显示套餐数量,单套餐时显示详细信息 */} + {hasMultiplePlans ? ( +
+ + {t("usage.multiplePlans", { + count: usage?.data?.length || 0, + defaultValue: `${usage?.data?.length || 0} 个套餐`, + })} + +
+ ) : ( + + )} + {/* 展开/折叠按钮 - 仅在有多套餐时显示 */} + {hasMultiplePlans && ( + + )} +
-
+ {/* 操作按钮区域 - 绝对定位在右侧,hover 时滑入 */} +
onTest(provider) : undefined} onConfigureUsage={() => onConfigureUsage(provider)} onDelete={() => onDelete(provider)} - onResetCircuitBreaker={ - isProxyRunning && provider.isProxyTarget - ? handleResetCircuitBreaker - : undefined - } - isProxyTarget={provider.isProxyTarget} - consecutiveFailures={health?.consecutive_failures ?? 0} />
+ + {/* 展开的完整套餐列表 */} + {isExpanded && hasMultiplePlans && ( +
+ +
+ )}
); } diff --git a/src/components/proxy/AutoFailoverConfigPanel.tsx b/src/components/proxy/AutoFailoverConfigPanel.tsx index c7fab569d..055fcf25f 100644 --- a/src/components/proxy/AutoFailoverConfigPanel.tsx +++ b/src/components/proxy/AutoFailoverConfigPanel.tsx @@ -54,6 +54,7 @@ export function AutoFailoverConfigPanel({ }); toast.success( t("proxy.autoFailover.configSaved", "自动故障转移配置已保存"), + { closeButton: true }, ); } catch (e) { toast.error( @@ -100,7 +101,7 @@ export function AutoFailoverConfigPanel({ {t( "proxy.autoFailover.info", - "当启用多个代理目标时,系统会按优先级顺序依次尝试。当某个供应商连续失败达到阈值时,熔断器会自动打开,跳过该供应商。", + "当故障转移队列中配置了多个供应商时,系统会在请求失败时按优先级顺序依次尝试。当某个供应商连续失败达到阈值时,熔断器会打开并在一段时间内跳过该供应商。", )} diff --git a/src/components/proxy/CircuitBreakerConfigPanel.tsx b/src/components/proxy/CircuitBreakerConfigPanel.tsx index 3014f532a..93f48d594 100644 --- a/src/components/proxy/CircuitBreakerConfigPanel.tsx +++ b/src/components/proxy/CircuitBreakerConfigPanel.tsx @@ -34,7 +34,7 @@ export function CircuitBreakerConfigPanel() { const handleSave = async () => { try { await updateConfig.mutateAsync(formData); - toast.success("熔断器配置已保存"); + toast.success("熔断器配置已保存", { closeButton: true }); } catch (error) { toast.error("保存失败: " + String(error)); } diff --git a/src/components/proxy/FailoverQueueManager.tsx b/src/components/proxy/FailoverQueueManager.tsx new file mode 100644 index 000000000..11da074e8 --- /dev/null +++ b/src/components/proxy/FailoverQueueManager.tsx @@ -0,0 +1,414 @@ +/** + * 故障转移队列管理组件 + * + * 允许用户管理代理模式下的故障转移队列,支持: + * - 拖拽排序 + * - 添加/移除供应商 + * - 启用/禁用队列项 + */ + +import { useState, useCallback, useMemo } from "react"; +import { useTranslation } from "react-i18next"; +import { CSS } from "@dnd-kit/utilities"; +import { DndContext, closestCenter } from "@dnd-kit/core"; +import { + SortableContext, + useSortable, + verticalListSortingStrategy, +} from "@dnd-kit/sortable"; +import { + KeyboardSensor, + PointerSensor, + useSensor, + useSensors, + type DragEndEvent, +} from "@dnd-kit/core"; +import { arrayMove, sortableKeyboardCoordinates } from "@dnd-kit/sortable"; +import { toast } from "sonner"; +import { + GripVertical, + Plus, + Trash2, + Loader2, + Info, + AlertTriangle, +} from "lucide-react"; +import { Button } from "@/components/ui/button"; +import { Switch } from "@/components/ui/switch"; +import { Alert, AlertDescription } from "@/components/ui/alert"; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from "@/components/ui/select"; +import { cn } from "@/lib/utils"; +import type { FailoverQueueItem } from "@/types/proxy"; +import type { AppId } from "@/lib/api"; +import { + useFailoverQueue, + useAvailableProvidersForFailover, + useAddToFailoverQueue, + useRemoveFromFailoverQueue, + useReorderFailoverQueue, + useSetFailoverItemEnabled, +} from "@/lib/query/failover"; + +interface FailoverQueueManagerProps { + appType: AppId; + disabled?: boolean; +} + +export function FailoverQueueManager({ + appType, + disabled = false, +}: FailoverQueueManagerProps) { + const { t } = useTranslation(); + const [selectedProviderId, setSelectedProviderId] = useState(""); + + // 查询数据 + const { + data: queue, + isLoading: isQueueLoading, + error: queueError, + } = useFailoverQueue(appType); + const { data: availableProviders, isLoading: isProvidersLoading } = + useAvailableProvidersForFailover(appType); + + // Mutations + const addToQueue = useAddToFailoverQueue(); + const removeFromQueue = useRemoveFromFailoverQueue(); + const reorderQueue = useReorderFailoverQueue(); + const setItemEnabled = useSetFailoverItemEnabled(); + + // 拖拽配置 + const sensors = useSensors( + useSensor(PointerSensor, { + activationConstraint: { distance: 8 }, + }), + useSensor(KeyboardSensor, { + coordinateGetter: sortableKeyboardCoordinates, + }), + ); + + // 排序后的队列 + const sortedQueue = useMemo(() => { + if (!queue) return []; + return [...queue].sort((a, b) => a.queueOrder - b.queueOrder); + }, [queue]); + + // 处理拖拽结束 + const handleDragEnd = useCallback( + async (event: DragEndEvent) => { + const { active, over } = event; + if (!over || active.id === over.id || !sortedQueue) return; + + const oldIndex = sortedQueue.findIndex( + (item) => item.providerId === active.id, + ); + const newIndex = sortedQueue.findIndex( + (item) => item.providerId === over.id, + ); + + if (oldIndex === -1 || newIndex === -1) return; + + const reordered = arrayMove(sortedQueue, oldIndex, newIndex); + const providerIds = reordered.map((item) => item.providerId); + + try { + await reorderQueue.mutateAsync({ appType, providerIds }); + toast.success( + t("proxy.failoverQueue.reorderSuccess", "队列顺序已更新"), + { closeButton: true }, + ); + } catch (error) { + toast.error( + t("proxy.failoverQueue.reorderFailed", "更新顺序失败") + + ": " + + String(error), + ); + } + }, + [sortedQueue, appType, reorderQueue, t], + ); + + // 添加供应商到队列 + const handleAddProvider = async () => { + if (!selectedProviderId) return; + + try { + await addToQueue.mutateAsync({ + appType, + providerId: selectedProviderId, + }); + setSelectedProviderId(""); + toast.success( + t("proxy.failoverQueue.addSuccess", "已添加到故障转移队列"), + { closeButton: true }, + ); + } catch (error) { + toast.error( + t("proxy.failoverQueue.addFailed", "添加失败") + ": " + String(error), + ); + } + }; + + // 从队列移除供应商 + const handleRemoveProvider = async (providerId: string) => { + try { + await removeFromQueue.mutateAsync({ appType, providerId }); + toast.success( + t("proxy.failoverQueue.removeSuccess", "已从故障转移队列移除"), + { closeButton: true }, + ); + } catch (error) { + toast.error( + t("proxy.failoverQueue.removeFailed", "移除失败") + + ": " + + String(error), + ); + } + }; + + // 切换启用状态 + const handleToggleEnabled = async (providerId: string, enabled: boolean) => { + try { + await setItemEnabled.mutateAsync({ appType, providerId, enabled }); + } catch (error) { + toast.error( + t("proxy.failoverQueue.toggleFailed", "状态更新失败") + + ": " + + String(error), + ); + } + }; + + if (isQueueLoading) { + return ( +
+ +
+ ); + } + + if (queueError) { + return ( + + + {String(queueError)} + + ); + } + + return ( +
+ {/* 说明信息 */} + + + + {t( + "proxy.failoverQueue.info", + "当前激活的供应商始终优先。当请求失败时,系统会按队列顺序依次尝试其他供应商。", + )} + + + + {/* 添加供应商 */} +
+ + +
+ + {/* 队列列表 */} + {sortedQueue.length === 0 ? ( +
+

+ {t( + "proxy.failoverQueue.empty", + "故障转移队列为空。添加供应商以启用自动故障转移。", + )} +

+
+ ) : ( + + item.providerId)} + strategy={verticalListSortingStrategy} + > +
+ {sortedQueue.map((item, index) => ( + + ))} +
+
+
+ )} + + {/* 队列说明 */} + {sortedQueue.length > 0 && ( +

+ {t( + "proxy.failoverQueue.dragHint", + "拖拽供应商可调整故障转移顺序,序号越小优先级越高。", + )} +

+ )} +
+ ); +} + +interface SortableQueueItemProps { + item: FailoverQueueItem; + index: number; + disabled: boolean; + onToggleEnabled: (providerId: string, enabled: boolean) => void; + onRemove: (providerId: string) => void; + isRemoving: boolean; + isToggling: boolean; +} + +function SortableQueueItem({ + item, + index, + disabled, + onToggleEnabled, + onRemove, + isRemoving, + isToggling, +}: SortableQueueItemProps) { + const { t } = useTranslation(); + const { + setNodeRef, + attributes, + listeners, + transform, + transition, + isDragging, + } = useSortable({ id: item.providerId, disabled }); + + const style = { + transform: CSS.Transform.toString(transform), + transition, + }; + + return ( +
+ {/* 拖拽手柄 */} + + + {/* 序号 */} +
+ {index + 1} +
+ + {/* 供应商名称 */} +
+ + {item.providerName} + +
+ + {/* 启用开关 */} + onToggleEnabled(item.providerId, checked)} + disabled={disabled || isToggling} + aria-label={t("proxy.failoverQueue.toggleEnabled", "启用/禁用")} + /> + + {/* 删除按钮 */} + +
+ ); +} diff --git a/src/components/proxy/ProxyPanel.tsx b/src/components/proxy/ProxyPanel.tsx index bb0dbe271..64521a1e6 100644 --- a/src/components/proxy/ProxyPanel.tsx +++ b/src/components/proxy/ProxyPanel.tsx @@ -11,7 +11,7 @@ import { Button } from "@/components/ui/button"; import { useProxyStatus } from "@/hooks/useProxyStatus"; import { ProxySettingsDialog } from "./ProxySettingsDialog"; import { toast } from "sonner"; -import { useProxyTargets } from "@/lib/query/failover"; +import { useFailoverQueue } from "@/lib/query/failover"; import { ProviderHealthBadge } from "@/components/providers/ProviderHealthBadge"; import { useProviderHealth } from "@/lib/query/failover"; import type { ProxyStatus } from "@/types/proxy"; @@ -20,10 +20,11 @@ export function ProxyPanel() { const { status, isRunning } = useProxyStatus(); const [showSettings, setShowSettings] = useState(false); - // 获取所有三个应用类型的代理目标列表 - const { data: claudeTargets = [] } = useProxyTargets("claude"); - const { data: codexTargets = [] } = useProxyTargets("codex"); - const { data: geminiTargets = [] } = useProxyTargets("gemini"); + // 获取所有三个应用类型的故障转移队列(不包含当前供应商) + // 当前供应商始终优先,队列仅用于失败后的备用顺序 + const { data: claudeQueue = [] } = useFailoverQueue("claude"); + const { data: codexQueue = [] } = useFailoverQueue("codex"); + const { data: geminiQueue = [] } = useFailoverQueue("gemini"); const formatUptime = (seconds: number): string => { const hours = Math.floor(seconds / 3600); @@ -69,7 +70,7 @@ export function ProxyPanel() { navigator.clipboard.writeText( `http://${status.address}:${status.port}`, ); - toast.success("地址已复制"); + toast.success("地址已复制", { closeButton: true }); }} > 复制 @@ -113,9 +114,9 @@ export function ProxyPanel() {
{/* 供应商队列 - 按应用类型分组展示 */} - {(claudeTargets.length > 0 || - codexTargets.length > 0 || - geminiTargets.length > 0) && ( + {(claudeQueue.length > 0 || + codexQueue.length > 0 || + geminiQueue.length > 0) && (
@@ -125,31 +126,49 @@ export function ProxyPanel() {
{/* Claude 队列 */} - {claudeTargets.length > 0 && ( + {claudeQueue.length > 0 && ( item.enabled) + .sort((a, b) => a.queueOrder - b.queueOrder) + .map((item) => ({ + id: item.providerId, + name: item.providerName, + }))} status={status} /> )} {/* Codex 队列 */} - {codexTargets.length > 0 && ( + {codexQueue.length > 0 && ( item.enabled) + .sort((a, b) => a.queueOrder - b.queueOrder) + .map((item) => ({ + id: item.providerId, + name: item.providerName, + }))} status={status} /> )} {/* Gemini 队列 */} - {geminiTargets.length > 0 && ( + {geminiQueue.length > 0 && ( item.enabled) + .sort((a, b) => a.queueOrder - b.queueOrder) + .map((item) => ({ + id: item.providerId, + name: item.providerName, + }))} status={status} /> )} diff --git a/src/components/settings/AboutSection.tsx b/src/components/settings/AboutSection.tsx index ba38430b4..cf9465ec4 100644 --- a/src/components/settings/AboutSection.tsx +++ b/src/components/settings/AboutSection.tsx @@ -142,7 +142,7 @@ export function AboutSection({ isPortable }: AboutSectionProps) { try { const available = await checkUpdate(); if (!available) { - toast.success(t("settings.upToDate")); + toast.success(t("settings.upToDate"), { closeButton: true }); } } catch (error) { console.error("[AboutSection] Check update failed", error); diff --git a/src/components/settings/SettingsPage.tsx b/src/components/settings/SettingsPage.tsx index 42ece5a2e..1760382bf 100644 --- a/src/components/settings/SettingsPage.tsx +++ b/src/components/settings/SettingsPage.tsx @@ -7,7 +7,9 @@ import { Coins, Database, Server, + ChevronDown, } from "lucide-react"; +import * as AccordionPrimitive from "@radix-ui/react-accordion"; import { toast } from "sonner"; import { Dialog, @@ -35,6 +37,7 @@ import { ProxyPanel } from "@/components/proxy"; import { PricingConfigPanel } from "@/components/usage/PricingConfigPanel"; import { ModelTestConfigPanel } from "@/components/usage/ModelTestConfigPanel"; import { AutoFailoverConfigPanel } from "@/components/proxy/AutoFailoverConfigPanel"; +import { FailoverQueueManager } from "@/components/proxy/FailoverQueueManager"; import { UsageDashboard } from "@/components/usage/UsageDashboard"; import { useSettings } from "@/hooks/useSettings"; import { useImportExport } from "@/hooks/useImportExport"; @@ -135,7 +138,7 @@ export function SettingsPage({ const handleRestartNow = useCallback(async () => { setShowRestartPrompt(false); if (import.meta.env.DEV) { - toast.success(t("settings.devModeRestartHint")); + toast.success(t("settings.devModeRestartHint"), { closeButton: true }); closeAfterSave(); return; } @@ -278,10 +281,10 @@ export function SettingsPage({ - -
+ +
@@ -293,27 +296,26 @@ export function SettingsPage({

-
e.stopPropagation()} + + + +
+ - - - {isRunning ? "运行中" : "已停止"} - - -
+ {isRunning ? "运行中" : "已停止"} + +
- +
@@ -343,10 +345,10 @@ export function SettingsPage({ - -
+ +
@@ -354,29 +356,70 @@ export function SettingsPage({ 自动故障转移

- 配置自动故障转移和熔断策略 + 配置故障转移队列和熔断策略

-
e.stopPropagation()} - > -
- {/* Removed status text as requested */} - + + + +
+ +
+ + +
+ {/* 故障转移队列管理 */} +
+
+

+ {t("proxy.failoverQueue.title", "故障转移队列")} +

+

+ {t( + "proxy.failoverQueue.description", + "管理各应用的供应商故障转移顺序", + )} +

+ + + Claude + Codex + Gemini + + + + + + + + + + + +
+ + {/* 熔断器配置 */} +
+
- - - diff --git a/src/components/skills/SkillsPage.tsx b/src/components/skills/SkillsPage.tsx index 2bb9e186e..2566c92bf 100644 --- a/src/components/skills/SkillsPage.tsx +++ b/src/components/skills/SkillsPage.tsx @@ -103,7 +103,9 @@ export const SkillsPage = forwardRef( const handleInstall = async (directory: string) => { try { await skillsApi.install(directory, selectedApp); - toast.success(t("skills.installSuccess", { name: directory })); + toast.success(t("skills.installSuccess", { name: directory }), { + closeButton: true, + }); await loadSkills(); } catch (error) { const errorMessage = @@ -132,7 +134,9 @@ export const SkillsPage = forwardRef( const handleUninstall = async (directory: string) => { try { await skillsApi.uninstall(directory, selectedApp); - toast.success(t("skills.uninstallSuccess", { name: directory })); + toast.success(t("skills.uninstallSuccess", { name: directory }), { + closeButton: true, + }); await loadSkills(); } catch (error) { const errorMessage = @@ -180,12 +184,15 @@ export const SkillsPage = forwardRef( name: repo.name, count: repoSkillCount, }), + { closeButton: true }, ); }; const handleRemoveRepo = async (owner: string, name: string) => { await skillsApi.removeRepo(owner, name); - toast.success(t("skills.repo.removeSuccess", { owner, name })); + toast.success(t("skills.repo.removeSuccess", { owner, name }), { + closeButton: true, + }); await Promise.all([loadRepos(), loadSkills()]); }; diff --git a/src/components/usage/ModelTestConfigPanel.tsx b/src/components/usage/ModelTestConfigPanel.tsx index 847e5e6c4..d6c972705 100644 --- a/src/components/usage/ModelTestConfigPanel.tsx +++ b/src/components/usage/ModelTestConfigPanel.tsx @@ -47,7 +47,9 @@ export function ModelTestConfigPanel() { try { setIsSaving(true); await saveStreamCheckConfig(config); - toast.success(t("streamCheck.configSaved", "健康检查配置已保存")); + toast.success(t("streamCheck.configSaved", "健康检查配置已保存"), { + closeButton: true, + }); } catch (e) { toast.error( t("streamCheck.configSaveFailed", "保存失败") + ": " + String(e), diff --git a/src/components/usage/PricingEditModal.tsx b/src/components/usage/PricingEditModal.tsx index c759ac01d..79958ed98 100644 --- a/src/components/usage/PricingEditModal.tsx +++ b/src/components/usage/PricingEditModal.tsx @@ -76,6 +76,7 @@ export function PricingEditModal({ isNew ? t("usage.pricingAdded", "定价已添加") : t("usage.pricingUpdated", "定价已更新"), + { closeButton: true }, ); onClose(); diff --git a/src/hooks/useDragSort.ts b/src/hooks/useDragSort.ts index a2097c625..1f26573dc 100644 --- a/src/hooks/useDragSort.ts +++ b/src/hooks/useDragSort.ts @@ -87,6 +87,7 @@ export function useDragSort(providers: Record, appId: AppId) { t("provider.sortUpdated", { defaultValue: "排序已更新", }), + { closeButton: true }, ); } catch (error) { console.error("Failed to update provider sort order", error); diff --git a/src/hooks/useImportExport.ts b/src/hooks/useImportExport.ts index be9cb267a..3fa573aaa 100644 --- a/src/hooks/useImportExport.ts +++ b/src/hooks/useImportExport.ts @@ -113,6 +113,7 @@ export function useImportExport( t("settings.importSuccess", { defaultValue: "配置导入成功", }), + { closeButton: true }, ); successTimerRef.current = window.setTimeout(() => { @@ -170,6 +171,7 @@ export function useImportExport( t("settings.configExported", { defaultValue: "配置已导出", }) + `\n${displayPath}`, + { closeButton: true }, ); } else { toast.error( diff --git a/src/hooks/usePromptActions.ts b/src/hooks/usePromptActions.ts index fc6e4ecd8..538d6ee40 100644 --- a/src/hooks/usePromptActions.ts +++ b/src/hooks/usePromptActions.ts @@ -36,7 +36,7 @@ export function usePromptActions(appId: AppId) { try { await promptsApi.upsertPrompt(appId, id, prompt); await reload(); - toast.success(t("prompts.saveSuccess")); + toast.success(t("prompts.saveSuccess"), { closeButton: true }); } catch (error) { toast.error(t("prompts.saveFailed")); throw error; @@ -50,7 +50,7 @@ export function usePromptActions(appId: AppId) { try { await promptsApi.deletePrompt(appId, id); await reload(); - toast.success(t("prompts.deleteSuccess")); + toast.success(t("prompts.deleteSuccess"), { closeButton: true }); } catch (error) { toast.error(t("prompts.deleteFailed")); throw error; @@ -64,7 +64,7 @@ export function usePromptActions(appId: AppId) { try { await promptsApi.enablePrompt(appId, id); await reload(); - toast.success(t("prompts.enableSuccess")); + toast.success(t("prompts.enableSuccess"), { closeButton: true }); } catch (error) { toast.error(t("prompts.enableFailed")); throw error; @@ -104,14 +104,14 @@ export function usePromptActions(appId: AppId) { try { if (enabled) { await promptsApi.enablePrompt(appId, id); - toast.success(t("prompts.enableSuccess")); + toast.success(t("prompts.enableSuccess"), { closeButton: true }); } else { // 禁用提示词 - 需要后端支持 await promptsApi.upsertPrompt(appId, id, { ...prompts[id], enabled: false, }); - toast.success(t("prompts.disableSuccess")); + toast.success(t("prompts.disableSuccess"), { closeButton: true }); } await reload(); } catch (error) { @@ -130,7 +130,7 @@ export function usePromptActions(appId: AppId) { try { const id = await promptsApi.importFromFile(appId); await reload(); - toast.success(t("prompts.importSuccess")); + toast.success(t("prompts.importSuccess"), { closeButton: true }); return id; } catch (error) { toast.error(t("prompts.importFailed")); diff --git a/src/hooks/useProviderActions.ts b/src/hooks/useProviderActions.ts index fb683a3f3..18739ca58 100644 --- a/src/hooks/useProviderActions.ts +++ b/src/hooks/useProviderActions.ts @@ -119,6 +119,7 @@ export function useProviderActions(activeApp: AppId) { t("provider.usageSaved", { defaultValue: "用量查询配置已保存", }), + { closeButton: true }, ); } catch (error) { const detail = diff --git a/src/hooks/useProxyConfig.ts b/src/hooks/useProxyConfig.ts index 234017c22..effb64b1e 100644 --- a/src/hooks/useProxyConfig.ts +++ b/src/hooks/useProxyConfig.ts @@ -24,7 +24,7 @@ export function useProxyConfig() { mutationFn: (newConfig: ProxyConfig) => invoke("update_proxy_config", { config: newConfig }), onSuccess: () => { - toast.success("代理配置已保存"); + toast.success("代理配置已保存", { closeButton: true }); queryClient.invalidateQueries({ queryKey: ["proxyConfig"] }); queryClient.invalidateQueries({ queryKey: ["proxyStatus"] }); }, diff --git a/src/hooks/useProxyStatus.ts b/src/hooks/useProxyStatus.ts index 12d7ac6a0..1f13547b2 100644 --- a/src/hooks/useProxyStatus.ts +++ b/src/hooks/useProxyStatus.ts @@ -40,6 +40,7 @@ export function useProxyStatus() { t("proxy.startedWithTakeover", { defaultValue: `代理模式已启用 - ${info.address}:${info.port}`, }), + { closeButton: true }, ); queryClient.invalidateQueries({ queryKey: ["proxyStatus"] }); queryClient.invalidateQueries({ queryKey: ["proxyTakeoverActive"] }); @@ -62,9 +63,12 @@ export function useProxyStatus() { t("proxy.stoppedWithRestore", { defaultValue: "代理模式已关闭,配置已恢复", }), + { closeButton: true }, ); queryClient.invalidateQueries({ queryKey: ["proxyStatus"] }); queryClient.invalidateQueries({ queryKey: ["proxyTakeoverActive"] }); + // 清除所有供应商健康状态缓存(后端已清空数据库记录) + queryClient.invalidateQueries({ queryKey: ["providerHealth"] }); }, onError: (error: Error) => { const detail = extractErrorMessage(error) || "未知错误"; diff --git a/src/hooks/useSettings.ts b/src/hooks/useSettings.ts index d6dd74583..149d16075 100644 --- a/src/hooks/useSettings.ts +++ b/src/hooks/useSettings.ts @@ -309,6 +309,7 @@ export function useSettings(): UseSettingsResult { t("notifications.settingsSaved", { defaultValue: "设置已保存", }), + { closeButton: true }, ); } diff --git a/src/hooks/useStreamCheck.ts b/src/hooks/useStreamCheck.ts index 0685b6031..d180c2b9f 100644 --- a/src/hooks/useStreamCheck.ts +++ b/src/hooks/useStreamCheck.ts @@ -28,6 +28,7 @@ export function useStreamCheck(appId: AppId) { time: result.responseTimeMs, defaultValue: `${providerName} 运行正常 (${result.responseTimeMs}ms)`, }), + { closeButton: true }, ); } else if (result.status === "degraded") { toast.warning( diff --git a/src/i18n/locales/en.json b/src/i18n/locales/en.json index 92820c5ab..2226b977b 100644 --- a/src/i18n/locales/en.json +++ b/src/i18n/locales/en.json @@ -112,7 +112,7 @@ "providerAdded": "Provider added", "providerSaved": "Provider configuration saved", "providerDeleted": "Provider deleted successfully", - "switchSuccess": "Switch successful! Please restart {{appName}} terminal to take effect", + "switchSuccess": "Switch successful!", "switchFailedTitle": "Switch failed", "switchFailed": "Switch failed: {{error}}", "autoImported": "Default provider created from existing configuration", @@ -365,7 +365,10 @@ "justNow": "Just now", "minutesAgo": "{{count}} min ago", "hoursAgo": "{{count}} hr ago", - "daysAgo": "{{count}} day ago" + "daysAgo": "{{count}} day ago", + "multiplePlans": "{{count}} plans", + "expand": "Expand", + "collapse": "Collapse" }, "usageScript": { "title": "Configure Usage Query", @@ -378,6 +381,10 @@ "templateGeneral": "General", "templateNewAPI": "NewAPI", "credentialsConfig": "Credentials", + "credentialsHint": "Leave empty to use provider config", + "optional": "optional", + "apiKeyPlaceholder": "Leave empty to use provider's API Key", + "baseUrlPlaceholder": "Leave empty to use provider's base URL", "baseUrl": "Base URL", "accessToken": "Access Token", "accessTokenPlaceholder": "Generate in 'Security Settings'", @@ -534,9 +541,6 @@ "env": "Environment (one per line, KEY=VALUE)", "envPlaceholder": "FOO=bar\nHELLO=world", "reset": "Reset", - "notice": { - "restartClaude": "Written. Restart Claude to take effect." - }, "msg": { "saved": "Saved", "deleted": "Deleted", @@ -840,5 +844,50 @@ }, "agents": { "title": "Agents" + }, + "proxy": { + "failoverQueue": { + "title": "Failover Queue", + "description": "Manage failover order for each app's providers", + "info": "The current active provider always takes priority. When requests fail, the system will try other providers in queue order.", + "selectProvider": "Select a provider to add to queue", + "noAvailableProviders": "No providers available to add", + "empty": "Failover queue is empty. Add providers to enable automatic failover.", + "dragHint": "Drag providers to adjust failover order. Lower numbers have higher priority.", + "toggleEnabled": "Enable/Disable", + "addSuccess": "Added to failover queue", + "addFailed": "Failed to add", + "removeSuccess": "Removed from failover queue", + "removeFailed": "Failed to remove", + "reorderSuccess": "Queue order updated", + "reorderFailed": "Failed to update order", + "toggleFailed": "Failed to update status" + }, + "autoFailover": { + "info": "When the failover queue has multiple providers, the system will try them in priority order when requests fail. When a provider reaches the consecutive failure threshold, the circuit breaker will open and skip it temporarily.", + "configSaved": "Auto failover config saved", + "configSaveFailed": "Failed to save", + "retrySettings": "Retry & Timeout Settings", + "failureThreshold": "Failure Threshold", + "failureThresholdHint": "Open circuit breaker after this many consecutive failures (recommended: 3-10)", + "timeout": "Recovery Wait Time (seconds)", + "timeoutHint": "Wait this long before trying to recover after circuit opens (recommended: 30-120)", + "circuitBreakerSettings": "Circuit Breaker Advanced Settings", + "successThreshold": "Recovery Success Threshold", + "successThresholdHint": "Close circuit breaker after this many successes in half-open state", + "errorRate": "Error Rate Threshold (%)", + "errorRateHint": "Open circuit breaker when error rate exceeds this value", + "minRequests": "Minimum Requests", + "minRequestsHint": "Minimum requests before calculating error rate", + "explanationTitle": "How It Works", + "failureThresholdLabel": "Failure Threshold", + "failureThresholdExplain": "Circuit breaker opens after this many consecutive failures, making the provider temporarily unavailable", + "timeoutLabel": "Recovery Wait Time", + "timeoutExplain": "After circuit opens, wait this long before trying half-open state", + "successThresholdLabel": "Recovery Success Threshold", + "successThresholdExplain": "In half-open state, close circuit breaker after this many successes, making provider available again", + "errorRateLabel": "Error Rate Threshold", + "errorRateExplain": "Open circuit breaker when error rate exceeds this value, even if failure threshold not reached" + } } } diff --git a/src/i18n/locales/ja.json b/src/i18n/locales/ja.json index 0fab2caa1..07f1a32ca 100644 --- a/src/i18n/locales/ja.json +++ b/src/i18n/locales/ja.json @@ -112,7 +112,7 @@ "providerAdded": "プロバイダーを追加しました", "providerSaved": "プロバイダー設定を保存しました", "providerDeleted": "プロバイダーを削除しました", - "switchSuccess": "切り替え成功! {{appName}} ターミナルを再起動すると反映されます", + "switchSuccess": "切り替え成功!", "switchFailedTitle": "切り替えに失敗しました", "switchFailed": "切り替えに失敗しました: {{error}}", "autoImported": "既存設定からデフォルトプロバイダーを自動作成しました", @@ -365,7 +365,10 @@ "justNow": "たった今", "minutesAgo": "{{count}} 分前", "hoursAgo": "{{count}} 時間前", - "daysAgo": "{{count}} 日前" + "daysAgo": "{{count}} 日前", + "multiplePlans": "{{count}} プラン", + "expand": "展開", + "collapse": "折りたたむ" }, "usageScript": { "title": "利用状況を設定", @@ -378,6 +381,10 @@ "templateGeneral": "General", "templateNewAPI": "NewAPI", "credentialsConfig": "認証情報", + "credentialsHint": "空欄の場合はプロバイダー設定を使用", + "optional": "オプション", + "apiKeyPlaceholder": "空欄の場合はプロバイダーの API Key を使用", + "baseUrlPlaceholder": "空欄の場合はプロバイダーの Base URL を使用", "baseUrl": "Base URL", "accessToken": "Access Token", "accessTokenPlaceholder": "「Security Settings」で生成", @@ -534,9 +541,6 @@ "env": "環境変数(1 行に 1 件、KEY=VALUE)", "envPlaceholder": "FOO=bar\nHELLO=world", "reset": "リセット", - "notice": { - "restartClaude": "書き込みました。Claude を再起動すると反映されます。" - }, "msg": { "saved": "保存しました", "deleted": "削除しました", diff --git a/src/i18n/locales/zh.json b/src/i18n/locales/zh.json index 6eacfbfde..c63bbc74c 100644 --- a/src/i18n/locales/zh.json +++ b/src/i18n/locales/zh.json @@ -112,7 +112,7 @@ "providerAdded": "供应商已添加", "providerSaved": "供应商配置已保存", "providerDeleted": "供应商删除成功", - "switchSuccess": "切换成功!请重启 {{appName}} 终端以生效", + "switchSuccess": "切换成功!", "switchFailedTitle": "切换失败", "switchFailed": "切换失败:{{error}}", "autoImported": "已从现有配置创建默认供应商", @@ -365,7 +365,10 @@ "justNow": "刚刚", "minutesAgo": "{{count}} 分钟前", "hoursAgo": "{{count}} 小时前", - "daysAgo": "{{count}} 天前" + "daysAgo": "{{count}} 天前", + "multiplePlans": "{{count}} 个套餐", + "expand": "展开", + "collapse": "收起" }, "usageScript": { "title": "配置用量查询", @@ -378,6 +381,10 @@ "templateGeneral": "通用模板", "templateNewAPI": "NewAPI", "credentialsConfig": "凭证配置", + "credentialsHint": "留空则自动使用供应商配置", + "optional": "可选", + "apiKeyPlaceholder": "留空则使用供应商的 API Key", + "baseUrlPlaceholder": "留空则使用供应商的请求地址", "baseUrl": "请求地址", "accessToken": "访问令牌(在个人安全设置里获取)", "accessTokenPlaceholder": "在'安全设置'里生成", @@ -534,9 +541,6 @@ "env": "环境变量 (一行一个,KEY=VALUE)", "envPlaceholder": "FOO=bar\nHELLO=world", "reset": "重置", - "notice": { - "restartClaude": "已写入配置,重启 Claude 生效" - }, "msg": { "saved": "已保存", "deleted": "已删除", @@ -840,5 +844,50 @@ }, "agents": { "title": "智能体" + }, + "proxy": { + "failoverQueue": { + "title": "故障转移队列", + "description": "管理各应用的供应商故障转移顺序", + "info": "当前激活的供应商始终优先。当请求失败时,系统会按队列顺序依次尝试其他供应商。", + "selectProvider": "选择供应商添加到队列", + "noAvailableProviders": "没有可添加的供应商", + "empty": "故障转移队列为空。添加供应商以启用自动故障转移。", + "dragHint": "拖拽供应商可调整故障转移顺序,序号越小优先级越高。", + "toggleEnabled": "启用/禁用", + "addSuccess": "已添加到故障转移队列", + "addFailed": "添加失败", + "removeSuccess": "已从故障转移队列移除", + "removeFailed": "移除失败", + "reorderSuccess": "队列顺序已更新", + "reorderFailed": "更新顺序失败", + "toggleFailed": "状态更新失败" + }, + "autoFailover": { + "info": "当故障转移队列中配置了多个供应商时,系统会在请求失败时按优先级顺序依次尝试。当某个供应商连续失败达到阈值时,熔断器会打开并在一段时间内跳过该供应商。", + "configSaved": "自动故障转移配置已保存", + "configSaveFailed": "保存失败", + "retrySettings": "重试与超时设置", + "failureThreshold": "失败阈值", + "failureThresholdHint": "连续失败多少次后打开熔断器(建议: 3-10)", + "timeout": "恢复等待时间(秒)", + "timeoutHint": "熔断器打开后,等待多久后尝试恢复(建议: 30-120)", + "circuitBreakerSettings": "熔断器高级设置", + "successThreshold": "恢复成功阈值", + "successThresholdHint": "半开状态下成功多少次后关闭熔断器", + "errorRate": "错误率阈值 (%)", + "errorRateHint": "错误率超过此值时打开熔断器", + "minRequests": "最小请求数", + "minRequestsHint": "计算错误率前的最小请求数", + "explanationTitle": "工作原理", + "failureThresholdLabel": "失败阈值", + "failureThresholdExplain": "连续失败达到此次数时,熔断器打开,该供应商暂时不可用", + "timeoutLabel": "恢复等待时间", + "timeoutExplain": "熔断器打开后,等待此时间后尝试半开状态", + "successThresholdLabel": "恢复成功阈值", + "successThresholdExplain": "半开状态下,成功达到此次数时关闭熔断器,供应商恢复可用", + "errorRateLabel": "错误率阈值", + "errorRateExplain": "错误率超过此值时,即使未达到失败阈值也会打开熔断器" + } } } diff --git a/src/lib/api/deeplink.ts b/src/lib/api/deeplink.ts index 003c19f68..dd1d33669 100644 --- a/src/lib/api/deeplink.ts +++ b/src/lib/api/deeplink.ts @@ -38,6 +38,15 @@ export interface DeepLinkImportRequest { config?: string; configFormat?: string; configUrl?: string; + + // Usage script fields (v3.9+) + usageEnabled?: boolean; + usageScript?: string; + usageApiKey?: string; + usageBaseUrl?: string; + usageAccessToken?: string; + usageUserId?: string; + usageAutoInterval?: number; } export interface McpImportResult { diff --git a/src/lib/api/failover.ts b/src/lib/api/failover.ts index aed49451a..58fc26e21 100644 --- a/src/lib/api/failover.ts +++ b/src/lib/api/failover.ts @@ -3,6 +3,7 @@ import type { ProviderHealth, CircuitBreakerConfig, CircuitBreakerStats, + FailoverQueueItem, } from "@/types/proxy"; export interface Provider { @@ -17,23 +18,10 @@ export interface Provider { meta?: unknown; icon?: string; iconColor?: string; - isProxyTarget?: boolean; } export const failoverApi = { - // 获取代理目标列表 - async getProxyTargets(appType: string): Promise { - return invoke("get_proxy_targets", { appType }); - }, - - // 设置代理目标 - async setProxyTarget( - providerId: string, - appType: string, - enabled: boolean, - ): Promise { - return invoke("set_proxy_target", { providerId, appType, enabled }); - }, + // ========== 熔断器 API ========== // 获取供应商健康状态 async getProviderHealth( @@ -70,4 +58,50 @@ export const failoverApi = { ): Promise { return invoke("get_circuit_breaker_stats", { providerId, appType }); }, + + // ========== 故障转移队列 API(新) ========== + + // 获取故障转移队列 + async getFailoverQueue(appType: string): Promise { + return invoke("get_failover_queue", { appType }); + }, + + // 获取可添加到队列的供应商(不在队列中的) + async getAvailableProvidersForFailover(appType: string): Promise { + return invoke("get_available_providers_for_failover", { appType }); + }, + + // 添加供应商到故障转移队列 + async addToFailoverQueue(appType: string, providerId: string): Promise { + return invoke("add_to_failover_queue", { appType, providerId }); + }, + + // 从故障转移队列移除供应商 + async removeFromFailoverQueue( + appType: string, + providerId: string, + ): Promise { + return invoke("remove_from_failover_queue", { appType, providerId }); + }, + + // 重新排序故障转移队列 + async reorderFailoverQueue( + appType: string, + providerIds: string[], + ): Promise { + return invoke("reorder_failover_queue", { appType, providerIds }); + }, + + // 设置故障转移队列项的启用状态 + async setFailoverItemEnabled( + appType: string, + providerId: string, + enabled: boolean, + ): Promise { + return invoke("set_failover_item_enabled", { + appType, + providerId, + enabled, + }); + }, }; diff --git a/src/lib/api/providers.ts b/src/lib/api/providers.ts index 4da270598..14c8e115a 100644 --- a/src/lib/api/providers.ts +++ b/src/lib/api/providers.ts @@ -38,10 +38,6 @@ export const providersApi = { return await invoke("switch_provider", { id, app: appId }); }, - async setProxyTarget(id: string, appId: AppId): Promise { - return await invoke("set_proxy_target_provider", { id, app: appId }); - }, - async importDefault(appId: AppId): Promise { return await invoke("import_default_config", { app: appId }); }, diff --git a/src/lib/query/failover.ts b/src/lib/query/failover.ts index 7c9aea6c3..9f0c7d9fb 100644 --- a/src/lib/query/failover.ts +++ b/src/lib/query/failover.ts @@ -1,16 +1,7 @@ import { useQuery, useMutation, useQueryClient } from "@tanstack/react-query"; import { failoverApi } from "@/lib/api/failover"; -/** - * 获取代理目标列表 - */ -export function useProxyTargets(appType: string) { - return useQuery({ - queryKey: ["proxyTargets", appType], - queryFn: () => failoverApi.getProxyTargets(appType), - enabled: !!appType, - }); -} +// ========== 熔断器 Hooks ========== /** * 获取供应商健康状态 @@ -25,37 +16,6 @@ export function useProviderHealth(providerId: string, appType: string) { }); } -/** - * 设置代理目标 - */ -export function useSetProxyTarget() { - const queryClient = useQueryClient(); - - return useMutation({ - mutationFn: ({ - providerId, - appType, - enabled, - }: { - providerId: string; - appType: string; - enabled: boolean; - }) => failoverApi.setProxyTarget(providerId, appType, enabled), - onSuccess: (_, variables) => { - // 刷新代理目标列表 - queryClient.invalidateQueries({ - queryKey: ["proxyTargets", variables.appType], - }); - // 刷新供应商列表 - queryClient.invalidateQueries({ queryKey: ["providers"] }); - // 刷新健康状态 - queryClient.invalidateQueries({ - queryKey: ["providerHealth", variables.providerId, variables.appType], - }); - }, - }); -} - /** * 重置熔断器 */ @@ -114,3 +74,123 @@ export function useCircuitBreakerStats(providerId: string, appType: string) { refetchInterval: 5000, // 每 5 秒刷新一次 }); } + +// ========== 故障转移队列 Hooks(新) ========== + +/** + * 获取故障转移队列 + */ +export function useFailoverQueue(appType: string) { + return useQuery({ + queryKey: ["failoverQueue", appType], + queryFn: () => failoverApi.getFailoverQueue(appType), + enabled: !!appType, + }); +} + +/** + * 获取可添加到队列的供应商 + */ +export function useAvailableProvidersForFailover(appType: string) { + return useQuery({ + queryKey: ["availableProvidersForFailover", appType], + queryFn: () => failoverApi.getAvailableProvidersForFailover(appType), + enabled: !!appType, + }); +} + +/** + * 添加供应商到故障转移队列 + */ +export function useAddToFailoverQueue() { + const queryClient = useQueryClient(); + + return useMutation({ + mutationFn: ({ + appType, + providerId, + }: { + appType: string; + providerId: string; + }) => failoverApi.addToFailoverQueue(appType, providerId), + onSuccess: (_, variables) => { + queryClient.invalidateQueries({ + queryKey: ["failoverQueue", variables.appType], + }); + queryClient.invalidateQueries({ + queryKey: ["availableProvidersForFailover", variables.appType], + }); + }, + }); +} + +/** + * 从故障转移队列移除供应商 + */ +export function useRemoveFromFailoverQueue() { + const queryClient = useQueryClient(); + + return useMutation({ + mutationFn: ({ + appType, + providerId, + }: { + appType: string; + providerId: string; + }) => failoverApi.removeFromFailoverQueue(appType, providerId), + onSuccess: (_, variables) => { + queryClient.invalidateQueries({ + queryKey: ["failoverQueue", variables.appType], + }); + queryClient.invalidateQueries({ + queryKey: ["availableProvidersForFailover", variables.appType], + }); + }, + }); +} + +/** + * 重新排序故障转移队列 + */ +export function useReorderFailoverQueue() { + const queryClient = useQueryClient(); + + return useMutation({ + mutationFn: ({ + appType, + providerIds, + }: { + appType: string; + providerIds: string[]; + }) => failoverApi.reorderFailoverQueue(appType, providerIds), + onSuccess: (_, variables) => { + queryClient.invalidateQueries({ + queryKey: ["failoverQueue", variables.appType], + }); + }, + }); +} + +/** + * 设置故障转移队列项的启用状态 + */ +export function useSetFailoverItemEnabled() { + const queryClient = useQueryClient(); + + return useMutation({ + mutationFn: ({ + appType, + providerId, + enabled, + }: { + appType: string; + providerId: string; + enabled: boolean; + }) => failoverApi.setFailoverItemEnabled(appType, providerId, enabled), + onSuccess: (_, variables) => { + queryClient.invalidateQueries({ + queryKey: ["failoverQueue", variables.appType], + }); + }, + }); +} diff --git a/src/lib/query/mutations.ts b/src/lib/query/mutations.ts index 00f0b54ea..b5118680c 100644 --- a/src/lib/query/mutations.ts +++ b/src/lib/query/mutations.ts @@ -180,34 +180,6 @@ export const useSwitchProviderMutation = (appId: AppId) => { }); }; -export const useSetProxyTargetMutation = (appId: AppId) => { - const queryClient = useQueryClient(); - const { t } = useTranslation(); - - return useMutation({ - mutationFn: async (providerId: string) => { - return await providersApi.setProxyTarget(providerId, appId); - }, - onSuccess: async () => { - await queryClient.invalidateQueries({ queryKey: ["providers", appId] }); - toast.success( - t("notifications.proxyTargetSet", { - defaultValue: "已设置代理目标", - }), - ); - }, - onError: (error: Error) => { - const detail = extractErrorMessage(error) || t("common.unknown"); - toast.error( - t("notifications.setProxyTargetFailed", { - defaultValue: "设置代理目标失败: {{error}}", - error: detail, - }), - ); - }, - }); -}; - export const useSaveSettingsMutation = () => { const queryClient = useQueryClient(); diff --git a/src/types.ts b/src/types.ts index 570d7a2d6..bdf46fcf9 100644 --- a/src/types.ts +++ b/src/types.ts @@ -23,8 +23,6 @@ export interface Provider { // 图标配置 icon?: string; // 图标名称(如 "openai", "anthropic") iconColor?: string; // 图标颜色(Hex 格式,如 "#00A67E") - // 新增:是否为代理目标 - isProxyTarget?: boolean; } export interface AppConfig { diff --git a/src/types/proxy.ts b/src/types/proxy.ts index 133295c2e..393e0e662 100644 --- a/src/types/proxy.ts +++ b/src/types/proxy.ts @@ -93,3 +93,12 @@ export interface ProxyUsageRecord { error: string | null; timestamp: string; } + +// 故障转移队列条目 +export interface FailoverQueueItem { + providerId: string; + providerName: string; + queueOrder: number; + enabled: boolean; + createdAt: number; +} diff --git a/src/utils/formatters.ts b/src/utils/formatters.ts index 0f5a52c62..04e23a951 100644 --- a/src/utils/formatters.ts +++ b/src/utils/formatters.ts @@ -1,6 +1,6 @@ /** * 格式化 JSON 字符串 - * @param value - 原始 JSON 字符串(支持带键名包装的格式) + * @param value - 原始 JSON 字符串 * @returns 格式化后的 JSON 字符串(2 空格缩进) * @throws 如果 JSON 格式无效 */ @@ -9,9 +9,8 @@ export function formatJSON(value: string): string { if (!trimmed) { return ""; } - // 使用智能解析器来处理可能的片段格式 - const result = parseSmartMcpJson(trimmed); - return result.formattedConfig; + const parsed = JSON.parse(trimmed); + return JSON.stringify(parsed, null, 2); } /** diff --git a/tests/components/SettingsDialog.test.tsx b/tests/components/SettingsDialog.test.tsx index d27827e7c..22b33144f 100644 --- a/tests/components/SettingsDialog.test.tsx +++ b/tests/components/SettingsDialog.test.tsx @@ -19,6 +19,23 @@ vi.mock("react-i18next", () => ({ useTranslation: () => ({ t: tMock }), })); +vi.mock("@/hooks/useProxyStatus", () => ({ + useProxyStatus: () => ({ + status: null, + isLoading: false, + isRunning: false, + isTakeoverActive: false, + startWithTakeover: vi.fn(), + stopWithRestore: vi.fn(), + switchProxyProvider: vi.fn(), + checkRunning: vi.fn(), + checkTakeoverActive: vi.fn(), + isStarting: false, + isStopping: false, + isPending: false, + }), +})); + interface SettingsMock { settings: any; isLoading: boolean; @@ -288,6 +305,7 @@ describe("SettingsPage Component", () => { }); fireEvent.click(screen.getByText("settings.tabAdvanced")); + fireEvent.click(screen.getByText("数据管理")); // 有文件时,点击导入按钮执行 importConfig fireEvent.click( @@ -363,6 +381,7 @@ describe("SettingsPage Component", () => { await waitFor(() => { expect(toastSuccessMock).toHaveBeenCalledWith( "settings.devModeRestartHint", + expect.objectContaining({ closeButton: true }), ); }); }); @@ -393,6 +412,7 @@ describe("SettingsPage Component", () => { render(); fireEvent.click(screen.getByText("settings.tabAdvanced")); + fireEvent.click(screen.getByText("配置文件目录")); fireEvent.click(screen.getByText("browse-directory")); expect(settingsMock.browseDirectory).toHaveBeenCalledWith("claude"); diff --git a/tests/hooks/useImportExport.extra.test.tsx b/tests/hooks/useImportExport.extra.test.tsx index aa5dd3c69..049794132 100644 --- a/tests/hooks/useImportExport.extra.test.tsx +++ b/tests/hooks/useImportExport.extra.test.tsx @@ -120,6 +120,7 @@ describe("useImportExport Hook (edge cases)", () => { expect(exportConfigMock).toHaveBeenCalledWith("/exports/config.json"); expect(toastSuccessMock).toHaveBeenCalledWith( expect.stringContaining("/final/config.json"), + expect.objectContaining({ closeButton: true }), ); }); }); diff --git a/tests/hooks/useImportExport.test.tsx b/tests/hooks/useImportExport.test.tsx index 4b2cfa0aa..030c228aa 100644 --- a/tests/hooks/useImportExport.test.tsx +++ b/tests/hooks/useImportExport.test.tsx @@ -180,6 +180,7 @@ describe("useImportExport Hook", () => { expect(exportConfigMock).toHaveBeenCalledWith("/export.json"); expect(toastSuccessMock).toHaveBeenCalledWith( expect.stringContaining("/backup/export.json"), + expect.objectContaining({ closeButton: true }), ); }); diff --git a/tests/integration/SettingsDialog.test.tsx b/tests/integration/SettingsDialog.test.tsx index 9b015b361..0ebd3eb09 100644 --- a/tests/integration/SettingsDialog.test.tsx +++ b/tests/integration/SettingsDialog.test.tsx @@ -150,6 +150,7 @@ describe("SettingsPage integration", () => { expect(screen.getByText("language:zh")).toBeInTheDocument(), ); fireEvent.click(screen.getByText("settings.tabAdvanced")); + fireEvent.click(screen.getByText("配置文件目录")); const appInput = await screen.findByPlaceholderText( "settings.browsePlaceholderApp", ); @@ -165,6 +166,7 @@ describe("SettingsPage integration", () => { ); fireEvent.click(screen.getByText("settings.tabAdvanced")); + fireEvent.click(screen.getByText("数据管理")); fireEvent.click(screen.getByText("settings.selectConfigFile")); await waitFor(() => expect(screen.getByTestId("selected-file").textContent).toContain( @@ -188,6 +190,7 @@ describe("SettingsPage integration", () => { ); fireEvent.click(screen.getByText("settings.tabAdvanced")); + fireEvent.click(screen.getByText("配置文件目录")); const appInput = await screen.findByPlaceholderText( "settings.browsePlaceholderApp", ); @@ -214,6 +217,7 @@ describe("SettingsPage integration", () => { ); fireEvent.click(screen.getByText("settings.tabAdvanced")); + fireEvent.click(screen.getByText("配置文件目录")); const browseButtons = screen.getAllByTitle("settings.browseDirectory"); const resetButtons = screen.getAllByTitle("settings.resetDefault"); @@ -253,6 +257,7 @@ describe("SettingsPage integration", () => { expect(screen.getByText("language:zh")).toBeInTheDocument(), ); fireEvent.click(screen.getByText("settings.tabAdvanced")); + fireEvent.click(screen.getByText("数据管理")); server.use( http.post("http://tauri.local/save_file_dialog", () => diff --git a/tests/msw/handlers.ts b/tests/msw/handlers.ts index 4672fe9e3..95937534b 100644 --- a/tests/msw/handlers.ts +++ b/tests/msw/handlers.ts @@ -236,4 +236,64 @@ export const handlers = [ http.post(`${TAURI_ENDPOINT}/sync_current_providers_live`, () => success({ success: true }), ), + + // Proxy status (for SettingsPage / ProxyPanel hooks) + http.post(`${TAURI_ENDPOINT}/get_proxy_status`, () => + success({ + running: false, + address: "127.0.0.1", + port: 0, + active_connections: 0, + total_requests: 0, + success_requests: 0, + failed_requests: 0, + success_rate: 0, + uptime_seconds: 0, + current_provider: null, + current_provider_id: null, + last_request_at: null, + last_error: null, + failover_count: 0, + active_targets: [], + }), + ), + + http.post(`${TAURI_ENDPOINT}/is_live_takeover_active`, () => success(false)), + + // Failover / circuit breaker defaults + http.post(`${TAURI_ENDPOINT}/get_failover_queue`, () => success([])), + http.post(`${TAURI_ENDPOINT}/get_available_providers_for_failover`, () => + success([]), + ), + http.post(`${TAURI_ENDPOINT}/add_to_failover_queue`, () => success(true)), + http.post(`${TAURI_ENDPOINT}/remove_from_failover_queue`, () => success(true)), + http.post(`${TAURI_ENDPOINT}/reorder_failover_queue`, () => success(true)), + http.post(`${TAURI_ENDPOINT}/set_failover_item_enabled`, () => success(true)), + + http.post(`${TAURI_ENDPOINT}/get_circuit_breaker_config`, () => + success({ + failureThreshold: 3, + successThreshold: 2, + timeoutSeconds: 60, + errorRateThreshold: 50, + minRequests: 5, + }), + ), + http.post(`${TAURI_ENDPOINT}/update_circuit_breaker_config`, () => + success(true), + ), + http.post(`${TAURI_ENDPOINT}/get_provider_health`, () => + success({ + provider_id: "mock-provider", + app_type: "claude", + is_healthy: true, + consecutive_failures: 0, + last_success_at: null, + last_failure_at: null, + last_error: null, + updated_at: new Date().toISOString(), + }), + ), + http.post(`${TAURI_ENDPOINT}/reset_circuit_breaker`, () => success(true)), + http.post(`${TAURI_ENDPOINT}/get_circuit_breaker_stats`, () => success(null)), ];