Merge remote-tracking branch 'origin/main' into feature/error-request-logging

# Conflicts:
#	src-tauri/src/proxy/handlers.rs
#	src-tauri/src/proxy/mod.rs
#	src-tauri/src/proxy/provider_router.rs
#	src-tauri/src/services/proxy.rs
#	src/components/providers/ProviderActions.tsx
This commit is contained in:
YoVinchen
2025-12-16 17:40:49 +08:00
82 changed files with 3861 additions and 2008 deletions
+84
View File
@@ -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<Vec<FailoverQueueItem>, 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<Vec<Provider>, 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<String>,
) -> 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())
}
+2
View File
@@ -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::*;
-13
View File
@@ -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<bool, String> {
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<bool, AppError> {
ProviderService::import_default_config(state, app_type)
}
+17 -47
View File
@@ -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<Vec<Provider>, 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(())
}
/// 获取熔断器统计信息(仅当代理服务器运行时)
+21 -2
View File
@@ -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<HashSet<String>> = 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)
+248
View File
@@ -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<Vec<FailoverQueueItem>, 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::<Result<Vec<_>, _>>()
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(items)
}
/// 获取故障转移队列中的供应商(完整 Provider 信息,按顺序)
pub fn get_failover_providers(&self, app_type: &str) -> Result<Vec<Provider>, 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<i32> = 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<bool, AppError> {
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<Vec<Provider>, 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<Provider> = all_providers
.into_values()
.filter(|p| !queue_ids.contains(&p.id))
.collect();
Ok(available)
}
}
+3
View File
@@ -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;
+11 -189
View File
@@ -17,7 +17,7 @@ impl Database {
) -> Result<IndexMap<String, Provider>, 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<String> = row.get(8)?;
let icon_color: Option<String> = 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<Option<Provider>, 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<String> = row.get(7)?;
let icon_color: Option<String> = 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<Option<String>, 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<bool> = 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<IndexMap<String, Provider>, 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<String> = row.get(3)?;
let category: Option<String> = row.get(4)?;
let created_at: Option<i64> = row.get(5)?;
let sort_index: Option<usize> = row.get(6)?;
let notes: Option<String> = row.get(7)?;
let icon: Option<String> = row.get(8)?;
let icon_color: Option<String> = 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<i64> = 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<Vec<(String, String, String)>, 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::<Result<Vec<_>, _>>()
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(targets)
}
/// 更新供应商的 settings_config(仅更新配置,不改变其他字段)
pub fn update_provider_settings_config(
&self,
+36 -5
View File
@@ -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<String>,
) -> 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<String>,
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 ====================
/// 获取熔断器配置
+3
View File
@@ -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;
+24 -7
View File
@@ -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(())
}
-1
View File
@@ -245,7 +245,6 @@ fn dry_run_validates_schema_compatibility() {
meta: None,
icon: None,
icon_color: None,
is_proxy_target: Some(false),
},
);
+23
View File
@@ -113,4 +113,27 @@ pub struct DeepLinkImportRequest {
/// Remote config URL
#[serde(skip_serializing_if = "Option::is_none")]
pub config_url: Option<String>,
// ============ 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<bool>,
/// Base64 encoded usage query script code
#[serde(skip_serializing_if = "Option::is_none")]
pub usage_script: Option<String>,
/// Usage query API key (if different from provider API key)
#[serde(skip_serializing_if = "Option::is_none")]
pub usage_api_key: Option<String>,
/// Usage query base URL (if different from provider endpoint)
#[serde(skip_serializing_if = "Option::is_none")]
pub usage_base_url: Option<String>,
/// Usage query access token (for NewAPI template)
#[serde(skip_serializing_if = "Option::is_none")]
pub usage_access_token: Option<String>,
/// Usage query user ID (for NewAPI template)
#[serde(skip_serializing_if = "Option::is_none")]
pub usage_user_id: Option<String>,
/// Auto query interval in minutes (0 to disable)
#[serde(skip_serializing_if = "Option::is_none")]
pub usage_auto_interval: Option<u64>,
}
+41
View File
@@ -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::<bool>().ok());
// Extract usage script fields (v3.9+)
let usage_enabled = params
.get("usageEnabled")
.and_then(|v| v.parse::<bool>().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::<u64>().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,
})
}
+56 -3
View File
@@ -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<Option<ProviderMeta>, 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();
+28
View File
@@ -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();
+42 -6
View File
@@ -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::<AppState>();
// 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);
});
-5
View File
@@ -36,10 +36,6 @@ pub struct Provider {
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "iconColor")]
pub icon_color: Option<String>,
/// 是否为代理目标(数据库专用字段,不写入配置文件)
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "isProxyTarget")]
pub is_proxy_target: Option<bool>,
}
impl Provider {
@@ -62,7 +58,6 @@ impl Provider {
meta: None,
icon: None,
icon_color: None,
is_proxy_target: None,
}
}
}
+104 -18
View File
@@ -72,8 +72,10 @@ pub struct CircuitBreaker {
failed_requests: Arc<AtomicU32>,
/// 上次打开时间
last_opened_at: Arc<RwLock<Option<Instant>>>,
/// 配置
config: CircuitBreakerConfig,
/// 配置(支持热更新)
config: Arc<RwLock<CircuitBreakerConfig>>,
/// 半开状态已放行的请求数(用于限流)
half_open_requests: Arc<AtomicU32>,
}
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);
}
/// 转换到关闭状态
+148
View File
@@ -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<RwLock<HashSet<String>>>,
db: Arc<Database>,
}
impl FailoverSwitchManager {
pub fn new(db: Arc<Database>) -> 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<bool, AppError> {
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<bool, AppError> {
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::<crate::store::AppState>() {
// 更新 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)
}
}
+144 -29
View File
@@ -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<NewProviderRouter>,
#[allow(dead_code)]
/// 共享的 ProviderRouter(持有熔断器状态)
router: Arc<ProviderRouter>,
/// 单个 Provider 内的最大重试次数
max_retries: u8,
status: Arc<RwLock<ProxyStatus>>,
current_providers: Arc<RwLock<std::collections::HashMap<String, (String, String)>>>,
/// 故障转移切换管理器
failover_manager: Arc<FailoverSwitchManager>,
/// AppHandle,用于发射事件和更新托盘
app_handle: Option<tauri::AppHandle>,
}
impl RequestForwarder {
pub fn new(
db: Arc<Database>,
router: Arc<ProviderRouter>,
timeout_secs: u64,
max_retries: u8,
status: Arc<RwLock<ProxyStatus>>,
current_providers: Arc<RwLock<std::collections::HashMap<String, (String, String)>>>,
failover_manager: Arc<FailoverSwitchManager>,
app_handle: Option<tauri::AppHandle>,
) -> 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<Response, ProxyError> {
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<Provider>,
) -> Result<Response, ProxyError> {
// 获取适配器
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,
}
}
+164
View File
@@ -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<TokenUsage>;
pub type ResponseUsageParser = fn(&Value) -> Option<TokenUsage>;
/// 模型提取器类型别名
/// 参数: (流式事件列表, 请求中的模型名称) -> 最终使用的模型名称
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,
};
+151
View File
@@ -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<Provider>,
/// 请求中的模型名称
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<Self, ProxyError> {
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<Provider> {
self.providers.clone()
}
/// 计算请求延迟(毫秒)
#[inline]
pub fn latency_ms(&self) -> u64 {
self.start_time.elapsed().as_millis() as u64
}
}
File diff suppressed because it is too large Load Diff
+4 -1
View File
@@ -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;
+167 -56
View File
@@ -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<Vec<Provider>, 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(&current_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(&current_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<String>,
) -> 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 后立刻进入 HalfOpentimeout_seconds=0),并用 1 次失败就打开熔断器
db.update_circuit_breaker_config(&CircuitBreakerConfig {
failure_threshold: 1,
timeout_seconds: 0,
..Default::default()
})
.await
.unwrap();
// 准备 2 个 ProviderA(当前)+ 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);
}
}
-1
View File
@@ -253,7 +253,6 @@ mod tests {
meta: None,
icon: None,
icon_color: None,
is_proxy_target: None,
}
}
-1
View File
@@ -174,7 +174,6 @@ mod tests {
meta: None,
icon: None,
icon_color: None,
is_proxy_target: None,
}
}
+7 -12
View File
@@ -120,11 +120,7 @@ impl GeminiAdapter {
/// 从 Provider 配置中提取原始 API Key
fn extract_key_raw(&self, provider: &Provider) -> Option<String> {
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!(
+3 -4
View File
@@ -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\"}"
}
}));
@@ -394,7 +394,6 @@ mod tests {
meta: None,
icon: None,
icon_color: None,
is_proxy_target: None,
}
}
+411
View File
@@ -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<Response, ProxyError> {
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::<Value>(&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<Response, ProxyError> {
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<dyn Fn(Vec<Value>, Option<u64>) + Send + Sync + 'static>;
/// SSE 使用量收集器
#[derive(Clone)]
pub struct SseUsageCollector {
inner: Arc<SseUsageCollectorInner>,
}
struct SseUsageCollectorInner {
events: Mutex<Vec<Value>>,
first_event_time: Mutex<Option<std::time::Instant>>,
start_time: std::time::Instant,
on_complete: UsageCallbackWithTiming,
finished: AtomicBool,
}
impl SseUsageCollector {
/// 创建新的使用量收集器
pub fn new(
start_time: std::time::Instant,
callback: impl Fn(Vec<Value>, Option<u64>) + 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<u64>,
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<Item = Result<Bytes, std::io::Error>> + Send + 'static,
tag: &'static str,
usage_collector: Option<SseUsageCollector>,
) -> impl Stream<Item = Result<Bytes, std::io::Error>> + 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::<Value>(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;
}
}
}
-70
View File
@@ -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<Database>,
}
impl ProviderRouter {
pub fn new(db: Arc<Database>) -> Self {
Self { db }
}
/// 选择Provider(只使用标记为代理目标的 Provider)
pub async fn select_provider(
&self,
app_type: &AppType,
_failed_ids: &[String],
) -> Result<Provider, ProxyError> {
// 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<String>,
) {
// 不再记录健康状态
}
}
+62 -5
View File
@@ -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<RwLock<Option<std::time::Instant>>>,
/// 每个应用类型当前使用的 provider (app_type -> (provider_id, provider_name))
pub current_providers: Arc<RwLock<std::collections::HashMap<String, (String, String)>>>,
/// 共享的 ProviderRouter(持有熔断器状态,跨请求保持)
pub provider_router: Arc<ProviderRouter>,
/// AppHandle,用于发射事件和更新托盘菜单
pub app_handle: Option<tauri::AppHandle>,
/// 故障转移切换管理器
pub failover_manager: Arc<FailoverSwitchManager>,
}
/// 代理HTTP服务器
@@ -29,22 +39,37 @@ pub struct ProxyServer {
config: ProxyConfig,
state: ProxyState,
shutdown_tx: Arc<RwLock<Option<oneshot::Sender<()>>>>,
/// 服务器任务句柄,用于等待服务器实际关闭
server_handle: Arc<RwLock<Option<JoinHandle<()>>>>,
}
impl ProxyServer {
pub fn new(config: ProxyConfig, db: Arc<Database>) -> Self {
pub fn new(
config: ProxyConfig,
db: Arc<Database>,
app_handle: Option<tauri::AppHandle>,
) -> 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;
}
}
+14 -16
View File
@@ -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)
+45 -3
View File
@@ -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<String> {
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<String> {
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(),
)
+177 -50
View File
@@ -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<Database>,
server: Arc<RwLock<Option<ProxyServer>>>,
/// AppHandle,用于传递给 ProxyServer 以支持故障转移时的 UI 更新
app_handle: Arc<RwLock<Option<tauri::AppHandle>>>,
}
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<ProxyServerInfo, String> {
// 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<ProxyServerInfo, String> {
// 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_urlauth.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::<DocumentMut>()
.unwrap_or_else(|_| DocumentMut::new());
doc["base_url"] = toml_edit::value(new_url);
doc.to_string()
}
fn read_claude_live(&self) -> Result<Value, String> {
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(())
}
}
+13 -5
View File
@@ -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");
+9 -5
View File
@@ -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<AppState, Box<dyn std::error::Error>> {
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<AppState, Box<dyn std::error::Error>> {
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 })
}