diff --git a/package.json b/package.json index c50034dbf..e023effeb 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.17.1", + "version": "0.17.2", "type": "module", "repository": { "type": "git", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index c5e66f3f5..2999bad60 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -3377,7 +3377,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.17.1" +version = "0.17.2" dependencies = [ "anyhow", "async-stream", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 62e5322e5..ade297e3d 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "proxycast" -version = "0.17.1" +version = "0.17.2" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" diff --git a/src-tauri/src/commands/router_cmd.rs b/src-tauri/src/commands/router_cmd.rs index 2eec53078..ecfe0a3f2 100644 --- a/src-tauri/src/commands/router_cmd.rs +++ b/src-tauri/src/commands/router_cmd.rs @@ -186,6 +186,26 @@ pub struct RecommendedPreset { pub description: String, pub aliases: Vec, pub rules: Vec, + /// 客户端路由配置 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub endpoint_providers: Option, +} + +/// 端点 Provider 配置 DTO +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct EndpointProvidersConfigDto { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cursor: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub claude_code: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub codex: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub windsurf: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub kiro: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub other: Option, } /// 获取推荐配置列表 @@ -260,6 +280,7 @@ pub async fn get_recommended_presets() -> Result, String> enabled: true, }, ], + endpoint_providers: None, }, RecommendedPreset { id: "gemini-optimized".to_string(), @@ -308,6 +329,7 @@ pub async fn get_recommended_presets() -> Result, String> enabled: true, }, ], + endpoint_providers: None, }, RecommendedPreset { id: "multi-provider".to_string(), @@ -398,6 +420,7 @@ pub async fn get_recommended_presets() -> Result, String> enabled: true, }, ], + endpoint_providers: None, }, RecommendedPreset { id: "coding-assistant".to_string(), @@ -443,6 +466,7 @@ pub async fn get_recommended_presets() -> Result, String> enabled: true, }, ], + endpoint_providers: None, }, RecommendedPreset { id: "cost-effective".to_string(), @@ -478,6 +502,24 @@ pub async fn get_recommended_presets() -> Result, String> enabled: true, }, ], + endpoint_providers: None, + }, + // 客户端路由预设 + RecommendedPreset { + id: "client-routing".to_string(), + name: "客户端路由配置".to_string(), + description: "为不同的 IDE 客户端配置不同的 Provider,Cursor/Windsurf 使用 Kiro,Claude Code 使用 Kiro,Codex 使用 OpenAI" + .to_string(), + aliases: vec![], + rules: vec![], + endpoint_providers: Some(EndpointProvidersConfigDto { + cursor: Some("kiro".to_string()), + claude_code: Some("kiro".to_string()), + codex: Some("openai".to_string()), + windsurf: Some("kiro".to_string()), + kiro: Some("kiro".to_string()), + other: None, + }), }, ]) } diff --git a/src-tauri/src/config/hot_reload.rs b/src-tauri/src/config/hot_reload.rs index e224bca15..f95c02488 100644 --- a/src-tauri/src/config/hot_reload.rs +++ b/src-tauri/src/config/hot_reload.rs @@ -5,7 +5,7 @@ //! - 支持原子性配置更新 //! - 失败时自动回滚到之前的配置 -use super::types::Config; +use super::types::{is_default_api_key, Config}; use super::yaml::ConfigManager; use notify::{Event, EventKind, RecommendedWatcher, RecursiveMode, Watcher}; use parking_lot::RwLock; @@ -414,7 +414,7 @@ impl HotReloadManager { } if (!is_localhost || config.remote_management.allow_remote) - && crate::config::is_default_api_key(&config.server.api_key) + && is_default_api_key(&config.server.api_key) { return Err(HotReloadError::ValidationError( "非本地访问场景下禁止使用默认 API Key,请设置强口令".to_string(), diff --git a/src-tauri/src/config/mod.rs b/src-tauri/src/config/mod.rs index b5f7e691f..860580481 100644 --- a/src-tauri/src/config/mod.rs +++ b/src-tauri/src/config/mod.rs @@ -17,11 +17,11 @@ pub use hot_reload::{ pub use import::{ImportOptions, ImportService, ValidationResult}; pub use path_utils::{collapse_tilde, contains_tilde, expand_tilde}; pub use types::{ - AmpConfig, AmpModelMapping, ApiKeyEntry, Config, CredentialEntry, CredentialPoolConfig, - CustomProviderConfig, GeminiApiKeyEntry, IFlowCredentialEntry, InjectionRuleConfig, - InjectionSettings, LoggingConfig, ProviderConfig, ProvidersConfig, QuotaExceededConfig, - RemoteManagementConfig, RetrySettings, RoutingConfig, ServerConfig, TlsConfig, - VertexApiKeyEntry, VertexModelAlias, + generate_secure_api_key, AmpConfig, AmpModelMapping, ApiKeyEntry, Config, CredentialEntry, + CredentialPoolConfig, CustomProviderConfig, EndpointProvidersConfig, GeminiApiKeyEntry, + IFlowCredentialEntry, InjectionRuleConfig, InjectionSettings, LoggingConfig, ProviderConfig, + ProvidersConfig, QuotaExceededConfig, RemoteManagementConfig, RetrySettings, RoutingConfig, + ServerConfig, TlsConfig, VertexApiKeyEntry, VertexModelAlias, DEFAULT_API_KEY, }; pub use yaml::{load_config, save_config, ConfigError, ConfigManager, YamlService}; diff --git a/src-tauri/src/config/tests.rs b/src-tauri/src/config/tests.rs index 5d6a6c9ac..ab7227b4a 100644 --- a/src-tauri/src/config/tests.rs +++ b/src-tauri/src/config/tests.rs @@ -217,6 +217,7 @@ fn arb_config() -> impl Strategy { quota_exceeded: crate::config::QuotaExceededConfig::default(), proxy_url: None, ampcode: crate::config::AmpConfig::default(), + endpoint_providers: crate::config::EndpointProvidersConfig::default(), }) } @@ -488,6 +489,7 @@ fn arb_valid_config() -> impl Strategy { quota_exceeded: crate::config::QuotaExceededConfig::default(), proxy_url: None, ampcode: crate::config::AmpConfig::default(), + endpoint_providers: crate::config::EndpointProvidersConfig::default(), }) } @@ -531,6 +533,7 @@ fn arb_invalid_config() -> impl Strategy { quota_exceeded: crate::config::QuotaExceededConfig::default(), proxy_url: None, ampcode: crate::config::AmpConfig::default(), + endpoint_providers: crate::config::EndpointProvidersConfig::default(), }; // 根据类型使配置无效 match invalid_type { @@ -2470,3 +2473,351 @@ proptest! { } } } + +// ============================================================================ +// Property 3: EndpointProvidersConfig 序列化往返一致性 +// ============================================================================ + +use crate::config::EndpointProvidersConfig; + +/// 生成随机的 Provider 名称 +fn arb_provider_name() -> impl Strategy { + prop_oneof![ + Just("kiro".to_string()), + Just("gemini".to_string()), + Just("qwen".to_string()), + Just("openai".to_string()), + Just("claude".to_string()), + Just("codex".to_string()), + Just("iflow".to_string()), + ] +} + +/// 生成随机的可选 Provider 名称 +fn arb_optional_provider() -> impl Strategy> { + proptest::option::of(arb_provider_name()) +} + +/// 生成随机的 EndpointProvidersConfig +fn arb_endpoint_providers_config() -> impl Strategy { + ( + arb_optional_provider(), + arb_optional_provider(), + arb_optional_provider(), + arb_optional_provider(), + arb_optional_provider(), + arb_optional_provider(), + ) + .prop_map(|(cursor, claude_code, codex, windsurf, kiro, other)| { + EndpointProvidersConfig { + cursor, + claude_code, + codex, + windsurf, + kiro, + other, + } + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: endpoint-provider-config, Property 3: 配置序列化往返一致性** + /// *对于任意* 有效的 EndpointProvidersConfig 对象,序列化后再反序列化应产生等价的对象。 + /// **Validates: Requirements 1.1** + #[test] + fn prop_endpoint_providers_config_yaml_roundtrip(config in arb_endpoint_providers_config()) { + // 序列化为 YAML + let yaml = serde_yaml::to_string(&config) + .expect("YAML 序列化应成功"); + + // 反序列化回 EndpointProvidersConfig + let parsed: EndpointProvidersConfig = serde_yaml::from_str(&yaml) + .expect("YAML 反序列化应成功"); + + // 验证往返一致性 + prop_assert_eq!( + config.cursor, + parsed.cursor, + "cursor 字段往返不一致" + ); + prop_assert_eq!( + config.claude_code, + parsed.claude_code, + "claude_code 字段往返不一致" + ); + prop_assert_eq!( + config.codex, + parsed.codex, + "codex 字段往返不一致" + ); + prop_assert_eq!( + config.windsurf, + parsed.windsurf, + "windsurf 字段往返不一致" + ); + prop_assert_eq!( + config.kiro, + parsed.kiro, + "kiro 字段往返不一致" + ); + prop_assert_eq!( + config.other, + parsed.other, + "other 字段往返不一致" + ); + } + + /// **Feature: endpoint-provider-config, Property 3: 配置序列化往返一致性(JSON)** + /// *对于任意* 有效的 EndpointProvidersConfig 对象,JSON 序列化后再反序列化应产生等价的对象。 + /// **Validates: Requirements 1.1** + #[test] + fn prop_endpoint_providers_config_json_roundtrip(config in arb_endpoint_providers_config()) { + // 序列化为 JSON + let json = serde_json::to_string(&config) + .expect("JSON 序列化应成功"); + + // 反序列化回 EndpointProvidersConfig + let parsed: EndpointProvidersConfig = serde_json::from_str(&json) + .expect("JSON 反序列化应成功"); + + // 验证往返一致性 + prop_assert_eq!( + config, + parsed, + "EndpointProvidersConfig JSON 往返不一致" + ); + } + + /// **Feature: endpoint-provider-config, Property 3: 配置序列化往返一致性(完整配置)** + /// *对于任意* 包含 EndpointProvidersConfig 的完整配置,序列化后再反序列化应保持 endpoint_providers 一致。 + /// **Validates: Requirements 1.1** + #[test] + fn prop_config_with_endpoint_providers_roundtrip( + endpoint_providers in arb_endpoint_providers_config() + ) { + // 创建包含 endpoint_providers 的完整配置 + let config = Config { + endpoint_providers: endpoint_providers.clone(), + ..Config::default() + }; + + // 序列化为 YAML + let yaml = ConfigManager::to_yaml(&config) + .expect("序列化应成功"); + + // 反序列化回 Config + let parsed = ConfigManager::parse_yaml(&yaml) + .expect("反序列化应成功"); + + // 验证 endpoint_providers 往返一致性 + prop_assert_eq!( + endpoint_providers, + parsed.endpoint_providers, + "endpoint_providers 往返不一致" + ); + } +} + +// ============================================================================ +// Property 4: Provider 类型验证 +// ============================================================================ + +use crate::ProviderType; + +/// 生成有效的 Provider 类型字符串 +fn arb_valid_provider_type() -> impl Strategy { + prop_oneof![ + Just("kiro".to_string()), + Just("gemini".to_string()), + Just("qwen".to_string()), + Just("openai".to_string()), + Just("claude".to_string()), + Just("antigravity".to_string()), + Just("vertex".to_string()), + Just("gemini_api_key".to_string()), + Just("codex".to_string()), + Just("claude_oauth".to_string()), + Just("iflow".to_string()), + ] +} + +/// 生成无效的 Provider 类型字符串 +fn arb_invalid_provider_type() -> impl Strategy { + // 生成不在有效列表中的字符串 + "[a-z]{3,15}".prop_filter("排除有效的 Provider 类型", |s| { + !matches!( + s.as_str(), + "kiro" + | "gemini" + | "qwen" + | "openai" + | "claude" + | "antigravity" + | "vertex" + | "gemini_api_key" + | "codex" + | "claude_oauth" + | "iflow" + ) + }) +} + +/// 生成有效的客户端类型字符串 +fn arb_valid_client_type() -> impl Strategy { + prop_oneof![ + Just("cursor".to_string()), + Just("claude_code".to_string()), + Just("codex".to_string()), + Just("windsurf".to_string()), + Just("kiro".to_string()), + Just("other".to_string()), + ] +} + +/// 生成无效的客户端类型字符串 +fn arb_invalid_client_type() -> impl Strategy { + // 生成不在有效列表中的字符串 + "[a-z]{3,15}".prop_filter("排除有效的客户端类型", |s| { + !matches!( + s.as_str(), + "cursor" | "claude_code" | "codex" | "windsurf" | "kiro" | "other" + ) + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: endpoint-provider-config, Property 4: Provider 类型验证** + /// *对于任意* 有效的 Provider 类型字符串,解析应成功并返回正确的 ProviderType。 + /// **Validates: Requirements 5.1, 5.2** + #[test] + fn prop_valid_provider_type_parsing(provider in arb_valid_provider_type()) { + // 解析 Provider 类型 + let result: Result = provider.parse(); + + // 验证解析成功 + prop_assert!( + result.is_ok(), + "有效的 Provider 类型应解析成功: {}", + provider + ); + + // 验证往返一致性 + let parsed = result.unwrap(); + prop_assert_eq!( + parsed.to_string(), + provider, + "Provider 类型往返不一致" + ); + } + + /// **Feature: endpoint-provider-config, Property 4: Provider 类型验证** + /// *对于任意* 无效的 Provider 类型字符串,解析应失败并返回描述性错误消息。 + /// **Validates: Requirements 5.1, 5.2** + #[test] + fn prop_invalid_provider_type_parsing(provider in arb_invalid_provider_type()) { + // 解析 Provider 类型 + let result: Result = provider.parse(); + + // 验证解析失败 + prop_assert!( + result.is_err(), + "无效的 Provider 类型应解析失败: {}", + provider + ); + + // 验证错误消息包含描述性信息 + let error = result.unwrap_err(); + prop_assert!( + error.contains("Invalid provider") || error.contains(&provider), + "错误消息应包含描述性信息: {}", + error + ); + } + + /// **Feature: endpoint-provider-config, Property 4: Provider 类型验证** + /// *对于任意* 有效的客户端类型,set_provider 应成功设置 Provider。 + /// **Validates: Requirements 5.1, 5.2** + #[test] + fn prop_valid_client_type_set_provider( + client_type in arb_valid_client_type(), + provider in arb_valid_provider_type() + ) { + let mut config = EndpointProvidersConfig::default(); + + // 设置 Provider + let result = config.set_provider(&client_type, Some(provider.clone())); + + // 验证设置成功 + prop_assert!( + result, + "有效的客户端类型应设置成功: {}", + client_type + ); + + // 验证 Provider 已正确设置 + let stored = config.get_provider(&client_type); + prop_assert_eq!( + stored, + Some(&provider), + "Provider 应正确存储" + ); + } + + /// **Feature: endpoint-provider-config, Property 4: Provider 类型验证** + /// *对于任意* 无效的客户端类型,set_provider 应返回 false。 + /// **Validates: Requirements 5.1, 5.2** + #[test] + fn prop_invalid_client_type_set_provider( + client_type in arb_invalid_client_type(), + provider in arb_valid_provider_type() + ) { + let mut config = EndpointProvidersConfig::default(); + + // 设置 Provider + let result = config.set_provider(&client_type, Some(provider)); + + // 验证设置失败 + prop_assert!( + !result, + "无效的客户端类型应设置失败: {}", + client_type + ); + } + + /// **Feature: endpoint-provider-config, Property 4: Provider 类型验证** + /// *对于任意* 有效的客户端类型,使用 None 或空字符串应清除 Provider 配置。 + /// **Validates: Requirements 5.1, 5.2** + #[test] + fn prop_clear_provider_config( + client_type in arb_valid_client_type(), + provider in arb_valid_provider_type() + ) { + let mut config = EndpointProvidersConfig::default(); + + // 先设置 Provider + config.set_provider(&client_type, Some(provider)); + + // 使用 None 清除 + let result = config.set_provider(&client_type, None); + prop_assert!(result, "清除操作应成功"); + prop_assert_eq!( + config.get_provider(&client_type), + None, + "Provider 应被清除" + ); + + // 重新设置后使用空字符串清除 + config.set_provider(&client_type, Some("kiro".to_string())); + let result = config.set_provider(&client_type, Some("".to_string())); + prop_assert!(result, "空字符串清除操作应成功"); + prop_assert_eq!( + config.get_provider(&client_type), + None, + "Provider 应被清除(空字符串)" + ); + } +} diff --git a/src-tauri/src/config/types.rs b/src-tauri/src/config/types.rs index 5e636855e..92c8a1a4c 100644 --- a/src-tauri/src/config/types.rs +++ b/src-tauri/src/config/types.rs @@ -166,6 +166,98 @@ fn default_auth_dir() -> String { "~/.proxycast/auth".to_string() } +/// 端点 Provider 配置 +/// +/// 允许为不同的客户端端点配置不同的 Provider +/// 例如:Cursor 使用 Qwen,Claude Code 使用 Kiro,Codex 使用 Codex +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)] +pub struct EndpointProvidersConfig { + /// Cursor 客户端使用的 Provider + /// 如果为空,则使用 default_provider + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cursor: Option, + /// Claude Code 客户端使用的 Provider + /// 如果为空,则使用 default_provider + #[serde(default, skip_serializing_if = "Option::is_none")] + pub claude_code: Option, + /// Codex 客户端使用的 Provider + /// 如果为空,则使用 default_provider + #[serde(default, skip_serializing_if = "Option::is_none")] + pub codex: Option, + /// Windsurf 客户端使用的 Provider + /// 如果为空,则使用 default_provider + #[serde(default, skip_serializing_if = "Option::is_none")] + pub windsurf: Option, + /// Kiro 客户端使用的 Provider + /// 如果为空,则使用 default_provider + #[serde(default, skip_serializing_if = "Option::is_none")] + pub kiro: Option, + /// 其他客户端使用的 Provider + /// 如果为空,则使用 default_provider + #[serde(default, skip_serializing_if = "Option::is_none")] + pub other: Option, +} + +impl EndpointProvidersConfig { + /// 根据客户端类型获取配置的 Provider + /// + /// # 参数 + /// - `client_type`: 客户端类型的配置键名(cursor, claude_code, codex, windsurf, kiro, other) + /// + /// # 返回 + /// 如果配置了对应的 Provider,返回 Some(&String);否则返回 None + pub fn get_provider(&self, client_type: &str) -> Option<&String> { + match client_type { + "cursor" => self.cursor.as_ref(), + "claude_code" => self.claude_code.as_ref(), + "codex" => self.codex.as_ref(), + "windsurf" => self.windsurf.as_ref(), + "kiro" => self.kiro.as_ref(), + "other" => self.other.as_ref(), + _ => None, + } + } + + /// 设置客户端类型的 Provider + /// + /// # 参数 + /// - `client_type`: 客户端类型的配置键名(cursor, claude_code, codex, windsurf, kiro, other) + /// - `provider`: 要设置的 Provider 名称,None 或空字符串表示清除配置 + /// + /// # 返回 + /// 如果客户端类型有效,返回 true;否则返回 false + pub fn set_provider(&mut self, client_type: &str, provider: Option) -> bool { + let provider = provider.filter(|p| !p.is_empty()); + match client_type { + "cursor" => { + self.cursor = provider; + true + } + "claude_code" => { + self.claude_code = provider; + true + } + "codex" => { + self.codex = provider; + true + } + "windsurf" => { + self.windsurf = provider; + true + } + "kiro" => { + self.kiro = provider; + true + } + "other" => { + self.other = provider; + true + } + _ => false, + } + } +} + /// 主配置结构 /// /// 支持两种格式: @@ -212,6 +304,10 @@ pub struct Config { /// Amp CLI 配置 #[serde(default)] pub ampcode: AmpConfig, + /// 端点 Provider 配置 + /// 允许为不同的客户端端点(CC/Codex)配置不同的 Provider + #[serde(default)] + pub endpoint_providers: EndpointProvidersConfig, } /// 服务器配置 @@ -674,6 +770,7 @@ impl Default for Config { quota_exceeded: QuotaExceededConfig::default(), proxy_url: None, ampcode: AmpConfig::default(), + endpoint_providers: EndpointProvidersConfig::default(), } } } @@ -834,4 +931,125 @@ mod unit_tests { assert!(config.model_aliases.is_empty()); assert!(config.exclusions.is_empty()); } + + #[test] + fn test_endpoint_providers_config_default() { + let config = EndpointProvidersConfig::default(); + assert!(config.cursor.is_none()); + assert!(config.claude_code.is_none()); + assert!(config.codex.is_none()); + assert!(config.windsurf.is_none()); + assert!(config.kiro.is_none()); + assert!(config.other.is_none()); + } + + #[test] + fn test_endpoint_providers_config_get_provider() { + let config = EndpointProvidersConfig { + cursor: Some("qwen".to_string()), + claude_code: Some("kiro".to_string()), + codex: Some("codex".to_string()), + windsurf: None, + kiro: Some("gemini".to_string()), + other: None, + }; + + assert_eq!(config.get_provider("cursor"), Some(&"qwen".to_string())); + assert_eq!( + config.get_provider("claude_code"), + Some(&"kiro".to_string()) + ); + assert_eq!(config.get_provider("codex"), Some(&"codex".to_string())); + assert_eq!(config.get_provider("windsurf"), None); + assert_eq!(config.get_provider("kiro"), Some(&"gemini".to_string())); + assert_eq!(config.get_provider("other"), None); + assert_eq!(config.get_provider("invalid"), None); + } + + #[test] + fn test_endpoint_providers_config_set_provider() { + let mut config = EndpointProvidersConfig::default(); + + // 设置有效的客户端类型 + assert!(config.set_provider("cursor", Some("qwen".to_string()))); + assert_eq!(config.cursor, Some("qwen".to_string())); + + assert!(config.set_provider("claude_code", Some("kiro".to_string()))); + assert_eq!(config.claude_code, Some("kiro".to_string())); + + assert!(config.set_provider("codex", Some("codex".to_string()))); + assert_eq!(config.codex, Some("codex".to_string())); + + assert!(config.set_provider("windsurf", Some("gemini".to_string()))); + assert_eq!(config.windsurf, Some("gemini".to_string())); + + assert!(config.set_provider("kiro", Some("openai".to_string()))); + assert_eq!(config.kiro, Some("openai".to_string())); + + assert!(config.set_provider("other", Some("claude".to_string()))); + assert_eq!(config.other, Some("claude".to_string())); + + // 设置无效的客户端类型 + assert!(!config.set_provider("invalid", Some("test".to_string()))); + } + + #[test] + fn test_endpoint_providers_config_set_provider_clear() { + let mut config = EndpointProvidersConfig { + cursor: Some("qwen".to_string()), + claude_code: Some("kiro".to_string()), + codex: None, + windsurf: None, + kiro: None, + other: None, + }; + + // 使用 None 清除配置 + assert!(config.set_provider("cursor", None)); + assert_eq!(config.cursor, None); + + // 使用空字符串清除配置 + assert!(config.set_provider("claude_code", Some("".to_string()))); + assert_eq!(config.claude_code, None); + } + + #[test] + fn test_endpoint_providers_config_serialization() { + let config = EndpointProvidersConfig { + cursor: Some("qwen".to_string()), + claude_code: Some("kiro".to_string()), + codex: None, + windsurf: None, + kiro: None, + other: None, + }; + + let yaml = serde_yaml::to_string(&config).unwrap(); + assert!(yaml.contains("cursor: qwen")); + assert!(yaml.contains("claude_code: kiro")); + // None 值应该被跳过 + assert!(!yaml.contains("codex")); + assert!(!yaml.contains("windsurf")); + assert!(!yaml.contains("kiro:")); + assert!(!yaml.contains("other")); + + let parsed: EndpointProvidersConfig = serde_yaml::from_str(&yaml).unwrap(); + assert_eq!(parsed, config); + } + + #[test] + fn test_endpoint_providers_config_json_serialization() { + let config = EndpointProvidersConfig { + cursor: Some("qwen".to_string()), + claude_code: None, + codex: Some("codex".to_string()), + windsurf: None, + kiro: None, + other: Some("openai".to_string()), + }; + + let json = serde_json::to_string(&config).unwrap(); + let parsed: EndpointProvidersConfig = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed, config); + } } diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 0ce234f36..769bdbfa2 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -312,6 +312,63 @@ async fn set_default_provider( Ok(provider) } +/// 获取端点 Provider 配置 +#[tauri::command] +async fn get_endpoint_providers( + state: tauri::State<'_, AppState>, +) -> Result { + let s = state.read().await; + let ep = &s.config.endpoint_providers; + Ok(serde_json::json!({ + "cursor": ep.cursor.clone(), + "claude_code": ep.claude_code.clone(), + "codex": ep.codex.clone(), + "windsurf": ep.windsurf.clone(), + "kiro": ep.kiro.clone(), + "other": ep.other.clone() + })) +} + +/// 设置端点 Provider 配置 +#[tauri::command] +async fn set_endpoint_provider( + state: tauri::State<'_, AppState>, + logs: tauri::State<'_, LogState>, + endpoint: String, + provider: Option, +) -> Result { + // 验证 provider(如果提供) + if let Some(ref p) = provider { + if !p.is_empty() { + let _: ProviderType = p.parse().map_err(|e: String| e)?; + } + } + + let mut s = state.write().await; + + // 使用 set_provider 方法设置对应的 provider + if !s + .config + .endpoint_providers + .set_provider(&endpoint, provider.clone()) + { + return Err(format!("未知的客户端类型: {}", endpoint)); + } + + config::save_config(&s.config).map_err(|e| e.to_string())?; + + let provider_display = provider.as_deref().unwrap_or("默认"); + logs.write().await.add( + "info", + &format!( + "客户端 {} 的 Provider 已设置为: {}", + endpoint, provider_display + ), + ); + + Ok(provider_display.to_string()) +} + #[tauri::command] async fn refresh_kiro_token( state: tauri::State<'_, AppState>, @@ -1541,7 +1598,7 @@ pub fn run() { let flow_replayer_state = FlowReplayerState(flow_replayer); // 初始化会话管理器 - let db_path = database::get_db_path(); + let db_path = database::get_db_path().expect("Failed to get database path"); let session_manager = Arc::new(SessionManager::new(db_path.clone()).expect("Failed to create SessionManager")); let session_manager_state = SessionManagerState(session_manager); @@ -1754,6 +1811,8 @@ pub fn run() { save_config, get_default_provider, set_default_provider, + get_endpoint_providers, + set_endpoint_provider, // Unified OAuth commands (new) commands::oauth_cmd::get_oauth_credentials, commands::oauth_cmd::reload_oauth_credentials, diff --git a/src-tauri/src/server/client_detector.rs b/src-tauri/src/server/client_detector.rs new file mode 100644 index 000000000..c3dd3c01e --- /dev/null +++ b/src-tauri/src/server/client_detector.rs @@ -0,0 +1,450 @@ +//! 客户端类型检测模块 +//! +//! 通过解析 HTTP 请求的 User-Agent 头来识别客户端类型。 + +use serde::{Deserialize, Serialize}; + +/// 客户端类型枚举 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ClientType { + /// Cursor 编辑器 + Cursor, + /// Claude Code 客户端 + ClaudeCode, + /// OpenAI Codex CLI + Codex, + /// Windsurf 编辑器 + Windsurf, + /// Kiro IDE + Kiro, + /// 未识别的客户端 + Other, +} + +impl ClientType { + /// 从 User-Agent 字符串检测客户端类型 + /// + /// 支持大小写不敏感匹配。 + /// + /// # 参数 + /// - `user_agent`: HTTP 请求的 User-Agent 头值 + /// + /// # 返回 + /// 检测到的客户端类型 + /// + /// # 示例 + /// ```ignore + /// use proxycast_lib::server::client_detector::ClientType; + /// + /// assert_eq!(ClientType::from_user_agent("Cursor/1.0"), ClientType::Cursor); + /// assert_eq!(ClientType::from_user_agent("claude-code/2.0"), ClientType::ClaudeCode); + /// assert_eq!(ClientType::from_user_agent("Unknown"), ClientType::Other); + /// ``` + pub fn from_user_agent(user_agent: &str) -> Self { + let ua_lower = user_agent.to_lowercase(); + + if ua_lower.contains("cursor") { + ClientType::Cursor + } else if ua_lower.contains("claude-code") || ua_lower.contains("claude_code") { + ClientType::ClaudeCode + } else if ua_lower.contains("codex") { + ClientType::Codex + } else if ua_lower.contains("windsurf") { + ClientType::Windsurf + } else if ua_lower.contains("kiro") { + ClientType::Kiro + } else { + ClientType::Other + } + } + + /// 获取配置键名 + /// + /// 返回用于配置文件中的键名。 + /// + /// # 返回 + /// 配置键名字符串 + pub fn config_key(&self) -> &'static str { + match self { + ClientType::Cursor => "cursor", + ClientType::ClaudeCode => "claude_code", + ClientType::Codex => "codex", + ClientType::Windsurf => "windsurf", + ClientType::Kiro => "kiro", + ClientType::Other => "other", + } + } + + /// 获取所有客户端类型 + /// + /// 返回所有支持的客户端类型列表。 + pub fn all() -> &'static [ClientType] { + &[ + ClientType::Cursor, + ClientType::ClaudeCode, + ClientType::Codex, + ClientType::Windsurf, + ClientType::Kiro, + ClientType::Other, + ] + } + + /// 从配置键名解析客户端类型 + /// + /// # 参数 + /// - `key`: 配置键名 + /// + /// # 返回 + /// 如果键名有效,返回对应的客户端类型;否则返回 None + pub fn from_config_key(key: &str) -> Option { + match key { + "cursor" => Some(ClientType::Cursor), + "claude_code" => Some(ClientType::ClaudeCode), + "codex" => Some(ClientType::Codex), + "windsurf" => Some(ClientType::Windsurf), + "kiro" => Some(ClientType::Kiro), + "other" => Some(ClientType::Other), + _ => None, + } + } +} + +impl std::fmt::Display for ClientType { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}", self.config_key()) + } +} + +/// 根据客户端类型和端点配置选择 Provider +/// +/// **Validates: Requirements 1.3, 1.4, 3.4** +/// +/// 优先级:端点 Provider 配置 > 默认 Provider +/// +/// # 参数 +/// - `client_type`: 检测到的客户端类型 +/// - `endpoint_provider`: 端点配置中该客户端类型对应的 Provider(可选) +/// - `default_provider`: 默认 Provider +/// +/// # 返回 +/// 选择的 Provider 名称 +pub fn select_provider( + client_type: ClientType, + endpoint_provider: Option<&String>, + default_provider: &str, +) -> String { + match endpoint_provider { + Some(provider) => provider.clone(), + None => default_provider.to_string(), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_from_user_agent_cursor() { + assert_eq!( + ClientType::from_user_agent("Cursor/1.0"), + ClientType::Cursor + ); + assert_eq!(ClientType::from_user_agent("cursor"), ClientType::Cursor); + assert_eq!(ClientType::from_user_agent("CURSOR"), ClientType::Cursor); + assert_eq!( + ClientType::from_user_agent("Mozilla/5.0 Cursor"), + ClientType::Cursor + ); + } + + #[test] + fn test_from_user_agent_claude_code() { + assert_eq!( + ClientType::from_user_agent("Claude-Code/2.0"), + ClientType::ClaudeCode + ); + assert_eq!( + ClientType::from_user_agent("claude-code"), + ClientType::ClaudeCode + ); + assert_eq!( + ClientType::from_user_agent("CLAUDE-CODE"), + ClientType::ClaudeCode + ); + assert_eq!( + ClientType::from_user_agent("claude_code"), + ClientType::ClaudeCode + ); + assert_eq!( + ClientType::from_user_agent("CLAUDE_CODE"), + ClientType::ClaudeCode + ); + } + + #[test] + fn test_from_user_agent_codex() { + assert_eq!(ClientType::from_user_agent("Codex/1.0"), ClientType::Codex); + assert_eq!(ClientType::from_user_agent("codex"), ClientType::Codex); + assert_eq!(ClientType::from_user_agent("CODEX"), ClientType::Codex); + } + + #[test] + fn test_from_user_agent_windsurf() { + assert_eq!( + ClientType::from_user_agent("Windsurf/1.0"), + ClientType::Windsurf + ); + assert_eq!( + ClientType::from_user_agent("windsurf"), + ClientType::Windsurf + ); + assert_eq!( + ClientType::from_user_agent("WINDSURF"), + ClientType::Windsurf + ); + } + + #[test] + fn test_from_user_agent_kiro() { + assert_eq!(ClientType::from_user_agent("Kiro/1.0"), ClientType::Kiro); + assert_eq!(ClientType::from_user_agent("kiro"), ClientType::Kiro); + assert_eq!(ClientType::from_user_agent("KIRO"), ClientType::Kiro); + } + + #[test] + fn test_from_user_agent_other() { + assert_eq!(ClientType::from_user_agent("Unknown"), ClientType::Other); + assert_eq!(ClientType::from_user_agent(""), ClientType::Other); + assert_eq!( + ClientType::from_user_agent("Mozilla/5.0"), + ClientType::Other + ); + } + + #[test] + fn test_config_key() { + assert_eq!(ClientType::Cursor.config_key(), "cursor"); + assert_eq!(ClientType::ClaudeCode.config_key(), "claude_code"); + assert_eq!(ClientType::Codex.config_key(), "codex"); + assert_eq!(ClientType::Windsurf.config_key(), "windsurf"); + assert_eq!(ClientType::Kiro.config_key(), "kiro"); + assert_eq!(ClientType::Other.config_key(), "other"); + } + + #[test] + fn test_from_config_key() { + assert_eq!( + ClientType::from_config_key("cursor"), + Some(ClientType::Cursor) + ); + assert_eq!( + ClientType::from_config_key("claude_code"), + Some(ClientType::ClaudeCode) + ); + assert_eq!( + ClientType::from_config_key("codex"), + Some(ClientType::Codex) + ); + assert_eq!( + ClientType::from_config_key("windsurf"), + Some(ClientType::Windsurf) + ); + assert_eq!(ClientType::from_config_key("kiro"), Some(ClientType::Kiro)); + assert_eq!( + ClientType::from_config_key("other"), + Some(ClientType::Other) + ); + assert_eq!(ClientType::from_config_key("invalid"), None); + } + + #[test] + fn test_all_client_types() { + let all = ClientType::all(); + assert_eq!(all.len(), 6); + assert!(all.contains(&ClientType::Cursor)); + assert!(all.contains(&ClientType::ClaudeCode)); + assert!(all.contains(&ClientType::Codex)); + assert!(all.contains(&ClientType::Windsurf)); + assert!(all.contains(&ClientType::Kiro)); + assert!(all.contains(&ClientType::Other)); + } + + #[test] + fn test_display() { + assert_eq!(format!("{}", ClientType::Cursor), "cursor"); + assert_eq!(format!("{}", ClientType::ClaudeCode), "claude_code"); + } + + #[test] + fn test_serialization() { + let cursor = ClientType::Cursor; + let json = serde_json::to_string(&cursor).unwrap(); + assert_eq!(json, "\"cursor\""); + + let claude_code = ClientType::ClaudeCode; + let json = serde_json::to_string(&claude_code).unwrap(); + assert_eq!(json, "\"claude_code\""); + } + + #[test] + fn test_deserialization() { + let cursor: ClientType = serde_json::from_str("\"cursor\"").unwrap(); + assert_eq!(cursor, ClientType::Cursor); + + let claude_code: ClientType = serde_json::from_str("\"claude_code\"").unwrap(); + assert_eq!(claude_code, ClientType::ClaudeCode); + } +} + +// ============================================================================ +// Property 2: Provider 选择优先级属性测试 +// ============================================================================ + +#[cfg(test)] +mod property_tests { + use super::*; + use crate::config::EndpointProvidersConfig; + use proptest::prelude::*; + + /// 生成随机的客户端类型 + fn arb_client_type() -> impl Strategy { + prop_oneof![ + Just(ClientType::Cursor), + Just(ClientType::ClaudeCode), + Just(ClientType::Codex), + Just(ClientType::Windsurf), + Just(ClientType::Kiro), + Just(ClientType::Other), + ] + } + + /// 生成随机的 Provider 名称 + fn arb_provider_name() -> impl Strategy { + prop_oneof![ + Just("kiro".to_string()), + Just("gemini".to_string()), + Just("qwen".to_string()), + Just("openai".to_string()), + Just("claude".to_string()), + Just("codex".to_string()), + ] + } + + /// 生成可选的 Provider 名称 + fn arb_optional_provider() -> impl Strategy> { + prop_oneof![Just(None), arb_provider_name().prop_map(Some),] + } + + /// 生成随机的 EndpointProvidersConfig + fn arb_endpoint_providers_config() -> impl Strategy { + ( + arb_optional_provider(), + arb_optional_provider(), + arb_optional_provider(), + arb_optional_provider(), + arb_optional_provider(), + arb_optional_provider(), + ) + .prop_map(|(cursor, claude_code, codex, windsurf, kiro, other)| { + EndpointProvidersConfig { + cursor, + claude_code, + codex, + windsurf, + kiro, + other, + } + }) + } + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: endpoint-provider-config, Property 2: Provider 选择优先级** + /// *对于任意* 客户端类型和配置: + /// - 当 endpoint_providers[client_type] 有值时,应使用该 Provider + /// - 当 endpoint_providers[client_type] 为空时,应使用 default_provider + /// **Validates: Requirements 1.3, 1.4, 3.4** + #[test] + fn prop_provider_selection_priority( + client_type in arb_client_type(), + endpoint_config in arb_endpoint_providers_config(), + default_provider in arb_provider_name() + ) { + // 获取端点配置中该客户端类型对应的 Provider + let endpoint_provider = endpoint_config.get_provider(client_type.config_key()); + + // 调用 select_provider 函数 + let selected = select_provider(client_type, endpoint_provider, &default_provider); + + // 验证选择逻辑 + match endpoint_provider { + Some(provider) => { + // 当端点配置有值时,应使用端点配置的 Provider + prop_assert_eq!( + selected, + provider.clone(), + "当端点配置有值时,应使用端点配置的 Provider" + ); + } + None => { + // 当端点配置为空时,应使用默认 Provider + prop_assert_eq!( + selected, + default_provider, + "当端点配置为空时,应使用默认 Provider" + ); + } + } + } + + /// **Feature: endpoint-provider-config, Property 2: Provider 选择优先级(端点配置优先)** + /// *对于任意* 客户端类型,当端点配置有值时,应始终使用端点配置的 Provider, + /// 而不是默认 Provider。 + /// **Validates: Requirements 1.3, 3.4** + #[test] + fn prop_endpoint_config_takes_priority( + client_type in arb_client_type(), + endpoint_provider in arb_provider_name(), + default_provider in arb_provider_name() + ) { + // 调用 select_provider 函数,端点配置有值 + let selected = select_provider( + client_type, + Some(&endpoint_provider), + &default_provider + ); + + // 验证:端点配置优先于默认配置 + prop_assert_eq!( + selected, + endpoint_provider, + "端点配置应优先于默认配置" + ); + } + + /// **Feature: endpoint-provider-config, Property 2: Provider 选择优先级(回退到默认)** + /// *对于任意* 客户端类型,当端点配置为空时,应使用默认 Provider。 + /// **Validates: Requirements 1.4** + #[test] + fn prop_fallback_to_default_provider( + client_type in arb_client_type(), + default_provider in arb_provider_name() + ) { + // 调用 select_provider 函数,端点配置为空 + let selected = select_provider( + client_type, + None, + &default_provider + ); + + // 验证:回退到默认 Provider + prop_assert_eq!( + selected, + default_provider, + "当端点配置为空时,应回退到默认 Provider" + ); + } + } +} diff --git a/src-tauri/src/server/handlers/api.rs b/src-tauri/src/server/handlers/api.rs index e6f1a1dda..00da50ca2 100644 --- a/src-tauri/src/server/handlers/api.rs +++ b/src-tauri/src/server/handlers/api.rs @@ -14,21 +14,32 @@ //! - 需求 5.3: 流中发生错误时发送错误事件并优雅关闭流 use axum::{ + body::Body, extract::State, - http::{HeaderMap, StatusCode}, + http::{header, HeaderMap, StatusCode}, response::{IntoResponse, Response}, Json, }; +use chrono::Utc; +use std::collections::HashMap; use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; +use crate::flow_monitor::{ + ClientInfo, FlowError, FlowErrorType, FlowMetadata, FlowType, InterceptAction, InterceptType, + LLMFlow, LLMRequest, LLMResponse, Message, MessageContent, MessageRole, RequestParameters, + RoutingInfo, TokenUsage, +}; use crate::models::anthropic::AnthropicMessagesRequest; use crate::models::openai::ChatCompletionRequest; use crate::processor::RequestContext; +use crate::server::client_detector::ClientType; use crate::server::{record_request_telemetry, record_token_usage, AppState}; use crate::server_utils::{ build_anthropic_response, build_anthropic_stream_response, message_content_len, parse_cw_response, safe_truncate, }; +use crate::streaming::StreamFormat as StreamingFormat; +use crate::ProviderType; use super::{call_provider_anthropic, call_provider_openai}; @@ -316,6 +327,46 @@ fn build_llm_response(status_code: u16, content: &str, usage: Option<(u32, u32)> } } +// ============================================================================ +// Provider 选择辅助函数 +// ============================================================================ + +/// 根据客户端类型和端点配置选择 Provider +/// +/// **Validates: Requirements 1.3, 1.4, 3.4** +/// +/// 优先级:端点 Provider 配置 > 默认 Provider +/// +/// # 参数 +/// - `headers`: HTTP 请求头,用于提取 User-Agent +/// - `state`: 应用状态,包含端点配置和默认 Provider +/// +/// # 返回 +/// 选择的 Provider 名称和检测到的客户端类型 +async fn select_provider_for_client(headers: &HeaderMap, state: &AppState) -> (String, ClientType) { + // 从 User-Agent 检测客户端类型 + let user_agent = headers + .get("user-agent") + .and_then(|v| v.to_str().ok()) + .unwrap_or(""); + let client_type = ClientType::from_user_agent(user_agent); + + // 获取端点 Provider 配置 + let endpoint_providers = state.endpoint_providers.read().await; + let endpoint_provider = endpoint_providers.get_provider(client_type.config_key()); + + // 获取默认 Provider + let default_provider = state.default_provider.read().await.clone(); + + // 选择 Provider:端点配置优先,否则使用默认 + let selected_provider = match endpoint_provider { + Some(provider) => provider.clone(), + None => default_provider, + }; + + (selected_provider, client_type) +} + // ============================================================================ // 拦截检查辅助函数 // ============================================================================ @@ -625,8 +676,18 @@ pub async fn chat_completions( } } - // 获取当前默认 provider(用于凭证池选择) - let default_provider = state.default_provider.read().await.clone(); + // 根据客户端类型选择 Provider + // **Validates: Requirements 3.1, 3.3, 3.4** + let (selected_provider, client_type) = select_provider_for_client(&headers, &state).await; + + // 记录客户端检测和 Provider 选择结果 + state.logs.write().await.add( + "info", + &format!( + "[CLIENT] request_id={} client_type={} selected_provider={}", + ctx.request_id, client_type, selected_provider + ), + ); // 记录路由结果 state.logs.write().await.add( @@ -641,7 +702,7 @@ pub async fn chat_completions( let credential = match &state.db { Some(db) => state .pool_service - .select_credential(db, &default_provider, Some(&request.model)) + .select_credential(db, &selected_provider, Some(&request.model)) .ok() .flatten(), None => None, @@ -815,12 +876,35 @@ pub async fn chat_completions( return response; } - // 回退到旧的单凭证模式 + // 回退到旧的单凭证模式(仅当选择的 Provider 是 Kiro 时) + // 如果选择的 Provider 不是 Kiro,且凭证池中没有找到凭证,返回错误 + // **Validates: Requirements 3.2** + if selected_provider.to_lowercase() != "kiro" { + state.logs.write().await.add( + "error", + &format!( + "[ROUTE] No pool credential found for '{}' (client_type={}), and legacy mode only supports Kiro", + selected_provider, client_type + ), + ); + return ( + StatusCode::SERVICE_UNAVAILABLE, + Json(serde_json::json!({ + "error": { + "message": format!("没有找到可用的 '{}' 凭证。请在凭证池中添加对应的凭证。", selected_provider), + "type": "no_credential_error", + "code": "no_credential" + } + })), + ) + .into_response(); + } + state.logs.write().await.add( "debug", &format!( "[ROUTE] No pool credential found for '{}', using legacy mode", - default_provider + selected_provider ), ); @@ -1414,8 +1498,18 @@ pub async fn anthropic_messages( } } - // 获取当前默认 provider(用于凭证池选择) - let default_provider = state.default_provider.read().await.clone(); + // 根据客户端类型选择 Provider + // **Validates: Requirements 3.1, 3.3, 3.4** + let (selected_provider, client_type) = select_provider_for_client(&headers, &state).await; + + // 记录客户端检测和 Provider 选择结果 + state.logs.write().await.add( + "info", + &format!( + "[CLIENT] request_id={} client_type={} selected_provider={}", + ctx.request_id, client_type, selected_provider + ), + ); // 记录路由结果 state.logs.write().await.add( @@ -1429,10 +1523,10 @@ pub async fn anthropic_messages( // 尝试从凭证池中选择凭证 let credential = match &state.db { Some(db) => { - // 根据 default_provider 配置选择凭证 + // 根据选择的 Provider 配置选择凭证 state .pool_service - .select_credential(db, &default_provider, Some(&request.model)) + .select_credential(db, &selected_provider, Some(&request.model)) .ok() .flatten() } @@ -1609,12 +1703,35 @@ pub async fn anthropic_messages( return response; } - // 回退到旧的单凭证模式 + // 回退到旧的单凭证模式(仅当选择的 Provider 是 Kiro 时) + // 如果选择的 Provider 不是 Kiro,且凭证池中没有找到凭证,返回错误 + // **Validates: Requirements 3.2** + if selected_provider.to_lowercase() != "kiro" { + state.logs.write().await.add( + "error", + &format!( + "[ROUTE] No pool credential found for '{}' (client_type={}), and legacy mode only supports Kiro", + selected_provider, client_type + ), + ); + return ( + StatusCode::SERVICE_UNAVAILABLE, + Json(serde_json::json!({ + "type": "error", + "error": { + "type": "no_credential_error", + "message": format!("没有找到可用的 '{}' 凭证。请在凭证池中添加对应的凭证。", selected_provider) + } + })), + ) + .into_response(); + } + state.logs.write().await.add( "debug", &format!( "[ROUTE] No pool credential found for '{}', using legacy mode", - default_provider + selected_provider ), ); diff --git a/src-tauri/src/server/handlers/provider_calls.rs b/src-tauri/src/server/handlers/provider_calls.rs index a3de83da5..b6f6d7b48 100644 --- a/src-tauri/src/server/handlers/provider_calls.rs +++ b/src-tauri/src/server/handlers/provider_calls.rs @@ -20,6 +20,7 @@ use axum::{ response::{IntoResponse, Response}, Json, }; +use futures::StreamExt; use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; use crate::converter::openai_to_antigravity::{ diff --git a/src-tauri/src/server/handlers/websocket.rs b/src-tauri/src/server/handlers/websocket.rs index 417410de1..2387eb349 100644 --- a/src-tauri/src/server/handlers/websocket.rs +++ b/src-tauri/src/server/handlers/websocket.rs @@ -196,7 +196,12 @@ pub async fn handle_websocket( handle_ws_message(&state, &conn_id, ws_msg, &flow_subscribed).await; if let Some(resp) = response { let resp_text = serde_json::to_string(&resp).unwrap_or_default(); - if sender.send(WsMessage::Text(resp_text)).await.is_err() { + let mut sender_guard = sender.lock().await; + if sender_guard + .send(WsMessage::Text(resp_text.into())) + .await + .is_err() + { break; } } @@ -208,7 +213,12 @@ pub async fn handle_websocket( e ))); let error_text = serde_json::to_string(&error).unwrap_or_default(); - if sender.send(WsMessage::Text(error_text)).await.is_err() { + let mut sender_guard = sender.lock().await; + if sender_guard + .send(WsMessage::Text(error_text.into())) + .await + .is_err() + { break; } } @@ -220,7 +230,12 @@ pub async fn handle_websocket( "Binary messages not supported", )); let error_text = serde_json::to_string(&error).unwrap_or_default(); - if sender.send(WsMessage::Text(error_text)).await.is_err() { + let mut sender_guard = sender.lock().await; + if sender_guard + .send(WsMessage::Text(error_text.into())) + .await + .is_err() + { break; } } diff --git a/src-tauri/src/server/mod.rs b/src-tauri/src/server/mod.rs index e1ca3330b..7b7d4e852 100644 --- a/src-tauri/src/server/mod.rs +++ b/src-tauri/src/server/mod.rs @@ -1,7 +1,10 @@ //! HTTP API 服务器 + +pub mod client_detector; + use crate::config::{ - Config, ConfigChangeEvent, ConfigChangeKind, ConfigManager, FileWatcher, HotReloadManager, - ReloadResult, + Config, ConfigChangeEvent, ConfigChangeKind, ConfigManager, EndpointProvidersConfig, + FileWatcher, HotReloadManager, ReloadResult, }; use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; use crate::credential::CredentialSyncService; @@ -381,6 +384,8 @@ pub struct AppState { pub flow_monitor: Arc, /// Flow 拦截器 pub flow_interceptor: Arc, + /// 端点 Provider 配置 + pub endpoint_providers: Arc>, } /// 启动配置文件监控 @@ -721,6 +726,14 @@ async fn run_server( let flow_interceptor = shared_flow_interceptor.unwrap_or_else(|| Arc::new(FlowInterceptor::default())); + // 初始化端点 Provider 配置 + let endpoint_providers = Arc::new(RwLock::new( + config + .as_ref() + .map(|c| c.endpoint_providers.clone()) + .unwrap_or_default(), + )); + let state = AppState { api_key: api_key.to_string(), base_url, @@ -743,6 +756,7 @@ async fn run_server( amp_router, flow_monitor, flow_interceptor, + endpoint_providers, }; // 启动配置文件监控 diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 9bfe51501..d1c472304 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.17.1", + "version": "0.17.2", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src/components/provider-pool/CredentialCard.tsx b/src/components/provider-pool/CredentialCard.tsx index dc22b6521..fd6160bde 100644 --- a/src/components/provider-pool/CredentialCard.tsx +++ b/src/components/provider-pool/CredentialCard.tsx @@ -20,6 +20,7 @@ import { Fingerprint, Copy, Check, + Timer, } from "lucide-react"; import type { CredentialDisplay, @@ -191,7 +192,7 @@ export function CredentialCard({ return (
-
+ {/* 第一行:状态图标 + 名称 + 标签 + 操作按钮 */} +
{/* Status Icon */}
{credential.is_disabled ? ( - + ) : isHealthy ? ( - + ) : ( - + )}
{/* Main Info */}
-
-

