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:
coso
2026-02-08 18:58:00 +08:00
parent e3c9efdf9f
commit a4ca895aa3
8 changed files with 132 additions and 517 deletions
+19
View File
@@ -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"
+2
View File
@@ -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
+35
View File
@@ -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
@@ -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!(
+21
View File
@@ -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() {
+5 -15
View File
@@ -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;