From a35d112cd49f2e37b5bf8e4c301c1997d2965398 Mon Sep 17 00:00:00 2001 From: YoVinchen Date: Mon, 1 Dec 2025 16:01:06 +0800 Subject: [PATCH] feat(proxy): implement provider adapter pattern with OpenRouter support MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit This major refactoring introduces a modular provider adapter architecture to support format transformation between different AI API formats. New features: - Add ProviderAdapter trait for unified provider abstraction - Implement Claude, Codex, and Gemini adapters with specific logic - Add Anthropic ↔ OpenAI format transformation for OpenRouter compatibility - Support model mapping from provider configuration (ANTHROPIC_MODEL, etc.) - Add OpenRouter preset to Claude provider presets Refactoring: - Extract authentication logic into auth.rs with AuthInfo and AuthStrategy - Move URL building and request transformation to individual adapters - Simplify ProviderRouter to only use proxy target providers - Refactor RequestForwarder to use adapter-based request/response handling - Use whitelist mode for header forwarding (only pass necessary headers) Architecture: - providers/adapter.rs: ProviderAdapter trait definition - providers/auth.rs: AuthInfo, AuthStrategy types - providers/claude.rs: Claude adapter with OpenRouter detection - providers/codex.rs: Codex (OpenAI) adapter - providers/gemini.rs: Gemini (Google) adapter - providers/models/: Anthropic and OpenAI API data models - providers/transform.rs: Bidirectional format transformation --- src-tauri/src/proxy/forwarder.rs | 315 +++------- src-tauri/src/proxy/mod.rs | 1 + src-tauri/src/proxy/providers/adapter.rs | 131 ++++ src-tauri/src/proxy/providers/auth.rs | 92 +++ src-tauri/src/proxy/providers/claude.rs | 276 +++++++++ src-tauri/src/proxy/providers/gemini.rs | 202 ++++++ src-tauri/src/proxy/providers/mod.rs | 38 ++ .../src/proxy/providers/models/anthropic.rs | 104 ++++ src-tauri/src/proxy/providers/models/mod.rs | 6 + .../src/proxy/providers/models/openai.rs | 113 ++++ src-tauri/src/proxy/providers/transform.rs | 586 ++++++++++++++++++ src-tauri/src/proxy/router.rs | 142 +---- src/App.tsx | 15 +- src/config/claudeProviderPresets.ts | 18 + 14 files changed, 1691 insertions(+), 348 deletions(-) create mode 100644 src-tauri/src/proxy/providers/adapter.rs create mode 100644 src-tauri/src/proxy/providers/auth.rs create mode 100644 src-tauri/src/proxy/providers/claude.rs create mode 100644 src-tauri/src/proxy/providers/gemini.rs create mode 100644 src-tauri/src/proxy/providers/mod.rs create mode 100644 src-tauri/src/proxy/providers/models/anthropic.rs create mode 100644 src-tauri/src/proxy/providers/models/mod.rs create mode 100644 src-tauri/src/proxy/providers/models/openai.rs create mode 100644 src-tauri/src/proxy/providers/transform.rs diff --git a/src-tauri/src/proxy/forwarder.rs b/src-tauri/src/proxy/forwarder.rs index e7415eb0f..9f9fcf6fc 100644 --- a/src-tauri/src/proxy/forwarder.rs +++ b/src-tauri/src/proxy/forwarder.rs @@ -2,7 +2,13 @@ //! //! 负责将请求转发到上游Provider,支持重试和故障转移 -use super::{error::*, router::ProviderRouter, types::ProxyStatus, ProxyError}; +use super::{ + error::*, + providers::{get_adapter, ProviderAdapter}, + router::ProviderRouter, + types::ProxyStatus, + ProxyError, +}; use crate::{app_config::AppType, database::Database, provider::Provider}; use reqwest::{Client, Response}; use serde_json::Value; @@ -52,6 +58,9 @@ impl RequestForwarder { let mut failed_ids = Vec::new(); let mut failover_happened = false; + // 获取适配器 + let adapter = get_adapter(app_type); + for attempt in 0..self.max_retries { // 选择Provider let provider = self.router.select_provider(app_type, &failed_ids).await?; @@ -78,7 +87,10 @@ impl RequestForwarder { let start = Instant::now(); // 转发请求 - match self.forward(&provider, endpoint, &body, &headers).await { + match self + .forward(&provider, endpoint, &body, &headers, adapter.as_ref()) + .await + { Ok(response) => { let _latency = start.elapsed().as_millis() as u64; @@ -168,51 +180,93 @@ impl RequestForwarder { Err(ProxyError::MaxRetriesExceeded) } - /// 转发单个请求 + /// 转发单个请求(使用适配器) async fn forward( &self, provider: &Provider, endpoint: &str, body: &Value, headers: &axum::http::HeaderMap, + adapter: &dyn ProviderAdapter, ) -> Result { - // 提取 base_url - let base_url = self.extract_base_url(provider)?; + // 使用适配器提取 base_url + let base_url = adapter.extract_base_url(provider)?; + log::info!("[{}] base_url: {}", adapter.name(), base_url); - // 使用辅助函数构建完整 URL(自动去重版本路径) - let url = self.build_full_url(&base_url, endpoint); + // 使用适配器构建 URL + let url = adapter.build_url(&base_url, endpoint); + + // 检查是否需要格式转换 + let needs_transform = adapter.needs_transform(provider); + + // 转换请求体(如果需要) + let request_body = if needs_transform { + log::info!("[{}] 转换请求格式 (Anthropic → OpenAI)", adapter.name()); + let transformed = adapter.transform_request(body.clone(), provider)?; + log::debug!( + "[{}] 转换后的请求: {}", + adapter.name(), + serde_json::to_string_pretty(&transformed).unwrap_or_default() + ); + transformed + } else { + body.clone() + }; + + log::info!( + "[{}] 转发请求: {} -> {}", + adapter.name(), + provider.name, + url + ); // 构建请求 let mut request = self.client.post(&url); - // 透传 Headers + // 只透传必要的 Headers(白名单模式) + let allowed_headers = [ + "accept", + "user-agent", + "x-request-id", + "x-stainless-arch", + "x-stainless-lang", + "x-stainless-os", + "x-stainless-package-version", + "x-stainless-runtime", + "x-stainless-runtime-version", + ]; + for (key, value) in headers { let key_str = key.as_str().to_lowercase(); - // 过滤掉一些不应该直接转发的 Header - if key_str == "host" - || key_str == "content-length" - || key_str == "accept-encoding" - // 过滤认证相关 Header - || key_str == "x-api-key" - || key_str == "authorization" - || key_str == "x-goog-api-key" - || key_str == "anthropic-version" - { - continue; + if allowed_headers.contains(&key_str.as_str()) { + request = request.header(key, value); } - - request = request.header(key, value); } // 确保 Content-Type 是 json request = request.header("Content-Type", "application/json"); - // 添加认证头 - request = self.add_auth_headers(request, provider)?; + // 使用适配器添加认证头 + if let Some(auth) = adapter.extract_auth(provider) { + log::debug!( + "[{}] 使用认证: {:?} (key: {})", + adapter.name(), + auth.strategy, + auth.masked_key() + ); + request = adapter.add_auth_headers(request, &auth); + } else { + log::error!( + "[{}] 未找到 API Key!Provider: {}", + adapter.name(), + provider.name + ); + } // 发送请求 - let response = request.json(body).send().await.map_err(|e| { - log::error!("Request Failed: {e}"); + log::info!("[{}] 发送请求到: {}", adapter.name(), url); + let response = request.json(&request_body).send().await.map_err(|e| { + log::error!("[{}] 请求失败: {}", adapter.name(), e); if e.is_timeout() { ProxyError::Timeout(format!("请求超时: {e}")) } else if e.is_connect() { @@ -224,12 +278,19 @@ impl RequestForwarder { // 检查响应状态 let status = response.status(); + log::info!("[{}] 响应状态: {}", adapter.name(), status); if status.is_success() { Ok(response) } else { let status_code = status.as_u16(); let body_text = response.text().await.ok(); + log::error!( + "[{}] 上游错误 ({}): {:?}", + adapter.name(), + status_code, + body_text + ); Err(ProxyError::UpstreamError { status: status_code, @@ -238,204 +299,6 @@ impl RequestForwarder { } } - /// 添加认证头 - fn add_auth_headers( - &self, - mut request: reqwest::RequestBuilder, - provider: &Provider, - ) -> Result { - // 提取 apiKey 和认证类型 - if let Some((api_key, auth_type)) = self.extract_api_key(provider) { - // 遮蔽 key 用于日志 - let _masked_key = if api_key.len() > 8 { - format!("{}...{}", &api_key[..4], &api_key[api_key.len() - 4..]) - } else { - "***".to_string() - }; - - match auth_type { - AuthType::Anthropic => { - request = request.header("x-api-key", api_key); - request = request.header("anthropic-version", "2023-06-01"); - } - AuthType::Gemini => { - request = request.header("x-goog-api-key", api_key); - } - AuthType::Bearer => { - request = request.header("Authorization", format!("Bearer {api_key}")); - } - } - } else { - log::error!("✗ 未找到 API Key!将发送未认证的请求(会失败)"); - log::error!("Provider 配置: {:?}", provider.settings_config); - } - - Ok(request) - } - - /// 构建完整 URL(智能去重版本路径) - fn build_full_url(&self, base_url: &str, endpoint: &str) -> String { - let base_trimmed = base_url.trim_end_matches('/'); - let endpoint_trimmed = endpoint.trim_start_matches('/'); - - // 检查是否存在版本路径重复 - let version_patterns = ["/v1beta", "/v1"]; - let mut final_url = format!("{base_trimmed}/{endpoint_trimmed}"); - - for pattern in &version_patterns { - let duplicate_pattern = format!("{pattern}{pattern}"); - if final_url.contains(&duplicate_pattern) { - final_url = final_url.replace(&duplicate_pattern, pattern); - log::debug!( - "URL 去重: 移除重复的 {pattern} (base: {base_url}, endpoint: {endpoint})" - ); - } - } - - final_url - } - - /// 从 Provider 配置中提取 base_url - fn extract_base_url(&self, provider: &Provider) -> Result { - log::debug!("Extracting base_url for provider: {}", provider.name); - - // 1. 尝试直接获取 base_url 字段 (Codex CLI 常用格式) - if let Some(url) = provider - .settings_config - .get("base_url") - .and_then(|v| v.as_str()) - { - log::debug!("Found base_url in direct field: {url}"); - return Ok(url.trim_end_matches('/').to_string()); - } - - // 2. 尝试从 env 中获取 (Claude / Gemini) - if let Some(env) = provider.settings_config.get("env") { - if let Some(url) = env.get("ANTHROPIC_BASE_URL").and_then(|v| v.as_str()) { - log::debug!("Found base_url in env.ANTHROPIC_BASE_URL: {url}"); - return Ok(url.trim_end_matches('/').to_string()); - } - if let Some(url) = env.get("GOOGLE_GEMINI_BASE_URL").and_then(|v| v.as_str()) { - log::debug!("Found base_url in env.GOOGLE_GEMINI_BASE_URL: {url}"); - return Ok(url.trim_end_matches('/').to_string()); - } - } - - // 3. 尝试其他通用字段 - if let Some(url) = provider - .settings_config - .get("baseURL") - .and_then(|v| v.as_str()) - { - log::debug!("Found base_url in baseURL: {url}"); - return Ok(url.trim_end_matches('/').to_string()); - } - if let Some(url) = provider - .settings_config - .get("apiEndpoint") - .and_then(|v| v.as_str()) - { - log::debug!("Found base_url in apiEndpoint: {url}"); - return Ok(url.trim_end_matches('/').to_string()); - } - - // 4. 尝试从 config 对象中获取 (Codex - JSON 格式) - if let Some(config) = provider.settings_config.get("config") { - // 如果 config 是一个对象 - if let Some(url) = config.get("base_url").and_then(|v| v.as_str()) { - log::debug!("Found base_url in config.base_url: {url}"); - return Ok(url.trim_end_matches('/').to_string()); - } - - // 如果 config 是一个字符串,尝试解析 - if let Some(config_str) = config.as_str() { - // 尝试双引号 - if let Some(start) = config_str.find("base_url = \"") { - let rest = &config_str[start + 12..]; - if let Some(end) = rest.find('"') { - let url = rest[..end].trim_end_matches('/').to_string(); - log::debug!("Found base_url in config string (double quotes): {url}"); - return Ok(url); - } - } - // 尝试单引号 - if let Some(start) = config_str.find("base_url = '") { - let rest = &config_str[start + 12..]; - if let Some(end) = rest.find('\'') { - let url = rest[..end].trim_end_matches('/').to_string(); - log::debug!("Found base_url in config string (single quotes): {url}"); - return Ok(url); - } - } - } - } - - log::error!( - "Failed to extract base_url from config: {:?}", - provider.settings_config - ); - Err(ProxyError::ConfigError( - "Provider缺少base_url配置".to_string(), - )) - } - - /// 从 Provider 配置中提取 api_key - fn extract_api_key(&self, provider: &Provider) -> Option<(String, AuthType)> { - // 1. 尝试从 env 中获取 - if let Some(env) = provider.settings_config.get("env") { - // Claude/Anthropic - if let Some(key) = env.get("ANTHROPIC_AUTH_TOKEN").and_then(|v| v.as_str()) { - return Some((key.to_string(), AuthType::Anthropic)); - } - - // Gemini (支持两种字段名,优先使用标准的 GOOGLE_GEMINI_API_KEY) - if let Some(key) = env - .get("GOOGLE_GEMINI_API_KEY") - .or_else(|| env.get("GEMINI_API_KEY")) - .and_then(|v| v.as_str()) - { - return Some((key.to_string(), AuthType::Gemini)); - } - - // OpenAI/Codex (env 中的 OPENAI_API_KEY) - if let Some(key) = env.get("OPENAI_API_KEY").and_then(|v| v.as_str()) { - return Some((key.to_string(), AuthType::Bearer)); - } - } - - // 2. 尝试从 auth 中获取 (Codex CLI 格式) - if let Some(auth) = provider.settings_config.get("auth") { - if let Some(key) = auth.get("OPENAI_API_KEY").and_then(|v| v.as_str()) { - return Some((key.to_string(), AuthType::Bearer)); - } - } - - // 3. 尝试直接获取 (支持 apiKey 和 api_key) - if let Some(key) = provider - .settings_config - .get("apiKey") - .or_else(|| provider.settings_config.get("api_key")) - .and_then(|v| v.as_str()) - { - return Some((key.to_string(), AuthType::Bearer)); - } - - // 4. 尝试从 config 对象中获取 - if let Some(config) = provider.settings_config.get("config") { - if let Some(key) = config - .get("api_key") - .or_else(|| config.get("apiKey")) - .and_then(|v| v.as_str()) - { - return Some((key.to_string(), AuthType::Bearer)); - } - } - - log::error!("✗ 所有位置都未找到 API Key!"); - log::error!("完整配置结构: {:?}", provider.settings_config); - None - } - /// 分类ProxyError fn categorize_proxy_error(&self, error: &ProxyError) -> ErrorCategory { match error { @@ -456,9 +319,3 @@ impl RequestForwarder { } } } - -enum AuthType { - Anthropic, - Gemini, - Bearer, -} diff --git a/src-tauri/src/proxy/mod.rs b/src-tauri/src/proxy/mod.rs index 6309cd2ef..8149aba85 100644 --- a/src-tauri/src/proxy/mod.rs +++ b/src-tauri/src/proxy/mod.rs @@ -6,6 +6,7 @@ pub mod error; mod forwarder; mod handlers; mod health; +pub mod providers; mod router; pub(crate) mod server; pub(crate) mod types; diff --git a/src-tauri/src/proxy/providers/adapter.rs b/src-tauri/src/proxy/providers/adapter.rs new file mode 100644 index 000000000..c0372e03f --- /dev/null +++ b/src-tauri/src/proxy/providers/adapter.rs @@ -0,0 +1,131 @@ +//! Provider Adapter Trait +//! +//! 定义供应商适配器的统一接口,抽象不同上游供应商的处理逻辑。 + +use super::auth::AuthInfo; +use crate::provider::Provider; +use crate::proxy::error::ProxyError; +use reqwest::RequestBuilder; +use serde_json::Value; + +/// 供应商适配器 Trait +/// +/// 所有供应商适配器都需要实现此 trait,提供统一的接口来处理: +/// - URL 构建 +/// - 认证信息提取和头部注入 +/// - 请求/响应格式转换(可选) +/// +/// # 示例 +/// +/// ```ignore +/// pub struct ClaudeAdapter; +/// +/// impl ProviderAdapter for ClaudeAdapter { +/// fn name(&self) -> &'static str { "Claude" } +/// +/// fn extract_base_url(&self, provider: &Provider) -> Result { +/// // 从 provider 配置中提取 base_url +/// } +/// +/// fn extract_auth(&self, provider: &Provider) -> Option { +/// // 从 provider 配置中提取认证信息 +/// } +/// +/// fn build_url(&self, base_url: &str, endpoint: &str) -> String { +/// format!("{}{}", base_url.trim_end_matches('/'), endpoint) +/// } +/// +/// fn add_auth_headers(&self, request: RequestBuilder, auth: &AuthInfo) -> RequestBuilder { +/// // 添加认证头 +/// } +/// } +/// ``` +pub trait ProviderAdapter: Send + Sync { + /// 适配器名称(用于日志和调试) + fn name(&self) -> &'static str; + + /// 从 Provider 配置中提取 base_url + /// + /// # Arguments + /// * `provider` - Provider 配置 + /// + /// # Returns + /// * `Ok(String)` - 提取到的 base_url(已去除尾部斜杠) + /// * `Err(ProxyError)` - 提取失败 + fn extract_base_url(&self, provider: &Provider) -> Result; + + /// 从 Provider 配置中提取认证信息 + /// + /// # Arguments + /// * `provider` - Provider 配置 + /// + /// # Returns + /// * `Some(AuthInfo)` - 提取到的认证信息 + /// * `None` - 未找到认证信息 + fn extract_auth(&self, provider: &Provider) -> Option; + + /// 构建请求 URL + /// + /// # Arguments + /// * `base_url` - 基础 URL + /// * `endpoint` - 请求端点(如 `/v1/messages`) + /// + /// # Returns + /// 完整的请求 URL + fn build_url(&self, base_url: &str, endpoint: &str) -> String; + + /// 添加认证头到请求 + /// + /// # Arguments + /// * `request` - reqwest RequestBuilder + /// * `auth` - 认证信息 + /// + /// # Returns + /// 添加了认证头的 RequestBuilder + fn add_auth_headers(&self, request: RequestBuilder, auth: &AuthInfo) -> RequestBuilder; + + /// 是否需要格式转换 + /// + /// 默认返回 `false`(透传模式)。 + /// 仅当供应商需要格式转换时(如 Claude + OpenRouter)才返回 `true`。 + /// + /// # Arguments + /// * `provider` - Provider 配置 + fn needs_transform(&self, _provider: &Provider) -> bool { + false + } + + /// 转换请求体 + /// + /// 将请求体从一种格式转换为另一种格式(如 Anthropic → OpenAI)。 + /// 默认实现直接返回原始请求体(透传)。 + /// + /// # Arguments + /// * `body` - 原始请求体 + /// * `provider` - Provider 配置(用于获取模型映射等) + /// + /// # Returns + /// * `Ok(Value)` - 转换后的请求体 + /// * `Err(ProxyError)` - 转换失败 + fn transform_request(&self, body: Value, _provider: &Provider) -> Result { + Ok(body) + } + + /// 转换响应体 + /// + /// 将响应体从一种格式转换为另一种格式(如 OpenAI → Anthropic)。 + /// 默认实现直接返回原始响应体(透传)。 + /// + /// # Arguments + /// * `body` - 原始响应体 + /// + /// # Returns + /// * `Ok(Value)` - 转换后的响应体 + /// * `Err(ProxyError)` - 转换失败 + /// + /// Note: 响应转换将在 handler 层集成,目前预留接口 + #[allow(dead_code)] + fn transform_response(&self, body: Value) -> Result { + Ok(body) + } +} diff --git a/src-tauri/src/proxy/providers/auth.rs b/src-tauri/src/proxy/providers/auth.rs new file mode 100644 index 000000000..9c293e40f --- /dev/null +++ b/src-tauri/src/proxy/providers/auth.rs @@ -0,0 +1,92 @@ +//! Authentication Types +//! +//! 定义认证信息和认证策略,支持多种上游供应商的认证方式。 + +/// 认证信息 +/// +/// 包含 API Key 和对应的认证策略 +#[derive(Debug, Clone)] +pub struct AuthInfo { + /// API Key + pub api_key: String, + /// 认证策略 + pub strategy: AuthStrategy, +} + +impl AuthInfo { + /// 创建新的认证信息 + pub fn new(api_key: String, strategy: AuthStrategy) -> Self { + Self { api_key, strategy } + } + + /// 返回遮蔽后的 API Key(用于日志输出) + /// + /// 显示前4位和后4位,中间用 `...` 代替 + /// 如果 key 长度不足8位,则返回 `***` + pub fn masked_key(&self) -> String { + if self.api_key.len() > 8 { + format!( + "{}...{}", + &self.api_key[..4], + &self.api_key[self.api_key.len() - 4..] + ) + } else { + "***".to_string() + } + } +} + +/// 认证策略 +/// +/// 不同供应商使用不同的认证方式 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum AuthStrategy { + /// Anthropic 认证方式 + /// - Header: `x-api-key: ` + /// - Header: `anthropic-version: 2023-06-01` + Anthropic, + + /// Bearer Token 认证方式(OpenAI 等) + /// - Header: `Authorization: Bearer ` + Bearer, + + /// Google 认证方式 + /// - Header: `x-goog-api-key: ` + Google, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_masked_key_long() { + let auth = AuthInfo::new("sk-1234567890abcdef".to_string(), AuthStrategy::Bearer); + assert_eq!(auth.masked_key(), "sk-1...cdef"); + } + + #[test] + fn test_masked_key_short() { + let auth = AuthInfo::new("short".to_string(), AuthStrategy::Bearer); + assert_eq!(auth.masked_key(), "***"); + } + + #[test] + fn test_masked_key_exactly_8() { + let auth = AuthInfo::new("12345678".to_string(), AuthStrategy::Bearer); + assert_eq!(auth.masked_key(), "***"); + } + + #[test] + fn test_masked_key_9_chars() { + let auth = AuthInfo::new("123456789".to_string(), AuthStrategy::Bearer); + assert_eq!(auth.masked_key(), "1234...6789"); + } + + #[test] + fn test_auth_strategy_equality() { + assert_eq!(AuthStrategy::Anthropic, AuthStrategy::Anthropic); + assert_ne!(AuthStrategy::Anthropic, AuthStrategy::Bearer); + assert_ne!(AuthStrategy::Bearer, AuthStrategy::Google); + } +} diff --git a/src-tauri/src/proxy/providers/claude.rs b/src-tauri/src/proxy/providers/claude.rs new file mode 100644 index 000000000..f01895917 --- /dev/null +++ b/src-tauri/src/proxy/providers/claude.rs @@ -0,0 +1,276 @@ +//! Claude (Anthropic) Provider Adapter +//! +//! 支持透传模式和 OpenRouter 兼容模式 + +use super::{AuthInfo, AuthStrategy, ProviderAdapter}; +use crate::provider::Provider; +use crate::proxy::error::ProxyError; +use reqwest::RequestBuilder; + +/// Claude 适配器 +pub struct ClaudeAdapter; + +impl ClaudeAdapter { + pub fn new() -> Self { + Self + } + + /// 检测是否使用 OpenRouter + fn is_openrouter(&self, provider: &Provider) -> bool { + if let Ok(base_url) = self.extract_base_url(provider) { + return base_url.contains("openrouter.ai"); + } + false + } + + /// 从 Provider 配置中提取 API Key + fn extract_key(&self, provider: &Provider) -> Option { + if let Some(env) = provider.settings_config.get("env") { + // Anthropic 标准 key + if let Some(key) = env + .get("ANTHROPIC_AUTH_TOKEN") + .and_then(|v| v.as_str()) + .filter(|s| !s.is_empty()) + { + log::debug!("[Claude] 使用 ANTHROPIC_AUTH_TOKEN"); + return Some(key.to_string()); + } + // OpenRouter key + if let Some(key) = env + .get("OPENROUTER_API_KEY") + .and_then(|v| v.as_str()) + .filter(|s| !s.is_empty()) + { + log::debug!("[Claude] 使用 OPENROUTER_API_KEY"); + return Some(key.to_string()); + } + // 备选 OpenAI key (用于 OpenRouter) + if let Some(key) = env + .get("OPENAI_API_KEY") + .and_then(|v| v.as_str()) + .filter(|s| !s.is_empty()) + { + log::debug!("[Claude] 使用 OPENAI_API_KEY"); + return Some(key.to_string()); + } + } + + // 尝试直接获取 + if let Some(key) = provider + .settings_config + .get("apiKey") + .or_else(|| provider.settings_config.get("api_key")) + .and_then(|v| v.as_str()) + .filter(|s| !s.is_empty()) + { + log::debug!("[Claude] 使用 apiKey/api_key"); + return Some(key.to_string()); + } + + log::warn!("[Claude] 未找到有效的 API Key"); + None + } +} + +impl Default for ClaudeAdapter { + fn default() -> Self { + Self::new() + } +} + +impl ProviderAdapter for ClaudeAdapter { + fn name(&self) -> &'static str { + "Claude" + } + + fn extract_base_url(&self, provider: &Provider) -> Result { + // 1. 从 env 中获取 + if let Some(env) = provider.settings_config.get("env") { + if let Some(url) = env.get("ANTHROPIC_BASE_URL").and_then(|v| v.as_str()) { + return Ok(url.trim_end_matches('/').to_string()); + } + } + + // 2. 尝试直接获取 + if let Some(url) = provider + .settings_config + .get("base_url") + .and_then(|v| v.as_str()) + { + return Ok(url.trim_end_matches('/').to_string()); + } + + if let Some(url) = provider + .settings_config + .get("baseURL") + .and_then(|v| v.as_str()) + { + return Ok(url.trim_end_matches('/').to_string()); + } + + if let Some(url) = provider + .settings_config + .get("apiEndpoint") + .and_then(|v| v.as_str()) + { + return Ok(url.trim_end_matches('/').to_string()); + } + + Err(ProxyError::ConfigError( + "Claude Provider 缺少 base_url 配置".to_string(), + )) + } + + fn extract_auth(&self, provider: &Provider) -> Option { + let is_openrouter = self.is_openrouter(provider); + let strategy = if is_openrouter { + AuthStrategy::Bearer + } else { + AuthStrategy::Anthropic + }; + + self.extract_key(provider) + .map(|key| AuthInfo::new(key, strategy)) + } + + fn build_url(&self, base_url: &str, endpoint: &str) -> String { + // OpenRouter 使用 /v1/chat/completions + if base_url.contains("openrouter.ai") { + return format!("{}/v1/chat/completions", base_url.trim_end_matches('/')); + } + + // Anthropic 直连 + format!( + "{}/{}", + base_url.trim_end_matches('/'), + endpoint.trim_start_matches('/') + ) + } + + fn add_auth_headers(&self, request: RequestBuilder, auth: &AuthInfo) -> RequestBuilder { + match auth.strategy { + AuthStrategy::Anthropic => request + .header("x-api-key", &auth.api_key) + .header("anthropic-version", "2023-06-01"), + AuthStrategy::Bearer => { + request.header("Authorization", format!("Bearer {}", auth.api_key)) + } + _ => request, + } + } + + fn needs_transform(&self, provider: &Provider) -> bool { + self.is_openrouter(provider) + } + + fn transform_request( + &self, + body: serde_json::Value, + provider: &Provider, + ) -> Result { + super::transform::anthropic_to_openai(body, provider) + } + + fn transform_response(&self, body: serde_json::Value) -> Result { + super::transform::openai_to_anthropic(body) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn create_provider(config: serde_json::Value) -> Provider { + Provider { + id: "test".to_string(), + name: "Test Claude".to_string(), + settings_config: config, + website_url: None, + category: Some("claude".to_string()), + created_at: None, + sort_index: None, + notes: None, + meta: None, + icon: None, + icon_color: None, + is_proxy_target: None, + } + } + + #[test] + fn test_extract_base_url_from_env() { + let adapter = ClaudeAdapter::new(); + let provider = create_provider(json!({ + "env": { + "ANTHROPIC_BASE_URL": "https://api.anthropic.com" + } + })); + + let url = adapter.extract_base_url(&provider).unwrap(); + assert_eq!(url, "https://api.anthropic.com"); + } + + #[test] + fn test_extract_auth_anthropic() { + let adapter = ClaudeAdapter::new(); + let provider = create_provider(json!({ + "env": { + "ANTHROPIC_BASE_URL": "https://api.anthropic.com", + "ANTHROPIC_AUTH_TOKEN": "sk-ant-test-key" + } + })); + + let auth = adapter.extract_auth(&provider).unwrap(); + assert_eq!(auth.api_key, "sk-ant-test-key"); + assert_eq!(auth.strategy, AuthStrategy::Anthropic); + } + + #[test] + fn test_extract_auth_openrouter() { + let adapter = ClaudeAdapter::new(); + let provider = create_provider(json!({ + "env": { + "ANTHROPIC_BASE_URL": "https://openrouter.ai/api", + "OPENROUTER_API_KEY": "sk-or-test-key" + } + })); + + let auth = adapter.extract_auth(&provider).unwrap(); + assert_eq!(auth.api_key, "sk-or-test-key"); + assert_eq!(auth.strategy, AuthStrategy::Bearer); + } + + #[test] + fn test_build_url_anthropic() { + let adapter = ClaudeAdapter::new(); + let url = adapter.build_url("https://api.anthropic.com", "/v1/messages"); + assert_eq!(url, "https://api.anthropic.com/v1/messages"); + } + + #[test] + fn test_build_url_openrouter() { + let adapter = ClaudeAdapter::new(); + let url = adapter.build_url("https://openrouter.ai/api", "/v1/messages"); + assert_eq!(url, "https://openrouter.ai/api/v1/chat/completions"); + } + + #[test] + fn test_needs_transform() { + let adapter = ClaudeAdapter::new(); + + let anthropic_provider = create_provider(json!({ + "env": { + "ANTHROPIC_BASE_URL": "https://api.anthropic.com" + } + })); + assert!(!adapter.needs_transform(&anthropic_provider)); + + let openrouter_provider = create_provider(json!({ + "env": { + "ANTHROPIC_BASE_URL": "https://openrouter.ai/api" + } + })); + assert!(adapter.needs_transform(&openrouter_provider)); + } +} diff --git a/src-tauri/src/proxy/providers/gemini.rs b/src-tauri/src/proxy/providers/gemini.rs new file mode 100644 index 000000000..3c3dcf43c --- /dev/null +++ b/src-tauri/src/proxy/providers/gemini.rs @@ -0,0 +1,202 @@ +//! Gemini (Google) Provider Adapter +//! +//! 仅透传模式,支持直连 Google Gemini API + +use super::{AuthInfo, AuthStrategy, ProviderAdapter}; +use crate::provider::Provider; +use crate::proxy::error::ProxyError; +use reqwest::RequestBuilder; + +/// Gemini 适配器 +pub struct GeminiAdapter; + +impl GeminiAdapter { + pub fn new() -> Self { + Self + } + + /// 从 Provider 配置中提取 API Key + fn extract_key(&self, provider: &Provider) -> Option { + if let Some(env) = provider.settings_config.get("env") { + // 优先使用 GOOGLE_GEMINI_API_KEY + if let Some(key) = env.get("GOOGLE_GEMINI_API_KEY").and_then(|v| v.as_str()) { + return Some(key.to_string()); + } + // 备选 GEMINI_API_KEY + if let Some(key) = env.get("GEMINI_API_KEY").and_then(|v| v.as_str()) { + return Some(key.to_string()); + } + } + + // 尝试直接获取 + if let Some(key) = provider + .settings_config + .get("apiKey") + .or_else(|| provider.settings_config.get("api_key")) + .and_then(|v| v.as_str()) + { + return Some(key.to_string()); + } + + None + } +} + +impl Default for GeminiAdapter { + fn default() -> Self { + Self::new() + } +} + +impl ProviderAdapter for GeminiAdapter { + fn name(&self) -> &'static str { + "Gemini" + } + + fn extract_base_url(&self, provider: &Provider) -> Result { + // 从 env 中获取 + if let Some(env) = provider.settings_config.get("env") { + if let Some(url) = env.get("GOOGLE_GEMINI_BASE_URL").and_then(|v| v.as_str()) { + return Ok(url.trim_end_matches('/').to_string()); + } + } + + // 尝试直接获取 + if let Some(url) = provider + .settings_config + .get("base_url") + .and_then(|v| v.as_str()) + { + return Ok(url.trim_end_matches('/').to_string()); + } + + if let Some(url) = provider + .settings_config + .get("baseURL") + .and_then(|v| v.as_str()) + { + return Ok(url.trim_end_matches('/').to_string()); + } + + Err(ProxyError::ConfigError( + "Gemini Provider 缺少 base_url 配置".to_string(), + )) + } + + fn extract_auth(&self, provider: &Provider) -> Option { + self.extract_key(provider) + .map(|key| AuthInfo::new(key, AuthStrategy::Google)) + } + + fn build_url(&self, base_url: &str, endpoint: &str) -> String { + let base_trimmed = base_url.trim_end_matches('/'); + let endpoint_trimmed = endpoint.trim_start_matches('/'); + + let mut url = format!("{base_trimmed}/{endpoint_trimmed}"); + + // 处理 /v1beta 路径去重 + let version_patterns = ["/v1beta", "/v1"]; + for pattern in &version_patterns { + let duplicate = format!("{pattern}{pattern}"); + if url.contains(&duplicate) { + url = url.replace(&duplicate, pattern); + } + } + + url + } + + fn add_auth_headers(&self, request: RequestBuilder, auth: &AuthInfo) -> RequestBuilder { + request.header("x-goog-api-key", &auth.api_key) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn create_provider(config: serde_json::Value) -> Provider { + Provider { + id: "test".to_string(), + name: "Test Gemini".to_string(), + settings_config: config, + website_url: None, + category: Some("gemini".to_string()), + created_at: None, + sort_index: None, + notes: None, + meta: None, + icon: None, + icon_color: None, + is_proxy_target: None, + } + } + + #[test] + fn test_extract_base_url_from_env() { + let adapter = GeminiAdapter::new(); + let provider = create_provider(json!({ + "env": { + "GOOGLE_GEMINI_BASE_URL": "https://generativelanguage.googleapis.com/v1beta" + } + })); + + let url = adapter.extract_base_url(&provider).unwrap(); + assert_eq!(url, "https://generativelanguage.googleapis.com/v1beta"); + } + + #[test] + fn test_extract_auth() { + let adapter = GeminiAdapter::new(); + let provider = create_provider(json!({ + "env": { + "GOOGLE_GEMINI_API_KEY": "AIza-test-key-12345678" + } + })); + + let auth = adapter.extract_auth(&provider).unwrap(); + assert_eq!(auth.api_key, "AIza-test-key-12345678"); + assert_eq!(auth.strategy, AuthStrategy::Google); + } + + #[test] + fn test_extract_auth_fallback() { + let adapter = GeminiAdapter::new(); + let provider = create_provider(json!({ + "env": { + "GEMINI_API_KEY": "AIza-fallback-key" + } + })); + + let auth = adapter.extract_auth(&provider).unwrap(); + assert_eq!(auth.api_key, "AIza-fallback-key"); + } + + #[test] + fn test_build_url_dedup() { + let adapter = GeminiAdapter::new(); + // 模拟 base_url 已包含 /v1beta,endpoint 也包含 /v1beta + let url = adapter.build_url( + "https://generativelanguage.googleapis.com/v1beta", + "/v1beta/models/gemini-pro:generateContent", + ); + assert_eq!( + url, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-pro:generateContent" + ); + } + + #[test] + fn test_build_url_normal() { + let adapter = GeminiAdapter::new(); + let url = adapter.build_url( + "https://generativelanguage.googleapis.com/v1beta", + "/models/gemini-pro:generateContent", + ); + assert_eq!( + url, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-pro:generateContent" + ); + } +} diff --git a/src-tauri/src/proxy/providers/mod.rs b/src-tauri/src/proxy/providers/mod.rs new file mode 100644 index 000000000..a0633ae65 --- /dev/null +++ b/src-tauri/src/proxy/providers/mod.rs @@ -0,0 +1,38 @@ +//! Provider Adapters Module +//! +//! 供应商适配器模块,提供统一的接口抽象不同上游供应商的处理逻辑。 +//! +//! ## 模块结构 +//! - `adapter`: 定义 `ProviderAdapter` trait +//! - `auth`: 认证类型和策略 +//! - `claude`: Claude (Anthropic) 适配器 +//! - `codex`: Codex (OpenAI) 适配器 +//! - `gemini`: Gemini (Google) 适配器 +//! - `models`: API 数据模型 +//! - `transform`: 格式转换 + +mod adapter; +mod auth; +mod claude; +mod codex; +mod gemini; +pub mod models; +pub mod transform; + +use crate::app_config::AppType; + +// 公开导出 +pub use adapter::ProviderAdapter; +pub use auth::{AuthInfo, AuthStrategy}; +pub use claude::ClaudeAdapter; +pub use codex::CodexAdapter; +pub use gemini::GeminiAdapter; + +/// 根据 AppType 获取对应的适配器 +pub fn get_adapter(app_type: &AppType) -> Box { + match app_type { + AppType::Claude => Box::new(ClaudeAdapter::new()), + AppType::Codex => Box::new(CodexAdapter::new()), + AppType::Gemini => Box::new(GeminiAdapter::new()), + } +} diff --git a/src-tauri/src/proxy/providers/models/anthropic.rs b/src-tauri/src/proxy/providers/models/anthropic.rs new file mode 100644 index 000000000..02791ecce --- /dev/null +++ b/src-tauri/src/proxy/providers/models/anthropic.rs @@ -0,0 +1,104 @@ +//! Anthropic API 数据模型 +//! +//! 用于 Anthropic Messages API 的请求/响应格式转换 + +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +/// Anthropic 请求 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AnthropicRequest { + pub model: String, + pub messages: Vec, + pub max_tokens: u32, + #[serde(skip_serializing_if = "Option::is_none")] + pub system: Option, // 可以是 String 或 Vec + #[serde(skip_serializing_if = "Option::is_none")] + pub temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub stream: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub tools: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_choice: Option, +} + +/// Anthropic 消息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AnthropicMessage { + pub role: String, + pub content: Value, // String 或 Vec +} + +/// Anthropic 内容块 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type")] +pub enum AnthropicContentBlock { + #[serde(rename = "text")] + Text { text: String }, + #[serde(rename = "image")] + Image { source: ImageSource }, + #[serde(rename = "tool_use")] + ToolUse { + id: String, + name: String, + input: Value, + }, + #[serde(rename = "tool_result")] + ToolResult { tool_use_id: String, content: Value }, +} + +/// 图片来源 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ImageSource { + #[serde(rename = "type")] + pub source_type: String, + pub media_type: String, + pub data: String, +} + +/// Anthropic 工具定义 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AnthropicTool { + pub name: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub description: Option, + pub input_schema: Value, +} + +/// Anthropic 响应 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AnthropicResponse { + pub id: String, + #[serde(rename = "type")] + pub response_type: String, + pub role: String, + pub content: Vec, + pub model: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub stop_reason: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub stop_sequence: Option, + pub usage: AnthropicUsage, +} + +/// Anthropic 响应内容 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type")] +pub enum AnthropicResponseContent { + #[serde(rename = "text")] + Text { text: String }, + #[serde(rename = "tool_use")] + ToolUse { + id: String, + name: String, + input: Value, + }, +} + +/// Anthropic 使用量 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AnthropicUsage { + pub input_tokens: u32, + pub output_tokens: u32, +} diff --git a/src-tauri/src/proxy/providers/models/mod.rs b/src-tauri/src/proxy/providers/models/mod.rs new file mode 100644 index 000000000..df3d9e109 --- /dev/null +++ b/src-tauri/src/proxy/providers/models/mod.rs @@ -0,0 +1,6 @@ +//! API 数据模型 +//! +//! 定义 Anthropic 和 OpenAI API 的请求/响应结构 + +pub mod anthropic; +pub mod openai; diff --git a/src-tauri/src/proxy/providers/models/openai.rs b/src-tauri/src/proxy/providers/models/openai.rs new file mode 100644 index 000000000..e04aefb4d --- /dev/null +++ b/src-tauri/src/proxy/providers/models/openai.rs @@ -0,0 +1,113 @@ +//! OpenAI API 数据模型 +//! +//! 用于 OpenAI Chat Completions API 的请求/响应格式转换 + +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +/// OpenAI 请求 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAIRequest { + pub model: String, + pub messages: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub stream: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub tools: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_choice: Option, +} + +/// OpenAI 消息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAIMessage { + pub role: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub content: Option, // String 或 Vec + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_calls: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_call_id: Option, +} + +/// OpenAI 内容部分 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type")] +pub enum OpenAIContentPart { + #[serde(rename = "text")] + Text { text: String }, + #[serde(rename = "image_url")] + ImageUrl { image_url: ImageUrl }, +} + +/// 图片 URL +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ImageUrl { + pub url: String, +} + +/// OpenAI 工具调用 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAIToolCall { + pub id: String, + #[serde(rename = "type")] + pub call_type: String, + pub function: OpenAIFunction, +} + +/// OpenAI 函数 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAIFunction { + pub name: String, + pub arguments: String, // JSON 字符串 +} + +/// OpenAI 工具定义 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAITool { + #[serde(rename = "type")] + pub tool_type: String, + pub function: OpenAIFunctionDef, +} + +/// OpenAI 函数定义 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAIFunctionDef { + pub name: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub description: Option, + pub parameters: Value, +} + +/// OpenAI 响应 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAIResponse { + pub id: String, + pub object: String, + pub created: u64, + pub model: String, + pub choices: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub usage: Option, +} + +/// OpenAI 选择 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAIChoice { + pub index: u32, + pub message: OpenAIMessage, + #[serde(skip_serializing_if = "Option::is_none")] + pub finish_reason: Option, +} + +/// OpenAI 使用量 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OpenAIUsage { + pub prompt_tokens: u32, + pub completion_tokens: u32, + pub total_tokens: u32, +} diff --git a/src-tauri/src/proxy/providers/transform.rs b/src-tauri/src/proxy/providers/transform.rs new file mode 100644 index 000000000..eaabb3e57 --- /dev/null +++ b/src-tauri/src/proxy/providers/transform.rs @@ -0,0 +1,586 @@ +//! 格式转换模块 +//! +//! 实现 Anthropic ↔ OpenAI 格式转换,用于 OpenRouter 支持 +//! 参考: anthropic-proxy-rs + +use crate::provider::Provider; +use crate::proxy::error::ProxyError; +use serde_json::{json, Value}; + +/// 从 Provider 配置中获取模型映射 +fn get_model_from_provider(model: &str, provider: &Provider) -> String { + let env = provider.settings_config.get("env"); + + // 根据请求的模型类型选择对应的配置模型 + let model_lower = model.to_lowercase(); + + if let Some(env) = env { + // 检查是否是 haiku 模型 + if model_lower.contains("haiku") { + if let Some(m) = env + .get("ANTHROPIC_DEFAULT_HAIKU_MODEL") + .and_then(|v| v.as_str()) + { + return m.to_string(); + } + } + // 检查是否是 opus 模型 + if model_lower.contains("opus") { + if let Some(m) = env + .get("ANTHROPIC_DEFAULT_OPUS_MODEL") + .and_then(|v| v.as_str()) + { + return m.to_string(); + } + } + // 检查是否是 sonnet 模型 + if model_lower.contains("sonnet") { + if let Some(m) = env + .get("ANTHROPIC_DEFAULT_SONNET_MODEL") + .and_then(|v| v.as_str()) + { + return m.to_string(); + } + } + // 默认使用 ANTHROPIC_MODEL + if let Some(m) = env.get("ANTHROPIC_MODEL").and_then(|v| v.as_str()) { + return m.to_string(); + } + } + + // 如果没有配置,返回原始模型名 + model.to_string() +} + +/// Anthropic 请求 → OpenAI 请求 +pub fn anthropic_to_openai(body: Value, provider: &Provider) -> Result { + let mut result = json!({}); + + // 模型映射:使用 Provider 配置中的模型 + if let Some(model) = body.get("model").and_then(|m| m.as_str()) { + let mapped_model = get_model_from_provider(model, provider); + result["model"] = json!(mapped_model); + } + + let mut messages = Vec::new(); + + // 处理 system prompt + if let Some(system) = body.get("system") { + if let Some(text) = system.as_str() { + // 单个字符串 + messages.push(json!({"role": "system", "content": text})); + } else if let Some(arr) = system.as_array() { + // 多个 system message + for msg in arr { + if let Some(text) = msg.get("text").and_then(|t| t.as_str()) { + messages.push(json!({"role": "system", "content": text})); + } + } + } + } + + // 转换 messages + if let Some(msgs) = body.get("messages").and_then(|m| m.as_array()) { + for msg in msgs { + let role = msg.get("role").and_then(|r| r.as_str()).unwrap_or("user"); + let content = msg.get("content"); + let converted = convert_message_to_openai(role, content)?; + messages.extend(converted); + } + } + + result["messages"] = json!(messages); + + // 转换参数 + if let Some(v) = body.get("max_tokens") { + result["max_tokens"] = v.clone(); + } + if let Some(v) = body.get("temperature") { + result["temperature"] = v.clone(); + } + if let Some(v) = body.get("top_p") { + result["top_p"] = v.clone(); + } + if let Some(v) = body.get("stop_sequences") { + result["stop"] = v.clone(); + } + if let Some(v) = body.get("stream") { + result["stream"] = v.clone(); + } + + // 转换 tools (过滤 BatchTool) + if let Some(tools) = body.get("tools").and_then(|t| t.as_array()) { + let openai_tools: Vec = tools + .iter() + .filter(|t| t.get("type").and_then(|v| v.as_str()) != Some("BatchTool")) + .map(|t| { + json!({ + "type": "function", + "function": { + "name": t.get("name").and_then(|n| n.as_str()).unwrap_or(""), + "description": t.get("description"), + "parameters": clean_schema(t.get("input_schema").cloned().unwrap_or(json!({}))) + } + }) + }) + .collect(); + + if !openai_tools.is_empty() { + result["tools"] = json!(openai_tools); + } + } + + if let Some(v) = body.get("tool_choice") { + result["tool_choice"] = v.clone(); + } + + Ok(result) +} + +/// 转换单条消息到 OpenAI 格式(可能产生多条消息) +fn convert_message_to_openai( + role: &str, + content: Option<&Value>, +) -> Result, ProxyError> { + let mut result = Vec::new(); + + let content = match content { + Some(c) => c, + None => { + result.push(json!({"role": role, "content": null})); + return Ok(result); + } + }; + + // 字符串内容 + if let Some(text) = content.as_str() { + result.push(json!({"role": role, "content": text})); + return Ok(result); + } + + // 数组内容(多模态/工具调用) + if let Some(blocks) = content.as_array() { + let mut content_parts = Vec::new(); + let mut tool_calls = Vec::new(); + + for block in blocks { + let block_type = block.get("type").and_then(|t| t.as_str()).unwrap_or(""); + + match block_type { + "text" => { + if let Some(text) = block.get("text").and_then(|t| t.as_str()) { + content_parts.push(json!({"type": "text", "text": text})); + } + } + "image" => { + if let Some(source) = block.get("source") { + let media_type = source + .get("media_type") + .and_then(|m| m.as_str()) + .unwrap_or("image/png"); + let data = source.get("data").and_then(|d| d.as_str()).unwrap_or(""); + content_parts.push(json!({ + "type": "image_url", + "image_url": {"url": format!("data:{};base64,{}", media_type, data)} + })); + } + } + "tool_use" => { + let id = block.get("id").and_then(|i| i.as_str()).unwrap_or(""); + let name = block.get("name").and_then(|n| n.as_str()).unwrap_or(""); + let input = block.get("input").cloned().unwrap_or(json!({})); + tool_calls.push(json!({ + "id": id, + "type": "function", + "function": { + "name": name, + "arguments": serde_json::to_string(&input).unwrap_or_default() + } + })); + } + "tool_result" => { + // tool_result 变成单独的 tool role 消息 + let tool_use_id = block + .get("tool_use_id") + .and_then(|i| i.as_str()) + .unwrap_or(""); + let content_val = block.get("content"); + let content_str = match content_val { + Some(Value::String(s)) => s.clone(), + Some(v) => serde_json::to_string(v).unwrap_or_default(), + None => String::new(), + }; + result.push(json!({ + "role": "tool", + "tool_call_id": tool_use_id, + "content": content_str + })); + } + "thinking" => { + // 跳过 thinking blocks + } + _ => {} + } + } + + // 添加带内容和/或工具调用的消息 + if !content_parts.is_empty() || !tool_calls.is_empty() { + let mut msg = json!({"role": role}); + + // 内容处理 + if content_parts.is_empty() { + msg["content"] = Value::Null; + } else if content_parts.len() == 1 { + if let Some(text) = content_parts[0].get("text") { + msg["content"] = text.clone(); + } else { + msg["content"] = json!(content_parts); + } + } else { + msg["content"] = json!(content_parts); + } + + // 工具调用 + if !tool_calls.is_empty() { + msg["tool_calls"] = json!(tool_calls); + } + + result.push(msg); + } + + return Ok(result); + } + + // 其他情况直接透传 + result.push(json!({"role": role, "content": content})); + Ok(result) +} + +/// 清理 JSON schema(移除不支持的 format) +fn clean_schema(mut schema: Value) -> Value { + if let Some(obj) = schema.as_object_mut() { + // 移除 "format": "uri" + if obj.get("format").and_then(|v| v.as_str()) == Some("uri") { + obj.remove("format"); + } + + // 递归清理嵌套 schema + if let Some(properties) = obj.get_mut("properties").and_then(|v| v.as_object_mut()) { + for (_, value) in properties.iter_mut() { + *value = clean_schema(value.clone()); + } + } + + if let Some(items) = obj.get_mut("items") { + *items = clean_schema(items.clone()); + } + } + schema +} + +/// OpenAI 响应 → Anthropic 响应 +pub fn openai_to_anthropic(body: Value) -> Result { + let choices = body + .get("choices") + .and_then(|c| c.as_array()) + .ok_or_else(|| ProxyError::TransformError("No choices in response".to_string()))?; + + let choice = choices + .first() + .ok_or_else(|| ProxyError::TransformError("Empty choices array".to_string()))?; + + let message = choice + .get("message") + .ok_or_else(|| ProxyError::TransformError("No message in choice".to_string()))?; + + let mut content = Vec::new(); + + // 文本内容 + if let Some(text) = message.get("content").and_then(|c| c.as_str()) { + if !text.is_empty() { + content.push(json!({"type": "text", "text": text})); + } + } + + // 工具调用 + if let Some(tool_calls) = message.get("tool_calls").and_then(|t| t.as_array()) { + for tc in tool_calls { + let id = tc.get("id").and_then(|i| i.as_str()).unwrap_or(""); + let empty_obj = json!({}); + let func = tc.get("function").unwrap_or(&empty_obj); + let name = func.get("name").and_then(|n| n.as_str()).unwrap_or(""); + let args_str = func + .get("arguments") + .and_then(|a| a.as_str()) + .unwrap_or("{}"); + let input: Value = serde_json::from_str(args_str).unwrap_or(json!({})); + + content.push(json!({ + "type": "tool_use", + "id": id, + "name": name, + "input": input + })); + } + } + + // 映射 finish_reason → stop_reason + let stop_reason = choice + .get("finish_reason") + .and_then(|r| r.as_str()) + .map(|r| match r { + "stop" => "end_turn", + "length" => "max_tokens", + "tool_calls" => "tool_use", + other => other, + }); + + // usage + let usage = body.get("usage").cloned().unwrap_or(json!({})); + let input_tokens = usage + .get("prompt_tokens") + .and_then(|v| v.as_u64()) + .unwrap_or(0) as u32; + let output_tokens = usage + .get("completion_tokens") + .and_then(|v| v.as_u64()) + .unwrap_or(0) as u32; + + let result = json!({ + "id": body.get("id").and_then(|i| i.as_str()).unwrap_or(""), + "type": "message", + "role": "assistant", + "content": content, + "model": body.get("model").and_then(|m| m.as_str()).unwrap_or(""), + "stop_reason": stop_reason, + "stop_sequence": null, + "usage": { + "input_tokens": input_tokens, + "output_tokens": output_tokens + } + }); + + Ok(result) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn create_provider(env_config: Value) -> Provider { + Provider { + id: "test".to_string(), + name: "Test Provider".to_string(), + settings_config: json!({"env": env_config}), + website_url: None, + category: None, + created_at: None, + sort_index: None, + notes: None, + meta: None, + icon: None, + icon_color: None, + is_proxy_target: None, + } + } + + fn create_openrouter_provider() -> Provider { + create_provider(json!({ + "ANTHROPIC_BASE_URL": "https://openrouter.ai/api", + "ANTHROPIC_MODEL": "anthropic/claude-sonnet-4.5", + "ANTHROPIC_DEFAULT_HAIKU_MODEL": "anthropic/claude-haiku-4.5", + "ANTHROPIC_DEFAULT_SONNET_MODEL": "anthropic/claude-sonnet-4.5", + "ANTHROPIC_DEFAULT_OPUS_MODEL": "anthropic/claude-opus-4.5" + })) + } + + #[test] + fn test_anthropic_to_openai_simple() { + let provider = create_openrouter_provider(); + let input = json!({ + "model": "claude-3-opus", + "max_tokens": 1024, + "messages": [{"role": "user", "content": "Hello"}] + }); + + let result = anthropic_to_openai(input, &provider).unwrap(); + // opus 模型映射到配置的 ANTHROPIC_DEFAULT_OPUS_MODEL + assert_eq!(result["model"], "anthropic/claude-opus-4.5"); + assert_eq!(result["max_tokens"], 1024); + assert_eq!(result["messages"][0]["role"], "user"); + assert_eq!(result["messages"][0]["content"], "Hello"); + } + + #[test] + fn test_anthropic_to_openai_with_system() { + let provider = create_openrouter_provider(); + let input = json!({ + "model": "claude-3-sonnet", + "max_tokens": 1024, + "system": "You are a helpful assistant.", + "messages": [{"role": "user", "content": "Hello"}] + }); + + let result = anthropic_to_openai(input, &provider).unwrap(); + assert_eq!(result["messages"][0]["role"], "system"); + assert_eq!( + result["messages"][0]["content"], + "You are a helpful assistant." + ); + assert_eq!(result["messages"][1]["role"], "user"); + } + + #[test] + fn test_anthropic_to_openai_with_tools() { + let provider = create_openrouter_provider(); + let input = json!({ + "model": "claude-3-opus", + "max_tokens": 1024, + "messages": [{"role": "user", "content": "What's the weather?"}], + "tools": [{ + "name": "get_weather", + "description": "Get weather info", + "input_schema": {"type": "object", "properties": {"location": {"type": "string"}}} + }] + }); + + let result = anthropic_to_openai(input, &provider).unwrap(); + assert_eq!(result["tools"][0]["type"], "function"); + assert_eq!(result["tools"][0]["function"]["name"], "get_weather"); + } + + #[test] + fn test_anthropic_to_openai_tool_use() { + let provider = create_openrouter_provider(); + let input = json!({ + "model": "claude-3-opus", + "max_tokens": 1024, + "messages": [{ + "role": "assistant", + "content": [ + {"type": "text", "text": "Let me check"}, + {"type": "tool_use", "id": "call_123", "name": "get_weather", "input": {"location": "Tokyo"}} + ] + }] + }); + + let result = anthropic_to_openai(input, &provider).unwrap(); + let msg = &result["messages"][0]; + assert_eq!(msg["role"], "assistant"); + assert!(msg.get("tool_calls").is_some()); + assert_eq!(msg["tool_calls"][0]["id"], "call_123"); + } + + #[test] + fn test_anthropic_to_openai_tool_result() { + let provider = create_openrouter_provider(); + let input = json!({ + "model": "claude-3-opus", + "max_tokens": 1024, + "messages": [{ + "role": "user", + "content": [ + {"type": "tool_result", "tool_use_id": "call_123", "content": "Sunny, 25°C"} + ] + }] + }); + + let result = anthropic_to_openai(input, &provider).unwrap(); + let msg = &result["messages"][0]; + assert_eq!(msg["role"], "tool"); + assert_eq!(msg["tool_call_id"], "call_123"); + assert_eq!(msg["content"], "Sunny, 25°C"); + } + + #[test] + fn test_openai_to_anthropic_simple() { + let input = json!({ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1234567890, + "model": "gpt-4", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello!"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15} + }); + + let result = openai_to_anthropic(input).unwrap(); + assert_eq!(result["id"], "chatcmpl-123"); + assert_eq!(result["type"], "message"); + assert_eq!(result["content"][0]["type"], "text"); + assert_eq!(result["content"][0]["text"], "Hello!"); + assert_eq!(result["stop_reason"], "end_turn"); + assert_eq!(result["usage"]["input_tokens"], 10); + assert_eq!(result["usage"]["output_tokens"], 5); + } + + #[test] + fn test_openai_to_anthropic_with_tool_calls() { + let input = json!({ + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1234567890, + "model": "gpt-4", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": null, + "tool_calls": [{ + "id": "call_123", + "type": "function", + "function": {"name": "get_weather", "arguments": "{\"location\": \"Tokyo\"}"} + }] + }, + "finish_reason": "tool_calls" + }], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15} + }); + + let result = openai_to_anthropic(input).unwrap(); + assert_eq!(result["content"][0]["type"], "tool_use"); + assert_eq!(result["content"][0]["id"], "call_123"); + assert_eq!(result["content"][0]["name"], "get_weather"); + assert_eq!(result["content"][0]["input"]["location"], "Tokyo"); + assert_eq!(result["stop_reason"], "tool_use"); + } + + #[test] + fn test_model_mapping_from_provider() { + let provider = create_openrouter_provider(); + + // sonnet 模型 + assert_eq!( + get_model_from_provider("claude-sonnet-4-5-20250929", &provider), + "anthropic/claude-sonnet-4.5" + ); + + // haiku 模型 + assert_eq!( + get_model_from_provider("claude-haiku-4-5-20250929", &provider), + "anthropic/claude-haiku-4.5" + ); + + // opus 模型 + assert_eq!( + get_model_from_provider("claude-opus-4-5", &provider), + "anthropic/claude-opus-4.5" + ); + } + + #[test] + fn test_anthropic_to_openai_model_mapping() { + let provider = create_openrouter_provider(); + let input = json!({ + "model": "claude-sonnet-4-5-20250929", + "max_tokens": 1024, + "messages": [{"role": "user", "content": "Hello"}] + }); + + let result = anthropic_to_openai(input, &provider).unwrap(); + assert_eq!(result["model"], "anthropic/claude-sonnet-4.5"); + } +} diff --git a/src-tauri/src/proxy/router.rs b/src-tauri/src/proxy/router.rs index 0f10bc9b3..47efa952b 100644 --- a/src-tauri/src/proxy/router.rs +++ b/src-tauri/src/proxy/router.rs @@ -1,6 +1,6 @@ //! Provider路由器 //! -//! 负责选择合适的Provider进行请求转发,支持健康检查和故障转移 +//! 负责选择合适的Provider进行请求转发 use super::ProxyError; use crate::{app_config::AppType, database::Database, provider::Provider}; @@ -15,137 +15,55 @@ impl ProviderRouter { Self { db } } - /// 选择Provider(带故障转移) - /// - /// 优先使用当前Provider,失败则尝试备用Provider + /// 选择Provider(只使用标记为代理目标的 Provider) pub async fn select_provider( &self, app_type: &AppType, - failed_ids: &[String], + _failed_ids: &[String], ) -> Result { - // 1. 尝试获取当前Provider - match self.get_current_provider(app_type, failed_ids).await { - Ok(provider) => return Ok(provider), - Err(e) => { - log::debug!("当前Provider不可用: {e:?}"); - } - } - - // 2. 尝试备用Provider - self.select_fallback(app_type, failed_ids).await - } - - /// 获取当前Provider - async fn get_current_provider( - &self, - app_type: &AppType, - failed_ids: &[String], - ) -> Result { - // 1. 尝试获取 Proxy Target Provider ID + // 1. 获取 Proxy Target Provider ID let proxy_target_id = self .db .get_proxy_target_provider(app_type.as_str()) .map_err(|e| ProxyError::DatabaseError(e.to_string()))?; - // 2. 获取 Current Provider ID (作为 fallback) - let current_id = self - .db - .get_current_provider(app_type.as_str()) - .map_err(|e| ProxyError::DatabaseError(e.to_string()))?; + let target_id = proxy_target_id.ok_or_else(|| { + log::warn!("[{}] 未设置代理目标 Provider", app_type.as_str()); + ProxyError::NoAvailableProvider + })?; - // 3. 确定使用的 ID (优先 proxy_target) - let target_id = proxy_target_id - .or(current_id) - .ok_or(ProxyError::NoAvailableProvider)?; - - // 4. 获取所有Provider + // 2. 获取所有 Provider let providers = self .db .get_all_providers(app_type.as_str()) .map_err(|e| ProxyError::DatabaseError(e.to_string()))?; - // 5. 找到目标Provider - let target = providers - .get(&target_id) - .ok_or(ProxyError::NoAvailableProvider)?; + // 3. 找到目标 Provider + let target = providers.get(&target_id).ok_or_else(|| { + log::warn!( + "[{}] 代理目标 Provider 不存在: {}", + app_type.as_str(), + target_id + ); + ProxyError::NoAvailableProvider + })?; - // 4. 检查是否在失败列表中 - if failed_ids.contains(&target.id) { - return Err(ProxyError::ProviderUnhealthy("Provider已失败".to_string())); - } - - // 5. 检查健康状态 - if self.is_provider_healthy(target, app_type).await { - Ok(target.clone()) - } else { - Err(ProxyError::ProviderUnhealthy(target.id.clone())) - } + log::info!( + "[{}] 使用代理目标 Provider: {}", + app_type.as_str(), + target.name + ); + Ok(target.clone()) } - /// 选择备用Provider - async fn select_fallback( - &self, - app_type: &AppType, - failed_ids: &[String], - ) -> Result { - let providers = self - .db - .get_all_providers(app_type.as_str()) - .map_err(|e| ProxyError::DatabaseError(e.to_string()))?; - - // 过滤失败的Provider,按sort_index排序 - let mut available: Vec<_> = providers - .into_values() - .filter(|p| !failed_ids.contains(&p.id)) - .collect(); - - available.sort_by_key(|p| p.sort_index.unwrap_or(9999)); - - // 寻找健康的Provider - for provider in available { - if self.is_provider_healthy(&provider, app_type).await { - log::info!("选择备用Provider: {}", provider.name); - return Ok(provider); - } - } - - log::warn!("无可用Provider"); - Err(ProxyError::NoAvailableProvider) - } - - /// 检查Provider是否健康 - async fn is_provider_healthy(&self, provider: &Provider, app_type: &AppType) -> bool { - // 从数据库查询健康状态 - match self - .db - .get_provider_health(&provider.id, app_type.as_str()) - .await - { - Ok(health) => { - // 连续失败3次以上视为不健康 - health.is_healthy && health.consecutive_failures < 3 - } - Err(_) => { - // 未记录状态时默认健康 - true - } - } - } - - /// 更新Provider健康状态 + /// 更新Provider健康状态(保留接口但不影响选择) pub async fn update_health( &self, - provider: &Provider, - app_type: &AppType, - success: bool, - error_msg: Option, + _provider: &Provider, + _app_type: &AppType, + _success: bool, + _error_msg: Option, ) { - if let Err(e) = self - .db - .update_provider_health(&provider.id, app_type.as_str(), success, error_msg) - .await - { - log::warn!("更新Provider健康状态失败: {e:?}"); - } + // 不再记录健康状态 } } diff --git a/src/App.tsx b/src/App.tsx index ed59b443c..562cf79cb 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -315,13 +315,14 @@ function App() { isLoading={isLoading} onSwitch={switchProvider} onSetProxyTarget={setProxyTarget} - onEdit={setEditingProvider} - onDelete={setConfirmDelete} - onDuplicate={handleDuplicateProvider} - onConfigureUsage={setUsageProvider} - onOpenWebsite={handleOpenWebsite} - onCreate={() => setIsAddOpen(true)} - /> + onEdit={setEditingProvider} + onDelete={setConfirmDelete} + onDuplicate={handleDuplicateProvider} + onConfigureUsage={setUsageProvider} + onOpenWebsite={handleOpenWebsite} + onCreate={() => setIsAddOpen(true)} + /> + ); diff --git a/src/config/claudeProviderPresets.ts b/src/config/claudeProviderPresets.ts index 23729c48d..bc998d0e5 100644 --- a/src/config/claudeProviderPresets.ts +++ b/src/config/claudeProviderPresets.ts @@ -365,4 +365,22 @@ export const providerPresets: ProviderPreset[] = [ partnerPromotionKey: "packycode", // 促销信息 i18n key icon: "packycode", }, + { + name: "OpenRouter", + websiteUrl: "https://openrouter.ai", + apiKeyUrl: "https://openrouter.ai/keys", + settingsConfig: { + env: { + ANTHROPIC_BASE_URL: "https://openrouter.ai/api", + ANTHROPIC_AUTH_TOKEN: "", + ANTHROPIC_MODEL: "anthropic/claude-sonnet-4.5", + ANTHROPIC_DEFAULT_HAIKU_MODEL: "anthropic/claude-haiku-4.5", + ANTHROPIC_DEFAULT_SONNET_MODEL: "anthropic/claude-sonnet-4.5", + ANTHROPIC_DEFAULT_OPUS_MODEL: "anthropic/claude-opus-4.5", + }, + }, + category: "aggregator", + icon: "openrouter", + iconColor: "#6366F1", + }, ];