diff --git a/package-lock.json b/package-lock.json index 9f64dfa20..25c3a4647 100644 --- a/package-lock.json +++ b/package-lock.json @@ -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", diff --git a/package.json b/package.json index e70cbe362..8820f831e 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.9.0", + "version": "0.10.0", "type": "module", "scripts": { "dev": "vite", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 26d0536f3..7713a7652 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -3274,7 +3274,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.9.0" +version = "0.10.0" dependencies = [ "anyhow", "async-stream", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 353fbe5b5..dfcac8e55 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -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" diff --git a/src-tauri/proptest-regressions/config/tests.txt b/src-tauri/proptest-regressions/config/tests.txt new file mode 100644 index 000000000..823f10fd2 --- /dev/null +++ b/src-tauri/proptest-regressions/config/tests.txt @@ -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\": }" diff --git a/src-tauri/proptest-regressions/websocket/tests.txt b/src-tauri/proptest-regressions/websocket/tests.txt new file mode 100644 index 000000000..212af077e --- /dev/null +++ b/src-tauri/proptest-regressions/websocket/tests.txt @@ -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 diff --git a/src-tauri/src/commands/config_cmd.rs b/src-tauri/src/commands/config_cmd.rs index b4a4c319a..95b98569e 100644 --- a/src-tauri/src/commands/config_cmd.rs +++ b/src-tauri/src/commands/config_cmd.rs @@ -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 { 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 { + 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 { + 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 { + 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 { + // 首先尝试解析为 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 { + 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 { + 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) +} diff --git a/src-tauri/src/commands/provider_pool_cmd.rs b/src-tauri/src/commands/provider_pool_cmd.rs index 7a887b2dd..5fe3f7f2d 100644 --- a/src-tauri/src/commands/provider_pool_cmd.rs +++ b/src-tauri/src/commands/provider_pool_cmd.rs @@ -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); +/// 凭证同步服务状态封装 +pub struct CredentialSyncServiceState(pub Option>); + /// 展开路径中的 ~ 为用户主目录 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 { - 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 { // 如果需要重新上传文件,先处理文件上传 - 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, ) -> Result { - 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::() { + if let Err(e) = sync.remove_credential(pool_type, &uuid) { + // 记录警告但不中断操作 + tracing::warn!("从 YAML 删除凭证失败: {}", e); + } + } + } + } + + Ok(result) } /// 切换凭证启用/禁用状态 diff --git a/src-tauri/src/commands/router_cmd.rs b/src-tauri/src/commands/router_cmd.rs index 36f6bb93c..2eec53078 100644 --- a/src-tauri/src/commands/router_cmd.rs +++ b/src-tauri/src/commands/router_cmd.rs @@ -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, + pub rules: Vec, +} + +/// 获取推荐配置列表 +#[tauri::command] +pub async fn get_recommended_presets() -> Result, 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(()) +} diff --git a/src-tauri/src/commands/telemetry_cmd.rs b/src-tauri/src/commands/telemetry_cmd.rs index 5be62b413..a4906f155 100644 --- a/src-tauri/src/commands/telemetry_cmd.rs +++ b/src-tauri/src/commands/telemetry_cmd.rs @@ -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, - pub stats: Arc, - pub tokens: Arc, + /// 统计聚合器(使用 RwLock 以支持与 RequestProcessor 共享) + pub stats: Arc>, + /// Token 追踪器(使用 RwLock 以支持与 RequestProcessor 共享) + pub tokens: Arc>, } impl TelemetryState { + /// 创建独立的遥测状态(使用自己的实例) pub fn new() -> Result { 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>, + tokens: Arc>, + logger: Option>, + ) -> Result { + 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, ) -> Result { 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, ) -> Result, 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, ) -> Result, 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, ) -> Result, 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 = state - .stats + let stats_guard = state.stats.read(); + let stats = stats_guard.summary(range); + let by_provider: HashMap = 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(); diff --git a/src-tauri/src/config/export.rs b/src-tauri/src/config/export.rs new file mode 100644 index 000000000..ed67ed06d --- /dev/null +++ b/src-tauri/src/config/export.rs @@ -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, + /// 应用版本 + pub app_version: String, + /// YAML 配置内容(如果包含配置) + #[serde(skip_serializing_if = "Option::is_none")] + pub config_yaml: Option, + /// OAuth Token 文件(base64 编码) + /// key: 相对于 auth_dir 的路径 + /// value: base64 编码的文件内容 + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + pub token_files: HashMap, + /// 是否已脱敏 + 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 { + serde_json::to_string_pretty(self).map_err(|e| ExportError::SerializeError(e.to_string())) + } + + /// 从 JSON 字符串反序列化 + pub fn from_json(json: &str) -> Result { + 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 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 { + 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 { + 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, 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, String> { + let data = data.trim_end_matches('='); + let mut result = Vec::new(); + + let decode_char = |c: char| -> Result { + 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 = 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 = (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 文件不存在")); + } +} diff --git a/src-tauri/src/config/import.rs b/src-tauri/src/config/import.rs new file mode 100644 index 000000000..5224c7c2e --- /dev/null +++ b/src-tauri/src/config/import.rs @@ -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, + /// 是否已脱敏 + pub redacted: bool, + /// 是否包含配置 + pub has_config: bool, + /// 是否包含凭证 + pub has_credentials: bool, + /// 错误信息列表 + pub errors: Vec, + /// 警告信息列表 + pub warnings: Vec, +} + +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) -> 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) { + self.errors.push(error.into()); + self.valid = false; + } + + /// 添加警告 + pub fn add_warning(&mut self, warning: impl Into) { + self.warnings.push(warning.into()); + } +} + +/// 导入结果 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ImportResult { + /// 是否成功 + pub success: bool, + /// 警告信息 + pub warnings: Vec, + /// 导入的配置 + 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) -> 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 for ImportError { + fn from(err: ConfigError) -> Self { + ImportError::ConfigError(err.to_string()) + } +} + +impl From 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 { + // 解析 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 { + 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 { + let mut result: Vec = 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 { + let mut result: Vec = 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)` - 警告信息列表 + fn restore_token_files( + token_files: &std::collections::HashMap, + auth_dir: &str, + ) -> Result, 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 { + 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("脱敏数据")); + } +} diff --git a/src-tauri/src/config/mod.rs b/src-tauri/src/config/mod.rs index 16a84c901..904491f32 100644 --- a/src-tauri/src/config/mod.rs +++ b/src-tauri/src/config/mod.rs @@ -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; diff --git a/src-tauri/src/config/path_utils.rs b/src-tauri/src/config/path_utils.rs new file mode 100644 index 000000000..0ccf9818e --- /dev/null +++ b/src-tauri/src/config/path_utils.rs @@ -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>(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>(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>(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); + } +} diff --git a/src-tauri/src/config/tests.rs b/src-tauri/src/config/tests.rs index ec87f1676..82adfa413 100644 --- a/src-tauri/src/config/tests.rs +++ b/src-tauri/src/config/tests.rs @@ -3,9 +3,9 @@ //! 使用 proptest 进行属性测试 use crate::config::{ - Config, ConfigManager, CustomProviderConfig, HotReloadManager, InjectionSettings, - LoggingConfig, ProviderConfig, ProvidersConfig, ReloadResult, RetrySettings, RoutingConfig, - ServerConfig, + collapse_tilde, contains_tilde, expand_tilde, Config, ConfigManager, CustomProviderConfig, + HotReloadManager, InjectionSettings, LoggingConfig, ProviderConfig, ProvidersConfig, + ReloadResult, RetrySettings, RoutingConfig, ServerConfig, YamlService, }; use proptest::prelude::*; use std::io::Write; @@ -210,6 +210,8 @@ fn arb_config() -> impl Strategy { retry, logging, injection: InjectionSettings::default(), + auth_dir: "~/.proxycast/auth".to_string(), + credential_pool: crate::config::CredentialPoolConfig::default(), }) } @@ -474,6 +476,8 @@ fn arb_valid_config() -> impl Strategy { retry, logging, injection: InjectionSettings::default(), + auth_dir: "~/.proxycast/auth".to_string(), + credential_pool: crate::config::CredentialPoolConfig::default(), }) } @@ -511,6 +515,8 @@ fn arb_invalid_config() -> impl Strategy { retry, logging, injection: InjectionSettings::default(), + auth_dir: "~/.proxycast/auth".to_string(), + credential_pool: crate::config::CredentialPoolConfig::default(), }; // 根据类型使配置无效 match invalid_type { @@ -708,3 +714,1470 @@ proptest! { } } } + +// ============================================================================ +// Property 3: Tilde Path Expansion +// ============================================================================ + +/// 生成有效的 tilde 路径(~/path 格式) +/// 排除 "." 和 ".." 路径段,因为这些会导致路径规范化问题 +fn arb_tilde_path() -> impl Strategy { + // 生成路径段:字母数字、下划线、连字符 + // 排除单独的 "." 和 ".." 以避免路径规范化问题 + let path_segment = "[a-zA-Z0-9_-]{1,20}"; + + // 生成 0-5 个路径段 + proptest::collection::vec(path_segment, 0..6).prop_map(|segments| { + if segments.is_empty() { + "~".to_string() + } else { + format!("~/{}", segments.join("/")) + } + }) +} + +/// 生成不包含 tilde 的绝对路径 +fn arb_absolute_path() -> impl Strategy { + let path_segment = "[a-zA-Z0-9_.-]{1,20}"; + + proptest::collection::vec(path_segment, 1..6) + .prop_map(|segments| format!("/{}", segments.join("/"))) +} + +/// 生成不包含 tilde 的相对路径 +fn arb_relative_path() -> impl Strategy { + let path_segment = "[a-zA-Z0-9_.-]{1,20}"; + + proptest::collection::vec(path_segment, 1..6).prop_map(|segments| segments.join("/")) +} + +/// 生成 ~user/path 格式的路径(不支持的格式) +fn arb_tilde_user_path() -> impl Strategy { + let username = "[a-z]{3,10}"; + let path_segment = "[a-zA-Z0-9_.-]{1,20}"; + + (username, proptest::collection::vec(path_segment, 0..4)).prop_map(|(user, segments)| { + if segments.is_empty() { + format!("~{}", user) + } else { + format!("~{}/{}", user, segments.join("/")) + } + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* valid tilde path (~/path format), expanding and then collapsing + /// should produce the original path. + /// **Validates: Requirements 2.3** + #[test] + fn prop_tilde_path_roundtrip(path in arb_tilde_path()) { + // 展开 tilde 路径 + let expanded = expand_tilde(&path); + + // 收缩回 tilde 格式 + let collapsed = collapse_tilde(&expanded); + + // 验证往返一致性 + prop_assert_eq!( + &collapsed, + &path, + "Tilde 路径往返不一致: 原始={}, 展开={:?}, 收缩={}", + path, + expanded, + collapsed + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* tilde path, the expanded path should start with the user's home directory. + /// **Validates: Requirements 2.3** + #[test] + fn prop_tilde_expansion_starts_with_home(path in arb_tilde_path()) { + let home_dir = dirs::home_dir().expect("应该能获取主目录"); + let expanded = expand_tilde(&path); + + prop_assert!( + expanded.starts_with(&home_dir), + "展开后的路径应以主目录开头: 路径={}, 展开={:?}, 主目录={:?}", + path, + expanded, + home_dir + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* tilde path, contains_tilde should return true before expansion + /// and false after expansion. + /// **Validates: Requirements 2.3** + #[test] + fn prop_contains_tilde_before_expansion(path in arb_tilde_path()) { + // 展开前应包含 tilde + prop_assert!( + contains_tilde(&path), + "展开前路径应包含 tilde: {}", + path + ); + + // 展开后不应包含 tilde + let expanded = expand_tilde(&path); + prop_assert!( + !contains_tilde(&expanded), + "展开后路径不应包含 tilde: {:?}", + expanded + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* absolute path (not starting with ~), expand_tilde should return + /// the path unchanged. + /// **Validates: Requirements 2.3** + #[test] + fn prop_absolute_path_unchanged(path in arb_absolute_path()) { + let expanded = expand_tilde(&path); + let expanded_str = expanded.to_string_lossy().to_string(); + + prop_assert_eq!( + &expanded_str, + &path, + "绝对路径应保持不变: 原始={}, 展开={:?}", + path, + expanded + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* relative path (not starting with ~ or /), expand_tilde should + /// return the path unchanged. + /// **Validates: Requirements 2.3** + #[test] + fn prop_relative_path_unchanged(path in arb_relative_path()) { + let expanded = expand_tilde(&path); + let expanded_str = expanded.to_string_lossy().to_string(); + + prop_assert_eq!( + &expanded_str, + &path, + "相对路径应保持不变: 原始={}, 展开={:?}", + path, + expanded + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* ~user/path format (unsupported), expand_tilde should return + /// the path unchanged. + /// **Validates: Requirements 2.3** + #[test] + fn prop_tilde_user_path_unchanged(path in arb_tilde_user_path()) { + let expanded = expand_tilde(&path); + let expanded_str = expanded.to_string_lossy().to_string(); + + prop_assert_eq!( + &expanded_str, + &path, + "~user/path 格式应保持不变: 原始={}, 展开={:?}", + path, + expanded + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* path under the home directory, collapse_tilde should produce + /// a path starting with ~. + /// **Validates: Requirements 2.3** + #[test] + fn prop_collapse_home_path_starts_with_tilde(subpath in arb_relative_path()) { + let home_dir = dirs::home_dir().expect("应该能获取主目录"); + let full_path = home_dir.join(&subpath); + + let collapsed = collapse_tilde(&full_path); + + prop_assert!( + collapsed.starts_with("~/"), + "主目录下的路径收缩后应以 ~/ 开头: 路径={:?}, 收缩={}", + full_path, + collapsed + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* path not under the home directory, collapse_tilde should return + /// the path unchanged. + /// **Validates: Requirements 2.3** + #[test] + fn prop_collapse_non_home_path_unchanged(path in arb_absolute_path()) { + // 确保路径不在主目录下(使用 /tmp 或类似路径) + let test_path = format!("/tmp{}", path); + let collapsed = collapse_tilde(&test_path); + + prop_assert_eq!( + &collapsed, + &test_path, + "非主目录路径应保持不变: 原始={}, 收缩={}", + test_path, + collapsed + ); + } +} + +// ============================================================================ +// Property 2: YAML Comment Preservation +// ============================================================================ + +/// 生成有效的 YAML 注释(以 # 开头) +fn arb_yaml_comment() -> impl Strategy { + // 生成注释内容:字母、数字、空格、中文字符 + "[a-zA-Z0-9 ]{1,50}".prop_map(|s| format!("# {}", s)) +} + +/// 生成带注释的 YAML 配置字符串 +fn arb_yaml_with_comments() -> impl Strategy)> { + ( + arb_valid_config(), + proptest::collection::vec(arb_yaml_comment(), 1..5), + ) + .prop_map(|(config, comments)| { + // 序列化配置 + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let lines: Vec<&str> = yaml.lines().collect(); + + // 在 YAML 中插入注释 + let mut result_lines: Vec = Vec::new(); + let mut comment_iter = comments.iter(); + + // 在文件开头添加一个注释 + if let Some(comment) = comment_iter.next() { + result_lines.push(comment.clone()); + } + + for (i, line) in lines.iter().enumerate() { + result_lines.push(line.to_string()); + + // 在某些行后添加注释 + if i % 5 == 0 { + if let Some(comment) = comment_iter.next() { + result_lines.push(comment.clone()); + } + } + } + + // 收集实际插入的注释 + let inserted_comments: Vec = result_lines + .iter() + .filter(|line| line.trim().starts_with('#')) + .cloned() + .collect(); + + (result_lines.join("\n"), inserted_comments) + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 2: YAML Comment Preservation** + /// *For any* YAML file with comments, saving configuration changes should preserve + /// all existing comments in their original positions. + /// **Validates: Requirements 1.3** + #[test] + fn prop_yaml_comment_preservation( + (yaml_with_comments, original_comments) in arb_yaml_with_comments(), + new_config in arb_valid_config() + ) { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 写入带注释的原始 YAML + std::fs::write(&config_path, &yaml_with_comments).expect("写入文件失败"); + + // 使用 YamlService 保存新配置(应保留注释) + YamlService::save_preserve_comments(&config_path, &new_config) + .expect("保存配置失败"); + + // 读取保存后的内容 + let saved_content = std::fs::read_to_string(&config_path).expect("读取文件失败"); + + // 提取保存后的注释 + let saved_comments: Vec = saved_content + .lines() + .filter(|line| line.trim().starts_with('#')) + .map(|s| s.to_string()) + .collect(); + + // 验证注释被保留 + // 注意:由于 YAML 结构可能变化,我们只验证注释内容被保留,不验证位置 + for original_comment in &original_comments { + let comment_content = original_comment.trim(); + let found = saved_comments.iter().any(|c| c.trim() == comment_content); + prop_assert!( + found, + "注释应被保留: 原始注释='{}', 保存后的注释={:?}", + comment_content, + saved_comments + ); + } + } + + /// **Feature: config-credential-export, Property 2: YAML Comment Preservation** + /// *For any* configuration saved with YamlService, the configuration should be + /// correctly parseable and equivalent to the original. + /// **Validates: Requirements 1.3** + #[test] + fn prop_yaml_save_preserves_config(config in arb_valid_config()) { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 使用 YamlService 保存配置 + YamlService::save_preserve_comments(&config_path, &config) + .expect("保存配置失败"); + + // 读取并解析保存后的配置 + let saved_content = std::fs::read_to_string(&config_path).expect("读取文件失败"); + let parsed_config = ConfigManager::parse_yaml(&saved_content).expect("解析配置失败"); + + // 验证配置一致性 + prop_assert_eq!( + config.server, + parsed_config.server, + "服务器配置应一致" + ); + prop_assert_eq!( + config.providers, + parsed_config.providers, + "Provider 配置应一致" + ); + prop_assert_eq!( + config.retry, + parsed_config.retry, + "重试配置应一致" + ); + prop_assert_eq!( + config.logging, + parsed_config.logging, + "日志配置应一致" + ); + } + + /// **Feature: config-credential-export, Property 2: YAML Comment Preservation** + /// *For any* YAML file with header comments, saving should preserve header comments. + /// **Validates: Requirements 1.3** + #[test] + fn prop_yaml_header_comment_preservation( + header_comment in arb_yaml_comment(), + config in arb_valid_config() + ) { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 创建带头部注释的 YAML + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let yaml_with_header = format!("{}\n{}", header_comment, yaml); + + // 写入文件 + std::fs::write(&config_path, &yaml_with_header).expect("写入文件失败"); + + // 使用 YamlService 保存新配置 + YamlService::save_preserve_comments(&config_path, &config) + .expect("保存配置失败"); + + // 读取保存后的内容 + let saved_content = std::fs::read_to_string(&config_path).expect("读取文件失败"); + + // 验证头部注释被保留 + let header_content = header_comment.trim(); + let has_header = saved_content.lines().any(|line| line.trim() == header_content); + + prop_assert!( + has_header, + "头部注释应被保留: 原始='{}', 保存后内容前100字符='{}'", + header_content, + &saved_content[..saved_content.len().min(100)] + ); + } +} + +// ============================================================================ +// Unit Tests for YamlService::update_field +// ============================================================================ + +#[test] +fn test_update_field_simple() { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 写入初始 YAML + let initial_yaml = r#"server: + host: 127.0.0.1 + port: 8999 + api_key: test_key +"#; + std::fs::write(&config_path, initial_yaml).expect("写入文件失败"); + + // 更新 port 字段 + YamlService::update_field(&config_path, &["server", "port"], "9000").expect("更新字段失败"); + + // 读取并验证 + let content = std::fs::read_to_string(&config_path).expect("读取文件失败"); + assert!(content.contains("port: 9000"), "端口应被更新为 9000"); + assert!(content.contains("host: 127.0.0.1"), "其他字段应保持不变"); +} + +#[test] +fn test_update_field_preserves_comments() { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 写入带注释的 YAML + let initial_yaml = r#"# 服务器配置 +server: + # 监听地址 + host: 127.0.0.1 + # 监听端口 + port: 8999 + api_key: test_key +"#; + std::fs::write(&config_path, initial_yaml).expect("写入文件失败"); + + // 更新 port 字段 + YamlService::update_field(&config_path, &["server", "port"], "9000").expect("更新字段失败"); + + // 读取并验证 + let content = std::fs::read_to_string(&config_path).expect("读取文件失败"); + assert!(content.contains("port: 9000"), "端口应被更新为 9000"); + assert!(content.contains("# 服务器配置"), "头部注释应保留"); + assert!(content.contains("# 监听地址"), "字段注释应保留"); + assert!(content.contains("# 监听端口"), "字段注释应保留"); +} + +#[test] +fn test_update_field_not_found() { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 写入初始 YAML + let initial_yaml = r#"server: + host: 127.0.0.1 + port: 8999 +"#; + std::fs::write(&config_path, initial_yaml).expect("写入文件失败"); + + // 尝试更新不存在的字段 + let result = YamlService::update_field(&config_path, &["server", "nonexistent"], "value"); + assert!(result.is_err(), "更新不存在的字段应返回错误"); +} + +#[test] +fn test_update_field_nested() { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 写入初始 YAML + let initial_yaml = r#"server: + host: 127.0.0.1 + port: 8999 +providers: + kiro: + enabled: true + region: us-east-1 +"#; + std::fs::write(&config_path, initial_yaml).expect("写入文件失败"); + + // 更新嵌套字段 + YamlService::update_field(&config_path, &["providers", "kiro", "region"], "us-west-2") + .expect("更新字段失败"); + + // 读取并验证 + let content = std::fs::read_to_string(&config_path).expect("读取文件失败"); + assert!( + content.contains("region: us-west-2"), + "region 应被更新为 us-west-2" + ); + assert!(content.contains("enabled: true"), "其他字段应保持不变"); +} + +// ============================================================================ +// Property 4: Export Scope Filtering +// ============================================================================ + +use crate::config::{ + ApiKeyEntry, CredentialEntry, CredentialPoolConfig, ExportOptions, ExportService, +}; + +/// 生成随机的 OAuth 凭证条目 +fn arb_credential_entry() -> impl Strategy { + ( + "[a-z]{3,10}-[0-9]{1,5}".prop_map(|s| s), + "[a-z]+/token-[0-9]{1,5}\\.json".prop_map(|s| s), + any::(), + ) + .prop_map(|(id, token_file, disabled)| CredentialEntry { + id, + token_file, + disabled, + }) +} + +/// 生成随机的 API Key 凭证条目 +fn arb_api_key_entry() -> impl Strategy { + ( + "[a-z]{3,10}-[0-9]{1,5}".prop_map(|s| s), + "sk-[a-zA-Z0-9]{20,40}".prop_map(|s| s), + proptest::option::of("https://api\\.[a-z]+\\.com/v[0-9]".prop_map(|s| s)), + any::(), + ) + .prop_map(|(id, api_key, base_url, disabled)| ApiKeyEntry { + id, + api_key, + base_url, + disabled, + }) +} + +/// 生成随机的凭证池配置 +fn arb_credential_pool_config() -> impl Strategy { + ( + proptest::collection::vec(arb_credential_entry(), 0..3), + proptest::collection::vec(arb_credential_entry(), 0..3), + proptest::collection::vec(arb_credential_entry(), 0..3), + proptest::collection::vec(arb_api_key_entry(), 0..3), + proptest::collection::vec(arb_api_key_entry(), 0..3), + ) + .prop_map( + |(kiro, gemini, qwen, openai, claude)| CredentialPoolConfig { + kiro, + gemini, + qwen, + openai, + claude, + }, + ) +} + +/// 生成带凭证池的配置 +fn arb_config_with_credentials() -> impl Strategy { + (arb_valid_config(), arb_credential_pool_config()).prop_map(|(mut config, pool)| { + config.credential_pool = pool; + config + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 4: Export Scope Filtering** + /// *For any* export operation with config-only scope, the resulting bundle + /// should contain only configuration data and no credential token files. + /// **Validates: Requirements 3.2** + #[test] + fn prop_export_scope_config_only(config in arb_config_with_credentials()) { + let options = ExportOptions::config_only(); + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 验证只包含配置 + prop_assert!( + bundle.has_config(), + "config-only 导出应包含配置" + ); + prop_assert!( + !bundle.has_credentials(), + "config-only 导出不应包含凭证 token 文件" + ); + } + + /// **Feature: config-credential-export, Property 4: Export Scope Filtering** + /// *For any* export operation with credentials-only scope, the resulting bundle + /// should contain only credential data and no configuration YAML. + /// **Validates: Requirements 3.2** + #[test] + fn prop_export_scope_credentials_only(config in arb_config_with_credentials()) { + let options = ExportOptions::credentials_only(); + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 验证不包含配置 + prop_assert!( + !bundle.has_config(), + "credentials-only 导出不应包含配置 YAML" + ); + // 注意:token_files 可能为空(如果没有实际的 token 文件存在) + // 但 config_yaml 必须为 None + prop_assert!( + bundle.config_yaml.is_none(), + "credentials-only 导出的 config_yaml 应为 None" + ); + } + + /// **Feature: config-credential-export, Property 4: Export Scope Filtering** + /// *For any* export operation with full scope, the resulting bundle + /// should contain both configuration and credential data. + /// **Validates: Requirements 3.2** + #[test] + fn prop_export_scope_full(config in arb_config_with_credentials()) { + let options = ExportOptions::full(); + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 验证包含配置 + prop_assert!( + bundle.has_config(), + "full 导出应包含配置" + ); + // token_files 可能为空(如果没有实际的 token 文件存在) + // 但 config_yaml 必须存在 + prop_assert!( + bundle.config_yaml.is_some(), + "full 导出的 config_yaml 应存在" + ); + } + + /// **Feature: config-credential-export, Property 4: Export Scope Filtering** + /// *For any* export operation, the bundle should correctly reflect the + /// include_config and include_credentials options. + /// **Validates: Requirements 3.2** + #[test] + fn prop_export_scope_matches_options( + config in arb_config_with_credentials(), + include_config in any::(), + include_credentials in any::() + ) { + let options = ExportOptions { + include_config, + include_credentials, + redact_secrets: false, + }; + + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 验证配置包含状态与选项一致 + prop_assert_eq!( + bundle.has_config(), + include_config, + "配置包含状态应与 include_config 选项一致" + ); + + // 验证 config_yaml 存在性与选项一致 + prop_assert_eq!( + bundle.config_yaml.is_some(), + include_config, + "config_yaml 存在性应与 include_config 选项一致" + ); + } +} + +// ============================================================================ +// Property 5: Redaction Completeness +// ============================================================================ + +use crate::config::REDACTED_PLACEHOLDER; + +/// 生成包含敏感信息的配置 +fn arb_config_with_secrets() -> impl Strategy { + ( + arb_valid_config(), + arb_credential_pool_config(), + // 生成看起来像真实 API 密钥的字符串 + "sk-[a-zA-Z0-9]{20,40}".prop_map(|s| s), + proptest::option::of("sk-[a-zA-Z0-9]{20,40}".prop_map(|s| s)), + proptest::option::of("sk-ant-[a-zA-Z0-9]{20,40}".prop_map(|s| s)), + ) + .prop_map(|(mut config, pool, server_key, openai_key, claude_key)| { + config.server.api_key = server_key; + config.providers.openai.api_key = openai_key; + config.providers.claude.api_key = claude_key; + config.credential_pool = pool; + config + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 5: Redaction Completeness** + /// *For any* export with redaction enabled, all sensitive values (API keys, tokens, + /// secrets) should be replaced with placeholder markers, and no original sensitive + /// data should remain. + /// **Validates: Requirements 3.4** + #[test] + fn prop_redaction_removes_all_secrets(config in arb_config_with_secrets()) { + // 脱敏配置 + let redacted = ExportService::redact_config(&config); + + // 验证脱敏后不包含敏感信息 + prop_assert!( + !ExportService::contains_secrets(&redacted), + "脱敏后的配置不应包含敏感信息" + ); + + // 验证服务器 API 密钥已脱敏 + prop_assert_eq!( + &redacted.server.api_key, + REDACTED_PLACEHOLDER, + "服务器 API 密钥应被脱敏" + ); + + // 验证 OpenAI API 密钥已脱敏(如果存在) + if config.providers.openai.api_key.is_some() { + prop_assert_eq!( + redacted.providers.openai.api_key.as_deref(), + Some(REDACTED_PLACEHOLDER), + "OpenAI API 密钥应被脱敏" + ); + } + + // 验证 Claude API 密钥已脱敏(如果存在) + if config.providers.claude.api_key.is_some() { + prop_assert_eq!( + redacted.providers.claude.api_key.as_deref(), + Some(REDACTED_PLACEHOLDER), + "Claude API 密钥应被脱敏" + ); + } + } + + /// **Feature: config-credential-export, Property 5: Redaction Completeness** + /// *For any* configuration with API keys in credential pool, redaction should + /// replace all API keys with placeholder markers. + /// **Validates: Requirements 3.4** + #[test] + fn prop_redaction_credential_pool_api_keys(config in arb_config_with_secrets()) { + let redacted = ExportService::redact_config(&config); + + // 验证 OpenAI 凭证池中的 API 密钥已脱敏 + for (i, entry) in redacted.credential_pool.openai.iter().enumerate() { + prop_assert_eq!( + &entry.api_key, + REDACTED_PLACEHOLDER, + "OpenAI 凭证池条目 {} 的 API 密钥应被脱敏", + i + ); + } + + // 验证 Claude 凭证池中的 API 密钥已脱敏 + for (i, entry) in redacted.credential_pool.claude.iter().enumerate() { + prop_assert_eq!( + &entry.api_key, + REDACTED_PLACEHOLDER, + "Claude 凭证池条目 {} 的 API 密钥应被脱敏", + i + ); + } + } + + /// **Feature: config-credential-export, Property 5: Redaction Completeness** + /// *For any* export with redaction enabled, the exported YAML should not contain + /// any original sensitive values. + /// **Validates: Requirements 3.4** + #[test] + fn prop_redaction_yaml_no_secrets(config in arb_config_with_secrets()) { + // 导出带脱敏的 YAML + let yaml = ExportService::export_yaml(&config, true) + .expect("导出应成功"); + + // 验证 YAML 中不包含原始敏感值 + // 检查原始服务器 API 密钥 + if !config.server.api_key.is_empty() && config.server.api_key != REDACTED_PLACEHOLDER { + prop_assert!( + !yaml.contains(&config.server.api_key), + "YAML 不应包含原始服务器 API 密钥: {}", + config.server.api_key + ); + } + + // 检查原始 OpenAI API 密钥 + if let Some(ref key) = config.providers.openai.api_key { + if !key.is_empty() && key != REDACTED_PLACEHOLDER { + prop_assert!( + !yaml.contains(key), + "YAML 不应包含原始 OpenAI API 密钥" + ); + } + } + + // 检查原始 Claude API 密钥 + if let Some(ref key) = config.providers.claude.api_key { + if !key.is_empty() && key != REDACTED_PLACEHOLDER { + prop_assert!( + !yaml.contains(key), + "YAML 不应包含原始 Claude API 密钥" + ); + } + } + + // 检查凭证池中的 API 密钥 + for entry in &config.credential_pool.openai { + if !entry.api_key.is_empty() && entry.api_key != REDACTED_PLACEHOLDER { + prop_assert!( + !yaml.contains(&entry.api_key), + "YAML 不应包含原始 OpenAI 凭证池 API 密钥" + ); + } + } + + for entry in &config.credential_pool.claude { + if !entry.api_key.is_empty() && entry.api_key != REDACTED_PLACEHOLDER { + prop_assert!( + !yaml.contains(&entry.api_key), + "YAML 不应包含原始 Claude 凭证池 API 密钥" + ); + } + } + } + + /// **Feature: config-credential-export, Property 5: Redaction Completeness** + /// *For any* configuration, redaction should preserve non-sensitive data unchanged. + /// **Validates: Requirements 3.4** + #[test] + fn prop_redaction_preserves_non_sensitive_data(config in arb_config_with_secrets()) { + let redacted = ExportService::redact_config(&config); + + // 验证非敏感数据保持不变 + prop_assert_eq!( + config.server.host, + redacted.server.host, + "服务器主机应保持不变" + ); + prop_assert_eq!( + config.server.port, + redacted.server.port, + "服务器端口应保持不变" + ); + prop_assert_eq!( + config.providers.kiro.enabled, + redacted.providers.kiro.enabled, + "Kiro 启用状态应保持不变" + ); + prop_assert_eq!( + config.routing.default_provider, + redacted.routing.default_provider, + "默认 Provider 应保持不变" + ); + prop_assert_eq!( + config.retry, + redacted.retry, + "重试配置应保持不变" + ); + prop_assert_eq!( + config.logging, + redacted.logging, + "日志配置应保持不变" + ); + + // 验证 OAuth 凭证条目保持不变(它们不包含敏感信息) + prop_assert_eq!( + config.credential_pool.kiro, + redacted.credential_pool.kiro, + "Kiro 凭证条目应保持不变" + ); + prop_assert_eq!( + config.credential_pool.gemini, + redacted.credential_pool.gemini, + "Gemini 凭证条目应保持不变" + ); + prop_assert_eq!( + config.credential_pool.qwen, + redacted.credential_pool.qwen, + "Qwen 凭证条目应保持不变" + ); + } + + /// **Feature: config-credential-export, Property 5: Redaction Completeness** + /// *For any* export bundle with redaction, the redacted flag should be true. + /// **Validates: Requirements 3.4** + #[test] + fn prop_redaction_bundle_flag(config in arb_config_with_secrets()) { + let options = ExportOptions::redacted(); + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + prop_assert!( + bundle.is_redacted(), + "脱敏导出的 bundle 应标记为已脱敏" + ); + } +} + +// ============================================================================ +// Property 6: Import Validation +// ============================================================================ + +use crate::config::{ExportBundle, ImportService}; + +/// 生成有效的导出包 +fn arb_valid_export_bundle() -> impl Strategy { + ( + arb_valid_config(), + any::(), // redacted + "[0-9]+\\.[0-9]+\\.[0-9]+".prop_map(|s| s), // app_version + ) + .prop_map(|(config, redacted, app_version)| { + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let mut bundle = ExportBundle::new(&app_version); + bundle.config_yaml = Some(yaml); + bundle.redacted = redacted; + bundle + }) +} + +/// 生成无效的导入内容(既不是有效的 JSON ExportBundle,也不是有效的 YAML Config) +/// 注意:YAML 解析器非常宽松,大多数内容都可以解析为某种 YAML 结构 +/// 因此我们只测试语法错误的内容 +fn arb_invalid_import_content() -> impl Strategy { + prop_oneof![ + // 无效的 JSON/YAML 语法 + Just("{invalid json".to_string()), + Just("invalid: yaml: content: [".to_string()), + Just(" - bad\n indentation".to_string()), + Just("key: value\n invalid: indent".to_string()), + // 有效的 JSON 但不是 ExportBundle 或 Config(数组类型) + Just("[1, 2, 3]".to_string()), + ] +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 6: Import Validation** + /// *For any* valid export bundle, validation should return valid=true and + /// correctly identify format, version, and redaction status. + /// **Validates: Requirements 4.1, 4.2** + #[test] + fn prop_import_validation_valid_bundle(bundle in arb_valid_export_bundle()) { + let json = bundle.to_json().expect("序列化应成功"); + let result = ImportService::validate(&json); + + // 验证结果应为有效 + prop_assert!( + result.valid, + "有效的导出包应通过验证: errors={:?}", + result.errors + ); + + // 验证版本被正确识别 + prop_assert_eq!( + result.version, + Some(bundle.version.clone()), + "版本应被正确识别" + ); + + // 验证脱敏状态被正确识别 + prop_assert_eq!( + result.redacted, + bundle.redacted, + "脱敏状态应被正确识别" + ); + + // 验证配置存在性被正确识别 + prop_assert_eq!( + result.has_config, + bundle.has_config(), + "配置存在性应被正确识别" + ); + } + + /// **Feature: config-credential-export, Property 6: Import Validation** + /// *For any* valid YAML configuration, validation should return valid=true + /// and identify it as config-only (no credentials). + /// **Validates: Requirements 4.1, 4.2** + #[test] + fn prop_import_validation_valid_yaml(config in arb_valid_config()) { + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let result = ImportService::validate(&yaml); + + // 验证结果应为有效 + prop_assert!( + result.valid, + "有效的 YAML 配置应通过验证: errors={:?}", + result.errors + ); + + // 验证识别为配置 + prop_assert!( + result.has_config, + "应识别为包含配置" + ); + + // YAML 配置不包含凭证 token 文件 + prop_assert!( + !result.has_credentials, + "YAML 配置不应包含凭证 token 文件" + ); + + // YAML 配置不是脱敏的 + prop_assert!( + !result.redacted, + "YAML 配置不应标记为脱敏" + ); + } + + /// **Feature: config-credential-export, Property 6: Import Validation** + /// *For any* invalid import content (neither valid ExportBundle JSON nor valid Config YAML), + /// validation should return valid=false with appropriate error messages. + /// **Validates: Requirements 4.1, 4.2** + #[test] + fn prop_import_validation_invalid_content(content in arb_invalid_import_content()) { + let result = ImportService::validate(&content); + + // 无效内容应验证失败 + prop_assert!( + !result.valid, + "无效的导入内容应验证失败: content={}", content + ); + + // 应有错误信息 + prop_assert!( + !result.errors.is_empty(), + "应有错误信息" + ); + } + + /// **Feature: config-credential-export, Property 6: Import Validation** + /// *For any* redacted export bundle, validation should warn about + /// credentials that cannot be restored. + /// **Validates: Requirements 4.1, 4.2** + #[test] + fn prop_import_validation_redacted_warning(config in arb_valid_config()) { + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let mut bundle = ExportBundle::new("1.0.0"); + bundle.config_yaml = Some(yaml); + bundle.redacted = true; + + let json = bundle.to_json().expect("序列化应成功"); + let result = ImportService::validate(&json); + + // 验证结果应为有效(脱敏不影响有效性) + prop_assert!( + result.valid, + "脱敏的导出包应通过验证" + ); + + // 应有脱敏警告 + prop_assert!( + !result.warnings.is_empty(), + "脱敏的导出包应有警告信息" + ); + + // 警告应提及脱敏 + let has_redaction_warning = result.warnings.iter().any(|w| + w.contains("脱敏") || w.contains("redact") + ); + prop_assert!( + has_redaction_warning, + "应有关于脱敏的警告: {:?}", + result.warnings + ); + } +} + +// ============================================================================ +// Property 7: Import Merge vs Replace +// ============================================================================ + +use crate::config::ImportOptions; + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 7: Import Merge vs Replace** + /// *For any* import operation in replace mode, the resulting configuration + /// should be exactly the imported configuration (not merged with current). + /// **Validates: Requirements 4.3** + #[test] + fn prop_import_replace_mode( + current_config in arb_config_with_credentials(), + imported_config in arb_config_with_credentials() + ) { + let yaml = ConfigManager::to_yaml(&imported_config).expect("序列化应成功"); + let options = ImportOptions::replace(); + + let result = ImportService::import_yaml(&yaml, ¤t_config, &options) + .expect("导入应成功"); + + // 替换模式下,结果应等于导入的配置 + prop_assert_eq!( + result.config.server, + imported_config.server, + "替换模式下服务器配置应等于导入的配置" + ); + prop_assert_eq!( + result.config.providers, + imported_config.providers, + "替换模式下 Provider 配置应等于导入的配置" + ); + prop_assert_eq!( + result.config.routing.default_provider, + imported_config.routing.default_provider, + "替换模式下默认 Provider 应等于导入的配置" + ); + prop_assert_eq!( + result.config.retry, + imported_config.retry, + "替换模式下重试配置应等于导入的配置" + ); + prop_assert_eq!( + result.config.logging, + imported_config.logging, + "替换模式下日志配置应等于导入的配置" + ); + } + + /// **Feature: config-credential-export, Property 7: Import Merge vs Replace** + /// *For any* import operation in merge mode, the resulting configuration + /// should combine new data with existing data. + /// **Validates: Requirements 4.3** + #[test] + fn prop_import_merge_mode_combines_credentials( + current_config in arb_config_with_credentials(), + imported_config in arb_config_with_credentials() + ) { + let yaml = ConfigManager::to_yaml(&imported_config).expect("序列化应成功"); + let options = ImportOptions::merge(); + + let result = ImportService::import_yaml(&yaml, ¤t_config, &options) + .expect("导入应成功"); + + // 合并模式下,凭证池应包含两边的凭证(按 ID 去重) + // 计算预期的凭证数量(去重后) + let expected_kiro_ids: std::collections::HashSet<_> = current_config + .credential_pool + .kiro + .iter() + .chain(imported_config.credential_pool.kiro.iter()) + .map(|e| e.id.clone()) + .collect(); + + prop_assert_eq!( + result.config.credential_pool.kiro.len(), + expected_kiro_ids.len(), + "合并模式下 Kiro 凭证数量应为去重后的总数" + ); + + let expected_openai_ids: std::collections::HashSet<_> = current_config + .credential_pool + .openai + .iter() + .chain(imported_config.credential_pool.openai.iter()) + .filter(|e| e.api_key != REDACTED_PLACEHOLDER) + .map(|e| e.id.clone()) + .collect(); + + // OpenAI 凭证数量应包含所有非脱敏的凭证 + prop_assert!( + result.config.credential_pool.openai.len() >= expected_openai_ids.len().saturating_sub( + imported_config.credential_pool.openai.iter() + .filter(|e| e.api_key == REDACTED_PLACEHOLDER) + .count() + ), + "合并模式下 OpenAI 凭证应包含所有非脱敏的凭证" + ); + } + + /// **Feature: config-credential-export, Property 7: Import Merge vs Replace** + /// *For any* import operation in merge mode, imported values should override + /// current values for the same keys. + /// **Validates: Requirements 4.3** + #[test] + fn prop_import_merge_mode_overrides_config( + current_config in arb_config_with_credentials(), + imported_config in arb_config_with_credentials() + ) { + let yaml = ConfigManager::to_yaml(&imported_config).expect("序列化应成功"); + let options = ImportOptions::merge(); + + let result = ImportService::import_yaml(&yaml, ¤t_config, &options) + .expect("导入应成功"); + + // 合并模式下,配置值应被导入的值覆盖 + prop_assert_eq!( + result.config.server, + imported_config.server, + "合并模式下服务器配置应被导入的值覆盖" + ); + prop_assert_eq!( + result.config.providers, + imported_config.providers, + "合并模式下 Provider 配置应被导入的值覆盖" + ); + prop_assert_eq!( + result.config.retry, + imported_config.retry, + "合并模式下重试配置应被导入的值覆盖" + ); + } + + /// **Feature: config-credential-export, Property 7: Import Merge vs Replace** + /// *For any* import operation in replace mode with empty imported credentials, + /// the result should have empty credentials (not preserve current). + /// **Validates: Requirements 4.3** + #[test] + fn prop_import_replace_mode_clears_credentials( + current_config in arb_config_with_credentials() + ) { + // 创建一个没有凭证的配置 + let mut imported_config = Config::default(); + imported_config.server.port = 9999; // 修改一个值以区分 + + let yaml = ConfigManager::to_yaml(&imported_config).expect("序列化应成功"); + let options = ImportOptions::replace(); + + let result = ImportService::import_yaml(&yaml, ¤t_config, &options) + .expect("导入应成功"); + + // 替换模式下,凭证池应为空(因为导入的配置没有凭证) + prop_assert!( + result.config.credential_pool.kiro.is_empty(), + "替换模式下 Kiro 凭证应为空" + ); + prop_assert!( + result.config.credential_pool.openai.is_empty(), + "替换模式下 OpenAI 凭证应为空" + ); + prop_assert_eq!( + result.config.server.port, + 9999, + "替换模式下端口应为导入的值" + ); + } +} + +// ============================================================================ +// Property 8: Export-Import Round Trip +// ============================================================================ + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 8: Export-Import Round Trip** + /// *For any* valid configuration, exporting (without redaction) and then importing + /// should produce an equivalent configuration. + /// **Validates: Requirements 5.5** + #[test] + fn prop_export_import_roundtrip_yaml(config in arb_config_with_credentials()) { + // 导出为 YAML(不脱敏) + let yaml = ExportService::export_yaml(&config, false) + .expect("导出应成功"); + + // 导入 YAML(替换模式) + let empty_config = Config::default(); + let options = ImportOptions::replace(); + let result = ImportService::import_yaml(&yaml, &empty_config, &options) + .expect("导入应成功"); + + // 验证往返一致性 + prop_assert_eq!( + config.server, + result.config.server, + "服务器配置往返不一致" + ); + prop_assert_eq!( + config.providers, + result.config.providers, + "Provider 配置往返不一致" + ); + prop_assert_eq!( + config.routing.default_provider, + result.config.routing.default_provider, + "默认 Provider 往返不一致" + ); + prop_assert_eq!( + config.retry, + result.config.retry, + "重试配置往返不一致" + ); + prop_assert_eq!( + config.logging, + result.config.logging, + "日志配置往返不一致" + ); + prop_assert_eq!( + config.auth_dir, + result.config.auth_dir, + "auth_dir 往返不一致" + ); + + // 验证凭证池往返一致性 + prop_assert_eq!( + config.credential_pool.kiro, + result.config.credential_pool.kiro, + "Kiro 凭证池往返不一致" + ); + prop_assert_eq!( + config.credential_pool.gemini, + result.config.credential_pool.gemini, + "Gemini 凭证池往返不一致" + ); + prop_assert_eq!( + config.credential_pool.qwen, + result.config.credential_pool.qwen, + "Qwen 凭证池往返不一致" + ); + prop_assert_eq!( + config.credential_pool.openai, + result.config.credential_pool.openai, + "OpenAI 凭证池往返不一致" + ); + prop_assert_eq!( + config.credential_pool.claude, + result.config.credential_pool.claude, + "Claude 凭证池往返不一致" + ); + } + + /// **Feature: config-credential-export, Property 8: Export-Import Round Trip** + /// *For any* valid configuration, exporting as a bundle (without redaction) and + /// then importing should produce an equivalent configuration. + /// **Validates: Requirements 5.5** + #[test] + fn prop_export_import_roundtrip_bundle(config in arb_config_with_credentials()) { + // 导出为 bundle(不脱敏,仅配置) + let options = ExportOptions { + include_config: true, + include_credentials: false, // 不包含 token 文件,因为测试环境没有实际文件 + redact_secrets: false, + }; + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 序列化为 JSON + let json = bundle.to_json().expect("序列化应成功"); + + // 反序列化 + let parsed_bundle = ExportBundle::from_json(&json).expect("反序列化应成功"); + + // 导入 bundle + let empty_config = Config::default(); + let import_options = ImportOptions::replace(); + let result = ImportService::import( + &parsed_bundle, + &empty_config, + &import_options, + &config.auth_dir, + ) + .expect("导入应成功"); + + // 验证往返一致性 + prop_assert_eq!( + config.server, + result.config.server, + "服务器配置往返不一致" + ); + prop_assert_eq!( + config.providers, + result.config.providers, + "Provider 配置往返不一致" + ); + prop_assert_eq!( + config.retry, + result.config.retry, + "重试配置往返不一致" + ); + prop_assert_eq!( + config.logging, + result.config.logging, + "日志配置往返不一致" + ); + } + + /// **Feature: config-credential-export, Property 8: Export-Import Round Trip** + /// *For any* configuration with API keys, exporting with redaction and then + /// importing should NOT restore the original API keys. + /// **Validates: Requirements 5.5** + #[test] + fn prop_export_import_redacted_loses_secrets(config in arb_config_with_secrets()) { + // 导出为 YAML(脱敏) + let yaml = ExportService::export_yaml(&config, true) + .expect("导出应成功"); + + // 导入 YAML + let empty_config = Config::default(); + let options = ImportOptions::replace(); + let result = ImportService::import_yaml(&yaml, &empty_config, &options) + .expect("导入应成功"); + + // 清理脱敏数据 + let mut imported = result.config; + ImportService::import( + &ExportBundle::new("1.0.0"), + &imported, + &ImportOptions::merge(), + &config.auth_dir, + ).ok(); // 忽略结果,只是为了触发清理 + + // 验证脱敏后的配置不包含原始敏感信息 + // 服务器 API 密钥应为脱敏占位符或默认值 + prop_assert!( + imported.server.api_key == REDACTED_PLACEHOLDER || + imported.server.api_key == "proxy_cast", + "脱敏后服务器 API 密钥应为占位符或默认值: {}", + imported.server.api_key + ); + + // 如果原始配置有 OpenAI API 密钥,导入后应为脱敏占位符 + if config.providers.openai.api_key.is_some() { + prop_assert_eq!( + imported.providers.openai.api_key, + Some(REDACTED_PLACEHOLDER.to_string()), + "脱敏后 OpenAI API 密钥应为占位符" + ); + } + } + + /// **Feature: config-credential-export, Property 8: Export-Import Round Trip** + /// *For any* configuration, the export bundle should be valid JSON that can + /// be parsed back. + /// **Validates: Requirements 5.5** + #[test] + fn prop_export_bundle_json_roundtrip(config in arb_config_with_credentials()) { + let options = ExportOptions { + include_config: true, + include_credentials: false, + redact_secrets: false, + }; + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 序列化为 JSON + let json = bundle.to_json().expect("序列化应成功"); + + // 反序列化 + let parsed = ExportBundle::from_json(&json).expect("反序列化应成功"); + + // 验证往返一致性 + prop_assert_eq!( + bundle.version, + parsed.version, + "版本往返不一致" + ); + prop_assert_eq!( + bundle.app_version, + parsed.app_version, + "应用版本往返不一致" + ); + prop_assert_eq!( + bundle.redacted, + parsed.redacted, + "脱敏状态往返不一致" + ); + prop_assert_eq!( + bundle.config_yaml, + parsed.config_yaml, + "配置 YAML 往返不一致" + ); + prop_assert_eq!( + bundle.token_files, + parsed.token_files, + "Token 文件往返不一致" + ); + } +} diff --git a/src-tauri/src/config/types.rs b/src-tauri/src/config/types.rs index 2df50e957..23a68cc9c 100644 --- a/src-tauri/src/config/types.rs +++ b/src-tauri/src/config/types.rs @@ -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, + /// Gemini 凭证列表(OAuth) + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub gemini: Vec, + /// Qwen 凭证列表(OAuth) + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub qwen: Vec, + /// OpenAI 凭证列表(API Key) + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub openai: Vec, + /// Claude 凭证列表(API Key) + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub claude: Vec, +} + +/// 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, + /// 是否禁用 + #[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] diff --git a/src-tauri/src/config/yaml.rs b/src-tauri/src/config/yaml.rs index 674a0d06d..b00fa58c8 100644 --- a/src-tauri/src/config/yaml.rs +++ b/src-tauri/src/config/yaml.rs @@ -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, +} + +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 { + yaml.lines() + .filter(|line| line.trim().starts_with('#')) + .map(|s| s.to_string()) + .collect() + } + + /// 从 YAML 行中提取注释 + fn extract_comments(lines: &[&str]) -> Vec { + let mut comments = Vec::new(); + let mut current_key_path: Vec = Vec::new(); + let mut indent_stack: Vec = 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 { + 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, + indent_stack: &mut Vec, + ) { + 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 { + let mut positions = HashMap::new(); + let mut current_key_path: Vec = Vec::new(); + let mut indent_stack: Vec = 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 { + let mut result_lines: Vec = new_lines.iter().map(|s| s.to_string()).collect(); + let mut insertions: Vec<(usize, String)> = Vec::new(); + let mut unmatched_comments: Vec = 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 = 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 = Vec::new(); + let mut indent_stack: Vec = 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 配置文件路径(向后兼容) diff --git a/src-tauri/src/credential/mod.rs b/src-tauri/src/credential/mod.rs index 03ebd150f..6adb7c2df 100644 --- a/src-tauri/src/credential/mod.rs +++ b/src-tauri/src/credential/mod.rs @@ -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)] diff --git a/src-tauri/src/credential/sync.rs b/src-tauri/src/credential/sync.rs new file mode 100644 index 000000000..2056419b2 --- /dev/null +++ b/src-tauri/src/credential/sync.rs @@ -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 for SyncError { + fn from(err: ConfigError) -> Self { + SyncError::ConfigError(err.to_string()) + } +} + +impl From for SyncError { + fn from(err: std::io::Error) -> Self { + SyncError::IoError(err.to_string()) + } +} + +/// 凭证同步服务 +/// +/// 负责将凭证池变更同步到 YAML 配置文件 +pub struct CredentialSyncService { + /// 配置管理器 + config_manager: Arc>, +} + +impl CredentialSyncService { + /// 创建新的凭证同步服务 + pub fn new(config_manager: Arc>) -> Self { + Self { config_manager } + } + + /// 获取当前配置 + fn get_config(&self) -> Result { + 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 { + let config = self.get_config()?; + Ok(expand_tilde(&config.auth_dir)) + } + + /// 确保 auth_dir 目录存在 + pub fn ensure_auth_dir(&self) -> Result { + 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 { + 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)` - 加载的凭证列表 + /// * `Err(SyncError)` - 加载失败 + pub fn load_from_config(&self) -> Result, 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 { + 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 { + 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) + } +} diff --git a/src-tauri/src/credential/tests.rs b/src-tauri/src/credential/tests.rs index c07b1ed6b..404931aed 100644 --- a/src-tauri/src/credential/tests.rs +++ b/src-tauri/src/credential/tests.rs @@ -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>) { + 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 { + 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 { + 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 文件内容应该已更新" + ); + } +} diff --git a/src-tauri/src/injection/types.rs b/src-tauri/src/injection/types.rs index eb0255ff2..300ef203e 100644 --- a/src-tauri/src/injection/types.rs +++ b/src-tauri/src/injection/types.rs @@ -184,6 +184,11 @@ impl Injector { self.rules.iter().filter(|r| r.matches(model)).collect() } + /// 清空所有规则 + pub fn clear(&mut self) { + self.rules.clear(); + } + /// 注入参数到请求 /// /// 按规则优先级顺序应用注入: diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 8653243d8..62eec7a7c 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -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, diff --git a/src-tauri/src/processor/context.rs b/src-tauri/src/processor/context.rs new file mode 100644 index 000000000..30f15b871 --- /dev/null +++ b/src-tauri/src/processor/context.rs @@ -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, + /// 原始模型名称(请求中的模型) + pub original_model: String, + /// 解析后的模型名称(经过别名映射) + pub resolved_model: String, + /// 选择的 Provider + pub provider: Option, + /// 使用的凭证 ID + pub credential_id: Option, + /// 重试次数 + pub retry_count: u32, + /// 是否为流式请求 + pub is_stream: bool, + /// 插件上下文 + pub plugin_ctx: Option, + /// 元数据 + pub metadata: std::collections::HashMap, +} + +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")); + } +} diff --git a/src-tauri/src/processor/error.rs b/src-tauri/src/processor/error.rs new file mode 100644 index 000000000..7b5cf461b --- /dev/null +++ b/src-tauri/src/processor/error.rs @@ -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); + } +} diff --git a/src-tauri/src/processor/mod.rs b/src-tauri/src/processor/mod.rs new file mode 100644 index 000000000..8ec3dfaa7 --- /dev/null +++ b/src-tauri/src/processor/mod.rs @@ -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>, + /// 模型映射器 + pub mapper: Arc>, + /// 参数注入器 + pub injector: Arc>, + /// 重试器 + pub retrier: Arc, + /// 故障转移器 + pub failover: Arc, + /// 超时控制器 + pub timeout: Arc, + /// 插件管理器 + pub plugins: Arc, + /// 统计聚合器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享) + pub stats: Arc>, + /// Token 追踪器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享) + pub tokens: Arc>, + /// 凭证池服务 + pub pool_service: Arc, +} + +impl RequestProcessor { + /// 创建新的请求处理器 + pub fn new( + router: Arc>, + mapper: Arc>, + injector: Arc>, + retrier: Arc, + failover: Arc, + timeout: Arc, + plugins: Arc, + stats: Arc>, + tokens: Arc>, + pool_service: Arc, + ) -> Self { + Self { + router, + mapper, + injector, + retrier, + failover, + timeout, + plugins, + stats, + tokens, + pool_service, + } + } + + /// 使用默认配置创建请求处理器 + pub fn with_defaults(pool_service: Arc) -> 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, + stats: Arc>, + tokens: Arc>, + ) -> 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; diff --git a/src-tauri/src/processor/steps/auth.rs b/src-tauri/src/processor/steps/auth.rs new file mode 100644 index 000000000..d67559b79 --- /dev/null +++ b/src-tauri/src/processor/steps/auth.rs @@ -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()); + } +} diff --git a/src-tauri/src/processor/steps/injection.rs b/src-tauri/src/processor/steps/injection.rs new file mode 100644 index 000000000..a26b57b93 --- /dev/null +++ b/src-tauri/src/processor/steps/injection.rs @@ -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>, + /// 是否启用 + enabled: Arc>, +} + +impl InjectionStep { + /// 创建新的注入步骤 + pub fn new(injector: Arc>) -> Self { + Self { + injector, + enabled: Arc::new(RwLock::new(true)), + } + } + + /// 设置是否启用 + pub fn with_enabled(self, enabled: Arc>) -> 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()); + } +} diff --git a/src-tauri/src/processor/steps/mod.rs b/src-tauri/src/processor/steps/mod.rs new file mode 100644 index 000000000..7ff3bd813 --- /dev/null +++ b/src-tauri/src/processor/steps/mod.rs @@ -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}; diff --git a/src-tauri/src/processor/steps/plugin.rs b/src-tauri/src/processor/steps/plugin.rs new file mode 100644 index 000000000..b1b30ab1b --- /dev/null +++ b/src-tauri/src/processor/steps/plugin.rs @@ -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, +} + +impl PluginPreStep { + /// 创建新的插件前置步骤 + pub fn new(plugins: Arc) -> 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::>()), + ); + } + + Ok(()) + } + + fn name(&self) -> &str { + "plugin_pre" + } +} + +/// 插件后置钩子步骤 +/// +/// 在 Provider 调用后执行所有启用插件的 on_response 钩子 +pub struct PluginPostStep { + /// 插件管理器 + plugins: Arc, +} + +impl PluginPostStep { + /// 创建新的插件后置步骤 + pub fn new(plugins: Arc) -> 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::>()), + ); + } + + 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()); + } +} diff --git a/src-tauri/src/processor/steps/provider.rs b/src-tauri/src/processor/steps/provider.rs new file mode 100644 index 000000000..e2a52f70a --- /dev/null +++ b/src-tauri/src/processor/steps/provider.rs @@ -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, +} + +/// Provider 调用错误 +#[derive(Debug, Clone)] +pub struct ProviderCallError { + /// 错误消息 + pub message: String, + /// HTTP 状态码(如果有) + pub status_code: Option, + /// 是否可重试 + pub retryable: bool, + /// 是否应触发故障转移 + pub should_failover: bool, +} + +impl ProviderCallError { + /// 创建可重试错误 + pub fn retryable(message: impl Into, status_code: Option) -> Self { + Self { + message: message.into(), + status_code, + retryable: true, + should_failover: false, + } + } + + /// 创建需要故障转移的错误 + pub fn failover(message: impl Into, status_code: Option) -> Self { + Self { + message: message.into(), + status_code, + retryable: false, + should_failover: true, + } + } + + /// 创建不可恢复错误 + pub fn fatal(message: impl Into, status_code: Option) -> 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, + /// 故障转移器 + failover: Arc, + /// 超时控制器 + timeout: Arc, + /// 凭证池服务 + pool_service: Arc, +} + +impl ProviderStep { + /// 创建新的 Provider 步骤 + pub fn new( + retrier: Arc, + failover: Arc, + timeout: Arc, + pool_service: Arc, + ) -> Self { + Self { + retrier, + failover, + timeout, + pool_service, + } + } + + /// 使用默认配置创建 + pub fn with_defaults(pool_service: Arc) -> 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, + ) -> 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( + &self, + ctx: &mut RequestContext, + mut operation: F, + ) -> Result + where + F: FnMut() -> Fut, + Fut: Future>, + { + 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( + &self, + ctx: &RequestContext, + operation: F, + ) -> Result + where + F: Future>, + { + 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 { + 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( + &self, + ctx: &mut RequestContext, + mut operation_factory: F, + available_providers: &[ProviderType], + ) -> Result + where + F: FnMut(ProviderType) -> Fut, + Fut: Future>, + { + 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 = 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); + } +} diff --git a/src-tauri/src/processor/steps/routing.rs b/src-tauri/src/processor/steps/routing.rs new file mode 100644 index 000000000..6e665df03 --- /dev/null +++ b/src-tauri/src/processor/steps/routing.rs @@ -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>, + /// 模型映射器 + mapper: Arc>, + /// 默认 Provider + default_provider: Arc>, +} + +impl RoutingStep { + /// 创建新的路由步骤 + pub fn new( + router: Arc>, + mapper: Arc>, + default_provider: Arc>, + ) -> 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 { + 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"); + } +} diff --git a/src-tauri/src/processor/steps/telemetry.rs b/src-tauri/src/processor/steps/telemetry.rs new file mode 100644 index 000000000..a51bd91d1 --- /dev/null +++ b/src-tauri/src/processor/steps/telemetry.rs @@ -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>, + /// Token 追踪器(使用 parking_lot::RwLock 以支持与 TelemetryState 共享) + tokens: Arc>, +} + +impl TelemetryStep { + /// 创建新的统计记录步骤 + pub fn new(stats: Arc>, tokens: Arc>) -> Self { + Self { stats, tokens } + } + + /// 记录请求日志 + /// + /// 请求完成后记录统计,按 Provider 和模型分组 + /// _需求: 4.1_ + pub fn record_request( + &self, + ctx: &RequestContext, + status: RequestStatus, + error_message: Option, + ) { + 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, + output_tokens: Option, + 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); + } +} diff --git a/src-tauri/src/processor/steps/traits.rs b/src-tauri/src/processor/steps/traits.rs new file mode 100644 index 000000000..8cb9f9d89 --- /dev/null +++ b/src-tauri/src/processor/steps/traits.rs @@ -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 + } +} diff --git a/src-tauri/src/processor/tests.rs b/src-tauri/src/processor/tests.rs new file mode 100644 index 000000000..f484f230b --- /dev/null +++ b/src-tauri/src/processor/tests.rs @@ -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 { + prop_oneof![ + Just(ProviderType::Kiro), + Just(ProviderType::Gemini), + Just(ProviderType::Qwen), + Just(ProviderType::OpenAI), + Just(ProviderType::Claude), + ] +} + +/// 生成随机的 RequestStatus +fn arb_request_status() -> impl Strategy { + prop_oneof![ + Just(RequestStatus::Success), + Just(RequestStatus::Failed), + Just(RequestStatus::Timeout), + Just(RequestStatus::Cancelled), + ] +} + +/// 生成随机的模型名称 +fn arb_model_name() -> impl Strategy { + 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 { + ( + "[a-zA-Z0-9_-]{8,16}", // id + arb_provider_type(), + arb_model_name(), + any::(), // 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 = 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 { + "[a-z0-9]{8,16}".prop_map(|s| s.to_string()) +} + +/// 生成随机的失败次数 +fn arb_failure_count() -> impl Strategy { + 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 = 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 { + 1u32..10000u32 +} + +/// 生成随机的 OpenAI 格式响应 +fn arb_openai_response() -> impl Strategy { + (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 { + (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 数应一致" + ); + } +} diff --git a/src-tauri/src/router/rules.rs b/src-tauri/src/router/rules.rs index d297372ac..3b1a9178f 100644 --- a/src-tauri/src/router/rules.rs +++ b/src-tauri/src/router/rules.rs @@ -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 diff --git a/src-tauri/src/server.rs b/src-tauri/src/server.rs index c9be66aa3..80cc0991b 100644 --- a/src-tauri/src/server.rs +++ b/src-tauri/src/server.rs @@ -1,15 +1,21 @@ //! HTTP API 服务器 -use crate::config::Config; +use crate::config::{ + Config, ConfigChangeEvent, ConfigChangeKind, ConfigManager, FileWatcher, HotReloadManager, + ReloadResult, +}; use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; use crate::converter::openai_to_antigravity::{ convert_antigravity_to_openai_response, convert_openai_to_antigravity, }; +use crate::credential::CredentialSyncService; +use crate::database::dao::provider_pool::ProviderPoolDao; use crate::database::DbConnection; use crate::injection::Injector; use crate::logger::LogStore; use crate::models::anthropic::*; use crate::models::openai::*; use crate::models::route_model::{RouteInfo, RouteListResponse}; +use crate::processor::{RequestContext, RequestProcessor}; use crate::providers::antigravity::AntigravityProvider; use crate::providers::claude_custom::ClaudeCustomProvider; use crate::providers::gemini::GeminiProvider; @@ -18,6 +24,8 @@ use crate::providers::openai_custom::OpenAICustomProvider; use crate::providers::qwen::QwenProvider; use crate::services::provider_pool_service::ProviderPoolService; use crate::services::token_cache_service::TokenCacheService; +use crate::telemetry::{RequestLog, RequestStatus}; +use crate::websocket::{WsConfig, WsConnectionManager, WsStats}; use axum::{ body::Body, extract::{Path, State}, @@ -28,8 +36,9 @@ use axum::{ }; use futures::stream; use serde::{Deserialize, Serialize}; +use std::path::PathBuf; use std::sync::Arc; -use tokio::sync::{oneshot, RwLock}; +use tokio::sync::{mpsc, oneshot, RwLock}; /// 安全截断字符串到指定字符数,避免 UTF-8 边界问题 fn safe_truncate(s: &str, max_chars: usize) -> String { @@ -41,6 +50,124 @@ fn safe_truncate(s: &str, max_chars: usize) -> String { } } +/// 计算 MessageContent 的字符长度 +fn message_content_len(content: &crate::models::openai::MessageContent) -> usize { + use crate::models::openai::{ContentPart, MessageContent}; + match content { + MessageContent::Text(s) => s.len(), + MessageContent::Parts(parts) => parts + .iter() + .filter_map(|p| { + if let ContentPart::Text { text } = p { + Some(text.len()) + } else { + None + } + }) + .sum(), + } +} + +/// 记录请求统计到遥测系统 +fn record_request_telemetry( + state: &AppState, + ctx: &RequestContext, + status: crate::telemetry::RequestStatus, + error_message: Option, +) { + use crate::telemetry::RequestLog; + + let provider = ctx.provider.unwrap_or(crate::ProviderType::Kiro); + let mut log = RequestLog::new( + ctx.request_id.clone(), + provider, + ctx.resolved_model.clone(), + ctx.is_stream, + ); + + // 设置状态和持续时间 + match status { + crate::telemetry::RequestStatus::Success => log.mark_success(ctx.elapsed_ms(), 200), + crate::telemetry::RequestStatus::Failed => log.mark_failed( + ctx.elapsed_ms(), + None, + error_message.clone().unwrap_or_default(), + ), + crate::telemetry::RequestStatus::Timeout => log.mark_timeout(ctx.elapsed_ms()), + crate::telemetry::RequestStatus::Cancelled => log.mark_cancelled(ctx.elapsed_ms()), + crate::telemetry::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; + + // 记录到统计聚合器 + { + let stats = state.processor.stats.write(); + stats.record(log.clone()); + } + + // 记录到请求日志记录器(用于前端日志列表显示) + if let Some(logger) = &state.request_logger { + let _ = logger.record(log.clone()); + } + + tracing::info!( + "[TELEMETRY] request_id={} provider={:?} model={} status={:?} duration_ms={}", + ctx.request_id, + provider, + ctx.resolved_model, + status, + ctx.elapsed_ms() + ); +} + +/// 记录 Token 使用量到遥测系统 +fn record_token_usage( + state: &AppState, + ctx: &RequestContext, + input_tokens: Option, + output_tokens: Option, +) { + use crate::telemetry::{TokenSource, TokenUsageRecord}; + + // 只有当至少有一个 Token 值时才记录 + if input_tokens.is_none() && output_tokens.is_none() { + return; + } + + let provider = ctx.provider.unwrap_or(crate::ProviderType::Kiro); + 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), + TokenSource::Actual, + ) + .with_request_id(ctx.request_id.clone()); + + // 记录到 Token 追踪器 + { + let tokens = state.processor.tokens.write(); + tokens.record(record); + } + + tracing::debug!( + "[TOKEN] request_id={} input={} output={}", + ctx.request_id, + input_tokens.unwrap_or(0), + output_tokens.unwrap_or(0) + ); +} + #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ServerStatus { pub running: bool, @@ -104,6 +231,24 @@ impl ServerState { pool_service: Arc, token_cache: Arc, db: Option, + ) -> Result<(), Box> { + self.start_with_telemetry(logs, pool_service, token_cache, db, None, None, None) + .await + } + + /// 启动服务器(使用共享的遥测实例) + /// + /// 这允许服务器与 TelemetryState 共享同一个 StatsAggregator、TokenTracker 和 RequestLogger, + /// 使得请求处理过程中记录的统计数据能够在前端监控页面中显示。 + pub async fn start_with_telemetry( + &mut self, + logs: Arc>, + pool_service: Arc, + token_cache: Arc, + db: Option, + shared_stats: Option>>, + shared_tokens: Option>>, + shared_logger: Option>, ) -> Result<(), Box> { if self.running { return Ok(()); @@ -132,6 +277,10 @@ impl ServerState { .collect(), ); + // 获取配置和配置路径用于热重载 + let config = self.config.clone(); + let config_path = crate::config::ConfigManager::default_config_path(); + tokio::spawn(async move { if let Err(e) = run_server( &host, @@ -146,6 +295,11 @@ impl ServerState { db, injector, injection_enabled, + shared_stats, + shared_tokens, + shared_logger, + Some(config), + Some(config_path), ) .await { @@ -195,6 +349,275 @@ struct AppState { injector: Arc>, /// 是否启用参数注入 injection_enabled: Arc>, + /// 请求处理器 + processor: Arc, + /// WebSocket 连接管理器 + ws_manager: Arc, + /// WebSocket 统计信息 + ws_stats: Arc, + /// 热重载管理器 + hot_reload_manager: Option>, + /// 请求日志记录器(与 TelemetryState 共享) + request_logger: Option>, +} + +/// 启动配置文件监控 +/// +/// 监控配置文件变化并触发热重载。 +/// +/// # 连接保持 +/// +/// 热重载过程不会中断现有连接: +/// - 配置更新在独立的 tokio 任务中异步执行 +/// - 使用 RwLock 进行原子性更新,不会阻塞正在处理的请求 +/// - 服务器继续运行,不需要重启 +/// - HTTP 和 WebSocket 连接保持活跃 +async fn start_config_watcher( + config_path: PathBuf, + hot_reload_manager: Option>, + processor: Arc, + logs: Arc>, + db: Option, + config_manager: Option>>, +) -> Option { + let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::(); + + // 创建文件监控器 + let mut watcher = match FileWatcher::new(&config_path, tx) { + Ok(w) => w, + Err(e) => { + tracing::error!("[HOT_RELOAD] 创建文件监控器失败: {}", e); + return None; + } + }; + + // 启动监控 + if let Err(e) = watcher.start() { + tracing::error!("[HOT_RELOAD] 启动文件监控失败: {}", e); + return None; + } + + tracing::info!("[HOT_RELOAD] 配置文件监控已启动: {:?}", config_path); + + // 启动事件处理任务 + let hot_reload_manager_clone = hot_reload_manager.clone(); + let processor_clone = processor.clone(); + let logs_clone = logs.clone(); + let db_clone = db.clone(); + let config_manager_clone = config_manager.clone(); + + tokio::spawn(async move { + while let Some(event) = rx.recv().await { + // 只处理修改事件 + if event.kind != ConfigChangeKind::Modified { + continue; + } + + tracing::info!("[HOT_RELOAD] 检测到配置文件变更: {:?}", event.path); + logs_clone.write().await.add( + "info", + &format!("[HOT_RELOAD] 检测到配置文件变更: {:?}", event.path), + ); + + // 执行热重载 + if let Some(ref manager) = hot_reload_manager_clone { + let result = manager.reload(); + match &result { + ReloadResult::Success { .. } => { + tracing::info!("[HOT_RELOAD] 配置热重载成功"); + logs_clone + .write() + .await + .add("info", "[HOT_RELOAD] 配置热重载成功"); + + // 更新处理器中的组件 + let new_config = manager.config(); + update_processor_config(&processor_clone, &new_config).await; + + // 同步凭证池 + if let (Some(ref db), Some(ref cfg_manager)) = + (&db_clone, &config_manager_clone) + { + match sync_credential_pool_from_config(db, cfg_manager, &logs_clone) + .await + { + Ok(count) => { + tracing::info!( + "[HOT_RELOAD] 凭证池同步完成,共 {} 个凭证", + count + ); + logs_clone.write().await.add( + "info", + &format!( + "[HOT_RELOAD] 凭证池同步完成,共 {} 个凭证", + count + ), + ); + } + Err(e) => { + tracing::warn!("[HOT_RELOAD] 凭证池同步失败: {}", e); + logs_clone.write().await.add( + "warn", + &format!("[HOT_RELOAD] 凭证池同步失败: {}", e), + ); + } + } + } + } + ReloadResult::RolledBack { error, .. } => { + tracing::warn!("[HOT_RELOAD] 配置热重载失败,已回滚: {}", error); + logs_clone.write().await.add( + "warn", + &format!("[HOT_RELOAD] 配置热重载失败,已回滚: {}", error), + ); + } + ReloadResult::Failed { + error, + rollback_error, + .. + } => { + tracing::error!( + "[HOT_RELOAD] 配置热重载失败: {}, 回滚错误: {:?}", + error, + rollback_error + ); + logs_clone.write().await.add( + "error", + &format!( + "[HOT_RELOAD] 配置热重载失败: {}, 回滚错误: {:?}", + error, rollback_error + ), + ); + } + } + } + } + }); + + Some(watcher) +} + +/// 更新处理器配置 +/// +/// 当配置热重载成功后,更新 RequestProcessor 中的各个组件。 +/// +/// # 原子性更新 +/// +/// 每个组件的更新都是原子性的,使用 RwLock 确保: +/// - 正在处理的请求不会看到部分更新的状态 +/// - 更新过程不会阻塞新请求的处理 +/// - 现有连接不受影响 +async fn update_processor_config(processor: &RequestProcessor, config: &Config) { + // 更新注入器规则 + { + let mut injector = processor.injector.write().await; + injector.clear(); + for rule in &config.injection.rules { + injector.add_rule(rule.clone().into()); + } + tracing::debug!( + "[HOT_RELOAD] 注入器规则已更新: {} 条规则", + config.injection.rules.len() + ); + } + + // 更新路由器规则 + { + let mut router = processor.router.write().await; + router.clear_rules(); + for rule in &config.routing.rules { + // 解析 provider 字符串为 ProviderType + if let Ok(provider_type) = rule.provider.parse::() { + router.add_rule(crate::router::RoutingRule { + pattern: rule.pattern.clone(), + target_provider: provider_type, + priority: rule.priority, + enabled: true, + }); + } else { + tracing::warn!("[HOT_RELOAD] 无法解析 provider: {}", rule.provider); + } + } + tracing::debug!( + "[HOT_RELOAD] 路由规则已更新: {} 条规则", + config.routing.rules.len() + ); + } + + // 更新模型映射器 + { + let mut mapper = processor.mapper.write().await; + mapper.clear(); + for (alias, model) in &config.routing.model_aliases { + mapper.add_alias(alias, model); + } + tracing::debug!( + "[HOT_RELOAD] 模型别名已更新: {} 个别名", + config.routing.model_aliases.len() + ); + } + + // 注意:重试配置目前不支持热更新,因为 Retrier 是不可变的 + // 如果需要更新重试配置,需要重启服务器 + tracing::debug!( + "[HOT_RELOAD] 重试配置: max_retries={}, base_delay={}ms (需重启生效)", + config.retry.max_retries, + config.retry.base_delay_ms + ); + + tracing::info!("[HOT_RELOAD] 处理器配置更新完成"); +} + +/// 从配置同步凭证池 +/// +/// 当配置热重载成功后,从 YAML 配置中加载凭证并同步到数据库。 +/// +/// # 同步策略 +/// +/// - 从配置中加载所有凭证 +/// - 对于配置中存在但数据库中不存在的凭证,添加到数据库 +/// - 对于配置中存在且数据库中也存在的凭证,更新数据库中的记录 +/// - 对于数据库中存在但配置中不存在的凭证,保留(不删除,避免丢失运行时状态) +async fn sync_credential_pool_from_config( + db: &DbConnection, + config_manager: &Arc>, + _logs: &Arc>, +) -> Result { + // 创建凭证同步服务 + let sync_service = CredentialSyncService::new(config_manager.clone()); + + // 从配置加载凭证 + let credentials = sync_service.load_from_config().map_err(|e| e.to_string())?; + + let conn = db.lock().map_err(|e| e.to_string())?; + let mut synced_count = 0; + + for cred in &credentials { + // 检查凭证是否已存在 + let existing = + ProviderPoolDao::get_by_uuid(&conn, &cred.uuid).map_err(|e| e.to_string())?; + + if existing.is_some() { + // 更新现有凭证 + ProviderPoolDao::update(&conn, cred).map_err(|e| e.to_string())?; + tracing::debug!( + "[HOT_RELOAD] 更新凭证: {} ({})", + cred.uuid, + cred.provider_type + ); + } else { + // 添加新凭证 + ProviderPoolDao::insert(&conn, cred).map_err(|e| e.to_string())?; + tracing::debug!( + "[HOT_RELOAD] 添加凭证: {} ({})", + cred.uuid, + cred.provider_type + ); + } + synced_count += 1; + } + + Ok(synced_count) } async fn run_server( @@ -210,8 +633,53 @@ async fn run_server( db: Option, injector: Injector, injection_enabled: bool, + shared_stats: Option>>, + shared_tokens: Option>>, + shared_logger: Option>, + config: Option, + config_path: Option, ) -> Result<(), Box> { let base_url = format!("http://{}:{}", host, port); + + // 创建请求处理器(使用共享的遥测实例或默认实例) + let processor = match (shared_stats, shared_tokens) { + (Some(stats), Some(tokens)) => Arc::new(RequestProcessor::with_shared_telemetry( + pool_service.clone(), + stats, + tokens, + )), + _ => Arc::new(RequestProcessor::with_defaults(pool_service.clone())), + }; + + // 将注入器规则同步到处理器 + { + let mut proc_injector = processor.injector.write().await; + for rule in injector.rules() { + proc_injector.add_rule(rule.clone()); + } + } + + // 初始化 WebSocket 管理器 + let ws_manager = Arc::new(WsConnectionManager::new(WsConfig::default())); + let ws_stats = ws_manager.stats().clone(); + + // 初始化热重载管理器 + let hot_reload_manager = match (&config, &config_path) { + (Some(cfg), Some(path)) => Some(Arc::new(HotReloadManager::new(cfg.clone(), path.clone()))), + _ => None, + }; + + // 初始化配置管理器(用于凭证池同步) + let config_manager: Option>> = + match (&config, &config_path) { + (Some(cfg), Some(path)) => Some(Arc::new(std::sync::RwLock::new( + ConfigManager::with_config(cfg.clone(), path.clone()), + ))), + _ => None, + }; + + let logs_clone = logs.clone(); + let db_clone = db.clone(); let state = AppState { api_key: api_key.to_string(), base_url, @@ -226,6 +694,26 @@ async fn run_server( db, injector: Arc::new(RwLock::new(injector)), injection_enabled: Arc::new(RwLock::new(injection_enabled)), + processor: processor.clone(), + ws_manager, + ws_stats, + hot_reload_manager: hot_reload_manager.clone(), + request_logger: shared_logger, + }; + + // 启动配置文件监控 + let _file_watcher = if let Some(path) = config_path { + start_config_watcher( + path, + hot_reload_manager, + processor, + logs_clone, + db_clone, + config_manager, + ) + .await + } else { + None }; let app = Router::new() @@ -235,6 +723,9 @@ async fn run_server( .route("/v1/chat/completions", post(chat_completions)) .route("/v1/messages", post(anthropic_messages)) .route("/v1/messages/count_tokens", post(count_tokens)) + // WebSocket 路由 + .route("/v1/ws", get(ws_upgrade_handler)) + .route("/ws", get(ws_upgrade_handler)) // 多供应商路由 .route( "/:selector/v1/messages", @@ -263,7 +754,7 @@ async fn run_server( async fn health() -> impl IntoResponse { Json(serde_json::json!({ "status": "healthy", - "version": "0.9.0" + "version": "0.10.0" })) } @@ -333,26 +824,44 @@ async fn chat_completions( return e.into_response(); } + // 创建请求上下文 + let mut ctx = RequestContext::new(request.model.clone()).with_stream(request.stream); + state.logs.write().await.add( "info", &format!( - "POST /v1/chat/completions model={} stream={}", - request.model, request.stream + "POST /v1/chat/completions request_id={} model={} stream={}", + ctx.request_id, request.model, request.stream ), ); + // 使用 RequestProcessor 解析模型别名和路由 + let provider = state.processor.resolve_and_route(&mut ctx).await; + + // 更新请求中的模型名为解析后的模型 + if ctx.resolved_model != ctx.original_model { + request.model = ctx.resolved_model.clone(); + state.logs.write().await.add( + "info", + &format!( + "[MAPPER] request_id={} alias={} -> model={}", + ctx.request_id, ctx.original_model, ctx.resolved_model + ), + ); + } + // 应用参数注入 let injection_enabled = *state.injection_enabled.read().await; if injection_enabled { - let injector = state.injector.read().await; + let injector = state.processor.injector.read().await; let mut payload = serde_json::to_value(&request).unwrap_or_default(); let result = injector.inject(&request.model, &mut payload); if result.has_injections() { state.logs.write().await.add( "info", &format!( - "[INJECT] Applied rules: {:?}, injected params: {:?}", - result.applied_rules, result.injected_params + "[INJECT] request_id={} applied_rules={:?} injected_params={:?}", + ctx.request_id, result.applied_rules, result.injected_params ), ); // 更新请求 @@ -362,9 +871,18 @@ async fn chat_completions( } } - // 获取当前默认 provider + // 获取当前默认 provider(用于凭证池选择) let default_provider = state.default_provider.read().await.clone(); + // 记录路由结果 + state.logs.write().await.add( + "info", + &format!( + "[ROUTE] request_id={} model={} provider={}", + ctx.request_id, ctx.resolved_model, provider + ), + ); + // 尝试从凭证池中选择凭证 let credential = match &state.db { Some(db) => state @@ -386,7 +904,41 @@ async fn chat_completions( &cred.uuid[..8] ), ); - return call_provider_openai(&state, &cred, &request).await; + let response = call_provider_openai(&state, &cred, &request).await; + + // 记录请求统计 + let is_success = response.status().is_success(); + let status = if is_success { + crate::telemetry::RequestStatus::Success + } else { + crate::telemetry::RequestStatus::Failed + }; + record_request_telemetry(&state, &ctx, status, None); + + // 如果成功,记录估算的 Token 使用量 + if is_success { + let estimated_input_tokens = request + .messages + .iter() + .map(|m| { + let content_len = match &m.content { + Some(c) => message_content_len(c), + None => 0, + }; + content_len / 4 + }) + .sum::() as u32; + // 输出 Token 使用估算值(假设平均响应长度) + let estimated_output_tokens = 100u32; + record_token_usage( + &state, + &ctx, + Some(estimated_input_tokens), + Some(estimated_output_tokens), + ); + } + + return response; } // 回退到旧的单凭证模式 @@ -462,6 +1014,22 @@ async fn chat_completions( }) }; + // 估算 Token 数量(基于字符数,约 4 字符 = 1 token) + let estimated_output_tokens = (parsed.content.len() / 4) as u32; + // 估算输入 Token(基于请求消息) + let estimated_input_tokens = request + .messages + .iter() + .map(|m| { + let content_len = match &m.content { + Some(c) => message_content_len(c), + None => 0, + }; + content_len / 4 + }) + .sum::() + as u32; + let response = serde_json::json!({ "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), "object": "chat.completion", @@ -476,18 +1044,41 @@ async fn chat_completions( "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } }], "usage": { - "prompt_tokens": 0, - "completion_tokens": 0, - "total_tokens": 0 + "prompt_tokens": estimated_input_tokens, + "completion_tokens": estimated_output_tokens, + "total_tokens": estimated_input_tokens + estimated_output_tokens } }); + // 记录成功请求统计 + record_request_telemetry( + &state, + &ctx, + crate::telemetry::RequestStatus::Success, + None, + ); + // 记录 Token 使用量 + record_token_usage( + &state, + &ctx, + Some(estimated_input_tokens), + Some(estimated_output_tokens), + ); Json(response).into_response() } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), + Err(e) => { + // 记录失败请求统计 + record_request_telemetry( + &state, + &ctx, + crate::telemetry::RequestStatus::Failed, + Some(e.to_string()), + ); + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } } } else if status.as_u16() == 403 || status.as_u16() == 402 { // Token 过期或账户问题,尝试重新加载凭证并刷新 @@ -644,6 +1235,9 @@ async fn anthropic_messages( return e.into_response(); } + // 创建请求上下文 + let mut ctx = RequestContext::new(request.model.clone()).with_stream(request.stream); + // 详细记录请求信息 let msg_count = request.messages.len(); let has_tools = request.tools.as_ref().map(|t| t.len()).unwrap_or(0); @@ -651,11 +1245,26 @@ async fn anthropic_messages( state.logs.write().await.add( "info", &format!( - "[REQ] POST /v1/messages model={} stream={} messages={} tools={} has_system={}", - request.model, request.stream, msg_count, has_tools, has_system + "[REQ] POST /v1/messages request_id={} model={} stream={} messages={} tools={} has_system={}", + ctx.request_id, request.model, request.stream, msg_count, has_tools, has_system ), ); + // 使用 RequestProcessor 解析模型别名和路由 + let provider = state.processor.resolve_and_route(&mut ctx).await; + + // 更新请求中的模型名为解析后的模型 + if ctx.resolved_model != ctx.original_model { + request.model = ctx.resolved_model.clone(); + state.logs.write().await.add( + "info", + &format!( + "[MAPPER] request_id={} alias={} -> model={}", + ctx.request_id, ctx.original_model, ctx.resolved_model + ), + ); + } + // 记录最后一条消息的角色和内容预览 if let Some(last_msg) = request.messages.last() { let content_preview = match &last_msg.content { @@ -676,8 +1285,8 @@ async fn anthropic_messages( state.logs.write().await.add( "debug", &format!( - "[REQ] Last message: role={} content={}", - last_msg.role, content_preview + "[REQ] request_id={} last_message: role={} content={}", + ctx.request_id, last_msg.role, content_preview ), ); } @@ -685,15 +1294,15 @@ async fn anthropic_messages( // 应用参数注入 let injection_enabled = *state.injection_enabled.read().await; if injection_enabled { - let injector = state.injector.read().await; + let injector = state.processor.injector.read().await; let mut payload = serde_json::to_value(&request).unwrap_or_default(); let result = injector.inject(&request.model, &mut payload); if result.has_injections() { state.logs.write().await.add( "info", &format!( - "[INJECT] Applied rules: {:?}, injected params: {:?}", - result.applied_rules, result.injected_params + "[INJECT] request_id={} applied_rules={:?} injected_params={:?}", + ctx.request_id, result.applied_rules, result.injected_params ), ); // 更新请求 @@ -703,9 +1312,18 @@ async fn anthropic_messages( } } - // 获取当前默认 provider + // 获取当前默认 provider(用于凭证池选择) let default_provider = state.default_provider.read().await.clone(); + // 记录路由结果 + state.logs.write().await.add( + "info", + &format!( + "[ROUTE] request_id={} model={} provider={}", + ctx.request_id, ctx.resolved_model, provider + ), + ); + // 尝试从凭证池中选择凭证 let credential = match &state.db { Some(db) => { @@ -730,7 +1348,46 @@ async fn anthropic_messages( &cred.uuid[..8] ), ); - return call_provider_anthropic(&state, &cred, &request).await; + let response = call_provider_anthropic(&state, &cred, &request).await; + + // 记录请求统计 + let is_success = response.status().is_success(); + let status = if is_success { + crate::telemetry::RequestStatus::Success + } else { + crate::telemetry::RequestStatus::Failed + }; + record_request_telemetry(&state, &ctx, status, None); + + // 如果成功,记录估算的 Token 使用量 + if is_success { + let estimated_input_tokens = request + .messages + .iter() + .map(|m| { + let content_len = match &m.content { + serde_json::Value::String(s) => s.len(), + serde_json::Value::Array(arr) => arr + .iter() + .filter_map(|v| v.get("text").and_then(|t| t.as_str())) + .map(|s| s.len()) + .sum(), + _ => 0, + }; + content_len / 4 + }) + .sum::() as u32; + // 输出 Token 使用估算值 + let estimated_output_tokens = 100u32; + record_token_usage( + &state, + &ctx, + Some(estimated_input_tokens), + Some(estimated_output_tokens), + ); + } + + return response; } // 回退到旧的单凭证模式 @@ -1888,6 +2545,12 @@ async fn call_provider_anthropic( // 回退到从源文件加载 let mut kiro = KiroProvider::new(); if let Err(e) = kiro.load_credentials_from_path(creds_file_path).await { + // 记录凭证加载失败 + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&format!("Failed to load credentials: {}", e)), + ); return ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": format!("Failed to load Kiro credentials: {}", e)}})), @@ -1895,6 +2558,12 @@ async fn call_provider_anthropic( .into_response(); } if let Err(e) = kiro.refresh_token().await { + // 记录 Token 刷新失败 + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&format!("Token refresh failed: {}", e)), + ); return ( StatusCode::UNAUTHORIZED, Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})), @@ -1915,6 +2584,12 @@ async fn call_provider_anthropic( let resp = match kiro.call_api(&openai_request).await { Ok(r) => r, Err(e) => { + // 记录 API 调用失败 + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); return ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}})), @@ -1929,17 +2604,31 @@ async fn call_provider_anthropic( Ok(bytes) => { let body = String::from_utf8_lossy(&bytes).to_string(); let parsed = parse_cw_response(&body); + // 记录成功 + let _ = state.pool_service.mark_healthy( + db, + &credential.uuid, + Some(&request.model), + ); + let _ = state.pool_service.record_usage(db, &credential.uuid); if request.stream { build_anthropic_stream_response(&request.model, &parsed) } else { build_anthropic_response(&request.model, &parsed) } } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), + Err(e) => { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } } } else if status.as_u16() == 401 || status.as_u16() == 403 { // Token 过期,强制刷新并重试 @@ -1956,6 +2645,12 @@ async fn call_provider_anthropic( { Ok(t) => t, Err(e) => { + // 记录 Token 刷新失败 + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&format!("Token refresh failed: {}", e)), + ); return ( StatusCode::UNAUTHORIZED, Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})), @@ -1973,20 +2668,39 @@ async fn call_provider_anthropic( Ok(bytes) => { let body = String::from_utf8_lossy(&bytes).to_string(); let parsed = parse_cw_response(&body); + // 记录重试成功 + let _ = state.pool_service.mark_healthy( + db, + &credential.uuid, + Some(&request.model), + ); + let _ = state.pool_service.record_usage(db, &credential.uuid); if request.stream { build_anthropic_stream_response(&request.model, &parsed) } else { build_anthropic_response(&request.model, &parsed) } } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), + Err(e) => { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } } } else { let body = retry_resp.text().await.unwrap_or_default(); + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&format!("Retry failed: {}", body)), + ); ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": format!("Retry failed: {}", body)}})), @@ -1994,14 +2708,24 @@ async fn call_provider_anthropic( .into_response() } } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), + Err(e) => { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } } } else { let body = resp.text().await.unwrap_or_default(); + let _ = state + .pool_service + .mark_unhealthy(db, &credential.uuid, Some(&body)); ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": body}})), @@ -2034,6 +2758,14 @@ async fn call_provider_anthropic( .load_credentials_from_path(creds_file_path) .await { + // 记录凭证加载失败 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&format!("Failed to load credentials: {}", e)), + ); + } return ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": format!("Failed to load Antigravity credentials: {}", e)}})), @@ -2044,6 +2776,14 @@ async fn call_provider_anthropic( // 检查并刷新 token if antigravity.is_token_expiring_soon() { if let Err(e) = antigravity.refresh_token().await { + // 记录 Token 刷新失败 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&format!("Token refresh failed: {}", e)), + ); + } return ( StatusCode::UNAUTHORIZED, Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})), @@ -2078,17 +2818,36 @@ async fn call_provider_anthropic( usage_credits: 0.0, context_usage_percentage: 0.0, }; + // 记录成功 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_healthy( + db, + &credential.uuid, + Some(&request.model), + ); + let _ = state.pool_service.record_usage(db, &credential.uuid); + } if request.stream { build_anthropic_stream_response(&request.model, &parsed) } else { build_anthropic_response(&request.model, &parsed) } } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), + Err(e) => { + // 记录 API 调用失败 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + } + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } } } CredentialData::OpenAIKey { api_key, base_url } => { @@ -2111,12 +2870,30 @@ async fn call_provider_anthropic( usage_credits: 0.0, context_usage_percentage: 0.0, }; + // 记录成功 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_healthy( + db, + &credential.uuid, + Some(&request.model), + ); + let _ = + state.pool_service.record_usage(db, &credential.uuid); + } if request.stream { build_anthropic_stream_response(&request.model, &parsed) } else { build_anthropic_response(&request.model, &parsed) } } else { + // 记录解析失败 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some("Failed to parse OpenAI response"), + ); + } ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": "Failed to parse OpenAI response"}})), @@ -2124,14 +2901,30 @@ async fn call_provider_anthropic( .into_response() } } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), + Err(e) => { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + } + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } } } else { let body = resp.text().await.unwrap_or_default(); + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&body), + ); + } ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": body}})), @@ -2139,11 +2932,20 @@ async fn call_provider_anthropic( .into_response() } } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), + Err(e) => { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + } + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } } } CredentialData::ClaudeKey { api_key, base_url } => { @@ -2154,6 +2956,15 @@ async fn call_provider_anthropic( match resp.text().await { Ok(body) => { if status.is_success() { + // 记录成功 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_healthy( + db, + &credential.uuid, + Some(&request.model), + ); + let _ = state.pool_service.record_usage(db, &credential.uuid); + } Response::builder() .status(StatusCode::OK) .header(header::CONTENT_TYPE, "application/json") @@ -2166,6 +2977,13 @@ async fn call_provider_anthropic( .into_response() }) } else { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&body), + ); + } ( StatusCode::from_u16(status.as_u16()) .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), @@ -2174,18 +2992,36 @@ async fn call_provider_anthropic( .into_response() } } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), + Err(e) => { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + } + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } } } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), + Err(e) => { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + } + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } } } } @@ -2193,14 +3029,20 @@ async fn call_provider_anthropic( /// 根据凭证调用 Provider (OpenAI 格式) async fn call_provider_openai( - _state: &AppState, + state: &AppState, credential: &ProviderCredential, request: &ChatCompletionRequest, ) -> Response { + let start_time = std::time::Instant::now(); + match &credential.credential { CredentialData::KiroOAuth { creds_file_path } => { let mut kiro = KiroProvider::new(); if let Err(e) = kiro.load_credentials_from_path(creds_file_path).await { + // 记录凭证加载失败 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&e.to_string())); + } return ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": format!("Failed to load Kiro credentials: {}", e)}})), @@ -2208,6 +3050,10 @@ async fn call_provider_openai( .into_response(); } if let Err(e) = kiro.refresh_token().await { + // 记录 Token 刷新失败 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&format!("Token refresh failed: {}", e))); + } return ( StatusCode::UNAUTHORIZED, Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})), @@ -2217,7 +3063,13 @@ async fn call_provider_openai( match kiro.call_api(request).await { Ok(resp) => { - if resp.status().is_success() { + let status = resp.status(); + if status.is_success() { + // 记录成功 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_healthy(db, &credential.uuid, Some(&request.model)); + let _ = state.pool_service.record_usage(db, &credential.uuid); + } match resp.text().await { Ok(body) => { let parsed = parse_cw_response(&body); @@ -2273,7 +3125,11 @@ async fn call_provider_openai( .into_response(), } } else { + // 记录 API 调用失败 let body = resp.text().await.unwrap_or_default(); + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&format!("HTTP {}: {}", status, safe_truncate(&body, 100)))); + } ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": body}})), @@ -2281,11 +3137,17 @@ async fn call_provider_openai( .into_response() } } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), + Err(e) => { + // 记录请求错误 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&e.to_string())); + } + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } } } CredentialData::GeminiOAuth { .. } => { @@ -2397,3 +3259,803 @@ async fn call_provider_openai( } } } + +// ========== WebSocket 处理 ========== + +use crate::websocket::{ + WsApiRequest, WsApiResponse, WsEndpoint, WsError, WsMessage as WsProtoMessage, +}; +use axum::extract::ws::{Message as WsMessage, WebSocket, WebSocketUpgrade}; +use futures::{SinkExt, StreamExt as FuturesStreamExt}; + +/// WebSocket 升级处理器 +async fn ws_upgrade_handler( + ws: WebSocketUpgrade, + State(state): State, + headers: HeaderMap, +) -> impl IntoResponse { + // 验证 API 密钥 + let auth = headers + .get("authorization") + .or_else(|| headers.get("x-api-key")) + .and_then(|v| v.to_str().ok()); + + let key = match auth { + Some(s) if s.starts_with("Bearer ") => &s[7..], + Some(s) => s, + None => { + return axum::http::Response::builder() + .status(401) + .body(Body::from("No API key provided")) + .unwrap() + .into_response(); + } + }; + + if key != state.api_key { + return axum::http::Response::builder() + .status(401) + .body(Body::from("Invalid API key")) + .unwrap() + .into_response(); + } + + // 获取客户端信息 + let client_info = headers + .get("user-agent") + .and_then(|v| v.to_str().ok()) + .map(|s| s.to_string()); + + ws.on_upgrade(move |socket| handle_websocket(socket, state, client_info)) +} + +/// 处理 WebSocket 连接 +async fn handle_websocket(socket: WebSocket, state: AppState, client_info: Option) { + let conn_id = uuid::Uuid::new_v4().to_string(); + + // 注册连接 + if let Err(e) = state + .ws_manager + .register(conn_id.clone(), client_info.clone()) + { + state.logs.write().await.add( + "error", + &format!("[WS] Failed to register connection: {}", e.message), + ); + return; + } + + state.logs.write().await.add( + "info", + &format!( + "[WS] New connection: {} (client: {:?})", + &conn_id[..8], + client_info + ), + ); + + let (mut sender, mut receiver) = socket.split(); + + // 消息处理循环 + while let Some(msg) = receiver.next().await { + match msg { + Ok(WsMessage::Text(text)) => { + state.ws_manager.on_message(); + state.ws_manager.increment_request_count(&conn_id); + + match serde_json::from_str::(&text) { + Ok(ws_msg) => { + let response = handle_ws_message(&state, &conn_id, ws_msg).await; + if let Some(resp) = response { + let resp_text = serde_json::to_string(&resp).unwrap_or_default(); + if sender + .send(WsMessage::Text(resp_text.into())) + .await + .is_err() + { + break; + } + } + } + Err(e) => { + state.ws_manager.on_error(); + let error = WsProtoMessage::Error(WsError::invalid_message(format!( + "Failed to parse message: {}", + e + ))); + let error_text = serde_json::to_string(&error).unwrap_or_default(); + if sender + .send(WsMessage::Text(error_text.into())) + .await + .is_err() + { + break; + } + } + } + } + Ok(WsMessage::Binary(_)) => { + state.ws_manager.on_error(); + let error = WsProtoMessage::Error(WsError::invalid_message( + "Binary messages not supported", + )); + let error_text = serde_json::to_string(&error).unwrap_or_default(); + if sender + .send(WsMessage::Text(error_text.into())) + .await + .is_err() + { + break; + } + } + Ok(WsMessage::Ping(data)) => { + if sender.send(WsMessage::Pong(data)).await.is_err() { + break; + } + } + Ok(WsMessage::Pong(_)) => { + // 收到 pong,连接正常 + } + Ok(WsMessage::Close(_)) => { + break; + } + Err(e) => { + state.logs.write().await.add( + "error", + &format!("[WS] Connection {} error: {}", &conn_id[..8], e), + ); + break; + } + } + } + + // 清理连接 + state.ws_manager.unregister(&conn_id); + state.logs.write().await.add( + "info", + &format!("[WS] Connection closed: {}", &conn_id[..8]), + ); +} + +/// 处理 WebSocket 消息 +async fn handle_ws_message( + state: &AppState, + conn_id: &str, + msg: WsProtoMessage, +) -> Option { + match msg { + WsProtoMessage::Ping { timestamp } => Some(WsProtoMessage::Pong { timestamp }), + WsProtoMessage::Pong { .. } => None, + WsProtoMessage::Request(request) => { + state.logs.write().await.add( + "info", + &format!( + "[WS] Request from {}: id={} endpoint={:?}", + &conn_id[..8], + request.request_id, + request.endpoint + ), + ); + + // 处理 API 请求 + let response = handle_ws_api_request(state, &request).await; + Some(response) + } + WsProtoMessage::Response(_) + | WsProtoMessage::StreamChunk(_) + | WsProtoMessage::StreamEnd(_) => Some(WsProtoMessage::Error(WsError::invalid_request( + None, + "Invalid message type from client", + ))), + WsProtoMessage::Error(_) => None, + } +} + +/// 处理 WebSocket API 请求 +async fn handle_ws_api_request(state: &AppState, request: &WsApiRequest) -> WsProtoMessage { + match request.endpoint { + WsEndpoint::Models => { + // 返回模型列表 + let models = serde_json::json!({ + "object": "list", + "data": [ + {"id": "claude-sonnet-4-5", "object": "model", "owned_by": "anthropic"}, + {"id": "claude-sonnet-4-5-20250929", "object": "model", "owned_by": "anthropic"}, + {"id": "claude-3-7-sonnet-20250219", "object": "model", "owned_by": "anthropic"}, + {"id": "gemini-2.5-flash", "object": "model", "owned_by": "google"}, + {"id": "gemini-2.5-pro", "object": "model", "owned_by": "google"}, + {"id": "qwen3-coder-plus", "object": "model", "owned_by": "alibaba"}, + ] + }); + WsProtoMessage::Response(WsApiResponse { + request_id: request.request_id.clone(), + payload: models, + }) + } + WsEndpoint::ChatCompletions => { + // 解析 ChatCompletionRequest + match serde_json::from_value::(request.payload.clone()) { + Ok(chat_request) => { + handle_ws_chat_completions(state, &request.request_id, chat_request).await + } + Err(e) => WsProtoMessage::Error(WsError::invalid_request( + Some(request.request_id.clone()), + format!("Invalid chat completion request: {}", e), + )), + } + } + WsEndpoint::Messages => { + // 解析 AnthropicMessagesRequest + match serde_json::from_value::(request.payload.clone()) { + Ok(messages_request) => { + handle_ws_anthropic_messages(state, &request.request_id, messages_request).await + } + Err(e) => WsProtoMessage::Error(WsError::invalid_request( + Some(request.request_id.clone()), + format!("Invalid messages request: {}", e), + )), + } + } + } +} + +/// 处理 WebSocket chat completions 请求 +async fn handle_ws_chat_completions( + state: &AppState, + request_id: &str, + mut request: ChatCompletionRequest, +) -> WsProtoMessage { + // 创建请求上下文 + let mut ctx = RequestContext::new(request.model.clone()).with_stream(request.stream); + + // 使用 RequestProcessor 解析模型别名和路由 + let _provider = state.processor.resolve_and_route(&mut ctx).await; + + // 更新请求中的模型名为解析后的模型 + if ctx.resolved_model != ctx.original_model { + request.model = ctx.resolved_model.clone(); + } + + // 应用参数注入 + let injection_enabled = *state.injection_enabled.read().await; + if injection_enabled { + let injector = state.processor.injector.read().await; + let mut payload = serde_json::to_value(&request).unwrap_or_default(); + let result = injector.inject(&request.model, &mut payload); + if result.has_injections() { + if let Ok(updated) = serde_json::from_value(payload) { + request = updated; + } + } + } + + // 获取默认 provider + let default_provider = state.default_provider.read().await.clone(); + + // 尝试从凭证池中选择凭证 + let credential = match &state.db { + Some(db) => state + .pool_service + .select_credential(db, &default_provider, Some(&request.model)) + .ok() + .flatten(), + None => None, + }; + + // 如果找到凭证,使用它调用 API + if let Some(cred) = credential { + // 简化实现:直接调用 provider 并返回结果 + // 实际实现应该复用 call_provider_openai 的逻辑 + match call_provider_openai_for_ws(state, &cred, &request).await { + Ok(response) => WsProtoMessage::Response(WsApiResponse { + request_id: request_id.to_string(), + payload: response, + }), + Err(e) => WsProtoMessage::Error(WsError::upstream(Some(request_id.to_string()), e)), + } + } else { + // 回退到 Kiro provider + let kiro = state.kiro.read().await; + match kiro.call_api(&request).await { + Ok(resp) => { + if resp.status().is_success() { + match resp.text().await { + Ok(body) => { + let parsed = parse_cw_response(&body); + let has_tool_calls = !parsed.tool_calls.is_empty(); + + let message = if has_tool_calls { + serde_json::json!({ + "role": "assistant", + "content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) }, + "tool_calls": parsed.tool_calls.iter().map(|tc| { + serde_json::json!({ + "id": tc.id, + "type": "function", + "function": { + "name": tc.function.name, + "arguments": tc.function.arguments + } + }) + }).collect::>() + }) + } else { + serde_json::json!({ + "role": "assistant", + "content": parsed.content + }) + }; + + let response = serde_json::json!({ + "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), + "object": "chat.completion", + "created": std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + "model": request.model, + "choices": [{ + "index": 0, + "message": message, + "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } + }], + "usage": { + "prompt_tokens": 0, + "completion_tokens": 0, + "total_tokens": 0 + } + }); + + WsProtoMessage::Response(WsApiResponse { + request_id: request_id.to_string(), + payload: response, + }) + } + Err(e) => WsProtoMessage::Error(WsError::internal( + Some(request_id.to_string()), + e.to_string(), + )), + } + } else { + let body = resp.text().await.unwrap_or_default(); + WsProtoMessage::Error(WsError::upstream( + Some(request_id.to_string()), + format!("Upstream error: {}", body), + )) + } + } + Err(e) => WsProtoMessage::Error(WsError::internal( + Some(request_id.to_string()), + e.to_string(), + )), + } + } +} + +/// 处理 WebSocket anthropic messages 请求 +async fn handle_ws_anthropic_messages( + state: &AppState, + request_id: &str, + mut request: AnthropicMessagesRequest, +) -> WsProtoMessage { + // 创建请求上下文 + let mut ctx = RequestContext::new(request.model.clone()).with_stream(request.stream); + + // 使用 RequestProcessor 解析模型别名和路由 + let _provider = state.processor.resolve_and_route(&mut ctx).await; + + // 更新请求中的模型名为解析后的模型 + if ctx.resolved_model != ctx.original_model { + request.model = ctx.resolved_model.clone(); + } + + // 应用参数注入 + let injection_enabled = *state.injection_enabled.read().await; + if injection_enabled { + let injector = state.processor.injector.read().await; + let mut payload = serde_json::to_value(&request).unwrap_or_default(); + let result = injector.inject(&request.model, &mut payload); + if result.has_injections() { + if let Ok(updated) = serde_json::from_value(payload) { + request = updated; + } + } + } + + // 获取默认 provider + let default_provider = state.default_provider.read().await.clone(); + + // 尝试从凭证池中选择凭证 + let credential = match &state.db { + Some(db) => state + .pool_service + .select_credential(db, &default_provider, Some(&request.model)) + .ok() + .flatten(), + None => None, + }; + + // 如果找到凭证,使用它调用 API + if let Some(cred) = credential { + match call_provider_anthropic_for_ws(state, &cred, &request).await { + Ok(response) => WsProtoMessage::Response(WsApiResponse { + request_id: request_id.to_string(), + payload: response, + }), + Err(e) => WsProtoMessage::Error(WsError::upstream(Some(request_id.to_string()), e)), + } + } else { + // 回退到 Kiro provider + let kiro = state.kiro.read().await; + + // 转换为 OpenAI 格式 + let openai_request = convert_anthropic_to_openai(&request); + + match kiro.call_api(&openai_request).await { + Ok(resp) => { + if resp.status().is_success() { + match resp.text().await { + Ok(body) => { + let parsed = parse_cw_response(&body); + + // 转换为 Anthropic 格式响应 + let response = serde_json::json!({ + "id": format!("msg_{}", uuid::Uuid::new_v4()), + "type": "message", + "role": "assistant", + "content": [{ + "type": "text", + "text": parsed.content + }], + "model": request.model, + "stop_reason": "end_turn", + "usage": { + "input_tokens": 0, + "output_tokens": 0 + } + }); + + WsProtoMessage::Response(WsApiResponse { + request_id: request_id.to_string(), + payload: response, + }) + } + Err(e) => WsProtoMessage::Error(WsError::internal( + Some(request_id.to_string()), + e.to_string(), + )), + } + } else { + let body = resp.text().await.unwrap_or_default(); + WsProtoMessage::Error(WsError::upstream( + Some(request_id.to_string()), + format!("Upstream error: {}", body), + )) + } + } + Err(e) => WsProtoMessage::Error(WsError::internal( + Some(request_id.to_string()), + e.to_string(), + )), + } + } +} + +/// WebSocket 专用的 OpenAI 格式 Provider 调用 +async fn call_provider_openai_for_ws( + state: &AppState, + credential: &ProviderCredential, + request: &ChatCompletionRequest, +) -> Result { + use crate::models::provider_pool_model::CredentialData; + + match &credential.credential { + CredentialData::KiroOAuth { creds_file_path } => { + let mut kiro = KiroProvider::new(); + if let Err(e) = kiro.load_credentials_from_path(creds_file_path).await { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&format!("Failed to load credentials: {}", e)), + ); + } + return Err(e.to_string()); + } + if let Err(e) = kiro.refresh_token().await { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&format!("Token refresh failed: {}", e)), + ); + } + return Err(e.to_string()); + } + + let resp = match kiro.call_api(request).await { + Ok(r) => r, + Err(e) => { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + } + return Err(e.to_string()); + } + }; + if resp.status().is_success() { + let body = resp.text().await.map_err(|e| e.to_string())?; + let parsed = parse_cw_response(&body); + let has_tool_calls = !parsed.tool_calls.is_empty(); + + // 记录成功 + if let Some(db) = &state.db { + let _ = + state + .pool_service + .mark_healthy(db, &credential.uuid, Some(&request.model)); + let _ = state.pool_service.record_usage(db, &credential.uuid); + } + + let message = if has_tool_calls { + serde_json::json!({ + "role": "assistant", + "content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) }, + "tool_calls": parsed.tool_calls.iter().map(|tc| { + serde_json::json!({ + "id": tc.id, + "type": "function", + "function": { + "name": tc.function.name, + "arguments": tc.function.arguments + } + }) + }).collect::>() + }) + } else { + serde_json::json!({ + "role": "assistant", + "content": parsed.content + }) + }; + + Ok(serde_json::json!({ + "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), + "object": "chat.completion", + "created": std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_secs(), + "model": request.model, + "choices": [{ + "index": 0, + "message": message, + "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } + }], + "usage": { + "prompt_tokens": 0, + "completion_tokens": 0, + "total_tokens": 0 + } + })) + } else { + let body = resp.text().await.unwrap_or_default(); + if let Some(db) = &state.db { + let _ = state + .pool_service + .mark_unhealthy(db, &credential.uuid, Some(&body)); + } + Err(format!("Upstream error: {}", body)) + } + } + CredentialData::OpenAIKey { api_key, base_url } => { + let provider = OpenAICustomProvider::with_config(api_key.clone(), base_url.clone()); + let resp = match provider.call_api(request).await { + Ok(r) => r, + Err(e) => { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + } + return Err(e.to_string()); + } + }; + if resp.status().is_success() { + // 记录成功 + if let Some(db) = &state.db { + let _ = + state + .pool_service + .mark_healthy(db, &credential.uuid, Some(&request.model)); + let _ = state.pool_service.record_usage(db, &credential.uuid); + } + resp.json::() + .await + .map_err(|e| e.to_string()) + } else { + let body = resp.text().await.unwrap_or_default(); + if let Some(db) = &state.db { + let _ = state + .pool_service + .mark_unhealthy(db, &credential.uuid, Some(&body)); + } + Err(format!("Upstream error: {}", body)) + } + } + CredentialData::ClaudeKey { api_key, base_url } => { + let provider = ClaudeCustomProvider::with_config(api_key.clone(), base_url.clone()); + match provider.call_openai_api(request).await { + Ok(result) => { + // 记录成功 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_healthy( + db, + &credential.uuid, + Some(&request.model), + ); + let _ = state.pool_service.record_usage(db, &credential.uuid); + } + Ok(result) + } + Err(e) => { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + } + Err(e.to_string()) + } + } + } + CredentialData::AntigravityOAuth { + creds_file_path, .. + } => { + let mut antigravity = AntigravityProvider::new(); + if let Err(e) = antigravity + .load_credentials_from_path(creds_file_path) + .await + { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&format!("Failed to load credentials: {}", e)), + ); + } + return Err(e.to_string()); + } + if !antigravity.is_token_valid() { + if let Err(e) = antigravity.refresh_token().await { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&format!("Token refresh failed: {}", e)), + ); + } + return Err(e.to_string()); + } + } + let antigravity_request = convert_openai_to_antigravity(request); + match antigravity + .call_api("generateContent", &antigravity_request) + .await + { + Ok(resp) => { + // 记录成功 + if let Some(db) = &state.db { + let _ = state.pool_service.mark_healthy( + db, + &credential.uuid, + Some(&request.model), + ); + let _ = state.pool_service.record_usage(db, &credential.uuid); + } + Ok(convert_antigravity_to_openai_response( + &resp, + &request.model, + )) + } + Err(e) => { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + } + Err(e.to_string()) + } + } + } + // GeminiOAuth 和 QwenOAuth 暂不支持 WebSocket,需要使用 HTTP 端点 + _ => Err( + "This credential type is not yet supported via WebSocket. Please use HTTP endpoints." + .to_string(), + ), + } +} + +/// WebSocket 专用的 Anthropic 格式 Provider 调用 +async fn call_provider_anthropic_for_ws( + state: &AppState, + credential: &ProviderCredential, + request: &AnthropicMessagesRequest, +) -> Result { + use crate::models::provider_pool_model::CredentialData; + + match &credential.credential { + CredentialData::ClaudeKey { api_key, base_url } => { + let provider = ClaudeCustomProvider::with_config(api_key.clone(), base_url.clone()); + let resp = match provider.call_api(request).await { + Ok(r) => r, + Err(e) => { + if let Some(db) = &state.db { + let _ = state.pool_service.mark_unhealthy( + db, + &credential.uuid, + Some(&e.to_string()), + ); + } + return Err(e.to_string()); + } + }; + if resp.status().is_success() { + // 记录成功 + if let Some(db) = &state.db { + let _ = + state + .pool_service + .mark_healthy(db, &credential.uuid, Some(&request.model)); + let _ = state.pool_service.record_usage(db, &credential.uuid); + } + resp.json::() + .await + .map_err(|e| e.to_string()) + } else { + let body = resp.text().await.unwrap_or_default(); + if let Some(db) = &state.db { + let _ = state + .pool_service + .mark_unhealthy(db, &credential.uuid, Some(&body)); + } + Err(format!("Upstream error: {}", body)) + } + } + _ => { + // 转换为 OpenAI 格式并调用(健康状态更新在 call_provider_openai_for_ws 中处理) + let openai_request = convert_anthropic_to_openai(request); + let result = call_provider_openai_for_ws(state, credential, &openai_request).await?; + + // 转换响应为 Anthropic 格式 + Ok(serde_json::json!({ + "id": format!("msg_{}", uuid::Uuid::new_v4()), + "type": "message", + "role": "assistant", + "content": [{ + "type": "text", + "text": result.get("choices") + .and_then(|c| c.get(0)) + .and_then(|c| c.get("message")) + .and_then(|m| m.get("content")) + .and_then(|c| c.as_str()) + .unwrap_or("") + }], + "model": request.model, + "stop_reason": "end_turn", + "usage": { + "input_tokens": 0, + "output_tokens": 0 + } + })) + } + } +} diff --git a/src-tauri/src/websocket/tests.rs b/src-tauri/src/websocket/tests.rs index 8034ee0a4..7916d533b 100644 --- a/src-tauri/src/websocket/tests.rs +++ b/src-tauri/src/websocket/tests.rs @@ -460,3 +460,169 @@ proptest! { } } } + +// ============ SSE 到 WebSocket 转换属性测试 ============ + +use super::stream::StreamForwarder; + +/// 生成任意的 SSE 数据内容(非空且不全是空格) +fn arb_sse_data() -> impl Strategy { + 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 { + arb_sse_data().prop_map(|data| format!("data: {}", data)) +} + +/// 生成任意的 SSE 响应体(多行) +fn arb_sse_body() -> impl Strategy, String)> { + prop::collection::vec(arb_sse_data(), 1..10).prop_map(|data_items| { + let lines: Vec = data_items.clone(); + let body = data_items + .iter() + .map(|d| format!("data: {}\n\n", d)) + .collect::>() + .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()); + } +} diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index dbd7f5eda..e531b9d54 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -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", diff --git a/src/components/config/AuthDirSettings.tsx b/src/components/config/AuthDirSettings.tsx new file mode 100644 index 000000000..e482068dc --- /dev/null +++ b/src/components/config/AuthDirSettings.tsx @@ -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(null); + const [expandedPath, setExpandedPath] = useState(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("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 ( +
+
+

+ + 认证目录设置 +

+

+ 配置 OAuth Token 文件的存储目录。支持使用 ~ 表示用户主目录。 +

+
+ +
+
+ +
+
+ + 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} + /> +
+ + +
+
+ + {/* Expanded path preview */} + {expandedPath && ( +
+ 展开后路径: + + {expandedPath} + +
+ )} + + {/* Error display */} + {error && ( +
+ + {error} +
+ )} + + {/* Success message */} + {saveSuccess && ( +
+ + 设置已保存 +
+ )} + + {/* Save button */} +
+ +
+
+ + {/* Help text */} +
+

说明

+
    +
  • 认证目录用于存储 OAuth Token 文件(Kiro、Gemini、Qwen 等)
  • +
  • + 使用 ~{" "} + 表示用户主目录,例如{" "} + ~/.proxycast/auth +
  • +
  • 修改此设置后,现有的 Token 文件不会自动迁移,需要手动移动
  • +
  • 导出配置时,Token 文件会从此目录读取并包含在导出包中
  • +
+
+
+ ); +} diff --git a/src/components/config/ConfigPage.tsx b/src/components/config/ConfigPage.tsx index c27fc5640..50dc32ada 100644 --- a/src/components/config/ConfigPage.tsx +++ b/src/components/config/ConfigPage.tsx @@ -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((_props, ref) => { const [activeTab, setActiveTab] = useState("editor"); @@ -62,9 +68,10 @@ export const ConfigPage = forwardRef((_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: }, ]; if (isLoading) { @@ -135,12 +142,13 @@ export const ConfigPage = forwardRef((_props, ref) => { ))} @@ -154,6 +162,12 @@ export const ConfigPage = forwardRef((_props, ref) => { {activeTab === "import-export" && ( )} + {activeTab === "settings" && ( + + )} ); diff --git a/src/components/config/ImportExport.tsx b/src/components/config/ImportExport.tsx index b7e314bf8..460628da8 100644 --- a/src/components/config/ImportExport.tsx +++ b/src/components/config/ImportExport.tsx @@ -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("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(null); - const [error, setError] = useState(null); - const [redactSecrets, setRedactSecrets] = useState(true); + const [validationResult, setValidationResult] = + useState(null); const [mergeConfig, setMergeConfig] = useState(true); + const [error, setError] = useState(null); + const fileInputRef = useRef(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) => { + const handleFileSelect = async (e: React.ChangeEvent) => { 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 ; } + return ; }; return ( @@ -115,6 +198,47 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) {
+ {/* Export Scope Selection */} +
+ +
+ + + +
+
+ + {/* Redaction Option */} + {/* Security hint */} + {!redactSecrets && + (exportScope === "credentials" || exportScope === "full") && ( +
+ + + 未脱敏的导出文件将包含明文 API 密钥和 Token,请妥善保管。 + +
+ )} +

- 导出当前配置为 YAML 文件,可用于备份或迁移到其他设备。 + {exportScope === "config" && + "导出当前配置为 YAML 文件,可用于备份或迁移。"} + {exportScope === "credentials" && + "导出凭证池中的所有凭证,包括 OAuth Token 文件。"} + {exportScope === "full" && + "导出完整的配置和凭证包,可用于完整迁移到其他设备。"}

@@ -154,7 +294,7 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) { @@ -168,55 +308,180 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) {

- 从 YAML 文件导入配置,支持合并或替换现有配置。 + 支持导入 YAML 配置文件或 JSON 导出包,支持合并或替换现有配置。

+ {/* Security Warning Dialog */} + {showSecurityWarning && ( +
+
+
+
+ +
+

安全警告

+
+ +

+ 您即将导出未脱敏的凭证数据,导出文件将包含明文 API 密钥和 OAuth + Token。 请确保: +

+
    +
  • 不要将此文件分享给他人
  • +
  • 不要上传到公共代码仓库
  • +
  • 妥善保管导出文件
  • +
+ +
+ + +
+
+
+ )} + {/* Import Dialog */} {showImportDialog && (
-
-

导入配置

+
+

+ {getFileTypeIcon()} + 导入配置 - {importFileName} +

+ + {/* Validation Result */} + {validationResult && ( +
+ {validationResult.valid ? ( +
+ + 文件格式有效 +
+ ) : ( +
+ +
+

文件格式无效

+
    + {validationResult.errors.map((err, i) => ( +
  • {err}
  • + ))} +
+
+
+ )} + + {/* Content info */} + {validationResult.valid && ( +
+ {validationResult.version && ( + + 版本: {validationResult.version} + + )} + {validationResult.has_config && ( + + 包含配置 + + )} + {validationResult.has_credentials && ( + + 包含凭证 + + )} + {validationResult.redacted && ( + + 已脱敏 + + )} +
+ )} + + {/* Redaction warning */} + {validationResult.redacted && ( +
+ + + 此导出包已脱敏,凭证数据(API 密钥、Token)无法恢复。 + +
+ )} + + {/* Validation warnings */} + {validationResult.warnings.length > 0 && ( +
+

+ 警告 +

+
    + {validationResult.warnings.map((warning, i) => ( +
  • {warning}
  • + ))} +
+
+ )} +
+ )} {/* Preview */}
- -
+              
+              
                 {importContent.slice(0, 2000)}
                 {importContent.length > 2000 && "\n..."}
               
- {/* Options */} + {/* Import Options */}
- - + +
+ + +
- {/* Warnings */} + {/* Import Result Warnings */} {importResult?.warnings && importResult.warnings.length > 0 && (

- 警告 + 导入警告

    {importResult.warnings.map((warning, i) => ( @@ -245,32 +510,19 @@ export function ImportExport({ config, onConfigImported }: ImportExportProps) { {/* Actions */}
    - {!importResult?.success && ( - <> - - - + {!importResult?.success && validationResult?.valid && ( + )}
diff --git a/src/components/config/index.ts b/src/components/config/index.ts index 33cc7d4fe..a559acfcb 100644 --- a/src/components/config/index.ts +++ b/src/components/config/index.ts @@ -1,3 +1,4 @@ export { ConfigPage } from "./ConfigPage"; export { ConfigEditor } from "./ConfigEditor"; export { ImportExport } from "./ImportExport"; +export { AuthDirSettings } from "./AuthDirSettings"; diff --git a/src/components/provider-pool/ProviderPoolPage.tsx b/src/components/provider-pool/ProviderPoolPage.tsx index 8f94c8b9a..db9fb69ad 100644 --- a/src/components/provider-pool/ProviderPoolPage.tsx +++ b/src/components/provider-pool/ProviderPoolPage.tsx @@ -87,7 +87,8 @@ export const ProviderPoolPage = forwardRef( 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 { diff --git a/src/components/routing/RoutingPage.tsx b/src/components/routing/RoutingPage.tsx index 2f2c7cf69..a65e771da 100644 --- a/src/components/routing/RoutingPage.tsx +++ b/src/components/routing/RoutingPage.tsx @@ -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((_props, ref) => { const [injectionRules, setInjectionRules] = useState([]); const [injectionEnabled, setInjectionEnabled] = useState(false); + // Presets state + const [presets, setPresets] = useState([]); + const [showPresets, setShowPresets] = useState(false); + const [applyingPreset, setApplyingPreset] = useState(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((_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((_props, ref) => { 配置模型映射、路由规则和排除列表

- +
+ + + +
+ {/* Presets Modal */} + {showPresets && ( +
+
+
+

+ + 推荐配置 +

+ +
+

+ 选择一个预设配置快速设置路由规则和模型别名 +

+
+ {presets.map((preset) => ( +
+
+
+

{preset.name}

+

+ {preset.description} +

+
+ {preset.aliases.length} 个别名 + {preset.rules.length} 条规则 +
+
+
+ + +
+
+
+ ))} +
+
+
+ )} +
  • diff --git a/src/hooks/useProviderPool.ts b/src/hooks/useProviderPool.ts index 0dd8883e4..bbe5f0da4 100644 --- a/src/hooks/useProviderPool.ts +++ b/src/hooks/useProviderPool.ts @@ -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(); }; diff --git a/src/lib/api/config.ts b/src/lib/api/config.ts index 6e49f5bf7..6f3185705 100644 --- a/src/lib/api/config.ts +++ b/src/lib/api/config.ts @@ -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 { + return invoke("export_bundle", { config, options }); + }, + + // Validate import content (JSON bundle or YAML config) + async validateImport(content: string): Promise { + return invoke("validate_import", { content }); + }, + // Validate YAML config async validateConfigYaml(yamlContent: string): Promise { 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 { + return invoke("import_bundle", { currentConfig, content, merge }); + }, + // Get config file paths async getConfigPaths(): Promise { return invoke("get_config_paths"); diff --git a/src/lib/api/providerPool.ts b/src/lib/api/providerPool.ts index 966abc34c..0dac322ab 100644 --- a/src/lib/api/providerPool.ts +++ b/src/lib/api/providerPool.ts @@ -195,8 +195,11 @@ export const providerPoolApi = { }, // Delete a credential - async deleteCredential(uuid: string): Promise { - return invoke("delete_provider_pool_credential", { uuid }); + async deleteCredential( + uuid: string, + providerType?: PoolProviderType, + ): Promise { + return invoke("delete_provider_pool_credential", { uuid, providerType }); }, // Toggle credential enabled/disabled diff --git a/src/lib/api/router.ts b/src/lib/api/router.ts index c4e82f3ba..cd33a6ea3 100644 --- a/src/lib/api/router.ts +++ b/src/lib/api/router.ts @@ -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 { return invoke("set_router_default_provider", { provider }); }, + + // Recommended presets + async getRecommendedPresets(): Promise { + return invoke("get_recommended_presets"); + }, + + async applyRecommendedPreset( + presetId: string, + merge: boolean = false, + ): Promise { + return invoke("apply_recommended_preset", { presetId, merge }); + }, + + async clearAllRoutingConfig(): Promise { + return invoke("clear_all_routing_config"); + }, };