Files
CC-Switch/src-tauri/src/proxy/handlers.rs
T
YoVinchen b1103c8a59 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.
2025-12-05 11:26:41 +08:00

1126 lines
39 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! 请求处理器
//!
//! 处理各种API端点的HTTP请求
use super::{
forwarder::RequestForwarder,
providers::{get_adapter, transform, ProviderType},
server::ProxyState,
session::ProxySession,
types::*,
usage::{logger::UsageLogger, parser::TokenUsage},
ProxyError,
};
use crate::app_config::AppType;
use axum::{extract::State, http::StatusCode, response::IntoResponse, Json};
use bytes::Bytes;
use futures::stream::{Stream, StreamExt};
use rust_decimal::Decimal;
use serde_json::{json, Value};
use std::{
str::FromStr,
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
};
use tokio::sync::Mutex;
/// 记录请求使用量(带 ProxySession 支持)
#[allow(dead_code, clippy::too_many_arguments)]
async fn log_usage_with_session(
state: &ProxyState,
session: &ProxySession,
provider_id: &str,
app_type: &str,
usage: TokenUsage,
latency_ms: u64,
first_token_ms: Option<u64>,
status_code: u16,
provider_type: Option<&ProviderType>,
) {
let logger = UsageLogger::new(&state.db);
// 获取 provider 的 cost_multiplier
let multiplier = match state.db.get_provider_by_id(provider_id, app_type) {
Ok(Some(p)) => {
if let Some(meta) = p.meta {
if let Some(cm) = meta.cost_multiplier {
Decimal::from_str(&cm).unwrap_or(Decimal::from(1))
} else {
Decimal::from(1)
}
} else {
Decimal::from(1)
}
}
_ => Decimal::from(1),
};
let model = session
.model
.clone()
.unwrap_or_else(|| "unknown".to_string());
let provider_type_str = provider_type.map(|pt| pt.as_str().to_string());
if let Err(e) = logger.log_with_calculation(
session.session_id.clone(),
provider_id.to_string(),
app_type.to_string(),
model,
usage,
multiplier,
latency_ms,
first_token_ms,
status_code,
Some(session.session_id.clone()),
provider_type_str,
session.is_streaming,
) {
log::warn!("记录使用量失败: {e}");
}
}
/// 记录请求使用量(兼容旧接口)
#[allow(clippy::too_many_arguments)]
async fn log_usage(
state: &ProxyState,
provider_id: &str,
app_type: &str,
model: &str,
usage: TokenUsage,
latency_ms: u64,
first_token_ms: Option<u64>,
is_streaming: bool,
status_code: u16,
) {
let logger = UsageLogger::new(&state.db);
// 获取 provider 的 cost_multiplier
let multiplier = match state.db.get_provider_by_id(provider_id, app_type) {
Ok(Some(p)) => {
if let Some(meta) = p.meta {
if let Some(cm) = meta.cost_multiplier {
Decimal::from_str(&cm).unwrap_or(Decimal::from(1))
} else {
Decimal::from(1)
}
} else {
Decimal::from(1)
}
}
_ => Decimal::from(1),
};
let request_id = uuid::Uuid::new_v4().to_string();
if let Err(e) = logger.log_with_calculation(
request_id,
provider_id.to_string(),
app_type.to_string(),
model.to_string(),
usage,
multiplier,
latency_ms,
first_token_ms,
status_code,
None,
None, // provider_type
is_streaming,
) {
log::warn!("记录使用量失败: {e}");
}
}
type UsageCallbackWithTiming = Arc<dyn Fn(Vec<Value>, Option<u64>) + Send + Sync + 'static>;
#[derive(Clone)]
struct SseUsageCollector {
inner: Arc<SseUsageCollectorInner>,
}
struct SseUsageCollectorInner {
events: Mutex<Vec<Value>>,
first_event_time: Mutex<Option<std::time::Instant>>,
start_time: std::time::Instant,
on_complete: UsageCallbackWithTiming,
finished: AtomicBool,
}
impl SseUsageCollector {
fn new(
start_time: std::time::Instant,
callback: impl Fn(Vec<Value>, Option<u64>) + Send + Sync + 'static,
) -> Self {
let on_complete: UsageCallbackWithTiming = Arc::new(callback);
Self {
inner: Arc::new(SseUsageCollectorInner {
events: Mutex::new(Vec::new()),
first_event_time: Mutex::new(None),
start_time,
on_complete,
finished: AtomicBool::new(false),
}),
}
}
async fn push(&self, event: Value) {
// 记录首个事件时间
{
let mut first_time = self.inner.first_event_time.lock().await;
if first_time.is_none() {
*first_time = Some(std::time::Instant::now());
}
}
let mut events = self.inner.events.lock().await;
events.push(event);
}
async fn finish(&self) {
if self.inner.finished.swap(true, Ordering::SeqCst) {
return;
}
let events = {
let mut guard = self.inner.events.lock().await;
std::mem::take(&mut *guard)
};
let first_token_ms = {
let first_time = self.inner.first_event_time.lock().await;
first_time.map(|t| (t - self.inner.start_time).as_millis() as u64)
};
(self.inner.on_complete)(events, first_token_ms);
}
}
/// 创建带日志记录的透传流
fn create_logged_passthrough_stream(
stream: impl Stream<Item = Result<Bytes, std::io::Error>> + Send + 'static,
tag: &'static str,
usage_collector: Option<SseUsageCollector>,
) -> impl Stream<Item = Result<Bytes, std::io::Error>> + Send {
async_stream::stream! {
let mut buffer = String::new();
let mut collector = usage_collector;
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);
// 尝试解析并记录完整的 SSE 事件
while let Some(pos) = buffer.find("\n\n") {
let event_text = buffer[..pos].to_string();
buffer = buffer[pos + 2..].to_string();
if !event_text.trim().is_empty() {
// 提取 data 部分并尝试解析为 JSON
for line in event_text.lines() {
if let Some(data) = line.strip_prefix("data: ") {
if data.trim() != "[DONE]" {
if let Ok(json_value) = serde_json::from_str::<Value>(data) {
if let Some(c) = &collector {
c.push(json_value.clone()).await;
}
log::info!(
"[{}] <<< SSE 事件:\n{}",
tag,
serde_json::to_string_pretty(&json_value).unwrap_or_else(|_| data.to_string())
);
} else {
log::info!("[{tag}] <<< SSE 数据: {data}");
}
} else {
log::info!("[{tag}] <<< SSE: [DONE]");
}
}
}
}
}
yield Ok(bytes);
}
Err(e) => {
log::error!("[{tag}] 流错误: {e}");
yield Err(std::io::Error::other(e.to_string()));
break;
}
}
}
log::info!("[{}] ====== 流结束 ======", tag);
if let Some(c) = collector.take() {
c.finish().await;
}
}
}
/// 健康检查
pub async fn health_check() -> (StatusCode, Json<Value>) {
(
StatusCode::OK,
Json(json!({
"status": "healthy",
"timestamp": chrono::Utc::now().to_rfc3339(),
})),
)
}
/// 获取服务状态
pub async fn get_status(State(state): State<ProxyState>) -> Result<Json<ProxyStatus>, ProxyError> {
let status = state.status.read().await.clone();
Ok(Json(status))
}
/// 处理 /v1/messages 请求(Claude API
pub async fn handle_messages(
State(state): State<ProxyState>,
headers: axum::http::HeaderMap,
Json(body): Json<Value>,
) -> Result<axum::response::Response, ProxyError> {
let start_time = std::time::Instant::now();
let config = state.config.read().await.clone();
let request_model = body
.get("model")
.and_then(|m| m.as_str())
.unwrap_or("unknown")
.to_string();
// 选择目标 Provider
let router = super::router::ProviderRouter::new(state.db.clone());
let failed_ids = Vec::new();
let provider = router
.select_provider(&AppType::Claude, &failed_ids)
.await?;
// 检查是否需要转换(OpenRouter
let adapter = get_adapter(&AppType::Claude);
let needs_transform = adapter.needs_transform(&provider);
// 检查是否是流式请求
let is_stream = body
.get("stream")
.and_then(|s| s.as_bool())
.unwrap_or(false);
log::info!(
"[Claude] Provider: {}, needs_transform: {}, is_stream: {}",
provider.name,
needs_transform,
is_stream
);
let forwarder = RequestForwarder::new(
state.db.clone(),
config.request_timeout,
config.max_retries,
state.status.clone(),
);
let response = forwarder
.forward_with_retry(&AppType::Claude, "/v1/messages", body, headers)
.await?;
let status = response.status();
log::info!("[Claude] 上游响应状态: {status}");
// 如果需要转换
if needs_transform {
if is_stream {
// 流式响应转换
log::info!("[Claude] 开始流式响应转换 (OpenAI SSE → Anthropic SSE)");
let stream = response.bytes_stream();
let sse_stream = super::providers::streaming::create_anthropic_sse_stream(stream);
let usage_collector = {
let state = state.clone();
let provider_id = provider.id.clone();
let model = request_model.clone();
let status_code = status.as_u16();
let start_time_clone = start_time;
SseUsageCollector::new(start_time, move |events, first_token_ms| {
if let Some(usage) = TokenUsage::from_claude_stream_events(&events) {
let latency_ms = start_time_clone.elapsed().as_millis() as u64;
let state = state.clone();
let provider_id = provider_id.clone();
let model = model.clone();
tokio::spawn(async move {
log_usage(
&state,
&provider_id,
"claude",
&model,
usage,
latency_ms,
first_token_ms,
true, // is_streaming
status_code,
)
.await;
});
} else {
log::debug!("[Claude] OpenRouter 流式响应缺少 usage 统计,跳过消费记录");
}
})
};
let logged_stream = create_logged_passthrough_stream(
sse_stream,
"Claude/OpenRouter",
Some(usage_collector),
);
let mut headers = axum::http::HeaderMap::new();
headers.insert(
"Content-Type",
axum::http::HeaderValue::from_static("text/event-stream"),
);
headers.insert(
"Cache-Control",
axum::http::HeaderValue::from_static("no-cache"),
);
headers.insert(
"Connection",
axum::http::HeaderValue::from_static("keep-alive"),
);
let body = axum::body::Body::from_stream(logged_stream);
log::info!("[Claude] ====== 请求结束 (流式转换) ======");
return Ok((headers, body).into_response());
} else {
// 非流式响应转换
log::info!("[Claude] 开始转换响应 (OpenAI → Anthropic)");
let response_headers = response.headers().clone();
// 读取响应体
let body_bytes = response.bytes().await.map_err(|e| {
log::error!("[Claude] 读取响应体失败: {e}");
ProxyError::ForwardFailed(format!("Failed to read response body: {e}"))
})?;
let body_str = String::from_utf8_lossy(&body_bytes);
log::info!("[Claude] OpenAI 响应长度: {} bytes", body_bytes.len());
log::debug!("[Claude] OpenAI 原始响应: {body_str}");
// 解析并转换
let openai_response: Value = serde_json::from_slice(&body_bytes).map_err(|e| {
log::error!("[Claude] 解析 OpenAI 响应失败: {e}, body: {body_str}");
ProxyError::TransformError(format!("Failed to parse OpenAI response: {e}"))
})?;
log::info!("[Claude] 解析 OpenAI 响应成功");
log::info!(
"[Claude] <<< OpenAI 响应 JSON:\n{}",
serde_json::to_string_pretty(&openai_response).unwrap_or_default()
);
let anthropic_response =
transform::openai_to_anthropic(openai_response).map_err(|e| {
log::error!("[Claude] 转换响应失败: {e}");
e
})?;
log::info!("[Claude] 转换响应成功");
log::info!(
"[Claude] <<< Anthropic 响应 JSON:\n{}",
serde_json::to_string_pretty(&anthropic_response).unwrap_or_default()
);
// 记录使用量
if let Some(usage) = TokenUsage::from_claude_response(&anthropic_response) {
let model = anthropic_response
.get("model")
.and_then(|m| m.as_str())
.unwrap_or("unknown");
let latency_ms = start_time.elapsed().as_millis() as u64;
tokio::spawn({
let state = state.clone();
let provider_id = provider.id.clone();
let model = model.to_string();
async move {
log_usage(
&state,
&provider_id,
"claude",
&model,
usage,
latency_ms,
None,
false,
status.as_u16(),
)
.await;
}
});
}
log::info!("[Claude] ====== 请求结束 ======");
// 构建响应
let mut builder = axum::response::Response::builder().status(status);
// 复制响应头(排除 content-length,因为内容已改变)
for (key, value) in response_headers.iter() {
if key.as_str().to_lowercase() != "content-length"
&& key.as_str().to_lowercase() != "transfer-encoding"
{
builder = builder.header(key, value);
}
}
builder = builder.header("content-type", "application/json");
let response_body = serde_json::to_vec(&anthropic_response).map_err(|e| {
log::error!("[Claude] 序列化响应失败: {e}");
ProxyError::TransformError(format!("Failed to serialize response: {e}"))
})?;
log::info!(
"[Claude] 返回转换后的响应, 长度: {} bytes",
response_body.len()
);
let body = axum::body::Body::from(response_body);
return Ok(builder.body(body).unwrap());
}
}
// 透传响应(直连 Anthropic
log::info!("[Claude] 透传响应模式");
// 检查是否流式响应
let content_type = response
.headers()
.get("content-type")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
let is_sse = content_type.contains("text/event-stream");
if is_sse {
// 流式透传:使用包装流记录 SSE 事件
log::info!("[Claude] 流式透传响应 (SSE)");
let mut builder = axum::response::Response::builder().status(status);
for (key, value) in response.headers() {
builder = builder.header(key, value);
}
let stream = response
.bytes_stream()
.map(|chunk| chunk.map_err(|e| std::io::Error::other(e.to_string())));
let usage_collector = {
let state = state.clone();
let provider_id = provider.id.clone();
let model = request_model.clone();
let status_code = status.as_u16();
let start_time_clone = start_time;
SseUsageCollector::new(start_time, move |events, first_token_ms| {
if let Some(usage) = TokenUsage::from_claude_stream_events(&events) {
let latency_ms = start_time_clone.elapsed().as_millis() as u64;
let state = state.clone();
let provider_id = provider_id.clone();
let model = model.clone();
tokio::spawn(async move {
log_usage(
&state,
&provider_id,
"claude",
&model,
usage,
latency_ms,
first_token_ms,
true,
status_code,
)
.await;
});
} else {
log::debug!("[Claude] 流式响应缺少 usage 统计,跳过消费记录");
}
})
};
let logged_stream =
create_logged_passthrough_stream(stream, "Claude", Some(usage_collector));
let body = axum::body::Body::from_stream(logged_stream);
log::info!("[Claude] ====== 请求结束 (流式) ======");
Ok(builder.body(body).unwrap())
} else {
// 非流式透传:读取完整响应并记录
let response_headers = response.headers().clone();
let status = response.status();
let body_bytes = response.bytes().await.map_err(|e| {
log::error!("[Claude] 读取透传响应失败: {e}");
ProxyError::ForwardFailed(format!("Failed to read response body: {e}"))
})?;
// 记录响应 JSON
if let Ok(json_value) = serde_json::from_slice::<Value>(&body_bytes) {
log::info!(
"[Claude] <<< Anthropic 透传响应 JSON:\n{}",
serde_json::to_string_pretty(&json_value).unwrap_or_default()
);
// 记录使用量
if let Some(usage) = TokenUsage::from_claude_response(&json_value) {
let model = json_value
.get("model")
.and_then(|m| m.as_str())
.unwrap_or("unknown");
let latency_ms = start_time.elapsed().as_millis() as u64;
tokio::spawn({
let state = state.clone();
let provider_id = provider.id.clone();
let model = model.to_string();
async move {
log_usage(
&state,
&provider_id,
"claude",
&model,
usage,
latency_ms,
None,
false,
status.as_u16(),
)
.await;
}
});
}
} else {
log::info!(
"[Claude] <<< 透传响应 (非 JSON): {} bytes",
body_bytes.len()
);
}
log::info!("[Claude] ====== 请求结束 ======");
let mut builder = axum::response::Response::builder().status(status);
for (key, value) in response_headers.iter() {
builder = builder.header(key, value);
}
let body = axum::body::Body::from(body_bytes);
Ok(builder.body(body).unwrap())
}
}
/// 处理 Gemini API 请求(透传,包括查询参数)
pub async fn handle_gemini(
State(state): State<ProxyState>,
uri: axum::http::Uri,
headers: axum::http::HeaderMap,
Json(body): Json<Value>,
) -> Result<axum::response::Response, ProxyError> {
let start_time = std::time::Instant::now();
let config = state.config.read().await.clone();
// 选择目标 Provider
let router = super::router::ProviderRouter::new(state.db.clone());
let failed_ids = Vec::new();
let provider = router
.select_provider(&AppType::Gemini, &failed_ids)
.await?;
let forwarder = RequestForwarder::new(
state.db.clone(),
config.request_timeout,
config.max_retries,
state.status.clone(),
);
// 提取完整的路径和查询参数
let endpoint = uri
.path_and_query()
.map(|pq| pq.as_str())
.unwrap_or(uri.path());
let gemini_model = endpoint
.split('/')
.find(|s| s.starts_with("models/"))
.and_then(|s| s.strip_prefix("models/"))
.map(|s| s.split(':').next().unwrap_or(s))
.unwrap_or("unknown")
.to_string();
log::info!("[Gemini] 请求端点: {endpoint}");
let response = forwarder
.forward_with_retry(&AppType::Gemini, endpoint, body, headers)
.await?;
let status = response.status();
log::info!("[Gemini] 上游响应状态: {status}");
// 检查是否流式响应
let content_type = response
.headers()
.get("content-type")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
let is_sse = content_type.contains("text/event-stream");
if is_sse {
// 流式透传
log::info!("[Gemini] 流式透传响应 (SSE)");
let mut builder = axum::response::Response::builder().status(status);
for (key, value) in response.headers() {
builder = builder.header(key, value);
}
let stream = response
.bytes_stream()
.map(|chunk| chunk.map_err(|e| std::io::Error::other(e.to_string())));
let usage_collector = {
let state = state.clone();
let provider_id = provider.id.clone();
let fallback_model = gemini_model.clone();
let status_code = status.as_u16();
let start_time_clone = start_time;
SseUsageCollector::new(start_time, move |events, first_token_ms| {
if let Some(usage) = TokenUsage::from_gemini_stream_chunks(&events) {
// 优先使用响应中的实际模型名称,否则使用从 URI 提取的模型名称
let model = usage
.model
.clone()
.unwrap_or_else(|| fallback_model.clone());
let latency_ms = start_time_clone.elapsed().as_millis() as u64;
let state = state.clone();
let provider_id = provider_id.clone();
tokio::spawn(async move {
log_usage(
&state,
&provider_id,
"gemini",
&model,
usage,
latency_ms,
first_token_ms,
true,
status_code,
)
.await;
});
} else {
log::debug!("[Gemini] 流式响应缺少 usage 统计,跳过消费记录");
}
})
};
let logged_stream =
create_logged_passthrough_stream(stream, "Gemini", Some(usage_collector));
let body = axum::body::Body::from_stream(logged_stream);
Ok(builder.body(body).unwrap())
} else {
// 非流式透传
let response_headers = response.headers().clone();
let status = response.status();
let body_bytes = response.bytes().await.map_err(|e| {
log::error!("[Gemini] 读取响应失败: {e}");
ProxyError::ForwardFailed(format!("Failed to read response body: {e}"))
})?;
// 记录响应 JSON
if let Ok(json_value) = serde_json::from_slice::<Value>(&body_bytes) {
log::info!(
"[Gemini] <<< 响应 JSON:\n{}",
serde_json::to_string_pretty(&json_value).unwrap_or_default()
);
// 记录使用量
if let Some(usage) = TokenUsage::from_gemini_response(&json_value) {
// 优先使用响应中的实际模型名称,否则使用从 URI 提取的模型名称
let model = usage.model.clone().unwrap_or_else(|| gemini_model.clone());
let latency_ms = start_time.elapsed().as_millis() as u64;
tokio::spawn({
let state = state.clone();
let provider_id = provider.id.clone();
async move {
log_usage(
&state,
&provider_id,
"gemini",
&model,
usage,
latency_ms,
None,
false,
status.as_u16(),
)
.await;
}
});
}
} else {
log::info!("[Gemini] <<< 响应 (非 JSON): {} bytes", body_bytes.len());
}
log::info!("[Gemini] ====== 请求结束 ======");
let mut builder = axum::response::Response::builder().status(status);
for (key, value) in response_headers.iter() {
builder = builder.header(key, value);
}
let body = axum::body::Body::from(body_bytes);
Ok(builder.body(body).unwrap())
}
}
/// 处理 /v1/responses 请求(OpenAI Responses API - Codex CLI 透传)
pub async fn handle_responses(
State(state): State<ProxyState>,
headers: axum::http::HeaderMap,
Json(body): Json<Value>,
) -> Result<axum::response::Response, ProxyError> {
let start_time = std::time::Instant::now();
let config = state.config.read().await.clone();
let request_model = body
.get("model")
.and_then(|m| m.as_str())
.unwrap_or("unknown")
.to_string();
// 选择目标 Provider
let router = super::router::ProviderRouter::new(state.db.clone());
let failed_ids = Vec::new();
let provider = router.select_provider(&AppType::Codex, &failed_ids).await?;
let forwarder = RequestForwarder::new(
state.db.clone(),
config.request_timeout,
config.max_retries,
state.status.clone(),
);
let response = forwarder
.forward_with_retry(&AppType::Codex, "/v1/responses", body, headers)
.await?;
let status = response.status();
log::info!("[Codex] 上游响应状态: {status}");
// 检查是否流式响应
let content_type = response
.headers()
.get("content-type")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
let is_sse = content_type.contains("text/event-stream");
if is_sse {
// 流式透传
log::info!("[Codex] 流式透传响应 (SSE)");
let mut builder = axum::response::Response::builder().status(status);
for (key, value) in response.headers() {
builder = builder.header(key, value);
}
let stream = response
.bytes_stream()
.map(|chunk| chunk.map_err(|e| std::io::Error::other(e.to_string())));
let usage_collector = {
let state = state.clone();
let provider_id = provider.id.clone();
let request_model = request_model.clone();
let status_code = status.as_u16();
let start_time_clone = start_time;
SseUsageCollector::new(start_time, move |events, first_token_ms| {
if let Some(usage) = TokenUsage::from_codex_stream_events(&events) {
// 尝试从事件中提取模型,回退到请求模型
let model = events
.iter()
.find_map(|e| {
if e.get("type")?.as_str()? == "response.completed" {
e.get("response")?.get("model")?.as_str()
} else {
None
}
})
.unwrap_or(&request_model)
.to_string();
let latency_ms = start_time_clone.elapsed().as_millis() as u64;
let state = state.clone();
let provider_id = provider_id.clone();
tokio::spawn(async move {
log_usage(
&state,
&provider_id,
"codex",
&model,
usage,
latency_ms,
first_token_ms,
true,
status_code,
)
.await;
});
} else {
log::debug!("[Codex] 流式响应缺少 usage 统计,跳过消费记录");
}
})
};
let logged_stream =
create_logged_passthrough_stream(stream, "Codex", Some(usage_collector));
let body = axum::body::Body::from_stream(logged_stream);
Ok(builder.body(body).unwrap())
} else {
// 非流式透传
let response_headers = response.headers().clone();
let status = response.status();
let body_bytes = response.bytes().await.map_err(|e| {
log::error!("[Codex] 读取响应失败: {e}");
ProxyError::ForwardFailed(format!("Failed to read response body: {e}"))
})?;
// 记录响应 JSON
if let Ok(json_value) = serde_json::from_slice::<Value>(&body_bytes) {
log::info!(
"[Codex] <<< 响应 JSON:\n{}",
serde_json::to_string_pretty(&json_value).unwrap_or_default()
);
// 记录使用量
if let Some(usage) = TokenUsage::from_codex_response(&json_value) {
let model = json_value
.get("model")
.and_then(|m| m.as_str())
.unwrap_or("unknown");
let latency_ms = start_time.elapsed().as_millis() as u64;
log::info!(
"[Codex] 解析到 usage: input={}, output={}",
usage.input_tokens,
usage.output_tokens
);
tokio::spawn({
let state = state.clone();
let provider_id = provider.id.clone();
let model = model.to_string();
async move {
log_usage(
&state,
&provider_id,
"codex",
&model,
usage,
latency_ms,
None,
false,
status.as_u16(),
)
.await;
}
});
} else {
log::warn!("[Codex] 未能解析 usage 信息,跳过记录");
}
} else {
log::info!("[Codex] <<< 响应 (非 JSON): {} bytes", body_bytes.len());
}
log::info!("[Codex] ====== 请求结束 ======");
let mut builder = axum::response::Response::builder().status(status);
for (key, value) in response_headers.iter() {
builder = builder.header(key, value);
}
let body = axum::body::Body::from(body_bytes);
Ok(builder.body(body).unwrap())
}
}
/// 处理 /v1/chat/completions 请求(OpenAI Chat Completions API - Codex CLI
pub async fn handle_chat_completions(
State(state): State<ProxyState>,
headers: axum::http::HeaderMap,
Json(body): Json<Value>,
) -> Result<axum::response::Response, ProxyError> {
let start_time = std::time::Instant::now();
log::info!("[Codex] ====== /v1/chat/completions 请求开始 ======");
let config = state.config.read().await.clone();
let request_model = body
.get("model")
.and_then(|m| m.as_str())
.unwrap_or("unknown")
.to_string();
let is_stream = body
.get("stream")
.and_then(|v| v.as_bool())
.unwrap_or(false);
log::info!("[Codex] 请求模型: {request_model}, 流式: {is_stream}");
// 选择目标 Provider
let router = super::router::ProviderRouter::new(state.db.clone());
let failed_ids = Vec::new();
let provider = router.select_provider(&AppType::Codex, &failed_ids).await?;
log::info!("[Codex] 选择 Provider: {}", provider.id);
let forwarder = RequestForwarder::new(
state.db.clone(),
config.request_timeout,
config.max_retries,
state.status.clone(),
);
let response = forwarder
.forward_with_retry(&AppType::Codex, "/v1/chat/completions", body, headers)
.await?;
let status = response.status();
log::info!("[Codex] 上游响应状态: {status}");
// 检查是否流式响应
let content_type = response
.headers()
.get("content-type")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
let is_sse = content_type.contains("text/event-stream");
if is_sse {
// 流式透传
log::info!("[Codex] 流式透传响应 (SSE)");
let mut builder = axum::response::Response::builder().status(status);
for (key, value) in response.headers() {
builder = builder.header(key, value);
}
let stream = response
.bytes_stream()
.map(|chunk| chunk.map_err(|e| std::io::Error::other(e.to_string())));
let usage_collector = {
let state = state.clone();
let provider_id = provider.id.clone();
let request_model = request_model.clone();
let status_code = status.as_u16();
let start_time_clone = start_time;
SseUsageCollector::new(start_time, move |events, first_token_ms| {
if let Some(usage) = TokenUsage::from_openai_stream_events(&events) {
let model = events
.iter()
.find_map(|e| e.get("model")?.as_str())
.unwrap_or(&request_model)
.to_string();
let latency_ms = start_time_clone.elapsed().as_millis() as u64;
let state = state.clone();
let provider_id = provider_id.clone();
tokio::spawn(async move {
log_usage(
&state,
&provider_id,
"codex",
&model,
usage,
latency_ms,
first_token_ms,
true,
status_code,
)
.await;
});
} else {
log::debug!("[Codex] 流式响应缺少 usage 统计,跳过消费记录");
}
})
};
let logged_stream =
create_logged_passthrough_stream(stream, "Codex", Some(usage_collector));
let body = axum::body::Body::from_stream(logged_stream);
Ok(builder.body(body).unwrap())
} else {
// 非流式透传
let response_headers = response.headers().clone();
let status = response.status();
let body_bytes = response.bytes().await.map_err(|e| {
log::error!("[Codex] 读取响应失败: {e}");
ProxyError::ForwardFailed(format!("Failed to read response body: {e}"))
})?;
// 记录响应 JSON
if let Ok(json_value) = serde_json::from_slice::<Value>(&body_bytes) {
log::info!(
"[Codex] <<< 响应 JSON:\n{}",
serde_json::to_string_pretty(&json_value).unwrap_or_default()
);
// 记录使用量 (OpenAI 格式: prompt_tokens, completion_tokens)
if let Some(usage) = TokenUsage::from_openai_response(&json_value) {
let model = json_value
.get("model")
.and_then(|m| m.as_str())
.unwrap_or("unknown");
let latency_ms = start_time.elapsed().as_millis() as u64;
log::info!(
"[Codex] 解析到 usage: input={}, output={}",
usage.input_tokens,
usage.output_tokens
);
tokio::spawn({
let state = state.clone();
let provider_id = provider.id.clone();
let model = model.to_string();
async move {
log_usage(
&state,
&provider_id,
"codex",
&model,
usage,
latency_ms,
None,
false,
status.as_u16(),
)
.await;
}
});
} else {
log::warn!("[Codex] 未能解析 usage 信息,跳过记录");
}
} else {
log::info!("[Codex] <<< 响应 (非 JSON): {} bytes", body_bytes.len());
}
log::info!("[Codex] ====== 请求结束 ======");
let mut builder = axum::response::Response::builder().status(status);
for (key, value) in response_headers.iter() {
builder = builder.header(key, value);
}
let body = axum::body::Body::from(body_bytes);
Ok(builder.body(body).unwrap())
}
}