mirror of
https://github.com/farion1231/cc-switch.git
synced 2026-07-25 13:45:03 +08:00
Feat/proxy server (#355)
* feat(proxy): implement local HTTP proxy server with multi-provider failover Add a complete HTTP proxy server implementation built on Axum framework, enabling local API request forwarding with automatic provider failover and load balancing capabilities. Backend Implementation (Rust): - Add proxy server module with 7 core components: * server.rs: Axum HTTP server lifecycle management (start/stop/status) * router.rs: API routing configuration for Claude/OpenAI/Gemini endpoints * handlers.rs: Request/response handling and transformation * forwarder.rs: Upstream forwarding logic with retry mechanism (652 lines) * error.rs: Comprehensive error handling and HTTP status mapping * types.rs: Shared types (ProxyConfig, ProxyStatus, ProxyServerInfo) * health.rs: Provider health check infrastructure Service Layer: - Add ProxyService (services/proxy.rs, 157 lines): * Manage proxy server lifecycle * Handle configuration updates * Track runtime status and metrics Database Layer: - Add proxy configuration DAO (dao/proxy.rs, 242 lines): * Persist proxy settings (listen address, port, timeout) * Store provider priority and availability flags - Update schema with proxy_config table (schema.rs): * Support runtime configuration persistence Tauri Commands: - Add 6 command endpoints (commands/proxy.rs): * start_proxy_server: Launch proxy server * stop_proxy_server: Gracefully shutdown server * get_proxy_status: Query runtime status * get_proxy_config: Retrieve current configuration * update_proxy_config: Modify settings without restart * is_proxy_running: Check server state Frontend Implementation (React + TypeScript): - Add ProxyPanel component (222 lines): * Real-time server status display * Start/stop controls * Provider availability monitoring - Add ProxySettingsDialog component (420 lines): * Configuration editor (address, port, timeout) * Provider priority management * Settings validation - Add React hooks: * useProxyConfig: Manage proxy configuration state * useProxyStatus: Poll and display server status - Add TypeScript types (types/proxy.ts): * Define ProxyConfig, ProxyStatus interfaces Provider Integration: - Extend Provider model with availability field (providers.rs): * Track provider health for failover logic - Update ProviderCard UI to display proxy status - Integrate proxy controls in Settings page Dependencies: - Add Axum 0.7 (async web framework) - Add Tower 0.4 (middleware and service abstractions) - Add Tower-HTTP (CORS layer) - Add Tokio sync primitives (oneshot, RwLock) Technical Details: - Graceful shutdown via oneshot channel - Shared state with Arc<RwLock<T>> for thread-safe config updates - CORS enabled for cross-origin frontend access - Request/response streaming support - Automatic retry with exponential backoff (forwarder) - API key extraction from multiple config formats (Claude/Codex/Gemini) File Statistics: - 41 files changed - 3491 insertions(+), 41 deletions(-) - Core modules: 1393 lines (server + forwarder + handlers) - Frontend UI: 642 lines (ProxyPanel + ProxySettingsDialog) - Database/DAO: 326 lines This implementation provides the foundation for advanced features like: - Multi-provider load balancing - Automatic failover on provider errors - Request logging and analytics - Usage tracking and cost monitoring * fix(proxy): resolve UI/UX issues and database constraint error Simplify proxy control interface and fix database persistence issues: Backend Fixes: - Fix NOT NULL constraint error in proxy_config.created_at field * Use COALESCE to preserve created_at on updates * Ensure proper INSERT OR REPLACE behavior - Remove redundant enabled field validation on startup * Auto-enable when user clicks start button * Persist enabled state after successful start - Preserve enabled state during config updates * Prevent accidental service shutdown on config save Frontend Improvements: - Remove duplicate proxy enable switch from settings dialog * Keep only runtime toggle in ProxyPanel * Simplify user experience with single control point - Hide proxy target button when proxy service is stopped * Add isProxyRunning prop to ProviderCard * Conditionally render proxy controls based on service status - Update form schema to omit enabled field * Managed automatically by backend Files: 5 changed, 81 insertions(+), 94 deletions(-) * fix(proxy): improve URL building and Gemini request handling - Refactor URL construction with version path deduplication (/v1, /v1beta) - Preserve query parameters for Gemini API requests - Support GOOGLE_GEMINI_API_KEY field name (with fallback) - Change default proxy port from 5000 to 15721 - Fix test: use Option type for is_proxy_target field * refactor(proxy): remove unused request handlers and routes - Remove unused GET/DELETE request forwarding methods - Remove count_tokens, get/delete response handlers - Simplify router by removing unused endpoints - Keep only essential routes: /v1/messages, /v1/responses, /v1beta/* * Merge branch 'main' into feat/proxy-server * fix(proxy): resolve clippy warnings for dead code and uninlined format args - Add #[allow(dead_code)] to unused ProviderUnhealthy variant - Inline format string arguments in handlers.rs and codex.rs log macros - Refactor error response handling to properly pass through upstream errors - Add URL deduplication logic for /v1/v1 paths in CodexAdapter * feat(proxy): implement provider adapter pattern with OpenRouter support 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 * feat(proxy): add streaming SSE transform and thinking parameter support New features: - Add OpenAI → Anthropic SSE streaming response transformation - Support thinking parameter detection for reasoning model selection - Add ANTHROPIC_REASONING_MODEL config option for extended thinking Changes: - streaming.rs: Implement SSE event parsing and Anthropic format conversion - transform.rs: Add thinking detection logic and reasoning model mapping - handlers.rs: Integrate streaming transform for OpenRouter compatibility - Cargo.toml: Add async-stream and bytes dependencies * feat(db): add usage tracking schema and types Add database tables for proxy request logs and model pricing. Extend Provider and error types to support usage statistics. * feat(proxy): implement usage tracking subsystem Add request logger with automatic cost calculation. Implement token parser for Claude/OpenAI/Gemini responses. Add cost calculator based on model pricing configuration. * feat(proxy): integrate usage logging into request handlers Add usage logging to forwarder and streaming handlers. Track token usage and costs for each proxy request. * feat(commands): add usage statistics Tauri commands Register usage commands for summary, trends, logs, and pricing. Expose usage stats service through Tauri command layer. * feat(api): add frontend usage API and query hooks Add TypeScript types for usage statistics. Implement usage API with Tauri invoke calls. Add TanStack Query hooks for usage data fetching. * feat(ui): add usage dashboard components Add UsageDashboard with summary cards, trend chart, and data tables. Implement model pricing configuration panel. Add request log viewer with filtering and detail panel. * fix(ui): integrate usage dashboard and fix type errors Add usage dashboard tab to settings page. Fix UsageScriptModal TypeScript type annotations. * deps: add recharts for charts and rust_decimal/uuid for usage tracking - recharts: Chart visualization for usage trends - rust_decimal: Precise cost calculations - uuid: Request ID generation * feat(proxy): add ProviderType enum for fine-grained provider detection Introduce ProviderType enum to distinguish between different provider implementations (Claude, ClaudeAuth, Codex, Gemini, GeminiCli, OpenRouter). This enables proper authentication handling and request transformation based on the actual provider type rather than just AppType. - Add ProviderType enum with detection logic from config - Enhance Claude adapter with OpenRouter detection - Enhance Gemini adapter with CLI mode detection - Add helper methods for provider type inference * feat(database): extend schema with streaming and timing fields Add new columns to proxy_request_logs table for enhanced usage tracking: - first_token_ms and duration_ms for performance metrics - provider_type and is_streaming for request classification - cost_multiplier for flexible pricing Update model pricing with accurate rates for Claude/GPT/Gemini models. Add ensure_model_pricing_seeded() call on database initialization. Add test for model pricing auto-seeding verification. * feat(proxy/usage): enhance token parser and logger for multi-format support Parser enhancements: - Add OpenAI Chat Completions format parsing (prompt_tokens/completion_tokens) - Add model field to TokenUsage for actual model name extraction - Add from_codex_response_adjusted() for proper cache token handling - Add debug logging for better stream event tracing Logger enhancements: - Add first_token_ms, provider_type, is_streaming, cost_multiplier fields - Extend RequestLog struct with full metadata tracking - Update log_with_calculation() signature for new fields Calculator: Update tests with model field in TokenUsage. * feat(proxy): enhance proxy server with session tracking and OpenAI route Error handling: - Add StreamIdleTimeout and AuthError variants for better error classification Module exports: - Export ResponseType, StreamHandler, NonStreamHandler from response_handler - Export ProxySession, ClientFormat from session module Server routing: - Add /v1/chat/completions route for OpenAI Chat Completions API Handlers: - Add log_usage_with_session() for enhanced usage tracking with session context - Add first_token_ms timing measurement for streaming responses - Use SseUsageCollector with start_time for accurate latency calculation - Track is_streaming flag in usage logs * feat(services): add pagination and enhanced filtering for request logs Usage stats service: - Change get_request_logs() from limit/offset to page/page_size pagination - Return PaginatedLogs with total count, page, and page_size - Add appType and providerName filters with LIKE search - Add is_streaming, first_token_ms, duration_ms to RequestLogDetail - Join with providers table for provider name lookup Commands: - Update get_request_logs command signature for pagination params Module exports: - Export PaginatedLogs struct * feat(frontend): update usage types and API for pagination support Types (usage.ts): - Add isStreaming, firstTokenMs, durationMs to RequestLog - Add PaginatedLogs interface with data, total, page, pageSize - Change LogFilters: providerId -> appType + providerName API (usage.ts): - Change getRequestLogs params from limit/offset to page/pageSize - Return PaginatedLogs instead of RequestLog[] - Pass filters object directly to backend Query (usage.ts): - Update usageKeys.logs key generation for pagination - Update useRequestLogs hook signature * refactor(ui): enhance RequestLogTable with filtering and pagination UI improvements: - Add filter bar with app type, provider name, model, status selectors - Add date range picker (startDate/endDate) - Add search/reset/refresh buttons Pagination: - Implement proper page-based pagination with page info display - Show total count and current page range - Add prev/next navigation buttons Features: - Default to last 24 hours filter - Streamlined table columns layout - Query invalidation on refresh * style(config): format mcpPresets code style Apply consistent formatting to createNpxCommand function and sequential-thinking server configuration. * fix(ui): update SettingsPage tab styles for improved appearance (#342) * feat(model-test): add provider model availability testing Implement standalone model testing feature to verify provider API connectivity: - Add ModelTestService for Claude/Codex/Gemini endpoint testing - Create model_test_logs table for test result persistence - Add test button to ProviderCard with loading state - Include ModelTestConfigPanel for customizing test parameters * fix(proxy): resolve token parsing for OpenRouter streaming responses Problem: - OpenRouter and similar third-party services return streaming responses where input_tokens appear in message_delta instead of message_start - The previous implementation only extracted input_tokens from message_start, causing input_tokens to be recorded as 0 for these providers Changes: - streaming.rs: Add prompt_tokens field to Usage struct and include input_tokens in the transformed message_delta event when converting OpenAI format to Anthropic format - parser.rs: Update from_claude_stream_events() to handle input_tokens from both message_start (native Claude API) and message_delta (OpenRouter) - Use if-let pattern instead of direct unwrap for safer parsing - Only update input_tokens from message_delta if not already set - logger.rs: Adjust test parameters to match updated function signature Tests: - Add test_openrouter_stream_parsing() for OpenRouter format validation - Add test_native_claude_stream_parsing() for native Claude API validation * fix(pricing): standardize model ID format for pricing lookup Normalize model IDs by removing vendor prefixes and converting dots to hyphens to ensure consistent pricing lookups across different API response formats. Changes: - Update seed data to use hyphen format (e.g., gpt-5-1, gemini-2-5-pro) - Add normalize_model_id() function to strip vendor prefixes (anthropic/, openai/) - Convert dots to hyphens in model IDs (claude-haiku-4.5 → claude-haiku-4-5) - Try both original and normalized IDs for exact matching - Use normalized ID for suffix-based fallback matching - Add comprehensive test cases for prefix and dot handling - Add warning log when no pricing found This ensures pricing lookups work correctly for: - Models with vendor prefixes: anthropic/claude-haiku-4.5 - Models with dots in version: claude-sonnet-4.5 - Models with date suffixes: claude-haiku-4-5-20240229 * style(rust): apply clippy formatting suggestions Apply automatic clippy fixes for uninlined_format_args warnings across Rust codebase. Replace format string placeholders with inline variable syntax for improved readability. Changes: - Convert format!("{}", var) to format!("{var}") - Apply to model_test.rs, parser.rs, and usage_stats.rs - Fix line length issues by breaking long function calls - Improve code formatting consistency All changes are automatic formatting with no functional impact. * fix(ui): restore card borders in usage statistics panels Restore proper card styling for ModelTestConfigPanel and PricingConfigPanel by adding back border and rounded-lg classes. The transparent background styling was causing visual inconsistency. Changes: - Replace border-none bg-transparent shadow-none with border rounded-lg - Apply to both loading and error states for consistency - Format TypeScript code for better readability - Break long function signatures across multiple lines This ensures the usage statistics panels have consistent visual appearance with proper borders and rounded corners. * feat(pricing): add GPT-5 Codex model pricing presets Add pricing configuration for GPT-5 Codex variants to support cost tracking for Codex-specific models. Changes: - Add gpt-5-codex model with standard GPT-5 pricing - Add gpt-5-1-codex model with standard GPT-5.1 pricing - Input: $1.25/M tokens, Output: $10/M tokens - Cache read: $0.125/M tokens, Cache creation: $0 This ensures accurate cost calculation for Codex API requests using GPT-5 Codex models.
This commit is contained in:
@@ -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<String, ProxyError> {
|
||||
/// // 从 provider 配置中提取 base_url
|
||||
/// }
|
||||
///
|
||||
/// fn extract_auth(&self, provider: &Provider) -> Option<AuthInfo> {
|
||||
/// // 从 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<String, ProxyError>;
|
||||
|
||||
/// 从 Provider 配置中提取认证信息
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `provider` - Provider 配置
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Some(AuthInfo)` - 提取到的认证信息
|
||||
/// * `None` - 未找到认证信息
|
||||
fn extract_auth(&self, provider: &Provider) -> Option<AuthInfo>;
|
||||
|
||||
/// 构建请求 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<Value, ProxyError> {
|
||||
Ok(body)
|
||||
}
|
||||
|
||||
/// 转换响应体
|
||||
///
|
||||
/// 将响应体从一种格式转换为另一种格式(如 OpenAI → Anthropic)。
|
||||
/// 默认实现直接返回原始响应体(透传)。
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `body` - 原始响应体
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(Value)` - 转换后的响应体
|
||||
/// * `Err(ProxyError)` - 转换失败
|
||||
///
|
||||
/// Note: 响应转换将在 handler 层集成,目前预留接口
|
||||
#[allow(dead_code)]
|
||||
fn transform_response(&self, body: Value) -> Result<Value, ProxyError> {
|
||||
Ok(body)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,210 @@
|
||||
//! Authentication Types
|
||||
//!
|
||||
//! 定义认证信息和认证策略,支持多种上游供应商的认证方式。
|
||||
|
||||
/// 认证信息
|
||||
///
|
||||
/// 包含 API Key 和对应的认证策略
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct AuthInfo {
|
||||
/// API Key
|
||||
pub api_key: String,
|
||||
/// 认证策略
|
||||
pub strategy: AuthStrategy,
|
||||
/// OAuth access_token(用于 GoogleOAuth 策略)
|
||||
pub access_token: Option<String>,
|
||||
}
|
||||
|
||||
impl AuthInfo {
|
||||
/// 创建新的认证信息
|
||||
pub fn new(api_key: String, strategy: AuthStrategy) -> Self {
|
||||
Self {
|
||||
api_key,
|
||||
strategy,
|
||||
access_token: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建带有 access_token 的认证信息(用于 OAuth)
|
||||
pub fn with_access_token(api_key: String, access_token: String) -> Self {
|
||||
Self {
|
||||
api_key,
|
||||
strategy: AuthStrategy::GoogleOAuth,
|
||||
access_token: Some(access_token),
|
||||
}
|
||||
}
|
||||
|
||||
/// 返回遮蔽后的 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()
|
||||
}
|
||||
}
|
||||
|
||||
/// 返回遮蔽后的 access_token(用于日志输出)
|
||||
#[allow(dead_code)]
|
||||
pub fn masked_access_token(&self) -> Option<String> {
|
||||
self.access_token.as_ref().map(|token| {
|
||||
if token.len() > 8 {
|
||||
format!("{}...{}", &token[..4], &token[token.len() - 4..])
|
||||
} else {
|
||||
"***".to_string()
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// 认证策略
|
||||
///
|
||||
/// 不同供应商使用不同的认证方式
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum AuthStrategy {
|
||||
/// Anthropic 认证方式
|
||||
/// - Header: `x-api-key: <api_key>`
|
||||
/// - Header: `anthropic-version: 2023-06-01`
|
||||
Anthropic,
|
||||
|
||||
/// Claude 中转服务认证方式(仅 Bearer,无 x-api-key)
|
||||
///
|
||||
/// - Header: `Authorization: Bearer <api_key>`
|
||||
///
|
||||
/// 用于不支持 x-api-key 的中转服务
|
||||
ClaudeAuth,
|
||||
|
||||
/// Bearer Token 认证方式(OpenAI 等)
|
||||
///
|
||||
/// - Header: `Authorization: Bearer <api_key>`
|
||||
Bearer,
|
||||
|
||||
/// Google API Key 认证方式
|
||||
///
|
||||
/// - Header: `x-goog-api-key: <api_key>`
|
||||
Google,
|
||||
|
||||
/// Google OAuth 认证方式
|
||||
///
|
||||
/// - Header: `Authorization: Bearer <access_token>`
|
||||
///
|
||||
/// 用于 Gemini CLI 等需要 OAuth 的场景
|
||||
GoogleOAuth,
|
||||
}
|
||||
|
||||
#[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);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_auth_info_new_has_no_access_token() {
|
||||
let auth = AuthInfo::new("api-key".to_string(), AuthStrategy::Bearer);
|
||||
assert!(auth.access_token.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_auth_info_with_access_token() {
|
||||
let auth = AuthInfo::with_access_token(
|
||||
"refresh-token".to_string(),
|
||||
"ya29.access-token-12345".to_string(),
|
||||
);
|
||||
assert_eq!(auth.api_key, "refresh-token");
|
||||
assert_eq!(auth.strategy, AuthStrategy::GoogleOAuth);
|
||||
assert_eq!(
|
||||
auth.access_token,
|
||||
Some("ya29.access-token-12345".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_masked_access_token_long() {
|
||||
let auth =
|
||||
AuthInfo::with_access_token("refresh".to_string(), "ya29.1234567890abcdef".to_string());
|
||||
assert_eq!(auth.masked_access_token(), Some("ya29...cdef".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_masked_access_token_short() {
|
||||
let auth = AuthInfo::with_access_token("refresh".to_string(), "short".to_string());
|
||||
assert_eq!(auth.masked_access_token(), Some("***".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_masked_access_token_none() {
|
||||
let auth = AuthInfo::new("api-key".to_string(), AuthStrategy::Bearer);
|
||||
assert!(auth.masked_access_token().is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_claude_auth_strategy() {
|
||||
let auth = AuthInfo::new("sk-test".to_string(), AuthStrategy::ClaudeAuth);
|
||||
assert_eq!(auth.strategy, AuthStrategy::ClaudeAuth);
|
||||
assert_ne!(auth.strategy, AuthStrategy::Anthropic);
|
||||
assert_ne!(auth.strategy, AuthStrategy::Bearer);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_google_oauth_strategy() {
|
||||
let auth = AuthInfo::new("refresh-token".to_string(), AuthStrategy::GoogleOAuth);
|
||||
assert_eq!(auth.strategy, AuthStrategy::GoogleOAuth);
|
||||
assert_ne!(auth.strategy, AuthStrategy::Google);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_all_strategies_are_distinct() {
|
||||
let strategies = [
|
||||
AuthStrategy::Anthropic,
|
||||
AuthStrategy::ClaudeAuth,
|
||||
AuthStrategy::Bearer,
|
||||
AuthStrategy::Google,
|
||||
AuthStrategy::GoogleOAuth,
|
||||
];
|
||||
|
||||
for (i, s1) in strategies.iter().enumerate() {
|
||||
for (j, s2) in strategies.iter().enumerate() {
|
||||
if i == j {
|
||||
assert_eq!(s1, s2);
|
||||
} else {
|
||||
assert_ne!(s1, s2);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,403 @@
|
||||
//! Claude (Anthropic) Provider Adapter
|
||||
//!
|
||||
//! 支持透传模式和 OpenRouter 兼容模式
|
||||
//!
|
||||
//! ## 认证模式
|
||||
//! - **Claude**: Anthropic 官方 API (x-api-key + anthropic-version)
|
||||
//! - **ClaudeAuth**: 中转服务 (仅 Bearer 认证,无 x-api-key)
|
||||
//! - **OpenRouter**: 需要 Anthropic ↔ OpenAI 格式转换
|
||||
|
||||
use super::{AuthInfo, AuthStrategy, ProviderAdapter, ProviderType};
|
||||
use crate::provider::Provider;
|
||||
use crate::proxy::error::ProxyError;
|
||||
use reqwest::RequestBuilder;
|
||||
|
||||
/// Claude 适配器
|
||||
pub struct ClaudeAdapter;
|
||||
|
||||
impl ClaudeAdapter {
|
||||
pub fn new() -> Self {
|
||||
Self
|
||||
}
|
||||
|
||||
/// 获取供应商类型
|
||||
///
|
||||
/// 根据 base_url 和 auth_mode 检测具体的供应商类型:
|
||||
/// - OpenRouter: base_url 包含 openrouter.ai
|
||||
/// - ClaudeAuth: auth_mode 为 bearer_only
|
||||
/// - Claude: 默认 Anthropic 官方
|
||||
pub fn provider_type(&self, provider: &Provider) -> ProviderType {
|
||||
// 检测 OpenRouter
|
||||
if let Ok(base_url) = self.extract_base_url(provider) {
|
||||
if base_url.contains("openrouter.ai") {
|
||||
return ProviderType::OpenRouter;
|
||||
}
|
||||
}
|
||||
|
||||
// 检测 ClaudeAuth (仅 Bearer 认证)
|
||||
if self.is_bearer_only_mode(provider) {
|
||||
return ProviderType::ClaudeAuth;
|
||||
}
|
||||
|
||||
ProviderType::Claude
|
||||
}
|
||||
|
||||
/// 检测是否使用 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
|
||||
}
|
||||
|
||||
/// 检测是否为仅 Bearer 认证模式
|
||||
fn is_bearer_only_mode(&self, provider: &Provider) -> bool {
|
||||
// 检查 settings_config 中的 auth_mode
|
||||
if let Some(auth_mode) = provider
|
||||
.settings_config
|
||||
.get("auth_mode")
|
||||
.and_then(|v| v.as_str())
|
||||
{
|
||||
if auth_mode == "bearer_only" {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
// 检查 env 中的 AUTH_MODE
|
||||
if let Some(env) = provider.settings_config.get("env") {
|
||||
if let Some(auth_mode) = env.get("AUTH_MODE").and_then(|v| v.as_str()) {
|
||||
if auth_mode == "bearer_only" {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
false
|
||||
}
|
||||
|
||||
/// 从 Provider 配置中提取 API Key
|
||||
fn extract_key(&self, provider: &Provider) -> Option<String> {
|
||||
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<String, ProxyError> {
|
||||
// 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<AuthInfo> {
|
||||
let provider_type = self.provider_type(provider);
|
||||
let strategy = match provider_type {
|
||||
ProviderType::OpenRouter => AuthStrategy::Bearer,
|
||||
ProviderType::ClaudeAuth => AuthStrategy::ClaudeAuth,
|
||||
_ => 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 {
|
||||
// Anthropic 官方: Authorization Bearer + x-api-key + anthropic-version
|
||||
AuthStrategy::Anthropic => request
|
||||
.header("Authorization", format!("Bearer {}", auth.api_key))
|
||||
.header("x-api-key", &auth.api_key)
|
||||
.header("anthropic-version", "2023-06-01"),
|
||||
// ClaudeAuth 中转服务: 仅 Bearer,无 x-api-key
|
||||
AuthStrategy::ClaudeAuth => {
|
||||
request.header("Authorization", format!("Bearer {}", auth.api_key))
|
||||
}
|
||||
// OpenRouter: Bearer
|
||||
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<serde_json::Value, ProxyError> {
|
||||
super::transform::anthropic_to_openai(body, provider)
|
||||
}
|
||||
|
||||
fn transform_response(&self, body: serde_json::Value) -> Result<serde_json::Value, ProxyError> {
|
||||
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_extract_auth_claude_auth_mode() {
|
||||
let adapter = ClaudeAdapter::new();
|
||||
let provider = create_provider(json!({
|
||||
"env": {
|
||||
"ANTHROPIC_BASE_URL": "https://some-proxy.com",
|
||||
"ANTHROPIC_AUTH_TOKEN": "sk-proxy-key"
|
||||
},
|
||||
"auth_mode": "bearer_only"
|
||||
}));
|
||||
|
||||
let auth = adapter.extract_auth(&provider).unwrap();
|
||||
assert_eq!(auth.api_key, "sk-proxy-key");
|
||||
assert_eq!(auth.strategy, AuthStrategy::ClaudeAuth);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_auth_claude_auth_env_mode() {
|
||||
let adapter = ClaudeAdapter::new();
|
||||
let provider = create_provider(json!({
|
||||
"env": {
|
||||
"ANTHROPIC_BASE_URL": "https://some-proxy.com",
|
||||
"ANTHROPIC_AUTH_TOKEN": "sk-proxy-key",
|
||||
"AUTH_MODE": "bearer_only"
|
||||
}
|
||||
}));
|
||||
|
||||
let auth = adapter.extract_auth(&provider).unwrap();
|
||||
assert_eq!(auth.api_key, "sk-proxy-key");
|
||||
assert_eq!(auth.strategy, AuthStrategy::ClaudeAuth);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_type_detection() {
|
||||
let adapter = ClaudeAdapter::new();
|
||||
|
||||
// Anthropic 官方
|
||||
let anthropic = create_provider(json!({
|
||||
"env": {
|
||||
"ANTHROPIC_BASE_URL": "https://api.anthropic.com",
|
||||
"ANTHROPIC_AUTH_TOKEN": "sk-ant-test"
|
||||
}
|
||||
}));
|
||||
assert_eq!(adapter.provider_type(&anthropic), ProviderType::Claude);
|
||||
|
||||
// OpenRouter
|
||||
let openrouter = create_provider(json!({
|
||||
"env": {
|
||||
"ANTHROPIC_BASE_URL": "https://openrouter.ai/api",
|
||||
"OPENROUTER_API_KEY": "sk-or-test"
|
||||
}
|
||||
}));
|
||||
assert_eq!(adapter.provider_type(&openrouter), ProviderType::OpenRouter);
|
||||
|
||||
// ClaudeAuth
|
||||
let claude_auth = create_provider(json!({
|
||||
"env": {
|
||||
"ANTHROPIC_BASE_URL": "https://some-proxy.com",
|
||||
"ANTHROPIC_AUTH_TOKEN": "sk-test"
|
||||
},
|
||||
"auth_mode": "bearer_only"
|
||||
}));
|
||||
assert_eq!(
|
||||
adapter.provider_type(&claude_auth),
|
||||
ProviderType::ClaudeAuth
|
||||
);
|
||||
}
|
||||
|
||||
#[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));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,265 @@
|
||||
//! Codex (OpenAI) Provider Adapter
|
||||
//!
|
||||
//! 仅透传模式,支持直连 OpenAI API
|
||||
//!
|
||||
//! ## 客户端检测
|
||||
//! 支持检测官方 Codex 客户端 (codex_vscode, codex_cli_rs)
|
||||
|
||||
use super::{AuthInfo, AuthStrategy, ProviderAdapter};
|
||||
use crate::provider::Provider;
|
||||
use crate::proxy::error::ProxyError;
|
||||
use regex::Regex;
|
||||
use reqwest::RequestBuilder;
|
||||
use std::sync::LazyLock;
|
||||
|
||||
/// 官方 Codex 客户端 User-Agent 正则
|
||||
#[allow(dead_code)]
|
||||
static CODEX_CLIENT_REGEX: LazyLock<Regex> =
|
||||
LazyLock::new(|| Regex::new(r"^(codex_vscode|codex_cli_rs)/[\d.]+").unwrap());
|
||||
|
||||
/// Codex 适配器
|
||||
pub struct CodexAdapter;
|
||||
|
||||
impl CodexAdapter {
|
||||
pub fn new() -> Self {
|
||||
Self
|
||||
}
|
||||
|
||||
/// 检测是否为官方 Codex 客户端
|
||||
///
|
||||
/// 匹配 User-Agent 模式: `^(codex_vscode|codex_cli_rs)/[\d.]+`
|
||||
#[allow(dead_code)]
|
||||
pub fn is_official_client(user_agent: &str) -> bool {
|
||||
CODEX_CLIENT_REGEX.is_match(user_agent)
|
||||
}
|
||||
|
||||
/// 从 Provider 配置中提取 API Key
|
||||
fn extract_key(&self, provider: &Provider) -> Option<String> {
|
||||
// 1. 尝试从 env 中获取
|
||||
if let Some(env) = provider.settings_config.get("env") {
|
||||
if let Some(key) = env.get("OPENAI_API_KEY").and_then(|v| v.as_str()) {
|
||||
return Some(key.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
// 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());
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 尝试直接获取
|
||||
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());
|
||||
}
|
||||
|
||||
// 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());
|
||||
}
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for CodexAdapter {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl ProviderAdapter for CodexAdapter {
|
||||
fn name(&self) -> &'static str {
|
||||
"Codex"
|
||||
}
|
||||
|
||||
fn extract_base_url(&self, provider: &Provider) -> Result<String, ProxyError> {
|
||||
// 1. 尝试直接获取 base_url 字段
|
||||
if let Some(url) = provider
|
||||
.settings_config
|
||||
.get("base_url")
|
||||
.and_then(|v| v.as_str())
|
||||
{
|
||||
return Ok(url.trim_end_matches('/').to_string());
|
||||
}
|
||||
|
||||
// 2. 尝试 baseURL
|
||||
if let Some(url) = provider
|
||||
.settings_config
|
||||
.get("baseURL")
|
||||
.and_then(|v| v.as_str())
|
||||
{
|
||||
return Ok(url.trim_end_matches('/').to_string());
|
||||
}
|
||||
|
||||
// 3. 尝试从 config 对象中获取
|
||||
if let Some(config) = provider.settings_config.get("config") {
|
||||
if let Some(url) = config.get("base_url").and_then(|v| v.as_str()) {
|
||||
return Ok(url.trim_end_matches('/').to_string());
|
||||
}
|
||||
|
||||
// 尝试解析 TOML 字符串格式
|
||||
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('"') {
|
||||
return Ok(rest[..end].trim_end_matches('/').to_string());
|
||||
}
|
||||
}
|
||||
if let Some(start) = config_str.find("base_url = '") {
|
||||
let rest = &config_str[start + 12..];
|
||||
if let Some(end) = rest.find('\'') {
|
||||
return Ok(rest[..end].trim_end_matches('/').to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Err(ProxyError::ConfigError(
|
||||
"Codex Provider 缺少 base_url 配置".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
fn extract_auth(&self, provider: &Provider) -> Option<AuthInfo> {
|
||||
self.extract_key(provider)
|
||||
.map(|key| AuthInfo::new(key, AuthStrategy::Bearer))
|
||||
}
|
||||
|
||||
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}");
|
||||
|
||||
// 去除重复的 /v1/v1
|
||||
if url.contains("/v1/v1") {
|
||||
url = url.replace("/v1/v1", "/v1");
|
||||
}
|
||||
|
||||
url
|
||||
}
|
||||
|
||||
fn add_auth_headers(&self, request: RequestBuilder, auth: &AuthInfo) -> RequestBuilder {
|
||||
request.header("Authorization", format!("Bearer {}", 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 Codex".to_string(),
|
||||
settings_config: config,
|
||||
website_url: None,
|
||||
category: Some("codex".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_direct() {
|
||||
let adapter = CodexAdapter::new();
|
||||
let provider = create_provider(json!({
|
||||
"base_url": "https://api.openai.com/v1"
|
||||
}));
|
||||
|
||||
let url = adapter.extract_base_url(&provider).unwrap();
|
||||
assert_eq!(url, "https://api.openai.com/v1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_auth_from_auth_field() {
|
||||
let adapter = CodexAdapter::new();
|
||||
let provider = create_provider(json!({
|
||||
"auth": {
|
||||
"OPENAI_API_KEY": "sk-test-key-12345678"
|
||||
}
|
||||
}));
|
||||
|
||||
let auth = adapter.extract_auth(&provider).unwrap();
|
||||
assert_eq!(auth.api_key, "sk-test-key-12345678");
|
||||
assert_eq!(auth.strategy, AuthStrategy::Bearer);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_auth_from_env() {
|
||||
let adapter = CodexAdapter::new();
|
||||
let provider = create_provider(json!({
|
||||
"env": {
|
||||
"OPENAI_API_KEY": "sk-env-key-12345678"
|
||||
}
|
||||
}));
|
||||
|
||||
let auth = adapter.extract_auth(&provider).unwrap();
|
||||
assert_eq!(auth.api_key, "sk-env-key-12345678");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_url() {
|
||||
let adapter = CodexAdapter::new();
|
||||
let url = adapter.build_url("https://api.openai.com/v1", "/responses");
|
||||
assert_eq!(url, "https://api.openai.com/v1/responses");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_url_dedup_v1() {
|
||||
let adapter = CodexAdapter::new();
|
||||
// base_url 已包含 /v1,endpoint 也包含 /v1
|
||||
let url = adapter.build_url("https://www.packyapi.com/v1", "/v1/responses");
|
||||
assert_eq!(url, "https://www.packyapi.com/v1/responses");
|
||||
}
|
||||
|
||||
// 官方客户端检测测试
|
||||
#[test]
|
||||
fn test_is_official_client_vscode() {
|
||||
assert!(CodexAdapter::is_official_client("codex_vscode/1.0.0"));
|
||||
assert!(CodexAdapter::is_official_client("codex_vscode/2.3.4"));
|
||||
assert!(CodexAdapter::is_official_client("codex_vscode/0.1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_official_client_cli() {
|
||||
assert!(CodexAdapter::is_official_client("codex_cli_rs/1.0.0"));
|
||||
assert!(CodexAdapter::is_official_client("codex_cli_rs/0.5.2"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_not_official_client() {
|
||||
assert!(!CodexAdapter::is_official_client("Mozilla/5.0"));
|
||||
assert!(!CodexAdapter::is_official_client("curl/7.68.0"));
|
||||
assert!(!CodexAdapter::is_official_client("python-requests/2.25.1"));
|
||||
assert!(!CodexAdapter::is_official_client("codex_other/1.0.0"));
|
||||
assert!(!CodexAdapter::is_official_client(""));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_official_client_partial_match() {
|
||||
// 必须从开头匹配
|
||||
assert!(!CodexAdapter::is_official_client("some codex_vscode/1.0.0"));
|
||||
assert!(!CodexAdapter::is_official_client(
|
||||
"prefix_codex_cli_rs/1.0.0"
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,426 @@
|
||||
//! Gemini (Google) Provider Adapter
|
||||
//!
|
||||
//! 支持 API Key 和 OAuth 两种认证方式
|
||||
//!
|
||||
//! ## 认证模式
|
||||
//! - **Gemini**: API Key 认证 (x-goog-api-key)
|
||||
//! - **GeminiCli**: OAuth Bearer 认证 (用于 Gemini CLI)
|
||||
|
||||
use super::{AuthInfo, AuthStrategy, ProviderAdapter, ProviderType};
|
||||
use crate::provider::Provider;
|
||||
use crate::proxy::error::ProxyError;
|
||||
use reqwest::RequestBuilder;
|
||||
|
||||
/// Gemini 适配器
|
||||
pub struct GeminiAdapter;
|
||||
|
||||
/// OAuth 凭证结构
|
||||
#[derive(Debug, Clone)]
|
||||
#[allow(dead_code)]
|
||||
pub struct OAuthCredentials {
|
||||
pub access_token: String,
|
||||
pub refresh_token: Option<String>,
|
||||
pub client_id: Option<String>,
|
||||
pub client_secret: Option<String>,
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
impl OAuthCredentials {
|
||||
/// 检查是否需要刷新 token(有 refresh_token 但没有有效的 access_token)
|
||||
pub fn needs_refresh(&self) -> bool {
|
||||
self.refresh_token.is_some() && self.access_token.is_empty()
|
||||
}
|
||||
|
||||
/// 检查是否可以刷新 token
|
||||
pub fn can_refresh(&self) -> bool {
|
||||
self.refresh_token.is_some() && self.client_id.is_some() && self.client_secret.is_some()
|
||||
}
|
||||
}
|
||||
|
||||
impl GeminiAdapter {
|
||||
pub fn new() -> Self {
|
||||
Self
|
||||
}
|
||||
|
||||
/// 获取供应商类型
|
||||
///
|
||||
/// 根据 API Key 格式检测:
|
||||
/// - GeminiCli: access_token (ya29. 开头) 或 JSON 格式凭证
|
||||
/// - Gemini: 普通 API Key
|
||||
pub fn provider_type(&self, provider: &Provider) -> ProviderType {
|
||||
if let Some(key) = self.extract_key_raw(provider) {
|
||||
// OAuth access_token 以 ya29. 开头
|
||||
if key.starts_with("ya29.") {
|
||||
return ProviderType::GeminiCli;
|
||||
}
|
||||
// JSON 格式的 OAuth 凭证
|
||||
if key.starts_with('{') {
|
||||
return ProviderType::GeminiCli;
|
||||
}
|
||||
}
|
||||
ProviderType::Gemini
|
||||
}
|
||||
|
||||
/// 检测认证类型
|
||||
pub fn detect_auth_type(&self, provider: &Provider) -> AuthStrategy {
|
||||
match self.provider_type(provider) {
|
||||
ProviderType::GeminiCli => AuthStrategy::GoogleOAuth,
|
||||
_ => AuthStrategy::Google,
|
||||
}
|
||||
}
|
||||
|
||||
/// 解析 OAuth 凭证
|
||||
pub fn parse_oauth_credentials(&self, key: &str) -> Option<OAuthCredentials> {
|
||||
// 直接是 access_token
|
||||
if key.starts_with("ya29.") {
|
||||
return Some(OAuthCredentials {
|
||||
access_token: key.to_string(),
|
||||
refresh_token: None,
|
||||
client_id: None,
|
||||
client_secret: None,
|
||||
});
|
||||
}
|
||||
|
||||
// JSON 格式
|
||||
if key.starts_with('{') {
|
||||
if let Ok(json) = serde_json::from_str::<serde_json::Value>(key) {
|
||||
let access_token = json
|
||||
.get("access_token")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string())
|
||||
.unwrap_or_default();
|
||||
let refresh_token = json
|
||||
.get("refresh_token")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
let client_id = json
|
||||
.get("client_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
let client_secret = json
|
||||
.get("client_secret")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
|
||||
// 如果有 access_token 或 refresh_token,返回凭证
|
||||
if !access_token.is_empty() || refresh_token.is_some() {
|
||||
return Some(OAuthCredentials {
|
||||
access_token,
|
||||
refresh_token,
|
||||
client_id,
|
||||
client_secret,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
None
|
||||
}
|
||||
|
||||
/// 从 Provider 配置中提取原始 API Key
|
||||
fn extract_key_raw(&self, provider: &Provider) -> Option<String> {
|
||||
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<String, ProxyError> {
|
||||
// 从 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<AuthInfo> {
|
||||
let key = self.extract_key_raw(provider)?;
|
||||
let strategy = self.detect_auth_type(provider);
|
||||
|
||||
match strategy {
|
||||
AuthStrategy::GoogleOAuth => {
|
||||
// 解析 OAuth 凭证
|
||||
if let Some(creds) = self.parse_oauth_credentials(&key) {
|
||||
Some(AuthInfo::with_access_token(key, creds.access_token))
|
||||
} else {
|
||||
// 回退到普通 API Key
|
||||
Some(AuthInfo::new(key, AuthStrategy::Google))
|
||||
}
|
||||
}
|
||||
_ => Some(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 {
|
||||
match auth.strategy {
|
||||
// OAuth Bearer 认证
|
||||
AuthStrategy::GoogleOAuth => {
|
||||
let token = auth.access_token.as_ref().unwrap_or(&auth.api_key);
|
||||
request
|
||||
.header("Authorization", format!("Bearer {token}"))
|
||||
.header("x-goog-api-client", "GeminiCLI/1.0")
|
||||
}
|
||||
// API Key 认证
|
||||
_ => 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_api_key() {
|
||||
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);
|
||||
assert!(auth.access_token.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_auth_oauth_access_token() {
|
||||
let adapter = GeminiAdapter::new();
|
||||
let provider = create_provider(json!({
|
||||
"env": {
|
||||
"GOOGLE_GEMINI_API_KEY": "ya29.test-access-token-12345"
|
||||
}
|
||||
}));
|
||||
|
||||
let auth = adapter.extract_auth(&provider).unwrap();
|
||||
assert_eq!(auth.strategy, AuthStrategy::GoogleOAuth);
|
||||
assert_eq!(
|
||||
auth.access_token,
|
||||
Some("ya29.test-access-token-12345".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_auth_oauth_json() {
|
||||
let adapter = GeminiAdapter::new();
|
||||
let provider = create_provider(json!({
|
||||
"env": {
|
||||
"GOOGLE_GEMINI_API_KEY": "{\"access_token\":\"ya29.test-token\",\"refresh_token\":\"1//refresh\"}"
|
||||
}
|
||||
}));
|
||||
|
||||
let auth = adapter.extract_auth(&provider).unwrap();
|
||||
assert_eq!(auth.strategy, AuthStrategy::GoogleOAuth);
|
||||
assert_eq!(auth.access_token, Some("ya29.test-token".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_type_detection() {
|
||||
let adapter = GeminiAdapter::new();
|
||||
|
||||
// API Key
|
||||
let api_key_provider = create_provider(json!({
|
||||
"env": {
|
||||
"GOOGLE_GEMINI_API_KEY": "AIza-test-key"
|
||||
}
|
||||
}));
|
||||
assert_eq!(
|
||||
adapter.provider_type(&api_key_provider),
|
||||
ProviderType::Gemini
|
||||
);
|
||||
|
||||
// OAuth access_token
|
||||
let oauth_provider = create_provider(json!({
|
||||
"env": {
|
||||
"GOOGLE_GEMINI_API_KEY": "ya29.test-token"
|
||||
}
|
||||
}));
|
||||
assert_eq!(
|
||||
adapter.provider_type(&oauth_provider),
|
||||
ProviderType::GeminiCli
|
||||
);
|
||||
|
||||
// OAuth JSON
|
||||
let oauth_json_provider = create_provider(json!({
|
||||
"env": {
|
||||
"GOOGLE_GEMINI_API_KEY": "{\"access_token\":\"ya29.test\"}"
|
||||
}
|
||||
}));
|
||||
assert_eq!(
|
||||
adapter.provider_type(&oauth_json_provider),
|
||||
ProviderType::GeminiCli
|
||||
);
|
||||
}
|
||||
|
||||
#[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"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_oauth_credentials_direct_token() {
|
||||
let adapter = GeminiAdapter::new();
|
||||
let creds = adapter
|
||||
.parse_oauth_credentials("ya29.test-access-token")
|
||||
.unwrap();
|
||||
assert_eq!(creds.access_token, "ya29.test-access-token");
|
||||
assert!(creds.refresh_token.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_oauth_credentials_json() {
|
||||
let adapter = GeminiAdapter::new();
|
||||
let creds = adapter
|
||||
.parse_oauth_credentials(
|
||||
"{\"access_token\":\"ya29.test\",\"refresh_token\":\"1//refresh\"}",
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(creds.access_token, "ya29.test");
|
||||
assert_eq!(creds.refresh_token, Some("1//refresh".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_oauth_credentials_invalid() {
|
||||
let adapter = GeminiAdapter::new();
|
||||
assert!(adapter.parse_oauth_credentials("AIza-api-key").is_none());
|
||||
assert!(adapter.parse_oauth_credentials("invalid-json{").is_none());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,424 @@
|
||||
//! 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 streaming;
|
||||
pub mod transform;
|
||||
|
||||
use crate::app_config::AppType;
|
||||
use crate::provider::Provider;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
// 公开导出
|
||||
pub use adapter::ProviderAdapter;
|
||||
pub use auth::{AuthInfo, AuthStrategy};
|
||||
pub use claude::ClaudeAdapter;
|
||||
pub use codex::CodexAdapter;
|
||||
pub use gemini::GeminiAdapter;
|
||||
|
||||
/// 供应商类型枚举
|
||||
///
|
||||
/// 区分不同供应商的具体实现方式,决定认证和请求处理逻辑。
|
||||
/// 比 AppType 更细粒度,支持同一 AppType 下的多种变体。
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ProviderType {
|
||||
/// Anthropic 官方 API (x-api-key + anthropic-version)
|
||||
Claude,
|
||||
/// Claude 中转服务 (仅 Bearer 认证,无 x-api-key)
|
||||
ClaudeAuth,
|
||||
/// OpenAI Codex Response API
|
||||
Codex,
|
||||
/// Google Gemini API (x-goog-api-key)
|
||||
Gemini,
|
||||
/// Google Gemini CLI (OAuth Bearer)
|
||||
GeminiCli,
|
||||
/// OpenRouter (需要 Anthropic ↔ OpenAI 格式转换)
|
||||
OpenRouter,
|
||||
}
|
||||
|
||||
impl ProviderType {
|
||||
/// 是否需要格式转换
|
||||
///
|
||||
/// OpenRouter 需要将 Anthropic 格式转换为 OpenAI 格式
|
||||
#[allow(dead_code)]
|
||||
pub fn needs_transform(&self) -> bool {
|
||||
matches!(self, ProviderType::OpenRouter)
|
||||
}
|
||||
|
||||
/// 获取默认端点
|
||||
#[allow(dead_code)]
|
||||
pub fn default_endpoint(&self) -> &'static str {
|
||||
match self {
|
||||
ProviderType::Claude | ProviderType::ClaudeAuth => "https://api.anthropic.com",
|
||||
ProviderType::Codex => "https://api.openai.com",
|
||||
ProviderType::Gemini | ProviderType::GeminiCli => {
|
||||
"https://generativelanguage.googleapis.com"
|
||||
}
|
||||
ProviderType::OpenRouter => "https://openrouter.ai/api",
|
||||
}
|
||||
}
|
||||
|
||||
/// 从 AppType 和 Provider 配置推断供应商类型
|
||||
///
|
||||
/// 根据配置中的 base_url、auth_mode、api_key 格式等信息推断具体的供应商类型
|
||||
#[allow(dead_code)]
|
||||
pub fn from_app_type_and_config(app_type: &AppType, provider: &Provider) -> Self {
|
||||
match app_type {
|
||||
AppType::Claude => {
|
||||
// 检测是否为 OpenRouter
|
||||
let adapter = ClaudeAdapter::new();
|
||||
if let Ok(base_url) = adapter.extract_base_url(provider) {
|
||||
if base_url.contains("openrouter.ai") {
|
||||
return ProviderType::OpenRouter;
|
||||
}
|
||||
}
|
||||
// 检测是否为中转服务(仅 Bearer 认证)
|
||||
// 注意:ProviderMeta 没有直接的 auth_mode 字段,
|
||||
// 我们通过检查 settings_config 中的配置来判断
|
||||
// 检查 settings_config 中的 auth_mode
|
||||
if let Some(auth_mode) = provider
|
||||
.settings_config
|
||||
.get("auth_mode")
|
||||
.and_then(|v| v.as_str())
|
||||
{
|
||||
if auth_mode == "bearer_only" {
|
||||
return ProviderType::ClaudeAuth;
|
||||
}
|
||||
}
|
||||
// 检查 env 中的 auth_mode
|
||||
if let Some(env) = provider.settings_config.get("env") {
|
||||
if let Some(auth_mode) = env.get("AUTH_MODE").and_then(|v| v.as_str()) {
|
||||
if auth_mode == "bearer_only" {
|
||||
return ProviderType::ClaudeAuth;
|
||||
}
|
||||
}
|
||||
}
|
||||
ProviderType::Claude
|
||||
}
|
||||
AppType::Codex => ProviderType::Codex,
|
||||
AppType::Gemini => {
|
||||
// 检测是否为 CLI 模式(OAuth)
|
||||
let adapter = GeminiAdapter::new();
|
||||
if let Some(auth) = adapter.extract_auth(provider) {
|
||||
let key = &auth.api_key;
|
||||
// OAuth access_token 以 ya29. 开头
|
||||
if key.starts_with("ya29.") {
|
||||
return ProviderType::GeminiCli;
|
||||
}
|
||||
// JSON 格式的 OAuth 凭证
|
||||
if key.starts_with('{') {
|
||||
return ProviderType::GeminiCli;
|
||||
}
|
||||
}
|
||||
ProviderType::Gemini
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 转换为字符串表示
|
||||
pub fn as_str(&self) -> &'static str {
|
||||
match self {
|
||||
ProviderType::Claude => "claude",
|
||||
ProviderType::ClaudeAuth => "claude_auth",
|
||||
ProviderType::Codex => "codex",
|
||||
ProviderType::Gemini => "gemini",
|
||||
ProviderType::GeminiCli => "gemini_cli",
|
||||
ProviderType::OpenRouter => "openrouter",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for ProviderType {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
write!(f, "{}", self.as_str())
|
||||
}
|
||||
}
|
||||
|
||||
impl std::str::FromStr for ProviderType {
|
||||
type Err = String;
|
||||
|
||||
fn from_str(s: &str) -> Result<Self, Self::Err> {
|
||||
match s.to_lowercase().as_str() {
|
||||
"claude" => Ok(ProviderType::Claude),
|
||||
"claude_auth" | "claude-auth" => Ok(ProviderType::ClaudeAuth),
|
||||
"codex" => Ok(ProviderType::Codex),
|
||||
"gemini" => Ok(ProviderType::Gemini),
|
||||
"gemini_cli" | "gemini-cli" => Ok(ProviderType::GeminiCli),
|
||||
"openrouter" => Ok(ProviderType::OpenRouter),
|
||||
_ => Err(format!("Invalid provider type: {s}")),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 根据 AppType 获取对应的适配器
|
||||
pub fn get_adapter(app_type: &AppType) -> Box<dyn ProviderAdapter> {
|
||||
match app_type {
|
||||
AppType::Claude => Box::new(ClaudeAdapter::new()),
|
||||
AppType::Codex => Box::new(CodexAdapter::new()),
|
||||
AppType::Gemini => Box::new(GeminiAdapter::new()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 根据 ProviderType 获取对应的适配器
|
||||
#[allow(dead_code)]
|
||||
pub fn get_adapter_for_provider_type(provider_type: &ProviderType) -> Box<dyn ProviderAdapter> {
|
||||
match provider_type {
|
||||
ProviderType::Claude | ProviderType::ClaudeAuth | ProviderType::OpenRouter => {
|
||||
Box::new(ClaudeAdapter::new())
|
||||
}
|
||||
ProviderType::Codex => Box::new(CodexAdapter::new()),
|
||||
ProviderType::Gemini | ProviderType::GeminiCli => Box::new(GeminiAdapter::new()),
|
||||
}
|
||||
}
|
||||
|
||||
#[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 Provider".to_string(),
|
||||
settings_config: 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,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_type_needs_transform() {
|
||||
assert!(!ProviderType::Claude.needs_transform());
|
||||
assert!(!ProviderType::ClaudeAuth.needs_transform());
|
||||
assert!(!ProviderType::Codex.needs_transform());
|
||||
assert!(!ProviderType::Gemini.needs_transform());
|
||||
assert!(!ProviderType::GeminiCli.needs_transform());
|
||||
assert!(ProviderType::OpenRouter.needs_transform());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_type_default_endpoint() {
|
||||
assert_eq!(
|
||||
ProviderType::Claude.default_endpoint(),
|
||||
"https://api.anthropic.com"
|
||||
);
|
||||
assert_eq!(
|
||||
ProviderType::ClaudeAuth.default_endpoint(),
|
||||
"https://api.anthropic.com"
|
||||
);
|
||||
assert_eq!(
|
||||
ProviderType::Codex.default_endpoint(),
|
||||
"https://api.openai.com"
|
||||
);
|
||||
assert_eq!(
|
||||
ProviderType::Gemini.default_endpoint(),
|
||||
"https://generativelanguage.googleapis.com"
|
||||
);
|
||||
assert_eq!(
|
||||
ProviderType::GeminiCli.default_endpoint(),
|
||||
"https://generativelanguage.googleapis.com"
|
||||
);
|
||||
assert_eq!(
|
||||
ProviderType::OpenRouter.default_endpoint(),
|
||||
"https://openrouter.ai/api"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_type_from_str() {
|
||||
assert_eq!(
|
||||
"claude".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::Claude
|
||||
);
|
||||
assert_eq!(
|
||||
"claude_auth".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::ClaudeAuth
|
||||
);
|
||||
assert_eq!(
|
||||
"claude-auth".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::ClaudeAuth
|
||||
);
|
||||
assert_eq!(
|
||||
"codex".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::Codex
|
||||
);
|
||||
assert_eq!(
|
||||
"gemini".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::Gemini
|
||||
);
|
||||
assert_eq!(
|
||||
"gemini_cli".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::GeminiCli
|
||||
);
|
||||
assert_eq!(
|
||||
"gemini-cli".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::GeminiCli
|
||||
);
|
||||
assert_eq!(
|
||||
"openrouter".parse::<ProviderType>().unwrap(),
|
||||
ProviderType::OpenRouter
|
||||
);
|
||||
assert!("invalid".parse::<ProviderType>().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_type_as_str() {
|
||||
assert_eq!(ProviderType::Claude.as_str(), "claude");
|
||||
assert_eq!(ProviderType::ClaudeAuth.as_str(), "claude_auth");
|
||||
assert_eq!(ProviderType::Codex.as_str(), "codex");
|
||||
assert_eq!(ProviderType::Gemini.as_str(), "gemini");
|
||||
assert_eq!(ProviderType::GeminiCli.as_str(), "gemini_cli");
|
||||
assert_eq!(ProviderType::OpenRouter.as_str(), "openrouter");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_provider_type_serde() {
|
||||
// Test serialization
|
||||
let claude = ProviderType::Claude;
|
||||
let serialized = serde_json::to_string(&claude).unwrap();
|
||||
assert_eq!(serialized, "\"claude\"");
|
||||
|
||||
let claude_auth = ProviderType::ClaudeAuth;
|
||||
let serialized = serde_json::to_string(&claude_auth).unwrap();
|
||||
assert_eq!(serialized, "\"claude_auth\"");
|
||||
|
||||
// Test deserialization
|
||||
let deserialized: ProviderType = serde_json::from_str("\"claude\"").unwrap();
|
||||
assert_eq!(deserialized, ProviderType::Claude);
|
||||
|
||||
let deserialized: ProviderType = serde_json::from_str("\"gemini_cli\"").unwrap();
|
||||
assert_eq!(deserialized, ProviderType::GeminiCli);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_from_app_type_claude_direct() {
|
||||
let provider = create_provider(json!({
|
||||
"env": {
|
||||
"ANTHROPIC_BASE_URL": "https://api.anthropic.com",
|
||||
"ANTHROPIC_AUTH_TOKEN": "sk-ant-test"
|
||||
}
|
||||
}));
|
||||
|
||||
let provider_type = ProviderType::from_app_type_and_config(&AppType::Claude, &provider);
|
||||
assert_eq!(provider_type, ProviderType::Claude);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_from_app_type_claude_openrouter() {
|
||||
let provider = create_provider(json!({
|
||||
"env": {
|
||||
"ANTHROPIC_BASE_URL": "https://openrouter.ai/api",
|
||||
"OPENROUTER_API_KEY": "sk-or-test"
|
||||
}
|
||||
}));
|
||||
|
||||
let provider_type = ProviderType::from_app_type_and_config(&AppType::Claude, &provider);
|
||||
assert_eq!(provider_type, ProviderType::OpenRouter);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_from_app_type_claude_auth() {
|
||||
let provider = create_provider(json!({
|
||||
"env": {
|
||||
"ANTHROPIC_BASE_URL": "https://some-proxy.com",
|
||||
"ANTHROPIC_AUTH_TOKEN": "sk-test"
|
||||
},
|
||||
"auth_mode": "bearer_only"
|
||||
}));
|
||||
|
||||
let provider_type = ProviderType::from_app_type_and_config(&AppType::Claude, &provider);
|
||||
assert_eq!(provider_type, ProviderType::ClaudeAuth);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_from_app_type_codex() {
|
||||
let provider = create_provider(json!({
|
||||
"env": {
|
||||
"OPENAI_API_KEY": "sk-test"
|
||||
}
|
||||
}));
|
||||
|
||||
let provider_type = ProviderType::from_app_type_and_config(&AppType::Codex, &provider);
|
||||
assert_eq!(provider_type, ProviderType::Codex);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_from_app_type_gemini_api_key() {
|
||||
let provider = create_provider(json!({
|
||||
"env": {
|
||||
"GOOGLE_GEMINI_API_KEY": "AIza-test-key"
|
||||
}
|
||||
}));
|
||||
|
||||
let provider_type = ProviderType::from_app_type_and_config(&AppType::Gemini, &provider);
|
||||
assert_eq!(provider_type, ProviderType::Gemini);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_from_app_type_gemini_cli_oauth() {
|
||||
let provider = create_provider(json!({
|
||||
"env": {
|
||||
"GOOGLE_GEMINI_API_KEY": "ya29.test-access-token"
|
||||
}
|
||||
}));
|
||||
|
||||
let provider_type = ProviderType::from_app_type_and_config(&AppType::Gemini, &provider);
|
||||
assert_eq!(provider_type, ProviderType::GeminiCli);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_from_app_type_gemini_cli_json() {
|
||||
let provider = create_provider(json!({
|
||||
"env": {
|
||||
"GOOGLE_GEMINI_API_KEY": "{\"access_token\":\"ya29.test\",\"refresh_token\":\"1//test\"}"
|
||||
}
|
||||
}));
|
||||
|
||||
let provider_type = ProviderType::from_app_type_and_config(&AppType::Gemini, &provider);
|
||||
assert_eq!(provider_type, ProviderType::GeminiCli);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_get_adapter_for_provider_type() {
|
||||
let adapter = get_adapter_for_provider_type(&ProviderType::Claude);
|
||||
assert_eq!(adapter.name(), "Claude");
|
||||
|
||||
let adapter = get_adapter_for_provider_type(&ProviderType::ClaudeAuth);
|
||||
assert_eq!(adapter.name(), "Claude");
|
||||
|
||||
let adapter = get_adapter_for_provider_type(&ProviderType::OpenRouter);
|
||||
assert_eq!(adapter.name(), "Claude");
|
||||
|
||||
let adapter = get_adapter_for_provider_type(&ProviderType::Codex);
|
||||
assert_eq!(adapter.name(), "Codex");
|
||||
|
||||
let adapter = get_adapter_for_provider_type(&ProviderType::Gemini);
|
||||
assert_eq!(adapter.name(), "Gemini");
|
||||
|
||||
let adapter = get_adapter_for_provider_type(&ProviderType::GeminiCli);
|
||||
assert_eq!(adapter.name(), "Gemini");
|
||||
}
|
||||
}
|
||||
@@ -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<AnthropicMessage>,
|
||||
pub max_tokens: u32,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub system: Option<Value>, // 可以是 String 或 Vec<SystemBlock>
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub temperature: Option<f32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub stream: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tools: Option<Vec<AnthropicTool>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tool_choice: Option<Value>,
|
||||
}
|
||||
|
||||
/// Anthropic 消息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct AnthropicMessage {
|
||||
pub role: String,
|
||||
pub content: Value, // String 或 Vec<ContentBlock>
|
||||
}
|
||||
|
||||
/// 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<String>,
|
||||
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<AnthropicResponseContent>,
|
||||
pub model: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub stop_reason: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub stop_sequence: Option<String>,
|
||||
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,
|
||||
}
|
||||
@@ -0,0 +1,6 @@
|
||||
//! API 数据模型
|
||||
//!
|
||||
//! 定义 Anthropic 和 OpenAI API 的请求/响应结构
|
||||
|
||||
pub mod anthropic;
|
||||
pub mod openai;
|
||||
@@ -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<OpenAIMessage>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub max_tokens: Option<u32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub temperature: Option<f32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub stream: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tools: Option<Vec<OpenAITool>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tool_choice: Option<Value>,
|
||||
}
|
||||
|
||||
/// OpenAI 消息
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct OpenAIMessage {
|
||||
pub role: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub content: Option<Value>, // String 或 Vec<ContentPart>
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tool_calls: Option<Vec<OpenAIToolCall>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tool_call_id: Option<String>,
|
||||
}
|
||||
|
||||
/// 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<String>,
|
||||
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<OpenAIChoice>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub usage: Option<OpenAIUsage>,
|
||||
}
|
||||
|
||||
/// 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<String>,
|
||||
}
|
||||
|
||||
/// OpenAI 使用量
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct OpenAIUsage {
|
||||
pub prompt_tokens: u32,
|
||||
pub completion_tokens: u32,
|
||||
pub total_tokens: u32,
|
||||
}
|
||||
@@ -0,0 +1,338 @@
|
||||
//! 流式响应转换模块
|
||||
//!
|
||||
//! 实现 OpenAI SSE → Anthropic SSE 格式转换
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures::stream::{Stream, StreamExt};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::json;
|
||||
|
||||
/// OpenAI 流式响应数据结构
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct OpenAIStreamChunk {
|
||||
id: String,
|
||||
model: String,
|
||||
choices: Vec<StreamChoice>,
|
||||
#[serde(default)]
|
||||
usage: Option<Usage>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct StreamChoice {
|
||||
delta: Delta,
|
||||
#[serde(default)]
|
||||
finish_reason: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct Delta {
|
||||
#[serde(default)]
|
||||
content: Option<String>,
|
||||
#[serde(default)]
|
||||
reasoning: Option<String>, // OpenRouter 的推理内容
|
||||
#[serde(default)]
|
||||
tool_calls: Option<Vec<DeltaToolCall>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize, Serialize)]
|
||||
struct DeltaToolCall {
|
||||
index: usize,
|
||||
#[serde(default)]
|
||||
id: Option<String>,
|
||||
#[serde(rename = "type", default)]
|
||||
call_type: Option<String>,
|
||||
#[serde(default)]
|
||||
function: Option<DeltaFunction>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize, Serialize)]
|
||||
struct DeltaFunction {
|
||||
#[serde(default)]
|
||||
name: Option<String>,
|
||||
#[serde(default)]
|
||||
arguments: Option<String>,
|
||||
}
|
||||
|
||||
/// OpenAI 流式响应的 usage 信息(完整版)
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct Usage {
|
||||
#[serde(default)]
|
||||
prompt_tokens: u32,
|
||||
#[serde(default)]
|
||||
completion_tokens: u32,
|
||||
}
|
||||
|
||||
/// 创建 Anthropic SSE 流
|
||||
pub fn create_anthropic_sse_stream(
|
||||
stream: impl Stream<Item = Result<Bytes, reqwest::Error>> + Send + 'static,
|
||||
) -> impl Stream<Item = Result<Bytes, std::io::Error>> + Send {
|
||||
async_stream::stream! {
|
||||
let mut buffer = String::new();
|
||||
let mut message_id = None;
|
||||
let mut current_model = None;
|
||||
let mut content_index = 0;
|
||||
let mut has_sent_message_start = false;
|
||||
let mut current_block_type: Option<String> = None;
|
||||
let mut tool_call_id = None;
|
||||
|
||||
log::info!("[Claude/OpenRouter] ====== 开始流式响应转换 ======");
|
||||
|
||||
tokio::pin!(stream);
|
||||
|
||||
while let Some(chunk) = stream.next().await {
|
||||
match chunk {
|
||||
Ok(bytes) => {
|
||||
let text = String::from_utf8_lossy(&bytes);
|
||||
buffer.push_str(&text);
|
||||
|
||||
while let Some(pos) = buffer.find("\n\n") {
|
||||
let line = buffer[..pos].to_string();
|
||||
buffer = buffer[pos + 2..].to_string();
|
||||
|
||||
if line.trim().is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
for l in line.lines() {
|
||||
if let Some(data) = l.strip_prefix("data: ") {
|
||||
if data.trim() == "[DONE]" {
|
||||
log::info!("[Claude/OpenRouter] <<< OpenAI SSE: [DONE]");
|
||||
let event = json!({"type": "message_stop"});
|
||||
let sse_data = format!("event: message_stop\ndata: {}\n\n",
|
||||
serde_json::to_string(&event).unwrap_or_default());
|
||||
log::info!("[Claude/OpenRouter] >>> Anthropic SSE: message_stop");
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
continue;
|
||||
}
|
||||
|
||||
if let Ok(chunk) = serde_json::from_str::<OpenAIStreamChunk>(data) {
|
||||
// 记录原始 OpenAI 事件(格式化显示)
|
||||
if let Ok(json_value) = serde_json::from_str::<serde_json::Value>(data) {
|
||||
log::info!(
|
||||
"[Claude/OpenRouter] <<< OpenAI SSE 事件:\n{}",
|
||||
serde_json::to_string_pretty(&json_value).unwrap_or_else(|_| data.to_string())
|
||||
);
|
||||
} else {
|
||||
log::info!("[Claude/OpenRouter] <<< OpenAI SSE 数据: {data}");
|
||||
}
|
||||
|
||||
if message_id.is_none() {
|
||||
message_id = Some(chunk.id.clone());
|
||||
}
|
||||
if current_model.is_none() {
|
||||
current_model = Some(chunk.model.clone());
|
||||
}
|
||||
|
||||
if let Some(choice) = chunk.choices.first() {
|
||||
if !has_sent_message_start {
|
||||
let event = json!({
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": message_id.clone().unwrap_or_default(),
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": current_model.clone().unwrap_or_default(),
|
||||
"usage": {
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0
|
||||
}
|
||||
}
|
||||
});
|
||||
let sse_data = format!("event: message_start\ndata: {}\n\n",
|
||||
serde_json::to_string(&event).unwrap_or_default());
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
has_sent_message_start = true;
|
||||
}
|
||||
|
||||
// 处理 reasoning(thinking)
|
||||
if let Some(reasoning) = &choice.delta.reasoning {
|
||||
if current_block_type.is_none() {
|
||||
let event = json!({
|
||||
"type": "content_block_start",
|
||||
"index": content_index,
|
||||
"content_block": {
|
||||
"type": "thinking",
|
||||
"thinking": ""
|
||||
}
|
||||
});
|
||||
let sse_data = format!("event: content_block_start\ndata: {}\n\n",
|
||||
serde_json::to_string(&event).unwrap_or_default());
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
current_block_type = Some("thinking".to_string());
|
||||
}
|
||||
|
||||
let event = json!({
|
||||
"type": "content_block_delta",
|
||||
"index": content_index,
|
||||
"delta": {
|
||||
"type": "thinking_delta",
|
||||
"thinking": reasoning
|
||||
}
|
||||
});
|
||||
let sse_data = format!("event: content_block_delta\ndata: {}\n\n",
|
||||
serde_json::to_string(&event).unwrap_or_default());
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
}
|
||||
|
||||
// 处理文本内容
|
||||
if let Some(content) = &choice.delta.content {
|
||||
if !content.is_empty() {
|
||||
if current_block_type.as_deref() != Some("text") {
|
||||
if current_block_type.is_some() {
|
||||
let event = json!({
|
||||
"type": "content_block_stop",
|
||||
"index": content_index
|
||||
});
|
||||
let sse_data = format!("event: content_block_stop\ndata: {}\n\n",
|
||||
serde_json::to_string(&event).unwrap_or_default());
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
content_index += 1;
|
||||
}
|
||||
|
||||
let event = json!({
|
||||
"type": "content_block_start",
|
||||
"index": content_index,
|
||||
"content_block": {
|
||||
"type": "text",
|
||||
"text": ""
|
||||
}
|
||||
});
|
||||
let sse_data = format!("event: content_block_start\ndata: {}\n\n",
|
||||
serde_json::to_string(&event).unwrap_or_default());
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
current_block_type = Some("text".to_string());
|
||||
}
|
||||
|
||||
let event = json!({
|
||||
"type": "content_block_delta",
|
||||
"index": content_index,
|
||||
"delta": {
|
||||
"type": "text_delta",
|
||||
"text": content
|
||||
}
|
||||
});
|
||||
let sse_data = format!("event: content_block_delta\ndata: {}\n\n",
|
||||
serde_json::to_string(&event).unwrap_or_default());
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
}
|
||||
}
|
||||
|
||||
// 处理工具调用
|
||||
if let Some(tool_calls) = &choice.delta.tool_calls {
|
||||
for tool_call in tool_calls {
|
||||
if let Some(id) = &tool_call.id {
|
||||
if current_block_type.is_some() {
|
||||
let event = json!({
|
||||
"type": "content_block_stop",
|
||||
"index": content_index
|
||||
});
|
||||
let sse_data = format!("event: content_block_stop\ndata: {}\n\n",
|
||||
serde_json::to_string(&event).unwrap_or_default());
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
content_index += 1;
|
||||
}
|
||||
|
||||
tool_call_id = Some(id.clone());
|
||||
}
|
||||
|
||||
if let Some(function) = &tool_call.function {
|
||||
if let Some(name) = &function.name {
|
||||
let event = json!({
|
||||
"type": "content_block_start",
|
||||
"index": content_index,
|
||||
"content_block": {
|
||||
"type": "tool_use",
|
||||
"id": tool_call_id.clone().unwrap_or_default(),
|
||||
"name": name
|
||||
}
|
||||
});
|
||||
let sse_data = format!("event: content_block_start\ndata: {}\n\n",
|
||||
serde_json::to_string(&event).unwrap_or_default());
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
current_block_type = Some("tool_use".to_string());
|
||||
}
|
||||
|
||||
if let Some(args) = &function.arguments {
|
||||
let event = json!({
|
||||
"type": "content_block_delta",
|
||||
"index": content_index,
|
||||
"delta": {
|
||||
"type": "input_json_delta",
|
||||
"partial_json": args
|
||||
}
|
||||
});
|
||||
let sse_data = format!("event: content_block_delta\ndata: {}\n\n",
|
||||
serde_json::to_string(&event).unwrap_or_default());
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 处理 finish_reason
|
||||
if let Some(finish_reason) = &choice.finish_reason {
|
||||
if current_block_type.is_some() {
|
||||
let event = json!({
|
||||
"type": "content_block_stop",
|
||||
"index": content_index
|
||||
});
|
||||
let sse_data = format!("event: content_block_stop\ndata: {}\n\n",
|
||||
serde_json::to_string(&event).unwrap_or_default());
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
}
|
||||
|
||||
let stop_reason = map_stop_reason(Some(finish_reason));
|
||||
// 构建 usage 信息,包含 input_tokens 和 output_tokens
|
||||
let usage_json = chunk.usage.as_ref().map(|u| json!({
|
||||
"input_tokens": u.prompt_tokens,
|
||||
"output_tokens": u.completion_tokens
|
||||
}));
|
||||
let event = json!({
|
||||
"type": "message_delta",
|
||||
"delta": {
|
||||
"stop_reason": stop_reason,
|
||||
"stop_sequence": null
|
||||
},
|
||||
"usage": usage_json
|
||||
});
|
||||
let sse_data = format!("event: message_delta\ndata: {}\n\n",
|
||||
serde_json::to_string(&event).unwrap_or_default());
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
log::error!("Stream error: {e}");
|
||||
let error_event = json!({
|
||||
"type": "error",
|
||||
"error": {
|
||||
"type": "stream_error",
|
||||
"message": format!("Stream error: {e}")
|
||||
}
|
||||
});
|
||||
let sse_data = format!("event: error\ndata: {}\n\n",
|
||||
serde_json::to_string(&error_event).unwrap_or_default());
|
||||
yield Ok(Bytes::from(sse_data));
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 映射停止原因
|
||||
fn map_stop_reason(finish_reason: Option<&str>) -> Option<String> {
|
||||
finish_reason.map(|r| {
|
||||
match r {
|
||||
"tool_calls" => "tool_use",
|
||||
"stop" => "end_turn",
|
||||
"length" => "max_tokens",
|
||||
_ => "end_turn",
|
||||
}
|
||||
.to_string()
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,640 @@
|
||||
//! 格式转换模块
|
||||
//!
|
||||
//! 实现 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, body: &Value) -> String {
|
||||
let env = provider.settings_config.get("env");
|
||||
let model_lower = model.to_lowercase();
|
||||
|
||||
// 检测 thinking 参数
|
||||
let has_thinking = body
|
||||
.get("thinking")
|
||||
.and_then(|v| v.as_object())
|
||||
.and_then(|o| o.get("type"))
|
||||
.and_then(|t| t.as_str())
|
||||
== Some("enabled");
|
||||
|
||||
if let Some(env) = env {
|
||||
// 如果启用 thinking,优先使用推理模型
|
||||
if has_thinking {
|
||||
if let Some(m) = env
|
||||
.get("ANTHROPIC_REASONING_MODEL")
|
||||
.and_then(|v| v.as_str())
|
||||
{
|
||||
log::debug!("[Transform] 使用推理模型: {m}");
|
||||
return m.to_string();
|
||||
}
|
||||
}
|
||||
|
||||
// 根据模型类型选择配置模型
|
||||
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();
|
||||
}
|
||||
}
|
||||
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();
|
||||
}
|
||||
}
|
||||
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<Value, ProxyError> {
|
||||
let mut result = json!({});
|
||||
|
||||
// 模型映射:使用 Provider 配置中的模型(支持 thinking 参数)
|
||||
if let Some(model) = body.get("model").and_then(|m| m.as_str()) {
|
||||
let mapped_model = get_model_from_provider(model, provider, &body);
|
||||
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<Value> = 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<Vec<Value>, 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<Value, ProxyError> {
|
||||
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();
|
||||
let body = json!({"model": "test"});
|
||||
|
||||
// sonnet 模型
|
||||
assert_eq!(
|
||||
get_model_from_provider("claude-sonnet-4-5-20250929", &provider, &body),
|
||||
"anthropic/claude-sonnet-4.5"
|
||||
);
|
||||
|
||||
// haiku 模型
|
||||
assert_eq!(
|
||||
get_model_from_provider("claude-haiku-4-5-20250929", &provider, &body),
|
||||
"anthropic/claude-haiku-4.5"
|
||||
);
|
||||
|
||||
// opus 模型
|
||||
assert_eq!(
|
||||
get_model_from_provider("claude-opus-4-5", &provider, &body),
|
||||
"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");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_thinking_parameter_detection() {
|
||||
let mut provider = create_openrouter_provider();
|
||||
// 添加推理模型配置
|
||||
if let Some(env) = provider.settings_config.get_mut("env") {
|
||||
env["ANTHROPIC_REASONING_MODEL"] = json!("anthropic/claude-sonnet-4.5:extended");
|
||||
}
|
||||
|
||||
let input = json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"max_tokens": 1024,
|
||||
"thinking": {"type": "enabled"},
|
||||
"messages": [{"role": "user", "content": "Solve this problem"}]
|
||||
});
|
||||
|
||||
let result = anthropic_to_openai(input, &provider).unwrap();
|
||||
// 应该使用推理模型
|
||||
assert_eq!(result["model"], "anthropic/claude-sonnet-4.5:extended");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_thinking_parameter_disabled() {
|
||||
let mut provider = create_openrouter_provider();
|
||||
if let Some(env) = provider.settings_config.get_mut("env") {
|
||||
env["ANTHROPIC_REASONING_MODEL"] = json!("anthropic/claude-sonnet-4.5:extended");
|
||||
}
|
||||
|
||||
let input = json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"max_tokens": 1024,
|
||||
"thinking": {"type": "disabled"},
|
||||
"messages": [{"role": "user", "content": "Hello"}]
|
||||
});
|
||||
|
||||
let result = anthropic_to_openai(input, &provider).unwrap();
|
||||
// 应该使用普通模型
|
||||
assert_eq!(result["model"], "anthropic/claude-sonnet-4.5");
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user