refactor(failover): make auto-failover toggle per-app independent

- Change setting key from 'auto_failover_enabled' to 'auto_failover_enabled_{app_type}'
- Update provider_router to check per-app failover setting
- When failover disabled, use current provider only; when enabled, use queue order
- Add unit tests for failover enabled/disabled behavior
This commit is contained in:
YoVinchen
2025-12-22 00:33:34 +08:00
parent c231c840a7
commit 59096d9c15
4 changed files with 177 additions and 91 deletions
+19 -12
View File
@@ -82,28 +82,35 @@ pub async fn set_failover_item_enabled(
.set_failover_item_enabled(&app_type, &provider_id, enabled)
.map_err(|e| e.to_string())
}
/// 获取自动故障转移总开关状态
/// 获取指定应用的自动故障转移开关状态
#[tauri::command]
pub async fn get_auto_failover_enabled(state: tauri::State<'_, AppState>) -> Result<bool, String> {
pub async fn get_auto_failover_enabled(
state: tauri::State<'_, AppState>,
app_type: String,
) -> Result<bool, String> {
let key = format!("auto_failover_enabled_{app_type}");
state
.db
.get_setting("auto_failover_enabled")
.get_setting(&key)
.map(|v| v.map(|s| s == "true").unwrap_or(false)) // 默认关闭
.map_err(|e| e.to_string())
}
/// 设置自动故障转移开关状态
/// 设置指定应用的自动故障转移开关状态
///
/// 注意:关闭故障转移时不会清除队列,队列内容会保留供下次开启时使用
#[tauri::command]
pub async fn set_auto_failover_enabled(
state: tauri::State<'_, AppState>,
app_type: String,
enabled: bool,
) -> Result<(), String> {
state
.db
.set_setting(
"auto_failover_enabled",
if enabled { "true" } else { "false" },
)
.map_err(|e| e.to_string())
let key = format!("auto_failover_enabled_{app_type}");
let value = if enabled { "true" } else { "false" };
log::info!(
"[Failover] Setting auto_failover_enabled: key='{key}', value='{value}', app_type='{app_type}'"
);
state.db.set_setting(&key, value).map_err(|e| e.to_string())
}
+117 -57
View File
@@ -30,54 +30,46 @@ impl ProviderRouter {
/// 选择可用的供应商(支持故障转移)
///
/// 返回按优先级排序的可用供应商列表:
/// 1. 当前供应商(is_current=true)始终第一位
/// 2. 故障转移队列中的其他供应商(按 queue_order 排序)- 仅当自动故障转移开关开启时
/// 3. 只返回熔断器未打开的供应商
/// - 故障转移关闭时:仅返回当前供应商
/// - 故障转移开启时:完全按照故障转移队列顺序返回,忽略当前供应商设置
pub async fn select_providers(&self, app_type: &str) -> Result<Vec<Provider>, AppError> {
let mut result = Vec::new();
let all_providers = self.db.get_all_providers(app_type)?;
// 检查自动故障转移开关是否开启
let auto_failover_enabled = self
.db
.get_setting("auto_failover_enabled")
.map(|v| v.map(|s| s == "true").unwrap_or(false)) // 默认关闭
.unwrap_or(false);
// 1. 当前供应商始终第一位
if let Some(current_id) = self.db.get_current_provider(app_type)? {
if let Some(current) = all_providers.get(&current_id) {
let circuit_key = format!("{}:{}", app_type, current.id);
let breaker = self.get_or_create_circuit_breaker(&circuit_key).await;
if breaker.is_available().await {
log::info!(
"[{}] Current provider available: {} ({})",
app_type,
current.name,
current.id
);
result.push(current.clone());
} else {
log::warn!(
"[{}] Current provider {} circuit breaker open, checking failover queue",
app_type,
current.name
);
}
// 检查该应用的自动故障转移开关是否开启
let failover_key = format!("auto_failover_enabled_{app_type}");
let auto_failover_enabled = match self.db.get_setting(&failover_key) {
Ok(Some(value)) => {
let enabled = value == "true";
log::info!(
"[{app_type}] Failover setting '{failover_key}' = '{value}', enabled: {enabled}"
);
enabled
}
}
Ok(None) => {
log::warn!(
"[{app_type}] Failover setting '{failover_key}' not found in database, defaulting to disabled"
);
false
}
Err(e) => {
log::error!(
"[{app_type}] Failed to read failover setting '{failover_key}': {e}, defaulting to disabled"
);
false
}
};
// 2. 获取故障转移队列中的供应商(仅当自动故障转移开关开启时)
if auto_failover_enabled {
// 故障转移开启:完全按照队列顺序,忽略当前供应商
let queue = self.db.get_failover_queue(app_type)?;
log::info!(
"[{}] Failover enabled, using queue order ({} items)",
app_type,
queue.len()
);
for item in queue {
// 跳过已添加的当前供应商
if result.iter().any(|p| p.id == item.provider_id) {
continue;
}
// 跳过禁用的队列项
if !item.enabled {
continue;
@@ -91,7 +83,7 @@ impl ProviderRouter {
if breaker.is_available().await {
log::info!(
"[{}] Failover provider available: {} ({}) at queue position {}",
"[{}] Queue provider available: {} ({}) at position {}",
app_type,
provider.name,
provider.id,
@@ -100,7 +92,7 @@ impl ProviderRouter {
result.push(provider.clone());
} else {
log::debug!(
"[{}] Failover provider {} circuit breaker open, skipping",
"[{}] Queue provider {} circuit breaker open, skipping",
app_type,
provider.name
);
@@ -108,7 +100,31 @@ impl ProviderRouter {
}
}
} else {
log::info!("[{app_type}] Auto-failover disabled, using only current provider");
// 故障转移关闭:仅使用当前供应商
log::info!("[{app_type}] Failover disabled, using current provider only");
if let Some(current_id) = self.db.get_current_provider(app_type)? {
if let Some(current) = all_providers.get(&current_id) {
let circuit_key = format!("{}:{}", app_type, current.id);
let breaker = self.get_or_create_circuit_breaker(&circuit_key).await;
if breaker.is_available().await {
log::info!(
"[{}] Current provider available: {} ({})",
app_type,
current.name,
current.id
);
result.push(current.clone());
} else {
log::warn!(
"[{}] Current provider {} circuit breaker open",
app_type,
current.name
);
}
}
}
}
if result.is_empty() {
@@ -118,7 +134,7 @@ impl ProviderRouter {
}
log::info!(
"[{}] Failover chain: {} provider(s) available",
"[{}] Provider chain: {} provider(s) available",
app_type,
result.len()
);
@@ -275,25 +291,14 @@ mod tests {
let db = Arc::new(Database::memory().unwrap());
let router = ProviderRouter::new(db);
// 测试创建熔断器
let breaker = router.get_or_create_circuit_breaker("claude:test").await;
assert!(breaker.allow_request().await.allowed);
}
#[tokio::test]
async fn select_providers_does_not_consume_half_open_permit() {
async fn test_failover_disabled_uses_current_provider() {
let db = Arc::new(Database::memory().unwrap());
// 配置:让熔断器 Open 后立刻进入 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 =
@@ -304,19 +309,74 @@ mod tests {
db.set_current_provider("claude", "a").unwrap();
db.add_to_failover_queue("claude", "b").unwrap();
let router = ProviderRouter::new(db.clone());
let providers = router.select_providers("claude").await.unwrap();
assert_eq!(providers.len(), 1);
assert_eq!(providers[0].id, "a");
}
#[tokio::test]
async fn test_failover_enabled_uses_queue_order() {
let db = Arc::new(Database::memory().unwrap());
let provider_a =
Provider::with_id("a".to_string(), "Provider A".to_string(), json!({}), None);
let provider_b =
Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None);
db.save_provider("claude", &provider_a).unwrap();
db.save_provider("claude", &provider_b).unwrap();
db.set_current_provider("claude", "a").unwrap();
db.add_to_failover_queue("claude", "b").unwrap();
db.add_to_failover_queue("claude", "a").unwrap();
db.set_setting("auto_failover_enabled_claude", "true")
.unwrap();
let router = ProviderRouter::new(db.clone());
let providers = router.select_providers("claude").await.unwrap();
assert_eq!(providers.len(), 2);
assert_eq!(providers[0].id, "b");
assert_eq!(providers[1].id, "a");
}
#[tokio::test]
async fn test_select_providers_does_not_consume_half_open_permit() {
let db = Arc::new(Database::memory().unwrap());
db.update_circuit_breaker_config(&CircuitBreakerConfig {
failure_threshold: 1,
timeout_seconds: 0,
..Default::default()
})
.await
.unwrap();
let provider_a =
Provider::with_id("a".to_string(), "Provider A".to_string(), json!({}), None);
let provider_b =
Provider::with_id("b".to_string(), "Provider B".to_string(), json!({}), None);
db.save_provider("claude", &provider_a).unwrap();
db.save_provider("claude", &provider_b).unwrap();
db.add_to_failover_queue("claude", "a").unwrap();
db.add_to_failover_queue("claude", "b").unwrap();
db.set_setting("auto_failover_enabled_claude", "true")
.unwrap();
let router = ProviderRouter::new(db.clone());
// 让 B 进入 Open 状态(failure_threshold=1
router
.record_result("b", "claude", false, false, Some("fail".to_string()))
.await
.unwrap();
// select_providers 只做“可用性判断”,不应占用 HalfOpen 探测名额
let providers = router.select_providers("claude").await.unwrap();
assert_eq!(providers.len(), 2);
// 如果 select_providers 错误地消耗了 HalfOpen 名额,这里会返回 false(被限流拒绝)
assert!(router.allow_provider_request("b", "claude").await.allowed);
}
}