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
This commit is contained in:
YoVinchen
2025-12-01 02:41:36 +08:00
parent 04a588694b
commit 4c94e70f97
41 changed files with 3491 additions and 41 deletions
+2
View File
@@ -9,6 +9,7 @@ mod misc;
mod plugin;
mod prompt;
mod provider;
mod proxy;
mod settings;
pub mod skill;
@@ -21,5 +22,6 @@ pub use misc::*;
pub use plugin::*;
pub use prompt::*;
pub use provider::*;
pub use proxy::*;
pub use settings::*;
pub use skill::*;
+13
View File
@@ -86,6 +86,19 @@ pub fn switch_provider(
.map_err(|e| e.to_string())
}
/// 设置代理目标供应商
#[tauri::command]
pub fn set_proxy_target_provider(
state: State<'_, AppState>,
app: String,
id: String,
) -> Result<bool, String> {
let app_type = AppType::from_str(&app).map_err(|e| e.to_string())?;
ProviderService::set_proxy_target(state.inner(), app_type, &id)
.map(|_| true)
.map_err(|e| e.to_string())
}
fn import_default_config_internal(state: &AppState, app_type: AppType) -> Result<bool, AppError> {
ProviderService::import_default_config(state, app_type)
}
+47
View File
@@ -0,0 +1,47 @@
//! 代理服务相关的 Tauri 命令
//!
//! 提供前端调用的 API 接口
use crate::proxy::types::*;
use crate::store::AppState;
/// 启动代理服务器
#[tauri::command]
pub async fn start_proxy_server(
state: tauri::State<'_, AppState>,
) -> Result<ProxyServerInfo, String> {
state.proxy_service.start().await
}
/// 停止代理服务器
#[tauri::command]
pub async fn stop_proxy_server(state: tauri::State<'_, AppState>) -> Result<(), String> {
state.proxy_service.stop().await
}
/// 获取代理服务器状态
#[tauri::command]
pub async fn get_proxy_status(state: tauri::State<'_, AppState>) -> Result<ProxyStatus, String> {
state.proxy_service.get_status().await
}
/// 获取代理配置
#[tauri::command]
pub async fn get_proxy_config(state: tauri::State<'_, AppState>) -> Result<ProxyConfig, String> {
state.proxy_service.get_config().await
}
/// 更新代理配置
#[tauri::command]
pub async fn update_proxy_config(
state: tauri::State<'_, AppState>,
config: ProxyConfig,
) -> Result<(), String> {
state.proxy_service.update_config(&config).await
}
/// 检查代理服务器是否正在运行
#[tauri::command]
pub async fn is_proxy_running(state: tauri::State<'_, AppState>) -> Result<bool, String> {
Ok(state.proxy_service.is_running().await)
}
+8 -7
View File
@@ -1,11 +1,12 @@
//! 数据访问对象 (DAO) 模块
//! Data Access Object layer
//!
//! 提供各类数据的 CRUD 操作。
//! Database access operations for each domain
mod mcp;
mod prompts;
mod providers;
mod settings;
mod skills;
pub mod mcp;
pub mod prompts;
pub mod providers;
pub mod proxy;
pub mod settings;
pub mod skills;
// 所有 DAO 方法都通过 Database impl 提供,无需单独导出
+76 -10
View File
@@ -17,7 +17,7 @@ impl Database {
) -> Result<IndexMap<String, Provider>, AppError> {
let conn = lock_conn!(self.conn);
let mut stmt = conn.prepare(
"SELECT id, name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta
"SELECT id, name, settings_config, website_url, category, created_at, sort_index, notes, icon, icon_color, meta, is_proxy_target
FROM providers WHERE app_type = ?1
ORDER BY COALESCE(sort_index, 999999), created_at ASC, id ASC"
).map_err(|e| AppError::Database(e.to_string()))?;
@@ -35,6 +35,7 @@ impl Database {
let icon: Option<String> = row.get(8)?;
let icon_color: Option<String> = row.get(9)?;
let meta_str: String = row.get(10)?;
let is_proxy_target: bool = row.get(11)?;
let settings_config =
serde_json::from_str(&settings_config_str).unwrap_or(serde_json::Value::Null);
@@ -54,6 +55,7 @@ impl Database {
meta: Some(meta),
icon,
icon_color,
is_proxy_target: Some(is_proxy_target),
},
))
})
@@ -121,6 +123,26 @@ impl Database {
}
}
/// 获取代理目标供应商 ID
pub fn get_proxy_target_provider(&self, app_type: &str) -> Result<Option<String>, AppError> {
let conn = lock_conn!(self.conn);
let mut stmt = conn
.prepare("SELECT id FROM providers WHERE app_type = ?1 AND is_proxy_target = 1 LIMIT 1")
.map_err(|e| AppError::Database(e.to_string()))?;
let mut rows = stmt
.query(params![app_type])
.map_err(|e| AppError::Database(e.to_string()))?;
if let Some(row) = rows.next().map_err(|e| AppError::Database(e.to_string()))? {
Ok(Some(
row.get(0).map_err(|e| AppError::Database(e.to_string()))?,
))
} else {
Ok(None)
}
}
/// 保存供应商(新增或更新)
///
/// 注意:更新模式下不同步 endpoints,因为编辑模式下端点通过单独的 API 管理
@@ -135,17 +157,17 @@ impl Database {
let mut meta_clone = provider.meta.clone().unwrap_or_default();
let endpoints = std::mem::take(&mut meta_clone.custom_endpoints);
// 检查是否存在(用于判断新增/更新,以及保留 is_current
let existing: Option<bool> = tx
// 检查是否存在(用于判断新增/更新,以及保留 is_current 和 is_proxy_target
let existing: Option<(bool, bool)> = tx
.query_row(
"SELECT is_current FROM providers WHERE id = ?1 AND app_type = ?2",
"SELECT is_current, is_proxy_target FROM providers WHERE id = ?1 AND app_type = ?2",
params![provider.id, app_type],
|row| row.get(0),
|row| Ok((row.get(0)?, row.get(1)?)),
)
.ok();
let is_update = existing.is_some();
let is_current = existing.unwrap_or(false);
let (is_current, is_proxy_target) = existing.unwrap_or((false, false));
if is_update {
// 更新模式:使用 UPDATE 避免触发 ON DELETE CASCADE
@@ -161,8 +183,9 @@ impl Database {
icon = ?8,
icon_color = ?9,
meta = ?10,
is_current = ?11
WHERE id = ?12 AND app_type = ?13",
is_current = ?11,
is_proxy_target = ?12
WHERE id = ?13 AND app_type = ?14",
params![
provider.name,
serde_json::to_string(&provider.settings_config).unwrap(),
@@ -175,6 +198,7 @@ impl Database {
provider.icon_color,
serde_json::to_string(&meta_clone).unwrap(),
is_current,
is_proxy_target,
provider.id,
app_type,
],
@@ -185,8 +209,8 @@ impl Database {
tx.execute(
"INSERT INTO providers (
id, app_type, name, settings_config, website_url, category,
created_at, sort_index, notes, icon, icon_color, meta, is_current
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13)",
created_at, sort_index, notes, icon, icon_color, meta, is_current, is_proxy_target
) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)",
params![
provider.id,
app_type,
@@ -201,6 +225,7 @@ impl Database {
provider.icon_color,
serde_json::to_string(&meta_clone).unwrap(),
is_current,
is_proxy_target,
],
)
.map_err(|e| AppError::Database(e.to_string()))?;
@@ -256,6 +281,47 @@ impl Database {
Ok(())
}
/// 设置代理目标供应商
pub fn set_proxy_target_provider(&self, app_type: &str, id: &str) -> Result<(), AppError> {
let mut conn = lock_conn!(self.conn);
let tx = conn
.transaction()
.map_err(|e| AppError::Database(e.to_string()))?;
// 重置所有为 0
tx.execute(
"UPDATE providers SET is_proxy_target = 0 WHERE app_type = ?1",
params![app_type],
)
.map_err(|e| AppError::Database(e.to_string()))?;
// 设置新的代理目标供应商
tx.execute(
"UPDATE providers SET is_proxy_target = 1 WHERE id = ?1 AND app_type = ?2",
params![id, app_type],
)
.map_err(|e| AppError::Database(e.to_string()))?;
tx.commit().map_err(|e| AppError::Database(e.to_string()))?;
Ok(())
}
/// 获取所有活跃的代理目标
pub fn get_all_proxy_targets(&self) -> Result<Vec<(String, String, String)>, AppError> {
let conn = lock_conn!(self.conn);
let mut stmt = conn
.prepare("SELECT app_type, name, id FROM providers WHERE is_proxy_target = 1")
.map_err(|e| AppError::Database(e.to_string()))?;
let targets = stmt
.query_map([], |row| Ok((row.get(0)?, row.get(1)?, row.get(2)?)))
.map_err(|e| AppError::Database(e.to_string()))?
.collect::<Result<Vec<_>, _>>()
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(targets)
}
/// 添加自定义端点
pub fn add_custom_endpoint(
&self,
+242
View File
@@ -0,0 +1,242 @@
//! 代理功能数据访问层
//!
//! 处理代理配置、Provider健康状态和使用统计的数据库操作
use crate::error::AppError;
use crate::proxy::types::*;
use super::super::{lock_conn, Database};
impl Database {
// ==================== Proxy Config ====================
/// 获取代理配置
pub async fn get_proxy_config(&self) -> Result<ProxyConfig, AppError> {
// 在一个作用域内获取锁并查询,确保锁在await之前释放
let result = {
let conn = lock_conn!(self.conn);
conn.query_row(
"SELECT enabled, listen_address, listen_port, max_retries,
request_timeout, enable_logging
FROM proxy_config WHERE id = 1",
[],
|row| {
Ok(ProxyConfig {
enabled: row.get::<_, i32>(0)? != 0,
listen_address: row.get(1)?,
listen_port: row.get::<_, i32>(2)? as u16,
max_retries: row.get::<_, i32>(3)? as u8,
request_timeout: row.get::<_, i32>(4)? as u64,
enable_logging: row.get::<_, i32>(5)? != 0,
})
},
)
}; // conn锁在这里释放
match result {
Ok(config) => Ok(config),
Err(rusqlite::Error::QueryReturnedNoRows) => {
// 如果不存在,插入默认配置
let default_config = ProxyConfig::default();
self.update_proxy_config(default_config.clone()).await?;
Ok(default_config)
}
Err(e) => Err(AppError::Database(e.to_string())),
}
}
/// 更新代理配置
pub async fn update_proxy_config(&self, config: ProxyConfig) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
conn.execute(
"INSERT OR REPLACE INTO proxy_config
(id, enabled, listen_address, listen_port, max_retries, request_timeout, enable_logging, target_app, updated_at)
VALUES (1, ?1, ?2, ?3, ?4, ?5, ?6, ?7, datetime('now'))",
rusqlite::params![
if config.enabled { 1 } else { 0 },
config.listen_address,
config.listen_port as i32,
config.max_retries as i32,
config.request_timeout as i32,
if config.enable_logging { 1 } else { 0 },
"claude", // 兼容旧字段,写入默认值
],
)
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(())
}
// ==================== Provider Health ====================
/// 获取Provider健康状态
pub async fn get_provider_health(
&self,
provider_id: &str,
app_type: &str,
) -> Result<ProviderHealth, AppError> {
let conn = lock_conn!(self.conn);
conn.query_row(
"SELECT provider_id, app_type, is_healthy, consecutive_failures,
last_success_at, last_failure_at, last_error, updated_at
FROM provider_health
WHERE provider_id = ?1 AND app_type = ?2",
rusqlite::params![provider_id, app_type],
|row| {
Ok(ProviderHealth {
provider_id: row.get(0)?,
app_type: row.get(1)?,
is_healthy: row.get::<_, i64>(2)? != 0,
consecutive_failures: row.get::<_, i64>(3)? as u32,
last_success_at: row.get(4)?,
last_failure_at: row.get(5)?,
last_error: row.get(6)?,
updated_at: row.get(7)?,
})
},
)
.map_err(|e| AppError::Database(e.to_string()))
}
/// 更新Provider健康状态
pub async fn update_provider_health(
&self,
provider_id: &str,
app_type: &str,
success: bool,
error_msg: Option<String>,
) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
let now = chrono::Utc::now().to_rfc3339();
// 先查询当前状态
let current = conn.query_row(
"SELECT consecutive_failures FROM provider_health
WHERE provider_id = ?1 AND app_type = ?2",
rusqlite::params![provider_id, app_type],
|row| Ok(row.get::<_, i64>(0)? as u32),
);
let (is_healthy, consecutive_failures) = if success {
// 成功:重置失败计数
(1, 0)
} else {
// 失败:增加失败计数
let failures = current.unwrap_or(0) + 1;
let healthy = if failures >= 3 { 0 } else { 1 };
(healthy, failures)
};
let (last_success_at, last_failure_at) = if success {
(Some(now.clone()), None)
} else {
(None, Some(now.clone()))
};
// UPSERT
conn.execute(
"INSERT OR REPLACE INTO provider_health
(provider_id, app_type, is_healthy, consecutive_failures,
last_success_at, last_failure_at, last_error, updated_at)
VALUES (?1, ?2, ?3, ?4,
COALESCE(?5, (SELECT last_success_at FROM provider_health
WHERE provider_id = ?1 AND app_type = ?2)),
COALESCE(?6, (SELECT last_failure_at FROM provider_health
WHERE provider_id = ?1 AND app_type = ?2)),
?7, ?8)",
rusqlite::params![
provider_id,
app_type,
is_healthy,
consecutive_failures as i64,
last_success_at,
last_failure_at,
error_msg,
&now,
],
)
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(())
}
// ==================== Proxy Usage (可选) ====================
/// 记录代理使用统计
#[allow(dead_code)]
pub async fn record_proxy_usage(&self, record: &ProxyUsageRecord) -> Result<(), AppError> {
let conn = lock_conn!(self.conn);
conn.execute(
"INSERT INTO proxy_usage
(provider_id, app_type, endpoint, request_tokens, response_tokens,
status_code, latency_ms, error, timestamp)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)",
rusqlite::params![
&record.provider_id,
&record.app_type,
&record.endpoint,
record.request_tokens,
record.response_tokens,
record.status_code as i64,
record.latency_ms as i64,
&record.error,
&record.timestamp,
],
)
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(())
}
/// 查询最近的使用统计
#[allow(dead_code)]
pub async fn get_recent_usage(
&self,
provider_id: &str,
app_type: &str,
limit: usize,
) -> Result<Vec<ProxyUsageRecord>, AppError> {
let conn = lock_conn!(self.conn);
let mut stmt = conn
.prepare(
"SELECT provider_id, app_type, endpoint, request_tokens, response_tokens,
status_code, latency_ms, error, timestamp
FROM proxy_usage
WHERE provider_id = ?1 AND app_type = ?2
ORDER BY timestamp DESC
LIMIT ?3",
)
.map_err(|e| AppError::Database(e.to_string()))?;
let rows = stmt
.query_map(
rusqlite::params![provider_id, app_type, limit as i64],
|row| {
Ok(ProxyUsageRecord {
provider_id: row.get(0)?,
app_type: row.get(1)?,
endpoint: row.get(2)?,
request_tokens: row.get(3)?,
response_tokens: row.get(4)?,
status_code: row.get::<_, i64>(5)? as u16,
latency_ms: row.get::<_, i64>(6)? as u64,
error: row.get(7)?,
timestamp: row.get(8)?,
})
},
)
.map_err(|e| AppError::Database(e.to_string()))?;
let mut records = Vec::new();
for row in rows {
records.push(row.map_err(|e| AppError::Database(e.to_string()))?);
}
Ok(records)
}
}
+84
View File
@@ -31,12 +31,19 @@ impl Database {
icon_color TEXT,
meta TEXT NOT NULL DEFAULT '{}',
is_current BOOLEAN NOT NULL DEFAULT 0,
is_proxy_target BOOLEAN NOT NULL DEFAULT 0,
PRIMARY KEY (id, app_type)
)",
[],
)
.map_err(|e| AppError::Database(e.to_string()))?;
// 尝试添加 is_proxy_target 列(如果表已存在但缺少该列)
let _ = conn.execute(
"ALTER TABLE providers ADD COLUMN is_proxy_target BOOLEAN NOT NULL DEFAULT 0",
[],
);
// 2. Provider Endpoints 表
conn.execute(
"CREATE TABLE IF NOT EXISTS provider_endpoints (
@@ -120,6 +127,83 @@ impl Database {
)
.map_err(|e| AppError::Database(e.to_string()))?;
// 8. Proxy Config 表 (代理服务器配置)
// 代理配置表(单例)
conn.execute(
"CREATE TABLE IF NOT EXISTS proxy_config (
id INTEGER PRIMARY KEY CHECK (id = 1),
enabled INTEGER NOT NULL DEFAULT 0,
listen_address TEXT NOT NULL DEFAULT '127.0.0.1',
listen_port INTEGER NOT NULL DEFAULT 5000,
max_retries INTEGER NOT NULL DEFAULT 3,
request_timeout INTEGER NOT NULL DEFAULT 300,
enable_logging INTEGER NOT NULL DEFAULT 1,
target_app TEXT NOT NULL DEFAULT 'claude',
created_at TEXT NOT NULL DEFAULT (datetime('now')),
updated_at TEXT NOT NULL DEFAULT (datetime('now'))
)",
[],
)
.map_err(|e| AppError::Database(e.to_string()))?;
// 尝试添加 target_app 列(如果表已存在但缺少该列)
// 忽略 "duplicate column name" 错误
let _ = conn.execute(
"ALTER TABLE proxy_config ADD COLUMN target_app TEXT NOT NULL DEFAULT 'claude'",
[],
);
// 9. Provider Health 表 (Provider健康状态)
conn.execute(
"CREATE TABLE IF NOT EXISTS provider_health (
provider_id TEXT NOT NULL,
app_type TEXT NOT NULL,
is_healthy INTEGER NOT NULL DEFAULT 1,
consecutive_failures INTEGER NOT NULL DEFAULT 0,
last_success_at TEXT,
last_failure_at TEXT,
last_error TEXT,
updated_at TEXT NOT NULL,
PRIMARY KEY (provider_id, app_type),
FOREIGN KEY (provider_id, app_type) REFERENCES providers(id, app_type) ON DELETE CASCADE
)",
[],
)
.map_err(|e| AppError::Database(e.to_string()))?;
// 10. Proxy Usage 表 (代理使用统计,可选)
conn.execute(
"CREATE TABLE IF NOT EXISTS proxy_usage (
id INTEGER PRIMARY KEY AUTOINCREMENT,
provider_id TEXT NOT NULL,
app_type TEXT NOT NULL,
endpoint TEXT NOT NULL,
request_tokens INTEGER,
response_tokens INTEGER,
status_code INTEGER NOT NULL,
latency_ms INTEGER NOT NULL,
error TEXT,
timestamp TEXT NOT NULL
)",
[],
)
.map_err(|e| AppError::Database(e.to_string()))?;
// 为 proxy_usage 创建索引
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_proxy_usage_timestamp
ON proxy_usage(timestamp)",
[],
)
.map_err(|e| AppError::Database(e.to_string()))?;
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_proxy_usage_provider
ON proxy_usage(provider_id, app_type)",
[],
)
.map_err(|e| AppError::Database(e.to_string()))?;
Ok(())
}
+1
View File
@@ -245,6 +245,7 @@ fn dry_run_validates_schema_compatibility() {
meta: None,
icon: None,
icon_color: None,
is_proxy_target: false,
},
);
+1
View File
@@ -129,6 +129,7 @@ pub(crate) fn build_provider_from_request(
meta: None,
icon: request.icon.clone(),
icon_color: None,
is_proxy_target: None,
};
Ok(provider)
+31
View File
@@ -17,6 +17,7 @@ mod prompt;
mod prompt_files;
mod provider;
mod provider_defaults;
mod proxy;
mod services;
mod settings;
mod store;
@@ -521,6 +522,28 @@ pub fn run() {
}
}
// 自动启动代理服务器
let app_handle = app.handle().clone();
tauri::async_runtime::spawn(async move {
let state = app_handle.state::<AppState>();
match state.db.get_proxy_config().await {
Ok(config) => {
if config.enabled {
log::info!("代理服务配置为启用,正在启动...");
match state.proxy_service.start().await {
Ok(info) => log::info!(
"代理服务器自动启动成功: {}:{}",
info.address,
info.port
),
Err(e) => log::error!("代理服务器自动启动失败: {e}"),
}
}
}
Err(e) => log::error!("启动时获取代理配置失败: {e}"),
}
});
Ok(())
})
.invoke_handler(tauri::generate_handler![
@@ -530,6 +553,7 @@ pub fn run() {
commands::update_provider,
commands::delete_provider,
commands::switch_provider,
commands::set_proxy_target_provider,
commands::import_default_config,
commands::get_claude_config_status,
commands::get_config_status,
@@ -619,6 +643,13 @@ pub fn run() {
// Auto launch
commands::set_auto_launch,
commands::get_auto_launch_status,
// Proxy server management
commands::start_proxy_server,
commands::stop_proxy_server,
commands::get_proxy_status,
commands::get_proxy_config,
commands::update_proxy_config,
commands::is_proxy_running,
]);
let app = builder
+5
View File
@@ -36,6 +36,10 @@ pub struct Provider {
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "iconColor")]
pub icon_color: Option<String>,
/// 是否为代理目标(数据库专用字段,不写入配置文件)
#[serde(skip_serializing_if = "Option::is_none")]
#[serde(rename = "isProxyTarget")]
pub is_proxy_target: Option<bool>,
}
impl Provider {
@@ -58,6 +62,7 @@ impl Provider {
meta: None,
icon: None,
icon_color: None,
is_proxy_target: None,
}
}
}
+119
View File
@@ -0,0 +1,119 @@
use axum::{
http::StatusCode,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
use thiserror::Error;
#[derive(Debug, Error)]
pub enum ProxyError {
#[error("服务器已在运行")]
AlreadyRunning,
#[error("服务器未运行")]
NotRunning,
#[error("地址绑定失败: {0}")]
BindFailed(String),
#[error("请求转发失败: {0}")]
ForwardFailed(String),
#[error("无可用的Provider")]
NoAvailableProvider,
#[error("Provider不健康: {0}")]
ProviderUnhealthy(String),
#[error("上游错误 (状态码 {status}): {body:?}")]
UpstreamError { status: u16, body: Option<String> },
#[error("超过最大重试次数")]
MaxRetriesExceeded,
#[error("数据库错误: {0}")]
DatabaseError(String),
#[error("配置错误: {0}")]
ConfigError(String),
#[allow(dead_code)]
#[error("格式转换错误: {0}")]
TransformError(String),
#[allow(dead_code)]
#[error("无效的请求: {0}")]
InvalidRequest(String),
#[error("超时: {0}")]
Timeout(String),
#[allow(dead_code)]
#[error("内部错误: {0}")]
Internal(String),
}
impl IntoResponse for ProxyError {
fn into_response(self) -> Response {
let (status, message) = match &self {
ProxyError::AlreadyRunning => (StatusCode::CONFLICT, self.to_string()),
ProxyError::NotRunning => (StatusCode::SERVICE_UNAVAILABLE, self.to_string()),
ProxyError::BindFailed(_) => (StatusCode::INTERNAL_SERVER_ERROR, self.to_string()),
ProxyError::ForwardFailed(_) => (StatusCode::BAD_GATEWAY, self.to_string()),
ProxyError::NoAvailableProvider => (StatusCode::SERVICE_UNAVAILABLE, self.to_string()),
ProxyError::ProviderUnhealthy(_) => (StatusCode::SERVICE_UNAVAILABLE, self.to_string()),
ProxyError::UpstreamError { status, .. } => (
StatusCode::from_u16(*status).unwrap_or(StatusCode::BAD_GATEWAY),
self.to_string(),
),
ProxyError::MaxRetriesExceeded => (StatusCode::SERVICE_UNAVAILABLE, self.to_string()),
ProxyError::DatabaseError(_) => (StatusCode::INTERNAL_SERVER_ERROR, self.to_string()),
ProxyError::ConfigError(_) => (StatusCode::BAD_REQUEST, self.to_string()),
ProxyError::TransformError(_) => (StatusCode::UNPROCESSABLE_ENTITY, self.to_string()),
ProxyError::InvalidRequest(_) => (StatusCode::BAD_REQUEST, self.to_string()),
ProxyError::Timeout(_) => (StatusCode::GATEWAY_TIMEOUT, self.to_string()),
ProxyError::Internal(_) => (StatusCode::INTERNAL_SERVER_ERROR, self.to_string()),
};
let body = Json(json!({
"error": {
"message": message,
"type": "proxy_error",
}
}));
(status, body).into_response()
}
}
/// 错误分类
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ErrorCategory {
/// 可重试错误(网络问题、5xx)
Retryable, // 网络超时、5xx 错误
/// 不可重试错误(4xx、认证失败)
NonRetryable, // 认证失败、参数错误、4xx 错误
#[allow(dead_code)]
ClientAbort, // 客户端主动中断
}
/// 判断错误是否可重试
#[allow(dead_code)]
pub fn categorize_error(error: &reqwest::Error) -> ErrorCategory {
if error.is_timeout() || error.is_connect() {
return ErrorCategory::Retryable;
}
if let Some(status) = error.status() {
if status.is_server_error() {
ErrorCategory::Retryable
} else if status.is_client_error() {
ErrorCategory::NonRetryable
} else {
ErrorCategory::Retryable
}
} else {
ErrorCategory::Retryable
}
}
+652
View File
@@ -0,0 +1,652 @@
//! 请求转发器
//!
//! 负责将请求转发到上游Provider,支持重试和故障转移
use super::{error::*, router::ProviderRouter, types::ProxyStatus, ProxyError};
use crate::{app_config::AppType, database::Database, provider::Provider};
use reqwest::{Client, Response};
use serde_json::Value;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::RwLock;
pub struct RequestForwarder {
client: Client,
router: ProviderRouter,
max_retries: u8,
status: Arc<RwLock<ProxyStatus>>,
}
impl RequestForwarder {
pub fn new(
db: Arc<Database>,
timeout_secs: u64,
max_retries: u8,
status: Arc<RwLock<ProxyStatus>>,
) -> Self {
let mut client_builder = Client::builder();
if timeout_secs > 0 {
client_builder = client_builder.timeout(Duration::from_secs(timeout_secs));
}
let client = client_builder
.build()
.expect("Failed to create HTTP client");
Self {
client,
router: ProviderRouter::new(db),
max_retries,
status,
}
}
/// 转发请求(带重试和故障转移)
pub async fn forward_with_retry(
&self,
app_type: &AppType,
endpoint: &str,
body: Value,
headers: axum::http::HeaderMap,
) -> Result<Response, ProxyError> {
let mut failed_ids = Vec::new();
let mut failover_happened = false;
for attempt in 0..self.max_retries {
// 选择Provider
let provider = self.router.select_provider(app_type, &failed_ids).await?;
log::debug!(
"尝试 {} - 使用Provider: {} ({})",
attempt + 1,
provider.name,
provider.id
);
// 更新状态中的当前Provider信息
{
let mut status = self.status.write().await;
status.current_provider = Some(provider.name.clone());
status.current_provider_id = Some(provider.id.clone());
status.total_requests += 1;
status.last_request_at = Some(chrono::Utc::now().to_rfc3339());
if attempt > 0 {
failover_happened = true;
}
}
let start = Instant::now();
// 转发请求
match self.forward(&provider, endpoint, &body, &headers).await {
Ok(response) => {
let _latency = start.elapsed().as_millis() as u64;
// 成功:更新健康状态
self.router
.update_health(&provider, app_type, true, None)
.await;
// 更新成功统计
{
let mut status = self.status.write().await;
status.success_requests += 1;
status.last_error = None;
if failover_happened {
status.failover_count += 1;
}
// 重新计算成功率
if status.total_requests > 0 {
status.success_rate = (status.success_requests as f32
/ status.total_requests as f32)
* 100.0;
}
}
return Ok(response);
}
Err(e) => {
let latency = start.elapsed().as_millis() as u64;
// 失败:分类错误
let category = self.categorize_proxy_error(&e);
match category {
ErrorCategory::Retryable => {
// 可重试:更新健康状态,添加到失败列表
self.router
.update_health(&provider, app_type, false, Some(e.to_string()))
.await;
failed_ids.push(provider.id.clone());
// 更新错误信息
{
let mut status = self.status.write().await;
status.last_error =
Some(format!("Provider {} 失败: {}", provider.name, e));
}
log::warn!(
"请求失败(可重试): Provider {} - {} - {}ms",
provider.name,
e,
latency
);
continue;
}
ErrorCategory::NonRetryable | ErrorCategory::ClientAbort => {
// 不可重试:更新失败统计并返回
{
let mut status = self.status.write().await;
status.failed_requests += 1;
status.last_error = Some(e.to_string());
if status.total_requests > 0 {
status.success_rate = (status.success_requests as f32
/ status.total_requests as f32)
* 100.0;
}
}
log::error!("请求失败(不可重试): {e}");
return Err(e);
}
}
}
}
}
// 所有重试都失败
{
let mut status = self.status.write().await;
status.failed_requests += 1;
status.last_error = Some("已达到最大重试次数".to_string());
if status.total_requests > 0 {
status.success_rate =
(status.success_requests as f32 / status.total_requests as f32) * 100.0;
}
}
Err(ProxyError::MaxRetriesExceeded)
}
/// 转发单个请求
async fn forward(
&self,
provider: &Provider,
endpoint: &str,
body: &Value,
headers: &axum::http::HeaderMap,
) -> Result<Response, ProxyError> {
// 提取 base_url
let base_url = self.extract_base_url(provider)?;
// 智能拼接 URL,避免重复的 /v1
let url = if base_url.ends_with("/v1") && endpoint.starts_with("/v1") {
format!("{}{}", base_url.trim_end_matches("/v1"), endpoint)
} else {
format!("{base_url}{endpoint}")
};
// 构建请求
let mut request = self.client.post(&url);
// 透传 Headers
for (key, value) in headers {
let key_str = key.as_str().to_lowercase();
// 过滤掉一些不应该直接转发的 Header
if key_str == "host"
|| key_str == "content-length"
|| key_str == "accept-encoding"
// 过滤认证相关 Header
|| key_str == "x-api-key"
|| key_str == "authorization"
|| key_str == "x-goog-api-key"
|| key_str == "anthropic-version"
{
continue;
}
request = request.header(key, value);
}
// 确保 Content-Type 是 json
request = request.header("Content-Type", "application/json");
// 添加认证头
request = self.add_auth_headers(request, provider)?;
// 发送请求
let response = request.json(body).send().await.map_err(|e| {
log::error!("Request Failed: {e}");
if e.is_timeout() {
ProxyError::Timeout(format!("请求超时: {e}"))
} else if e.is_connect() {
ProxyError::ForwardFailed(format!("连接失败: {e}"))
} else {
ProxyError::ForwardFailed(e.to_string())
}
})?;
// 检查响应状态
let status = response.status();
if status.is_success() {
Ok(response)
} else {
let status_code = status.as_u16();
let body_text = response.text().await.ok();
Err(ProxyError::UpstreamError {
status: status_code,
body: body_text,
})
}
}
/// 添加认证头
fn add_auth_headers(
&self,
mut request: reqwest::RequestBuilder,
provider: &Provider,
) -> Result<reqwest::RequestBuilder, ProxyError> {
// 提取 apiKey 和认证类型
if let Some((api_key, auth_type)) = self.extract_api_key(provider) {
// 遮蔽 key 用于日志
let _masked_key = if api_key.len() > 8 {
format!("{}...{}", &api_key[..4], &api_key[api_key.len() - 4..])
} else {
"***".to_string()
};
match auth_type {
AuthType::Anthropic => {
request = request.header("x-api-key", api_key);
request = request.header("anthropic-version", "2023-06-01");
}
AuthType::Gemini => {
request = request.header("x-goog-api-key", api_key);
}
AuthType::Bearer => {
request = request.header("Authorization", format!("Bearer {api_key}"));
}
}
} else {
log::error!("✗ 未找到 API Key!将发送未认证的请求(会失败)");
log::error!("Provider 配置: {:?}", provider.settings_config);
}
Ok(request)
}
/// 从 Provider 配置中提取 base_url
fn extract_base_url(&self, provider: &Provider) -> Result<String, ProxyError> {
log::debug!("Extracting base_url for provider: {}", provider.name);
// 1. 尝试直接获取 base_url 字段 (Codex CLI 常用格式)
if let Some(url) = provider
.settings_config
.get("base_url")
.and_then(|v| v.as_str())
{
log::debug!("Found base_url in direct field: {url}");
return Ok(url.trim_end_matches('/').to_string());
}
// 2. 尝试从 env 中获取 (Claude / Gemini)
if let Some(env) = provider.settings_config.get("env") {
if let Some(url) = env.get("ANTHROPIC_BASE_URL").and_then(|v| v.as_str()) {
log::debug!("Found base_url in env.ANTHROPIC_BASE_URL: {url}");
return Ok(url.trim_end_matches('/').to_string());
}
if let Some(url) = env.get("GOOGLE_GEMINI_BASE_URL").and_then(|v| v.as_str()) {
log::debug!("Found base_url in env.GOOGLE_GEMINI_BASE_URL: {url}");
return Ok(url.trim_end_matches('/').to_string());
}
}
// 3. 尝试其他通用字段
if let Some(url) = provider
.settings_config
.get("baseURL")
.and_then(|v| v.as_str())
{
log::debug!("Found base_url in baseURL: {url}");
return Ok(url.trim_end_matches('/').to_string());
}
if let Some(url) = provider
.settings_config
.get("apiEndpoint")
.and_then(|v| v.as_str())
{
log::debug!("Found base_url in apiEndpoint: {url}");
return Ok(url.trim_end_matches('/').to_string());
}
// 4. 尝试从 config 对象中获取 (Codex - JSON 格式)
if let Some(config) = provider.settings_config.get("config") {
// 如果 config 是一个对象
if let Some(url) = config.get("base_url").and_then(|v| v.as_str()) {
log::debug!("Found base_url in config.base_url: {url}");
return Ok(url.trim_end_matches('/').to_string());
}
// 如果 config 是一个字符串,尝试解析
if let Some(config_str) = config.as_str() {
// 尝试双引号
if let Some(start) = config_str.find("base_url = \"") {
let rest = &config_str[start + 12..];
if let Some(end) = rest.find('"') {
let url = rest[..end].trim_end_matches('/').to_string();
log::debug!("Found base_url in config string (double quotes): {url}");
return Ok(url);
}
}
// 尝试单引号
if let Some(start) = config_str.find("base_url = '") {
let rest = &config_str[start + 12..];
if let Some(end) = rest.find('\'') {
let url = rest[..end].trim_end_matches('/').to_string();
log::debug!("Found base_url in config string (single quotes): {url}");
return Ok(url);
}
}
}
}
log::error!(
"Failed to extract base_url from config: {:?}",
provider.settings_config
);
Err(ProxyError::ConfigError(
"Provider缺少base_url配置".to_string(),
))
}
/// 从 Provider 配置中提取 api_key
fn extract_api_key(&self, provider: &Provider) -> Option<(String, AuthType)> {
// 1. 尝试从 env 中获取
if let Some(env) = provider.settings_config.get("env") {
// Claude/Anthropic
if let Some(key) = env.get("ANTHROPIC_AUTH_TOKEN").and_then(|v| v.as_str()) {
return Some((key.to_string(), AuthType::Anthropic));
}
// Gemini
if let Some(key) = env.get("GEMINI_API_KEY").and_then(|v| v.as_str()) {
return Some((key.to_string(), AuthType::Gemini));
}
// OpenAI/Codex (env 中的 OPENAI_API_KEY)
if let Some(key) = env.get("OPENAI_API_KEY").and_then(|v| v.as_str()) {
return Some((key.to_string(), AuthType::Bearer));
}
}
// 2. 尝试从 auth 中获取 (Codex CLI 格式)
if let Some(auth) = provider.settings_config.get("auth") {
if let Some(key) = auth.get("OPENAI_API_KEY").and_then(|v| v.as_str()) {
return Some((key.to_string(), AuthType::Bearer));
}
}
// 3. 尝试直接获取 (支持 apiKey 和 api_key)
if let Some(key) = provider
.settings_config
.get("apiKey")
.or_else(|| provider.settings_config.get("api_key"))
.and_then(|v| v.as_str())
{
return Some((key.to_string(), AuthType::Bearer));
}
// 4. 尝试从 config 对象中获取
if let Some(config) = provider.settings_config.get("config") {
if let Some(key) = config
.get("api_key")
.or_else(|| config.get("apiKey"))
.and_then(|v| v.as_str())
{
return Some((key.to_string(), AuthType::Bearer));
}
}
log::error!("✗ 所有位置都未找到 API Key");
log::error!("完整配置结构: {:?}", provider.settings_config);
None
}
/// 分类ProxyError
fn categorize_proxy_error(&self, error: &ProxyError) -> ErrorCategory {
match error {
ProxyError::Timeout(_) => ErrorCategory::Retryable,
ProxyError::ForwardFailed(_) => ErrorCategory::Retryable,
ProxyError::UpstreamError { status, .. } => {
if *status >= 500 {
ErrorCategory::Retryable
} else if *status >= 400 && *status < 500 {
ErrorCategory::NonRetryable
} else {
ErrorCategory::Retryable
}
}
ProxyError::ProviderUnhealthy(_) => ErrorCategory::Retryable,
ProxyError::NoAvailableProvider => ErrorCategory::NonRetryable,
_ => ErrorCategory::NonRetryable,
}
}
}
enum AuthType {
Anthropic,
Gemini,
Bearer,
}
impl RequestForwarder {
/// 转发 GET 请求(带重试和故障转移)
pub async fn forward_get_request(
&self,
app_type: &AppType,
endpoint: &str,
headers: axum::http::HeaderMap,
) -> Result<Response, ProxyError> {
let mut failed_ids = Vec::new();
for attempt in 0..self.max_retries {
let provider = self.router.select_provider(app_type, &failed_ids).await?;
log::debug!(
"GET 尝试 {} - 使用Provider: {} ({})",
attempt + 1,
provider.name,
provider.id
);
match self.forward_get(&provider, endpoint, &headers).await {
Ok(response) => {
self.router
.update_health(&provider, app_type, true, None)
.await;
return Ok(response);
}
Err(e) => {
let category = self.categorize_proxy_error(&e);
match category {
ErrorCategory::Retryable => {
self.router
.update_health(&provider, app_type, false, Some(e.to_string()))
.await;
failed_ids.push(provider.id.clone());
continue;
}
_ => return Err(e),
}
}
}
}
Err(ProxyError::MaxRetriesExceeded)
}
/// 转发 DELETE 请求(带重试和故障转移)
pub async fn forward_delete_request(
&self,
app_type: &AppType,
endpoint: &str,
headers: axum::http::HeaderMap,
) -> Result<Response, ProxyError> {
let mut failed_ids = Vec::new();
for attempt in 0..self.max_retries {
let provider = self.router.select_provider(app_type, &failed_ids).await?;
log::debug!(
"DELETE 尝试 {} - 使用Provider: {} ({})",
attempt + 1,
provider.name,
provider.id
);
match self.forward_delete(&provider, endpoint, &headers).await {
Ok(response) => {
self.router
.update_health(&provider, app_type, true, None)
.await;
return Ok(response);
}
Err(e) => {
let category = self.categorize_proxy_error(&e);
match category {
ErrorCategory::Retryable => {
self.router
.update_health(&provider, app_type, false, Some(e.to_string()))
.await;
failed_ids.push(provider.id.clone());
continue;
}
_ => return Err(e),
}
}
}
}
Err(ProxyError::MaxRetriesExceeded)
}
/// 转发单个 GET 请求
async fn forward_get(
&self,
provider: &Provider,
endpoint: &str,
headers: &axum::http::HeaderMap,
) -> Result<Response, ProxyError> {
let base_url = self.extract_base_url(provider)?;
let url = if base_url.ends_with("/v1") && endpoint.starts_with("/v1") {
format!("{}{}", base_url.trim_end_matches("/v1"), endpoint)
} else {
format!("{base_url}{endpoint}")
};
log::info!("Proxy GET Request URL: {url}");
let mut request = self.client.get(&url);
// 透传 Headers
for (key, value) in headers {
let key_str = key.as_str().to_lowercase();
if key_str == "host"
|| key_str == "content-length"
|| key_str == "accept-encoding"
|| key_str == "x-api-key"
|| key_str == "authorization"
|| key_str == "x-goog-api-key"
|| key_str == "anthropic-version"
{
continue;
}
request = request.header(key, value);
}
request = self.add_auth_headers(request, provider)?;
let response = request.send().await.map_err(|e| {
if e.is_timeout() {
ProxyError::Timeout(format!("请求超时: {e}"))
} else if e.is_connect() {
ProxyError::ForwardFailed(format!("连接失败: {e}"))
} else {
ProxyError::ForwardFailed(e.to_string())
}
})?;
let status = response.status();
if status.is_success() {
Ok(response)
} else {
let status_code = status.as_u16();
let body_text = response.text().await.ok();
Err(ProxyError::UpstreamError {
status: status_code,
body: body_text,
})
}
}
/// 转发单个 DELETE 请求
async fn forward_delete(
&self,
provider: &Provider,
endpoint: &str,
headers: &axum::http::HeaderMap,
) -> Result<Response, ProxyError> {
let base_url = self.extract_base_url(provider)?;
let url = if base_url.ends_with("/v1") && endpoint.starts_with("/v1") {
format!("{}{}", base_url.trim_end_matches("/v1"), endpoint)
} else {
format!("{base_url}{endpoint}")
};
log::info!("Proxy DELETE Request URL: {url}");
let mut request = self.client.delete(&url);
// 透传 Headers
for (key, value) in headers {
let key_str = key.as_str().to_lowercase();
if key_str == "host"
|| key_str == "content-length"
|| key_str == "accept-encoding"
|| key_str == "x-api-key"
|| key_str == "authorization"
|| key_str == "x-goog-api-key"
|| key_str == "anthropic-version"
{
continue;
}
request = request.header(key, value);
}
request = self.add_auth_headers(request, provider)?;
let response = request.send().await.map_err(|e| {
if e.is_timeout() {
ProxyError::Timeout(format!("请求超时: {e}"))
} else if e.is_connect() {
ProxyError::ForwardFailed(format!("连接失败: {e}"))
} else {
ProxyError::ForwardFailed(e.to_string())
}
})?;
let status = response.status();
if status.is_success() {
Ok(response)
} else {
let status_code = status.as_u16();
let body_text = response.text().await.ok();
Err(ProxyError::UpstreamError {
status: status_code,
body: body_text,
})
}
}
}
+244
View File
@@ -0,0 +1,244 @@
//! 请求处理器
//!
//! 处理各种API端点的HTTP请求
use super::{forwarder::RequestForwarder, server::ProxyState, types::*, ProxyError};
use crate::app_config::AppType;
use axum::{extract::State, http::StatusCode, Json};
use serde_json::{json, Value};
/// 健康检查
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 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::Claude, &failed_ids)
.await?;
// 直接透传 Claude 请求
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 mut builder = axum::response::Response::builder().status(response.status());
// 复制响应头
for (key, value) in response.headers() {
builder = builder.header(key, value);
}
let body = axum::body::Body::from_stream(response.bytes_stream());
Ok(builder.body(body).unwrap())
}
/// 处理 /v1/messages/count_tokens 请求(透传)
pub async fn handle_count_tokens(
State(state): State<ProxyState>,
headers: axum::http::HeaderMap,
Json(body): Json<Value>,
) -> Result<axum::response::Response, ProxyError> {
let config = state.config.read().await.clone();
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/count_tokens", body, headers)
.await?;
// 透传响应
let mut builder = axum::response::Response::builder().status(response.status());
for (key, value) in response.headers() {
builder = builder.header(key, value);
}
let body = axum::body::Body::from_stream(response.bytes_stream());
Ok(builder.body(body).unwrap())
}
/// 处理 Gemini API 请求(透传)
pub async fn handle_gemini(
State(state): State<ProxyState>,
axum::extract::Path(path): axum::extract::Path<String>,
headers: axum::http::HeaderMap,
Json(body): Json<Value>,
) -> Result<axum::response::Response, ProxyError> {
let config = state.config.read().await.clone();
let forwarder = RequestForwarder::new(
state.db.clone(),
config.request_timeout,
config.max_retries,
state.status.clone(),
);
let endpoint = format!("/{path}");
let response = forwarder
.forward_with_retry(&AppType::Gemini, &endpoint, body, headers)
.await?;
// 透传响应
let mut builder = axum::response::Response::builder().status(response.status());
for (key, value) in response.headers() {
builder = builder.header(key, value);
}
let body = axum::body::Body::from_stream(response.bytes_stream());
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 config = state.config.read().await.clone();
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 mut builder = axum::response::Response::builder().status(response.status());
for (key, value) in response.headers() {
builder = builder.header(key, value);
}
let body = axum::body::Body::from_stream(response.bytes_stream());
Ok(builder.body(body).unwrap())
}
/// 获取单个 ResponseGET /v1/responses/:response_id 透传)
pub async fn handle_get_response(
State(state): State<ProxyState>,
axum::extract::Path(response_id): axum::extract::Path<String>,
headers: axum::http::HeaderMap,
) -> Result<axum::response::Response, ProxyError> {
let config = state.config.read().await.clone();
let forwarder = RequestForwarder::new(
state.db.clone(),
config.request_timeout,
config.max_retries,
state.status.clone(),
);
let endpoint = format!("/v1/responses/{response_id}");
let response = forwarder
.forward_get_request(&AppType::Codex, &endpoint, headers)
.await?;
// 透传响应
let mut builder = axum::response::Response::builder().status(response.status());
for (key, value) in response.headers() {
builder = builder.header(key, value);
}
let body = axum::body::Body::from_stream(response.bytes_stream());
Ok(builder.body(body).unwrap())
}
/// 删除 ResponseDELETE /v1/responses/:response_id 透传)
pub async fn handle_delete_response(
State(state): State<ProxyState>,
axum::extract::Path(response_id): axum::extract::Path<String>,
headers: axum::http::HeaderMap,
) -> Result<axum::response::Response, ProxyError> {
let config = state.config.read().await.clone();
let forwarder = RequestForwarder::new(
state.db.clone(),
config.request_timeout,
config.max_retries,
state.status.clone(),
);
let endpoint = format!("/v1/responses/{response_id}");
let response = forwarder
.forward_delete_request(&AppType::Codex, &endpoint, headers)
.await?;
// 透传响应
let mut builder = axum::response::Response::builder().status(response.status());
for (key, value) in response.headers() {
builder = builder.header(key, value);
}
let body = axum::body::Body::from_stream(response.bytes_stream());
Ok(builder.body(body).unwrap())
}
/// 获取 Response 的输入项(GET /v1/responses/:response_id/input_items 透传)
pub async fn handle_get_response_input_items(
State(state): State<ProxyState>,
axum::extract::Path(response_id): axum::extract::Path<String>,
headers: axum::http::HeaderMap,
) -> Result<axum::response::Response, ProxyError> {
let config = state.config.read().await.clone();
let forwarder = RequestForwarder::new(
state.db.clone(),
config.request_timeout,
config.max_retries,
state.status.clone(),
);
let endpoint = format!("/v1/responses/{response_id}/input_items");
let response = forwarder
.forward_get_request(&AppType::Codex, &endpoint, headers)
.await?;
// 透传响应
let mut builder = axum::response::Response::builder().status(response.status());
for (key, value) in response.headers() {
builder = builder.header(key, value);
}
let body = axum::body::Body::from_stream(response.bytes_stream());
Ok(builder.body(body).unwrap())
}
+7
View File
@@ -0,0 +1,7 @@
//! 健康检查器
//!
//! 负责定期检查Provider健康状态(占位实现)
// 占位实现,稍后添加完整逻辑
#[allow(dead_code)]
pub struct HealthChecker;
+22
View File
@@ -0,0 +1,22 @@
//! 代理服务器模块
//!
//! 提供本地HTTP代理服务,支持多Provider故障转移和请求透传
pub mod error;
mod forwarder;
mod handlers;
mod health;
mod router;
pub(crate) mod server;
pub(crate) mod types;
// 公开导出给外部使用(commands, services等模块需要)
#[allow(unused_imports)]
pub use error::ProxyError;
#[allow(unused_imports)]
pub use types::{ProxyConfig, ProxyServerInfo, ProxyStatus};
// 内部模块间共享(供子模块使用)
// 注意:这个导出用于模块内部,编译器可能警告未使用但实际被子模块使用
#[allow(unused_imports)]
pub(crate) use types::*;
+151
View File
@@ -0,0 +1,151 @@
//! Provider路由器
//!
//! 负责选择合适的Provider进行请求转发,支持健康检查和故障转移
use super::ProxyError;
use crate::{app_config::AppType, database::Database, provider::Provider};
use std::sync::Arc;
pub struct ProviderRouter {
db: Arc<Database>,
}
impl ProviderRouter {
pub fn new(db: Arc<Database>) -> Self {
Self { db }
}
/// 选择Provider(带故障转移)
///
/// 优先使用当前Provider,失败则尝试备用Provider
pub async fn select_provider(
&self,
app_type: &AppType,
failed_ids: &[String],
) -> Result<Provider, ProxyError> {
// 1. 尝试获取当前Provider
match self.get_current_provider(app_type, failed_ids).await {
Ok(provider) => return Ok(provider),
Err(e) => {
log::debug!("当前Provider不可用: {e:?}");
}
}
// 2. 尝试备用Provider
self.select_fallback(app_type, failed_ids).await
}
/// 获取当前Provider
async fn get_current_provider(
&self,
app_type: &AppType,
failed_ids: &[String],
) -> Result<Provider, ProxyError> {
// 1. 尝试获取 Proxy Target Provider ID
let proxy_target_id = self
.db
.get_proxy_target_provider(app_type.as_str())
.map_err(|e| ProxyError::DatabaseError(e.to_string()))?;
// 2. 获取 Current Provider ID (作为 fallback)
let current_id = self
.db
.get_current_provider(app_type.as_str())
.map_err(|e| ProxyError::DatabaseError(e.to_string()))?;
// 3. 确定使用的 ID (优先 proxy_target)
let target_id = proxy_target_id
.or(current_id)
.ok_or(ProxyError::NoAvailableProvider)?;
// 4. 获取所有Provider
let providers = self
.db
.get_all_providers(app_type.as_str())
.map_err(|e| ProxyError::DatabaseError(e.to_string()))?;
// 5. 找到目标Provider
let target = providers
.get(&target_id)
.ok_or(ProxyError::NoAvailableProvider)?;
// 4. 检查是否在失败列表中
if failed_ids.contains(&target.id) {
return Err(ProxyError::ProviderUnhealthy("Provider已失败".to_string()));
}
// 5. 检查健康状态
if self.is_provider_healthy(target, app_type).await {
Ok(target.clone())
} else {
Err(ProxyError::ProviderUnhealthy(target.id.clone()))
}
}
/// 选择备用Provider
async fn select_fallback(
&self,
app_type: &AppType,
failed_ids: &[String],
) -> Result<Provider, ProxyError> {
let providers = self
.db
.get_all_providers(app_type.as_str())
.map_err(|e| ProxyError::DatabaseError(e.to_string()))?;
// 过滤失败的Provider,按sort_index排序
let mut available: Vec<_> = providers
.into_values()
.filter(|p| !failed_ids.contains(&p.id))
.collect();
available.sort_by_key(|p| p.sort_index.unwrap_or(9999));
// 寻找健康的Provider
for provider in available {
if self.is_provider_healthy(&provider, app_type).await {
log::info!("选择备用Provider: {}", provider.name);
return Ok(provider);
}
}
log::warn!("无可用Provider");
Err(ProxyError::NoAvailableProvider)
}
/// 检查Provider是否健康
async fn is_provider_healthy(&self, provider: &Provider, app_type: &AppType) -> bool {
// 从数据库查询健康状态
match self
.db
.get_provider_health(&provider.id, app_type.as_str())
.await
{
Ok(health) => {
// 连续失败3次以上视为不健康
health.is_healthy && health.consecutive_failures < 3
}
Err(_) => {
// 未记录状态时默认健康
true
}
}
}
/// 更新Provider健康状态
pub async fn update_health(
&self,
provider: &Provider,
app_type: &AppType,
success: bool,
error_msg: Option<String>,
) {
if let Err(e) = self
.db
.update_provider_health(&provider.id, app_type.as_str(), success, error_msg)
.await
{
log::warn!("更新Provider健康状态失败: {e:?}");
}
}
}
+176
View File
@@ -0,0 +1,176 @@
//! HTTP代理服务器
//!
//! 基于Axum的HTTP服务器,处理代理请求
use super::{handlers, types::*, ProxyError};
use crate::database::Database;
use axum::{
routing::{get, post},
Router,
};
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::sync::{oneshot, RwLock};
use tower_http::cors::{Any, CorsLayer};
/// 代理服务器状态(共享)
#[derive(Clone)]
pub struct ProxyState {
pub db: Arc<Database>,
pub config: Arc<RwLock<ProxyConfig>>,
pub status: Arc<RwLock<ProxyStatus>>,
pub start_time: Arc<RwLock<Option<std::time::Instant>>>,
}
/// 代理HTTP服务器
pub struct ProxyServer {
config: ProxyConfig,
state: ProxyState,
shutdown_tx: Arc<RwLock<Option<oneshot::Sender<()>>>>,
}
impl ProxyServer {
pub fn new(config: ProxyConfig, db: Arc<Database>) -> Self {
let state = ProxyState {
db,
config: Arc::new(RwLock::new(config.clone())),
status: Arc::new(RwLock::new(ProxyStatus::default())),
start_time: Arc::new(RwLock::new(None)),
};
Self {
config,
state,
shutdown_tx: Arc::new(RwLock::new(None)),
}
}
pub async fn start(&self) -> Result<ProxyServerInfo, ProxyError> {
// 检查是否已在运行
if self.shutdown_tx.read().await.is_some() {
return Err(ProxyError::AlreadyRunning);
}
let addr: SocketAddr =
format!("{}:{}", self.config.listen_address, self.config.listen_port)
.parse()
.map_err(|e| ProxyError::BindFailed(format!("无效的地址: {e}")))?;
// 创建关闭通道
let (shutdown_tx, shutdown_rx) = oneshot::channel();
// 构建路由
let app = self.build_router();
// 绑定监听器
let listener = tokio::net::TcpListener::bind(&addr)
.await
.map_err(|e| ProxyError::BindFailed(e.to_string()))?;
log::info!("代理服务器启动于 {addr}");
// 保存关闭句柄
*self.shutdown_tx.write().await = Some(shutdown_tx);
// 更新状态
let mut status = self.state.status.write().await;
status.running = true;
status.address = self.config.listen_address.clone();
status.port = self.config.listen_port;
drop(status);
// 记录启动时间
*self.state.start_time.write().await = Some(std::time::Instant::now());
// 启动服务器
let state = self.state.clone();
tokio::spawn(async move {
axum::serve(listener, app)
.with_graceful_shutdown(async {
shutdown_rx.await.ok();
})
.await
.ok();
// 服务器停止后更新状态
state.status.write().await.running = false;
*state.start_time.write().await = None;
});
Ok(ProxyServerInfo {
address: self.config.listen_address.clone(),
port: self.config.listen_port,
started_at: chrono::Utc::now().to_rfc3339(),
})
}
pub async fn stop(&self) -> Result<(), ProxyError> {
if let Some(tx) = self.shutdown_tx.write().await.take() {
let _ = tx.send(());
Ok(())
} else {
Err(ProxyError::NotRunning)
}
}
pub async fn get_status(&self) -> ProxyStatus {
let mut status = self.state.status.read().await.clone();
// 计算运行时间
if let Some(start) = *self.state.start_time.read().await {
status.uptime_seconds = start.elapsed().as_secs();
}
// 获取所有活跃的代理目标
if let Ok(targets) = self.state.db.get_all_proxy_targets() {
status.active_targets = targets
.into_iter()
.map(|(app_type, name, id)| ActiveTarget {
app_type,
provider_name: name,
provider_id: id,
})
.collect();
}
status
}
fn build_router(&self) -> Router {
let cors = CorsLayer::new()
.allow_origin(Any)
.allow_methods(Any)
.allow_headers(Any);
Router::new()
// 健康检查
.route("/health", get(handlers::health_check))
.route("/status", get(handlers::get_status))
// Claude API
.route("/v1/messages", post(handlers::handle_messages))
.route(
"/v1/messages/count_tokens",
post(handlers::handle_count_tokens),
)
// OpenAI Responses API (Codex CLI)
.route("/v1/responses", post(handlers::handle_responses))
.route(
"/v1/responses/:response_id",
get(handlers::handle_get_response).delete(handlers::handle_delete_response),
)
.route(
"/v1/responses/:response_id/input_items",
get(handlers::handle_get_response_input_items),
)
// Gemini API (通配符路由)
.route("/v1/*path", post(handlers::handle_gemini))
.route("/v1beta/*path", post(handlers::handle_gemini))
.layer(cors)
.with_state(self.state.clone())
}
/// 在不重启服务的情况下更新运行时配置
pub async fn apply_runtime_config(&self, config: &ProxyConfig) {
*self.state.config.write().await = config.clone();
}
}
+119
View File
@@ -0,0 +1,119 @@
use serde::{Deserialize, Serialize};
/// 代理服务器配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProxyConfig {
/// 是否启用代理服务
pub enabled: bool,
/// 监听地址
pub listen_address: String,
/// 监听端口
pub listen_port: u16,
/// 最大重试次数
pub max_retries: u8,
/// 请求超时时间(秒)
pub request_timeout: u64,
/// 是否启用日志
pub enable_logging: bool,
}
impl Default for ProxyConfig {
fn default() -> Self {
Self {
enabled: false,
listen_address: "127.0.0.1".to_string(),
listen_port: 5000,
max_retries: 3,
request_timeout: 300,
enable_logging: true,
}
}
}
/// 代理服务器状态
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ProxyStatus {
/// 是否运行中
pub running: bool,
/// 监听地址
pub address: String,
/// 监听端口
pub port: u16,
/// 活跃连接数
pub active_connections: usize,
/// 总请求数
pub total_requests: u64,
/// 成功请求数
pub success_requests: u64,
/// 失败请求数
pub failed_requests: u64,
/// 成功率 (0-100)
pub success_rate: f32,
/// 运行时间(秒)
pub uptime_seconds: u64,
/// 当前使用的Provider名称
pub current_provider: Option<String>,
/// 当前Provider的ID
pub current_provider_id: Option<String>,
/// 最后一次请求时间
pub last_request_at: Option<String>,
/// 最后一次错误信息
pub last_error: Option<String>,
/// Provider故障转移次数
pub failover_count: u64,
/// 当前活跃的代理目标列表
#[serde(default)]
pub active_targets: Vec<ActiveTarget>,
}
/// 活跃的代理目标信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ActiveTarget {
pub app_type: String, // "Claude" | "Codex" | "Gemini"
pub provider_name: String,
pub provider_id: String,
}
/// 代理服务器信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProxyServerInfo {
pub address: String,
pub port: u16,
pub started_at: String,
}
/// API 格式类型(预留,当前不需要格式转换)
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[allow(dead_code)]
pub enum ApiFormat {
Claude,
OpenAI,
Gemini,
}
/// Provider健康状态
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProviderHealth {
pub provider_id: String,
pub app_type: String,
pub is_healthy: bool,
pub consecutive_failures: u32,
pub last_success_at: Option<String>,
pub last_failure_at: Option<String>,
pub last_error: Option<String>,
pub updated_at: String,
}
/// 使用统计记录
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProxyUsageRecord {
pub provider_id: String,
pub app_type: String,
pub endpoint: String,
pub request_tokens: Option<i32>,
pub response_tokens: Option<i32>,
pub status_code: u16,
pub latency_ms: u64,
pub error: Option<String>,
pub timestamp: String,
}
+2
View File
@@ -4,6 +4,7 @@ pub mod env_manager;
pub mod mcp;
pub mod prompt;
pub mod provider;
pub mod proxy;
pub mod skill;
pub mod speedtest;
@@ -11,5 +12,6 @@ pub use config::ConfigService;
pub use mcp::McpService;
pub use prompt::PromptService;
pub use provider::{ProviderService, ProviderSortUpdate};
pub use proxy::ProxyService;
pub use skill::{Skill, SkillRepo, SkillService};
pub use speedtest::{EndpointLatency, SpeedtestService};
+12
View File
@@ -217,6 +217,18 @@ impl ProviderService {
Ok(())
}
/// Set proxy target provider
pub fn set_proxy_target(state: &AppState, app_type: AppType, id: &str) -> Result<(), AppError> {
// Check if provider exists
let providers = state.db.get_all_providers(app_type.as_str())?;
if !providers.contains_key(id) {
return Err(AppError::Message(format!("供应商 {id} 不存在")));
}
state.db.set_proxy_target_provider(app_type.as_str(), id)?;
Ok(())
}
/// Sync current provider to live configuration (re-export)
pub fn sync_current_to_live(state: &AppState) -> Result<(), AppError> {
sync_current_to_live(state)
+157
View File
@@ -0,0 +1,157 @@
//! 代理服务业务逻辑层
//!
//! 提供代理服务器的启动、停止和配置管理
use crate::database::Database;
use crate::proxy::server::ProxyServer;
use crate::proxy::types::*;
use std::sync::Arc;
use tokio::sync::RwLock;
#[derive(Clone)]
pub struct ProxyService {
db: Arc<Database>,
server: Arc<RwLock<Option<ProxyServer>>>,
}
impl ProxyService {
pub fn new(db: Arc<Database>) -> Self {
Self {
db,
server: Arc::new(RwLock::new(None)),
}
}
/// 启动代理服务器
pub async fn start(&self) -> Result<ProxyServerInfo, String> {
// 1. 获取配置
let config = self
.db
.get_proxy_config()
.await
.map_err(|e| format!("获取代理配置失败: {e}"))?;
// 2. 检查是否启用
if !config.enabled {
return Err("代理服务未启用,请先在设置中启用".to_string());
}
// 3. 检查是否已在运行
if self.server.read().await.is_some() {
return Err("代理服务已在运行中".to_string());
}
// 4. 创建并启动服务器
let server = ProxyServer::new(config, self.db.clone());
let info = server
.start()
.await
.map_err(|e| format!("启动代理服务器失败: {e}"))?;
// 5. 保存服务器实例
*self.server.write().await = Some(server);
log::info!("代理服务器已启动: {}:{}", info.address, info.port);
Ok(info)
}
/// 停止代理服务器
pub async fn stop(&self) -> Result<(), String> {
if let Some(server) = self.server.write().await.take() {
server
.stop()
.await
.map_err(|e| format!("停止代理服务器失败: {e}"))?;
log::info!("代理服务器已停止");
Ok(())
} else {
Err("代理服务器未运行".to_string())
}
}
/// 获取服务器状态
pub async fn get_status(&self) -> Result<ProxyStatus, String> {
if let Some(server) = self.server.read().await.as_ref() {
Ok(server.get_status().await)
} else {
// 服务器未运行时返回默认状态
Ok(ProxyStatus {
running: false,
..Default::default()
})
}
}
/// 获取代理配置
pub async fn get_config(&self) -> Result<ProxyConfig, String> {
self.db
.get_proxy_config()
.await
.map_err(|e| format!("获取代理配置失败: {e}"))
}
/// 更新代理配置
pub async fn update_config(&self, config: &ProxyConfig) -> Result<(), String> {
// 记录旧配置用于判定是否需要重启
let previous = self
.db
.get_proxy_config()
.await
.map_err(|e| format!("获取代理配置失败: {e}"))?;
// 保存到数据库
self.db
.update_proxy_config(config.clone())
.await
.map_err(|e| format!("保存代理配置失败: {e}"))?;
// 检查服务器当前状态
let mut server_guard = self.server.write().await;
if server_guard.is_none() {
return Ok(());
}
// 如果关闭代理,直接停止
if !config.enabled {
if let Some(server) = server_guard.take() {
server
.stop()
.await
.map_err(|e| format!("停止代理服务器失败: {e}"))?;
log::info!("代理配置禁用了服务,已自动停止代理服务器");
}
return Ok(());
}
let require_restart = config.listen_address != previous.listen_address
|| config.listen_port != previous.listen_port;
if require_restart {
if let Some(server) = server_guard.take() {
server
.stop()
.await
.map_err(|e| format!("重启前停止代理服务器失败: {e}"))?;
}
let new_server = ProxyServer::new(config.clone(), self.db.clone());
new_server
.start()
.await
.map_err(|e| format!("重启代理服务器失败: {e}"))?;
*server_guard = Some(new_server);
log::info!("代理配置已更新,服务器已自动重启应用最新配置");
} else if let Some(server) = server_guard.as_ref() {
server.apply_runtime_config(config).await;
log::info!("代理配置已实时应用,无需重启代理服务器");
}
Ok(())
}
/// 检查服务器是否正在运行
pub async fn is_running(&self) -> bool {
self.server.read().await.is_some()
}
}
+5 -1
View File
@@ -1,14 +1,18 @@
use crate::database::Database;
use crate::services::ProxyService;
use std::sync::Arc;
/// 全局应用状态
pub struct AppState {
pub db: Arc<Database>,
pub proxy_service: ProxyService,
}
impl AppState {
/// 创建新的应用状态
pub fn new(db: Arc<Database>) -> Self {
Self { db }
let proxy_service = ProxyService::new(db.clone());
Self { db, proxy_service }
}
}