feat: v0.10.0 - 配置与凭证统一导入导出功能

主要更新:
- 实现以 YAML 配置文件为单一数据源的统一配置管理
- 支持凭证池配置同步到 YAML 文件
- 实现配置和凭证的统一导入导出
- 支持 OAuth Token 文件存储在可配置的 auth-dir 目录
- 实现配置热重载时自动同步凭证池
- 添加导出脱敏功能保护敏感信息
- 支持导入时合并或替换模式
This commit is contained in:
coso
2025-12-17 12:00:54 +08:00
parent a18384bcef
commit aca666ea1d
48 changed files with 11340 additions and 231 deletions
+2 -2
View File
@@ -1,12 +1,12 @@
{
"name": "proxycast",
"version": "0.9.0",
"version": "0.10.0",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"name": "proxycast",
"version": "0.9.0",
"version": "0.10.0",
"dependencies": {
"@radix-ui/react-dialog": "^1.1.2",
"@radix-ui/react-dropdown-menu": "^2.1.2",
+1 -1
View File
@@ -1,7 +1,7 @@
{
"name": "proxycast",
"private": true,
"version": "0.9.0",
"version": "0.10.0",
"type": "module",
"scripts": {
"dev": "vite",
+1 -1
View File
@@ -3274,7 +3274,7 @@ dependencies = [
[[package]]
name = "proxycast"
version = "0.9.0"
version = "0.10.0"
dependencies = [
"anyhow",
"async-stream",
+1 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "proxycast"
version = "0.9.0"
version = "0.10.0"
description = "AI API Proxy Desktop App"
authors = ["you"]
edition = "2021"
@@ -0,0 +1,9 @@
# Seeds for failure cases proptest has generated in the past. It is
# automatically read and these particular cases re-run before any
# novel cases are generated.
#
# It is recommended to check this file in to source control so that
# everyone who runs the test benefits from these saved cases.
cc c887db8633047b16f94b48762dc9bfe3c65bb683295d9047265207226e4e0f0b # shrinks to path = "~/."
cc 16635f3c212b2769f7fb51aec0142ac6e9e1a66d2f9e3118829cfba4ba86e4f8 # shrinks to (yaml_with_comments, original_comments) = ("# 0\nserver:\n# \n host: 127.0.0.1\n port: 1\n api_key: 0__AAA0_\nproviders:\n kiro:\n# a\n enabled: false\n region: us-east-1\n gemini:\n enabled: false\n credentials_path: 0-Aaa\n# 1\n qwen:\n enabled: false\n openai:\n enabled: false\n claude:\n enabled: false\ndefault_provider: kiro\nrouting:\n default_provider: kiro\n rules: []\n model_aliases: {}\n exclusions: {}\nretry:\n max_retries: 1\n base_delay_ms: 1\n max_delay_ms: 5000\n auto_switch_provider: false\nlogging:\n enabled: false\n level: debug\n retention_days: 1\n include_request_body: false\ninjection:\n enabled: false\n rules: []\nauth_dir: ~/.proxycast/auth\ncredential_pool: {}", ["# 0", "# ", "# a", "# 1"]), new_config = Config { server: ServerConfig { host: "127.0.0.1", port: 1, api_key: "a_a--a-a" }, providers: ProvidersConfig { kiro: ProviderConfig { enabled: false, credentials_path: None, region: None, project_id: None }, gemini: ProviderConfig { enabled: false, credentials_path: None, region: None, project_id: None }, qwen: ProviderConfig { enabled: false, credentials_path: None, region: None, project_id: Some("J3JRS6") }, openai: CustomProviderConfig { enabled: false, api_key: None, base_url: None }, claude: CustomProviderConfig { enabled: true, api_key: None, base_url: None } }, default_provider: "kiro", routing: RoutingConfig { default_provider: "kiro", rules: [RoutingRuleConfig { pattern: "cmbvvlhoabjthfkoczp-*", provider: "gemini", priority: 33 }], model_aliases: {}, exclusions: {} }, retry: RetrySettings { max_retries: 86, base_delay_ms: 4635, max_delay_ms: 9567, auto_switch_provider: false }, logging: LoggingConfig { enabled: false, level: "debug", retention_days: 16, include_request_body: false }, injection: InjectionSettings { enabled: false, rules: [] }, auth_dir: "~/.proxycast/auth", credential_pool: CredentialPoolConfig { kiro: [], gemini: [], qwen: [], openai: [], claude: [] } }
cc d22f0e24d166175ada35e5f5c4874b91f97fc63e0b40e6207ad9c70c07ecddda # shrinks to content = "{\"version\": }"
@@ -0,0 +1,8 @@
# Seeds for failure cases proptest has generated in the past. It is
# automatically read and these particular cases re-run before any
# novel cases are generated.
#
# It is recommended to check this file in to source control so that
# everyone who runs the test benefits from these saved cases.
cc bba7155a7ed5c22d19990ebdd9935e761dbdc41bc1e11076fbbf4218b82c8002 # shrinks to request_id = "0-AAaA-a", (original_data, sse_body) = ([" "], "data: \n\n")
cc 4f0d0f49ef961c396cdf1aa7ccbd2a01b525e96ba5a2048d18860dbe763abe37 # shrinks to request_id = "a00000aa", data = " ", index = 0
+218 -1
View File
@@ -1,4 +1,7 @@
use crate::config::{Config, ConfigManager};
use crate::config::{
Config, ConfigManager, ExportBundle, ExportOptions as ExportServiceOptions, ExportService,
ImportOptions as ImportServiceOptions, ImportService, ValidationResult,
};
use crate::models::AppType;
use serde::{Deserialize, Serialize};
use std::path::PathBuf;
@@ -356,3 +359,217 @@ pub fn get_config_paths() -> Result<ConfigPathInfo, String> {
json_exists: json_path.exists(),
})
}
// ============ Enhanced Export/Import Commands (using ExportService/ImportService) ============
/// 统一导出选项
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct UnifiedExportOptions {
/// 是否包含配置
pub include_config: bool,
/// 是否包含凭证
pub include_credentials: bool,
/// 是否脱敏敏感信息
pub redact_secrets: bool,
}
/// 统一导出结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct UnifiedExportResult {
/// 导出包内容(JSON 格式)
pub content: String,
/// 建议的文件名
pub suggested_filename: String,
/// 是否已脱敏
pub redacted: bool,
/// 是否包含配置
pub has_config: bool,
/// 是否包含凭证
pub has_credentials: bool,
}
/// 导出完整的配置和凭证包
///
/// # Arguments
/// * `config` - 当前配置
/// * `options` - 导出选项
///
/// # Requirements: 3.1, 3.2
#[tauri::command]
pub fn export_bundle(
config: Config,
options: UnifiedExportOptions,
) -> Result<UnifiedExportResult, String> {
let export_options = ExportServiceOptions {
include_config: options.include_config,
include_credentials: options.include_credentials,
redact_secrets: options.redact_secrets,
};
// 获取应用版本
let app_version = env!("CARGO_PKG_VERSION").to_string();
let bundle =
ExportService::export(&config, &export_options, &app_version).map_err(|e| e.to_string())?;
let content = bundle.to_json().map_err(|e| e.to_string())?;
// 生成带时间戳的文件名
let timestamp = chrono::Local::now().format("%Y%m%d_%H%M%S");
let suffix = if options.redact_secrets {
"_redacted"
} else {
""
};
let scope = match (options.include_config, options.include_credentials) {
(true, true) => "full",
(true, false) => "config",
(false, true) => "credentials",
(false, false) => "empty",
};
let suggested_filename = format!("proxycast_{}_{}{}.json", scope, timestamp, suffix);
Ok(UnifiedExportResult {
content,
suggested_filename,
redacted: bundle.redacted,
has_config: bundle.has_config(),
has_credentials: bundle.has_credentials(),
})
}
/// 仅导出配置为 YAML
///
/// # Arguments
/// * `config` - 当前配置
/// * `redact_secrets` - 是否脱敏敏感信息
///
/// # Requirements: 3.1, 5.1
#[tauri::command]
pub fn export_config_yaml(config: Config, redact_secrets: bool) -> Result<ExportResult, String> {
let content = ExportService::export_yaml(&config, redact_secrets).map_err(|e| e.to_string())?;
// 生成带时间戳的文件名
let timestamp = chrono::Local::now().format("%Y%m%d_%H%M%S");
let suffix = if redact_secrets { "_redacted" } else { "" };
let suggested_filename = format!("proxycast_config_{}{}.yaml", timestamp, suffix);
Ok(ExportResult {
content,
suggested_filename,
})
}
/// 验证导入内容
///
/// # Arguments
/// * `content` - 导入内容(JSON 导出包或 YAML 配置)
///
/// # Requirements: 4.1, 4.2
#[tauri::command]
pub fn validate_import(content: String) -> Result<ValidationResult, String> {
Ok(ImportService::validate(&content))
}
/// 导入完整的导出包
///
/// # Arguments
/// * `current_config` - 当前配置
/// * `content` - 导出包内容(JSON 格式)
/// * `merge` - 是否合并到现有配置
///
/// # Requirements: 4.1, 4.3
#[tauri::command]
pub fn import_bundle(
current_config: Config,
content: String,
merge: bool,
) -> Result<ImportResult, String> {
// 首先尝试解析为 ExportBundle
if let Ok(bundle) = ExportBundle::from_json(&content) {
let options = ImportServiceOptions { merge };
let result =
ImportService::import(&bundle, &current_config, &options, &current_config.auth_dir)
.map_err(|e| e.to_string())?;
return Ok(ImportResult {
success: result.success,
config: result.config,
warnings: result.warnings,
});
}
// 尝试解析为 YAML 配置
let options = ImportServiceOptions { merge };
let result = ImportService::import_yaml(&content, &current_config, &options)
.map_err(|e| e.to_string())?;
Ok(ImportResult {
success: result.success,
config: result.config,
warnings: result.warnings,
})
}
// ============ Path Utility Commands ============
/// 展开路径中的 tilde (~) 为用户主目录
///
/// # Arguments
/// * `path` - 要展开的路径字符串
///
/// # Returns
/// 展开后的完整路径字符串
///
/// # Requirements: 2.3
#[tauri::command]
pub fn expand_path(path: String) -> Result<String, String> {
use crate::config::expand_tilde;
let expanded = expand_tilde(&path);
Ok(expanded.to_string_lossy().to_string())
}
/// 打开认证目录
///
/// # Arguments
/// * `path` - 认证目录路径(支持 tilde 展开)
///
/// # Requirements: 2.2
#[tauri::command]
pub async fn open_auth_dir(path: String) -> Result<bool, String> {
use crate::config::expand_tilde;
let expanded = expand_tilde(&path);
// 确保目录存在
if !expanded.exists() {
std::fs::create_dir_all(&expanded).map_err(|e| e.to_string())?;
}
#[cfg(target_os = "macos")]
{
std::process::Command::new("open")
.arg(&expanded)
.spawn()
.map_err(|e| e.to_string())?;
}
#[cfg(target_os = "windows")]
{
std::process::Command::new("explorer")
.arg(&expanded)
.spawn()
.map_err(|e| e.to_string())?;
}
#[cfg(target_os = "linux")]
{
std::process::Command::new("xdg-open")
.arg(&expanded)
.spawn()
.map_err(|e| e.to_string())?;
}
Ok(true)
}
+62 -7
View File
@@ -1,10 +1,11 @@
//! Provider Pool Tauri 命令
use crate::credential::CredentialSyncService;
use crate::database::dao::provider_pool::ProviderPoolDao;
use crate::database::DbConnection;
use crate::models::provider_pool_model::{
AddCredentialRequest, CredentialData, CredentialDisplay, HealthCheckResult, OAuthStatus,
ProviderCredential, ProviderPoolOverview, UpdateCredentialRequest,
PoolProviderType, ProviderCredential, ProviderPoolOverview, UpdateCredentialRequest,
};
use crate::services::provider_pool_service::ProviderPoolService;
use chrono::Utc;
@@ -16,6 +17,9 @@ use uuid::Uuid;
pub struct ProviderPoolServiceState(pub Arc<ProviderPoolService>);
/// 凭证同步服务状态封装
pub struct CredentialSyncServiceState(pub Option<Arc<CredentialSyncService>>);
/// 展开路径中的 ~ 为用户主目录
fn expand_tilde(path: &str) -> String {
if path.starts_with("~/") {
@@ -121,32 +125,52 @@ pub fn get_provider_pool_credentials(
}
/// 添加凭证
///
/// 添加凭证到数据库,并同步到 YAML 配置文件
/// Requirements: 1.1, 1.2
#[tauri::command]
pub fn add_provider_pool_credential(
db: State<'_, DbConnection>,
pool_service: State<'_, ProviderPoolServiceState>,
sync_service: State<'_, CredentialSyncServiceState>,
request: AddCredentialRequest,
) -> Result<ProviderCredential, String> {
pool_service.0.add_credential(
// 添加到数据库
let credential = pool_service.0.add_credential(
&db,
&request.provider_type,
request.credential,
request.name,
request.check_health,
request.check_model_name,
)
)?;
// 同步到 YAML 配置(如果同步服务可用)
if let Some(ref sync) = sync_service.0 {
if let Err(e) = sync.add_credential(&credential) {
// 记录警告但不中断操作
tracing::warn!("同步凭证到 YAML 失败: {}", e);
}
}
Ok(credential)
}
/// 更新凭证
/// 更新凭证
///
/// 更新数据库中的凭证,并同步到 YAML 配置文件
/// Requirements: 1.1, 1.2
#[tauri::command]
pub fn update_provider_pool_credential(
db: State<'_, DbConnection>,
pool_service: State<'_, ProviderPoolServiceState>,
sync_service: State<'_, CredentialSyncServiceState>,
uuid: String,
request: UpdateCredentialRequest,
) -> Result<ProviderCredential, String> {
// 如果需要重新上传文件,先处理文件上传
if let Some(new_file_path) = request.new_creds_file_path {
let credential = if let Some(new_file_path) = request.new_creds_file_path {
// 获取当前凭证以确定类型
let conn = db.lock().map_err(|e| e.to_string())?;
let current_credential = ProviderPoolDao::get_by_uuid(&conn, &uuid)
@@ -238,7 +262,7 @@ pub fn update_provider_pool_credential(
// 保存到数据库
ProviderPoolDao::update(&conn, &updated_cred).map_err(|e| e.to_string())?;
Ok(updated_cred)
updated_cred
} else {
// 常规更新,不涉及文件
pool_service.0.update_credential(
@@ -249,18 +273,49 @@ pub fn update_provider_pool_credential(
request.check_health,
request.check_model_name,
request.not_supported_models,
)
)?
};
// 同步到 YAML 配置(如果同步服务可用)
if let Some(ref sync) = sync_service.0 {
if let Err(e) = sync.update_credential(&credential) {
// 记录警告但不中断操作
tracing::warn!("同步凭证更新到 YAML 失败: {}", e);
}
}
Ok(credential)
}
/// 删除凭证
/// 删除凭证
///
/// 从数据库删除凭证,并同步到 YAML 配置文件
/// Requirements: 1.1, 1.2
#[tauri::command]
pub fn delete_provider_pool_credential(
db: State<'_, DbConnection>,
pool_service: State<'_, ProviderPoolServiceState>,
sync_service: State<'_, CredentialSyncServiceState>,
uuid: String,
provider_type: Option<String>,
) -> Result<bool, String> {
pool_service.0.delete_credential(&db, &uuid)
// 从数据库删除
let result = pool_service.0.delete_credential(&db, &uuid)?;
// 同步到 YAML 配置(如果同步服务可用且提供了 provider_type)
if let Some(ref sync) = sync_service.0 {
if let Some(pt) = provider_type {
if let Ok(pool_type) = pt.parse::<PoolProviderType>() {
if let Err(e) = sync.remove_credential(pool_type, &uuid) {
// 记录警告但不中断操作
tracing::warn!("从 YAML 删除凭证失败: {}", e);
}
}
}
}
Ok(result)
}
/// 切换凭证启用/禁用状态
+367
View File
@@ -177,3 +177,370 @@ pub async fn set_router_default_provider(_provider: ProviderType) -> Result<(),
// For now, just acknowledge the request
Ok(())
}
/// 推荐配置预设
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RecommendedPreset {
pub id: String,
pub name: String,
pub description: String,
pub aliases: Vec<ModelAlias>,
pub rules: Vec<RoutingRuleDto>,
}
/// 获取推荐配置列表
#[tauri::command]
pub async fn get_recommended_presets() -> Result<Vec<RecommendedPreset>, String> {
Ok(vec![
RecommendedPreset {
id: "claude-optimized".to_string(),
name: "Claude 优化配置".to_string(),
description: "将所有 Claude 模型请求路由到 Kiro,适合主要使用 Claude 的用户"
.to_string(),
aliases: vec![
// Claude 4.5 系列 (最新)
ModelAlias {
alias: "claude".to_string(),
actual: "claude-opus-4-5".to_string(),
},
ModelAlias {
alias: "opus".to_string(),
actual: "claude-opus-4-5".to_string(),
},
ModelAlias {
alias: "sonnet".to_string(),
actual: "claude-sonnet-4-5".to_string(),
},
ModelAlias {
alias: "haiku".to_string(),
actual: "claude-haiku-4-5".to_string(),
},
// Claude 4 系列
ModelAlias {
alias: "opus-4".to_string(),
actual: "claude-opus-4".to_string(),
},
ModelAlias {
alias: "sonnet-4".to_string(),
actual: "claude-sonnet-4".to_string(),
},
// Claude 3.7/3.5 系列 (旧版)
ModelAlias {
alias: "sonnet-3.7".to_string(),
actual: "claude-3-7-sonnet-latest".to_string(),
},
ModelAlias {
alias: "sonnet-3.5".to_string(),
actual: "claude-3-5-sonnet-latest".to_string(),
},
],
rules: vec![
RoutingRuleDto {
pattern: "claude-*".to_string(),
target_provider: ProviderType::Kiro,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "*sonnet*".to_string(),
target_provider: ProviderType::Kiro,
priority: 2,
enabled: true,
},
RoutingRuleDto {
pattern: "*opus*".to_string(),
target_provider: ProviderType::Kiro,
priority: 2,
enabled: true,
},
RoutingRuleDto {
pattern: "*haiku*".to_string(),
target_provider: ProviderType::Kiro,
priority: 2,
enabled: true,
},
],
},
RecommendedPreset {
id: "gemini-optimized".to_string(),
name: "Gemini 优化配置".to_string(),
description: "将 Gemini 模型请求路由到 Gemini Provider,适合主要使用 Google AI 的用户"
.to_string(),
aliases: vec![
// Gemini 3 系列 (最新)
ModelAlias {
alias: "gemini".to_string(),
actual: "gemini-3-pro".to_string(),
},
ModelAlias {
alias: "gemini-pro".to_string(),
actual: "gemini-3-pro".to_string(),
},
ModelAlias {
alias: "gemini-3".to_string(),
actual: "gemini-3-pro".to_string(),
},
// Gemini 2.5 系列
ModelAlias {
alias: "flash".to_string(),
actual: "gemini-2.5-flash".to_string(),
},
ModelAlias {
alias: "flash-lite".to_string(),
actual: "gemini-2.5-flash-lite".to_string(),
},
ModelAlias {
alias: "gemini-2.5".to_string(),
actual: "gemini-2.5-pro".to_string(),
},
],
rules: vec![
RoutingRuleDto {
pattern: "gemini-*".to_string(),
target_provider: ProviderType::Gemini,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "*flash*".to_string(),
target_provider: ProviderType::Gemini,
priority: 2,
enabled: true,
},
],
},
RecommendedPreset {
id: "multi-provider".to_string(),
name: "多 Provider 均衡配置".to_string(),
description: "根据模型名称自动路由到对应的 Provider,适合同时使用多个 AI 服务的用户"
.to_string(),
aliases: vec![
// Claude (最新)
ModelAlias {
alias: "claude".to_string(),
actual: "claude-opus-4-5".to_string(),
},
ModelAlias {
alias: "sonnet".to_string(),
actual: "claude-sonnet-4-5".to_string(),
},
// Gemini (最新)
ModelAlias {
alias: "gemini".to_string(),
actual: "gemini-3-pro".to_string(),
},
ModelAlias {
alias: "flash".to_string(),
actual: "gemini-2.5-flash".to_string(),
},
// Qwen
ModelAlias {
alias: "qwen".to_string(),
actual: "qwen3-coder-plus".to_string(),
},
// OpenAI (最新)
ModelAlias {
alias: "gpt".to_string(),
actual: "gpt-5.2".to_string(),
},
ModelAlias {
alias: "gpt-5".to_string(),
actual: "gpt-5.2".to_string(),
},
ModelAlias {
alias: "gpt-4".to_string(),
actual: "gpt-4o".to_string(),
},
ModelAlias {
alias: "o1".to_string(),
actual: "o1".to_string(),
},
ModelAlias {
alias: "o3".to_string(),
actual: "o3".to_string(),
},
],
rules: vec![
RoutingRuleDto {
pattern: "claude-*".to_string(),
target_provider: ProviderType::Kiro,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "gemini-*".to_string(),
target_provider: ProviderType::Gemini,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "qwen*".to_string(),
target_provider: ProviderType::Qwen,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "gpt-*".to_string(),
target_provider: ProviderType::OpenAI,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "o1*".to_string(),
target_provider: ProviderType::OpenAI,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "o3*".to_string(),
target_provider: ProviderType::OpenAI,
priority: 1,
enabled: true,
},
],
},
RecommendedPreset {
id: "coding-assistant".to_string(),
name: "编程助手配置".to_string(),
description:
"针对编程场景优化,Claude Opus 4.5 用于复杂代码,Gemini Flash 用于快速响应"
.to_string(),
aliases: vec![
ModelAlias {
alias: "code".to_string(),
actual: "claude-opus-4-5".to_string(),
},
ModelAlias {
alias: "coder".to_string(),
actual: "qwen3-coder-plus".to_string(),
},
ModelAlias {
alias: "fast".to_string(),
actual: "gemini-2.5-flash".to_string(),
},
ModelAlias {
alias: "think".to_string(),
actual: "claude-sonnet-4-5".to_string(),
},
],
rules: vec![
RoutingRuleDto {
pattern: "*coder*".to_string(),
target_provider: ProviderType::Qwen,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "claude-*".to_string(),
target_provider: ProviderType::Kiro,
priority: 2,
enabled: true,
},
RoutingRuleDto {
pattern: "gemini-*".to_string(),
target_provider: ProviderType::Gemini,
priority: 2,
enabled: true,
},
],
},
RecommendedPreset {
id: "cost-effective".to_string(),
name: "性价比优先配置".to_string(),
description: "优先使用免费或低成本的模型,适合预算有限的用户".to_string(),
aliases: vec![
ModelAlias {
alias: "default".to_string(),
actual: "gemini-2.5-flash".to_string(),
},
ModelAlias {
alias: "cheap".to_string(),
actual: "gemini-2.5-flash-lite".to_string(),
},
ModelAlias {
alias: "free".to_string(),
actual: "gemini-2.5-flash".to_string(),
},
],
rules: vec![
// 默认路由到 Gemini(免费额度高)
RoutingRuleDto {
pattern: "*".to_string(),
target_provider: ProviderType::Gemini,
priority: 100,
enabled: true,
},
// Claude 请求仍然路由到 Kiro
RoutingRuleDto {
pattern: "claude-*".to_string(),
target_provider: ProviderType::Kiro,
priority: 1,
enabled: true,
},
],
},
])
}
/// 应用推荐配置
#[tauri::command]
pub async fn apply_recommended_preset(
state: tauri::State<'_, RouterConfigState>,
preset_id: String,
merge: bool,
) -> Result<(), String> {
let presets = get_recommended_presets().await?;
let preset = presets
.into_iter()
.find(|p| p.id == preset_id)
.ok_or_else(|| format!("未找到预设配置: {}", preset_id))?;
// 应用别名
{
let mut aliases = state.aliases.write().await;
if !merge {
aliases.clear();
}
for alias in preset.aliases {
aliases.insert(alias.alias, alias.actual);
}
}
// 应用规则
{
let mut rules = state.rules.write().await;
if !merge {
rules.clear();
}
for rule in preset.rules {
// 避免重复
if !rules.iter().any(|r| r.pattern == rule.pattern) {
rules.push(rule);
}
}
// 按优先级排序
rules.sort_by(|a, b| a.priority.cmp(&b.priority));
}
Ok(())
}
/// 清空所有路由配置
#[tauri::command]
pub async fn clear_all_routing_config(
state: tauri::State<'_, RouterConfigState>,
) -> Result<(), String> {
{
let mut aliases = state.aliases.write().await;
aliases.clear();
}
{
let mut rules = state.rules.write().await;
rules.clear();
}
{
let mut exclusions = state.exclusions.write().await;
exclusions.clear();
}
Ok(())
}
+61 -19
View File
@@ -8,26 +8,58 @@ use crate::telemetry::{
};
use crate::ProviderType;
use chrono::{DateTime, Utc};
use parking_lot::RwLock;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
/// 遥测服务状态
///
/// 支持两种模式:
/// 1. 独立模式:使用自己的 StatsAggregator 和 TokenTracker 实例
/// 2. 共享模式:使用外部传入的共享实例(与 RequestProcessor 共享)
pub struct TelemetryState {
pub logger: Arc<RequestLogger>,
pub stats: Arc<StatsAggregator>,
pub tokens: Arc<TokenTracker>,
/// 统计聚合器(使用 RwLock 以支持与 RequestProcessor 共享)
pub stats: Arc<RwLock<StatsAggregator>>,
/// Token 追踪器(使用 RwLock 以支持与 RequestProcessor 共享)
pub tokens: Arc<RwLock<TokenTracker>>,
}
impl TelemetryState {
/// 创建独立的遥测状态(使用自己的实例)
pub fn new() -> Result<Self, String> {
let logger = RequestLogger::with_defaults()
.map_err(|e| format!("Failed to create logger: {}", e))?;
Ok(Self {
logger: Arc::new(logger),
stats: Arc::new(StatsAggregator::with_defaults()),
tokens: Arc::new(TokenTracker::with_defaults()),
stats: Arc::new(RwLock::new(StatsAggregator::with_defaults())),
tokens: Arc::new(RwLock::new(TokenTracker::with_defaults())),
})
}
/// 创建共享的遥测状态(使用外部传入的实例)
///
/// 这允许 TelemetryState 与 RequestProcessor 共享同一个 StatsAggregator、TokenTracker 和 RequestLogger,
/// 使得请求处理过程中记录的统计数据能够在前端监控页面中显示。
pub fn with_shared(
stats: Arc<RwLock<StatsAggregator>>,
tokens: Arc<RwLock<TokenTracker>>,
logger: Option<Arc<RequestLogger>>,
) -> Result<Self, String> {
let logger = match logger {
Some(l) => l,
None => Arc::new(
RequestLogger::with_defaults()
.map_err(|e| format!("Failed to create logger: {}", e))?,
),
};
Ok(Self {
logger,
stats,
tokens,
})
}
}
@@ -150,7 +182,8 @@ pub async fn get_stats_summary(
time_range: Option<TimeRangeParam>,
) -> Result<StatsSummary, String> {
let range = time_range.map(|r| r.to_time_range()).transpose()?.flatten();
Ok(state.stats.summary(range))
let stats = state.stats.read();
Ok(stats.summary(range))
}
/// 按 Provider 分组统计
@@ -160,7 +193,8 @@ pub async fn get_stats_by_provider(
time_range: Option<TimeRangeParam>,
) -> Result<HashMap<String, ProviderStats>, String> {
let range = time_range.map(|r| r.to_time_range()).transpose()?.flatten();
let stats = state.stats.by_provider(range);
let stats_guard = state.stats.read();
let stats = stats_guard.by_provider(range);
// 转换 key 为 String
Ok(stats.into_iter().map(|(k, v)| (k.to_string(), v)).collect())
@@ -173,7 +207,8 @@ pub async fn get_stats_by_model(
time_range: Option<TimeRangeParam>,
) -> Result<HashMap<String, ModelStats>, String> {
let range = time_range.map(|r| r.to_time_range()).transpose()?.flatten();
Ok(state.stats.by_model(range))
let stats = state.stats.read();
Ok(stats.by_model(range))
}
// ========== Token 统计命令 ==========
@@ -194,7 +229,8 @@ pub async fn get_token_summary(
}
None => (None, None),
};
Ok(state.tokens.summary(start, end))
let tokens = state.tokens.read();
Ok(tokens.summary(start, end))
}
/// 按 Provider 分组 Token 统计
@@ -213,7 +249,8 @@ pub async fn get_token_stats_by_provider(
}
None => (None, None),
};
let stats = state.tokens.by_provider(start, end);
let tokens = state.tokens.read();
let stats = tokens.by_provider(start, end);
Ok(stats.into_iter().map(|(k, v)| (k.to_string(), v)).collect())
}
@@ -234,7 +271,8 @@ pub async fn get_token_stats_by_model(
}
None => (None, None),
};
Ok(state.tokens.by_model(start, end))
let tokens = state.tokens.read();
Ok(tokens.by_model(start, end))
}
/// 按天汇总 Token 统计
@@ -243,7 +281,8 @@ pub async fn get_token_stats_by_day(
state: tauri::State<'_, TelemetryState>,
days: Option<i64>,
) -> Result<Vec<crate::telemetry::PeriodTokenStats>, String> {
Ok(state.tokens.by_day(days.unwrap_or(7)))
let tokens = state.tokens.read();
Ok(tokens.by_day(days.unwrap_or(7)))
}
// ========== 仪表盘数据命令 ==========
@@ -269,18 +308,21 @@ pub async fn get_dashboard_data(
// 获取最近 24 小时的统计
let range = Some(TimeRange::last_hours(24));
let stats = state.stats.summary(range);
let tokens = state.tokens.summary(
Some(Utc::now() - chrono::Duration::hours(24)),
Some(Utc::now()),
);
let by_provider: HashMap<String, ProviderStats> = state
.stats
let stats_guard = state.stats.read();
let stats = stats_guard.summary(range);
let by_provider: HashMap<String, ProviderStats> = stats_guard
.by_provider(range)
.into_iter()
.map(|(k, v)| (k.to_string(), v))
.collect();
drop(stats_guard);
let tokens_guard = state.tokens.read();
let tokens = tokens_guard.summary(
Some(Utc::now() - chrono::Duration::hours(24)),
Some(Utc::now()),
);
drop(tokens_guard);
// 获取最近 20 条日志
let mut recent_logs = state.logger.get_all();
+728
View File
@@ -0,0 +1,728 @@
//! 配置导出服务
//!
//! 提供配置和凭证的统一导出功能,支持:
//! - 仅配置导出(YAML 格式)
//! - 仅凭证导出
//! - 完整导出(配置 + 凭证 + OAuth Token 文件)
//! - 敏感信息脱敏
use super::path_utils::expand_tilde;
use super::types::{ApiKeyEntry, Config, CredentialEntry, CredentialPoolConfig};
use super::yaml::{ConfigError, ConfigManager};
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
/// 导出选项
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExportOptions {
/// 是否包含配置
pub include_config: bool,
/// 是否包含凭证
pub include_credentials: bool,
/// 是否脱敏敏感信息
pub redact_secrets: bool,
}
impl Default for ExportOptions {
fn default() -> Self {
Self {
include_config: true,
include_credentials: true,
redact_secrets: false,
}
}
}
impl ExportOptions {
/// 创建仅配置导出选项
pub fn config_only() -> Self {
Self {
include_config: true,
include_credentials: false,
redact_secrets: false,
}
}
/// 创建仅凭证导出选项
pub fn credentials_only() -> Self {
Self {
include_config: false,
include_credentials: true,
redact_secrets: false,
}
}
/// 创建完整导出选项
pub fn full() -> Self {
Self {
include_config: true,
include_credentials: true,
redact_secrets: false,
}
}
/// 创建脱敏导出选项
pub fn redacted() -> Self {
Self {
include_config: true,
include_credentials: true,
redact_secrets: true,
}
}
}
/// 导出包
///
/// 包含配置和凭证的统一导出格式
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExportBundle {
/// 导出格式版本号
pub version: String,
/// 导出时间
pub exported_at: DateTime<Utc>,
/// 应用版本
pub app_version: String,
/// YAML 配置内容(如果包含配置)
#[serde(skip_serializing_if = "Option::is_none")]
pub config_yaml: Option<String>,
/// OAuth Token 文件(base64 编码)
/// key: 相对于 auth_dir 的路径
/// value: base64 编码的文件内容
#[serde(default, skip_serializing_if = "HashMap::is_empty")]
pub token_files: HashMap<String, String>,
/// 是否已脱敏
pub redacted: bool,
}
impl ExportBundle {
/// 当前导出格式版本
pub const CURRENT_VERSION: &'static str = "1.0";
/// 创建新的导出包
pub fn new(app_version: &str) -> Self {
Self {
version: Self::CURRENT_VERSION.to_string(),
exported_at: Utc::now(),
app_version: app_version.to_string(),
config_yaml: None,
token_files: HashMap::new(),
redacted: false,
}
}
/// 检查是否包含配置
pub fn has_config(&self) -> bool {
self.config_yaml.is_some()
}
/// 检查是否包含凭证
pub fn has_credentials(&self) -> bool {
!self.token_files.is_empty()
}
/// 检查是否已脱敏
pub fn is_redacted(&self) -> bool {
self.redacted
}
/// 序列化为 JSON 字符串
pub fn to_json(&self) -> Result<String, ExportError> {
serde_json::to_string_pretty(self).map_err(|e| ExportError::SerializeError(e.to_string()))
}
/// 从 JSON 字符串反序列化
pub fn from_json(json: &str) -> Result<Self, ExportError> {
serde_json::from_str(json).map_err(|e| ExportError::ParseError(e.to_string()))
}
}
/// 导出错误类型
#[derive(Debug, Clone)]
pub enum ExportError {
/// 配置错误
ConfigError(String),
/// 文件读取错误
ReadError(String),
/// 序列化错误
SerializeError(String),
/// 解析错误
ParseError(String),
/// Token 文件不存在
TokenFileNotFound(String),
}
impl std::fmt::Display for ExportError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ExportError::ConfigError(msg) => write!(f, "配置错误: {}", msg),
ExportError::ReadError(msg) => write!(f, "文件读取错误: {}", msg),
ExportError::SerializeError(msg) => write!(f, "序列化错误: {}", msg),
ExportError::ParseError(msg) => write!(f, "解析错误: {}", msg),
ExportError::TokenFileNotFound(path) => write!(f, "Token 文件不存在: {}", path),
}
}
}
impl std::error::Error for ExportError {}
impl From<ConfigError> for ExportError {
fn from(err: ConfigError) -> Self {
ExportError::ConfigError(err.to_string())
}
}
/// 脱敏占位符
pub const REDACTED_PLACEHOLDER: &str = "***REDACTED***";
/// 导出服务
///
/// 提供配置和凭证的统一导出功能
pub struct ExportService;
impl ExportService {
/// 导出配置为 YAML 字符串
///
/// # Arguments
/// * `config` - 要导出的配置
/// * `redact` - 是否脱敏敏感信息
///
/// # Returns
/// * `Ok(String)` - YAML 格式的配置字符串
/// * `Err(ExportError)` - 导出失败
pub fn export_yaml(config: &Config, redact: bool) -> Result<String, ExportError> {
let config_to_export = if redact {
Self::redact_config(config)
} else {
config.clone()
};
ConfigManager::to_yaml(&config_to_export).map_err(ExportError::from)
}
/// 导出完整的配置和凭证包
///
/// # Arguments
/// * `config` - 要导出的配置
/// * `options` - 导出选项
/// * `app_version` - 应用版本
///
/// # Returns
/// * `Ok(ExportBundle)` - 导出包
/// * `Err(ExportError)` - 导出失败
pub fn export(
config: &Config,
options: &ExportOptions,
app_version: &str,
) -> Result<ExportBundle, ExportError> {
let mut bundle = ExportBundle::new(app_version);
bundle.redacted = options.redact_secrets;
// 导出配置
if options.include_config {
let yaml = Self::export_yaml(config, options.redact_secrets)?;
bundle.config_yaml = Some(yaml);
}
// 导出凭证(OAuth Token 文件)
if options.include_credentials {
let token_files = Self::collect_token_files(config, options.redact_secrets)?;
bundle.token_files = token_files;
}
Ok(bundle)
}
/// 收集 OAuth Token 文件
///
/// 从 auth_dir 目录收集所有 OAuth 凭证的 token 文件
fn collect_token_files(
config: &Config,
redact: bool,
) -> Result<HashMap<String, String>, ExportError> {
let mut token_files = HashMap::new();
let auth_dir = expand_tilde(&config.auth_dir);
// 收集所有 OAuth 凭证的 token 文件
let oauth_credentials = Self::get_oauth_credentials(&config.credential_pool);
for entry in oauth_credentials {
let token_path = auth_dir.join(&entry.token_file);
if token_path.exists() {
let content = std::fs::read(&token_path)
.map_err(|e| ExportError::ReadError(format!("{}: {}", entry.token_file, e)))?;
let encoded = if redact {
// 脱敏:用占位符替换实际内容
base64::encode(REDACTED_PLACEHOLDER.as_bytes())
} else {
base64::encode(&content)
};
token_files.insert(entry.token_file.clone(), encoded);
}
// 如果文件不存在,跳过(不报错,只是不包含在导出中)
}
Ok(token_files)
}
/// 获取所有 OAuth 凭证条目
fn get_oauth_credentials(pool: &CredentialPoolConfig) -> Vec<&CredentialEntry> {
let mut credentials = Vec::new();
credentials.extend(pool.kiro.iter());
credentials.extend(pool.gemini.iter());
credentials.extend(pool.qwen.iter());
credentials
}
/// 脱敏配置
///
/// 将配置中的敏感信息替换为占位符
pub fn redact_config(config: &Config) -> Config {
let mut redacted = config.clone();
// 脱敏服务器 API 密钥
redacted.server.api_key = REDACTED_PLACEHOLDER.to_string();
// 脱敏 Provider API 密钥
if redacted.providers.openai.api_key.is_some() {
redacted.providers.openai.api_key = Some(REDACTED_PLACEHOLDER.to_string());
}
if redacted.providers.claude.api_key.is_some() {
redacted.providers.claude.api_key = Some(REDACTED_PLACEHOLDER.to_string());
}
// 脱敏凭证池中的 API Key
redacted.credential_pool = Self::redact_credential_pool(&config.credential_pool);
redacted
}
/// 脱敏凭证池
fn redact_credential_pool(pool: &CredentialPoolConfig) -> CredentialPoolConfig {
CredentialPoolConfig {
kiro: pool.kiro.clone(),
gemini: pool.gemini.clone(),
qwen: pool.qwen.clone(),
openai: pool
.openai
.iter()
.map(|entry| ApiKeyEntry {
id: entry.id.clone(),
api_key: REDACTED_PLACEHOLDER.to_string(),
base_url: entry.base_url.clone(),
disabled: entry.disabled,
})
.collect(),
claude: pool
.claude
.iter()
.map(|entry| ApiKeyEntry {
id: entry.id.clone(),
api_key: REDACTED_PLACEHOLDER.to_string(),
base_url: entry.base_url.clone(),
disabled: entry.disabled,
})
.collect(),
}
}
/// 检查配置是否包含敏感信息
///
/// 用于验证脱敏是否完整
pub fn contains_secrets(config: &Config) -> bool {
// 检查服务器 API 密钥
if !config.server.api_key.is_empty() && config.server.api_key != REDACTED_PLACEHOLDER {
return true;
}
// 检查 Provider API 密钥
if let Some(ref key) = config.providers.openai.api_key {
if !key.is_empty() && key != REDACTED_PLACEHOLDER {
return true;
}
}
if let Some(ref key) = config.providers.claude.api_key {
if !key.is_empty() && key != REDACTED_PLACEHOLDER {
return true;
}
}
// 检查凭证池中的 API Key
for entry in &config.credential_pool.openai {
if !entry.api_key.is_empty() && entry.api_key != REDACTED_PLACEHOLDER {
return true;
}
}
for entry in &config.credential_pool.claude {
if !entry.api_key.is_empty() && entry.api_key != REDACTED_PLACEHOLDER {
return true;
}
}
false
}
/// 检查 YAML 字符串是否包含敏感信息
pub fn yaml_contains_secrets(yaml: &str) -> bool {
// 检查是否包含看起来像 API 密钥的模式
let secret_patterns = [
"sk-", // OpenAI API key prefix
"sk-ant-", // Anthropic API key prefix
"api_key:", // API key field (if not redacted)
];
for pattern in &secret_patterns {
if yaml.contains(pattern) && !yaml.contains(REDACTED_PLACEHOLDER) {
// 进一步检查是否是实际的密钥值
for line in yaml.lines() {
if line.contains(pattern) && !line.contains(REDACTED_PLACEHOLDER) {
// 排除注释行
let trimmed = line.trim();
if !trimmed.starts_with('#') {
return true;
}
}
}
}
}
false
}
}
// 简单的 base64 编码/解码模块
mod base64 {
const ALPHABET: &[u8] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
pub fn encode(data: &[u8]) -> String {
let mut result = String::new();
let mut i = 0;
while i < data.len() {
let b0 = data[i] as u32;
let b1 = if i + 1 < data.len() {
data[i + 1] as u32
} else {
0
};
let b2 = if i + 2 < data.len() {
data[i + 2] as u32
} else {
0
};
let triple = (b0 << 16) | (b1 << 8) | b2;
result.push(ALPHABET[((triple >> 18) & 0x3F) as usize] as char);
result.push(ALPHABET[((triple >> 12) & 0x3F) as usize] as char);
if i + 1 < data.len() {
result.push(ALPHABET[((triple >> 6) & 0x3F) as usize] as char);
} else {
result.push('=');
}
if i + 2 < data.len() {
result.push(ALPHABET[(triple & 0x3F) as usize] as char);
} else {
result.push('=');
}
i += 3;
}
result
}
pub fn decode(data: &str) -> Result<Vec<u8>, String> {
let data = data.trim_end_matches('=');
let mut result = Vec::new();
let decode_char = |c: char| -> Result<u32, String> {
match c {
'A'..='Z' => Ok((c as u32) - ('A' as u32)),
'a'..='z' => Ok((c as u32) - ('a' as u32) + 26),
'0'..='9' => Ok((c as u32) - ('0' as u32) + 52),
'+' => Ok(62),
'/' => Ok(63),
_ => Err(format!("Invalid base64 character: {}", c)),
}
};
let chars: Vec<char> = data.chars().collect();
let mut i = 0;
while i < chars.len() {
let c0 = decode_char(chars[i])?;
let c1 = if i + 1 < chars.len() {
decode_char(chars[i + 1])?
} else {
0
};
let c2 = if i + 2 < chars.len() {
decode_char(chars[i + 2])?
} else {
0
};
let c3 = if i + 3 < chars.len() {
decode_char(chars[i + 3])?
} else {
0
};
let triple = (c0 << 18) | (c1 << 12) | (c2 << 6) | c3;
result.push(((triple >> 16) & 0xFF) as u8);
if i + 2 < chars.len() {
result.push(((triple >> 8) & 0xFF) as u8);
}
if i + 3 < chars.len() {
result.push((triple & 0xFF) as u8);
}
i += 4;
}
Ok(result)
}
}
pub use self::base64::{decode as base64_decode, encode as base64_encode};
#[cfg(test)]
mod unit_tests {
use super::*;
#[test]
fn test_export_options_default() {
let options = ExportOptions::default();
assert!(options.include_config);
assert!(options.include_credentials);
assert!(!options.redact_secrets);
}
#[test]
fn test_export_options_config_only() {
let options = ExportOptions::config_only();
assert!(options.include_config);
assert!(!options.include_credentials);
assert!(!options.redact_secrets);
}
#[test]
fn test_export_options_credentials_only() {
let options = ExportOptions::credentials_only();
assert!(!options.include_config);
assert!(options.include_credentials);
assert!(!options.redact_secrets);
}
#[test]
fn test_export_options_full() {
let options = ExportOptions::full();
assert!(options.include_config);
assert!(options.include_credentials);
assert!(!options.redact_secrets);
}
#[test]
fn test_export_options_redacted() {
let options = ExportOptions::redacted();
assert!(options.include_config);
assert!(options.include_credentials);
assert!(options.redact_secrets);
}
#[test]
fn test_export_bundle_new() {
let bundle = ExportBundle::new("1.0.0");
assert_eq!(bundle.version, ExportBundle::CURRENT_VERSION);
assert_eq!(bundle.app_version, "1.0.0");
assert!(!bundle.redacted);
assert!(bundle.config_yaml.is_none());
assert!(bundle.token_files.is_empty());
}
#[test]
fn test_export_bundle_has_config() {
let mut bundle = ExportBundle::new("1.0.0");
assert!(!bundle.has_config());
bundle.config_yaml = Some("server:\n port: 8999".to_string());
assert!(bundle.has_config());
}
#[test]
fn test_export_bundle_has_credentials() {
let mut bundle = ExportBundle::new("1.0.0");
assert!(!bundle.has_credentials());
bundle
.token_files
.insert("kiro/token.json".to_string(), "base64data".to_string());
assert!(bundle.has_credentials());
}
#[test]
fn test_export_bundle_json_roundtrip() {
let mut bundle = ExportBundle::new("1.0.0");
bundle.config_yaml = Some("server:\n port: 8999".to_string());
bundle
.token_files
.insert("kiro/token.json".to_string(), "dGVzdA==".to_string());
bundle.redacted = true;
let json = bundle.to_json().expect("序列化应成功");
let parsed = ExportBundle::from_json(&json).expect("反序列化应成功");
assert_eq!(parsed.version, bundle.version);
assert_eq!(parsed.app_version, bundle.app_version);
assert_eq!(parsed.config_yaml, bundle.config_yaml);
assert_eq!(parsed.token_files, bundle.token_files);
assert_eq!(parsed.redacted, bundle.redacted);
}
#[test]
fn test_export_yaml_without_redaction() {
let config = Config::default();
let yaml = ExportService::export_yaml(&config, false).expect("导出应成功");
assert!(yaml.contains("server:"));
assert!(yaml.contains("port: 8999"));
assert!(yaml.contains("api_key: proxy_cast"));
}
#[test]
fn test_export_yaml_with_redaction() {
let mut config = Config::default();
config.server.api_key = "secret-key".to_string();
config.providers.openai.api_key = Some("sk-openai-secret".to_string());
let yaml = ExportService::export_yaml(&config, true).expect("导出应成功");
assert!(yaml.contains(REDACTED_PLACEHOLDER));
assert!(!yaml.contains("secret-key"));
assert!(!yaml.contains("sk-openai-secret"));
}
#[test]
fn test_redact_config() {
let mut config = Config::default();
config.server.api_key = "secret-key".to_string();
config.providers.openai.api_key = Some("sk-openai-secret".to_string());
config.providers.claude.api_key = Some("sk-ant-claude-secret".to_string());
config.credential_pool.openai.push(ApiKeyEntry {
id: "openai-1".to_string(),
api_key: "sk-pool-key".to_string(),
base_url: None,
disabled: false,
});
let redacted = ExportService::redact_config(&config);
assert_eq!(redacted.server.api_key, REDACTED_PLACEHOLDER);
assert_eq!(
redacted.providers.openai.api_key,
Some(REDACTED_PLACEHOLDER.to_string())
);
assert_eq!(
redacted.providers.claude.api_key,
Some(REDACTED_PLACEHOLDER.to_string())
);
assert_eq!(
redacted.credential_pool.openai[0].api_key,
REDACTED_PLACEHOLDER
);
}
#[test]
fn test_contains_secrets() {
let mut config = Config::default();
config.server.api_key = "secret-key".to_string();
assert!(ExportService::contains_secrets(&config));
let redacted = ExportService::redact_config(&config);
assert!(!ExportService::contains_secrets(&redacted));
}
#[test]
fn test_contains_secrets_with_credential_pool() {
let mut config = Config::default();
config.server.api_key = REDACTED_PLACEHOLDER.to_string();
config.credential_pool.openai.push(ApiKeyEntry {
id: "openai-1".to_string(),
api_key: "sk-real-key".to_string(),
base_url: None,
disabled: false,
});
assert!(ExportService::contains_secrets(&config));
}
#[test]
fn test_export_config_only() {
let config = Config::default();
let options = ExportOptions::config_only();
let bundle = ExportService::export(&config, &options, "1.0.0").expect("导出应成功");
assert!(bundle.has_config());
assert!(!bundle.has_credentials());
assert!(!bundle.redacted);
}
#[test]
fn test_export_credentials_only() {
let config = Config::default();
let options = ExportOptions::credentials_only();
let bundle = ExportService::export(&config, &options, "1.0.0").expect("导出应成功");
assert!(!bundle.has_config());
// 默认配置没有凭证,所以 token_files 为空
assert!(!bundle.has_credentials());
}
#[test]
fn test_base64_encode_decode() {
let original = b"Hello, World!";
let encoded = base64_encode(original);
let decoded = base64_decode(&encoded).expect("解码应成功");
assert_eq!(decoded, original);
}
#[test]
fn test_base64_encode_empty() {
let original = b"";
let encoded = base64_encode(original);
let decoded = base64_decode(&encoded).expect("解码应成功");
assert_eq!(decoded, original);
}
#[test]
fn test_base64_encode_various_lengths() {
// 测试不同长度的输入(1, 2, 3 字节边界情况)
for len in 1..=10 {
let original: Vec<u8> = (0..len).map(|i| i as u8).collect();
let encoded = base64_encode(&original);
let decoded = base64_decode(&encoded).expect("解码应成功");
assert_eq!(decoded, original, "长度 {} 的数据往返失败", len);
}
}
#[test]
fn test_export_error_display() {
let err = ExportError::ConfigError("test error".to_string());
assert!(err.to_string().contains("配置错误"));
assert!(err.to_string().contains("test error"));
let err = ExportError::TokenFileNotFound("/path/to/token.json".to_string());
assert!(err.to_string().contains("Token 文件不存在"));
}
}
+770
View File
@@ -0,0 +1,770 @@
//! 配置导入服务
//!
//! 提供配置和凭证的统一导入功能,支持:
//! - YAML 配置导入
//! - 完整导入包导入(配置 + 凭证 + OAuth Token 文件)
//! - 导入验证(格式、版本、脱敏状态)
//! - 合并和替换模式
use super::export::{base64_decode, ExportBundle, REDACTED_PLACEHOLDER};
use super::path_utils::expand_tilde;
use super::types::{ApiKeyEntry, Config, CredentialEntry, CredentialPoolConfig};
use super::yaml::{ConfigError, ConfigManager, YamlService};
use serde::{Deserialize, Serialize};
use std::path::Path;
/// 导入选项
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImportOptions {
/// 是否合并(false 则替换)
pub merge: bool,
}
impl Default for ImportOptions {
fn default() -> Self {
Self { merge: true }
}
}
impl ImportOptions {
/// 创建合并模式选项
pub fn merge() -> Self {
Self { merge: true }
}
/// 创建替换模式选项
pub fn replace() -> Self {
Self { merge: false }
}
}
/// 验证结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ValidationResult {
/// 是否有效
pub valid: bool,
/// 格式版本
pub version: Option<String>,
/// 是否已脱敏
pub redacted: bool,
/// 是否包含配置
pub has_config: bool,
/// 是否包含凭证
pub has_credentials: bool,
/// 错误信息列表
pub errors: Vec<String>,
/// 警告信息列表
pub warnings: Vec<String>,
}
impl ValidationResult {
/// 创建有效的验证结果
pub fn valid() -> Self {
Self {
valid: true,
version: None,
redacted: false,
has_config: false,
has_credentials: false,
errors: Vec::new(),
warnings: Vec::new(),
}
}
/// 创建无效的验证结果
pub fn invalid(error: impl Into<String>) -> Self {
Self {
valid: false,
version: None,
redacted: false,
has_config: false,
has_credentials: false,
errors: vec![error.into()],
warnings: Vec::new(),
}
}
/// 添加错误
pub fn add_error(&mut self, error: impl Into<String>) {
self.errors.push(error.into());
self.valid = false;
}
/// 添加警告
pub fn add_warning(&mut self, warning: impl Into<String>) {
self.warnings.push(warning.into());
}
}
/// 导入结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImportResult {
/// 是否成功
pub success: bool,
/// 警告信息
pub warnings: Vec<String>,
/// 导入的配置
pub config: Config,
}
impl ImportResult {
/// 创建成功的导入结果
pub fn success(config: Config) -> Self {
Self {
success: true,
warnings: Vec::new(),
config,
}
}
/// 创建带警告的成功导入结果
pub fn success_with_warnings(config: Config, warnings: Vec<String>) -> Self {
Self {
success: true,
warnings,
config,
}
}
}
/// 导入错误类型
#[derive(Debug, Clone)]
pub enum ImportError {
/// 格式错误
FormatError(String),
/// 版本不兼容
VersionError(String),
/// 配置错误
ConfigError(String),
/// IO 错误
IoError(String),
/// 验证错误
ValidationError(String),
/// 脱敏数据无法导入
RedactedDataError(String),
}
impl std::fmt::Display for ImportError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ImportError::FormatError(msg) => write!(f, "格式错误: {}", msg),
ImportError::VersionError(msg) => write!(f, "版本不兼容: {}", msg),
ImportError::ConfigError(msg) => write!(f, "配置错误: {}", msg),
ImportError::IoError(msg) => write!(f, "IO 错误: {}", msg),
ImportError::ValidationError(msg) => write!(f, "验证错误: {}", msg),
ImportError::RedactedDataError(msg) => write!(f, "脱敏数据无法导入: {}", msg),
}
}
}
impl std::error::Error for ImportError {}
impl From<ConfigError> for ImportError {
fn from(err: ConfigError) -> Self {
ImportError::ConfigError(err.to_string())
}
}
impl From<std::io::Error> for ImportError {
fn from(err: std::io::Error) -> Self {
ImportError::IoError(err.to_string())
}
}
/// 导入服务
///
/// 提供配置和凭证的统一导入功能
pub struct ImportService;
impl ImportService {
/// 支持的导入格式版本
pub const SUPPORTED_VERSIONS: &'static [&'static str] = &["1.0"];
/// 验证导入内容
///
/// # Arguments
/// * `content` - 导入内容(JSON 格式的 ExportBundle 或 YAML 配置)
///
/// # Returns
/// * `ValidationResult` - 验证结果
pub fn validate(content: &str) -> ValidationResult {
// 首先尝试解析为 ExportBundle (JSON)
if let Ok(bundle) = ExportBundle::from_json(content) {
return Self::validate_bundle(&bundle);
}
// 尝试解析为 YAML 配置
if let Ok(_config) = ConfigManager::parse_yaml(content) {
let mut result = ValidationResult::valid();
result.has_config = true;
result.has_credentials = false;
result.version = Some("yaml".to_string());
return result;
}
ValidationResult::invalid(
"无法解析导入内容:既不是有效的 JSON 导出包,也不是有效的 YAML 配置",
)
}
/// 验证导出包
fn validate_bundle(bundle: &ExportBundle) -> ValidationResult {
let mut result = ValidationResult::valid();
result.version = Some(bundle.version.clone());
result.redacted = bundle.redacted;
result.has_config = bundle.has_config();
result.has_credentials = bundle.has_credentials();
// 检查版本兼容性
if !Self::SUPPORTED_VERSIONS.contains(&bundle.version.as_str()) {
result.add_warning(format!(
"导出包版本 {} 可能不完全兼容,支持的版本: {:?}",
bundle.version,
Self::SUPPORTED_VERSIONS
));
}
// 检查脱敏状态
if bundle.redacted {
result.add_warning("导出包已脱敏,凭证数据无法恢复");
}
// 验证配置内容(如果存在)
if let Some(ref yaml) = bundle.config_yaml {
if let Err(e) = ConfigManager::parse_yaml(yaml) {
result.add_error(format!("配置 YAML 解析失败: {}", e));
}
}
result
}
/// 导入 YAML 配置
///
/// # Arguments
/// * `yaml` - YAML 配置字符串
/// * `current_config` - 当前配置(用于合并模式)
/// * `options` - 导入选项
///
/// # Returns
/// * `Ok(ImportResult)` - 导入成功
/// * `Err(ImportError)` - 导入失败
pub fn import_yaml(
yaml: &str,
current_config: &Config,
options: &ImportOptions,
) -> Result<ImportResult, ImportError> {
// 解析 YAML
let imported_config = ConfigManager::parse_yaml(yaml)?;
// 根据选项合并或替换
let final_config = if options.merge {
Self::merge_configs(current_config, &imported_config)
} else {
imported_config
};
Ok(ImportResult::success(final_config))
}
/// 导入完整的导出包
///
/// # Arguments
/// * `bundle` - 导出包
/// * `current_config` - 当前配置(用于合并模式)
/// * `options` - 导入选项
/// * `auth_dir` - 认证目录路径(用于恢复 OAuth token 文件)
///
/// # Returns
/// * `Ok(ImportResult)` - 导入成功
/// * `Err(ImportError)` - 导入失败
pub fn import(
bundle: &ExportBundle,
current_config: &Config,
options: &ImportOptions,
auth_dir: &str,
) -> Result<ImportResult, ImportError> {
let mut warnings = Vec::new();
// 检查脱敏状态
if bundle.redacted {
warnings.push("导出包已脱敏,凭证数据将使用占位符".to_string());
}
// 导入配置
let mut config = if let Some(ref yaml) = bundle.config_yaml {
let imported = ConfigManager::parse_yaml(yaml)?;
if options.merge {
Self::merge_configs(current_config, &imported)
} else {
imported
}
} else if options.merge {
current_config.clone()
} else {
Config::default()
};
// 恢复 OAuth token 文件
if !bundle.token_files.is_empty() {
let token_warnings = Self::restore_token_files(&bundle.token_files, auth_dir)?;
warnings.extend(token_warnings);
}
// 如果是脱敏数据,清理凭证池中的占位符
if bundle.redacted {
Self::clean_redacted_credentials(&mut config);
}
Ok(ImportResult::success_with_warnings(config, warnings))
}
/// 合并配置
///
/// 将导入的配置合并到当前配置中
fn merge_configs(current: &Config, imported: &Config) -> Config {
let mut merged = current.clone();
// 合并服务器配置(导入的覆盖当前的)
merged.server = imported.server.clone();
// 合并 Provider 配置
merged.providers = imported.providers.clone();
// 合并路由配置
merged.routing = imported.routing.clone();
merged.default_provider = imported.default_provider.clone();
// 合并重试配置
merged.retry = imported.retry.clone();
// 合并日志配置
merged.logging = imported.logging.clone();
// 合并注入配置
merged.injection = imported.injection.clone();
// 合并 auth_dir
merged.auth_dir = imported.auth_dir.clone();
// 合并凭证池(添加新的,保留现有的)
merged.credential_pool =
Self::merge_credential_pools(&current.credential_pool, &imported.credential_pool);
merged
}
/// 合并凭证池
///
/// 将导入的凭证添加到现有凭证池中(按 ID 去重)
fn merge_credential_pools(
current: &CredentialPoolConfig,
imported: &CredentialPoolConfig,
) -> CredentialPoolConfig {
CredentialPoolConfig {
kiro: Self::merge_credential_entries(&current.kiro, &imported.kiro),
gemini: Self::merge_credential_entries(&current.gemini, &imported.gemini),
qwen: Self::merge_credential_entries(&current.qwen, &imported.qwen),
openai: Self::merge_api_key_entries(&current.openai, &imported.openai),
claude: Self::merge_api_key_entries(&current.claude, &imported.claude),
}
}
/// 合并 OAuth 凭证条目
fn merge_credential_entries(
current: &[CredentialEntry],
imported: &[CredentialEntry],
) -> Vec<CredentialEntry> {
let mut result: Vec<CredentialEntry> = current.to_vec();
for entry in imported {
// 如果 ID 已存在,更新;否则添加
if let Some(existing) = result.iter_mut().find(|e| e.id == entry.id) {
*existing = entry.clone();
} else {
result.push(entry.clone());
}
}
result
}
/// 合并 API Key 凭证条目
fn merge_api_key_entries(
current: &[ApiKeyEntry],
imported: &[ApiKeyEntry],
) -> Vec<ApiKeyEntry> {
let mut result: Vec<ApiKeyEntry> = current.to_vec();
for entry in imported {
// 跳过脱敏的条目
if entry.api_key == REDACTED_PLACEHOLDER {
continue;
}
// 如果 ID 已存在,更新;否则添加
if let Some(existing) = result.iter_mut().find(|e| e.id == entry.id) {
*existing = entry.clone();
} else {
result.push(entry.clone());
}
}
result
}
/// 恢复 OAuth token 文件到 auth_dir
///
/// # Arguments
/// * `token_files` - token 文件映射(相对路径 -> base64 编码内容)
/// * `auth_dir` - 认证目录路径
///
/// # Returns
/// * `Ok(Vec<String>)` - 警告信息列表
fn restore_token_files(
token_files: &std::collections::HashMap<String, String>,
auth_dir: &str,
) -> Result<Vec<String>, ImportError> {
let mut warnings = Vec::new();
let auth_path = expand_tilde(auth_dir);
// 确保 auth_dir 存在
std::fs::create_dir_all(&auth_path)?;
for (relative_path, base64_content) in token_files {
let token_path = auth_path.join(relative_path);
// 确保父目录存在
if let Some(parent) = token_path.parent() {
std::fs::create_dir_all(parent)?;
}
// 解码 base64 内容
match base64_decode(base64_content) {
Ok(content) => {
// 检查是否是脱敏内容
if content == REDACTED_PLACEHOLDER.as_bytes() {
warnings.push(format!("Token 文件 {} 已脱敏,无法恢复", relative_path));
continue;
}
// 写入文件
if let Err(e) = std::fs::write(&token_path, &content) {
warnings.push(format!("写入 token 文件 {} 失败: {}", relative_path, e));
}
}
Err(e) => {
warnings.push(format!("解码 token 文件 {} 失败: {}", relative_path, e));
}
}
}
Ok(warnings)
}
/// 清理脱敏的凭证数据
///
/// 移除凭证池中使用占位符的条目
fn clean_redacted_credentials(config: &mut Config) {
// 清理 OpenAI 凭证池中的脱敏条目
config
.credential_pool
.openai
.retain(|e| e.api_key != REDACTED_PLACEHOLDER);
// 清理 Claude 凭证池中的脱敏条目
config
.credential_pool
.claude
.retain(|e| e.api_key != REDACTED_PLACEHOLDER);
// 清理 Provider 配置中的脱敏 API 密钥
if config.providers.openai.api_key.as_deref() == Some(REDACTED_PLACEHOLDER) {
config.providers.openai.api_key = None;
}
if config.providers.claude.api_key.as_deref() == Some(REDACTED_PLACEHOLDER) {
config.providers.claude.api_key = None;
}
// 清理服务器 API 密钥(如果是脱敏的,恢复默认值)
if config.server.api_key == REDACTED_PLACEHOLDER {
config.server.api_key = "proxy_cast".to_string();
}
}
/// 从文件导入配置
///
/// # Arguments
/// * `path` - 文件路径
/// * `current_config` - 当前配置
/// * `options` - 导入选项
///
/// # Returns
/// * `Ok(ImportResult)` - 导入成功
/// * `Err(ImportError)` - 导入失败
pub fn import_from_file(
path: &Path,
current_config: &Config,
options: &ImportOptions,
) -> Result<ImportResult, ImportError> {
let content = std::fs::read_to_string(path)?;
// 首先尝试解析为 ExportBundle
if let Ok(bundle) = ExportBundle::from_json(&content) {
return Self::import(&bundle, current_config, options, &current_config.auth_dir);
}
// 尝试解析为 YAML
Self::import_yaml(&content, current_config, options)
}
/// 保存导入的配置到文件
///
/// # Arguments
/// * `config` - 要保存的配置
/// * `path` - 配置文件路径
///
/// # Returns
/// * `Ok(())` - 保存成功
/// * `Err(ImportError)` - 保存失败
pub fn save_config(config: &Config, path: &Path) -> Result<(), ImportError> {
YamlService::save_preserve_comments(path, config)?;
Ok(())
}
}
#[cfg(test)]
mod unit_tests {
use super::*;
#[test]
fn test_import_options_default() {
let options = ImportOptions::default();
assert!(options.merge);
}
#[test]
fn test_import_options_merge() {
let options = ImportOptions::merge();
assert!(options.merge);
}
#[test]
fn test_import_options_replace() {
let options = ImportOptions::replace();
assert!(!options.merge);
}
#[test]
fn test_validation_result_valid() {
let result = ValidationResult::valid();
assert!(result.valid);
assert!(result.errors.is_empty());
assert!(result.warnings.is_empty());
}
#[test]
fn test_validation_result_invalid() {
let result = ValidationResult::invalid("test error");
assert!(!result.valid);
assert_eq!(result.errors.len(), 1);
assert!(result.errors[0].contains("test error"));
}
#[test]
fn test_validation_result_add_error() {
let mut result = ValidationResult::valid();
result.add_error("error 1");
assert!(!result.valid);
assert_eq!(result.errors.len(), 1);
}
#[test]
fn test_validation_result_add_warning() {
let mut result = ValidationResult::valid();
result.add_warning("warning 1");
assert!(result.valid); // 警告不影响有效性
assert_eq!(result.warnings.len(), 1);
}
#[test]
fn test_validate_valid_yaml() {
let yaml = r#"
server:
host: 127.0.0.1
port: 8999
api_key: test_key
"#;
let result = ImportService::validate(yaml);
assert!(result.valid);
assert!(result.has_config);
assert!(!result.has_credentials);
}
#[test]
fn test_validate_invalid_content() {
let content = "this is not valid yaml or json {{{";
let result = ImportService::validate(content);
assert!(!result.valid);
assert!(!result.errors.is_empty());
}
#[test]
fn test_validate_export_bundle() {
let bundle = ExportBundle::new("1.0.0");
let json = bundle.to_json().expect("序列化应成功");
let result = ImportService::validate(&json);
assert!(result.valid);
assert_eq!(result.version, Some("1.0".to_string()));
}
#[test]
fn test_validate_redacted_bundle() {
let mut bundle = ExportBundle::new("1.0.0");
bundle.redacted = true;
let json = bundle.to_json().expect("序列化应成功");
let result = ImportService::validate(&json);
assert!(result.valid);
assert!(result.redacted);
assert!(!result.warnings.is_empty()); // 应有脱敏警告
}
#[test]
fn test_import_yaml_replace_mode() {
let current = Config::default();
let yaml = r#"
server:
host: 0.0.0.0
port: 9000
api_key: new_key
"#;
let options = ImportOptions::replace();
let result = ImportService::import_yaml(yaml, &current, &options).expect("导入应成功");
assert!(result.success);
assert_eq!(result.config.server.host, "0.0.0.0");
assert_eq!(result.config.server.port, 9000);
assert_eq!(result.config.server.api_key, "new_key");
}
#[test]
fn test_import_yaml_merge_mode() {
let mut current = Config::default();
current.credential_pool.openai.push(ApiKeyEntry {
id: "existing".to_string(),
api_key: "sk-existing".to_string(),
base_url: None,
disabled: false,
});
let yaml = r#"
server:
host: 0.0.0.0
port: 9000
api_key: new_key
credential_pool:
openai:
- id: new
api_key: sk-new
"#;
let options = ImportOptions::merge();
let result = ImportService::import_yaml(yaml, &current, &options).expect("导入应成功");
assert!(result.success);
// 服务器配置应被更新
assert_eq!(result.config.server.host, "0.0.0.0");
// 凭证池应合并
assert_eq!(result.config.credential_pool.openai.len(), 2);
}
#[test]
fn test_merge_credential_entries() {
let current = vec![CredentialEntry {
id: "id1".to_string(),
token_file: "old.json".to_string(),
disabled: false,
}];
let imported = vec![
CredentialEntry {
id: "id1".to_string(),
token_file: "new.json".to_string(),
disabled: true,
},
CredentialEntry {
id: "id2".to_string(),
token_file: "id2.json".to_string(),
disabled: false,
},
];
let merged = ImportService::merge_credential_entries(&current, &imported);
assert_eq!(merged.len(), 2);
// id1 应被更新
assert_eq!(merged[0].token_file, "new.json");
assert!(merged[0].disabled);
// id2 应被添加
assert_eq!(merged[1].id, "id2");
}
#[test]
fn test_merge_api_key_entries_skips_redacted() {
let current = vec![ApiKeyEntry {
id: "id1".to_string(),
api_key: "sk-real".to_string(),
base_url: None,
disabled: false,
}];
let imported = vec![ApiKeyEntry {
id: "id1".to_string(),
api_key: REDACTED_PLACEHOLDER.to_string(),
base_url: None,
disabled: false,
}];
let merged = ImportService::merge_api_key_entries(&current, &imported);
assert_eq!(merged.len(), 1);
// 脱敏的条目不应覆盖现有的
assert_eq!(merged[0].api_key, "sk-real");
}
#[test]
fn test_clean_redacted_credentials() {
let mut config = Config::default();
config.server.api_key = REDACTED_PLACEHOLDER.to_string();
config.providers.openai.api_key = Some(REDACTED_PLACEHOLDER.to_string());
config.credential_pool.openai.push(ApiKeyEntry {
id: "redacted".to_string(),
api_key: REDACTED_PLACEHOLDER.to_string(),
base_url: None,
disabled: false,
});
config.credential_pool.openai.push(ApiKeyEntry {
id: "real".to_string(),
api_key: "sk-real".to_string(),
base_url: None,
disabled: false,
});
ImportService::clean_redacted_credentials(&mut config);
// 服务器 API 密钥应恢复默认值
assert_eq!(config.server.api_key, "proxy_cast");
// Provider API 密钥应被清除
assert!(config.providers.openai.api_key.is_none());
// 凭证池中脱敏的条目应被移除
assert_eq!(config.credential_pool.openai.len(), 1);
assert_eq!(config.credential_pool.openai[0].id, "real");
}
#[test]
fn test_import_error_display() {
let err = ImportError::FormatError("test".to_string());
assert!(err.to_string().contains("格式错误"));
let err = ImportError::VersionError("test".to_string());
assert!(err.to_string().contains("版本不兼容"));
let err = ImportError::RedactedDataError("test".to_string());
assert!(err.to_string().contains("脱敏数据"));
}
}
+15 -3
View File
@@ -3,19 +3,31 @@
//! 提供 YAML 配置文件支持、热重载和配置导入导出功能
//! 同时保持与旧版 JSON 配置的向后兼容性
mod export;
mod hot_reload;
mod import;
mod path_utils;
mod types;
mod yaml;
pub use export::{
base64_decode, base64_encode, ExportBundle, ExportError, ExportOptions, ExportService,
REDACTED_PLACEHOLDER,
};
pub use hot_reload::{
ConfigChangeEvent, ConfigChangeKind, FileWatcher, HotReloadError, HotReloadManager,
HotReloadStatus, ReloadResult,
};
pub use import::{ImportError, ImportOptions, ImportResult, ImportService, ValidationResult};
pub use path_utils::{collapse_tilde, contains_tilde, expand_tilde};
pub use types::{
Config, CustomProviderConfig, InjectionRuleConfig, InjectionSettings, LoggingConfig,
ProviderConfig, ProvidersConfig, RetrySettings, RoutingConfig, RoutingRuleConfig, ServerConfig,
ApiKeyEntry, Config, CredentialEntry, CredentialPoolConfig, CustomProviderConfig,
InjectionRuleConfig, InjectionSettings, LoggingConfig, ProviderConfig, ProvidersConfig,
RetrySettings, RoutingConfig, RoutingRuleConfig, ServerConfig,
};
pub use yaml::{
load_config, save_config, save_config_yaml, ConfigError, ConfigManager, YamlService,
};
pub use yaml::{load_config, save_config, save_config_yaml, ConfigError, ConfigManager};
#[cfg(test)]
mod tests;
+186
View File
@@ -0,0 +1,186 @@
//! 路径工具模块
//!
//! 提供路径处理相关的工具函数,包括 tilde (~) 路径展开
use std::path::{Path, PathBuf};
/// 展开路径中的 tilde (~) 为用户主目录
///
/// 支持以下格式:
/// - `~` -> 用户主目录
/// - `~/path` -> 用户主目录/path
/// - `~user/path` -> 不支持,返回原路径
/// - 其他路径 -> 返回原路径
///
/// # Arguments
/// * `path` - 要展开的路径字符串
///
/// # Returns
/// 展开后的 PathBuf
///
/// # Examples
/// ```ignore
/// use proxycast_lib::config::expand_tilde;
///
/// let expanded = expand_tilde("~/.proxycast/auth");
/// // 返回类似 "/Users/username/.proxycast/auth" 的路径
/// ```
pub fn expand_tilde<P: AsRef<Path>>(path: P) -> PathBuf {
let path = path.as_ref();
let path_str = path.to_string_lossy();
// 如果路径不以 ~ 开头,直接返回原路径
if !path_str.starts_with('~') {
return path.to_path_buf();
}
// 获取用户主目录
let home_dir = match dirs::home_dir() {
Some(dir) => dir,
None => return path.to_path_buf(), // 无法获取主目录,返回原路径
};
// 处理不同的 tilde 格式
if path_str == "~" {
// 仅 ~
home_dir
} else if path_str.starts_with("~/") {
// ~/path 格式
let rest = &path_str[2..]; // 跳过 "~/"
home_dir.join(rest)
} else {
// ~user/path 格式,不支持,返回原路径
path.to_path_buf()
}
}
/// 将路径收缩为 tilde 格式(如果可能)
///
/// 如果路径以用户主目录开头,则将其替换为 ~
///
/// # Arguments
/// * `path` - 要收缩的路径
///
/// # Returns
/// 收缩后的路径字符串
///
/// # Examples
/// ```ignore
/// use proxycast_lib::config::collapse_tilde;
///
/// let collapsed = collapse_tilde("/Users/username/.proxycast/auth");
/// // 返回 "~/.proxycast/auth"
/// ```
pub fn collapse_tilde<P: AsRef<Path>>(path: P) -> String {
let path = path.as_ref();
// 获取用户主目录
let home_dir = match dirs::home_dir() {
Some(dir) => dir,
None => return path.to_string_lossy().to_string(),
};
// 检查路径是否以主目录开头
if let Ok(stripped) = path.strip_prefix(&home_dir) {
if stripped.as_os_str().is_empty() {
"~".to_string()
} else {
format!("~/{}", stripped.to_string_lossy())
}
} else {
path.to_string_lossy().to_string()
}
}
/// 检查路径是否包含 tilde
pub fn contains_tilde<P: AsRef<Path>>(path: P) -> bool {
path.as_ref().to_string_lossy().starts_with('~')
}
#[cfg(test)]
mod unit_tests {
use super::*;
#[test]
fn test_expand_tilde_only() {
let expanded = expand_tilde("~");
let home = dirs::home_dir().expect("应该能获取主目录");
assert_eq!(expanded, home);
}
#[test]
fn test_expand_tilde_with_path() {
let expanded = expand_tilde("~/.proxycast/auth");
let home = dirs::home_dir().expect("应该能获取主目录");
assert_eq!(expanded, home.join(".proxycast/auth"));
}
#[test]
fn test_expand_tilde_nested_path() {
let expanded = expand_tilde("~/a/b/c/d");
let home = dirs::home_dir().expect("应该能获取主目录");
assert_eq!(expanded, home.join("a/b/c/d"));
}
#[test]
fn test_expand_tilde_no_tilde() {
let path = "/absolute/path/to/file";
let expanded = expand_tilde(path);
assert_eq!(expanded, PathBuf::from(path));
}
#[test]
fn test_expand_tilde_relative_path() {
let path = "relative/path/to/file";
let expanded = expand_tilde(path);
assert_eq!(expanded, PathBuf::from(path));
}
#[test]
fn test_expand_tilde_user_format_not_supported() {
// ~user/path 格式不支持,应返回原路径
let path = "~otheruser/path";
let expanded = expand_tilde(path);
assert_eq!(expanded, PathBuf::from(path));
}
#[test]
fn test_collapse_tilde_home_dir() {
let home = dirs::home_dir().expect("应该能获取主目录");
let collapsed = collapse_tilde(&home);
assert_eq!(collapsed, "~");
}
#[test]
fn test_collapse_tilde_with_subpath() {
let home = dirs::home_dir().expect("应该能获取主目录");
let path = home.join(".proxycast/auth");
let collapsed = collapse_tilde(&path);
assert_eq!(collapsed, "~/.proxycast/auth");
}
#[test]
fn test_collapse_tilde_not_in_home() {
let path = "/tmp/some/path";
let collapsed = collapse_tilde(path);
assert_eq!(collapsed, path);
}
#[test]
fn test_contains_tilde() {
assert!(contains_tilde("~"));
assert!(contains_tilde("~/path"));
assert!(contains_tilde("~user/path"));
assert!(!contains_tilde("/absolute/path"));
assert!(!contains_tilde("relative/path"));
}
#[test]
fn test_expand_collapse_roundtrip() {
// 对于 ~/path 格式,展开后再收缩应该得到原路径
let original = "~/.proxycast/auth/token.json";
let expanded = expand_tilde(original);
let collapsed = collapse_tilde(&expanded);
assert_eq!(collapsed, original);
}
}
File diff suppressed because it is too large Load Diff
+161
View File
@@ -7,6 +7,66 @@ use crate::injection::{InjectionMode, InjectionRule};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
// ============ 凭证池配置类型 ============
/// 凭证池配置
///
/// 管理多个 Provider 的多个凭证,支持负载均衡
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)]
pub struct CredentialPoolConfig {
/// Kiro 凭证列表(OAuth)
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub kiro: Vec<CredentialEntry>,
/// Gemini 凭证列表(OAuth)
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub gemini: Vec<CredentialEntry>,
/// Qwen 凭证列表(OAuth)
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub qwen: Vec<CredentialEntry>,
/// OpenAI 凭证列表(API Key)
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub openai: Vec<ApiKeyEntry>,
/// Claude 凭证列表(API Key)
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub claude: Vec<ApiKeyEntry>,
}
/// OAuth 凭证条目
///
/// 用于 Kiro、Gemini、Qwen 等 OAuth 认证的 Provider
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct CredentialEntry {
/// 凭证 ID
pub id: String,
/// Token 文件路径(相对于 auth_dir)
pub token_file: String,
/// 是否禁用
#[serde(default)]
pub disabled: bool,
}
/// API Key 凭证条目
///
/// 用于 OpenAI、Claude 等 API Key 认证的 Provider
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct ApiKeyEntry {
/// 凭证 ID
pub id: String,
/// API Key
pub api_key: String,
/// 自定义 Base URL
#[serde(default, skip_serializing_if = "Option::is_none")]
pub base_url: Option<String>,
/// 是否禁用
#[serde(default)]
pub disabled: bool,
}
/// 默认 auth_dir 路径
fn default_auth_dir() -> String {
"~/.proxycast/auth".to_string()
}
/// 主配置结构
///
/// 支持两种格式:
@@ -35,6 +95,12 @@ pub struct Config {
/// 参数注入配置
#[serde(default)]
pub injection: InjectionSettings,
/// 认证目录路径(存储 OAuth Token 文件,支持 ~ 展开)
#[serde(default = "default_auth_dir")]
pub auth_dir: String,
/// 凭证池配置
#[serde(default)]
pub credential_pool: CredentialPoolConfig,
}
/// 服务器配置
@@ -372,6 +438,8 @@ impl Default for Config {
retry: RetrySettings::default(),
logging: LoggingConfig::default(),
injection: InjectionSettings::default(),
auth_dir: default_auth_dir(),
credential_pool: CredentialPoolConfig::default(),
}
}
}
@@ -394,6 +462,99 @@ mod unit_tests {
assert!(config.logging.enabled);
assert!(!config.injection.enabled);
assert!(config.injection.rules.is_empty());
// 新增字段测试
assert_eq!(config.auth_dir, "~/.proxycast/auth");
assert!(config.credential_pool.kiro.is_empty());
assert!(config.credential_pool.openai.is_empty());
}
#[test]
fn test_credential_pool_config_default() {
let pool = CredentialPoolConfig::default();
assert!(pool.kiro.is_empty());
assert!(pool.gemini.is_empty());
assert!(pool.qwen.is_empty());
assert!(pool.openai.is_empty());
assert!(pool.claude.is_empty());
}
#[test]
fn test_credential_entry_serialization() {
let entry = CredentialEntry {
id: "kiro-main".to_string(),
token_file: "kiro/main-token.json".to_string(),
disabled: false,
};
let yaml = serde_yaml::to_string(&entry).unwrap();
assert!(yaml.contains("id: kiro-main"));
assert!(yaml.contains("token_file: kiro/main-token.json"));
let parsed: CredentialEntry = serde_yaml::from_str(&yaml).unwrap();
assert_eq!(parsed, entry);
}
#[test]
fn test_api_key_entry_serialization() {
let entry = ApiKeyEntry {
id: "openai-main".to_string(),
api_key: "sk-test-key".to_string(),
base_url: Some("https://api.openai.com/v1".to_string()),
disabled: false,
};
let yaml = serde_yaml::to_string(&entry).unwrap();
assert!(yaml.contains("id: openai-main"));
assert!(yaml.contains("api_key: sk-test-key"));
let parsed: ApiKeyEntry = serde_yaml::from_str(&yaml).unwrap();
assert_eq!(parsed, entry);
}
#[test]
fn test_api_key_entry_without_base_url() {
let entry = ApiKeyEntry {
id: "claude-main".to_string(),
api_key: "sk-ant-test".to_string(),
base_url: None,
disabled: true,
};
let yaml = serde_yaml::to_string(&entry).unwrap();
// base_url should be skipped when None
assert!(!yaml.contains("base_url"));
assert!(yaml.contains("disabled: true"));
let parsed: ApiKeyEntry = serde_yaml::from_str(&yaml).unwrap();
assert_eq!(parsed, entry);
}
#[test]
fn test_credential_pool_config_serialization() {
let pool = CredentialPoolConfig {
kiro: vec![CredentialEntry {
id: "kiro-1".to_string(),
token_file: "kiro/token-1.json".to_string(),
disabled: false,
}],
gemini: vec![],
qwen: vec![],
openai: vec![ApiKeyEntry {
id: "openai-1".to_string(),
api_key: "sk-xxx".to_string(),
base_url: None,
disabled: false,
}],
claude: vec![],
};
let yaml = serde_yaml::to_string(&pool).unwrap();
// Empty vecs should be skipped
assert!(!yaml.contains("gemini"));
assert!(!yaml.contains("qwen"));
assert!(!yaml.contains("claude"));
assert!(yaml.contains("kiro"));
assert!(yaml.contains("openai"));
let parsed: CredentialPoolConfig = serde_yaml::from_str(&yaml).unwrap();
assert_eq!(parsed, pool);
}
#[test]
+405
View File
@@ -1,8 +1,10 @@
//! YAML 配置文件支持
//!
//! 提供 YAML 配置的加载、保存和管理功能
//! 支持保留注释的配置保存
use super::types::Config;
use std::collections::HashMap;
use std::path::{Path, PathBuf};
/// 配置错误类型
@@ -54,6 +56,14 @@ impl ConfigManager {
}
}
/// 使用指定配置创建配置管理器
pub fn with_config(config: Config, config_path: PathBuf) -> Self {
Self {
config,
config_path,
}
}
/// 从文件加载配置
///
/// 如果文件不存在,返回默认配置
@@ -238,6 +248,401 @@ impl Default for ConfigManager {
}
}
// ============ YAML 注释保留功能 ============
/// YAML 服务 - 提供保留注释的 YAML 操作
pub struct YamlService;
/// 注释信息
#[derive(Debug, Clone)]
struct CommentInfo {
/// 行号(0-indexed)
line: usize,
/// 注释内容(包含 # 符号)
content: String,
/// 是否是行尾注释
is_inline: bool,
/// 关联的键路径(如果有)
key_path: Option<String>,
}
impl YamlService {
/// 保存配置到 YAML,保留原文件中的注释和格式
///
/// # Arguments
/// * `path` - 配置文件路径
/// * `config` - 要保存的配置
///
/// # Returns
/// * `Ok(())` - 保存成功
/// * `Err(ConfigError)` - 保存失败
pub fn save_preserve_comments(path: &Path, config: &Config) -> Result<(), ConfigError> {
// 读取原文件内容(如果存在)
let original_content = if path.exists() {
std::fs::read_to_string(path).ok()
} else {
None
};
// 序列化新配置
let new_yaml = ConfigManager::to_yaml(config)?;
// 如果原文件存在,尝试保留注释
let final_content = if let Some(original) = original_content {
Self::merge_comments(&original, &new_yaml)
} else {
new_yaml
};
// 确保父目录存在
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent).map_err(|e| ConfigError::WriteError(e.to_string()))?;
}
// 写入文件
std::fs::write(path, final_content).map_err(|e| ConfigError::WriteError(e.to_string()))
}
/// 合并原文件的注释到新 YAML 内容中
///
/// 策略:
/// 1. 提取原文件中的所有注释(独立行注释和行尾注释)
/// 2. 尝试将注释与键路径关联
/// 3. 在新 YAML 中找到对应位置插入注释
/// 4. 无法关联的注释放在文件头部
fn merge_comments(original: &str, new_yaml: &str) -> String {
let original_lines: Vec<&str> = original.lines().collect();
let new_lines: Vec<&str> = new_yaml.lines().collect();
// 提取原文件中的注释
let comments = Self::extract_comments(&original_lines);
// 如果没有注释,直接返回新内容
if comments.is_empty() {
return new_yaml.to_string();
}
// 构建新 YAML 的键位置映射
let new_key_positions = Self::build_key_positions(&new_lines);
// 合并注释到新内容
Self::insert_comments(&new_lines, &comments, &new_key_positions)
}
/// 提取所有独立行注释(不尝试关联键路径)
pub fn extract_all_comments(yaml: &str) -> Vec<String> {
yaml.lines()
.filter(|line| line.trim().starts_with('#'))
.map(|s| s.to_string())
.collect()
}
/// 从 YAML 行中提取注释
fn extract_comments(lines: &[&str]) -> Vec<CommentInfo> {
let mut comments = Vec::new();
let mut current_key_path: Vec<String> = Vec::new();
let mut indent_stack: Vec<usize> = vec![0];
for (line_num, line) in lines.iter().enumerate() {
let trimmed = line.trim();
// 跳过空行
if trimmed.is_empty() {
continue;
}
// 计算当前行的缩进
let indent = line.len() - line.trim_start().len();
// 检查是否是纯注释行
if trimmed.starts_with('#') {
// 确定注释关联的键路径
let key_path = if current_key_path.is_empty() {
None
} else {
Some(current_key_path.join("."))
};
comments.push(CommentInfo {
line: line_num,
content: line.to_string(),
is_inline: false,
key_path,
});
continue;
}
// 检查是否有行尾注释
if let Some(comment_pos) = Self::find_comment_position(line) {
let comment_content = line[comment_pos..].to_string();
// 更新键路径
Self::update_key_path(line, indent, &mut current_key_path, &mut indent_stack);
comments.push(CommentInfo {
line: line_num,
content: comment_content,
is_inline: true,
key_path: Some(current_key_path.join(".")),
});
} else {
// 普通键值行,更新键路径
Self::update_key_path(line, indent, &mut current_key_path, &mut indent_stack);
}
}
comments
}
/// 查找行中注释的位置(考虑字符串内的 # 符号)
fn find_comment_position(line: &str) -> Option<usize> {
let mut in_single_quote = false;
let mut in_double_quote = false;
let mut prev_char = ' ';
for (i, c) in line.char_indices() {
match c {
'\'' if !in_double_quote && prev_char != '\\' => {
in_single_quote = !in_single_quote;
}
'"' if !in_single_quote && prev_char != '\\' => {
in_double_quote = !in_double_quote;
}
'#' if !in_single_quote && !in_double_quote => {
// 确保 # 前面有空格(YAML 注释规则)
if i == 0
|| line
.chars()
.nth(i - 1)
.map(|c| c.is_whitespace())
.unwrap_or(false)
{
return Some(i);
}
}
_ => {}
}
prev_char = c;
}
None
}
/// 更新当前键路径
fn update_key_path(
line: &str,
indent: usize,
current_key_path: &mut Vec<String>,
indent_stack: &mut Vec<usize>,
) {
let trimmed = line.trim();
// 提取键名
if let Some(colon_pos) = trimmed.find(':') {
let key = trimmed[..colon_pos].trim();
// 跳过列表项
if key.starts_with('-') {
return;
}
// 根据缩进调整键路径
while indent_stack.len() > 1 && indent <= indent_stack[indent_stack.len() - 1] {
indent_stack.pop();
current_key_path.pop();
}
// 添加新键
current_key_path.push(key.to_string());
indent_stack.push(indent);
}
}
/// 构建新 YAML 的键位置映射
fn build_key_positions(lines: &[&str]) -> HashMap<String, usize> {
let mut positions = HashMap::new();
let mut current_key_path: Vec<String> = Vec::new();
let mut indent_stack: Vec<usize> = vec![0];
for (line_num, line) in lines.iter().enumerate() {
let trimmed = line.trim();
if trimmed.is_empty() || trimmed.starts_with('#') {
continue;
}
let indent = line.len() - line.trim_start().len();
if let Some(colon_pos) = trimmed.find(':') {
let key = trimmed[..colon_pos].trim();
if key.starts_with('-') {
continue;
}
// 根据缩进调整键路径
while indent_stack.len() > 1 && indent <= indent_stack[indent_stack.len() - 1] {
indent_stack.pop();
current_key_path.pop();
}
current_key_path.push(key.to_string());
indent_stack.push(indent);
// 记录位置
positions.insert(current_key_path.join("."), line_num);
}
}
positions
}
/// 将注释插入到新 YAML 内容中
fn insert_comments(
new_lines: &[&str],
comments: &[CommentInfo],
key_positions: &HashMap<String, usize>,
) -> String {
let mut result_lines: Vec<String> = new_lines.iter().map(|s| s.to_string()).collect();
let mut insertions: Vec<(usize, String)> = Vec::new();
let mut unmatched_comments: Vec<String> = Vec::new();
for comment in comments {
if comment.is_inline {
// 行尾注释:找到对应的键并追加
if let Some(key_path) = &comment.key_path {
if let Some(&line_num) = key_positions.get(key_path) {
if line_num < result_lines.len() {
// 追加行尾注释
let existing = &result_lines[line_num];
if !existing.contains('#') {
result_lines[line_num] =
format!("{} {}", existing, comment.content);
}
}
} else {
// 无法匹配的行尾注释,作为独立注释保留
unmatched_comments.push(comment.content.clone());
}
}
} else {
// 独立行注释:尝试在对应位置插入
if let Some(key_path) = &comment.key_path {
// 找到下一个键的位置
if let Some(&next_line) = key_positions.get(key_path) {
insertions.push((next_line, comment.content.clone()));
} else {
// 无法匹配的注释,放到头部
unmatched_comments.push(comment.content.clone());
}
} else {
// 文件头部注释
insertions.push((0, comment.content.clone()));
}
}
}
// 按位置倒序插入,避免索引偏移
insertions.sort_by(|a, b| b.0.cmp(&a.0));
for (pos, content) in insertions {
if pos <= result_lines.len() {
result_lines.insert(pos, content);
}
}
// 将无法匹配的注释放在文件头部
if !unmatched_comments.is_empty() {
let mut final_lines = unmatched_comments;
final_lines.extend(result_lines);
return final_lines.join("\n");
}
result_lines.join("\n")
}
/// 更新 YAML 中的特定字段
///
/// # Arguments
/// * `path` - 配置文件路径
/// * `field_path` - 字段路径,如 ["server", "port"]
/// * `value` - 新值(YAML 格式的字符串)
///
/// # Returns
/// * `Ok(())` - 更新成功
/// * `Err(ConfigError)` - 更新失败
pub fn update_field(path: &Path, field_path: &[&str], value: &str) -> Result<(), ConfigError> {
// 读取原文件
let content =
std::fs::read_to_string(path).map_err(|e| ConfigError::ReadError(e.to_string()))?;
let lines: Vec<&str> = content.lines().collect();
let mut result_lines: Vec<String> = Vec::new();
let target_key = field_path.last().copied().unwrap_or("");
let parent_path = &field_path[..field_path.len().saturating_sub(1)];
let mut current_path: Vec<String> = Vec::new();
let mut indent_stack: Vec<usize> = vec![0];
let mut found = false;
for line in lines {
let trimmed = line.trim();
if trimmed.is_empty() || trimmed.starts_with('#') {
result_lines.push(line.to_string());
continue;
}
let indent = line.len() - line.trim_start().len();
if let Some(colon_pos) = trimmed.find(':') {
let key = trimmed[..colon_pos].trim();
if !key.starts_with('-') {
// 根据缩进调整路径
while indent_stack.len() > 1 && indent <= indent_stack[indent_stack.len() - 1] {
indent_stack.pop();
current_path.pop();
}
// 检查是否匹配目标字段
let parent_matches = current_path.len() == parent_path.len()
&& current_path
.iter()
.zip(parent_path.iter())
.all(|(a, b)| a == *b);
if parent_matches && key == target_key {
// 找到目标字段,替换值
let new_line = format!("{}{}: {}", " ".repeat(indent), key, value);
result_lines.push(new_line);
found = true;
current_path.push(key.to_string());
indent_stack.push(indent);
continue;
}
current_path.push(key.to_string());
indent_stack.push(indent);
}
}
result_lines.push(line.to_string());
}
if !found {
return Err(ConfigError::ValidationError(format!(
"字段 {} 未找到",
field_path.join(".")
)));
}
// 写入文件
std::fs::write(path, result_lines.join("\n"))
.map_err(|e| ConfigError::WriteError(e.to_string()))
}
}
// ============ 向后兼容的 JSON 配置函数 ============
/// 获取 JSON 配置文件路径(向后兼容)
+2
View File
@@ -5,11 +5,13 @@
mod balancer;
mod health;
mod pool;
mod sync;
mod types;
pub use balancer::{BalanceStrategy, CooldownInfo, LoadBalancer};
pub use health::{HealthCheckConfig, HealthCheckResult, HealthChecker, HealthStatus};
pub use pool::{CredentialPool, PoolError, PoolStatus};
pub use sync::{CredentialSyncService, SyncError};
pub use types::{Credential, CredentialData, CredentialStats, CredentialStatus};
#[cfg(test)]
+538
View File
@@ -0,0 +1,538 @@
//! 凭证同步服务
//!
//! 负责将凭证池变更同步到 YAML 配置文件
//! 实现凭证的添加、删除、更新操作与配置文件的同步
use crate::config::{
expand_tilde, ApiKeyEntry, Config, ConfigError, ConfigManager, CredentialEntry, YamlService,
};
use crate::models::provider_pool_model::{CredentialData, PoolProviderType, ProviderCredential};
use std::path::PathBuf;
use std::sync::{Arc, RwLock};
/// 凭证同步服务错误类型
#[derive(Debug, Clone)]
pub enum SyncError {
/// 配置错误
ConfigError(String),
/// IO 错误
IoError(String),
/// 凭证不存在
CredentialNotFound(String),
/// 无效的凭证类型
InvalidCredentialType(String),
}
impl std::fmt::Display for SyncError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SyncError::ConfigError(msg) => write!(f, "配置错误: {}", msg),
SyncError::IoError(msg) => write!(f, "IO 错误: {}", msg),
SyncError::CredentialNotFound(id) => write!(f, "凭证不存在: {}", id),
SyncError::InvalidCredentialType(msg) => write!(f, "无效的凭证类型: {}", msg),
}
}
}
impl std::error::Error for SyncError {}
impl From<ConfigError> for SyncError {
fn from(err: ConfigError) -> Self {
SyncError::ConfigError(err.to_string())
}
}
impl From<std::io::Error> for SyncError {
fn from(err: std::io::Error) -> Self {
SyncError::IoError(err.to_string())
}
}
/// 凭证同步服务
///
/// 负责将凭证池变更同步到 YAML 配置文件
pub struct CredentialSyncService {
/// 配置管理器
config_manager: Arc<RwLock<ConfigManager>>,
}
impl CredentialSyncService {
/// 创建新的凭证同步服务
pub fn new(config_manager: Arc<RwLock<ConfigManager>>) -> Self {
Self { config_manager }
}
/// 获取当前配置
fn get_config(&self) -> Result<Config, SyncError> {
let manager = self
.config_manager
.read()
.map_err(|e| SyncError::ConfigError(format!("获取配置锁失败: {}", e)))?;
Ok(manager.config().clone())
}
/// 更新配置并保存
fn update_config(&self, config: Config) -> Result<(), SyncError> {
let mut manager = self
.config_manager
.write()
.map_err(|e| SyncError::ConfigError(format!("获取配置写锁失败: {}", e)))?;
let config_path = manager.config_path().to_path_buf();
manager.set_config(config.clone());
// 使用 YamlService 保存配置,保留注释
YamlService::save_preserve_comments(&config_path, &config)?;
Ok(())
}
/// 获取 auth_dir 的绝对路径
pub fn get_auth_dir(&self) -> Result<PathBuf, SyncError> {
let config = self.get_config()?;
Ok(expand_tilde(&config.auth_dir))
}
/// 确保 auth_dir 目录存在
pub fn ensure_auth_dir(&self) -> Result<PathBuf, SyncError> {
let auth_dir = self.get_auth_dir()?;
std::fs::create_dir_all(&auth_dir)?;
Ok(auth_dir)
}
/// 添加凭证并同步到配置
///
/// # 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")?;
let entry = CredentialEntry {
id: credential.uuid.clone(),
token_file,
disabled: credential.is_disabled,
};
config.credential_pool.kiro.push(entry);
}
CredentialData::GeminiOAuth {
creds_file_path, ..
} => {
let token_file =
self.save_oauth_token_file(creds_file_path, &credential.uuid, "gemini")?;
let entry = CredentialEntry {
id: credential.uuid.clone(),
token_file,
disabled: credential.is_disabled,
};
config.credential_pool.gemini.push(entry);
}
CredentialData::QwenOAuth { creds_file_path } => {
let token_file =
self.save_oauth_token_file(creds_file_path, &credential.uuid, "qwen")?;
let entry = CredentialEntry {
id: credential.uuid.clone(),
token_file,
disabled: credential.is_disabled,
};
config.credential_pool.qwen.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(),
api_key: api_key.clone(),
base_url: base_url.clone(),
disabled: credential.is_disabled,
};
config.credential_pool.openai.push(entry);
}
CredentialData::ClaudeKey { api_key, base_url } => {
let entry = ApiKeyEntry {
id: credential.uuid.clone(),
api_key: api_key.clone(),
base_url: base_url.clone(),
disabled: credential.is_disabled,
};
config.credential_pool.claude.push(entry);
}
}
self.update_config(config)
}
/// 保存 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,
credential_id: &str,
provider: &str,
) -> Result<String, SyncError> {
let auth_dir = self.ensure_auth_dir()?;
let provider_dir = auth_dir.join(provider);
std::fs::create_dir_all(&provider_dir)?;
// 生成 token 文件名
let token_filename = format!("{}.json", credential_id);
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))
}
/// 删除凭证并同步到配置
///
/// # Arguments
/// * `provider_type` - Provider 类型
/// * `credential_id` - 凭证 ID
///
/// # Returns
/// * `Ok(())` - 删除成功
/// * `Err(SyncError)` - 删除失败
pub fn remove_credential(
&self,
provider_type: PoolProviderType,
credential_id: &str,
) -> Result<(), SyncError> {
let mut config = self.get_config()?;
let mut found = false;
match provider_type {
PoolProviderType::Kiro => {
if let Some(pos) = config
.credential_pool
.kiro
.iter()
.position(|e| e.id == credential_id)
{
let entry = config.credential_pool.kiro.remove(pos);
self.delete_oauth_token_file(&entry.token_file)?;
found = true;
}
}
PoolProviderType::Gemini => {
if let Some(pos) = config
.credential_pool
.gemini
.iter()
.position(|e| e.id == credential_id)
{
let entry = config.credential_pool.gemini.remove(pos);
self.delete_oauth_token_file(&entry.token_file)?;
found = true;
}
}
PoolProviderType::Qwen => {
if let Some(pos) = config
.credential_pool
.qwen
.iter()
.position(|e| e.id == credential_id)
{
let entry = config.credential_pool.qwen.remove(pos);
self.delete_oauth_token_file(&entry.token_file)?;
found = true;
}
}
PoolProviderType::OpenAI => {
if let Some(pos) = config
.credential_pool
.openai
.iter()
.position(|e| e.id == credential_id)
{
config.credential_pool.openai.remove(pos);
found = true;
}
}
PoolProviderType::Claude => {
if let Some(pos) = config
.credential_pool
.claude
.iter()
.position(|e| e.id == credential_id)
{
config.credential_pool.claude.remove(pos);
found = true;
}
}
PoolProviderType::Antigravity => {
return Err(SyncError::InvalidCredentialType(
"Antigravity 凭证暂不支持同步到配置".to_string(),
));
}
}
if !found {
return Err(SyncError::CredentialNotFound(credential_id.to_string()));
}
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;
match &credential.credential {
CredentialData::KiroOAuth { creds_file_path } => {
if let Some(entry) = config
.credential_pool
.kiro
.iter_mut()
.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;
found = true;
}
}
CredentialData::GeminiOAuth {
creds_file_path, ..
} => {
if let Some(entry) = config
.credential_pool
.gemini
.iter_mut()
.find(|e| e.id == credential.uuid)
{
entry.disabled = credential.is_disabled;
let new_token_file =
self.save_oauth_token_file(creds_file_path, &credential.uuid, "gemini")?;
entry.token_file = new_token_file;
found = true;
}
}
CredentialData::QwenOAuth { creds_file_path } => {
if let Some(entry) = config
.credential_pool
.qwen
.iter_mut()
.find(|e| e.id == credential.uuid)
{
entry.disabled = credential.is_disabled;
let new_token_file =
self.save_oauth_token_file(creds_file_path, &credential.uuid, "qwen")?;
entry.token_file = new_token_file;
found = true;
}
}
CredentialData::AntigravityOAuth { .. } => {
return Err(SyncError::InvalidCredentialType(
"Antigravity 凭证暂不支持同步到配置".to_string(),
));
}
CredentialData::OpenAIKey { api_key, base_url } => {
if let Some(entry) = config
.credential_pool
.openai
.iter_mut()
.find(|e| e.id == credential.uuid)
{
entry.api_key = api_key.clone();
entry.base_url = base_url.clone();
entry.disabled = credential.is_disabled;
found = true;
}
}
CredentialData::ClaudeKey { api_key, base_url } => {
if let Some(entry) = config
.credential_pool
.claude
.iter_mut()
.find(|e| e.id == credential.uuid)
{
entry.api_key = api_key.clone();
entry.base_url = base_url.clone();
entry.disabled = credential.is_disabled;
found = true;
}
}
}
if !found {
return Err(SyncError::CredentialNotFound(credential.uuid.clone()));
}
self.update_config(config)
}
/// 从配置加载凭证到池中
///
/// 启动时从 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()?;
let mut credentials = Vec::new();
// 加载 Kiro 凭证
for entry in &config.credential_pool.kiro {
let token_path = auth_dir.join(&entry.token_file);
let 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);
}
// 加载 Gemini 凭证
for entry in &config.credential_pool.gemini {
let token_path = auth_dir.join(&entry.token_file);
let 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);
}
// 加载 Qwen 凭证
for entry in &config.credential_pool.qwen {
let token_path = auth_dir.join(&entry.token_file);
let cred = ProviderCredential::new(
PoolProviderType::Qwen,
CredentialData::QwenOAuth {
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);
}
// 加载 OpenAI 凭证
for entry in &config.credential_pool.openai {
let 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);
}
// 加载 Claude 凭证
for entry in &config.credential_pool.claude {
let 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);
}
Ok(credentials)
}
/// 获取 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() {
std::fs::create_dir_all(parent)?;
}
std::fs::write(&path, content).map_err(SyncError::from)
}
}
+362
View File
@@ -664,3 +664,365 @@ proptest! {
);
}
}
// ============ 凭证同步服务属性测试 ============
use crate::config::{Config, ConfigManager};
use crate::credential::CredentialSyncService;
use crate::models::provider_pool_model::{
CredentialData as PoolCredentialData, PoolProviderType, ProviderCredential,
};
use std::sync::RwLock;
use tempfile::TempDir;
/// 创建临时测试环境
fn create_test_env() -> (TempDir, Arc<RwLock<ConfigManager>>) {
let temp_dir = TempDir::new().expect("创建临时目录失败");
let config_path = temp_dir.path().join("config.yaml");
// 创建配置管理器
let mut config = Config::default();
config.auth_dir = temp_dir.path().join("auth").to_string_lossy().to_string();
let mut manager = ConfigManager::new(config_path);
manager.set_config(config);
manager.save().expect("保存配置失败");
(temp_dir, Arc::new(RwLock::new(manager)))
}
/// 生成随机的 PoolProviderType(仅支持同步的类型)
fn arb_sync_provider_type() -> impl Strategy<Value = PoolProviderType> {
prop_oneof![
Just(PoolProviderType::Kiro),
Just(PoolProviderType::Gemini),
Just(PoolProviderType::Qwen),
Just(PoolProviderType::OpenAI),
Just(PoolProviderType::Claude),
]
}
/// 生成随机的 API Key 凭证数据
fn arb_api_key_credential() -> impl Strategy<Value = (PoolProviderType, PoolCredentialData)> {
prop_oneof![
(
"[a-zA-Z0-9]{20,50}",
prop::option::of("https://[a-z]+\\.[a-z]+/v1")
)
.prop_map(|(api_key, base_url)| {
(
PoolProviderType::OpenAI,
PoolCredentialData::OpenAIKey { api_key, base_url },
)
}),
(
"[a-zA-Z0-9]{20,50}",
prop::option::of("https://[a-z]+\\.[a-z]+")
)
.prop_map(|(api_key, base_url)| {
(
PoolProviderType::Claude,
PoolCredentialData::ClaudeKey { api_key, base_url },
)
}),
]
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
/// **Feature: config-credential-export, Property 1: Credential Sync Round Trip**
/// *For any* credential added to the credential pool, saving to YAML and then loading
/// from YAML should produce an equivalent credential configuration.
/// **Validates: Requirements 1.1, 1.2, 1.5**
#[test]
fn prop_credential_sync_round_trip(
(provider_type, cred_data) in arb_api_key_credential(),
is_disabled in proptest::bool::ANY
) {
let (_temp_dir, config_manager) = create_test_env();
let sync_service = CredentialSyncService::new(config_manager.clone());
// 创建凭证
let mut credential = ProviderCredential::new(provider_type, cred_data.clone());
credential.is_disabled = is_disabled;
let original_uuid = credential.uuid.clone();
// 添加凭证
let add_result = sync_service.add_credential(&credential);
prop_assert!(add_result.is_ok(), "添加凭证应该成功: {:?}", add_result);
// 从配置加载凭证
let loaded = sync_service.load_from_config();
prop_assert!(loaded.is_ok(), "加载凭证应该成功: {:?}", loaded);
let loaded_creds = loaded.unwrap();
// 查找对应的凭证
let found = loaded_creds.iter().find(|c| c.uuid == original_uuid);
prop_assert!(found.is_some(), "应该能找到添加的凭证");
let loaded_cred = found.unwrap();
// 验证凭证属性
prop_assert_eq!(
&loaded_cred.uuid,
&original_uuid,
"UUID 应该一致"
);
prop_assert_eq!(
loaded_cred.provider_type,
provider_type,
"Provider 类型应该一致"
);
prop_assert_eq!(
loaded_cred.is_disabled,
is_disabled,
"禁用状态应该一致"
);
// 验证凭证数据
match (&loaded_cred.credential, &cred_data) {
(
PoolCredentialData::OpenAIKey { api_key: loaded_key, base_url: loaded_url },
PoolCredentialData::OpenAIKey { api_key: orig_key, base_url: orig_url },
) => {
prop_assert_eq!(loaded_key, orig_key, "API Key 应该一致");
prop_assert_eq!(loaded_url, orig_url, "Base URL 应该一致");
}
(
PoolCredentialData::ClaudeKey { api_key: loaded_key, base_url: loaded_url },
PoolCredentialData::ClaudeKey { api_key: orig_key, base_url: orig_url },
) => {
prop_assert_eq!(loaded_key, orig_key, "API Key 应该一致");
prop_assert_eq!(loaded_url, orig_url, "Base URL 应该一致");
}
_ => {
prop_assert!(false, "凭证类型不匹配");
}
}
}
/// **Feature: config-credential-export, Property 1: Credential Sync Round Trip (Multiple)**
/// *For any* set of credentials, adding them all and then loading should preserve all.
/// **Validates: Requirements 1.1, 1.2, 1.5**
#[test]
fn prop_credential_sync_round_trip_multiple(
cred_count in 1usize..=5usize
) {
let (_temp_dir, config_manager) = create_test_env();
let sync_service = CredentialSyncService::new(config_manager.clone());
// 创建多个凭证
let mut original_uuids = Vec::new();
for i in 0..cred_count {
let cred_data = if i % 2 == 0 {
PoolCredentialData::OpenAIKey {
api_key: format!("sk-test-key-{}", i),
base_url: Some("https://api.openai.com/v1".to_string()),
}
} else {
PoolCredentialData::ClaudeKey {
api_key: format!("sk-ant-test-key-{}", i),
base_url: None,
}
};
let provider_type = if i % 2 == 0 {
PoolProviderType::OpenAI
} else {
PoolProviderType::Claude
};
let credential = ProviderCredential::new(provider_type, cred_data);
original_uuids.push(credential.uuid.clone());
let add_result = sync_service.add_credential(&credential);
prop_assert!(add_result.is_ok(), "添加凭证 {} 应该成功", i);
}
// 从配置加载凭证
let loaded = sync_service.load_from_config().unwrap();
// 验证所有凭证都被加载
prop_assert_eq!(
loaded.len(),
cred_count,
"加载的凭证数量应该与添加的一致"
);
// 验证每个 UUID 都存在
for uuid in &original_uuids {
let found = loaded.iter().any(|c| &c.uuid == uuid);
prop_assert!(found, "应该能找到 UUID 为 {} 的凭证", uuid);
}
}
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
/// **Feature: config-credential-export, Property 9: OAuth Token File Handling**
/// *For any* OAuth credential, the token file should be stored in auth-dir on add,
/// included in export bundles, and restored to auth-dir on import.
/// **Validates: Requirements 2.1, 2.4, 3.3, 4.4**
#[test]
fn prop_oauth_token_file_handling(
token_content in "[a-zA-Z0-9]{50,200}",
provider_idx in 0usize..3usize
) {
let (temp_dir, config_manager) = create_test_env();
let sync_service = CredentialSyncService::new(config_manager.clone());
// 创建源 token 文件
let source_token_dir = temp_dir.path().join("source_tokens");
std::fs::create_dir_all(&source_token_dir).expect("创建源目录失败");
let source_token_path = source_token_dir.join("token.json");
let token_json = format!(r#"{{"access_token": "{}", "refresh_token": "refresh-{}", "expires_at": "2025-12-31T23:59:59Z"}}"#, token_content, token_content);
std::fs::write(&source_token_path, &token_json).expect("写入源 token 文件失败");
// 根据索引选择 provider 类型
let (provider_type, cred_data) = match provider_idx {
0 => (
PoolProviderType::Kiro,
PoolCredentialData::KiroOAuth {
creds_file_path: source_token_path.to_string_lossy().to_string(),
},
),
1 => (
PoolProviderType::Gemini,
PoolCredentialData::GeminiOAuth {
creds_file_path: source_token_path.to_string_lossy().to_string(),
project_id: None,
},
),
_ => (
PoolProviderType::Qwen,
PoolCredentialData::QwenOAuth {
creds_file_path: source_token_path.to_string_lossy().to_string(),
},
),
};
// 创建凭证
let credential = ProviderCredential::new(provider_type, cred_data);
let original_uuid = credential.uuid.clone();
// 添加凭证(应该复制 token 文件到 auth_dir)
let add_result = sync_service.add_credential(&credential);
prop_assert!(add_result.is_ok(), "添加 OAuth 凭证应该成功: {:?}", add_result);
// 验证 token 文件已复制到 auth_dir
let auth_dir = sync_service.get_auth_dir().expect("获取 auth_dir 失败");
let provider_name = match provider_type {
PoolProviderType::Kiro => "kiro",
PoolProviderType::Gemini => "gemini",
PoolProviderType::Qwen => "qwen",
_ => "unknown",
};
let expected_token_path = auth_dir.join(provider_name).join(format!("{}.json", original_uuid));
prop_assert!(
expected_token_path.exists(),
"Token 文件应该存在于 auth_dir: {:?}",
expected_token_path
);
// 验证 token 文件内容一致
let copied_content = std::fs::read_to_string(&expected_token_path)
.expect("读取复制的 token 文件失败");
prop_assert_eq!(
copied_content,
token_json,
"Token 文件内容应该一致"
);
// 从配置加载凭证
let loaded = sync_service.load_from_config().expect("加载凭证失败");
let loaded_cred = loaded.iter().find(|c| c.uuid == original_uuid);
prop_assert!(loaded_cred.is_some(), "应该能找到加载的凭证");
// 验证加载的凭证指向正确的 token 文件路径
let loaded_cred = loaded_cred.unwrap();
let loaded_path = match &loaded_cred.credential {
PoolCredentialData::KiroOAuth { creds_file_path } => creds_file_path.clone(),
PoolCredentialData::GeminiOAuth { creds_file_path, .. } => creds_file_path.clone(),
PoolCredentialData::QwenOAuth { creds_file_path } => creds_file_path.clone(),
_ => String::new(),
};
prop_assert_eq!(
loaded_path,
expected_token_path.to_string_lossy().to_string(),
"加载的凭证应该指向 auth_dir 中的 token 文件"
);
// 删除凭证(应该删除 token 文件)
let remove_result = sync_service.remove_credential(provider_type, &original_uuid);
prop_assert!(remove_result.is_ok(), "删除凭证应该成功: {:?}", remove_result);
// 验证 token 文件已被删除
prop_assert!(
!expected_token_path.exists(),
"删除凭证后 token 文件应该被删除"
);
}
/// **Feature: config-credential-export, Property 9: OAuth Token File Update**
/// *For any* OAuth credential update, the token file should be updated in auth-dir.
/// **Validates: Requirements 2.1, 2.4**
#[test]
fn prop_oauth_token_file_update(
initial_content in "[a-zA-Z0-9]{50,100}",
updated_content in "[a-zA-Z0-9]{50,100}"
) {
let (temp_dir, config_manager) = create_test_env();
let sync_service = CredentialSyncService::new(config_manager.clone());
// 创建初始 token 文件
let source_token_dir = temp_dir.path().join("source_tokens");
std::fs::create_dir_all(&source_token_dir).expect("创建源目录失败");
let source_token_path = source_token_dir.join("token.json");
let initial_json = format!(r#"{{"access_token": "{}"}}"#, initial_content);
std::fs::write(&source_token_path, &initial_json).expect("写入初始 token 文件失败");
// 创建凭证
let credential = ProviderCredential::new(
PoolProviderType::Kiro,
PoolCredentialData::KiroOAuth {
creds_file_path: source_token_path.to_string_lossy().to_string(),
},
);
let original_uuid = credential.uuid.clone();
// 添加凭证
sync_service.add_credential(&credential).expect("添加凭证失败");
// 更新源 token 文件内容
let updated_json = format!(r#"{{"access_token": "{}"}}"#, updated_content);
std::fs::write(&source_token_path, &updated_json).expect("更新源 token 文件失败");
// 更新凭证
let mut updated_credential = credential.clone();
updated_credential.credential = PoolCredentialData::KiroOAuth {
creds_file_path: source_token_path.to_string_lossy().to_string(),
};
let update_result = sync_service.update_credential(&updated_credential);
prop_assert!(update_result.is_ok(), "更新凭证应该成功: {:?}", update_result);
// 验证 auth_dir 中的 token 文件已更新
let auth_dir = sync_service.get_auth_dir().expect("获取 auth_dir 失败");
let token_path = auth_dir.join("kiro").join(format!("{}.json", original_uuid));
let stored_content = std::fs::read_to_string(&token_path)
.expect("读取存储的 token 文件失败");
prop_assert_eq!(
stored_content,
updated_json,
"存储的 token 文件内容应该已更新"
);
}
}
+5
View File
@@ -184,6 +184,11 @@ impl Injector {
self.rules.iter().filter(|r| r.matches(model)).collect()
}
/// 清空所有规则
pub fn clear(&mut self) {
self.rules.clear();
}
/// 注入参数到请求
///
/// 按规则优先级顺序应用注入:
+55 -6
View File
@@ -7,9 +7,10 @@ pub mod injection;
mod logger;
mod models;
pub mod plugin;
pub mod processor;
mod providers;
pub mod resilience;
mod router;
pub mod router;
mod server;
mod services;
pub mod telemetry;
@@ -19,7 +20,7 @@ use serde::{Deserialize, Serialize};
use std::sync::Arc;
use tokio::sync::RwLock;
use commands::provider_pool_cmd::ProviderPoolServiceState;
use commands::provider_pool_cmd::{CredentialSyncServiceState, ProviderPoolServiceState};
use commands::resilience_cmd::ResilienceConfigState;
use commands::router_cmd::RouterConfigState;
use commands::skill_cmd::SkillServiceState;
@@ -1310,6 +1311,11 @@ pub fn run() {
let provider_pool_service = ProviderPoolService::new();
let provider_pool_service_state = ProviderPoolServiceState(Arc::new(provider_pool_service));
// Initialize CredentialSyncService (optional - only if config manager is available)
// For now, we initialize it as None since ConfigManager requires async setup
// This can be enhanced later to properly initialize with ConfigManager
let credential_sync_service_state = CredentialSyncServiceState(None);
// Initialize TokenCacheService
let token_cache_service = TokenCacheService::new();
let token_cache_service_state = TokenCacheServiceState(Arc::new(token_cache_service));
@@ -1320,8 +1326,25 @@ pub fn run() {
// Initialize ResilienceConfigState
let resilience_config_state = ResilienceConfigState::default();
// Initialize TelemetryState
let telemetry_state = commands::telemetry_cmd::TelemetryState::default();
// Initialize shared telemetry instances for both TelemetryState and RequestProcessor
// This allows the frontend monitoring page to display data recorded by the request processor
let shared_stats = Arc::new(parking_lot::RwLock::new(
telemetry::StatsAggregator::with_defaults(),
));
let shared_tokens = Arc::new(parking_lot::RwLock::new(
telemetry::TokenTracker::with_defaults(),
));
let shared_logger = Arc::new(
telemetry::RequestLogger::with_defaults().expect("Failed to create RequestLogger"),
);
// Initialize TelemetryState with shared instances
let telemetry_state = commands::telemetry_cmd::TelemetryState::with_shared(
shared_stats.clone(),
shared_tokens.clone(),
Some(shared_logger.clone()),
)
.expect("Failed to create TelemetryState");
// Initialize default skill repos
{
@@ -1336,6 +1359,9 @@ pub fn run() {
let db_clone = db.clone();
let pool_service_clone = provider_pool_service_state.0.clone();
let token_cache_clone = token_cache_service_state.0.clone();
let shared_stats_clone = shared_stats.clone();
let shared_tokens_clone = shared_tokens.clone();
let shared_logger_clone = shared_logger.clone();
tauri::Builder::default()
.plugin(tauri_plugin_shell::init())
@@ -1349,6 +1375,7 @@ pub fn run() {
.manage(db)
.manage(skill_service_state)
.manage(provider_pool_service_state)
.manage(credential_sync_service_state)
.manage(token_cache_service_state)
.manage(router_config_state)
.manage(resilience_config_state)
@@ -1360,6 +1387,9 @@ pub fn run() {
let db = db_clone.clone();
let pool_service = pool_service_clone.clone();
let token_cache = token_cache_clone.clone();
let shared_stats = shared_stats_clone.clone();
let shared_tokens = shared_tokens_clone.clone();
let shared_logger = shared_logger_clone.clone();
tauri::async_runtime::spawn(async move {
// 先加载凭证
{
@@ -1372,14 +1402,22 @@ pub fn run() {
logs.write().await.add("info", "[启动] Kiro 凭证已加载");
}
}
// 启动服务器
// 启动服务器(使用共享的遥测实例)
{
let mut s = state.write().await;
logs.write()
.await
.add("info", "[启动] 正在自动启动服务器...");
match s
.start(logs.clone(), pool_service, token_cache, Some(db))
.start_with_telemetry(
logs.clone(),
pool_service,
token_cache,
Some(db),
Some(shared_stats),
Some(shared_tokens),
Some(shared_logger),
)
.await
{
Ok(_) => {
@@ -1470,6 +1508,14 @@ pub fn run() {
commands::config_cmd::validate_config_yaml,
commands::config_cmd::import_config,
commands::config_cmd::get_config_paths,
// Enhanced export/import commands (using ExportService/ImportService)
commands::config_cmd::export_bundle,
commands::config_cmd::export_config_yaml,
commands::config_cmd::validate_import,
commands::config_cmd::import_bundle,
// Path utility commands
commands::config_cmd::expand_path,
commands::config_cmd::open_auth_dir,
// MCP commands
commands::mcp_cmd::get_mcp_servers,
commands::mcp_cmd::add_mcp_server,
@@ -1535,6 +1581,9 @@ pub fn run() {
commands::router_cmd::add_exclusion,
commands::router_cmd::remove_exclusion,
commands::router_cmd::set_router_default_provider,
commands::router_cmd::get_recommended_presets,
commands::router_cmd::apply_recommended_preset,
commands::router_cmd::clear_all_routing_config,
// Resilience config commands
commands::resilience_cmd::get_retry_config,
commands::resilience_cmd::update_retry_config,
+169
View File
@@ -0,0 +1,169 @@
//! 请求上下文
//!
//! 定义请求处理过程中的上下文信息
use crate::plugin::PluginContext;
use crate::ProviderType;
use chrono::{DateTime, Utc};
use std::time::Instant;
/// 请求上下文
///
/// 在请求处理管道中传递的上下文信息
#[derive(Debug, Clone)]
pub struct RequestContext {
/// 请求唯一标识
pub request_id: String,
/// 请求开始时间
pub start_time: Instant,
/// 请求时间戳
pub timestamp: DateTime<Utc>,
/// 原始模型名称(请求中的模型)
pub original_model: String,
/// 解析后的模型名称(经过别名映射)
pub resolved_model: String,
/// 选择的 Provider
pub provider: Option<ProviderType>,
/// 使用的凭证 ID
pub credential_id: Option<String>,
/// 重试次数
pub retry_count: u32,
/// 是否为流式请求
pub is_stream: bool,
/// 插件上下文
pub plugin_ctx: Option<PluginContext>,
/// 元数据
pub metadata: std::collections::HashMap<String, serde_json::Value>,
}
impl RequestContext {
/// 创建新的请求上下文
pub fn new(model: String) -> Self {
let request_id = uuid::Uuid::new_v4().to_string();
Self {
request_id: request_id.clone(),
start_time: Instant::now(),
timestamp: Utc::now(),
original_model: model.clone(),
resolved_model: model,
provider: None,
credential_id: None,
retry_count: 0,
is_stream: false,
plugin_ctx: None,
metadata: std::collections::HashMap::new(),
}
}
/// 设置流式请求标志
pub fn with_stream(mut self, is_stream: bool) -> Self {
self.is_stream = is_stream;
self
}
/// 设置 Provider
pub fn set_provider(&mut self, provider: ProviderType) {
self.provider = Some(provider);
}
/// 设置凭证 ID
pub fn set_credential_id(&mut self, credential_id: String) {
self.credential_id = Some(credential_id);
}
/// 设置解析后的模型名称
pub fn set_resolved_model(&mut self, model: String) {
self.resolved_model = model;
}
/// 增加重试计数
pub fn increment_retry(&mut self) {
self.retry_count += 1;
}
/// 获取已耗时(毫秒)
pub fn elapsed_ms(&self) -> u64 {
self.start_time.elapsed().as_millis() as u64
}
/// 初始化插件上下文
pub fn init_plugin_context(&mut self, provider: ProviderType) {
self.plugin_ctx = Some(PluginContext::new(
self.request_id.clone(),
provider,
self.resolved_model.clone(),
));
}
/// 获取插件上下文的可变引用
pub fn plugin_context_mut(&mut self) -> Option<&mut PluginContext> {
self.plugin_ctx.as_mut()
}
/// 添加元数据
pub fn set_metadata(&mut self, key: &str, value: serde_json::Value) {
self.metadata.insert(key.to_string(), value);
}
/// 获取元数据
pub fn get_metadata(&self, key: &str) -> Option<&serde_json::Value> {
self.metadata.get(key)
}
}
impl Default for RequestContext {
fn default() -> Self {
Self::new(String::new())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_request_context_new() {
let ctx = RequestContext::new("claude-sonnet-4-5".to_string());
assert!(!ctx.request_id.is_empty());
assert_eq!(ctx.original_model, "claude-sonnet-4-5");
assert_eq!(ctx.resolved_model, "claude-sonnet-4-5");
assert!(ctx.provider.is_none());
assert!(ctx.credential_id.is_none());
assert_eq!(ctx.retry_count, 0);
assert!(!ctx.is_stream);
}
#[test]
fn test_request_context_with_stream() {
let ctx = RequestContext::new("model".to_string()).with_stream(true);
assert!(ctx.is_stream);
}
#[test]
fn test_request_context_set_provider() {
let mut ctx = RequestContext::new("model".to_string());
ctx.set_provider(ProviderType::Kiro);
assert_eq!(ctx.provider, Some(ProviderType::Kiro));
}
#[test]
fn test_request_context_increment_retry() {
let mut ctx = RequestContext::new("model".to_string());
assert_eq!(ctx.retry_count, 0);
ctx.increment_retry();
assert_eq!(ctx.retry_count, 1);
ctx.increment_retry();
assert_eq!(ctx.retry_count, 2);
}
#[test]
fn test_request_context_metadata() {
let mut ctx = RequestContext::new("model".to_string());
ctx.set_metadata("key", serde_json::json!("value"));
let value = ctx.get_metadata("key");
assert!(value.is_some());
assert_eq!(value.unwrap(), &serde_json::json!("value"));
}
}
+194
View File
@@ -0,0 +1,194 @@
//! 处理错误类型
//!
//! 定义请求处理过程中可能发生的错误
use thiserror::Error;
/// 处理错误
#[derive(Error, Debug, Clone)]
pub enum ProcessError {
/// 认证失败
#[error("认证失败: {0}")]
AuthError(String),
/// 路由失败
#[error("路由失败: 无可用 Provider 处理模型 {model}")]
RoutingError { model: String },
/// Provider 调用失败
#[error("Provider 调用失败: {0}")]
ProviderError(String),
/// 重试耗尽
#[error("重试耗尽: 尝试 {attempts} 次后失败")]
RetriesExhausted { attempts: u32 },
/// 请求超时
#[error("请求超时: {timeout_ms}ms")]
Timeout { timeout_ms: u64 },
/// 流式响应空闲超时
#[error("流式响应空闲超时: {timeout_ms}ms")]
StreamIdleTimeout { timeout_ms: u64 },
/// 插件错误
#[error("插件错误: {plugin_name} - {message}")]
PluginError {
plugin_name: String,
message: String,
},
/// 参数注入错误
#[error("参数注入错误: {0}")]
InjectionError(String),
/// 凭证池错误
#[error("凭证池错误: {0}")]
CredentialPoolError(String),
/// 配置错误
#[error("配置错误: {0}")]
ConfigError(String),
/// 内部错误
#[error("内部错误: {0}")]
InternalError(String),
/// 请求被取消
#[error("请求已取消")]
Cancelled,
}
impl ProcessError {
/// 获取对应的 HTTP 状态码
pub fn status_code(&self) -> u16 {
match self {
ProcessError::AuthError(_) => 401,
ProcessError::RoutingError { .. } => 404,
ProcessError::ProviderError(_) => 502,
ProcessError::RetriesExhausted { .. } => 503,
ProcessError::Timeout { .. } => 408,
ProcessError::StreamIdleTimeout { .. } => 408,
ProcessError::PluginError { .. } => 500,
ProcessError::InjectionError(_) => 400,
ProcessError::CredentialPoolError(_) => 503,
ProcessError::ConfigError(_) => 500,
ProcessError::InternalError(_) => 500,
ProcessError::Cancelled => 499,
}
}
/// 检查是否为可重试错误
pub fn is_retryable(&self) -> bool {
matches!(
self,
ProcessError::ProviderError(_)
| ProcessError::Timeout { .. }
| ProcessError::StreamIdleTimeout { .. }
)
}
/// 检查是否应该触发故障转移
pub fn should_failover(&self) -> bool {
matches!(
self,
ProcessError::ProviderError(_)
| ProcessError::RetriesExhausted { .. }
| ProcessError::CredentialPoolError(_)
)
}
/// 转换为 JSON 错误响应
pub fn to_json(&self) -> serde_json::Value {
serde_json::json!({
"error": {
"message": self.to_string(),
"type": self.error_type(),
"code": self.status_code()
}
})
}
/// 获取错误类型字符串
pub fn error_type(&self) -> &'static str {
match self {
ProcessError::AuthError(_) => "authentication_error",
ProcessError::RoutingError { .. } => "routing_error",
ProcessError::ProviderError(_) => "provider_error",
ProcessError::RetriesExhausted { .. } => "retries_exhausted",
ProcessError::Timeout { .. } => "timeout_error",
ProcessError::StreamIdleTimeout { .. } => "stream_idle_timeout",
ProcessError::PluginError { .. } => "plugin_error",
ProcessError::InjectionError(_) => "injection_error",
ProcessError::CredentialPoolError(_) => "credential_pool_error",
ProcessError::ConfigError(_) => "config_error",
ProcessError::InternalError(_) => "internal_error",
ProcessError::Cancelled => "cancelled",
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_process_error_status_codes() {
assert_eq!(
ProcessError::AuthError("test".to_string()).status_code(),
401
);
assert_eq!(
ProcessError::RoutingError {
model: "test".to_string()
}
.status_code(),
404
);
assert_eq!(
ProcessError::ProviderError("test".to_string()).status_code(),
502
);
assert_eq!(
ProcessError::RetriesExhausted { attempts: 3 }.status_code(),
503
);
assert_eq!(
ProcessError::Timeout { timeout_ms: 5000 }.status_code(),
408
);
assert_eq!(ProcessError::Cancelled.status_code(), 499);
}
#[test]
fn test_process_error_is_retryable() {
assert!(ProcessError::ProviderError("test".to_string()).is_retryable());
assert!(ProcessError::Timeout { timeout_ms: 5000 }.is_retryable());
assert!(!ProcessError::AuthError("test".to_string()).is_retryable());
assert!(!ProcessError::RoutingError {
model: "test".to_string()
}
.is_retryable());
}
#[test]
fn test_process_error_should_failover() {
assert!(ProcessError::ProviderError("test".to_string()).should_failover());
assert!(ProcessError::RetriesExhausted { attempts: 3 }.should_failover());
assert!(!ProcessError::AuthError("test".to_string()).should_failover());
assert!(!ProcessError::Timeout { timeout_ms: 5000 }.should_failover());
}
#[test]
fn test_process_error_to_json() {
let error = ProcessError::AuthError("Invalid API key".to_string());
let json = error.to_json();
assert!(json["error"]["message"]
.as_str()
.unwrap()
.contains("Invalid API key"));
assert_eq!(json["error"]["type"], "authentication_error");
assert_eq!(json["error"]["code"], 401);
}
}
+236
View File
@@ -0,0 +1,236 @@
//! 请求处理器模块
//!
//! 提供统一的请求处理管道,集成路由、容错、监控、插件等功能模块。
//!
//! # 架构
//!
//! 请求处理流程:
//! 1. 认证 (AuthStep)
//! 2. 参数注入 (InjectionStep)
//! 3. 路由解析 (RoutingStep)
//! 4. 插件前置钩子 (PluginPreStep)
//! 5. Provider 调用 (ProviderStep) - 包含重试和故障转移
//! 6. 插件后置钩子 (PluginPostStep)
//! 7. 统计记录 (TelemetryStep)
mod context;
mod error;
mod steps;
pub use context::RequestContext;
pub use error::ProcessError;
pub use steps::{
AuthStep, InjectionStep, PipelineStep, PluginPostStep, PluginPreStep, ProviderStep,
RoutingStep, TelemetryStep,
};
use crate::injection::Injector;
use crate::plugin::PluginManager;
use crate::resilience::{Failover, Retrier, TimeoutController};
use crate::router::{ModelMapper, Router};
use crate::services::provider_pool_service::ProviderPoolService;
use crate::telemetry::{StatsAggregator, TokenTracker};
use parking_lot::RwLock as ParkingLotRwLock;
use std::sync::Arc;
use tokio::sync::RwLock;
/// 统一的请求处理器
///
/// 集成所有功能模块,提供完整的请求处理管道
pub struct RequestProcessor {
/// 路由器
pub router: Arc<RwLock<Router>>,
/// 模型映射器
pub mapper: Arc<RwLock<ModelMapper>>,
/// 参数注入器
pub injector: Arc<RwLock<Injector>>,
/// 重试器
pub retrier: Arc<Retrier>,
/// 故障转移器
pub failover: Arc<Failover>,
/// 超时控制器
pub timeout: Arc<TimeoutController>,
/// 插件管理器
pub plugins: Arc<PluginManager>,
/// 统计聚合器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享)
pub stats: Arc<ParkingLotRwLock<StatsAggregator>>,
/// Token 追踪器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享)
pub tokens: Arc<ParkingLotRwLock<TokenTracker>>,
/// 凭证池服务
pub pool_service: Arc<ProviderPoolService>,
}
impl RequestProcessor {
/// 创建新的请求处理器
pub fn new(
router: Arc<RwLock<Router>>,
mapper: Arc<RwLock<ModelMapper>>,
injector: Arc<RwLock<Injector>>,
retrier: Arc<Retrier>,
failover: Arc<Failover>,
timeout: Arc<TimeoutController>,
plugins: Arc<PluginManager>,
stats: Arc<ParkingLotRwLock<StatsAggregator>>,
tokens: Arc<ParkingLotRwLock<TokenTracker>>,
pool_service: Arc<ProviderPoolService>,
) -> Self {
Self {
router,
mapper,
injector,
retrier,
failover,
timeout,
plugins,
stats,
tokens,
pool_service,
}
}
/// 使用默认配置创建请求处理器
pub fn with_defaults(pool_service: Arc<ProviderPoolService>) -> Self {
use crate::ProviderType;
Self {
router: Arc::new(RwLock::new(Router::new(ProviderType::Kiro))),
mapper: Arc::new(RwLock::new(ModelMapper::new())),
injector: Arc::new(RwLock::new(Injector::new())),
retrier: Arc::new(Retrier::with_defaults()),
failover: Arc::new(Failover::with_defaults()),
timeout: Arc::new(TimeoutController::with_defaults()),
plugins: Arc::new(PluginManager::with_defaults()),
stats: Arc::new(ParkingLotRwLock::new(StatsAggregator::with_defaults())),
tokens: Arc::new(ParkingLotRwLock::new(TokenTracker::with_defaults())),
pool_service,
}
}
/// 使用共享的统计和 Token 追踪器创建请求处理器
///
/// 这允许 RequestProcessor 与 TelemetryState 共享同一个 StatsAggregator 和 TokenTracker,
/// 使得请求处理过程中记录的统计数据能够在前端监控页面中显示。
pub fn with_shared_telemetry(
pool_service: Arc<ProviderPoolService>,
stats: Arc<ParkingLotRwLock<StatsAggregator>>,
tokens: Arc<ParkingLotRwLock<TokenTracker>>,
) -> Self {
use crate::ProviderType;
Self {
router: Arc::new(RwLock::new(Router::new(ProviderType::Kiro))),
mapper: Arc::new(RwLock::new(ModelMapper::new())),
injector: Arc::new(RwLock::new(Injector::new())),
retrier: Arc::new(Retrier::with_defaults()),
failover: Arc::new(Failover::with_defaults()),
timeout: Arc::new(TimeoutController::with_defaults()),
plugins: Arc::new(PluginManager::with_defaults()),
stats,
tokens,
pool_service,
}
}
/// 解析模型别名
///
/// 使用 ModelMapper 将模型别名解析为实际模型名称
///
/// # Arguments
/// * `model` - 原始模型名称(可能是别名)
///
/// # Returns
/// 解析后的实际模型名称
pub async fn resolve_model(&self, model: &str) -> String {
let mapper = self.mapper.read().await;
mapper.resolve(model)
}
/// 解析模型别名并更新请求上下文
///
/// # Arguments
/// * `ctx` - 请求上下文
///
/// # Returns
/// 解析后的模型名称
pub async fn resolve_model_for_context(&self, ctx: &mut RequestContext) -> String {
let resolved = self.resolve_model(&ctx.original_model).await;
ctx.set_resolved_model(resolved.clone());
tracing::debug!(
"[MAPPER] request_id={} original_model={} resolved_model={}",
ctx.request_id,
ctx.original_model,
resolved
);
resolved
}
/// 根据模型选择 Provider
///
/// 使用 Router 根据路由规则选择合适的 Provider
///
/// # Arguments
/// * `model` - 模型名称(应该是解析后的实际模型名)
///
/// # Returns
/// 选择的 Provider 类型和是否使用默认 Provider
pub async fn route_model(&self, model: &str) -> (crate::ProviderType, bool) {
let router = self.router.read().await;
let result = router.route(model);
(result.provider, result.is_default)
}
/// 根据模型选择 Provider 并更新请求上下文
///
/// # Arguments
/// * `ctx` - 请求上下文
///
/// # Returns
/// 选择的 Provider 类型
pub async fn route_for_context(&self, ctx: &mut RequestContext) -> crate::ProviderType {
let (provider, is_default) = self.route_model(&ctx.resolved_model).await;
ctx.set_provider(provider);
tracing::info!(
"[ROUTE] request_id={} model={} provider={} is_default={}",
ctx.request_id,
ctx.resolved_model,
provider,
is_default
);
provider
}
/// 执行完整的路由解析流程
///
/// 包括模型别名解析和 Provider 选择
///
/// # Arguments
/// * `ctx` - 请求上下文
///
/// # Returns
/// 选择的 Provider 类型
pub async fn resolve_and_route(&self, ctx: &mut RequestContext) -> crate::ProviderType {
// 1. 解析模型别名
self.resolve_model_for_context(ctx).await;
// 2. 根据解析后的模型选择 Provider
self.route_for_context(ctx).await
}
/// 检查模型是否被指定 Provider 排除
///
/// # Arguments
/// * `provider` - Provider 类型
/// * `model` - 模型名称
///
/// # Returns
/// 如果模型被排除返回 true
pub async fn is_model_excluded(&self, provider: crate::ProviderType, model: &str) -> bool {
let router = self.router.read().await;
router.is_excluded(provider, model)
}
}
#[cfg(test)]
mod tests;
+105
View File
@@ -0,0 +1,105 @@
//! 认证步骤
//!
//! 验证请求的 API Key
use super::traits::{PipelineStep, StepError};
use crate::processor::RequestContext;
use async_trait::async_trait;
/// 认证步骤
///
/// 验证请求中的 API Key 是否有效
pub struct AuthStep {
/// 期望的 API Key
expected_key: String,
/// 是否启用
enabled: bool,
}
impl AuthStep {
/// 创建新的认证步骤
pub fn new(expected_key: String) -> Self {
Self {
expected_key,
enabled: true,
}
}
/// 设置是否启用
pub fn with_enabled(mut self, enabled: bool) -> Self {
self.enabled = enabled;
self
}
/// 验证 API Key
pub fn verify(&self, provided_key: Option<&str>) -> Result<(), StepError> {
match provided_key {
Some(key) if key == self.expected_key => Ok(()),
Some(_) => Err(StepError::Auth("Invalid API key".to_string())),
None => Err(StepError::Auth("No API key provided".to_string())),
}
}
}
#[async_trait]
impl PipelineStep for AuthStep {
async fn execute(
&self,
ctx: &mut RequestContext,
_payload: &mut serde_json::Value,
) -> Result<(), StepError> {
// 从元数据中获取 API Key
let api_key = ctx
.get_metadata("api_key")
.and_then(|v| v.as_str())
.map(|s| s.to_string());
self.verify(api_key.as_deref())
}
fn name(&self) -> &str {
"auth"
}
fn is_enabled(&self) -> bool {
self.enabled
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_auth_step_verify_success() {
let step = AuthStep::new("test-key".to_string());
assert!(step.verify(Some("test-key")).is_ok());
}
#[test]
fn test_auth_step_verify_invalid_key() {
let step = AuthStep::new("test-key".to_string());
let result = step.verify(Some("wrong-key"));
assert!(result.is_err());
assert!(matches!(result.unwrap_err(), StepError::Auth(_)));
}
#[test]
fn test_auth_step_verify_no_key() {
let step = AuthStep::new("test-key".to_string());
let result = step.verify(None);
assert!(result.is_err());
assert!(matches!(result.unwrap_err(), StepError::Auth(_)));
}
#[tokio::test]
async fn test_auth_step_execute() {
let step = AuthStep::new("test-key".to_string());
let mut ctx = RequestContext::new("model".to_string());
ctx.set_metadata("api_key", serde_json::json!("test-key"));
let mut payload = serde_json::json!({});
let result = step.execute(&mut ctx, &mut payload).await;
assert!(result.is_ok());
}
}
+124
View File
@@ -0,0 +1,124 @@
//! 参数注入步骤
//!
//! 根据配置的规则注入请求参数
use super::traits::{PipelineStep, StepError};
use crate::injection::Injector;
use crate::processor::RequestContext;
use async_trait::async_trait;
use std::sync::Arc;
use tokio::sync::RwLock;
/// 参数注入步骤
///
/// 根据模型匹配规则注入请求参数
pub struct InjectionStep {
/// 注入器
injector: Arc<RwLock<Injector>>,
/// 是否启用
enabled: Arc<RwLock<bool>>,
}
impl InjectionStep {
/// 创建新的注入步骤
pub fn new(injector: Arc<RwLock<Injector>>) -> Self {
Self {
injector,
enabled: Arc::new(RwLock::new(true)),
}
}
/// 设置是否启用
pub fn with_enabled(self, enabled: Arc<RwLock<bool>>) -> Self {
Self { enabled, ..self }
}
/// 检查是否启用
pub async fn is_injection_enabled(&self) -> bool {
*self.enabled.read().await
}
}
#[async_trait]
impl PipelineStep for InjectionStep {
async fn execute(
&self,
ctx: &mut RequestContext,
payload: &mut serde_json::Value,
) -> Result<(), StepError> {
if !self.is_injection_enabled().await {
return Ok(());
}
let injector = self.injector.read().await;
let result = injector.inject(&ctx.resolved_model, payload);
if result.has_injections() {
tracing::info!(
"[INJECT] request_id={} applied_rules={:?} injected_params={:?}",
ctx.request_id,
result.applied_rules,
result.injected_params
);
// 记录注入信息到元数据
ctx.set_metadata(
"injection_result",
serde_json::json!({
"applied_rules": result.applied_rules,
"injected_params": result.injected_params
}),
);
}
Ok(())
}
fn name(&self) -> &str {
"injection"
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::injection::InjectionRule;
#[tokio::test]
async fn test_injection_step_execute() {
let mut injector = Injector::new();
injector.add_rule(InjectionRule::new(
"test-rule",
"claude-*",
serde_json::json!({"temperature": 0.7}),
));
let step = InjectionStep::new(Arc::new(RwLock::new(injector)));
let mut ctx = RequestContext::new("claude-sonnet-4-5".to_string());
let mut payload = serde_json::json!({"model": "claude-sonnet-4-5"});
let result = step.execute(&mut ctx, &mut payload).await;
assert!(result.is_ok());
assert_eq!(payload["temperature"], 0.7);
}
#[tokio::test]
async fn test_injection_step_disabled() {
let mut injector = Injector::new();
injector.add_rule(InjectionRule::new(
"test-rule",
"claude-*",
serde_json::json!({"temperature": 0.7}),
));
let step = InjectionStep::new(Arc::new(RwLock::new(injector)))
.with_enabled(Arc::new(RwLock::new(false)));
let mut ctx = RequestContext::new("claude-sonnet-4-5".to_string());
let mut payload = serde_json::json!({"model": "claude-sonnet-4-5"});
let result = step.execute(&mut ctx, &mut payload).await;
assert!(result.is_ok());
// 参数不应该被注入
assert!(payload.get("temperature").is_none());
}
}
+19
View File
@@ -0,0 +1,19 @@
//! 管道步骤模块
//!
//! 定义请求处理管道中的各个步骤
mod auth;
mod injection;
mod plugin;
mod provider;
mod routing;
mod telemetry;
mod traits;
pub use auth::AuthStep;
pub use injection::InjectionStep;
pub use plugin::{PluginPostStep, PluginPreStep};
pub use provider::ProviderStep;
pub use routing::RoutingStep;
pub use telemetry::TelemetryStep;
pub use traits::{PipelineStep, StepError};
+170
View File
@@ -0,0 +1,170 @@
//! 插件钩子步骤
//!
//! 执行插件的前置和后置钩子
use super::traits::{PipelineStep, StepError};
use crate::plugin::PluginManager;
use crate::processor::RequestContext;
use crate::ProviderType;
use async_trait::async_trait;
use std::sync::Arc;
/// 插件前置钩子步骤
///
/// 在 Provider 调用前执行所有启用插件的 on_request 钩子
pub struct PluginPreStep {
/// 插件管理器
plugins: Arc<PluginManager>,
}
impl PluginPreStep {
/// 创建新的插件前置步骤
pub fn new(plugins: Arc<PluginManager>) -> Self {
Self { plugins }
}
}
#[async_trait]
impl PipelineStep for PluginPreStep {
async fn execute(
&self,
ctx: &mut RequestContext,
payload: &mut serde_json::Value,
) -> Result<(), StepError> {
// 初始化插件上下文
let provider = ctx.provider.unwrap_or(ProviderType::Kiro);
ctx.init_plugin_context(provider);
// 获取插件上下文的可变引用
if let Some(plugin_ctx) = ctx.plugin_context_mut() {
let results = self.plugins.run_on_request(plugin_ctx, payload).await;
// 检查是否有失败的钩子
for result in &results {
if !result.success {
tracing::warn!("[PLUGIN] on_request hook failed: {:?}", result.error);
// 插件失败不阻止请求继续,只记录警告
}
}
// 记录插件执行结果到元数据
ctx.set_metadata(
"plugin_pre_results",
serde_json::json!(results
.iter()
.map(|r| serde_json::json!({
"success": r.success,
"modified": r.modified,
"duration_ms": r.duration_ms
}))
.collect::<Vec<_>>()),
);
}
Ok(())
}
fn name(&self) -> &str {
"plugin_pre"
}
}
/// 插件后置钩子步骤
///
/// 在 Provider 调用后执行所有启用插件的 on_response 钩子
pub struct PluginPostStep {
/// 插件管理器
plugins: Arc<PluginManager>,
}
impl PluginPostStep {
/// 创建新的插件后置步骤
pub fn new(plugins: Arc<PluginManager>) -> Self {
Self { plugins }
}
/// 执行错误钩子
pub async fn run_on_error(&self, ctx: &mut RequestContext, error: &str) {
if let Some(plugin_ctx) = ctx.plugin_context_mut() {
let results = self.plugins.run_on_error(plugin_ctx, error).await;
for result in &results {
if !result.success {
tracing::warn!("[PLUGIN] on_error hook failed: {:?}", result.error);
}
}
}
}
}
#[async_trait]
impl PipelineStep for PluginPostStep {
async fn execute(
&self,
ctx: &mut RequestContext,
payload: &mut serde_json::Value,
) -> Result<(), StepError> {
if let Some(plugin_ctx) = ctx.plugin_context_mut() {
let results = self.plugins.run_on_response(plugin_ctx, payload).await;
// 检查是否有失败的钩子
for result in &results {
if !result.success {
tracing::warn!("[PLUGIN] on_response hook failed: {:?}", result.error);
}
}
// 记录插件执行结果到元数据
ctx.set_metadata(
"plugin_post_results",
serde_json::json!(results
.iter()
.map(|r| serde_json::json!({
"success": r.success,
"modified": r.modified,
"duration_ms": r.duration_ms
}))
.collect::<Vec<_>>()),
);
}
Ok(())
}
fn name(&self) -> &str {
"plugin_post"
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_plugin_pre_step_execute() {
let plugins = Arc::new(PluginManager::with_defaults());
let step = PluginPreStep::new(plugins);
let mut ctx = RequestContext::new("model".to_string());
ctx.set_provider(ProviderType::Kiro);
let mut payload = serde_json::json!({"model": "model"});
let result = step.execute(&mut ctx, &mut payload).await;
assert!(result.is_ok());
assert!(ctx.plugin_ctx.is_some());
}
#[tokio::test]
async fn test_plugin_post_step_execute() {
let plugins = Arc::new(PluginManager::with_defaults());
let step = PluginPostStep::new(plugins);
let mut ctx = RequestContext::new("model".to_string());
ctx.set_provider(ProviderType::Kiro);
ctx.init_plugin_context(ProviderType::Kiro);
let mut payload = serde_json::json!({"response": "test"});
let result = step.execute(&mut ctx, &mut payload).await;
assert!(result.is_ok());
}
}
+729
View File
@@ -0,0 +1,729 @@
//! Provider 调用步骤
//!
//! 集成重试、故障转移和超时控制
use super::traits::{PipelineStep, StepError};
use crate::processor::RequestContext;
use crate::resilience::{
Failover, FailoverConfig, FailoverManager, Retrier, RetryConfig, TimeoutConfig,
TimeoutController, TimeoutError,
};
use crate::services::provider_pool_service::ProviderPoolService;
use crate::ProviderType;
use async_trait::async_trait;
use std::future::Future;
use std::sync::Arc;
/// Provider 调用结果
#[derive(Debug, Clone)]
pub struct ProviderCallResult {
/// 响应内容
pub response: serde_json::Value,
/// HTTP 状态码
pub status_code: u16,
/// 延迟(毫秒)
pub latency_ms: u64,
/// 使用的凭证 ID
pub credential_id: Option<String>,
}
/// Provider 调用错误
#[derive(Debug, Clone)]
pub struct ProviderCallError {
/// 错误消息
pub message: String,
/// HTTP 状态码(如果有)
pub status_code: Option<u16>,
/// 是否可重试
pub retryable: bool,
/// 是否应触发故障转移
pub should_failover: bool,
}
impl ProviderCallError {
/// 创建可重试错误
pub fn retryable(message: impl Into<String>, status_code: Option<u16>) -> Self {
Self {
message: message.into(),
status_code,
retryable: true,
should_failover: false,
}
}
/// 创建需要故障转移的错误
pub fn failover(message: impl Into<String>, status_code: Option<u16>) -> Self {
Self {
message: message.into(),
status_code,
retryable: false,
should_failover: true,
}
}
/// 创建不可恢复错误
pub fn fatal(message: impl Into<String>, status_code: Option<u16>) -> Self {
Self {
message: message.into(),
status_code,
retryable: false,
should_failover: false,
}
}
/// 检查是否为配额超限错误
pub fn is_quota_exceeded(&self) -> bool {
Failover::is_quota_exceeded(self.status_code, &self.message)
}
}
/// Provider 调用步骤
///
/// 包含重试、故障转移和超时控制的 Provider 调用
pub struct ProviderStep {
/// 重试器
retrier: Arc<Retrier>,
/// 故障转移器
failover: Arc<Failover>,
/// 超时控制器
timeout: Arc<TimeoutController>,
/// 凭证池服务
pool_service: Arc<ProviderPoolService>,
}
impl ProviderStep {
/// 创建新的 Provider 步骤
pub fn new(
retrier: Arc<Retrier>,
failover: Arc<Failover>,
timeout: Arc<TimeoutController>,
pool_service: Arc<ProviderPoolService>,
) -> Self {
Self {
retrier,
failover,
timeout,
pool_service,
}
}
/// 使用默认配置创建
pub fn with_defaults(pool_service: Arc<ProviderPoolService>) -> Self {
Self {
retrier: Arc::new(Retrier::with_defaults()),
failover: Arc::new(Failover::new(FailoverConfig::default())),
timeout: Arc::new(TimeoutController::with_defaults()),
pool_service,
}
}
/// 使用自定义配置创建
pub fn with_config(
retry_config: RetryConfig,
failover_config: FailoverConfig,
timeout_config: TimeoutConfig,
pool_service: Arc<ProviderPoolService>,
) -> Self {
Self {
retrier: Arc::new(Retrier::new(retry_config)),
failover: Arc::new(Failover::new(failover_config)),
timeout: Arc::new(TimeoutController::new(timeout_config)),
pool_service,
}
}
/// 获取重试器
pub fn retrier(&self) -> &Retrier {
&self.retrier
}
/// 获取故障转移器
pub fn failover(&self) -> &Failover {
&self.failover
}
/// 获取超时控制器
pub fn timeout(&self) -> &TimeoutController {
&self.timeout
}
/// 获取凭证池服务
pub fn pool_service(&self) -> &ProviderPoolService {
&self.pool_service
}
/// 带重试执行 Provider 调用
///
/// 使用 Retrier 包装 Provider 调用,自动处理可重试错误
///
/// # Arguments
/// * `ctx` - 请求上下文
/// * `operation` - Provider 调用操作
///
/// # Returns
/// 成功返回调用结果,失败返回错误
pub async fn execute_with_retry<F, Fut>(
&self,
ctx: &mut RequestContext,
mut operation: F,
) -> Result<ProviderCallResult, ProviderCallError>
where
F: FnMut() -> Fut,
Fut: Future<Output = Result<ProviderCallResult, ProviderCallError>>,
{
let max_retries = self.retrier.config().max_retries;
let mut attempts = 0u32;
loop {
attempts += 1;
match operation().await {
Ok(result) => return Ok(result),
Err(err) => {
// 增加重试计数
ctx.increment_retry();
tracing::warn!(
"[RETRY] request_id={} attempt={}/{} error={} status={:?} retryable={}",
ctx.request_id,
attempts,
max_retries + 1,
err.message,
err.status_code,
err.retryable
);
// 如果不可重试,立即返回
if !err.retryable {
return Err(err);
}
// 检查状态码是否可重试
let should_retry = err
.status_code
.map_or(true, |code| self.retrier.config().is_retryable(code));
let should_failover = err.should_failover || err.is_quota_exceeded();
if !should_retry || attempts > max_retries {
return Err(ProviderCallError {
message: err.message,
status_code: err.status_code,
retryable: false,
should_failover,
});
}
// 等待退避时间
let delay = self.retrier.backoff_delay(attempts - 1);
tokio::time::sleep(delay).await;
}
}
}
}
/// 带超时执行 Provider 调用
///
/// 使用 TimeoutController 包装 Provider 调用,自动处理超时
///
/// # Arguments
/// * `ctx` - 请求上下文
/// * `operation` - Provider 调用操作
///
/// # Returns
/// 成功返回调用结果,失败返回错误
pub async fn execute_with_timeout<F>(
&self,
ctx: &RequestContext,
operation: F,
) -> Result<ProviderCallResult, ProviderCallError>
where
F: Future<Output = Result<ProviderCallResult, ProviderCallError>>,
{
let timeout_result = self.timeout.execute_with_timeout(operation).await;
match timeout_result {
Ok(call_result) => call_result,
Err(timeout_err) => {
let timeout_ms = match &timeout_err {
TimeoutError::RequestTimeout { timeout_ms, .. } => *timeout_ms,
TimeoutError::StreamIdleTimeout { timeout_ms, .. } => *timeout_ms,
TimeoutError::Cancelled => 0,
};
tracing::warn!(
"[TIMEOUT] request_id={} error={} timeout_ms={}",
ctx.request_id,
timeout_err,
timeout_ms
);
Err(ProviderCallError {
message: timeout_err.to_string(),
status_code: Some(408),
retryable: true,
should_failover: false,
})
}
}
}
/// 带故障转移执行 Provider 调用
///
/// 使用 Failover 处理 Provider 失败,自动切换到其他 Provider
///
/// # Arguments
/// * `ctx` - 请求上下文
/// * `error` - Provider 调用错误
/// * `available_providers` - 可用的 Provider 列表
///
/// # Returns
/// 如果可以故障转移,返回新的 Provider;否则返回 None
pub fn handle_failover(
&self,
ctx: &RequestContext,
error: &ProviderCallError,
available_providers: &[ProviderType],
) -> Option<ProviderType> {
let current_provider = ctx.provider?;
let result = self.failover.handle_failure(
current_provider,
error.status_code,
&error.message,
available_providers,
);
if result.switched {
tracing::info!(
"[FAILOVER] request_id={} from={} to={:?} reason={:?}",
ctx.request_id,
current_provider,
result.new_provider,
result.failure_type
);
result.new_provider
} else {
tracing::warn!(
"[FAILOVER] request_id={} provider={} no_switch reason={}",
ctx.request_id,
current_provider,
result.message
);
None
}
}
/// 带重试、超时和故障转移执行完整的 Provider 调用
///
/// 这是主要的调用入口,集成了所有容错机制
///
/// # Arguments
/// * `ctx` - 请求上下文
/// * `operation` - Provider 调用操作工厂
/// * `available_providers` - 可用的 Provider 列表
///
/// # Returns
/// 成功返回调用结果,失败返回 StepError
pub async fn execute_with_resilience<F, Fut>(
&self,
ctx: &mut RequestContext,
mut operation_factory: F,
available_providers: &[ProviderType],
) -> Result<ProviderCallResult, StepError>
where
F: FnMut(ProviderType) -> Fut,
Fut: Future<Output = Result<ProviderCallResult, ProviderCallError>>,
{
let mut failover_manager = FailoverManager::new(self.failover.config().clone());
let mut current_provider = ctx.provider.unwrap_or(ProviderType::Kiro);
let max_failover_attempts = available_providers.len();
let mut failover_attempts = 0;
let max_retries = self.retrier.config().max_retries;
'failover: loop {
// 更新上下文中的 Provider
ctx.set_provider(current_provider);
ctx.retry_count = 0; // 重置重试计数
tracing::info!(
"[PROVIDER] request_id={} provider={} model={} failover_attempt={}",
ctx.request_id,
current_provider,
ctx.resolved_model,
failover_attempts
);
// 重试循环
let mut retry_attempts = 0u32;
let result: Result<ProviderCallResult, ProviderCallError> = loop {
retry_attempts += 1;
// 带超时执行调用
let call_result = self
.execute_with_timeout(ctx, operation_factory(current_provider))
.await;
match call_result {
Ok(result) => break Ok(result),
Err(err) => {
ctx.increment_retry();
tracing::warn!(
"[RETRY] request_id={} attempt={}/{} error={} status={:?} retryable={}",
ctx.request_id,
retry_attempts,
max_retries + 1,
err.message,
err.status_code,
err.retryable
);
// 如果不可重试,立即返回错误
if !err.retryable {
break Err(err);
}
// 检查状态码是否可重试
let should_retry = err
.status_code
.map_or(true, |code| self.retrier.config().is_retryable(code));
let should_failover = err.should_failover || err.is_quota_exceeded();
if !should_retry || retry_attempts > max_retries {
break Err(ProviderCallError {
message: err.message,
status_code: err.status_code,
retryable: false,
should_failover,
});
}
// 等待退避时间
let delay = self.retrier.backoff_delay(retry_attempts - 1);
tokio::time::sleep(delay).await;
}
}
};
match result {
Ok(call_result) => {
return Ok(call_result);
}
Err(err) => {
// 检查是否应该故障转移
if err.should_failover || err.is_quota_exceeded() {
failover_attempts += 1;
if failover_attempts >= max_failover_attempts {
tracing::error!(
"[PROVIDER] request_id={} all_providers_failed attempts={}",
ctx.request_id,
failover_attempts
);
return Err(StepError::Provider(format!(
"所有 Provider 都失败: {}",
err.message
)));
}
// 尝试故障转移
let failover_result = failover_manager.handle_failure_and_switch(
current_provider,
err.status_code,
&err.message,
available_providers,
);
if let Some(new_provider) = failover_result.new_provider {
tracing::info!(
"[FAILOVER] request_id={} from={} to={} reason={:?}",
ctx.request_id,
current_provider,
new_provider,
failover_result.failure_type
);
current_provider = new_provider;
continue 'failover;
}
}
// 无法恢复,返回错误
return Err(StepError::Provider(err.message));
}
}
}
}
/// 检查错误是否为配额超限
pub fn is_quota_exceeded_error(&self, error: &ProviderCallError) -> bool {
error.is_quota_exceeded()
}
/// 检查状态码是否可重试
pub fn is_retryable_status(&self, status_code: u16) -> bool {
self.retrier.config().is_retryable(status_code)
}
}
#[async_trait]
impl PipelineStep for ProviderStep {
async fn execute(
&self,
ctx: &mut RequestContext,
_payload: &mut serde_json::Value,
) -> Result<(), StepError> {
// 注意:实际的 Provider 调用逻辑在 server.rs 中实现
// 这里的 execute 方法主要用于管道步骤的统一接口
// 实际调用应使用 execute_with_resilience 方法
tracing::info!(
"[PROVIDER] request_id={} provider={:?} model={} retry_count={}",
ctx.request_id,
ctx.provider,
ctx.resolved_model,
ctx.retry_count
);
// 占位实现 - 实际调用通过 execute_with_resilience 进行
Ok(())
}
fn name(&self) -> &str {
"provider"
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[tokio::test]
async fn test_provider_step_new() {
let pool_service = Arc::new(ProviderPoolService::new());
let step = ProviderStep::with_defaults(pool_service);
assert_eq!(step.name(), "provider");
assert!(step.is_enabled());
}
#[tokio::test]
async fn test_provider_step_execute() {
let pool_service = Arc::new(ProviderPoolService::new());
let step = ProviderStep::with_defaults(pool_service);
let mut ctx = RequestContext::new("claude-sonnet-4-5".to_string());
let mut payload = serde_json::json!({"model": "claude-sonnet-4-5"});
let result = step.execute(&mut ctx, &mut payload).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_provider_step_with_config() {
let pool_service = Arc::new(ProviderPoolService::new());
let retry_config = RetryConfig::new(5, 500, 10000);
let failover_config = FailoverConfig::new(true, true);
let timeout_config = TimeoutConfig::new(60000, 15000);
let step = ProviderStep::with_config(
retry_config.clone(),
failover_config.clone(),
timeout_config.clone(),
pool_service,
);
assert_eq!(step.retrier().config().max_retries, 5);
assert_eq!(step.retrier().config().base_delay_ms, 500);
assert!(step.failover().config().auto_switch);
assert_eq!(step.timeout().config().request_timeout_ms, 60000);
}
#[test]
fn test_provider_call_error_retryable() {
let err = ProviderCallError::retryable("Connection timeout", Some(408));
assert!(err.retryable);
assert!(!err.should_failover);
assert_eq!(err.status_code, Some(408));
}
#[test]
fn test_provider_call_error_failover() {
let err = ProviderCallError::failover("Rate limit exceeded", Some(429));
assert!(!err.retryable);
assert!(err.should_failover);
assert!(err.is_quota_exceeded());
}
#[test]
fn test_provider_call_error_fatal() {
let err = ProviderCallError::fatal("Invalid API key", Some(401));
assert!(!err.retryable);
assert!(!err.should_failover);
}
#[test]
fn test_is_quota_exceeded_by_status() {
let err = ProviderCallError::retryable("Error", Some(429));
assert!(err.is_quota_exceeded());
}
#[test]
fn test_is_quota_exceeded_by_message() {
let err = ProviderCallError::retryable("Rate limit exceeded", Some(400));
assert!(err.is_quota_exceeded());
let err2 = ProviderCallError::retryable("Quota exceeded for this API", Some(400));
assert!(err2.is_quota_exceeded());
}
#[test]
fn test_is_retryable_status() {
let pool_service = Arc::new(ProviderPoolService::new());
let step = ProviderStep::with_defaults(pool_service);
// 可重试状态码
assert!(step.is_retryable_status(408));
assert!(step.is_retryable_status(429));
assert!(step.is_retryable_status(500));
assert!(step.is_retryable_status(502));
assert!(step.is_retryable_status(503));
assert!(step.is_retryable_status(504));
// 不可重试状态码
assert!(!step.is_retryable_status(200));
assert!(!step.is_retryable_status(400));
assert!(!step.is_retryable_status(401));
assert!(!step.is_retryable_status(403));
assert!(!step.is_retryable_status(404));
}
#[tokio::test]
async fn test_execute_with_retry_success() {
let pool_service = Arc::new(ProviderPoolService::new());
let step = ProviderStep::with_defaults(pool_service);
let mut ctx = RequestContext::new("test-model".to_string());
let result = step
.execute_with_retry(&mut ctx, || async {
Ok(ProviderCallResult {
response: serde_json::json!({"content": "Hello"}),
status_code: 200,
latency_ms: 100,
credential_id: Some("cred-1".to_string()),
})
})
.await;
assert!(result.is_ok());
let call_result = result.unwrap();
assert_eq!(call_result.status_code, 200);
}
#[tokio::test]
async fn test_execute_with_retry_non_retryable_error() {
let pool_service = Arc::new(ProviderPoolService::new());
let step = ProviderStep::with_defaults(pool_service);
let mut ctx = RequestContext::new("test-model".to_string());
let result = step
.execute_with_retry(&mut ctx, || async {
Err(ProviderCallError::fatal("Invalid API key", Some(401)))
})
.await;
assert!(result.is_err());
let err = result.unwrap_err();
assert_eq!(err.status_code, Some(401)); // 保留原始状态码
assert!(!err.retryable);
}
#[tokio::test]
async fn test_handle_failover() {
let pool_service = Arc::new(ProviderPoolService::new());
let step = ProviderStep::with_defaults(pool_service);
let mut ctx = RequestContext::new("test-model".to_string());
ctx.set_provider(ProviderType::Kiro);
let error = ProviderCallError::failover("Rate limit exceeded", Some(429));
let available = vec![ProviderType::Kiro, ProviderType::Gemini, ProviderType::Qwen];
let new_provider = step.handle_failover(&ctx, &error, &available);
assert!(new_provider.is_some());
assert_eq!(new_provider.unwrap(), ProviderType::Gemini);
}
#[tokio::test]
async fn test_handle_failover_no_alternative() {
let pool_service = Arc::new(ProviderPoolService::new());
let step = ProviderStep::with_defaults(pool_service);
let mut ctx = RequestContext::new("test-model".to_string());
ctx.set_provider(ProviderType::Kiro);
let error = ProviderCallError::failover("Rate limit exceeded", Some(429));
let available = vec![ProviderType::Kiro]; // 只有一个 Provider
let new_provider = step.handle_failover(&ctx, &error, &available);
assert!(new_provider.is_none());
}
#[tokio::test]
async fn test_execute_with_timeout_success() {
let pool_service = Arc::new(ProviderPoolService::new());
let timeout_config = TimeoutConfig::new(5000, 1000);
let step = ProviderStep::with_config(
RetryConfig::default(),
FailoverConfig::default(),
timeout_config,
pool_service,
);
let ctx = RequestContext::new("test-model".to_string());
let result = step
.execute_with_timeout(&ctx, async {
Ok(ProviderCallResult {
response: serde_json::json!({"content": "Hello"}),
status_code: 200,
latency_ms: 50,
credential_id: None,
})
})
.await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_execute_with_timeout_timeout() {
let pool_service = Arc::new(ProviderPoolService::new());
let timeout_config = TimeoutConfig::new(50, 0); // 50ms 超时
let step = ProviderStep::with_config(
RetryConfig::default(),
FailoverConfig::default(),
timeout_config,
pool_service,
);
let ctx = RequestContext::new("test-model".to_string());
let result = step
.execute_with_timeout(&ctx, async {
tokio::time::sleep(Duration::from_millis(200)).await;
Ok(ProviderCallResult {
response: serde_json::json!({}),
status_code: 200,
latency_ms: 200,
credential_id: None,
})
})
.await;
assert!(result.is_err());
let err = result.unwrap_err();
assert_eq!(err.status_code, Some(408));
assert!(err.retryable);
}
}
+158
View File
@@ -0,0 +1,158 @@
//! 路由解析步骤
//!
//! 解析模型别名并选择 Provider
use super::traits::{PipelineStep, StepError};
use crate::processor::RequestContext;
use crate::router::{ModelMapper, Router};
use crate::ProviderType;
use async_trait::async_trait;
use std::sync::Arc;
use tokio::sync::RwLock;
/// 路由解析步骤
///
/// 解析模型别名并根据路由规则选择 Provider
pub struct RoutingStep {
/// 路由器
router: Arc<RwLock<Router>>,
/// 模型映射器
mapper: Arc<RwLock<ModelMapper>>,
/// 默认 Provider
default_provider: Arc<RwLock<String>>,
}
impl RoutingStep {
/// 创建新的路由步骤
pub fn new(
router: Arc<RwLock<Router>>,
mapper: Arc<RwLock<ModelMapper>>,
default_provider: Arc<RwLock<String>>,
) -> Self {
Self {
router,
mapper,
default_provider,
}
}
/// 解析模型别名
pub async fn resolve_model(&self, model: &str) -> String {
let mapper = self.mapper.read().await;
mapper.resolve(model)
}
/// 根据模型选择 Provider
pub async fn select_provider(&self, model: &str) -> Result<ProviderType, StepError> {
let router = self.router.read().await;
// 使用路由规则(如果没有匹配的规则,会返回默认 Provider)
let result = router.route(model);
Ok(result.provider)
}
}
#[async_trait]
impl PipelineStep for RoutingStep {
async fn execute(
&self,
ctx: &mut RequestContext,
payload: &mut serde_json::Value,
) -> Result<(), StepError> {
// 解析模型别名
let resolved_model = self.resolve_model(&ctx.original_model).await;
ctx.set_resolved_model(resolved_model.clone());
// 更新 payload 中的模型名
if let Some(obj) = payload.as_object_mut() {
obj.insert("model".to_string(), serde_json::json!(resolved_model));
}
// 选择 Provider
let provider = self.select_provider(&ctx.resolved_model).await?;
ctx.set_provider(provider);
tracing::info!(
"[ROUTE] request_id={} original_model={} resolved_model={} provider={}",
ctx.request_id,
ctx.original_model,
ctx.resolved_model,
provider
);
Ok(())
}
fn name(&self) -> &str {
"routing"
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::router::RoutingRule;
#[tokio::test]
async fn test_routing_step_resolve_model() {
let mut mapper = ModelMapper::new();
mapper.add_alias("gpt-4", "claude-sonnet-4-5");
let step = RoutingStep::new(
Arc::new(RwLock::new(Router::new(ProviderType::Kiro))),
Arc::new(RwLock::new(mapper)),
Arc::new(RwLock::new("kiro".to_string())),
);
// 别名应解析为实际模型
let resolved = step.resolve_model("gpt-4").await;
assert_eq!(resolved, "claude-sonnet-4-5");
// 非别名应返回原值
let resolved = step.resolve_model("unknown-model").await;
assert_eq!(resolved, "unknown-model");
}
#[tokio::test]
async fn test_routing_step_select_provider() {
let mut router = Router::new(ProviderType::Kiro);
router.add_rule(RoutingRule::new("gemini-*", ProviderType::Gemini, 10));
let step = RoutingStep::new(
Arc::new(RwLock::new(router)),
Arc::new(RwLock::new(ModelMapper::new())),
Arc::new(RwLock::new("kiro".to_string())),
);
// 匹配路由规则
let provider = step.select_provider("gemini-2.5-flash").await;
assert!(provider.is_ok());
assert_eq!(provider.unwrap(), ProviderType::Gemini);
// 使用默认 Provider
let provider = step.select_provider("claude-sonnet-4-5").await;
assert!(provider.is_ok());
assert_eq!(provider.unwrap(), ProviderType::Kiro);
}
#[tokio::test]
async fn test_routing_step_execute() {
let mut mapper = ModelMapper::new();
mapper.add_alias("gpt-4", "claude-sonnet-4-5");
let step = RoutingStep::new(
Arc::new(RwLock::new(Router::new(ProviderType::Kiro))),
Arc::new(RwLock::new(mapper)),
Arc::new(RwLock::new("kiro".to_string())),
);
let mut ctx = RequestContext::new("gpt-4".to_string());
let mut payload = serde_json::json!({"model": "gpt-4"});
let result = step.execute(&mut ctx, &mut payload).await;
assert!(result.is_ok());
assert_eq!(ctx.resolved_model, "claude-sonnet-4-5");
assert_eq!(ctx.provider, Some(ProviderType::Kiro));
assert_eq!(payload["model"], "claude-sonnet-4-5");
}
}
+247
View File
@@ -0,0 +1,247 @@
//! 统计记录步骤
//!
//! 记录请求统计和 Token 使用
use super::traits::{PipelineStep, StepError};
use crate::processor::RequestContext;
use crate::telemetry::{
RequestLog, RequestStatus, StatsAggregator, TokenSource, TokenTracker, TokenUsageRecord,
};
use crate::ProviderType;
use async_trait::async_trait;
use parking_lot::RwLock;
use std::sync::Arc;
/// 统计记录步骤
///
/// 记录请求统计和 Token 使用信息
pub struct TelemetryStep {
/// 统计聚合器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享)
stats: Arc<RwLock<StatsAggregator>>,
/// Token 追踪器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享)
tokens: Arc<RwLock<TokenTracker>>,
}
impl TelemetryStep {
/// 创建新的统计记录步骤
pub fn new(stats: Arc<RwLock<StatsAggregator>>, tokens: Arc<RwLock<TokenTracker>>) -> Self {
Self { stats, tokens }
}
/// 记录请求日志
///
/// 请求完成后记录统计,按 Provider 和模型分组
/// _需求: 4.1_
pub fn record_request(
&self,
ctx: &RequestContext,
status: RequestStatus,
error_message: Option<String>,
) {
let provider = ctx.provider.unwrap_or(ProviderType::Kiro);
let mut log = RequestLog::new(
ctx.request_id.clone(),
provider,
ctx.resolved_model.clone(),
ctx.is_stream,
);
// 设置状态和持续时间
match status {
RequestStatus::Success => log.mark_success(ctx.elapsed_ms(), 200),
RequestStatus::Failed => {
log.mark_failed(ctx.elapsed_ms(), None, error_message.unwrap_or_default())
}
RequestStatus::Timeout => log.mark_timeout(ctx.elapsed_ms()),
RequestStatus::Cancelled => log.mark_cancelled(ctx.elapsed_ms()),
RequestStatus::Retrying => {
log.duration_ms = ctx.elapsed_ms();
}
}
// 设置凭证 ID
if let Some(cred_id) = &ctx.credential_id {
log.set_credential_id(cred_id.clone());
}
// 设置重试次数
log.retry_count = ctx.retry_count;
// 使用 parking_lot::RwLock 的同步写锁
let stats = self.stats.write();
stats.record(log);
}
/// 记录 Token 使用
///
/// 从响应提取 Token 数,无 Token 时使用估算
/// _需求: 4.2, 4.3_
pub fn record_tokens(
&self,
ctx: &RequestContext,
input_tokens: Option<u32>,
output_tokens: Option<u32>,
source: TokenSource,
) {
let provider = ctx.provider.unwrap_or(ProviderType::Kiro);
// 只有当至少有一个 Token 值时才记录
if input_tokens.is_some() || output_tokens.is_some() {
let record = TokenUsageRecord::new(
uuid::Uuid::new_v4().to_string(),
provider,
ctx.resolved_model.clone(),
input_tokens.unwrap_or(0),
output_tokens.unwrap_or(0),
source,
)
.with_request_id(ctx.request_id.clone());
// 使用 parking_lot::RwLock 的同步写锁
let tokens = self.tokens.write();
tokens.record(record);
}
}
/// 从响应中提取并记录 Token 使用
///
/// 支持 OpenAI 和 Anthropic 两种响应格式
pub fn record_tokens_from_response(&self, ctx: &RequestContext, response: &serde_json::Value) {
// 尝试从 OpenAI 格式响应中提取 Token
if let Some(usage) = response.get("usage") {
let input_tokens = usage
.get("prompt_tokens")
.and_then(|v| v.as_u64())
.map(|v| v as u32);
let output_tokens = usage
.get("completion_tokens")
.and_then(|v| v.as_u64())
.map(|v| v as u32);
if input_tokens.is_some() || output_tokens.is_some() {
self.record_tokens(ctx, input_tokens, output_tokens, TokenSource::Actual);
return;
}
}
// 尝试从 Anthropic 格式响应中提取 Token
if let Some(usage) = response.get("usage") {
let input_tokens = usage
.get("input_tokens")
.and_then(|v| v.as_u64())
.map(|v| v as u32);
let output_tokens = usage
.get("output_tokens")
.and_then(|v| v.as_u64())
.map(|v| v as u32);
if input_tokens.is_some() || output_tokens.is_some() {
self.record_tokens(ctx, input_tokens, output_tokens, TokenSource::Actual);
}
}
}
}
#[async_trait]
impl PipelineStep for TelemetryStep {
async fn execute(
&self,
ctx: &mut RequestContext,
payload: &mut serde_json::Value,
) -> Result<(), StepError> {
// 记录成功的请求(同步方法,使用 parking_lot::RwLock)
self.record_request(ctx, RequestStatus::Success, None);
// 从响应中提取并记录 Token(同步方法)
self.record_tokens_from_response(ctx, payload);
tracing::info!(
"[TELEMETRY] request_id={} provider={:?} model={} duration_ms={}",
ctx.request_id,
ctx.provider,
ctx.resolved_model,
ctx.elapsed_ms()
);
Ok(())
}
fn name(&self) -> &str {
"telemetry"
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_telemetry_step_record_request() {
let stats = Arc::new(RwLock::new(StatsAggregator::with_defaults()));
let tokens = Arc::new(RwLock::new(TokenTracker::with_defaults()));
let step = TelemetryStep::new(stats.clone(), tokens);
let ctx = RequestContext::new("claude-sonnet-4-5".to_string());
step.record_request(&ctx, RequestStatus::Success, None);
let stats_guard = stats.read();
assert_eq!(stats_guard.len(), 1);
}
#[test]
fn test_telemetry_step_record_tokens() {
let stats = Arc::new(RwLock::new(StatsAggregator::with_defaults()));
let tokens = Arc::new(RwLock::new(TokenTracker::with_defaults()));
let step = TelemetryStep::new(stats, tokens.clone());
let ctx = RequestContext::new("claude-sonnet-4-5".to_string());
step.record_tokens(&ctx, Some(100), Some(50), TokenSource::Actual);
let tokens_guard = tokens.read();
assert_eq!(tokens_guard.len(), 1);
}
#[test]
fn test_telemetry_step_record_tokens_from_response() {
let stats = Arc::new(RwLock::new(StatsAggregator::with_defaults()));
let tokens = Arc::new(RwLock::new(TokenTracker::with_defaults()));
let step = TelemetryStep::new(stats, tokens.clone());
let ctx = RequestContext::new("claude-sonnet-4-5".to_string());
let response = serde_json::json!({
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50
}
});
step.record_tokens_from_response(&ctx, &response);
let tokens_guard = tokens.read();
assert_eq!(tokens_guard.len(), 1);
}
#[tokio::test]
async fn test_telemetry_step_execute() {
let stats = Arc::new(RwLock::new(StatsAggregator::with_defaults()));
let tokens = Arc::new(RwLock::new(TokenTracker::with_defaults()));
let step = TelemetryStep::new(stats.clone(), tokens.clone());
let mut ctx = RequestContext::new("claude-sonnet-4-5".to_string());
let mut payload = serde_json::json!({
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50
}
});
let result = step.execute(&mut ctx, &mut payload).await;
assert!(result.is_ok());
let stats_guard = stats.read();
assert_eq!(stats_guard.len(), 1);
let tokens_guard = tokens.read();
assert_eq!(tokens_guard.len(), 1);
}
}
+90
View File
@@ -0,0 +1,90 @@
//! 管道步骤 trait 定义
//!
//! 定义所有管道步骤必须实现的接口
use crate::processor::RequestContext;
use async_trait::async_trait;
use thiserror::Error;
/// 步骤错误
#[derive(Error, Debug, Clone)]
pub enum StepError {
/// 认证错误
#[error("认证错误: {0}")]
Auth(String),
/// 路由错误
#[error("路由错误: {0}")]
Routing(String),
/// 注入错误
#[error("注入错误: {0}")]
Injection(String),
/// Provider 错误
#[error("Provider 错误: {0}")]
Provider(String),
/// 插件错误
#[error("插件错误: {plugin_name} - {message}")]
Plugin {
plugin_name: String,
message: String,
},
/// 遥测错误
#[error("遥测错误: {0}")]
Telemetry(String),
/// 超时错误
#[error("超时: {timeout_ms}ms")]
Timeout { timeout_ms: u64 },
/// 内部错误
#[error("内部错误: {0}")]
Internal(String),
}
impl StepError {
/// 获取对应的 HTTP 状态码
pub fn status_code(&self) -> u16 {
match self {
StepError::Auth(_) => 401,
StepError::Routing(_) => 404,
StepError::Injection(_) => 400,
StepError::Provider(_) => 502,
StepError::Plugin { .. } => 500,
StepError::Telemetry(_) => 500,
StepError::Timeout { .. } => 408,
StepError::Internal(_) => 500,
}
}
}
/// 管道步骤 trait
///
/// 所有管道步骤必须实现此 trait
#[async_trait]
pub trait PipelineStep: Send + Sync {
/// 执行步骤
///
/// # Arguments
/// * `ctx` - 请求上下文
/// * `payload` - 请求/响应负载
///
/// # Returns
/// 成功返回 `Ok(())`,失败返回 `Err(StepError)`
async fn execute(
&self,
ctx: &mut RequestContext,
payload: &mut serde_json::Value,
) -> Result<(), StepError>;
/// 获取步骤名称
fn name(&self) -> &str;
/// 检查步骤是否启用
fn is_enabled(&self) -> bool {
true
}
}
+930
View File
@@ -0,0 +1,930 @@
//! 处理器模块测试
use super::*;
use crate::router::RoutingRule;
use crate::services::provider_pool_service::ProviderPoolService;
use crate::ProviderType;
#[test]
fn test_request_processor_new() {
let pool_service = Arc::new(ProviderPoolService::new());
let processor = RequestProcessor::with_defaults(pool_service);
// 验证所有组件都已初始化
assert!(Arc::strong_count(&processor.router) >= 1);
assert!(Arc::strong_count(&processor.mapper) >= 1);
assert!(Arc::strong_count(&processor.injector) >= 1);
assert!(Arc::strong_count(&processor.retrier) >= 1);
assert!(Arc::strong_count(&processor.failover) >= 1);
assert!(Arc::strong_count(&processor.timeout) >= 1);
assert!(Arc::strong_count(&processor.plugins) >= 1);
assert!(Arc::strong_count(&processor.stats) >= 1);
assert!(Arc::strong_count(&processor.tokens) >= 1);
assert!(Arc::strong_count(&processor.pool_service) >= 1);
}
#[tokio::test]
async fn test_request_processor_components() {
let pool_service = Arc::new(ProviderPoolService::new());
let processor = RequestProcessor::with_defaults(pool_service);
// 验证路由器可以正常使用
{
let router = processor.router.read().await;
assert!(router.rules().is_empty());
}
// 验证映射器可以正常使用
{
let mapper = processor.mapper.read().await;
// resolve 返回原值如果没有别名
assert_eq!(mapper.resolve("unknown"), "unknown");
}
// 验证注入器可以正常使用
{
let injector = processor.injector.read().await;
assert!(injector.rules().is_empty());
}
// 验证统计聚合器可以正常使用(使用 parking_lot::RwLock)
{
let stats = processor.stats.read();
assert!(stats.is_empty());
}
// 验证 Token 追踪器可以正常使用(使用 parking_lot::RwLock)
{
let tokens = processor.tokens.read();
assert!(tokens.is_empty());
}
}
// ========== 模型映射测试 (需求 2.1) ==========
#[tokio::test]
async fn test_resolve_model_with_alias() {
let pool_service = Arc::new(ProviderPoolService::new());
let processor = RequestProcessor::with_defaults(pool_service);
// 添加别名映射
{
let mut mapper = processor.mapper.write().await;
mapper.add_alias("gpt-4", "claude-sonnet-4-5");
mapper.add_alias("gpt-3.5-turbo", "claude-3-haiku");
}
// 测试别名解析
let resolved = processor.resolve_model("gpt-4").await;
assert_eq!(resolved, "claude-sonnet-4-5");
let resolved = processor.resolve_model("gpt-3.5-turbo").await;
assert_eq!(resolved, "claude-3-haiku");
// 非别名应返回原值
let resolved = processor.resolve_model("claude-sonnet-4-5").await;
assert_eq!(resolved, "claude-sonnet-4-5");
}
#[tokio::test]
async fn test_resolve_model_for_context() {
let pool_service = Arc::new(ProviderPoolService::new());
let processor = RequestProcessor::with_defaults(pool_service);
// 添加别名映射
{
let mut mapper = processor.mapper.write().await;
mapper.add_alias("gpt-4", "claude-sonnet-4-5");
}
// 创建请求上下文
let mut ctx = RequestContext::new("gpt-4".to_string());
assert_eq!(ctx.original_model, "gpt-4");
assert_eq!(ctx.resolved_model, "gpt-4"); // 初始时相同
// 解析模型并更新上下文
let resolved = processor.resolve_model_for_context(&mut ctx).await;
assert_eq!(resolved, "claude-sonnet-4-5");
assert_eq!(ctx.original_model, "gpt-4"); // 原始模型不变
assert_eq!(ctx.resolved_model, "claude-sonnet-4-5"); // 解析后的模型已更新
}
// ========== 路由测试 (需求 2.2, 2.3) ==========
#[tokio::test]
async fn test_route_model_with_rules() {
let pool_service = Arc::new(ProviderPoolService::new());
let processor = RequestProcessor::with_defaults(pool_service);
// 添加路由规则
{
let mut router = processor.router.write().await;
router.add_rule(RoutingRule::new("gemini-*", ProviderType::Gemini, 10));
router.add_rule(RoutingRule::new("qwen-*", ProviderType::Qwen, 10));
}
// 测试路由
let (provider, is_default) = processor.route_model("gemini-2.5-flash").await;
assert_eq!(provider, ProviderType::Gemini);
assert!(!is_default);
let (provider, is_default) = processor.route_model("qwen-plus").await;
assert_eq!(provider, ProviderType::Qwen);
assert!(!is_default);
// 无匹配规则时使用默认 Provider
let (provider, is_default) = processor.route_model("claude-sonnet-4-5").await;
assert_eq!(provider, ProviderType::Kiro);
assert!(is_default);
}
#[tokio::test]
async fn test_route_for_context() {
let pool_service = Arc::new(ProviderPoolService::new());
let processor = RequestProcessor::with_defaults(pool_service);
// 添加路由规则
{
let mut router = processor.router.write().await;
router.add_rule(RoutingRule::new("gemini-*", ProviderType::Gemini, 10));
}
// 创建请求上下文
let mut ctx = RequestContext::new("gemini-2.5-flash".to_string());
ctx.set_resolved_model("gemini-2.5-flash".to_string());
// 路由并更新上下文
let provider = processor.route_for_context(&mut ctx).await;
assert_eq!(provider, ProviderType::Gemini);
assert_eq!(ctx.provider, Some(ProviderType::Gemini));
}
#[tokio::test]
async fn test_is_model_excluded() {
let pool_service = Arc::new(ProviderPoolService::new());
let processor = RequestProcessor::with_defaults(pool_service);
// 添加排除规则
{
let mut router = processor.router.write().await;
router.add_exclusion(ProviderType::Gemini, "*-preview");
router.add_exclusion(ProviderType::Gemini, "gemini-2.5-pro");
}
// 测试排除检查
assert!(
processor
.is_model_excluded(ProviderType::Gemini, "gemini-2.5-pro-preview")
.await
);
assert!(
processor
.is_model_excluded(ProviderType::Gemini, "gemini-2.5-pro")
.await
);
assert!(
!processor
.is_model_excluded(ProviderType::Gemini, "gemini-2.5-flash")
.await
);
assert!(
!processor
.is_model_excluded(ProviderType::Kiro, "gemini-2.5-pro")
.await
);
}
#[tokio::test]
async fn test_resolve_and_route() {
let pool_service = Arc::new(ProviderPoolService::new());
let processor = RequestProcessor::with_defaults(pool_service);
// 添加别名映射
{
let mut mapper = processor.mapper.write().await;
mapper.add_alias("gpt-4", "claude-sonnet-4-5");
}
// 添加路由规则
{
let mut router = processor.router.write().await;
router.add_rule(RoutingRule::new("gemini-*", ProviderType::Gemini, 10));
}
// 测试完整的解析和路由流程
let mut ctx = RequestContext::new("gpt-4".to_string());
let provider = processor.resolve_and_route(&mut ctx).await;
// gpt-4 -> claude-sonnet-4-5 -> Kiro (默认)
assert_eq!(ctx.original_model, "gpt-4");
assert_eq!(ctx.resolved_model, "claude-sonnet-4-5");
assert_eq!(provider, ProviderType::Kiro);
assert_eq!(ctx.provider, Some(ProviderType::Kiro));
// 测试 Gemini 模型
let mut ctx2 = RequestContext::new("gemini-2.5-flash".to_string());
let provider2 = processor.resolve_and_route(&mut ctx2).await;
assert_eq!(ctx2.original_model, "gemini-2.5-flash");
assert_eq!(ctx2.resolved_model, "gemini-2.5-flash");
assert_eq!(provider2, ProviderType::Gemini);
assert_eq!(ctx2.provider, Some(ProviderType::Gemini));
}
#[tokio::test]
async fn test_route_with_exclusion() {
let pool_service = Arc::new(ProviderPoolService::new());
let processor = RequestProcessor::with_defaults(pool_service);
// 添加路由规则和排除规则
{
let mut router = processor.router.write().await;
router.add_rule(RoutingRule::new("gemini-*", ProviderType::Gemini, 10));
router.add_exclusion(ProviderType::Gemini, "*-preview");
}
// 正常路由
let (provider, is_default) = processor.route_model("gemini-2.5-flash").await;
assert_eq!(provider, ProviderType::Gemini);
assert!(!is_default);
// 被排除的模型应使用默认 Provider
let (provider, is_default) = processor.route_model("gemini-2.5-pro-preview").await;
assert_eq!(provider, ProviderType::Kiro);
assert!(is_default);
}
// ========== 属性测试 (Property-Based Tests) ==========
use crate::telemetry::{RequestLog, RequestStatus};
use proptest::prelude::*;
/// 生成随机的 ProviderType
fn arb_provider_type() -> impl Strategy<Value = ProviderType> {
prop_oneof![
Just(ProviderType::Kiro),
Just(ProviderType::Gemini),
Just(ProviderType::Qwen),
Just(ProviderType::OpenAI),
Just(ProviderType::Claude),
]
}
/// 生成随机的 RequestStatus
fn arb_request_status() -> impl Strategy<Value = RequestStatus> {
prop_oneof![
Just(RequestStatus::Success),
Just(RequestStatus::Failed),
Just(RequestStatus::Timeout),
Just(RequestStatus::Cancelled),
]
}
/// 生成随机的模型名称
fn arb_model_name() -> impl Strategy<Value = String> {
prop_oneof![
Just("claude-sonnet-4".to_string()),
Just("claude-opus-4".to_string()),
Just("gemini-2.5-flash".to_string()),
Just("gemini-2.5-pro".to_string()),
Just("qwen3-coder-plus".to_string()),
Just("gpt-4o".to_string()),
]
}
/// 生成随机的请求日志
fn arb_request_log() -> impl Strategy<Value = RequestLog> {
(
"[a-zA-Z0-9_-]{8,16}", // id
arb_provider_type(),
arb_model_name(),
any::<bool>(), // is_streaming
arb_request_status(),
1u64..10000u64, // duration_ms
prop::option::of(100u16..600u16), // http_status
prop::option::of(1u32..10000u32), // input_tokens
prop::option::of(1u32..5000u32), // output_tokens
)
.prop_map(
|(
id,
provider,
model,
is_streaming,
status,
duration_ms,
http_status,
input_tokens,
output_tokens,
)| {
let mut log = RequestLog::new(id, provider, model, is_streaming);
match status {
RequestStatus::Success => {
log.mark_success(duration_ms, http_status.unwrap_or(200));
}
RequestStatus::Failed => {
log.mark_failed(duration_ms, http_status, "Test error".to_string());
}
RequestStatus::Timeout => {
log.mark_timeout(duration_ms);
}
RequestStatus::Cancelled => {
log.mark_cancelled(duration_ms);
}
RequestStatus::Retrying => {
// 保持默认状态
}
}
log.set_tokens(input_tokens, output_tokens);
log
},
)
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
/// **Feature: module-integration, Property 1: 请求完成后统计记录**
/// *对于任意* 成功或失败的请求,请求完成后 StatsAggregator 中应存在对应的记录
/// **Validates: Requirements 1.3, 4.1**
#[test]
fn prop_request_stats_recorded(
log in arb_request_log()
) {
let pool_service = Arc::new(ProviderPoolService::new());
let processor = RequestProcessor::with_defaults(pool_service);
let original_id = log.id.clone();
let original_provider = log.provider;
let original_model = log.model.clone();
let original_status = log.status;
// 记录请求日志到 StatsAggregator(使用 parking_lot::RwLock)
{
let stats = processor.stats.write();
stats.record(log);
}
// 验证:StatsAggregator 中应存在对应的记录
{
let stats = processor.stats.read();
let all_logs = stats.get_all();
// 查找对应的记录
let found = all_logs.iter().find(|l| l.id == original_id);
prop_assert!(
found.is_some(),
"请求 {} 完成后应在 StatsAggregator 中存在记录",
original_id
);
// 验证记录内容一致性
let found_log = found.unwrap();
prop_assert_eq!(
found_log.provider,
original_provider,
"记录的 Provider 应与原始一致"
);
prop_assert_eq!(
&found_log.model,
&original_model,
"记录的模型应与原始一致"
);
prop_assert_eq!(
found_log.status,
original_status,
"记录的状态应与原始一致"
);
}
}
/// **Feature: module-integration, Property 1: 请求完成后统计记录(批量)**
/// *对于任意* 多个请求,所有请求完成后 StatsAggregator 中应存在所有对应的记录
/// **Validates: Requirements 1.3, 4.1**
#[test]
fn prop_multiple_requests_stats_recorded(
logs in prop::collection::vec(arb_request_log(), 1..20)
) {
let pool_service = Arc::new(ProviderPoolService::new());
let processor = RequestProcessor::with_defaults(pool_service);
// 收集所有日志 ID
let log_ids: Vec<String> = logs.iter().map(|l| l.id.clone()).collect();
let expected_count = logs.len();
// 记录所有请求日志(使用 parking_lot::RwLock)
{
let stats = processor.stats.write();
for log in logs {
stats.record(log);
}
}
// 验证:所有记录都应存在
{
let stats = processor.stats.read();
let all_logs = stats.get_all();
// 验证记录数量
prop_assert_eq!(
all_logs.len(),
expected_count,
"StatsAggregator 中的记录数应等于请求数"
);
// 验证每个请求都有对应记录
for id in &log_ids {
let found = all_logs.iter().any(|l| &l.id == id);
prop_assert!(
found,
"请求 {} 应在 StatsAggregator 中存在记录",
id
);
}
}
}
/// **Feature: module-integration, Property 1: 请求完成后统计记录(统计准确性)**
/// *对于任意* 请求集合,统计摘要中的总请求数应等于记录的请求数
/// **Validates: Requirements 1.3, 4.1**
#[test]
fn prop_stats_summary_accuracy(
logs in prop::collection::vec(arb_request_log(), 1..50)
) {
let pool_service = Arc::new(ProviderPoolService::new());
let processor = RequestProcessor::with_defaults(pool_service);
let expected_total = logs.len();
let expected_success = logs.iter().filter(|l| l.status == RequestStatus::Success).count();
let expected_failed = logs.iter().filter(|l| l.status == RequestStatus::Failed).count();
let expected_timeout = logs.iter().filter(|l| l.status == RequestStatus::Timeout).count();
// 记录所有请求日志(使用 parking_lot::RwLock)
{
let stats = processor.stats.write();
for log in logs {
stats.record(log);
}
}
// 获取统计摘要
{
let stats = processor.stats.read();
let summary = stats.summary(None);
// 验证总请求数
prop_assert_eq!(
summary.total_requests as usize,
expected_total,
"统计的总请求数应等于记录的请求数"
);
// 验证成功请求数
prop_assert_eq!(
summary.successful_requests as usize,
expected_success,
"统计的成功请求数应正确"
);
// 验证失败请求数
prop_assert_eq!(
summary.failed_requests as usize,
expected_failed,
"统计的失败请求数应正确"
);
// 验证超时请求数
prop_assert_eq!(
summary.timeout_requests as usize,
expected_timeout,
"统计的超时请求数应正确"
);
// 验证成功率
if expected_total > 0 {
let expected_rate = expected_success as f64 / expected_total as f64;
prop_assert!(
(summary.success_rate - expected_rate).abs() < 0.001,
"成功率应正确计算"
);
}
}
}
}
// ========== 凭证失败计数属性测试 ==========
use crate::credential::{Credential, CredentialData, CredentialPool, HealthChecker};
/// 生成随机的凭证 ID
fn arb_credential_id() -> impl Strategy<Value = String> {
"[a-z0-9]{8,16}".prop_map(|s| s.to_string())
}
/// 生成随机的失败次数
fn arb_failure_count() -> impl Strategy<Value = u32> {
1u32..10u32
}
/// 创建测试凭证
fn create_test_credential(id: &str) -> Credential {
Credential::new(
id.to_string(),
ProviderType::Kiro,
CredentialData::ApiKey {
key: format!("test-key-{}", id),
base_url: None,
},
)
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
/// **Feature: module-integration, Property 3: 凭证失败计数更新**
/// *对于任意* 凭证调用失败,HealthChecker 中该凭证的失败计数应增加 1
/// **Validates: Requirements 8.2**
#[test]
fn prop_credential_failure_count_update(
credential_id in arb_credential_id(),
failure_count in arb_failure_count()
) {
// 创建凭证池和健康检查器
let pool = CredentialPool::new(ProviderType::Kiro);
let health_checker = HealthChecker::with_defaults();
// 添加凭证到池中
let credential = create_test_credential(&credential_id);
pool.add(credential).unwrap();
// 获取初始失败计数
let initial_failures = pool.get(&credential_id)
.map(|c| c.stats.consecutive_failures)
.unwrap_or(0);
// 记录多次失败
for i in 0..failure_count {
let _ = health_checker.record_failure(&pool, &credential_id);
// 验证每次失败后计数增加 1
let current_failures = pool.get(&credential_id)
.map(|c| c.stats.consecutive_failures)
.unwrap_or(0);
prop_assert_eq!(
current_failures,
initial_failures + i + 1,
"第 {} 次失败后,失败计数应为 {}",
i + 1,
initial_failures + i + 1
);
}
// 验证最终失败计数
let final_failures = pool.get(&credential_id)
.map(|c| c.stats.consecutive_failures)
.unwrap_or(0);
prop_assert_eq!(
final_failures,
initial_failures + failure_count,
"最终失败计数应等于初始计数 + 失败次数"
);
}
/// **Feature: module-integration, Property 3: 凭证失败计数更新(多凭证)**
/// *对于任意* 多个凭证的失败,每个凭证的失败计数应独立更新
/// **Validates: Requirements 8.2**
#[test]
fn prop_credential_failure_count_independent(
credential_ids in prop::collection::vec(arb_credential_id(), 2..5),
failure_counts in prop::collection::vec(arb_failure_count(), 2..5)
) {
// 确保凭证 ID 唯一
let unique_ids: Vec<String> = credential_ids.into_iter()
.enumerate()
.map(|(i, id)| format!("{}-{}", id, i))
.collect();
// 创建凭证池和健康检查器
let pool = CredentialPool::new(ProviderType::Kiro);
let health_checker = HealthChecker::with_defaults();
// 添加所有凭证
for id in &unique_ids {
let credential = create_test_credential(id);
pool.add(credential).unwrap();
}
// 为每个凭证记录不同次数的失败
let min_len = unique_ids.len().min(failure_counts.len());
for i in 0..min_len {
let id = &unique_ids[i];
let count = failure_counts[i];
for _ in 0..count {
let _ = health_checker.record_failure(&pool, id);
}
}
// 验证每个凭证的失败计数独立
for i in 0..min_len {
let id = &unique_ids[i];
let expected_count = failure_counts[i];
let actual_count = pool.get(id)
.map(|c| c.stats.consecutive_failures)
.unwrap_or(0);
prop_assert_eq!(
actual_count,
expected_count,
"凭证 {} 的失败计数应为 {}",
id,
expected_count
);
}
}
/// **Feature: module-integration, Property 3: 凭证失败计数更新(成功重置)**
/// *对于任意* 凭证,成功调用后连续失败计数应重置为 0
/// **Validates: Requirements 8.2**
#[test]
fn prop_credential_failure_count_reset_on_success(
credential_id in arb_credential_id(),
failure_count in arb_failure_count(),
latency_ms in 10u64..1000u64
) {
// 创建凭证池和健康检查器
let pool = CredentialPool::new(ProviderType::Kiro);
let health_checker = HealthChecker::with_defaults();
// 添加凭证
let credential = create_test_credential(&credential_id);
pool.add(credential).unwrap();
// 记录多次失败
for _ in 0..failure_count {
let _ = health_checker.record_failure(&pool, &credential_id);
}
// 验证失败计数已增加
let failures_before_success = pool.get(&credential_id)
.map(|c| c.stats.consecutive_failures)
.unwrap_or(0);
prop_assert_eq!(
failures_before_success,
failure_count,
"成功前失败计数应为 {}", failure_count
);
// 记录成功
let _ = health_checker.record_success(&pool, &credential_id, latency_ms);
// 验证失败计数已重置
let failures_after_success = pool.get(&credential_id)
.map(|c| c.stats.consecutive_failures)
.unwrap_or(0);
prop_assert_eq!(
failures_after_success,
0,
"成功后连续失败计数应重置为 0"
);
}
}
// ========== Token 响应记录一致性属性测试 ==========
use crate::processor::steps::TelemetryStep;
use crate::telemetry::{TokenSource, TokenTracker, TokenUsageRecord};
use parking_lot::RwLock as ParkingLotRwLock;
/// 生成随机的 Token 数量
fn arb_token_count() -> impl Strategy<Value = u32> {
1u32..10000u32
}
/// 生成随机的 OpenAI 格式响应
fn arb_openai_response() -> impl Strategy<Value = (u32, u32, serde_json::Value)> {
(arb_token_count(), arb_token_count()).prop_map(|(input, output)| {
let response = serde_json::json!({
"id": "chatcmpl-test",
"object": "chat.completion",
"created": 1234567890,
"model": "claude-sonnet-4-5",
"choices": [{
"index": 0,
"message": {
"role": "assistant",
"content": "Test response"
},
"finish_reason": "stop"
}],
"usage": {
"prompt_tokens": input,
"completion_tokens": output,
"total_tokens": input + output
}
});
(input, output, response)
})
}
/// 生成随机的 Anthropic 格式响应
fn arb_anthropic_response() -> impl Strategy<Value = (u32, u32, serde_json::Value)> {
(arb_token_count(), arb_token_count()).prop_map(|(input, output)| {
let response = serde_json::json!({
"id": "msg-test",
"type": "message",
"role": "assistant",
"content": [{
"type": "text",
"text": "Test response"
}],
"model": "claude-sonnet-4-5",
"stop_reason": "end_turn",
"usage": {
"input_tokens": input,
"output_tokens": output
}
});
(input, output, response)
})
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
/// **Feature: module-integration, Property 2: Token 响应记录一致性**
/// *对于任意* 包含 Token 使用信息的响应,TokenTracker 中记录的 Token 数应与响应中的值一致
/// **Validates: Requirements 4.2**
#[test]
fn prop_token_response_recording_consistency_openai(
(expected_input, expected_output, response) in arb_openai_response()
) {
// 创建 TelemetryStep 和 TokenTracker
let stats = Arc::new(ParkingLotRwLock::new(crate::telemetry::StatsAggregator::with_defaults()));
let tokens = Arc::new(ParkingLotRwLock::new(TokenTracker::with_defaults()));
let step = TelemetryStep::new(stats, tokens.clone());
// 创建请求上下文
let ctx = RequestContext::new("claude-sonnet-4-5".to_string());
// 从响应中提取并记录 Token
step.record_tokens_from_response(&ctx, &response);
// 验证:TokenTracker 中记录的 Token 数应与响应中的值一致
let tokens_guard = tokens.read();
let all_records = tokens_guard.get_all();
prop_assert_eq!(
all_records.len(),
1,
"应该有且仅有一条 Token 记录"
);
let record = &all_records[0];
prop_assert_eq!(
record.input_tokens,
expected_input,
"记录的输入 Token 数应与响应中的值一致"
);
prop_assert_eq!(
record.output_tokens,
expected_output,
"记录的输出 Token 数应与响应中的值一致"
);
prop_assert_eq!(
record.total_tokens,
expected_input + expected_output,
"记录的总 Token 数应等于输入 + 输出"
);
prop_assert_eq!(
record.source,
TokenSource::Actual,
"Token 来源应为 Actual"
);
prop_assert_eq!(
record.request_id.as_ref(),
Some(&ctx.request_id),
"记录的请求 ID 应与上下文一致"
);
}
/// **Feature: module-integration, Property 2: Token 响应记录一致性(Anthropic 格式)**
/// *对于任意* 包含 Token 使用信息的 Anthropic 格式响应,TokenTracker 中记录的 Token 数应与响应中的值一致
/// **Validates: Requirements 4.2**
#[test]
fn prop_token_response_recording_consistency_anthropic(
(expected_input, expected_output, response) in arb_anthropic_response()
) {
// 创建 TelemetryStep 和 TokenTracker
let stats = Arc::new(ParkingLotRwLock::new(crate::telemetry::StatsAggregator::with_defaults()));
let tokens = Arc::new(ParkingLotRwLock::new(TokenTracker::with_defaults()));
let step = TelemetryStep::new(stats, tokens.clone());
// 创建请求上下文
let ctx = RequestContext::new("claude-sonnet-4-5".to_string());
// 从响应中提取并记录 Token
step.record_tokens_from_response(&ctx, &response);
// 验证:TokenTracker 中记录的 Token 数应与响应中的值一致
let tokens_guard = tokens.read();
let all_records = tokens_guard.get_all();
prop_assert_eq!(
all_records.len(),
1,
"应该有且仅有一条 Token 记录"
);
let record = &all_records[0];
prop_assert_eq!(
record.input_tokens,
expected_input,
"记录的输入 Token 数应与响应中的值一致"
);
prop_assert_eq!(
record.output_tokens,
expected_output,
"记录的输出 Token 数应与响应中的值一致"
);
prop_assert_eq!(
record.total_tokens,
expected_input + expected_output,
"记录的总 Token 数应等于输入 + 输出"
);
prop_assert_eq!(
record.source,
TokenSource::Actual,
"Token 来源应为 Actual"
);
}
/// **Feature: module-integration, Property 2: Token 响应记录一致性(批量)**
/// *对于任意* 多个包含 Token 使用信息的响应,TokenTracker 中应记录所有响应的 Token 数
/// **Validates: Requirements 4.2**
#[test]
fn prop_token_response_recording_consistency_batch(
responses in prop::collection::vec(arb_openai_response(), 1..20)
) {
// 创建 TelemetryStep 和 TokenTracker
let stats = Arc::new(ParkingLotRwLock::new(crate::telemetry::StatsAggregator::with_defaults()));
let tokens = Arc::new(ParkingLotRwLock::new(TokenTracker::with_defaults()));
let step = TelemetryStep::new(stats, tokens.clone());
let expected_count = responses.len();
let mut expected_total_input: u64 = 0;
let mut expected_total_output: u64 = 0;
// 记录所有响应的 Token
for (input, output, response) in &responses {
let ctx = RequestContext::new("claude-sonnet-4-5".to_string());
step.record_tokens_from_response(&ctx, response);
expected_total_input += *input as u64;
expected_total_output += *output as u64;
}
// 验证:TokenTracker 中应记录所有响应的 Token 数
let tokens_guard = tokens.read();
let all_records = tokens_guard.get_all();
prop_assert_eq!(
all_records.len(),
expected_count,
"Token 记录数应等于响应数"
);
// 验证总 Token 数
let actual_total_input: u64 = all_records.iter().map(|r| r.input_tokens as u64).sum();
let actual_total_output: u64 = all_records.iter().map(|r| r.output_tokens as u64).sum();
prop_assert_eq!(
actual_total_input,
expected_total_input,
"总输入 Token 数应一致"
);
prop_assert_eq!(
actual_total_output,
expected_total_output,
"总输出 Token 数应一致"
);
}
}
+10
View File
@@ -168,6 +168,16 @@ impl Router {
self.default_provider
}
/// 清空所有路由规则
pub fn clear_rules(&mut self) {
self.rules.clear();
}
/// 清空所有排除规则
pub fn clear_exclusions(&mut self) {
self.exclusions.clear();
}
/// 添加排除模式
pub fn add_exclusion(&mut self, provider: ProviderType, pattern: &str) {
self.exclusions
+1736 -74
View File
File diff suppressed because it is too large Load Diff
+166
View File
@@ -460,3 +460,169 @@ proptest! {
}
}
}
// ============ SSE 到 WebSocket 转换属性测试 ============
use super::stream::StreamForwarder;
/// 生成任意的 SSE 数据内容(非空且不全是空格)
fn arb_sse_data() -> impl Strategy<Value = String> {
prop_oneof![
// JSON 对象
Just(r#"{"content": "hello"}"#.to_string()),
Just(r#"{"delta": {"content": "world"}}"#.to_string()),
Just(r#"{"choices": [{"delta": {"content": "test"}}]}"#.to_string()),
// 简单文本(确保至少有一个非空格字符)
"[a-zA-Z0-9]{1,50}".prop_map(|s| s),
]
}
/// 生成任意的 SSE 行
fn arb_sse_line() -> impl Strategy<Value = String> {
arb_sse_data().prop_map(|data| format!("data: {}", data))
}
/// 生成任意的 SSE 响应体(多行)
fn arb_sse_body() -> impl Strategy<Value = (Vec<String>, String)> {
prop::collection::vec(arb_sse_data(), 1..10).prop_map(|data_items| {
let lines: Vec<String> = data_items.clone();
let body = data_items
.iter()
.map(|d| format!("data: {}\n\n", d))
.collect::<Vec<_>>()
.join("");
(lines, body)
})
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(100))]
/// **Feature: module-integration, Property 4: SSE 到 WebSocket 转换完整性**
/// *对于任意* SSE 流式响应,转换为 WebSocket 消息后数据内容应保持不变
/// **Validates: Requirements 6.3**
#[test]
fn prop_sse_to_ws_conversion_integrity(
request_id in "[a-zA-Z0-9-]{8,16}",
(original_data, sse_body) in arb_sse_body()
) {
let forwarder = StreamForwarder::new(request_id.clone());
// 转换 SSE 响应体为 WebSocket 消息
let messages = forwarder.process_sse_body(&sse_body);
// 验证消息数量:应该有 N 个数据块 + 1 个结束消息
let expected_chunks = original_data.len();
prop_assert_eq!(
messages.len(),
expected_chunks + 1,
"消息数量应为数据块数 + 1 (结束消息)"
);
// 验证每个数据块的内容
for (i, msg) in messages.iter().take(expected_chunks).enumerate() {
match msg {
WsMessage::StreamChunk(chunk) => {
// 验证 request_id
prop_assert_eq!(
&chunk.request_id,
&request_id,
"StreamChunk 的 request_id 应与原始一致"
);
// 验证索引
prop_assert_eq!(
chunk.index as usize,
i,
"StreamChunk 的索引应正确"
);
// 验证数据内容
prop_assert_eq!(
&chunk.data,
&original_data[i],
"StreamChunk 的数据应与原始 SSE 数据一致"
);
}
_ => prop_assert!(false, "期望 StreamChunk 消息,但得到其他类型"),
}
}
// 验证结束消息
match messages.last() {
Some(WsMessage::StreamEnd(end)) => {
prop_assert_eq!(
&end.request_id,
&request_id,
"StreamEnd 的 request_id 应与原始一致"
);
prop_assert_eq!(
end.total_chunks as usize,
expected_chunks,
"StreamEnd 的 total_chunks 应等于数据块数"
);
}
_ => prop_assert!(false, "最后一条消息应为 StreamEnd"),
}
}
/// **Feature: module-integration, Property 4: SSE 到 WebSocket 转换完整性(单行)**
/// *对于任意* 单个 SSE 数据行,转换后应保持数据内容不变
/// **Validates: Requirements 6.3**
#[test]
fn prop_sse_line_conversion_integrity(
request_id in "[a-zA-Z0-9-]{8,16}",
data in arb_sse_data(),
index in 0u32..1000u32
) {
let forwarder = StreamForwarder::new(request_id.clone());
let sse_line = format!("data: {}", data);
// 转换单行 SSE
let result = forwarder.convert_sse_line(&sse_line, index);
// 验证转换结果
match result {
Some(WsMessage::StreamChunk(chunk)) => {
prop_assert_eq!(
&chunk.request_id,
&request_id,
"request_id 应保持不变"
);
prop_assert_eq!(
chunk.index,
index,
"index 应保持不变"
);
prop_assert_eq!(
&chunk.data,
&data,
"数据内容应保持不变"
);
}
_ => prop_assert!(false, "应返回 StreamChunk 消息"),
}
}
/// **Feature: module-integration, Property 4: SSE 到 WebSocket 转换完整性(特殊行处理)**
/// *对于任意* 空行、注释行或 [DONE] 标记,转换应返回 None
/// **Validates: Requirements 6.3**
#[test]
fn prop_sse_special_lines_filtered(
request_id in "[a-zA-Z0-9-]{8,16}",
index in 0u32..1000u32
) {
let forwarder = StreamForwarder::new(request_id);
// 空行应返回 None
prop_assert!(forwarder.convert_sse_line("", index).is_none());
prop_assert!(forwarder.convert_sse_line(" ", index).is_none());
// 注释行应返回 None
prop_assert!(forwarder.convert_sse_line(": comment", index).is_none());
prop_assert!(forwarder.convert_sse_line(":ping", index).is_none());
// [DONE] 标记应返回 None
prop_assert!(forwarder.convert_sse_line("data: [DONE]", index).is_none());
}
}
+1 -1
View File
@@ -1,7 +1,7 @@
{
"$schema": "https://schema.tauri.app/config/2",
"productName": "ProxyCast",
"version": "0.9.0",
"version": "0.10.0",
"identifier": "com.proxycast.app",
"build": {
"beforeDevCommand": "npm run dev",
+209
View File
@@ -0,0 +1,209 @@
import { useState, useEffect } from "react";
import {
Folder,
RotateCcw,
Check,
AlertCircle,
FolderOpen,
} from "lucide-react";
import { Config } from "@/lib/api/config";
import { invoke } from "@tauri-apps/api/core";
interface AuthDirSettingsProps {
config: Config | null;
onConfigChange: (config: Config) => void;
}
const DEFAULT_AUTH_DIR = "~/.proxycast/auth";
export function AuthDirSettings({
config,
onConfigChange,
}: AuthDirSettingsProps) {
const [authDir, setAuthDir] = useState(DEFAULT_AUTH_DIR);
const [isSaving, setIsSaving] = useState(false);
const [saveSuccess, setSaveSuccess] = useState(false);
const [error, setError] = useState<string | null>(null);
const [expandedPath, setExpandedPath] = useState<string | null>(null);
// Load auth_dir from config
useEffect(() => {
if (config?.auth_dir) {
setAuthDir(config.auth_dir);
}
}, [config]);
// Validate and expand path
const validatePath = async (path: string) => {
try {
const expanded = await invoke<string>("expand_path", { path });
setExpandedPath(expanded);
setError(null);
return true;
} catch (err) {
setError(`路径无效: ${err}`);
setExpandedPath(null);
return false;
}
};
// Handle path change
const handlePathChange = (newPath: string) => {
setAuthDir(newPath);
setSaveSuccess(false);
// Debounce validation
const timer = setTimeout(() => {
validatePath(newPath);
}, 300);
return () => clearTimeout(timer);
};
// Reset to default
const handleReset = () => {
setAuthDir(DEFAULT_AUTH_DIR);
setSaveSuccess(false);
validatePath(DEFAULT_AUTH_DIR);
};
// Save changes
const handleSave = async () => {
if (!config) return;
setIsSaving(true);
setError(null);
setSaveSuccess(false);
try {
// Validate path first
const isValid = await validatePath(authDir);
if (!isValid) {
setIsSaving(false);
return;
}
// Update config
const newConfig = {
...config,
auth_dir: authDir,
};
onConfigChange(newConfig);
setSaveSuccess(true);
// Clear success message after 3 seconds
setTimeout(() => setSaveSuccess(false), 3000);
} catch (err) {
setError(`保存失败: ${err}`);
} finally {
setIsSaving(false);
}
};
// Open folder in file manager
const handleOpenFolder = async () => {
try {
await invoke("open_auth_dir", { path: authDir });
} catch (err) {
setError(`打开文件夹失败: ${err}`);
}
};
const hasChanges = config?.auth_dir !== authDir;
return (
<div className="space-y-6">
<div>
<h3 className="text-lg font-medium flex items-center gap-2">
<Folder className="h-5 w-5" />
认证目录设置
</h3>
<p className="text-sm text-muted-foreground mt-1">
配置 OAuth Token 文件的存储目录。支持使用 ~ 表示用户主目录。
</p>
</div>
<div className="rounded-lg border p-4 space-y-4">
<div className="space-y-2">
<label className="text-sm font-medium">认证目录路径 (auth_dir)</label>
<div className="flex gap-2">
<div className="relative flex-1">
<Folder className="absolute left-3 top-1/2 -translate-y-1/2 h-4 w-4 text-muted-foreground" />
<input
type="text"
value={authDir}
onChange={(e) => handlePathChange(e.target.value)}
className="w-full pl-9 pr-3 py-2 rounded-lg border bg-background text-sm font-mono focus:ring-2 focus:ring-primary/20 focus:border-primary outline-none"
placeholder={DEFAULT_AUTH_DIR}
/>
</div>
<button
onClick={handleReset}
className="p-2 rounded-lg border hover:bg-muted text-muted-foreground"
title="重置为默认"
>
<RotateCcw className="h-4 w-4" />
</button>
<button
onClick={handleOpenFolder}
className="p-2 rounded-lg border hover:bg-muted text-muted-foreground"
title="打开文件夹"
>
<FolderOpen className="h-4 w-4" />
</button>
</div>
</div>
{/* Expanded path preview */}
{expandedPath && (
<div className="text-sm">
<span className="text-muted-foreground">展开后路径: </span>
<code className="rounded bg-muted px-2 py-0.5 text-xs">
{expandedPath}
</code>
</div>
)}
{/* Error display */}
{error && (
<div className="flex items-start gap-2 rounded-lg border border-red-200 bg-red-50 p-3 text-red-700 dark:border-red-800 dark:bg-red-950 dark:text-red-400">
<AlertCircle className="h-5 w-5 flex-shrink-0" />
<span className="text-sm">{error}</span>
</div>
)}
{/* Success message */}
{saveSuccess && (
<div className="flex items-center gap-2 rounded-lg border border-green-200 bg-green-50 p-3 text-green-700 dark:border-green-800 dark:bg-green-950 dark:text-green-400">
<Check className="h-5 w-5" />
<span className="text-sm">设置已保存</span>
</div>
)}
{/* Save button */}
<div className="flex justify-end">
<button
onClick={handleSave}
disabled={!hasChanges || isSaving}
className="flex items-center gap-2 rounded-lg bg-primary px-4 py-2 text-sm text-primary-foreground hover:bg-primary/90 disabled:opacity-50"
>
{isSaving ? "保存中..." : "保存设置"}
</button>
</div>
</div>
{/* Help text */}
<div className="rounded-lg border bg-muted/50 p-4 space-y-2">
<h4 className="text-sm font-medium">说明</h4>
<ul className="text-sm text-muted-foreground space-y-1 list-disc list-inside">
<li>认证目录用于存储 OAuth Token 文件(Kiro、Gemini、Qwen 等)</li>
<li>
使用 <code className="rounded bg-muted px-1">~</code>{" "}
表示用户主目录,例如{" "}
<code className="rounded bg-muted px-1">~/.proxycast/auth</code>
</li>
<li>修改此设置后,现有的 Token 文件不会自动迁移,需要手动移动</li>
<li>导出配置时,Token 文件会从此目录读取并包含在导出包中</li>
</ul>
</div>
</div>
);
}
+19 -5
View File
@@ -1,7 +1,13 @@
import { useState, useEffect, forwardRef, useImperativeHandle } from "react";
import { FileCode, RefreshCw, FolderOpen } from "lucide-react";
import React, {
useState,
useEffect,
forwardRef,
useImperativeHandle,
} from "react";
import { FileCode, RefreshCw, FolderOpen, Settings } from "lucide-react";
import { ConfigEditor } from "./ConfigEditor";
import { ImportExport } from "./ImportExport";
import { AuthDirSettings } from "./AuthDirSettings";
import { Config, configApi, ConfigPathInfo } from "@/lib/api/config";
import { invoke } from "@tauri-apps/api/core";
@@ -9,7 +15,7 @@ export interface ConfigPageRef {
refresh: () => void;
}
type TabType = "editor" | "import-export";
type TabType = "editor" | "import-export" | "settings";
export const ConfigPage = forwardRef<ConfigPageRef>((_props, ref) => {
const [activeTab, setActiveTab] = useState<TabType>("editor");
@@ -62,9 +68,10 @@ export const ConfigPage = forwardRef<ConfigPageRef>((_props, ref) => {
}
};
const tabs: { id: TabType; label: string }[] = [
const tabs: { id: TabType; label: string; icon?: React.ReactNode }[] = [
{ id: "editor", label: "YAML 编辑器" },
{ id: "import-export", label: "导入/导出" },
{ id: "settings", label: "设置", icon: <Settings className="h-4 w-4" /> },
];
if (isLoading) {
@@ -135,12 +142,13 @@ export const ConfigPage = forwardRef<ConfigPageRef>((_props, ref) => {
<button
key={tab.id}
onClick={() => setActiveTab(tab.id)}
className={`px-4 py-2 text-sm font-medium border-b-2 -mb-px ${
className={`flex items-center gap-1.5 px-4 py-2 text-sm font-medium border-b-2 -mb-px ${
activeTab === tab.id
? "border-primary text-primary"
: "border-transparent text-muted-foreground hover:text-foreground"
}`}
>
{tab.icon}
{tab.label}
</button>
))}
@@ -154,6 +162,12 @@ export const ConfigPage = forwardRef<ConfigPageRef>((_props, ref) => {
{activeTab === "import-export" && (
<ImportExport config={config} onConfigImported={handleConfigChange} />
)}
{activeTab === "settings" && (
<AuthDirSettings
config={config}
onConfigChange={handleConfigChange}
/>
)}
</div>
</div>
);
+337 -85
View File
@@ -1,5 +1,15 @@
import React, { useState, useRef } from "react";
import { Download, Upload, AlertCircle, Check, Shield } from "lucide-react";
import {
Download,
Upload,
AlertCircle,
Check,
Shield,
AlertTriangle,
FileJson,
FileText,
Package,
} from "lucide-react";
import { Config, configApi, ImportResult } from "@/lib/api/config";
interface ImportExportProps {
@@ -7,37 +17,78 @@ interface ImportExportProps {
onConfigImported: (config: Config) => void;
}
// Export scope options
type ExportScope = "config" | "credentials" | "full";
// Validation result from backend
interface ValidationResult {
valid: boolean;
version: string | null;
redacted: boolean;
has_config: boolean;
has_credentials: boolean;
errors: string[];
warnings: string[];
}
export function ImportExport({ config, onConfigImported }: ImportExportProps) {
// Export state
const [isExporting, setIsExporting] = useState(false);
const [exportScope, setExportScope] = useState<ExportScope>("config");
const [redactSecrets, setRedactSecrets] = useState(true);
const [showSecurityWarning, setShowSecurityWarning] = useState(false);
// Import state
const [isImporting, setIsImporting] = useState(false);
const [showImportDialog, setShowImportDialog] = useState(false);
const [importContent, setImportContent] = useState("");
const [importFileName, setImportFileName] = useState("");
const [importResult, setImportResult] = useState<ImportResult | null>(null);
const [error, setError] = useState<string | null>(null);
const [redactSecrets, setRedactSecrets] = useState(true);
const [validationResult, setValidationResult] =
useState<ValidationResult | null>(null);
const [mergeConfig, setMergeConfig] = useState(true);
const [error, setError] = useState<string | null>(null);
const fileInputRef = useRef<HTMLInputElement>(null);
// Export config (download directly)
const handleExport = async () => {
// Handle export with security check
const handleExportClick = () => {
if (
!redactSecrets &&
(exportScope === "credentials" || exportScope === "full")
) {
setShowSecurityWarning(true);
} else {
performExport();
}
};
// Perform the actual export
const performExport = async () => {
if (!config) return;
setIsExporting(true);
setError(null);
setShowSecurityWarning(false);
try {
const result = await configApi.exportConfig(config, redactSecrets);
// Create download link
const blob = new Blob([result.content], { type: "text/yaml" });
const url = URL.createObjectURL(blob);
const a = document.createElement("a");
a.href = url;
a.download = result.suggested_filename;
document.body.appendChild(a);
a.click();
document.body.removeChild(a);
URL.revokeObjectURL(url);
if (exportScope === "config") {
// Export config only as YAML
const result = await configApi.exportConfig(config, redactSecrets);
downloadFile(result.content, result.suggested_filename, "text/yaml");
} else {
// Export bundle (credentials or full)
const result = await configApi.exportBundle(config, {
include_config: exportScope === "full",
include_credentials: true,
redact_secrets: redactSecrets,
});
downloadFile(
result.content,
result.suggested_filename,
"application/json",
);
}
} catch (err) {
setError(`导出失败: ${err}`);
} finally {
@@ -45,18 +96,45 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) {
}
};
// Download file helper
const downloadFile = (
content: string,
filename: string,
mimeType: string,
) => {
const blob = new Blob([content], { type: mimeType });
const url = URL.createObjectURL(blob);
const a = document.createElement("a");
a.href = url;
a.download = filename;
document.body.appendChild(a);
a.click();
document.body.removeChild(a);
URL.revokeObjectURL(url);
};
// Handle file selection
const handleFileSelect = (e: React.ChangeEvent<HTMLInputElement>) => {
const handleFileSelect = async (e: React.ChangeEvent<HTMLInputElement>) => {
const file = e.target.files?.[0];
if (!file) return;
const reader = new FileReader();
reader.onload = (event) => {
reader.onload = async (event) => {
const content = event.target?.result as string;
setImportContent(content);
setImportFileName(file.name);
setShowImportDialog(true);
setImportResult(null);
setError(null);
// Validate the import content
try {
const validation = await configApi.validateImport(content);
setValidationResult(validation);
} catch (err) {
setError(`验证失败: ${err}`);
setValidationResult(null);
}
};
reader.onerror = () => {
setError("读取文件失败");
@@ -67,7 +145,7 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) {
e.target.value = "";
};
// Import config
// Import config/bundle
const handleImport = async () => {
if (!config || !importContent) return;
@@ -75,7 +153,7 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) {
setError(null);
try {
const result = await configApi.importConfig(
const result = await configApi.importBundle(
config,
importContent,
mergeConfig,
@@ -92,17 +170,22 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) {
}
};
// Validate before import
const handleValidate = async () => {
if (!importContent) return;
// Close import dialog
const closeImportDialog = () => {
setShowImportDialog(false);
setImportContent("");
setImportFileName("");
setImportResult(null);
setValidationResult(null);
setError(null);
try {
await configApi.validateConfigYaml(importContent);
setError(null);
} catch (err) {
setError(`验证失败: ${err}`);
};
// Get file type icon
const getFileTypeIcon = () => {
if (importFileName.endsWith(".json")) {
return <FileJson className="h-4 w-4" />;
}
return <FileText className="h-4 w-4" />;
};
return (
@@ -115,6 +198,47 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) {
</h3>
<div className="space-y-4">
{/* Export Scope Selection */}
<div className="space-y-2">
<label className="text-sm font-medium">导出范围</label>
<div className="flex flex-col gap-2">
<label className="flex items-center gap-2">
<input
type="radio"
name="exportScope"
checked={exportScope === "config"}
onChange={() => setExportScope("config")}
className="rounded-full border-gray-300"
/>
<FileText className="h-4 w-4 text-muted-foreground" />
<span className="text-sm">仅配置 (YAML)</span>
</label>
<label className="flex items-center gap-2">
<input
type="radio"
name="exportScope"
checked={exportScope === "credentials"}
onChange={() => setExportScope("credentials")}
className="rounded-full border-gray-300"
/>
<Shield className="h-4 w-4 text-muted-foreground" />
<span className="text-sm">仅凭证 (JSON)</span>
</label>
<label className="flex items-center gap-2">
<input
type="radio"
name="exportScope"
checked={exportScope === "full"}
onChange={() => setExportScope("full")}
className="rounded-full border-gray-300"
/>
<Package className="h-4 w-4 text-muted-foreground" />
<span className="text-sm">完整导出 (配置 + 凭证)</span>
</label>
</div>
</div>
{/* Redaction Option */}
<label className="flex items-center gap-2">
<input
type="checkbox"
@@ -123,22 +247,38 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) {
className="rounded border-gray-300"
/>
<Shield className="h-4 w-4 text-muted-foreground" />
<span className="text-sm">脱敏敏感信息(API 密钥等)</span>
<span className="text-sm">脱敏敏感信息(API 密钥、Token 等)</span>
</label>
{/* Security hint */}
{!redactSecrets &&
(exportScope === "credentials" || exportScope === "full") && (
<div className="flex items-start gap-2 rounded-lg border border-yellow-200 bg-yellow-50 p-3 text-yellow-700 dark:border-yellow-800 dark:bg-yellow-950 dark:text-yellow-400">
<AlertTriangle className="h-5 w-5 flex-shrink-0" />
<span className="text-sm">
未脱敏的导出文件将包含明文 API 密钥和 Token,请妥善保管。
</span>
</div>
)}
<div className="flex gap-2">
<button
onClick={handleExport}
onClick={handleExportClick}
disabled={!config || isExporting}
className="flex items-center gap-2 rounded-lg bg-primary px-4 py-2 text-sm text-primary-foreground hover:bg-primary/90 disabled:opacity-50"
>
<Download className="h-4 w-4" />
{isExporting ? "导出中..." : "导出配置"}
{isExporting ? "导出中..." : "导出"}
</button>
</div>
<p className="text-sm text-muted-foreground">
导出当前配置为 YAML 文件,可用于备份或迁移到其他设备。
{exportScope === "config" &&
"导出当前配置为 YAML 文件,可用于备份或迁移。"}
{exportScope === "credentials" &&
"导出凭证池中的所有凭证,包括 OAuth Token 文件。"}
{exportScope === "full" &&
"导出完整的配置和凭证包,可用于完整迁移到其他设备。"}
</p>
</div>
</div>
@@ -154,7 +294,7 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) {
<input
ref={fileInputRef}
type="file"
accept=".yaml,.yml"
accept=".yaml,.yml,.json"
onChange={handleFileSelect}
className="hidden"
/>
@@ -168,55 +308,180 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) {
</button>
<p className="text-sm text-muted-foreground">
从 YAML 文件导入配置,支持合并或替换现有配置。
支持导入 YAML 配置文件或 JSON 导出包,支持合并或替换现有配置。
</p>
</div>
</div>
{/* Security Warning Dialog */}
{showSecurityWarning && (
<div className="fixed inset-0 z-50 flex items-center justify-center bg-black/50">
<div className="w-full max-w-md rounded-lg bg-background p-6 shadow-lg">
<div className="flex items-center gap-3 mb-4">
<div className="flex h-10 w-10 items-center justify-center rounded-full bg-yellow-100 dark:bg-yellow-900">
<AlertTriangle className="h-5 w-5 text-yellow-600 dark:text-yellow-400" />
</div>
<h3 className="text-lg font-medium">安全警告</h3>
</div>
<p className="text-sm text-muted-foreground mb-4">
您即将导出未脱敏的凭证数据,导出文件将包含明文 API 密钥和 OAuth
Token。 请确保:
</p>
<ul className="list-disc list-inside text-sm text-muted-foreground mb-4 space-y-1">
<li>不要将此文件分享给他人</li>
<li>不要上传到公共代码仓库</li>
<li>妥善保管导出文件</li>
</ul>
<div className="flex justify-end gap-2">
<button
onClick={() => setShowSecurityWarning(false)}
className="rounded-lg border px-4 py-2 text-sm hover:bg-muted"
>
取消
</button>
<button
onClick={performExport}
className="rounded-lg bg-yellow-600 px-4 py-2 text-sm text-white hover:bg-yellow-700"
>
我已了解,继续导出
</button>
</div>
</div>
</div>
)}
{/* Import Dialog */}
{showImportDialog && (
<div className="fixed inset-0 z-50 flex items-center justify-center bg-black/50">
<div className="w-full max-w-2xl rounded-lg bg-background p-6 shadow-lg">
<h3 className="text-lg font-medium mb-4">导入配置</h3>
<div className="w-full max-w-2xl rounded-lg bg-background p-6 shadow-lg max-h-[90vh] overflow-y-auto">
<h3 className="text-lg font-medium mb-4 flex items-center gap-2">
{getFileTypeIcon()}
导入配置 - {importFileName}
</h3>
{/* Validation Result */}
{validationResult && (
<div className="mb-4 space-y-2">
{validationResult.valid ? (
<div className="flex items-center gap-2 rounded-lg border border-green-200 bg-green-50 p-3 text-green-700 dark:border-green-800 dark:bg-green-950 dark:text-green-400">
<Check className="h-5 w-5" />
<span className="text-sm">文件格式有效</span>
</div>
) : (
<div className="flex items-start gap-2 rounded-lg border border-red-200 bg-red-50 p-3 text-red-700 dark:border-red-800 dark:bg-red-950 dark:text-red-400">
<AlertCircle className="h-5 w-5 flex-shrink-0" />
<div className="text-sm">
<p className="font-medium">文件格式无效</p>
<ul className="mt-1 list-disc list-inside">
{validationResult.errors.map((err, i) => (
<li key={i}>{err}</li>
))}
</ul>
</div>
</div>
)}
{/* Content info */}
{validationResult.valid && (
<div className="flex flex-wrap gap-2 text-sm">
{validationResult.version && (
<span className="rounded bg-muted px-2 py-1">
版本: {validationResult.version}
</span>
)}
{validationResult.has_config && (
<span className="rounded bg-blue-100 px-2 py-1 text-blue-700 dark:bg-blue-900 dark:text-blue-300">
包含配置
</span>
)}
{validationResult.has_credentials && (
<span className="rounded bg-purple-100 px-2 py-1 text-purple-700 dark:bg-purple-900 dark:text-purple-300">
包含凭证
</span>
)}
{validationResult.redacted && (
<span className="rounded bg-yellow-100 px-2 py-1 text-yellow-700 dark:bg-yellow-900 dark:text-yellow-300">
已脱敏
</span>
)}
</div>
)}
{/* Redaction warning */}
{validationResult.redacted && (
<div className="flex items-start gap-2 rounded-lg border border-yellow-200 bg-yellow-50 p-3 text-yellow-700 dark:border-yellow-800 dark:bg-yellow-950 dark:text-yellow-400">
<AlertTriangle className="h-5 w-5 flex-shrink-0" />
<span className="text-sm">
此导出包已脱敏,凭证数据(API 密钥、Token)无法恢复。
</span>
</div>
)}
{/* Validation warnings */}
{validationResult.warnings.length > 0 && (
<div className="rounded-lg border border-yellow-200 bg-yellow-50 p-3 dark:border-yellow-800 dark:bg-yellow-950">
<p className="text-sm font-medium text-yellow-700 dark:text-yellow-400">
警告
</p>
<ul className="mt-1 list-disc list-inside text-sm text-yellow-600 dark:text-yellow-500">
{validationResult.warnings.map((warning, i) => (
<li key={i}>{warning}</li>
))}
</ul>
</div>
)}
</div>
)}
{/* Preview */}
<div className="mb-4">
<label className="text-sm font-medium">配置预览</label>
<pre className="mt-2 max-h-64 overflow-auto rounded-lg bg-muted p-4 text-sm">
<label className="text-sm font-medium">内容预览</label>
<pre className="mt-2 max-h-48 overflow-auto rounded-lg bg-muted p-4 text-xs font-mono">
{importContent.slice(0, 2000)}
{importContent.length > 2000 && "\n..."}
</pre>
</div>
{/* Options */}
{/* Import Options */}
<div className="mb-4 space-y-2">
<label className="flex items-center gap-2">
<input
type="radio"
name="importMode"
checked={mergeConfig}
onChange={() => setMergeConfig(true)}
className="rounded-full border-gray-300"
/>
<span className="text-sm">合并到现有配置</span>
</label>
<label className="flex items-center gap-2">
<input
type="radio"
name="importMode"
checked={!mergeConfig}
onChange={() => setMergeConfig(false)}
className="rounded-full border-gray-300"
/>
<span className="text-sm">替换现有配置</span>
</label>
<label className="text-sm font-medium">导入模式</label>
<div className="flex flex-col gap-2">
<label className="flex items-center gap-2">
<input
type="radio"
name="importMode"
checked={mergeConfig}
onChange={() => setMergeConfig(true)}
className="rounded-full border-gray-300"
/>
<span className="text-sm">合并到现有配置</span>
<span className="text-xs text-muted-foreground">
(保留现有数据,添加新数据)
</span>
</label>
<label className="flex items-center gap-2">
<input
type="radio"
name="importMode"
checked={!mergeConfig}
onChange={() => setMergeConfig(false)}
className="rounded-full border-gray-300"
/>
<span className="text-sm">替换现有配置</span>
<span className="text-xs text-muted-foreground">
(完全覆盖现有数据)
</span>
</label>
</div>
</div>
{/* Warnings */}
{/* Import Result Warnings */}
{importResult?.warnings && importResult.warnings.length > 0 && (
<div className="mb-4 rounded-lg border border-yellow-200 bg-yellow-50 p-3 dark:border-yellow-800 dark:bg-yellow-950">
<p className="text-sm font-medium text-yellow-700 dark:text-yellow-400">
警告
导入警告
</p>
<ul className="mt-1 list-disc list-inside text-sm text-yellow-600 dark:text-yellow-500">
{importResult.warnings.map((warning, i) => (
@@ -245,32 +510,19 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) {
{/* Actions */}
<div className="flex justify-end gap-2">
<button
onClick={() => {
setShowImportDialog(false);
setImportContent("");
setImportResult(null);
setError(null);
}}
onClick={closeImportDialog}
className="rounded-lg border px-4 py-2 text-sm hover:bg-muted"
>
{importResult?.success ? "关闭" : "取消"}
</button>
{!importResult?.success && (
<>
<button
onClick={handleValidate}
className="rounded-lg border px-4 py-2 text-sm hover:bg-muted"
>
验证
</button>
<button
onClick={handleImport}
disabled={isImporting}
className="rounded-lg bg-primary px-4 py-2 text-sm text-primary-foreground hover:bg-primary/90 disabled:opacity-50"
>
{isImporting ? "导入中..." : "导入"}
</button>
</>
{!importResult?.success && validationResult?.valid && (
<button
onClick={handleImport}
disabled={isImporting}
className="rounded-lg bg-primary px-4 py-2 text-sm text-primary-foreground hover:bg-primary/90 disabled:opacity-50"
>
{isImporting ? "导入中..." : "导入"}
</button>
)}
</div>
</div>
+1
View File
@@ -1,3 +1,4 @@
export { ConfigPage } from "./ConfigPage";
export { ConfigEditor } from "./ConfigEditor";
export { ImportExport } from "./ImportExport";
export { AuthDirSettings } from "./AuthDirSettings";
@@ -87,7 +87,8 @@ export const ProviderPoolPage = forwardRef<ProviderPoolPageRef>(
setDeleteConfirm(null);
setDeletingCredentials((prev) => new Set(prev).add(uuid));
try {
await deleteCredential(uuid);
// Pass activeTab (provider_type) to enable YAML config sync
await deleteCredential(uuid, activeTab);
} catch (e) {
showError(e instanceof Error ? e.message : String(e), "delete", uuid);
} finally {
+147 -17
View File
@@ -1,5 +1,5 @@
import { useState, useEffect, forwardRef, useImperativeHandle } from "react";
import { RefreshCw, Route } from "lucide-react";
import { RefreshCw, Route, Sparkles, Check, Trash2 } from "lucide-react";
import { ModelMapping } from "./ModelMapping";
import { RoutingRules } from "./RoutingRules";
import { ExclusionList } from "./ExclusionList";
@@ -7,7 +7,12 @@ import { InjectionRules } from "./InjectionRules";
import { HelpTip } from "@/components/HelpTip";
import { routerApi } from "@/lib/api/router";
import { injectionApi } from "@/lib/api/injection";
import type { ModelAlias, RoutingRule, ProviderType } from "@/lib/api/router";
import type {
ModelAlias,
RoutingRule,
ProviderType,
RecommendedPreset,
} from "@/lib/api/router";
import type { InjectionRule } from "@/lib/api/injection";
export interface RoutingPageRef {
@@ -30,22 +35,34 @@ export const RoutingPage = forwardRef<RoutingPageRef>((_props, ref) => {
const [injectionRules, setInjectionRules] = useState<InjectionRule[]>([]);
const [injectionEnabled, setInjectionEnabled] = useState(false);
// Presets state
const [presets, setPresets] = useState<RecommendedPreset[]>([]);
const [showPresets, setShowPresets] = useState(false);
const [applyingPreset, setApplyingPreset] = useState<string | null>(null);
const refresh = async () => {
setLoading(true);
setError(null);
try {
const [aliasesData, rulesData, exclusionsData, injectionConfig] =
await Promise.all([
routerApi.getModelAliases(),
routerApi.getRoutingRules(),
routerApi.getExclusions(),
injectionApi.getInjectionConfig(),
]);
const [
aliasesData,
rulesData,
exclusionsData,
injectionConfig,
presetsData,
] = await Promise.all([
routerApi.getModelAliases(),
routerApi.getRoutingRules(),
routerApi.getExclusions(),
injectionApi.getInjectionConfig(),
routerApi.getRecommendedPresets(),
]);
setAliases(aliasesData);
setRules(rulesData);
setExclusions(exclusionsData);
setInjectionRules(injectionConfig.rules);
setInjectionEnabled(injectionConfig.enabled);
setPresets(presetsData);
} catch (e) {
setError(e instanceof Error ? e.message : String(e));
} finally {
@@ -53,6 +70,32 @@ export const RoutingPage = forwardRef<RoutingPageRef>((_props, ref) => {
}
};
const handleApplyPreset = async (
presetId: string,
merge: boolean = false,
) => {
setApplyingPreset(presetId);
try {
await routerApi.applyRecommendedPreset(presetId, merge);
await refresh();
setShowPresets(false);
} catch (e) {
setError(e instanceof Error ? e.message : String(e));
} finally {
setApplyingPreset(null);
}
};
const handleClearAll = async () => {
if (!confirm("确定要清空所有路由配置吗?此操作不可撤销。")) return;
try {
await routerApi.clearAllRoutingConfig();
await refresh();
} catch (e) {
setError(e instanceof Error ? e.message : String(e));
}
};
useImperativeHandle(ref, () => ({
refresh,
}));
@@ -149,16 +192,103 @@ export const RoutingPage = forwardRef<RoutingPageRef>((_props, ref) => {
配置模型映射、路由规则和排除列表
</p>
</div>
<button
onClick={refresh}
disabled={loading}
className="flex items-center gap-2 rounded-lg border px-3 py-2 text-sm hover:bg-muted disabled:opacity-50"
>
<RefreshCw className={`h-4 w-4 ${loading ? "animate-spin" : ""}`} />
刷新
</button>
<div className="flex items-center gap-2">
<button
onClick={() => setShowPresets(true)}
className="flex items-center gap-2 rounded-lg bg-primary px-3 py-2 text-sm text-primary-foreground hover:bg-primary/90"
>
<Sparkles className="h-4 w-4" />
推荐配置
</button>
<button
onClick={handleClearAll}
disabled={loading || (aliases.length === 0 && rules.length === 0)}
className="flex items-center gap-2 rounded-lg border border-red-300 px-3 py-2 text-sm text-red-600 hover:bg-red-50 disabled:opacity-50 dark:border-red-800 dark:text-red-400 dark:hover:bg-red-950/30"
>
<Trash2 className="h-4 w-4" />
清空
</button>
<button
onClick={refresh}
disabled={loading}
className="flex items-center gap-2 rounded-lg border px-3 py-2 text-sm hover:bg-muted disabled:opacity-50"
>
<RefreshCw className={`h-4 w-4 ${loading ? "animate-spin" : ""}`} />
刷新
</button>
</div>
</div>
{/* Presets Modal */}
{showPresets && (
<div className="fixed inset-0 z-50 flex items-center justify-center bg-black/50">
<div className="w-full max-w-2xl rounded-lg bg-background p-6 shadow-xl max-h-[80vh] overflow-y-auto">
<div className="flex items-center justify-between mb-4">
<h3 className="text-lg font-semibold flex items-center gap-2">
<Sparkles className="h-5 w-5 text-primary" />
推荐配置
</h3>
<button
onClick={() => setShowPresets(false)}
className="text-muted-foreground hover:text-foreground"
>
✕
</button>
</div>
<p className="text-sm text-muted-foreground mb-4">
选择一个预设配置快速设置路由规则和模型别名
</p>
<div className="space-y-3">
{presets.map((preset) => (
<div
key={preset.id}
className="rounded-lg border p-4 hover:border-primary/50 transition-colors"
>
<div className="flex items-start justify-between">
<div className="flex-1">
<h4 className="font-medium">{preset.name}</h4>
<p className="text-sm text-muted-foreground mt-1">
{preset.description}
</p>
<div className="flex gap-4 mt-2 text-xs text-muted-foreground">
<span>{preset.aliases.length} 个别名</span>
<span>{preset.rules.length} 条规则</span>
</div>
</div>
<div className="flex gap-2 ml-4">
<button
onClick={() => handleApplyPreset(preset.id, true)}
disabled={applyingPreset !== null}
className="flex items-center gap-1 rounded px-3 py-1.5 text-sm border hover:bg-muted disabled:opacity-50"
>
{applyingPreset === preset.id ? (
<RefreshCw className="h-3 w-3 animate-spin" />
) : (
<Check className="h-3 w-3" />
)}
合并
</button>
<button
onClick={() => handleApplyPreset(preset.id, false)}
disabled={applyingPreset !== null}
className="flex items-center gap-1 rounded bg-primary px-3 py-1.5 text-sm text-primary-foreground hover:bg-primary/90 disabled:opacity-50"
>
{applyingPreset === preset.id ? (
<RefreshCw className="h-3 w-3 animate-spin" />
) : (
<Check className="h-3 w-3" />
)}
应用
</button>
</div>
</div>
</div>
))}
</div>
</div>
</div>
)}
<HelpTip title="智能路由说明" variant="blue">
<ul className="list-disc list-inside space-y-1 text-sm text-blue-700 dark:text-blue-400">
<li>
+5 -2
View File
@@ -85,8 +85,11 @@ export function useProviderPool() {
};
// Delete credential
const deleteCredential = async (uuid: string) => {
await providerPoolApi.deleteCredential(uuid);
const deleteCredential = async (
uuid: string,
providerType?: PoolProviderType,
) => {
await providerPoolApi.deleteCredential(uuid, providerType);
await fetchOverview();
};
+73
View File
@@ -55,6 +55,28 @@ export interface LoggingConfig {
include_request_body: boolean;
}
// Credential pool types
export interface CredentialEntry {
id: string;
token_file: string;
disabled: boolean;
}
export interface ApiKeyEntry {
id: string;
api_key: string;
base_url?: string;
disabled: boolean;
}
export interface CredentialPoolConfig {
kiro: CredentialEntry[];
gemini: CredentialEntry[];
qwen: CredentialEntry[];
openai: ApiKeyEntry[];
claude: ApiKeyEntry[];
}
export interface Config {
server: ServerConfig;
providers: ProvidersConfig;
@@ -62,6 +84,8 @@ export interface Config {
routing: RoutingConfig;
retry: RetrySettings;
logging: LoggingConfig;
auth_dir: string;
credential_pool: CredentialPoolConfig;
}
// Export result
@@ -70,6 +94,33 @@ export interface ExportResult {
suggested_filename: string;
}
// Unified export options
export interface UnifiedExportOptions {
include_config: boolean;
include_credentials: boolean;
redact_secrets: boolean;
}
// Unified export result
export interface UnifiedExportResult {
content: string;
suggested_filename: string;
redacted: boolean;
has_config: boolean;
has_credentials: boolean;
}
// Validation result
export interface ValidationResult {
valid: boolean;
version: string | null;
redacted: boolean;
has_config: boolean;
has_credentials: boolean;
errors: string[];
warnings: string[];
}
// Import result
export interface ImportResult {
success: boolean;
@@ -94,6 +145,19 @@ export const configApi = {
return invoke("export_config", { config, redactSecrets });
},
// Export bundle (config + credentials)
async exportBundle(
config: Config,
options: UnifiedExportOptions,
): Promise<UnifiedExportResult> {
return invoke("export_bundle", { config, options });
},
// Validate import content (JSON bundle or YAML config)
async validateImport(content: string): Promise<ValidationResult> {
return invoke("validate_import", { content });
},
// Validate YAML config
async validateConfigYaml(yamlContent: string): Promise<Config> {
return invoke("validate_config_yaml", { yamlContent });
@@ -108,6 +172,15 @@ export const configApi = {
return invoke("import_config", { currentConfig, yamlContent, merge });
},
// Import bundle (JSON bundle or YAML config)
async importBundle(
currentConfig: Config,
content: string,
merge: boolean,
): Promise<ImportResult> {
return invoke("import_bundle", { currentConfig, content, merge });
},
// Get config file paths
async getConfigPaths(): Promise<ConfigPathInfo> {
return invoke("get_config_paths");
+5 -2
View File
@@ -195,8 +195,11 @@ export const providerPoolApi = {
},
// Delete a credential
async deleteCredential(uuid: string): Promise<boolean> {
return invoke("delete_provider_pool_credential", { uuid });
async deleteCredential(
uuid: string,
providerType?: PoolProviderType,
): Promise<boolean> {
return invoke("delete_provider_pool_credential", { uuid, providerType });
},
// Toggle credential enabled/disabled
+25
View File
@@ -29,6 +29,15 @@ export interface ExclusionPattern {
pattern: string;
}
// Recommended preset
export interface RecommendedPreset {
id: string;
name: string;
description: string;
aliases: ModelAlias[];
rules: RoutingRule[];
}
// Router configuration
export interface RouterConfig {
default_provider: ProviderType;
@@ -93,4 +102,20 @@ export const routerApi = {
async setDefaultProvider(provider: ProviderType): Promise<void> {
return invoke("set_router_default_provider", { provider });
},
// Recommended presets
async getRecommendedPresets(): Promise<RecommendedPreset[]> {
return invoke("get_recommended_presets");
},
async applyRecommendedPreset(
presetId: string,
merge: boolean = false,
): Promise<void> {
return invoke("apply_recommended_preset", { presetId, merge });
},
async clearAllRoutingConfig(): Promise<void> {
return invoke("clear_all_routing_config");
},
};