mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
refactor: 迁移 credential 模块(balancer/quota/sync)到独立 proxycast-credential crate
- 创建 proxycast-credential crate,包含 balancer、quota、sync 三个模块 - 主 crate credential/mod.rs 改为 re-export 层 - 37 个 crate 单元测试 + 31 个主 crate 属性测试全部通过
This commit is contained in:
Generated
+19
@@ -6671,6 +6671,7 @@ dependencies = [
|
||||
"proptest",
|
||||
"proxycast-config",
|
||||
"proxycast-core",
|
||||
"proxycast-credential",
|
||||
"proxycast-infra",
|
||||
"proxycast-providers",
|
||||
"proxycast-services",
|
||||
@@ -6774,6 +6775,24 @@ dependencies = [
|
||||
"zip",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-credential"
|
||||
version = "0.60.0"
|
||||
dependencies = [
|
||||
"axum 0.7.9",
|
||||
"chrono",
|
||||
"dashmap 5.5.3",
|
||||
"proptest",
|
||||
"proxycast-core",
|
||||
"proxycast-infra",
|
||||
"reqwest 0.12.28",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tempfile",
|
||||
"tokio",
|
||||
"tracing",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "proxycast-infra"
|
||||
version = "0.60.0"
|
||||
|
||||
@@ -17,6 +17,7 @@ proxycast-infra = { path = "crates/infra" }
|
||||
proxycast-providers = { path = "crates/providers" }
|
||||
proxycast-services = { path = "crates/services" }
|
||||
proxycast-terminal = { path = "crates/terminal" }
|
||||
proxycast-credential = { path = "crates/credential" }
|
||||
proxycast-websocket = { path = "crates/websocket" }
|
||||
voice-core = { path = "crates/voice-core" }
|
||||
|
||||
@@ -194,6 +195,7 @@ proxycast-infra.workspace = true
|
||||
proxycast-providers.workspace = true
|
||||
proxycast-services.workspace = true
|
||||
proxycast-terminal.workspace = true
|
||||
proxycast-credential.workspace = true
|
||||
proxycast-websocket.workspace = true
|
||||
voice-core.workspace = true
|
||||
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
[package]
|
||||
name = "proxycast-credential"
|
||||
version.workspace = true
|
||||
edition.workspace = true
|
||||
authors.workspace = true
|
||||
|
||||
[dependencies]
|
||||
proxycast-core.workspace = true
|
||||
proxycast-infra.workspace = true
|
||||
|
||||
# 序列化
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
|
||||
# 异步运行时
|
||||
tokio.workspace = true
|
||||
|
||||
# 日志
|
||||
tracing.workspace = true
|
||||
|
||||
# HTTP 服务器(AllCredentialsExhaustedError 的 IntoResponse)
|
||||
axum.workspace = true
|
||||
|
||||
# HTTP 客户端
|
||||
reqwest.workspace = true
|
||||
|
||||
# 时间
|
||||
chrono.workspace = true
|
||||
|
||||
# 并发
|
||||
dashmap.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
proptest.workspace = true
|
||||
tempfile.workspace = true
|
||||
+11
-138
@@ -2,13 +2,13 @@
|
||||
//!
|
||||
//! 提供轮询负载均衡策略,支持凭证冷却和自动恢复
|
||||
|
||||
use crate::proxy::ProxyClientFactory;
|
||||
use crate::ProviderType;
|
||||
use chrono::{DateTime, Duration, Utc};
|
||||
use dashmap::DashMap;
|
||||
use proxycast_core::credential::health::{HealthCheckConfig, HealthChecker};
|
||||
use proxycast_core::credential::pool::{CredentialPool, PoolError};
|
||||
use proxycast_core::credential::types::Credential;
|
||||
use proxycast_core::ProviderType;
|
||||
use proxycast_infra::ProxyClientFactory;
|
||||
use reqwest::Client;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
@@ -143,16 +143,9 @@ impl LoadBalancer {
|
||||
}
|
||||
|
||||
/// 选择下一个可用凭证(使用当前策略)
|
||||
///
|
||||
/// # 错误
|
||||
/// - 如果 Provider 未注册,返回 `PoolError::EmptyPool`
|
||||
/// - 如果没有可用凭证,返回 `PoolError::NoAvailableCredential`
|
||||
pub fn select(&self, provider: ProviderType) -> Result<Credential, PoolError> {
|
||||
let pool = self.pools.get(&provider).ok_or(PoolError::EmptyPool)?;
|
||||
|
||||
// 先刷新冷却状态
|
||||
pool.refresh_cooldowns();
|
||||
|
||||
match self.strategy {
|
||||
BalanceStrategy::RoundRobin => self.select_round_robin(&pool, provider),
|
||||
BalanceStrategy::LeastUsed => self.select_least_used(&pool),
|
||||
@@ -161,39 +154,19 @@ impl LoadBalancer {
|
||||
}
|
||||
|
||||
/// 选择下一个可用凭证并创建配置了代理的 HTTP 客户端
|
||||
///
|
||||
/// 代理选择逻辑:
|
||||
/// 1. 如果凭证有 proxy_url,使用 Per-Key 代理
|
||||
/// 2. 否则,使用全局代理(如果配置了)
|
||||
/// 3. 否则,不使用代理
|
||||
///
|
||||
/// # 错误
|
||||
/// - 如果 Provider 未注册,返回 `PoolError::EmptyPool`
|
||||
/// - 如果没有可用凭证,返回 `PoolError::NoAvailableCredential`
|
||||
/// - 如果代理配置无效,返回 `PoolError::CredentialNotFound`(包含错误信息)
|
||||
pub fn select_with_client(
|
||||
&self,
|
||||
provider: ProviderType,
|
||||
) -> Result<CredentialSelection, PoolError> {
|
||||
let credential = self.select(provider)?;
|
||||
|
||||
// 使用凭证的 proxy_url 或回退到全局代理
|
||||
let client = self
|
||||
.proxy_factory
|
||||
.create_client(credential.proxy_url())
|
||||
.map_err(|e| PoolError::CredentialNotFound(format!("代理配置错误: {e}")))?;
|
||||
|
||||
Ok(CredentialSelection { credential, client })
|
||||
}
|
||||
|
||||
/// 为指定凭证创建配置了代理的 HTTP 客户端
|
||||
///
|
||||
/// # 参数
|
||||
/// - `credential`: 凭证引用
|
||||
///
|
||||
/// # 返回
|
||||
/// - `Ok(Client)`: 配置了代理的 HTTP 客户端
|
||||
/// - `Err(PoolError)`: 代理配置错误
|
||||
pub fn create_client_for_credential(
|
||||
&self,
|
||||
credential: &Credential,
|
||||
@@ -204,17 +177,6 @@ impl LoadBalancer {
|
||||
}
|
||||
|
||||
/// 选择下一个可用凭证,支持代理失败时的故障转移
|
||||
///
|
||||
/// 当代理连接失败时,自动尝试下一个可用凭证。
|
||||
/// 最多尝试 `max_attempts` 次(默认为池中凭证数量)。
|
||||
///
|
||||
/// # 参数
|
||||
/// - `provider`: Provider 类型
|
||||
/// - `max_attempts`: 最大尝试次数(None 表示尝试所有可用凭证)
|
||||
///
|
||||
/// # 返回
|
||||
/// - `Ok(CredentialSelection)`: 成功选择的凭证和客户端
|
||||
/// - `Err(PoolError)`: 所有凭证都失败
|
||||
pub fn select_with_failover(
|
||||
&self,
|
||||
provider: ProviderType,
|
||||
@@ -233,7 +195,6 @@ impl LoadBalancer {
|
||||
let mut tried_ids = std::collections::HashSet::new();
|
||||
|
||||
for _ in 0..attempts {
|
||||
// 选择下一个凭证
|
||||
let credential = match self.select(provider) {
|
||||
Ok(cred) => cred,
|
||||
Err(e) => {
|
||||
@@ -242,19 +203,16 @@ impl LoadBalancer {
|
||||
}
|
||||
};
|
||||
|
||||
// 避免重复尝试同一个凭证
|
||||
if tried_ids.contains(&credential.id) {
|
||||
continue;
|
||||
}
|
||||
tried_ids.insert(credential.id.clone());
|
||||
|
||||
// 尝试创建客户端
|
||||
match self.proxy_factory.create_client(credential.proxy_url()) {
|
||||
Ok(client) => {
|
||||
return Ok(CredentialSelection { credential, client });
|
||||
}
|
||||
Err(e) => {
|
||||
// 记录警告并继续尝试下一个凭证
|
||||
tracing::warn!(
|
||||
credential_id = %credential.id,
|
||||
proxy_url = ?credential.proxy_url(),
|
||||
@@ -273,33 +231,17 @@ impl LoadBalancer {
|
||||
}
|
||||
|
||||
/// 报告代理连接失败并尝试故障转移
|
||||
///
|
||||
/// 当代理连接失败时调用此方法,它会:
|
||||
/// 1. 记录失败
|
||||
/// 2. 尝试选择下一个可用凭证
|
||||
///
|
||||
/// # 参数
|
||||
/// - `provider`: Provider 类型
|
||||
/// - `failed_credential_id`: 失败的凭证 ID
|
||||
///
|
||||
/// # 返回
|
||||
/// - `Ok(CredentialSelection)`: 故障转移成功,返回新的凭证和客户端
|
||||
/// - `Err(PoolError)`: 故障转移失败
|
||||
pub fn failover_on_proxy_error(
|
||||
&self,
|
||||
provider: ProviderType,
|
||||
failed_credential_id: &str,
|
||||
) -> Result<CredentialSelection, PoolError> {
|
||||
// 记录失败
|
||||
let _ = self.report(provider, failed_credential_id, false, 0);
|
||||
|
||||
tracing::warn!(
|
||||
credential_id = %failed_credential_id,
|
||||
provider = %provider,
|
||||
"代理连接失败,执行故障转移"
|
||||
);
|
||||
|
||||
// 尝试选择下一个凭证
|
||||
self.select_with_client(provider)
|
||||
}
|
||||
|
||||
@@ -309,7 +251,6 @@ impl LoadBalancer {
|
||||
pool: &CredentialPool,
|
||||
provider: ProviderType,
|
||||
) -> Result<Credential, PoolError> {
|
||||
// 收集所有活跃凭证
|
||||
let active_creds: Vec<Credential> = pool
|
||||
.all()
|
||||
.into_iter()
|
||||
@@ -320,13 +261,11 @@ impl LoadBalancer {
|
||||
return Err(PoolError::NoAvailableCredential);
|
||||
}
|
||||
|
||||
// 获取或创建轮询索引
|
||||
let index_entry = self
|
||||
.round_robin_indices
|
||||
.entry(provider)
|
||||
.or_insert_with(|| AtomicUsize::new(0));
|
||||
|
||||
// 原子递增并取模
|
||||
let index = index_entry.fetch_add(1, Ordering::SeqCst) % active_creds.len();
|
||||
Ok(active_creds[index].clone())
|
||||
}
|
||||
@@ -352,21 +291,12 @@ impl LoadBalancer {
|
||||
return Err(PoolError::NoAvailableCredential);
|
||||
}
|
||||
|
||||
// 使用简单的伪随机(基于时间戳)
|
||||
let now = Utc::now().timestamp_nanos_opt().unwrap_or(0) as usize;
|
||||
let index = now % active_creds.len();
|
||||
Ok(active_creds[index].clone())
|
||||
}
|
||||
|
||||
/// 标记凭证为冷却状态
|
||||
///
|
||||
/// # 参数
|
||||
/// - `provider`: Provider 类型
|
||||
/// - `credential_id`: 凭证 ID
|
||||
/// - `duration`: 冷却时长
|
||||
///
|
||||
/// # 错误
|
||||
/// - 如果 Provider 未注册或凭证不存在
|
||||
pub fn mark_cooldown(
|
||||
&self,
|
||||
provider: ProviderType,
|
||||
@@ -374,7 +304,6 @@ impl LoadBalancer {
|
||||
duration: Duration,
|
||||
) -> Result<(), PoolError> {
|
||||
let pool = self.pools.get(&provider).ok_or(PoolError::EmptyPool)?;
|
||||
|
||||
pool.mark_cooldown(credential_id, duration)
|
||||
}
|
||||
|
||||
@@ -385,7 +314,6 @@ impl LoadBalancer {
|
||||
credential_id: &str,
|
||||
) -> Result<(), PoolError> {
|
||||
let pool = self.pools.get(&provider).ok_or(PoolError::EmptyPool)?;
|
||||
|
||||
pool.mark_active(credential_id)
|
||||
}
|
||||
|
||||
@@ -397,20 +325,6 @@ impl LoadBalancer {
|
||||
}
|
||||
|
||||
/// 报告凭证使用结果
|
||||
///
|
||||
/// 自动更新健康状态:
|
||||
/// - 成功时:如果凭证之前不健康,恢复为健康
|
||||
/// - 失败时:如果连续失败达到阈值(默认 3 次),标记为不健康
|
||||
///
|
||||
/// # 参数
|
||||
/// - `provider`: Provider 类型
|
||||
/// - `credential_id`: 凭证 ID
|
||||
/// - `success`: 是否成功
|
||||
/// - `latency_ms`: 延迟(毫秒)
|
||||
///
|
||||
/// # 返回
|
||||
/// - `Ok(true)` 如果健康状态发生变化
|
||||
/// - `Ok(false)` 如果健康状态未变化
|
||||
pub fn report(
|
||||
&self,
|
||||
provider: ProviderType,
|
||||
@@ -419,7 +333,6 @@ impl LoadBalancer {
|
||||
latency_ms: u64,
|
||||
) -> Result<bool, PoolError> {
|
||||
let pool = self.pools.get(&provider).ok_or(PoolError::EmptyPool)?;
|
||||
|
||||
if success {
|
||||
self.health_checker
|
||||
.record_success(&pool, credential_id, latency_ms)
|
||||
@@ -461,7 +374,7 @@ impl Default for LoadBalancer {
|
||||
#[cfg(test)]
|
||||
mod balancer_tests {
|
||||
use super::*;
|
||||
use crate::credential::CredentialData;
|
||||
use proxycast_core::credential::types::CredentialData;
|
||||
|
||||
fn create_test_credential(id: &str, provider: ProviderType) -> Credential {
|
||||
Credential::new(
|
||||
@@ -487,9 +400,7 @@ mod balancer_tests {
|
||||
let pool = Arc::new(CredentialPool::new(ProviderType::Kiro));
|
||||
pool.add(create_test_credential("cred-1", ProviderType::Kiro))
|
||||
.unwrap();
|
||||
|
||||
lb.register_pool(pool.clone());
|
||||
|
||||
assert!(lb.providers().contains(&ProviderType::Kiro));
|
||||
assert!(lb.get_pool(ProviderType::Kiro).is_some());
|
||||
}
|
||||
@@ -498,22 +409,17 @@ mod balancer_tests {
|
||||
fn test_load_balancer_select_round_robin() {
|
||||
let lb = LoadBalancer::round_robin();
|
||||
let pool = Arc::new(CredentialPool::new(ProviderType::Kiro));
|
||||
|
||||
pool.add(create_test_credential("cred-1", ProviderType::Kiro))
|
||||
.unwrap();
|
||||
pool.add(create_test_credential("cred-2", ProviderType::Kiro))
|
||||
.unwrap();
|
||||
pool.add(create_test_credential("cred-3", ProviderType::Kiro))
|
||||
.unwrap();
|
||||
|
||||
lb.register_pool(pool);
|
||||
|
||||
// 轮询应该依次返回不同的凭证
|
||||
let c1 = lb.select(ProviderType::Kiro).unwrap();
|
||||
let c2 = lb.select(ProviderType::Kiro).unwrap();
|
||||
let c3 = lb.select(ProviderType::Kiro).unwrap();
|
||||
|
||||
// 三次选择应该返回三个不同的凭证
|
||||
let ids: std::collections::HashSet<_> = [c1.id, c2.id, c3.id].into_iter().collect();
|
||||
assert_eq!(ids.len(), 3);
|
||||
}
|
||||
@@ -529,22 +435,16 @@ mod balancer_tests {
|
||||
fn test_load_balancer_cooldown() {
|
||||
let lb = LoadBalancer::round_robin();
|
||||
let pool = Arc::new(CredentialPool::new(ProviderType::Kiro));
|
||||
|
||||
pool.add(create_test_credential("cred-1", ProviderType::Kiro))
|
||||
.unwrap();
|
||||
pool.add(create_test_credential("cred-2", ProviderType::Kiro))
|
||||
.unwrap();
|
||||
|
||||
lb.register_pool(pool);
|
||||
|
||||
// 标记 cred-1 为冷却
|
||||
lb.mark_cooldown(ProviderType::Kiro, "cred-1", Duration::hours(1))
|
||||
.unwrap();
|
||||
|
||||
// 现在只有 cred-2 可用
|
||||
assert_eq!(lb.active_count(ProviderType::Kiro), 1);
|
||||
|
||||
// 选择应该只返回 cred-2
|
||||
let selected = lb.select(ProviderType::Kiro).unwrap();
|
||||
assert_eq!(selected.id, "cred-2");
|
||||
}
|
||||
@@ -553,24 +453,18 @@ mod balancer_tests {
|
||||
fn test_load_balancer_report() {
|
||||
let lb = LoadBalancer::round_robin();
|
||||
let pool = Arc::new(CredentialPool::new(ProviderType::Kiro));
|
||||
|
||||
pool.add(create_test_credential("cred-1", ProviderType::Kiro))
|
||||
.unwrap();
|
||||
|
||||
lb.register_pool(pool.clone());
|
||||
|
||||
// 报告成功
|
||||
let changed = lb.report(ProviderType::Kiro, "cred-1", true, 100).unwrap();
|
||||
assert!(!changed); // 状态未变化
|
||||
|
||||
assert!(!changed);
|
||||
let cred = pool.get("cred-1").unwrap();
|
||||
assert_eq!(cred.stats.total_requests, 1);
|
||||
assert_eq!(cred.stats.successful_requests, 1);
|
||||
|
||||
// 报告失败
|
||||
let changed = lb.report(ProviderType::Kiro, "cred-1", false, 0).unwrap();
|
||||
assert!(!changed); // 第一次失败,状态未变化
|
||||
|
||||
assert!(!changed);
|
||||
let cred = pool.get("cred-1").unwrap();
|
||||
assert_eq!(cred.stats.total_requests, 2);
|
||||
assert_eq!(cred.stats.consecutive_failures, 1);
|
||||
@@ -578,20 +472,17 @@ mod balancer_tests {
|
||||
|
||||
#[test]
|
||||
fn test_load_balancer_auto_unhealthy() {
|
||||
use crate::credential::CredentialStatus;
|
||||
use proxycast_core::credential::types::CredentialStatus;
|
||||
|
||||
let lb = LoadBalancer::round_robin();
|
||||
let pool = Arc::new(CredentialPool::new(ProviderType::Kiro));
|
||||
|
||||
pool.add(create_test_credential("cred-1", ProviderType::Kiro))
|
||||
.unwrap();
|
||||
|
||||
lb.register_pool(pool.clone());
|
||||
|
||||
// 连续 3 次失败应标记为不健康
|
||||
assert!(!lb.report(ProviderType::Kiro, "cred-1", false, 0).unwrap());
|
||||
assert!(!lb.report(ProviderType::Kiro, "cred-1", false, 0).unwrap());
|
||||
assert!(lb.report(ProviderType::Kiro, "cred-1", false, 0).unwrap()); // 第 3 次应返回 true
|
||||
assert!(lb.report(ProviderType::Kiro, "cred-1", false, 0).unwrap());
|
||||
|
||||
let cred = pool.get("cred-1").unwrap();
|
||||
assert!(matches!(cred.status, CredentialStatus::Unhealthy { .. }));
|
||||
@@ -599,20 +490,15 @@ mod balancer_tests {
|
||||
|
||||
#[test]
|
||||
fn test_load_balancer_auto_recovery() {
|
||||
use crate::credential::CredentialStatus;
|
||||
use proxycast_core::credential::types::CredentialStatus;
|
||||
|
||||
let lb = LoadBalancer::round_robin();
|
||||
let pool = Arc::new(CredentialPool::new(ProviderType::Kiro));
|
||||
|
||||
pool.add(create_test_credential("cred-1", ProviderType::Kiro))
|
||||
.unwrap();
|
||||
|
||||
lb.register_pool(pool.clone());
|
||||
|
||||
// 先标记为不健康
|
||||
pool.mark_unhealthy("cred-1", "test".to_string()).unwrap();
|
||||
|
||||
// 成功后应恢复
|
||||
let recovered = lb.report(ProviderType::Kiro, "cred-1", true, 100).unwrap();
|
||||
assert!(recovered);
|
||||
|
||||
@@ -624,37 +510,30 @@ mod balancer_tests {
|
||||
fn test_load_balancer_cooldown_recovery() {
|
||||
let lb = LoadBalancer::round_robin();
|
||||
let pool = Arc::new(CredentialPool::new(ProviderType::Kiro));
|
||||
|
||||
pool.add(create_test_credential("cred-1", ProviderType::Kiro))
|
||||
.unwrap();
|
||||
|
||||
lb.register_pool(pool.clone());
|
||||
|
||||
// 标记为已过期的冷却(使用负时长模拟过期)
|
||||
// 直接设置状态为过去的时间
|
||||
{
|
||||
let mut entry = pool.credentials.get_mut("cred-1").unwrap();
|
||||
entry.status = crate::credential::CredentialStatus::Cooldown {
|
||||
entry.status = proxycast_core::credential::types::CredentialStatus::Cooldown {
|
||||
until: Utc::now() - Duration::seconds(1),
|
||||
};
|
||||
}
|
||||
|
||||
// 此时凭证应该处于冷却状态
|
||||
let cred = pool.get("cred-1").unwrap();
|
||||
assert!(matches!(
|
||||
cred.status,
|
||||
crate::credential::CredentialStatus::Cooldown { .. }
|
||||
proxycast_core::credential::types::CredentialStatus::Cooldown { .. }
|
||||
));
|
||||
|
||||
// 调用 select 会触发 refresh_cooldowns,应该自动恢复
|
||||
let selected = lb.select(ProviderType::Kiro).unwrap();
|
||||
assert_eq!(selected.id, "cred-1");
|
||||
|
||||
// 验证状态已恢复为 Active
|
||||
let cred = pool.get("cred-1").unwrap();
|
||||
assert!(matches!(
|
||||
cred.status,
|
||||
crate::credential::CredentialStatus::Active
|
||||
proxycast_core::credential::types::CredentialStatus::Active
|
||||
));
|
||||
}
|
||||
|
||||
@@ -662,28 +541,22 @@ mod balancer_tests {
|
||||
fn test_load_balancer_earliest_recovery() {
|
||||
let lb = LoadBalancer::round_robin();
|
||||
let pool = Arc::new(CredentialPool::new(ProviderType::Kiro));
|
||||
|
||||
pool.add(create_test_credential("cred-1", ProviderType::Kiro))
|
||||
.unwrap();
|
||||
pool.add(create_test_credential("cred-2", ProviderType::Kiro))
|
||||
.unwrap();
|
||||
|
||||
lb.register_pool(pool);
|
||||
|
||||
// 没有冷却时应返回 None
|
||||
assert!(lb.earliest_recovery(ProviderType::Kiro).is_none());
|
||||
|
||||
// 标记两个凭证为不同的冷却时间
|
||||
lb.mark_cooldown(ProviderType::Kiro, "cred-1", Duration::hours(2))
|
||||
.unwrap();
|
||||
lb.mark_cooldown(ProviderType::Kiro, "cred-2", Duration::hours(1))
|
||||
.unwrap();
|
||||
|
||||
// 应返回最早的恢复时间(cred-2 的 1 小时后)
|
||||
let recovery = lb.earliest_recovery(ProviderType::Kiro);
|
||||
assert!(recovery.is_some());
|
||||
|
||||
// 验证恢复时间大约在 1 小时后(允许几秒误差)
|
||||
let expected = Utc::now() + Duration::hours(1);
|
||||
let diff = (recovery.unwrap() - expected).num_seconds().abs();
|
||||
assert!(
|
||||
@@ -0,0 +1,21 @@
|
||||
//! 凭证池管理 crate
|
||||
//!
|
||||
//! 提供负载均衡、配额管理和凭证同步功能
|
||||
//!
|
||||
//! ## 模块结构
|
||||
//!
|
||||
//! - `balancer` - 负载均衡策略(轮询、最少使用、随机)
|
||||
//! - `quota` - 配额超限检测、自动切换和冷却恢复
|
||||
//! - `sync` - 凭证与 YAML 配置文件的同步
|
||||
|
||||
mod balancer;
|
||||
mod quota;
|
||||
mod sync;
|
||||
|
||||
// 重新导出
|
||||
pub use balancer::{BalanceStrategy, CooldownInfo, CredentialSelection, LoadBalancer};
|
||||
pub use quota::{
|
||||
create_shared_quota_manager, start_quota_cleanup_task, AllCredentialsExhaustedError,
|
||||
QuotaAutoSwitchResult, QuotaExceededRecord, QuotaManager,
|
||||
};
|
||||
pub use sync::{CredentialSyncService, SyncError};
|
||||
@@ -2,10 +2,10 @@
|
||||
//!
|
||||
//! 提供配额超限检测、自动切换和冷却恢复功能
|
||||
|
||||
use crate::config::QuotaExceededConfig;
|
||||
use crate::resilience::{QUOTA_EXCEEDED_KEYWORDS, QUOTA_EXCEEDED_STATUS_CODES};
|
||||
use chrono::{DateTime, Duration, Utc};
|
||||
use dashmap::DashMap;
|
||||
use proxycast_core::config::QuotaExceededConfig;
|
||||
use proxycast_infra::resilience::{QUOTA_EXCEEDED_KEYWORDS, QUOTA_EXCEEDED_STATUS_CODES};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -23,12 +23,6 @@ pub struct QuotaExceededRecord {
|
||||
}
|
||||
|
||||
/// 配额管理器
|
||||
///
|
||||
/// 管理凭证的配额超限状态,支持:
|
||||
/// - 标记凭证为配额超限
|
||||
/// - 检查凭证是否可用
|
||||
/// - 自动清理过期的冷却状态
|
||||
/// - 预览模型回退
|
||||
#[derive(Debug)]
|
||||
pub struct QuotaManager {
|
||||
/// 配额超限配置
|
||||
@@ -67,13 +61,6 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 标记凭证为配额超限
|
||||
///
|
||||
/// # 参数
|
||||
/// - `credential_id`: 凭证 ID
|
||||
/// - `reason`: 超限原因
|
||||
///
|
||||
/// # 返回
|
||||
/// 配额超限记录
|
||||
pub fn mark_quota_exceeded(&self, credential_id: &str, reason: &str) -> QuotaExceededRecord {
|
||||
let now = Utc::now();
|
||||
let cooldown_until = now + self.cooldown_duration();
|
||||
@@ -99,20 +86,12 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 检查凭证是否可用(未超限或已过冷却期)
|
||||
///
|
||||
/// # 参数
|
||||
/// - `credential_id`: 凭证 ID
|
||||
///
|
||||
/// # 返回
|
||||
/// - `true`: 凭证可用
|
||||
/// - `false`: 凭证处于冷却期
|
||||
pub fn is_available(&self, credential_id: &str) -> bool {
|
||||
match self.exceeded_credentials.get(credential_id) {
|
||||
Some(record) => {
|
||||
let now = Utc::now();
|
||||
if now >= record.cooldown_until {
|
||||
// 冷却期已过,移除记录
|
||||
drop(record); // 释放读锁
|
||||
drop(record);
|
||||
self.exceeded_credentials.remove(credential_id);
|
||||
true
|
||||
} else {
|
||||
@@ -124,19 +103,19 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 获取凭证的冷却结束时间
|
||||
///
|
||||
/// # 参数
|
||||
/// - `credential_id`: 凭证 ID
|
||||
///
|
||||
/// # 返回
|
||||
/// - `Some(DateTime)`: 冷却结束时间
|
||||
/// - `None`: 凭证未处于冷却期
|
||||
pub fn get_cooldown_until(&self, credential_id: &str) -> Option<DateTime<Utc>> {
|
||||
self.exceeded_credentials
|
||||
.get(credential_id)
|
||||
.map(|r| r.cooldown_until)
|
||||
}
|
||||
|
||||
/// 设置凭证的冷却结束时间(用于测试)
|
||||
pub fn set_cooldown_until(&self, credential_id: &str, until: DateTime<Utc>) {
|
||||
if let Some(mut record) = self.exceeded_credentials.get_mut(credential_id) {
|
||||
record.cooldown_until = until;
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取凭证的超限记录
|
||||
pub fn get_record(&self, credential_id: &str) -> Option<QuotaExceededRecord> {
|
||||
self.exceeded_credentials
|
||||
@@ -145,14 +124,10 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 清理过期的冷却记录
|
||||
///
|
||||
/// # 返回
|
||||
/// 清理的记录数量
|
||||
pub fn cleanup_expired(&self) -> usize {
|
||||
let now = Utc::now();
|
||||
let mut cleaned = 0;
|
||||
|
||||
// 收集需要移除的 ID
|
||||
let expired_ids: Vec<String> = self
|
||||
.exceeded_credentials
|
||||
.iter()
|
||||
@@ -160,7 +135,6 @@ impl QuotaManager {
|
||||
.map(|r| r.credential_id.clone())
|
||||
.collect();
|
||||
|
||||
// 移除过期记录
|
||||
for id in expired_ids {
|
||||
self.exceeded_credentials.remove(&id);
|
||||
cleaned += 1;
|
||||
@@ -175,13 +149,6 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 手动恢复凭证(移除冷却状态)
|
||||
///
|
||||
/// # 参数
|
||||
/// - `credential_id`: 凭证 ID
|
||||
///
|
||||
/// # 返回
|
||||
/// - `true`: 成功移除冷却状态
|
||||
/// - `false`: 凭证未处于冷却期
|
||||
pub fn restore_credential(&self, credential_id: &str) -> bool {
|
||||
self.exceeded_credentials.remove(credential_id).is_some()
|
||||
}
|
||||
@@ -208,23 +175,13 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 检查是否为配额超限错误
|
||||
///
|
||||
/// # 参数
|
||||
/// - `status_code`: HTTP 状态码
|
||||
/// - `error_message`: 错误消息
|
||||
///
|
||||
/// # 返回
|
||||
/// - `true`: 是配额超限错误
|
||||
/// - `false`: 不是配额超限错误
|
||||
pub fn is_quota_exceeded_error(status_code: Option<u16>, error_message: &str) -> bool {
|
||||
// 检查状态码
|
||||
if let Some(code) = status_code {
|
||||
if QUOTA_EXCEEDED_STATUS_CODES.contains(&code) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
// 检查错误消息中的关键词
|
||||
let error_lower = error_message.to_lowercase();
|
||||
for keyword in QUOTA_EXCEEDED_KEYWORDS {
|
||||
if error_lower.contains(keyword) {
|
||||
@@ -236,65 +193,26 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 获取预览模型名称
|
||||
///
|
||||
/// 将模型名称映射到预览版本,例如:
|
||||
/// - `gemini-2.5-pro` → `gemini-2.5-pro-preview`
|
||||
/// - `claude-3-opus` → `claude-3-opus-preview`
|
||||
/// - `gpt-4` → `gpt-4-preview`
|
||||
///
|
||||
/// 特殊映射:
|
||||
/// - `gemini-2.5-pro` → `gemini-2.5-pro-preview-05-06` (如果存在特定日期版本)
|
||||
///
|
||||
/// # 参数
|
||||
/// - `model`: 原始模型名称
|
||||
///
|
||||
/// # 返回
|
||||
/// - `Some(String)`: 预览模型名称
|
||||
/// - `None`: 无法生成预览模型名称(已经是预览版本或功能禁用)
|
||||
pub fn get_preview_model(&self, model: &str) -> Option<String> {
|
||||
if !self.config.switch_preview_model {
|
||||
return None;
|
||||
}
|
||||
|
||||
// 如果已经是预览版本,返回 None
|
||||
if Self::is_preview_model(model) {
|
||||
return None;
|
||||
}
|
||||
|
||||
// 添加 -preview 后缀
|
||||
Some(format!("{model}-preview"))
|
||||
}
|
||||
|
||||
/// 检查模型是否为预览版本
|
||||
///
|
||||
/// # 参数
|
||||
/// - `model`: 模型名称
|
||||
///
|
||||
/// # 返回
|
||||
/// - `true`: 是预览版本
|
||||
/// - `false`: 不是预览版本
|
||||
pub fn is_preview_model(model: &str) -> bool {
|
||||
model.ends_with("-preview") || model.contains("-preview-")
|
||||
}
|
||||
|
||||
/// 获取原始模型名称(从预览版本)
|
||||
///
|
||||
/// 将预览模型名称映射回原始版本,例如:
|
||||
/// - `gemini-2.5-pro-preview` → `gemini-2.5-pro`
|
||||
/// - `gemini-2.5-pro-preview-05-06` → `gemini-2.5-pro`
|
||||
///
|
||||
/// # 参数
|
||||
/// - `model`: 预览模型名称
|
||||
///
|
||||
/// # 返回
|
||||
/// - `Some(String)`: 原始模型名称
|
||||
/// - `None`: 不是预览版本
|
||||
pub fn get_original_model(model: &str) -> Option<String> {
|
||||
if !Self::is_preview_model(model) {
|
||||
return None;
|
||||
}
|
||||
|
||||
// 移除 -preview 后缀或 -preview-xxx 部分
|
||||
model.find("-preview").map(|pos| model[..pos].to_string())
|
||||
}
|
||||
|
||||
@@ -309,10 +227,6 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 获取最早的恢复时间
|
||||
///
|
||||
/// # 返回
|
||||
/// - `Some(DateTime)`: 最早的冷却结束时间
|
||||
/// - `None`: 没有凭证处于冷却期
|
||||
pub fn earliest_recovery(&self) -> Option<DateTime<Utc>> {
|
||||
self.exceeded_credentials
|
||||
.iter()
|
||||
@@ -321,13 +235,6 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 获取剩余冷却时间(秒)
|
||||
///
|
||||
/// # 参数
|
||||
/// - `credential_id`: 凭证 ID
|
||||
///
|
||||
/// # 返回
|
||||
/// - `Some(i64)`: 剩余冷却秒数(如果为负数则表示已过期)
|
||||
/// - `None`: 凭证未处于冷却期
|
||||
pub fn remaining_cooldown_seconds(&self, credential_id: &str) -> Option<i64> {
|
||||
self.exceeded_credentials.get(credential_id).map(|r| {
|
||||
let now = Utc::now();
|
||||
@@ -348,15 +255,6 @@ pub fn create_shared_quota_manager(config: QuotaExceededConfig) -> Arc<QuotaMana
|
||||
}
|
||||
|
||||
/// 启动配额管理器的定期清理任务
|
||||
///
|
||||
/// 在后台定期清理过期的配额超限记录
|
||||
///
|
||||
/// # 参数
|
||||
/// - `manager`: 共享的配额管理器
|
||||
/// - `interval_secs`: 清理间隔(秒)
|
||||
///
|
||||
/// # 返回
|
||||
/// 取消句柄(drop 时停止清理任务)
|
||||
pub fn start_quota_cleanup_task(
|
||||
manager: Arc<QuotaManager>,
|
||||
interval_secs: u64,
|
||||
@@ -389,7 +287,6 @@ pub struct QuotaAutoSwitchResult {
|
||||
}
|
||||
|
||||
impl QuotaAutoSwitchResult {
|
||||
/// 创建成功切换的结果
|
||||
pub fn switched(new_credential_id: String) -> Self {
|
||||
let message = format!("已切换到凭证: {new_credential_id}");
|
||||
Self {
|
||||
@@ -401,7 +298,6 @@ impl QuotaAutoSwitchResult {
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建使用预览模型的结果
|
||||
pub fn preview_model(model: String) -> Self {
|
||||
let message = format!("已切换到预览模型: {model}");
|
||||
Self {
|
||||
@@ -413,7 +309,6 @@ impl QuotaAutoSwitchResult {
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建未切换的结果
|
||||
pub fn not_switched(message: &str) -> Self {
|
||||
Self {
|
||||
switched: false,
|
||||
@@ -424,7 +319,6 @@ impl QuotaAutoSwitchResult {
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建所有凭证耗尽的结果
|
||||
pub fn all_exhausted(earliest_recovery: Option<DateTime<Utc>>) -> Self {
|
||||
let message = match earliest_recovery {
|
||||
Some(time) => format!("所有凭证配额超限,最早恢复时间: {time}"),
|
||||
@@ -452,7 +346,6 @@ pub struct AllCredentialsExhaustedError {
|
||||
}
|
||||
|
||||
impl AllCredentialsExhaustedError {
|
||||
/// 创建新的错误
|
||||
pub fn new(earliest_recovery: Option<DateTime<Utc>>) -> Self {
|
||||
let retry_after_seconds = earliest_recovery.map(|time| {
|
||||
let now = Utc::now();
|
||||
@@ -478,12 +371,10 @@ impl AllCredentialsExhaustedError {
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取 HTTP 状态码
|
||||
pub fn status_code(&self) -> u16 {
|
||||
503 // Service Unavailable
|
||||
503
|
||||
}
|
||||
|
||||
/// 获取 Retry-After 头的值
|
||||
pub fn retry_after_header(&self) -> Option<String> {
|
||||
self.retry_after_seconds.map(|s| s.to_string())
|
||||
}
|
||||
@@ -498,11 +389,6 @@ impl std::fmt::Display for AllCredentialsExhaustedError {
|
||||
impl std::error::Error for AllCredentialsExhaustedError {}
|
||||
|
||||
/// 实现 IntoResponse 以便在 axum 处理器中直接返回 503 响应
|
||||
///
|
||||
/// 响应格式:
|
||||
/// - HTTP 状态码: 503 Service Unavailable
|
||||
/// - Retry-After 头: 如果有最早恢复时间,则包含等待秒数
|
||||
/// - 响应体: JSON 格式的错误信息
|
||||
impl axum::response::IntoResponse for AllCredentialsExhaustedError {
|
||||
fn into_response(self) -> axum::response::Response {
|
||||
use axum::http::{header, StatusCode};
|
||||
@@ -519,7 +405,6 @@ impl axum::response::IntoResponse for AllCredentialsExhaustedError {
|
||||
|
||||
let mut response = (StatusCode::SERVICE_UNAVAILABLE, Json(json_body)).into_response();
|
||||
|
||||
// 添加 Retry-After 头
|
||||
if let Some(retry_after) = self.retry_after_header() {
|
||||
if let Ok(header_value) = retry_after.parse() {
|
||||
response
|
||||
@@ -534,19 +419,6 @@ impl axum::response::IntoResponse for AllCredentialsExhaustedError {
|
||||
|
||||
impl QuotaManager {
|
||||
/// 处理配额超限并尝试自动切换
|
||||
///
|
||||
/// 当凭证配额超限时,根据配置执行以下策略:
|
||||
/// 1. 如果 switch_project 启用,尝试切换到下一个可用凭证
|
||||
/// 2. 如果 switch_preview_model 启用,尝试使用预览模型
|
||||
///
|
||||
/// # 参数
|
||||
/// - `failed_credential_id`: 失败的凭证 ID
|
||||
/// - `model`: 请求的模型名称
|
||||
/// - `available_credential_ids`: 所有可用的凭证 ID 列表
|
||||
/// - `error_message`: 错误消息
|
||||
///
|
||||
/// # 返回
|
||||
/// 自动切换结果
|
||||
pub fn handle_quota_exceeded(
|
||||
&self,
|
||||
failed_credential_id: &str,
|
||||
@@ -554,12 +426,9 @@ impl QuotaManager {
|
||||
available_credential_ids: &[String],
|
||||
error_message: &str,
|
||||
) -> QuotaAutoSwitchResult {
|
||||
// 标记当前凭证为配额超限
|
||||
self.mark_quota_exceeded(failed_credential_id, error_message);
|
||||
|
||||
// 如果启用了自动切换项目
|
||||
if self.config.switch_project {
|
||||
// 查找下一个可用的凭证(排除已超限的)
|
||||
for cred_id in available_credential_ids {
|
||||
if cred_id != failed_credential_id && self.is_available(cred_id) {
|
||||
tracing::info!(
|
||||
@@ -572,7 +441,6 @@ impl QuotaManager {
|
||||
}
|
||||
}
|
||||
|
||||
// 如果没有可用凭证,尝试使用预览模型
|
||||
if self.config.switch_preview_model {
|
||||
if let Some(preview) = self.get_preview_model(model) {
|
||||
tracing::info!(
|
||||
@@ -584,7 +452,6 @@ impl QuotaManager {
|
||||
}
|
||||
}
|
||||
|
||||
// 所有凭证都不可用
|
||||
let earliest = self.earliest_recovery();
|
||||
tracing::warn!(
|
||||
credential_id = %failed_credential_id,
|
||||
@@ -595,15 +462,6 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 选择下一个可用凭证
|
||||
///
|
||||
/// 从可用凭证列表中选择一个未处于配额超限状态的凭证
|
||||
///
|
||||
/// # 参数
|
||||
/// - `available_credential_ids`: 所有可用的凭证 ID 列表
|
||||
///
|
||||
/// # 返回
|
||||
/// - `Some(String)`: 可用的凭证 ID
|
||||
/// - `None`: 没有可用凭证
|
||||
pub fn select_available_credential(
|
||||
&self,
|
||||
available_credential_ids: &[String],
|
||||
@@ -617,12 +475,6 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 过滤出可用的凭证 ID 列表
|
||||
///
|
||||
/// # 参数
|
||||
/// - `credential_ids`: 所有凭证 ID 列表
|
||||
///
|
||||
/// # 返回
|
||||
/// 未处于配额超限状态的凭证 ID 列表
|
||||
pub fn filter_available_credentials(&self, credential_ids: &[String]) -> Vec<String> {
|
||||
credential_ids
|
||||
.iter()
|
||||
@@ -632,13 +484,6 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 检查是否所有凭证都已耗尽
|
||||
///
|
||||
/// # 参数
|
||||
/// - `credential_ids`: 所有凭证 ID 列表
|
||||
///
|
||||
/// # 返回
|
||||
/// - `Ok(())`: 有可用凭证
|
||||
/// - `Err(AllCredentialsExhaustedError)`: 所有凭证都已耗尽
|
||||
pub fn check_all_exhausted(
|
||||
&self,
|
||||
credential_ids: &[String],
|
||||
@@ -652,9 +497,6 @@ impl QuotaManager {
|
||||
}
|
||||
|
||||
/// 获取所有凭证耗尽时的错误响应
|
||||
///
|
||||
/// # 返回
|
||||
/// 包含 503 状态码和 Retry-After 头的错误
|
||||
pub fn get_exhausted_error(&self) -> AllCredentialsExhaustedError {
|
||||
AllCredentialsExhaustedError::new(self.earliest_recovery())
|
||||
}
|
||||
@@ -711,20 +553,17 @@ mod unit_tests {
|
||||
cooldown_seconds: 300,
|
||||
};
|
||||
let manager = QuotaManager::new(config);
|
||||
|
||||
let available = vec![
|
||||
"cred-1".to_string(),
|
||||
"cred-2".to_string(),
|
||||
"cred-3".to_string(),
|
||||
];
|
||||
|
||||
let result = manager.handle_quota_exceeded(
|
||||
"cred-1",
|
||||
"gemini-2.5-pro",
|
||||
&available,
|
||||
"Rate limit exceeded",
|
||||
);
|
||||
|
||||
assert!(result.switched);
|
||||
assert_eq!(result.new_credential_id, Some("cred-2".to_string()));
|
||||
assert!(!result.used_preview_model);
|
||||
@@ -738,16 +577,13 @@ mod unit_tests {
|
||||
cooldown_seconds: 300,
|
||||
};
|
||||
let manager = QuotaManager::new(config);
|
||||
|
||||
let available = vec!["cred-1".to_string()];
|
||||
|
||||
let result = manager.handle_quota_exceeded(
|
||||
"cred-1",
|
||||
"gemini-2.5-pro",
|
||||
&available,
|
||||
"Rate limit exceeded",
|
||||
);
|
||||
|
||||
assert!(!result.switched);
|
||||
assert!(result.used_preview_model);
|
||||
assert_eq!(
|
||||
@@ -764,20 +600,15 @@ mod unit_tests {
|
||||
cooldown_seconds: 300,
|
||||
};
|
||||
let manager = QuotaManager::new(config);
|
||||
|
||||
// 标记所有凭证为超限
|
||||
manager.mark_quota_exceeded("cred-1", "test");
|
||||
manager.mark_quota_exceeded("cred-2", "test");
|
||||
|
||||
let available = vec!["cred-1".to_string(), "cred-2".to_string()];
|
||||
|
||||
let result = manager.handle_quota_exceeded(
|
||||
"cred-1",
|
||||
"gemini-2.5-pro",
|
||||
&available,
|
||||
"Rate limit exceeded",
|
||||
);
|
||||
|
||||
assert!(!result.switched);
|
||||
assert!(!result.used_preview_model);
|
||||
assert!(result.message.contains("所有凭证配额超限"));
|
||||
@@ -786,16 +617,12 @@ mod unit_tests {
|
||||
#[test]
|
||||
fn test_select_available_credential() {
|
||||
let manager = QuotaManager::with_defaults();
|
||||
|
||||
// 标记 cred-1 为超限
|
||||
manager.mark_quota_exceeded("cred-1", "test");
|
||||
|
||||
let available = vec![
|
||||
"cred-1".to_string(),
|
||||
"cred-2".to_string(),
|
||||
"cred-3".to_string(),
|
||||
];
|
||||
|
||||
let selected = manager.select_available_credential(&available);
|
||||
assert_eq!(selected, Some("cred-2".to_string()));
|
||||
}
|
||||
@@ -803,18 +630,14 @@ mod unit_tests {
|
||||
#[test]
|
||||
fn test_filter_available_credentials() {
|
||||
let manager = QuotaManager::with_defaults();
|
||||
|
||||
// 标记 cred-1 和 cred-3 为超限
|
||||
manager.mark_quota_exceeded("cred-1", "test");
|
||||
manager.mark_quota_exceeded("cred-3", "test");
|
||||
|
||||
let all = vec![
|
||||
"cred-1".to_string(),
|
||||
"cred-2".to_string(),
|
||||
"cred-3".to_string(),
|
||||
"cred-4".to_string(),
|
||||
];
|
||||
|
||||
let available = manager.filter_available_credentials(&all);
|
||||
assert_eq!(available, vec!["cred-2".to_string(), "cred-4".to_string()]);
|
||||
}
|
||||
@@ -831,10 +654,8 @@ mod unit_tests {
|
||||
fn test_all_credentials_exhausted_error_with_recovery() {
|
||||
let recovery_time = Utc::now() + Duration::seconds(300);
|
||||
let error = AllCredentialsExhaustedError::new(Some(recovery_time));
|
||||
|
||||
assert_eq!(error.status_code(), 503);
|
||||
assert!(error.retry_after_header().is_some());
|
||||
|
||||
let retry_after = error.retry_after_seconds.unwrap();
|
||||
assert!(retry_after > 0);
|
||||
assert!(retry_after <= 300);
|
||||
@@ -843,12 +664,8 @@ mod unit_tests {
|
||||
#[test]
|
||||
fn test_check_all_exhausted_has_available() {
|
||||
let manager = QuotaManager::with_defaults();
|
||||
|
||||
// 标记部分凭证为超限
|
||||
manager.mark_quota_exceeded("cred-1", "test");
|
||||
|
||||
let all = vec!["cred-1".to_string(), "cred-2".to_string()];
|
||||
|
||||
let result = manager.check_all_exhausted(&all);
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
@@ -856,16 +673,11 @@ mod unit_tests {
|
||||
#[test]
|
||||
fn test_check_all_exhausted_none_available() {
|
||||
let manager = QuotaManager::with_defaults();
|
||||
|
||||
// 标记所有凭证为超限
|
||||
manager.mark_quota_exceeded("cred-1", "test");
|
||||
manager.mark_quota_exceeded("cred-2", "test");
|
||||
|
||||
let all = vec!["cred-1".to_string(), "cred-2".to_string()];
|
||||
|
||||
let result = manager.check_all_exhausted(&all);
|
||||
assert!(result.is_err());
|
||||
|
||||
let error = result.unwrap_err();
|
||||
assert_eq!(error.status_code(), 503);
|
||||
assert!(error.earliest_recovery.is_some());
|
||||
@@ -874,10 +686,7 @@ mod unit_tests {
|
||||
#[test]
|
||||
fn test_get_exhausted_error() {
|
||||
let manager = QuotaManager::with_defaults();
|
||||
|
||||
// 标记凭证为超限
|
||||
manager.mark_quota_exceeded("cred-1", "test");
|
||||
|
||||
let error = manager.get_exhausted_error();
|
||||
assert_eq!(error.status_code(), 503);
|
||||
assert!(error.earliest_recovery.is_some());
|
||||
@@ -890,8 +699,7 @@ mod unit_tests {
|
||||
switch_preview_model: true,
|
||||
cooldown_seconds: 300,
|
||||
};
|
||||
let manager = QuotaManager::new(config.clone());
|
||||
|
||||
let manager = QuotaManager::new(config);
|
||||
assert_eq!(manager.config().cooldown_seconds, 300);
|
||||
assert!(manager.config().switch_project);
|
||||
assert!(manager.config().switch_preview_model);
|
||||
@@ -901,9 +709,7 @@ mod unit_tests {
|
||||
#[test]
|
||||
fn test_quota_manager_mark_exceeded() {
|
||||
let manager = QuotaManager::with_defaults();
|
||||
|
||||
let record = manager.mark_quota_exceeded("cred-1", "Rate limit exceeded");
|
||||
|
||||
assert_eq!(record.credential_id, "cred-1");
|
||||
assert_eq!(record.reason, "Rate limit exceeded");
|
||||
assert!(record.cooldown_until > Utc::now());
|
||||
@@ -915,18 +721,12 @@ mod unit_tests {
|
||||
let config = QuotaExceededConfig {
|
||||
switch_project: true,
|
||||
switch_preview_model: true,
|
||||
cooldown_seconds: 1, // 1 秒冷却
|
||||
cooldown_seconds: 1,
|
||||
};
|
||||
let manager = QuotaManager::new(config);
|
||||
|
||||
// 未标记的凭证应该可用
|
||||
assert!(manager.is_available("cred-1"));
|
||||
|
||||
// 标记后应该不可用
|
||||
manager.mark_quota_exceeded("cred-1", "test");
|
||||
assert!(!manager.is_available("cred-1"));
|
||||
|
||||
// 等待冷却期过后应该可用
|
||||
std::thread::sleep(std::time::Duration::from_secs(2));
|
||||
assert!(manager.is_available("cred-1"));
|
||||
}
|
||||
@@ -936,21 +736,14 @@ mod unit_tests {
|
||||
let config = QuotaExceededConfig {
|
||||
switch_project: true,
|
||||
switch_preview_model: true,
|
||||
cooldown_seconds: 0, // 立即过期
|
||||
cooldown_seconds: 0,
|
||||
};
|
||||
let manager = QuotaManager::new(config);
|
||||
|
||||
// 标记多个凭证
|
||||
manager.mark_quota_exceeded("cred-1", "test");
|
||||
manager.mark_quota_exceeded("cred-2", "test");
|
||||
manager.mark_quota_exceeded("cred-3", "test");
|
||||
|
||||
assert_eq!(manager.exceeded_count(), 3);
|
||||
|
||||
// 等待一小段时间确保过期
|
||||
std::thread::sleep(std::time::Duration::from_millis(100));
|
||||
|
||||
// 清理过期记录
|
||||
let cleaned = manager.cleanup_expired();
|
||||
assert_eq!(cleaned, 3);
|
||||
assert_eq!(manager.exceeded_count(), 0);
|
||||
@@ -959,26 +752,18 @@ mod unit_tests {
|
||||
#[test]
|
||||
fn test_quota_manager_restore_credential() {
|
||||
let manager = QuotaManager::with_defaults();
|
||||
|
||||
manager.mark_quota_exceeded("cred-1", "test");
|
||||
assert!(!manager.is_available("cred-1"));
|
||||
|
||||
// 手动恢复
|
||||
let restored = manager.restore_credential("cred-1");
|
||||
assert!(restored);
|
||||
assert!(manager.is_available("cred-1"));
|
||||
|
||||
// 再次恢复应该返回 false
|
||||
let restored = manager.restore_credential("cred-1");
|
||||
assert!(!restored);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_quota_manager_is_quota_exceeded_error() {
|
||||
// 429 状态码
|
||||
assert!(QuotaManager::is_quota_exceeded_error(Some(429), ""));
|
||||
|
||||
// 关键词检测
|
||||
assert!(QuotaManager::is_quota_exceeded_error(
|
||||
Some(400),
|
||||
"Rate limit exceeded"
|
||||
@@ -991,8 +776,6 @@ mod unit_tests {
|
||||
Some(400),
|
||||
"Too many requests"
|
||||
));
|
||||
|
||||
// 非配额超限错误
|
||||
assert!(!QuotaManager::is_quota_exceeded_error(
|
||||
Some(400),
|
||||
"Bad Request"
|
||||
@@ -1006,8 +789,6 @@ mod unit_tests {
|
||||
#[test]
|
||||
fn test_quota_manager_get_preview_model() {
|
||||
let manager = QuotaManager::with_defaults();
|
||||
|
||||
// 正常模型应该返回预览版本
|
||||
assert_eq!(
|
||||
manager.get_preview_model("gemini-2.5-pro"),
|
||||
Some("gemini-2.5-pro-preview".to_string())
|
||||
@@ -1016,8 +797,6 @@ mod unit_tests {
|
||||
manager.get_preview_model("claude-3-opus"),
|
||||
Some("claude-3-opus-preview".to_string())
|
||||
);
|
||||
|
||||
// 已经是预览版本应该返回 None
|
||||
assert_eq!(manager.get_preview_model("gemini-2.5-pro-preview"), None);
|
||||
assert_eq!(
|
||||
manager.get_preview_model("claude-3-opus-preview-20240101"),
|
||||
@@ -1029,25 +808,20 @@ mod unit_tests {
|
||||
fn test_quota_manager_get_preview_model_disabled() {
|
||||
let config = QuotaExceededConfig {
|
||||
switch_project: true,
|
||||
switch_preview_model: false, // 禁用预览模型
|
||||
switch_preview_model: false,
|
||||
cooldown_seconds: 300,
|
||||
};
|
||||
let manager = QuotaManager::new(config);
|
||||
|
||||
// 禁用时应该返回 None
|
||||
assert_eq!(manager.get_preview_model("gemini-2.5-pro"), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_preview_model() {
|
||||
// 预览版本
|
||||
assert!(QuotaManager::is_preview_model("gemini-2.5-pro-preview"));
|
||||
assert!(QuotaManager::is_preview_model(
|
||||
"claude-3-opus-preview-20240101"
|
||||
));
|
||||
assert!(QuotaManager::is_preview_model("gpt-4-preview"));
|
||||
|
||||
// 非预览版本
|
||||
assert!(!QuotaManager::is_preview_model("gemini-2.5-pro"));
|
||||
assert!(!QuotaManager::is_preview_model("claude-3-opus"));
|
||||
assert!(!QuotaManager::is_preview_model("gpt-4"));
|
||||
@@ -1055,7 +829,6 @@ mod unit_tests {
|
||||
|
||||
#[test]
|
||||
fn test_get_original_model() {
|
||||
// 从预览版本获取原始版本
|
||||
assert_eq!(
|
||||
QuotaManager::get_original_model("gemini-2.5-pro-preview"),
|
||||
Some("gemini-2.5-pro".to_string())
|
||||
@@ -1068,8 +841,6 @@ mod unit_tests {
|
||||
QuotaManager::get_original_model("gpt-4-preview"),
|
||||
Some("gpt-4".to_string())
|
||||
);
|
||||
|
||||
// 非预览版本应该返回 None
|
||||
assert_eq!(QuotaManager::get_original_model("gemini-2.5-pro"), None);
|
||||
assert_eq!(QuotaManager::get_original_model("claude-3-opus"), None);
|
||||
}
|
||||
@@ -1082,11 +853,7 @@ mod unit_tests {
|
||||
cooldown_seconds: 300,
|
||||
};
|
||||
let manager = QuotaManager::new(config);
|
||||
|
||||
// 没有超限凭证时应该返回 None
|
||||
assert!(manager.earliest_recovery().is_none());
|
||||
|
||||
// 标记凭证后应该返回最早的恢复时间
|
||||
manager.mark_quota_exceeded("cred-1", "test");
|
||||
let recovery = manager.earliest_recovery();
|
||||
assert!(recovery.is_some());
|
||||
@@ -1100,11 +867,7 @@ mod unit_tests {
|
||||
cooldown_seconds: 300,
|
||||
};
|
||||
let manager = QuotaManager::new(config);
|
||||
|
||||
// 未标记的凭证应该返回 None
|
||||
assert!(manager.remaining_cooldown_seconds("cred-1").is_none());
|
||||
|
||||
// 标记后应该返回剩余秒数
|
||||
manager.mark_quota_exceeded("cred-1", "test");
|
||||
let remaining = manager.remaining_cooldown_seconds("cred-1");
|
||||
assert!(remaining.is_some());
|
||||
@@ -1117,23 +880,17 @@ mod unit_tests {
|
||||
use axum::http::{header, StatusCode};
|
||||
use axum::response::IntoResponse;
|
||||
|
||||
// 测试无恢复时间的情况
|
||||
let error = AllCredentialsExhaustedError::new(None);
|
||||
let response = error.into_response();
|
||||
|
||||
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
|
||||
assert!(response.headers().get(header::RETRY_AFTER).is_none());
|
||||
|
||||
// 测试有恢复时间的情况
|
||||
let recovery_time = Utc::now() + Duration::seconds(300);
|
||||
let error = AllCredentialsExhaustedError::new(Some(recovery_time));
|
||||
let response = error.into_response();
|
||||
|
||||
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
|
||||
let retry_after = response.headers().get(header::RETRY_AFTER);
|
||||
assert!(retry_after.is_some());
|
||||
|
||||
// 验证 Retry-After 值在合理范围内
|
||||
let retry_value: u64 = retry_after.unwrap().to_str().unwrap().parse().unwrap();
|
||||
assert!(retry_value > 0);
|
||||
assert!(retry_value <= 300);
|
||||
@@ -3,10 +3,12 @@
|
||||
//! 负责将凭证池变更同步到 YAML 配置文件
|
||||
//! 实现凭证的添加、删除、更新操作与配置文件的同步
|
||||
|
||||
use crate::config::{
|
||||
use proxycast_core::config::{
|
||||
expand_tilde, ApiKeyEntry, Config, ConfigError, ConfigManager, CredentialEntry, YamlService,
|
||||
};
|
||||
use crate::models::provider_pool_model::{CredentialData, PoolProviderType, ProviderCredential};
|
||||
use proxycast_core::models::provider_pool_model::{
|
||||
CredentialData, PoolProviderType, ProviderCredential,
|
||||
};
|
||||
use std::path::PathBuf;
|
||||
use std::sync::{Arc, RwLock};
|
||||
|
||||
@@ -49,8 +51,6 @@ impl From<std::io::Error> for SyncError {
|
||||
}
|
||||
|
||||
/// 凭证同步服务
|
||||
///
|
||||
/// 负责将凭证池变更同步到 YAML 配置文件
|
||||
pub struct CredentialSyncService {
|
||||
/// 配置管理器
|
||||
config_manager: Arc<RwLock<ConfigManager>>,
|
||||
@@ -80,8 +80,6 @@ impl CredentialSyncService {
|
||||
|
||||
let config_path = manager.config_path().to_path_buf();
|
||||
manager.set_config(config.clone());
|
||||
|
||||
// 使用 YamlService 保存配置,保留注释
|
||||
YamlService::save_preserve_comments(&config_path, &config)?;
|
||||
Ok(())
|
||||
}
|
||||
@@ -100,18 +98,10 @@ impl CredentialSyncService {
|
||||
}
|
||||
|
||||
/// 添加凭证并同步到配置
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `credential` - 要添加的凭证
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(())` - 添加成功
|
||||
/// * `Err(SyncError)` - 添加失败
|
||||
pub fn add_credential(&self, credential: &ProviderCredential) -> Result<(), SyncError> {
|
||||
let mut config = self.get_config()?;
|
||||
|
||||
match &credential.credential {
|
||||
// OAuth 凭证:保存 token 文件到 auth_dir,配置中只保存相对路径
|
||||
CredentialData::KiroOAuth { creds_file_path } => {
|
||||
let token_file =
|
||||
self.save_oauth_token_file(creds_file_path, &credential.uuid, "kiro")?;
|
||||
@@ -137,12 +127,10 @@ impl CredentialSyncService {
|
||||
config.credential_pool.gemini.push(entry);
|
||||
}
|
||||
CredentialData::AntigravityOAuth { .. } => {
|
||||
// Antigravity 暂不支持同步到配置
|
||||
return Err(SyncError::InvalidCredentialType(
|
||||
"Antigravity 凭证暂不支持同步到配置".to_string(),
|
||||
));
|
||||
}
|
||||
// API Key 凭证:直接保存到 YAML
|
||||
CredentialData::OpenAIKey { api_key, base_url } => {
|
||||
let entry = ApiKeyEntry {
|
||||
id: credential.uuid.clone(),
|
||||
@@ -168,7 +156,7 @@ impl CredentialSyncService {
|
||||
base_url,
|
||||
model_aliases,
|
||||
} => {
|
||||
use crate::config::VertexModelAlias;
|
||||
use proxycast_core::models::vertex_model::VertexModelAlias;
|
||||
let models: Vec<VertexModelAlias> = model_aliases
|
||||
.iter()
|
||||
.map(|(alias, name)| VertexModelAlias {
|
||||
@@ -176,7 +164,7 @@ impl CredentialSyncService {
|
||||
name: name.clone(),
|
||||
})
|
||||
.collect();
|
||||
let entry = crate::config::VertexApiKeyEntry {
|
||||
let entry = proxycast_core::models::vertex_model::VertexApiKeyEntry {
|
||||
id: credential.uuid.clone(),
|
||||
api_key: api_key.clone(),
|
||||
base_url: base_url.clone(),
|
||||
@@ -191,7 +179,7 @@ impl CredentialSyncService {
|
||||
base_url,
|
||||
excluded_models,
|
||||
} => {
|
||||
use crate::config::GeminiApiKeyEntry;
|
||||
use proxycast_core::config::GeminiApiKeyEntry;
|
||||
let entry = GeminiApiKeyEntry {
|
||||
id: credential.uuid.clone(),
|
||||
api_key: api_key.clone(),
|
||||
@@ -203,19 +191,16 @@ impl CredentialSyncService {
|
||||
config.credential_pool.gemini_api_keys.push(entry);
|
||||
}
|
||||
CredentialData::CodexOAuth { .. } => {
|
||||
// Codex 暂不支持同步到配置
|
||||
return Err(SyncError::InvalidCredentialType(
|
||||
"Codex 凭证暂不支持同步到配置".to_string(),
|
||||
));
|
||||
}
|
||||
CredentialData::ClaudeOAuth { .. } => {
|
||||
// Claude OAuth 暂不支持同步到配置
|
||||
return Err(SyncError::InvalidCredentialType(
|
||||
"Claude OAuth 凭证暂不支持同步到配置".to_string(),
|
||||
));
|
||||
}
|
||||
CredentialData::AnthropicKey { api_key, base_url } => {
|
||||
// Anthropic API Key 保存到 claude 配置(使用相同的 API 格式)
|
||||
let entry = ApiKeyEntry {
|
||||
id: credential.uuid.clone(),
|
||||
api_key: api_key.clone(),
|
||||
@@ -223,8 +208,6 @@ impl CredentialSyncService {
|
||||
disabled: credential.is_disabled,
|
||||
proxy_url: None,
|
||||
};
|
||||
// 注意:Anthropic 凭证保存到单独的 anthropic 配置(如果有的话)
|
||||
// 目前暂时保存到 claude 配置中
|
||||
config.credential_pool.claude.push(entry);
|
||||
}
|
||||
}
|
||||
@@ -233,14 +216,6 @@ impl CredentialSyncService {
|
||||
}
|
||||
|
||||
/// 保存 OAuth token 文件到 auth_dir
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `source_path` - 源 token 文件路径
|
||||
/// * `credential_id` - 凭证 ID
|
||||
/// * `provider` - Provider 名称
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(String)` - 相对于 auth_dir 的 token 文件路径
|
||||
fn save_oauth_token_file(
|
||||
&self,
|
||||
source_path: &str,
|
||||
@@ -251,29 +226,28 @@ impl CredentialSyncService {
|
||||
let provider_dir = auth_dir.join(provider);
|
||||
std::fs::create_dir_all(&provider_dir)?;
|
||||
|
||||
// 生成 token 文件名
|
||||
let token_filename = format!("{credential_id}.json");
|
||||
let token_path = provider_dir.join(&token_filename);
|
||||
|
||||
// 展开源路径并复制文件
|
||||
let source = expand_tilde(source_path);
|
||||
if source.exists() {
|
||||
std::fs::copy(&source, &token_path)?;
|
||||
}
|
||||
|
||||
// 返回相对路径
|
||||
Ok(format!("{provider}/{token_filename}"))
|
||||
}
|
||||
|
||||
/// 删除 OAuth token 文件
|
||||
fn delete_oauth_token_file(&self, token_file: &str) -> Result<(), SyncError> {
|
||||
let auth_dir = self.get_auth_dir()?;
|
||||
let token_path = auth_dir.join(token_file);
|
||||
if token_path.exists() {
|
||||
std::fs::remove_file(&token_path)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 删除凭证并同步到配置
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `provider_type` - Provider 类型
|
||||
/// * `credential_id` - 凭证 ID
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(())` - 删除成功
|
||||
/// * `Err(SyncError)` - 删除失败
|
||||
pub fn remove_credential(
|
||||
&self,
|
||||
provider_type: PoolProviderType,
|
||||
@@ -357,24 +331,20 @@ impl CredentialSyncService {
|
||||
}
|
||||
}
|
||||
PoolProviderType::Codex => {
|
||||
// Codex 暂不支持同步到配置
|
||||
return Err(SyncError::InvalidCredentialType(
|
||||
"Codex 凭证暂不支持同步到配置".to_string(),
|
||||
));
|
||||
}
|
||||
PoolProviderType::ClaudeOAuth => {
|
||||
// Claude OAuth 暂不支持同步到配置
|
||||
return Err(SyncError::InvalidCredentialType(
|
||||
"Claude OAuth 凭证暂不支持同步到配置".to_string(),
|
||||
));
|
||||
}
|
||||
// Anthropic 兼容格式 - 不支持同步到配置
|
||||
PoolProviderType::AnthropicCompatible => {
|
||||
return Err(SyncError::InvalidCredentialType(
|
||||
"Anthropic Compatible 凭证暂不支持同步到配置".to_string(),
|
||||
));
|
||||
}
|
||||
// API Key Provider 类型 - 不支持同步到配置
|
||||
PoolProviderType::Anthropic
|
||||
| PoolProviderType::AzureOpenai
|
||||
| PoolProviderType::AwsBedrock
|
||||
@@ -392,24 +362,7 @@ impl CredentialSyncService {
|
||||
self.update_config(config)
|
||||
}
|
||||
|
||||
/// 删除 OAuth token 文件
|
||||
fn delete_oauth_token_file(&self, token_file: &str) -> Result<(), SyncError> {
|
||||
let auth_dir = self.get_auth_dir()?;
|
||||
let token_path = auth_dir.join(token_file);
|
||||
if token_path.exists() {
|
||||
std::fs::remove_file(&token_path)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 更新凭证并同步到配置
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `credential` - 更新后的凭证
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(())` - 更新成功
|
||||
/// * `Err(SyncError)` - 更新失败
|
||||
pub fn update_credential(&self, credential: &ProviderCredential) -> Result<(), SyncError> {
|
||||
let mut config = self.get_config()?;
|
||||
let mut found = false;
|
||||
@@ -423,7 +376,6 @@ impl CredentialSyncService {
|
||||
.find(|e| e.id == credential.uuid)
|
||||
{
|
||||
entry.disabled = credential.is_disabled;
|
||||
// 如果源文件路径变化,更新 token 文件
|
||||
let new_token_file =
|
||||
self.save_oauth_token_file(creds_file_path, &credential.uuid, "kiro")?;
|
||||
entry.token_file = new_token_file;
|
||||
@@ -488,7 +440,7 @@ impl CredentialSyncService {
|
||||
.iter_mut()
|
||||
.find(|e| e.id == credential.uuid)
|
||||
{
|
||||
use crate::config::VertexModelAlias;
|
||||
use proxycast_core::models::vertex_model::VertexModelAlias;
|
||||
entry.api_key = api_key.clone();
|
||||
entry.base_url = base_url.clone();
|
||||
entry.models = model_aliases
|
||||
@@ -521,19 +473,16 @@ impl CredentialSyncService {
|
||||
}
|
||||
}
|
||||
CredentialData::CodexOAuth { .. } => {
|
||||
// Codex 暂不支持同步到配置
|
||||
return Err(SyncError::InvalidCredentialType(
|
||||
"Codex 凭证暂不支持同步到配置".to_string(),
|
||||
));
|
||||
}
|
||||
CredentialData::ClaudeOAuth { .. } => {
|
||||
// Claude OAuth 暂不支持同步到配置
|
||||
return Err(SyncError::InvalidCredentialType(
|
||||
"Claude OAuth 凭证暂不支持同步到配置".to_string(),
|
||||
));
|
||||
}
|
||||
CredentialData::AnthropicKey { api_key, base_url } => {
|
||||
// Anthropic API Key 更新到 claude 配置
|
||||
if let Some(entry) = config
|
||||
.credential_pool
|
||||
.claude
|
||||
@@ -556,12 +505,6 @@ impl CredentialSyncService {
|
||||
}
|
||||
|
||||
/// 从配置加载凭证到池中
|
||||
///
|
||||
/// 启动时从 YAML 配置加载凭证
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(Vec<ProviderCredential>)` - 加载的凭证列表
|
||||
/// * `Err(SyncError)` - 加载失败
|
||||
pub fn load_from_config(&self) -> Result<Vec<ProviderCredential>, SyncError> {
|
||||
let config = self.get_config()?;
|
||||
let auth_dir = self.get_auth_dir()?;
|
||||
@@ -570,13 +513,12 @@ impl CredentialSyncService {
|
||||
// 加载 Kiro 凭证
|
||||
for entry in &config.credential_pool.kiro {
|
||||
let token_path = auth_dir.join(&entry.token_file);
|
||||
let cred = ProviderCredential::new(
|
||||
let mut cred = ProviderCredential::new(
|
||||
PoolProviderType::Kiro,
|
||||
CredentialData::KiroOAuth {
|
||||
creds_file_path: token_path.to_string_lossy().to_string(),
|
||||
},
|
||||
);
|
||||
let mut cred = cred;
|
||||
cred.uuid = entry.id.clone();
|
||||
cred.is_disabled = entry.disabled;
|
||||
credentials.push(cred);
|
||||
@@ -585,14 +527,13 @@ impl CredentialSyncService {
|
||||
// 加载 Gemini 凭证
|
||||
for entry in &config.credential_pool.gemini {
|
||||
let token_path = auth_dir.join(&entry.token_file);
|
||||
let cred = ProviderCredential::new(
|
||||
let mut cred = ProviderCredential::new(
|
||||
PoolProviderType::Gemini,
|
||||
CredentialData::GeminiOAuth {
|
||||
creds_file_path: token_path.to_string_lossy().to_string(),
|
||||
project_id: None,
|
||||
},
|
||||
);
|
||||
let mut cred = cred;
|
||||
cred.uuid = entry.id.clone();
|
||||
cred.is_disabled = entry.disabled;
|
||||
credentials.push(cred);
|
||||
@@ -600,14 +541,13 @@ impl CredentialSyncService {
|
||||
|
||||
// 加载 OpenAI 凭证
|
||||
for entry in &config.credential_pool.openai {
|
||||
let cred = ProviderCredential::new(
|
||||
let mut cred = ProviderCredential::new(
|
||||
PoolProviderType::OpenAI,
|
||||
CredentialData::OpenAIKey {
|
||||
api_key: entry.api_key.clone(),
|
||||
base_url: entry.base_url.clone(),
|
||||
},
|
||||
);
|
||||
let mut cred = cred;
|
||||
cred.uuid = entry.id.clone();
|
||||
cred.is_disabled = entry.disabled;
|
||||
credentials.push(cred);
|
||||
@@ -615,14 +555,13 @@ impl CredentialSyncService {
|
||||
|
||||
// 加载 Claude 凭证
|
||||
for entry in &config.credential_pool.claude {
|
||||
let cred = ProviderCredential::new(
|
||||
let mut cred = ProviderCredential::new(
|
||||
PoolProviderType::Claude,
|
||||
CredentialData::ClaudeKey {
|
||||
api_key: entry.api_key.clone(),
|
||||
base_url: entry.base_url.clone(),
|
||||
},
|
||||
);
|
||||
let mut cred = cred;
|
||||
cred.uuid = entry.id.clone();
|
||||
cred.is_disabled = entry.disabled;
|
||||
credentials.push(cred);
|
||||
@@ -635,7 +574,7 @@ impl CredentialSyncService {
|
||||
.iter()
|
||||
.map(|m| (m.alias.clone(), m.name.clone()))
|
||||
.collect();
|
||||
let cred = ProviderCredential::new(
|
||||
let mut cred = ProviderCredential::new(
|
||||
PoolProviderType::Vertex,
|
||||
CredentialData::VertexKey {
|
||||
api_key: entry.api_key.clone(),
|
||||
@@ -643,7 +582,6 @@ impl CredentialSyncService {
|
||||
model_aliases,
|
||||
},
|
||||
);
|
||||
let mut cred = cred;
|
||||
cred.uuid = entry.id.clone();
|
||||
cred.is_disabled = entry.disabled;
|
||||
credentials.push(cred);
|
||||
@@ -651,7 +589,7 @@ impl CredentialSyncService {
|
||||
|
||||
// 加载 Gemini API Key 凭证
|
||||
for entry in &config.credential_pool.gemini_api_keys {
|
||||
let cred = ProviderCredential::new(
|
||||
let mut cred = ProviderCredential::new(
|
||||
PoolProviderType::GeminiApiKey,
|
||||
CredentialData::GeminiApiKey {
|
||||
api_key: entry.api_key.clone(),
|
||||
@@ -659,7 +597,6 @@ impl CredentialSyncService {
|
||||
excluded_models: entry.excluded_models.clone(),
|
||||
},
|
||||
);
|
||||
let mut cred = cred;
|
||||
cred.uuid = entry.id.clone();
|
||||
cred.is_disabled = entry.disabled;
|
||||
credentials.push(cred);
|
||||
@@ -669,37 +606,18 @@ impl CredentialSyncService {
|
||||
}
|
||||
|
||||
/// 获取 OAuth token 文件的完整路径
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `token_file` - 相对于 auth_dir 的 token 文件路径
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(PathBuf)` - 完整路径
|
||||
pub fn get_token_file_path(&self, token_file: &str) -> Result<PathBuf, SyncError> {
|
||||
let auth_dir = self.get_auth_dir()?;
|
||||
Ok(auth_dir.join(token_file))
|
||||
}
|
||||
|
||||
/// 读取 OAuth token 文件内容
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `token_file` - 相对于 auth_dir 的 token 文件路径
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(String)` - token 文件内容
|
||||
pub fn read_token_file(&self, token_file: &str) -> Result<String, SyncError> {
|
||||
let path = self.get_token_file_path(token_file)?;
|
||||
std::fs::read_to_string(&path).map_err(SyncError::from)
|
||||
}
|
||||
|
||||
/// 写入 OAuth token 文件内容
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `token_file` - 相对于 auth_dir 的 token 文件路径
|
||||
/// * `content` - token 文件内容
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(())` - 写入成功
|
||||
pub fn write_token_file(&self, token_file: &str, content: &str) -> Result<(), SyncError> {
|
||||
let path = self.get_token_file_path(token_file)?;
|
||||
if let Some(parent) = path.parent() {
|
||||
@@ -8,18 +8,13 @@
|
||||
//! - `pool` - 凭证池管理(来自 proxycast-core)
|
||||
//! - `health` - 健康检查(来自 proxycast-core)
|
||||
//! - `risk` - 风控模块(来自 proxycast-core)
|
||||
//! - `balancer` - 负载均衡策略(本地)
|
||||
//! - `quota` - 配额管理(本地)
|
||||
//! - `sync` - 数据库同步(本地)
|
||||
//! - `balancer` - 负载均衡策略(来自 proxycast-credential)
|
||||
//! - `quota` - 配额管理(来自 proxycast-credential)
|
||||
//! - `sync` - 数据库同步(来自 proxycast-credential)
|
||||
|
||||
// 从 proxycast-core 重新导出核心类型模块
|
||||
pub use proxycast_core::credential::{health, pool, risk, types};
|
||||
|
||||
// 本地模块(依赖 infra 或 Tauri)
|
||||
mod balancer;
|
||||
mod quota;
|
||||
mod sync;
|
||||
|
||||
// 重新导出 core 类型
|
||||
pub use proxycast_core::credential::{
|
||||
CooldownConfig, Credential, CredentialData, CredentialPool, CredentialStats, CredentialStatus,
|
||||
@@ -27,13 +22,8 @@ pub use proxycast_core::credential::{
|
||||
RateLimitEvent, RateLimitStats, RiskController, RiskLevel,
|
||||
};
|
||||
|
||||
// 重新导出本地类型
|
||||
pub use balancer::{BalanceStrategy, CooldownInfo, CredentialSelection, LoadBalancer};
|
||||
pub use quota::{
|
||||
create_shared_quota_manager, start_quota_cleanup_task, AllCredentialsExhaustedError,
|
||||
QuotaAutoSwitchResult, QuotaExceededRecord, QuotaManager,
|
||||
};
|
||||
pub use sync::{CredentialSyncService, SyncError};
|
||||
// 从 proxycast-credential crate 重新导出
|
||||
pub use proxycast_credential::*;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
Reference in New Issue
Block a user