- {credential.name || `凭证 #${credential.uuid.slice(0, 8)}`} -

-
- - {getCredentialTypeLabel(credential.credential_type)} - +

+ {credential.name || `凭证 #${credential.uuid.slice(0, 8)}`} +

+
+ + {getCredentialTypeLabel(credential.credential_type)} + + + + {sourceInfo.text} + + {credential.proxy_url && ( - - {sourceInfo.text} + + 代理 - {credential.proxy_url && ( - - - 代理 - - )} -
-
-

- {credential.uuid} -

-
- - {/* Stats */} -
-
- -
-
使用次数
-
{credential.usage_count}
-
-
-
- -
-
错误次数
-
{credential.error_count}
-
-
-
- -
-
最后使用
-
- {formatDate(credential.last_used)} -
-
-
-
- - {/* Health Check Info */} - {credential.last_health_check_time && ( -
-
检查: {formatDate(credential.last_health_check_time)}
- {credential.last_health_check_model && ( -
- ({credential.last_health_check_model}) -
)}
- )} +
{/* Actions */} -
+
+ {/* 第二行:统计信息 - 使用网格布局 */} +
+
+ {/* 使用次数 */} +
+ +
+
使用次数
+
+ {credential.usage_count} +
+
+
+ + {/* 错误次数 */} +
+ +
+
错误次数
+
+ {credential.error_count} +
+
+
+ + {/* 最后使用 */} +
+ +
+
最后使用
+
+ {formatDate(credential.last_used)} +
+
+
+ + {/* Token 有效期 - OAuth 凭证显示 */} + {isOAuth ? ( +
+ +
+
+ Token 有效期 +
+ {credential.token_cache_status?.expiry_time ? ( +
+ {formatDate(credential.token_cache_status.expiry_time)} +
+ ) : ( +
--
+ )} +
+
+ ) : ( +
/* 占位 */ + )} + + {/* 健康检查 */} + {credential.last_health_check_time ? ( +
+ +
+
健康检查
+
+ {formatDate(credential.last_health_check_time)} +
+
+
+ ) : ( +
/* 占位 */ + )} +
+
+ + {/* 第三行:UUID */} +
+

+ {credential.uuid} +

+
+ {/* Mobile Stats - shown on small screens */} -
-
-
- - - 使用: {credential.usage_count} - - - - 错误: {credential.error_count} - +
+
+
+ + 使用: + {credential.usage_count} +
+
+ + 错误: + {credential.error_count} +
+
+ + 最后使用: + {formatDate(credential.last_used)}
- - - {formatDate(credential.last_used)} -
{/* Error Message */} {credential.last_error_message && ( -
+
{credential.last_error_message.slice(0, 150)} {credential.last_error_message.length > 150 && "..."}
@@ -429,51 +487,51 @@ export function CredentialCard({ {/* 指纹信息展示区域 - 仅 Kiro 凭证 */} {isKiroCredential && fingerprintExpanded && ( -
-
- - +
+
+ + 设备指纹
{fingerprintLoading ? ( -
-
+
+
加载中...
) : fingerprintInfo ? ( -
+
- + Machine ID: - + {fingerprintInfo.machine_id_short}...
-
- +
+ 来源: - + 认证:
) : ( -
+
无法获取指纹信息
)} @@ -508,22 +566,22 @@ export function CredentialCard({ {/* 用量信息展示区域 - 仅 Kiro 凭证 */} {isKiroCredential && usageExpanded && ( -
-
- - +
+
+ + Kiro 用量
{usageError ? ( -
+
{usageError}
) : usageInfo ? ( diff --git a/src/components/routing/ClientRouting.tsx b/src/components/routing/ClientRouting.tsx new file mode 100644 index 000000000..44e7b179d --- /dev/null +++ b/src/components/routing/ClientRouting.tsx @@ -0,0 +1,213 @@ +import { useState, useEffect } from "react"; +import { Check, AlertTriangle, Monitor } from "lucide-react"; +import { + getEndpointProviders, + setEndpointProvider, + EndpointProvidersConfig, +} from "@/hooks/useTauri"; +import { providerPoolApi, ProviderPoolOverview } from "@/lib/api/providerPool"; + +const clientTypes = [ + { id: "cursor", label: "Cursor", description: "Cursor 编辑器" }, + { + id: "claude_code", + label: "Claude Code", + description: "Claude Code 客户端", + }, + { id: "codex", label: "Codex", description: "OpenAI Codex CLI" }, + { id: "windsurf", label: "Windsurf", description: "Windsurf 编辑器" }, + { id: "kiro", label: "Kiro", description: "Kiro IDE" }, + { id: "other", label: "其他", description: "未识别的客户端" }, +] as const; + +const providers = [ + { id: "kiro", label: "Kiro" }, + { id: "gemini", label: "Gemini" }, + { id: "qwen", label: "Qwen" }, + { id: "antigravity", label: "Antigravity" }, + { id: "openai", label: "OpenAI" }, + { id: "claude", label: "Claude" }, +] as const; + +interface ClientRoutingProps { + loading?: boolean; +} + +export function ClientRouting({ + loading: externalLoading, +}: ClientRoutingProps) { + const [endpointProviders, setEndpointProviders] = + useState({}); + const [poolOverview, setPoolOverview] = useState([]); + const [saveMsg, setSaveMsg] = useState(null); + const [loading, setLoading] = useState(false); + + // 加载数据 + const loadData = async () => { + setLoading(true); + try { + const [config, overview] = await Promise.all([ + getEndpointProviders(), + providerPoolApi.getOverview(), + ]); + setEndpointProviders(config); + setPoolOverview(overview); + } catch (e) { + console.error("Failed to load client routing config:", e); + } finally { + setLoading(false); + } + }; + + useEffect(() => { + loadData(); + }, []); + + // 自动清除保存消息 + useEffect(() => { + if (saveMsg) { + const timer = setTimeout(() => setSaveMsg(null), 3000); + return () => clearTimeout(timer); + } + }, [saveMsg]); + + // 处理配置变更 + const handleSetProvider = async ( + clientType: string, + provider: string | null, + ) => { + try { + await setEndpointProvider(clientType, provider); + setEndpointProviders((prev) => ({ + ...prev, + [clientType]: provider, + })); + const clientLabel = + clientTypes.find((c) => c.id === clientType)?.label || clientType; + const providerLabel = provider + ? providers.find((p) => p.id === provider)?.label || provider + : "默认 Provider"; + setSaveMsg(`${clientLabel} 已设置为 ${providerLabel}`); + } catch (e: unknown) { + const errMsg = e instanceof Error ? e.message : String(e); + setSaveMsg(`保存失败: ${errMsg}`); + } + }; + + // 检查是否有任何自定义配置 + const hasCustomConfig = Object.values(endpointProviders).some((v) => v); + + const isLoading = loading || externalLoading; + + return ( +
+
+
+

+ + 客户端路由 +

+

+ 根据客户端 User-Agent 自动选择不同的 Provider +

+
+ {hasCustomConfig && ( + + 已配置 {Object.values(endpointProviders).filter((v) => v).length} 项 + + )} +
+ + {/* 说明 */} +
+

工作原理:

+
    +
  • 系统根据请求的 User-Agent 头识别客户端类型
  • +
  • 选择"默认"时,使用 API Server 中配置的默认 Provider
  • +
  • 此配置优先级高于模型路由规则
  • +
+
+ + {/* 保存消息 */} + {saveMsg && ( +
+ {saveMsg.includes("失败") ? ( + + ) : ( + + )} + {saveMsg} +
+ )} + + {/* 客户端配置列表 */} + {isLoading ? ( +
+ 加载中... +
+ ) : ( +
+ {clientTypes.map((client) => { + const currentProvider = + endpointProviders[client.id as keyof EndpointProvidersConfig]; + // 检查配置的 Provider 是否有可用凭证 + const hasCredentials = currentProvider + ? poolOverview.some( + (o) => + o.provider_type === currentProvider && o.stats.total > 0, + ) + : true; + + return ( +
+
+
+ {client.label} + {!hasCredentials && currentProvider && ( + + + 无凭证 + + )} +
+ + {client.description} + +
+ +
+ ); + })} +
+ )} +
+ ); +} diff --git a/src/components/routing/RoutingPage.tsx b/src/components/routing/RoutingPage.tsx index b453eefe0..6f947dbf6 100644 --- a/src/components/routing/RoutingPage.tsx +++ b/src/components/routing/RoutingPage.tsx @@ -4,9 +4,11 @@ import { ModelMapping } from "./ModelMapping"; import { RoutingRules } from "./RoutingRules"; import { ExclusionList } from "./ExclusionList"; import { InjectionRules } from "./InjectionRules"; +import { ClientRouting } from "./ClientRouting"; import { HelpTip } from "@/components/HelpTip"; import { routerApi } from "@/lib/api/router"; import { injectionApi } from "@/lib/api/injection"; +import { setEndpointProvider } from "@/hooks/useTauri"; import type { ModelAlias, RoutingRule, @@ -19,7 +21,7 @@ export interface RoutingPageRef { refresh: () => void; } -type TabType = "aliases" | "rules" | "exclusions" | "injection"; +type TabType = "aliases" | "rules" | "exclusions" | "injection" | "clients"; interface RoutingPageProps { hideHeader?: boolean; @@ -27,7 +29,7 @@ interface RoutingPageProps { export const RoutingPage = forwardRef( ({ hideHeader = false }, ref) => { - const [activeTab, setActiveTab] = useState("aliases"); + const [activeTab, setActiveTab] = useState("clients"); const [loading, setLoading] = useState(false); const [error, setError] = useState(null); @@ -81,7 +83,32 @@ export const RoutingPage = forwardRef( ) => { setApplyingPreset(presetId); try { + // 先找到预设配置 + const preset = presets.find((p) => p.id === presetId); + + // 应用别名和规则 await routerApi.applyRecommendedPreset(presetId, merge); + + // 如果预设包含客户端路由配置,也应用它 + if (preset?.endpoint_providers) { + const ep = preset.endpoint_providers; + const clientTypes = [ + "cursor", + "claude_code", + "codex", + "windsurf", + "kiro", + "other", + ] as const; + for (const clientType of clientTypes) { + const provider = ep[clientType]; + // 只有在非合并模式或有值时才设置 + if (!merge || provider) { + await setEndpointProvider(clientType, provider || null); + } + } + } + await refresh(); setShowPresets(false); } catch (e) { @@ -178,6 +205,7 @@ export const RoutingPage = forwardRef( }; const tabs: { id: TabType; label: string; count: number }[] = [ + { id: "clients", label: "客户端路由", count: 0 }, { id: "aliases", label: "模型别名", count: aliases.length }, { id: "rules", label: "路由规则", count: rules.length }, { @@ -298,6 +326,16 @@ export const RoutingPage = forwardRef(
{preset.aliases.length} 个别名 {preset.rules.length} 条规则 + {preset.endpoint_providers && ( + + { + Object.values(preset.endpoint_providers).filter( + (v) => v, + ).length + }{" "} + 个客户端路由 + + )}
@@ -404,6 +442,7 @@ export const RoutingPage = forwardRef( loading={loading} /> )} + {activeTab === "clients" && } {activeTab === "exclusions" && ( { return invoke("check_api_compatibility", { provider }); } + +// ============ Endpoint Provider Configuration ============ + +/** + * 端点 Provider 配置 + * 为不同客户端类型配置不同的 LLM Provider + */ +export interface EndpointProvidersConfig { + /** Cursor 客户端使用的 Provider */ + cursor?: string | null; + /** Claude Code 客户端使用的 Provider */ + claude_code?: string | null; + /** Codex 客户端使用的 Provider */ + codex?: string | null; + /** Windsurf 客户端使用的 Provider */ + windsurf?: string | null; + /** Kiro 客户端使用的 Provider */ + kiro?: string | null; + /** 其他客户端使用的 Provider */ + other?: string | null; +} + +/** + * 获取端点 Provider 配置 + * @returns 端点 Provider 配置对象 + */ +export async function getEndpointProviders(): Promise { + return invoke("get_endpoint_providers"); +} + +/** + * 设置端点 Provider 配置 + * @param clientType 客户端类型 (cursor, claude_code, codex, windsurf, kiro, other) + * @param provider Provider 名称,传 null 表示使用默认 Provider + * @returns 设置后的 Provider 名称 + */ +export async function setEndpointProvider( + clientType: string, + provider: string | null, +): Promise { + return invoke("set_endpoint_provider", { endpoint: clientType, provider }); +} diff --git a/src/lib/api/router.ts b/src/lib/api/router.ts index cd33a6ea3..b7a025fcf 100644 --- a/src/lib/api/router.ts +++ b/src/lib/api/router.ts @@ -36,6 +36,14 @@ export interface RecommendedPreset { description: string; aliases: ModelAlias[]; rules: RoutingRule[]; + endpoint_providers?: { + cursor?: string; + claude_code?: string; + codex?: string; + windsurf?: string; + kiro?: string; + other?: string; + }; } // Router configuration