mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
feat: v0.10.0 - 配置与凭证统一导入导出功能
主要更新: - 实现以 YAML 配置文件为单一数据源的统一配置管理 - 支持凭证池配置同步到 YAML 文件 - 实现配置和凭证的统一导入导出 - 支持 OAuth Token 文件存储在可配置的 auth-dir 目录 - 实现配置热重载时自动同步凭证池 - 添加导出脱敏功能保护敏感信息 - 支持导入时合并或替换模式
This commit is contained in:
Generated
+2
-2
@@ -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
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"name": "proxycast",
|
||||
"private": true,
|
||||
"version": "0.9.0",
|
||||
"version": "0.10.0",
|
||||
"type": "module",
|
||||
"scripts": {
|
||||
"dev": "vite",
|
||||
|
||||
Generated
+1
-1
@@ -3274,7 +3274,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast"
|
||||
version = "0.9.0"
|
||||
version = "0.10.0"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"async-stream",
|
||||
|
||||
@@ -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
|
||||
@@ -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, ¤t_config, &options, ¤t_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, ¤t_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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
/// 切换凭证启用/禁用状态
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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 文件不存在"));
|
||||
}
|
||||
}
|
||||
@@ -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(¤t.credential_pool, &imported.credential_pool);
|
||||
|
||||
merged
|
||||
}
|
||||
|
||||
/// 合并凭证池
|
||||
///
|
||||
/// 将导入的凭证添加到现有凭证池中(按 ID 去重)
|
||||
fn merge_credential_pools(
|
||||
current: &CredentialPoolConfig,
|
||||
imported: &CredentialPoolConfig,
|
||||
) -> CredentialPoolConfig {
|
||||
CredentialPoolConfig {
|
||||
kiro: Self::merge_credential_entries(¤t.kiro, &imported.kiro),
|
||||
gemini: Self::merge_credential_entries(¤t.gemini, &imported.gemini),
|
||||
qwen: Self::merge_credential_entries(¤t.qwen, &imported.qwen),
|
||||
openai: Self::merge_api_key_entries(¤t.openai, &imported.openai),
|
||||
claude: Self::merge_api_key_entries(¤t.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, ¤t_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, ¤t, &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, ¤t, &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(¤t, &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(¤t, &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("脱敏数据"));
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
+1476
-3
File diff suppressed because it is too large
Load Diff
@@ -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]
|
||||
|
||||
@@ -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 配置文件路径(向后兼容)
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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 文件内容应该已更新"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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"));
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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};
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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 数应一致"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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,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",
|
||||
|
||||
@@ -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>
|
||||
);
|
||||
}
|
||||
@@ -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>
|
||||
);
|
||||
|
||||
@@ -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,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 {
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -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();
|
||||
};
|
||||
|
||||
|
||||
@@ -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");
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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");
|
||||
},
|
||||
};
|
||||
|
||||
Reference in New Issue
Block a user