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);
}
}
+9 -6
View File
@@ -105,13 +105,16 @@ export const failoverApi = {
});
},
// 获取自动故障转移开关状态
async getAutoFailoverEnabled(): Promise<boolean> {
return invoke("get_auto_failover_enabled");
// 获取指定应用的自动故障转移开关状态
async getAutoFailoverEnabled(appType: string): Promise<boolean> {
return invoke("get_auto_failover_enabled", { appType });
},
// 设置自动故障转移开关状态
async setAutoFailoverEnabled(enabled: boolean): Promise<void> {
return invoke("set_auto_failover_enabled", { enabled });
// 设置指定应用的自动故障转移开关状态
async setAutoFailoverEnabled(
appType: string,
enabled: boolean,
): Promise<void> {
return invoke("set_auto_failover_enabled", { appType, enabled });
},
};
+32 -16
View File
@@ -18,6 +18,9 @@ export function useProviderHealth(providerId: string, appType: string) {
/**
* 重置熔断器
*
* 重置后后端会检查是否应该切回优先级更高的供应商,
* 因此需要同时刷新供应商列表和代理状态。
*/
export function useResetCircuitBreaker() {
const queryClient = useQueryClient();
@@ -35,6 +38,14 @@ export function useResetCircuitBreaker() {
queryClient.invalidateQueries({
queryKey: ["providerHealth", variables.providerId, variables.appType],
});
// 刷新供应商列表(因为可能发生了自动恢复切换)
queryClient.invalidateQueries({
queryKey: ["providers", variables.appType],
});
// 刷新代理状态(更新 active_targets
queryClient.invalidateQueries({
queryKey: ["proxyStatus"],
});
},
});
}
@@ -236,55 +247,60 @@ export function useSetFailoverItemEnabled() {
});
}
// ========== 自动故障转移开关 Hooks ==========
// ========== 自动故障转移开关 Hooks ==========
/**
* 获取自动故障转移开关状态
* 获取指定应用的自动故障转移开关状态
*/
export function useAutoFailoverEnabled() {
export function useAutoFailoverEnabled(appType: string) {
return useQuery({
queryKey: ["autoFailoverEnabled"],
queryFn: () => failoverApi.getAutoFailoverEnabled(),
queryKey: ["autoFailoverEnabled", appType],
queryFn: () => failoverApi.getAutoFailoverEnabled(appType),
// 默认值为 false(与后端保持一致)
placeholderData: false,
});
}
/**
* 设置自动故障转移开关状态
* 设置指定应用的自动故障转移开关状态
*/
export function useSetAutoFailoverEnabled() {
const queryClient = useQueryClient();
return useMutation({
mutationFn: (enabled: boolean) =>
failoverApi.setAutoFailoverEnabled(enabled),
mutationFn: ({ appType, enabled }: { appType: string; enabled: boolean }) =>
failoverApi.setAutoFailoverEnabled(appType, enabled),
// 乐观更新
onMutate: async (enabled) => {
await queryClient.cancelQueries({ queryKey: ["autoFailoverEnabled"] });
onMutate: async ({ appType, enabled }) => {
await queryClient.cancelQueries({
queryKey: ["autoFailoverEnabled", appType],
});
const previousValue = queryClient.getQueryData<boolean>([
"autoFailoverEnabled",
appType,
]);
queryClient.setQueryData(["autoFailoverEnabled"], enabled);
queryClient.setQueryData(["autoFailoverEnabled", appType], enabled);
return { previousValue };
return { previousValue, appType };
},
// 错误时回滚
onError: (_error, _enabled, context) => {
onError: (_error, _variables, context) => {
if (context?.previousValue !== undefined) {
queryClient.setQueryData(
["autoFailoverEnabled"],
["autoFailoverEnabled", context.appType],
context.previousValue,
);
}
},
// 无论成功失败,都重新获取
onSettled: () => {
queryClient.invalidateQueries({ queryKey: ["autoFailoverEnabled"] });
onSettled: (_, __, variables) => {
queryClient.invalidateQueries({
queryKey: ["autoFailoverEnabled", variables.appType],
});
},
});
}