refactor: v0.15.1 - server.rs 模块化重构 (-68%)

主要变更:
- server.rs 从 5639 行减少到 1799 行 (-68.1%)
- 创建 server/handlers/ 模块结构
  - api.rs: chat_completions, anthropic_messages (983 行)
  - provider_calls.rs: Provider 调用处理 (925 行)
  - websocket.rs: WebSocket 连接处理 (845 行)
  - management.rs: 管理 API (562 行)
- 创建 server_utils.rs: 工具函数 (656 行)
- 创建 providers/traits.rs: CredentialProvider trait (115 行)
- 统一 ProviderType 枚举,消除重复定义

Co-authored-by: factory-droid[bot] <138933559+factory-droid[bot]@users.noreply.github.com>
This commit is contained in:
coso
2025-12-21 03:28:18 +08:00
co-authored by factory-droid[bot]
parent 2af7093b07
commit 6923f5062a
20 changed files with 6187 additions and 5704 deletions
+1 -1
View File
@@ -3367,7 +3367,7 @@ dependencies = [
[[package]]
name = "proxycast"
version = "0.15.0"
version = "0.15.1"
dependencies = [
"anyhow",
"async-stream",
+1 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "proxycast"
version = "0.15.0"
version = "0.15.1"
description = "AI API Proxy Desktop App"
authors = ["you"]
edition = "2021"
@@ -0,0 +1,7 @@
# Seeds for failure cases proptest has generated in the past. It is
# automatically read and these particular cases re-run before any
# novel cases are generated.
#
# It is recommended to check this file in to source control so that
# everyone who runs the test benefits from these saved cases.
cc 10c52f014c7e9b4cd049a9802452d417b5774e6f42173c2c9329544d8ac4340c # shrinks to lead_time_mins = 21, time_offset_secs = 1260
@@ -0,0 +1,7 @@
# Seeds for failure cases proptest has generated in the past. It is
# automatically read and these particular cases re-run before any
# novel cases are generated.
#
# It is recommended to check this file in to source control so that
# everyone who runs the test benefits from these saved cases.
cc 99af5a6dad66f0b5a2650223417a2e83a7a3bfa8b6eab8ad57a88023367e739c # shrinks to url = "http://08:1024"
+21
View File
@@ -14,6 +14,7 @@ pub mod proxy;
pub mod resilience;
pub mod router;
mod server;
mod server_utils;
mod services;
pub mod telemetry;
pub mod tray;
@@ -58,6 +59,14 @@ pub enum ProviderType {
/// Gemini API Key (multi-account load balancing)
#[serde(rename = "gemini_api_key")]
GeminiApiKey,
/// Codex (OpenAI OAuth)
Codex,
/// Claude OAuth (Anthropic OAuth)
#[serde(rename = "claude_oauth")]
ClaudeOAuth,
/// iFlow
#[serde(rename = "iflow")]
IFlow,
}
impl std::fmt::Display for ProviderType {
@@ -71,6 +80,9 @@ impl std::fmt::Display for ProviderType {
ProviderType::Antigravity => write!(f, "antigravity"),
ProviderType::Vertex => write!(f, "vertex"),
ProviderType::GeminiApiKey => write!(f, "gemini_api_key"),
ProviderType::Codex => write!(f, "codex"),
ProviderType::ClaudeOAuth => write!(f, "claude_oauth"),
ProviderType::IFlow => write!(f, "iflow"),
}
}
}
@@ -88,6 +100,9 @@ impl std::str::FromStr for ProviderType {
"antigravity" => Ok(ProviderType::Antigravity),
"vertex" => Ok(ProviderType::Vertex),
"gemini_api_key" => Ok(ProviderType::GeminiApiKey),
"codex" => Ok(ProviderType::Codex),
"claude_oauth" => Ok(ProviderType::ClaudeOAuth),
"iflow" => Ok(ProviderType::IFlow),
_ => Err(format!("Invalid provider: {s}")),
}
}
@@ -1011,6 +1026,12 @@ async fn check_api_compatibility(
("gemini-2.5-flash", "basic"),
("gemini-2.5-flash", "tool_call"),
],
ProviderType::Codex => vec![("gpt-4.1", "basic"), ("gpt-4.1", "tool_call")],
ProviderType::ClaudeOAuth => vec![
("claude-sonnet-4-5", "basic"),
("claude-sonnet-4-5", "tool_call"),
],
ProviderType::IFlow => vec![("gpt-4o", "basic"), ("gpt-4o", "tool_call")],
ProviderType::OpenAI | ProviderType::Claude => vec![],
};
+5 -63
View File
@@ -21,69 +21,11 @@ pub enum CredentialSource {
Private,
}
/// Provider 类型枚举
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum PoolProviderType {
Kiro,
Gemini,
Qwen,
#[serde(rename = "openai")]
OpenAI,
Claude,
Antigravity,
Vertex,
/// Gemini API Key (multi-account load balancing)
#[serde(rename = "gemini_api_key")]
GeminiApiKey,
/// Codex (OpenAI OAuth)
Codex,
/// Claude OAuth (Anthropic OAuth)
#[serde(rename = "claude_oauth")]
ClaudeOAuth,
/// iFlow
#[serde(rename = "iflow")]
IFlow,
}
impl std::fmt::Display for PoolProviderType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
PoolProviderType::Kiro => write!(f, "kiro"),
PoolProviderType::Gemini => write!(f, "gemini"),
PoolProviderType::Qwen => write!(f, "qwen"),
PoolProviderType::OpenAI => write!(f, "openai"),
PoolProviderType::Claude => write!(f, "claude"),
PoolProviderType::Antigravity => write!(f, "antigravity"),
PoolProviderType::Vertex => write!(f, "vertex"),
PoolProviderType::GeminiApiKey => write!(f, "gemini_api_key"),
PoolProviderType::Codex => write!(f, "codex"),
PoolProviderType::ClaudeOAuth => write!(f, "claude_oauth"),
PoolProviderType::IFlow => write!(f, "iflow"),
}
}
}
impl std::str::FromStr for PoolProviderType {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"kiro" => Ok(PoolProviderType::Kiro),
"gemini" => Ok(PoolProviderType::Gemini),
"qwen" => Ok(PoolProviderType::Qwen),
"openai" => Ok(PoolProviderType::OpenAI),
"claude" => Ok(PoolProviderType::Claude),
"antigravity" => Ok(PoolProviderType::Antigravity),
"vertex" => Ok(PoolProviderType::Vertex),
"gemini_api_key" => Ok(PoolProviderType::GeminiApiKey),
"codex" => Ok(PoolProviderType::Codex),
"claude_oauth" => Ok(PoolProviderType::ClaudeOAuth),
"iflow" => Ok(PoolProviderType::IFlow),
_ => Err(format!("Invalid provider type: {s}")),
}
}
}
/// Provider 类型别名
///
/// 为了向后兼容,PoolProviderType 是 crate::ProviderType 的类型别名。
/// 所有 Provider 类型定义已统一到 lib.rs 中的 ProviderType。
pub type PoolProviderType = crate::ProviderType;
/// 凭证数据,根据 Provider 类型不同而不同
#[derive(Debug, Clone, Serialize, Deserialize)]
+37
View File
@@ -2,6 +2,8 @@
//!
//! 支持 Gemini 3 Pro 等高级模型,通过 Google 内部 API 访问。
use super::traits::{CredentialProvider, ProviderResult};
use async_trait::async_trait;
use reqwest::Client;
use serde::{Deserialize, Serialize};
use std::error::Error;
@@ -1530,3 +1532,38 @@ pub async fn start_oauth_login(
}
}
}
// ============================================================================
// CredentialProvider Trait 实现
// ============================================================================
#[async_trait]
impl CredentialProvider for AntigravityProvider {
async fn load_credentials_from_path(&mut self, path: &str) -> ProviderResult<()> {
AntigravityProvider::load_credentials_from_path(self, path).await
}
async fn save_credentials(&self) -> ProviderResult<()> {
AntigravityProvider::save_credentials(self).await
}
fn is_token_valid(&self) -> bool {
AntigravityProvider::is_token_valid(self)
}
fn is_token_expiring_soon(&self) -> bool {
AntigravityProvider::is_token_expiring_soon(self)
}
async fn refresh_token(&mut self) -> ProviderResult<String> {
AntigravityProvider::refresh_token(self).await
}
fn get_access_token(&self) -> Option<&str> {
self.credentials.access_token.as_deref()
}
fn provider_type(&self) -> &'static str {
"antigravity"
}
}
+54
View File
@@ -6,6 +6,8 @@
use super::error::{
create_auth_error, create_config_error, create_token_refresh_error, ProviderError,
};
use super::traits::{CredentialProvider, ProviderResult};
use async_trait::async_trait;
use reqwest::Client;
use serde::{Deserialize, Serialize};
use std::error::Error;
@@ -1353,3 +1355,55 @@ pub async fn start_gemini_oauth_login(
// 等待回调
wait_future.await
}
// ============================================================================
// CredentialProvider Trait 实现
// ============================================================================
#[async_trait]
impl CredentialProvider for GeminiProvider {
async fn load_credentials_from_path(&mut self, path: &str) -> ProviderResult<()> {
GeminiProvider::load_credentials_from_path(self, path).await
}
async fn save_credentials(&self) -> ProviderResult<()> {
GeminiProvider::save_credentials(self).await
}
fn is_token_valid(&self) -> bool {
GeminiProvider::is_token_valid(self)
}
fn is_token_expiring_soon(&self) -> bool {
// Gemini 使用与 is_token_valid 相同的逻辑,但阈值为 10 分钟
if self.credentials.access_token.is_none() {
return true;
}
if let Some(expire_str) = &self.credentials.expire {
if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expire_str) {
let now = chrono::Utc::now();
return expires <= now + chrono::Duration::minutes(10);
}
}
if let Some(expiry) = self.credentials.expiry_date {
let now = chrono::Utc::now().timestamp_millis();
return expiry <= now + 600_000; // 10 分钟
}
false
}
async fn refresh_token(&mut self) -> ProviderResult<String> {
GeminiProvider::refresh_token(self).await
}
fn get_access_token(&self) -> Option<&str> {
self.credentials.access_token.as_deref()
}
fn provider_type(&self) -> &'static str {
"gemini"
}
}
+38
View File
@@ -1,6 +1,8 @@
//! Kiro/CodeWhisperer Provider
use crate::converter::openai_to_cw::convert_openai_to_codewhisperer;
use crate::models::openai::*;
use crate::providers::traits::{CredentialProvider, ProviderResult};
use async_trait::async_trait;
use reqwest::Client;
use serde::{Deserialize, Serialize};
use std::error::Error;
@@ -1075,3 +1077,39 @@ fn merge_credentials(target: &mut KiroCredentials, source: &KiroCredentials) {
}
// cred_type 使用默认值,不需要合并
}
// ============================================================================
// CredentialProvider Trait 实现
// ============================================================================
#[async_trait]
impl CredentialProvider for KiroProvider {
async fn load_credentials_from_path(&mut self, path: &str) -> ProviderResult<()> {
// 调用已有的实现
KiroProvider::load_credentials_from_path(self, path).await
}
async fn save_credentials(&self) -> ProviderResult<()> {
KiroProvider::save_credentials(self).await
}
fn is_token_valid(&self) -> bool {
!self.is_token_expired()
}
fn is_token_expiring_soon(&self) -> bool {
KiroProvider::is_token_expiring_soon(self)
}
async fn refresh_token(&mut self) -> ProviderResult<String> {
KiroProvider::refresh_token(self).await
}
fn get_access_token(&self) -> Option<&str> {
self.credentials.access_token.as_deref()
}
fn provider_type(&self) -> &'static str {
"kiro"
}
}
+6 -1
View File
@@ -8,11 +8,16 @@ pub mod iflow;
pub mod kiro;
pub mod openai_custom;
pub mod qwen;
pub mod traits;
pub mod vertex;
#[cfg(test)]
mod tests;
// Trait exports
#[allow(unused_imports)]
pub use traits::{CredentialProvider, ProviderResult, TokenManager};
#[allow(unused_imports)]
pub use antigravity::AntigravityProvider;
#[allow(unused_imports)]
@@ -22,7 +27,7 @@ pub use claude_oauth::ClaudeOAuthProvider;
#[allow(unused_imports)]
pub use codex::CodexProvider;
#[allow(unused_imports)]
pub use error::{ProviderError, ProviderResult};
pub use error::ProviderError;
#[allow(unused_imports)]
pub use gemini::{GeminiApiKeyCredential, GeminiApiKeyProvider, GeminiProvider};
#[allow(unused_imports)]
+49
View File
@@ -6,6 +6,8 @@
use super::error::{
create_auth_error, create_config_error, create_token_refresh_error, ProviderError,
};
use super::traits::{CredentialProvider, ProviderResult};
use async_trait::async_trait;
use reqwest::Client;
use serde::{Deserialize, Serialize};
use std::error::Error;
@@ -653,3 +655,50 @@ pub async fn start_qwen_device_code_and_get_info() -> Result<
Ok((device_response, wait_future))
}
// ============================================================================
// CredentialProvider Trait 实现
// ============================================================================
#[async_trait]
impl CredentialProvider for QwenProvider {
async fn load_credentials_from_path(&mut self, path: &str) -> ProviderResult<()> {
QwenProvider::load_credentials_from_path(self, path).await
}
async fn save_credentials(&self) -> ProviderResult<()> {
QwenProvider::save_credentials(self).await
}
fn is_token_valid(&self) -> bool {
QwenProvider::is_token_valid(self)
}
fn is_token_expiring_soon(&self) -> bool {
// Qwen 使用与 is_token_valid 相同的逻辑,但阈值为 10 分钟
if self.credentials.access_token.is_none() {
return true;
}
if let Some(expire_str) = &self.credentials.expire {
if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expire_str) {
let now = chrono::Utc::now();
return expires <= now + chrono::Duration::minutes(10);
}
}
false
}
async fn refresh_token(&mut self) -> ProviderResult<String> {
QwenProvider::refresh_token(self).await
}
fn get_access_token(&self) -> Option<&str> {
self.credentials.access_token.as_deref()
}
fn provider_type(&self) -> &'static str {
"qwen"
}
}
+182
View File
@@ -0,0 +1,182 @@
//! Provider Trait 定义
//!
//! 统一的 Provider 接口,用于凭证管理和 Token 生命周期管理。
use async_trait::async_trait;
use std::error::Error;
/// Provider 结果类型别名(与现有方法签名兼容)
pub type ProviderResult<T> = Result<T, Box<dyn Error + Send + Sync>>;
/// 凭证管理 Trait
///
/// 定义所有 OAuth Provider 必须实现的凭证管理接口
#[async_trait]
pub trait CredentialProvider: Send + Sync {
/// 从指定路径加载凭证
async fn load_credentials_from_path(&mut self, path: &str) -> ProviderResult<()>;
/// 保存凭证到文件
async fn save_credentials(&self) -> ProviderResult<()>;
/// 检查 Token 是否有效(未过期)
fn is_token_valid(&self) -> bool;
/// 检查 Token 是否即将过期(通常提前 5 分钟)
fn is_token_expiring_soon(&self) -> bool;
/// 刷新 Token
///
/// 返回新的 access_token
async fn refresh_token(&mut self) -> ProviderResult<String>;
/// 获取当前 access_token
fn get_access_token(&self) -> Option<&str>;
/// 获取 Provider 类型名称
fn provider_type(&self) -> &'static str;
}
/// Token 管理辅助 Trait
///
/// 提供带重试的 Token 刷新功能
#[async_trait]
pub trait TokenManager: CredentialProvider {
/// 带重试的 Token 刷新
///
/// # Arguments
/// * `max_retries` - 最大重试次数
/// * `retry_delay_ms` - 重试间隔(毫秒)
async fn refresh_token_with_retry(
&mut self,
max_retries: u32,
retry_delay_ms: u64,
) -> ProviderResult<String> {
let mut last_error = None;
for attempt in 0..=max_retries {
match self.refresh_token().await {
Ok(token) => return Ok(token),
Err(e) => {
tracing::warn!(
"[{}] Token refresh attempt {} failed: {}",
self.provider_type(),
attempt + 1,
e
);
last_error = Some(e);
if attempt < max_retries {
tokio::time::sleep(tokio::time::Duration::from_millis(retry_delay_ms))
.await;
}
}
}
}
Err(last_error.unwrap_or_else(|| "Token refresh failed".into()))
}
/// 确保 Token 有效(如需要则刷新)
async fn ensure_valid_token(&mut self) -> ProviderResult<String> {
if !self.is_token_valid() || self.is_token_expiring_soon() {
self.refresh_token().await
} else {
self.get_access_token()
.map(|s| s.to_string())
.ok_or_else(|| "No access token available".into())
}
}
}
// 为所有实现了 CredentialProvider 的类型自动实现 TokenManager
impl<T: CredentialProvider> TokenManager for T {}
#[cfg(test)]
mod tests {
use super::*;
// Mock Provider for testing
struct MockProvider {
token: Option<String>,
valid: bool,
expiring_soon: bool,
refresh_count: u32,
}
#[async_trait]
impl CredentialProvider for MockProvider {
async fn load_credentials_from_path(&mut self, _path: &str) -> ProviderResult<()> {
Ok(())
}
async fn save_credentials(&self) -> ProviderResult<()> {
Ok(())
}
fn is_token_valid(&self) -> bool {
self.valid
}
fn is_token_expiring_soon(&self) -> bool {
self.expiring_soon
}
async fn refresh_token(&mut self) -> ProviderResult<String> {
self.refresh_count += 1;
self.token = Some(format!("new_token_{}", self.refresh_count));
self.valid = true;
self.expiring_soon = false;
Ok(self.token.clone().unwrap())
}
fn get_access_token(&self) -> Option<&str> {
self.token.as_deref()
}
fn provider_type(&self) -> &'static str {
"mock"
}
}
#[tokio::test]
async fn test_ensure_valid_token_when_valid() {
let mut provider = MockProvider {
token: Some("existing_token".to_string()),
valid: true,
expiring_soon: false,
refresh_count: 0,
};
let token = provider.ensure_valid_token().await.unwrap();
assert_eq!(token, "existing_token");
assert_eq!(provider.refresh_count, 0);
}
#[tokio::test]
async fn test_ensure_valid_token_when_expiring() {
let mut provider = MockProvider {
token: Some("old_token".to_string()),
valid: true,
expiring_soon: true,
refresh_count: 0,
};
let token = provider.ensure_valid_token().await.unwrap();
assert_eq!(token, "new_token_1");
assert_eq!(provider.refresh_count, 1);
}
#[tokio::test]
async fn test_ensure_valid_token_when_invalid() {
let mut provider = MockProvider {
token: Some("invalid_token".to_string()),
valid: false,
expiring_soon: false,
refresh_count: 0,
};
let token = provider.ensure_valid_token().await.unwrap();
assert_eq!(token, "new_token_1");
assert_eq!(provider.refresh_count, 1);
}
}
File diff suppressed because it is too large Load Diff
+982
View File
@@ -0,0 +1,982 @@
//! API 端点处理器
//!
//! 处理 OpenAI 和 Anthropic 格式的 API 请求
use axum::{
body::Body,
extract::State,
http::{header, HeaderMap, StatusCode},
response::{IntoResponse, Response},
Json,
};
use futures::stream;
use crate::converter::anthropic_to_openai::convert_anthropic_to_openai;
use crate::converter::openai_to_antigravity::{
convert_antigravity_to_openai_response, convert_openai_to_antigravity_with_context,
};
use crate::models::anthropic::AnthropicMessagesRequest;
use crate::models::openai::ChatCompletionRequest;
use crate::processor::RequestContext;
use crate::providers::{AntigravityProvider, GeminiProvider, KiroProvider, QwenProvider};
use crate::server::{record_request_telemetry, record_token_usage, AppState};
use crate::server_utils::{
build_anthropic_response, build_anthropic_stream_response, message_content_len,
parse_cw_response, safe_truncate,
};
use crate::telemetry::RequestStatus;
use crate::ProviderType;
use super::{call_provider_anthropic, call_provider_openai};
/// OpenAI 格式的 API key 验证
pub async fn verify_api_key(
headers: &HeaderMap,
expected_key: &str,
) -> Result<(), (StatusCode, Json<serde_json::Value>)> {
let auth = headers
.get("authorization")
.or_else(|| headers.get("x-api-key"))
.and_then(|v| v.to_str().ok());
let key = match auth {
Some(s) if s.starts_with("Bearer ") => &s[7..],
Some(s) => s,
None => {
return Err((
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": "No API key provided"}})),
))
}
};
if key != expected_key {
return Err((
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": "Invalid API key"}})),
));
}
Ok(())
}
/// Anthropic 格式的 API key 验证
pub async fn verify_api_key_anthropic(
headers: &HeaderMap,
expected_key: &str,
) -> Result<(), (StatusCode, Json<serde_json::Value>)> {
let auth = headers
.get("x-api-key")
.or_else(|| headers.get("authorization"))
.and_then(|v| v.to_str().ok());
let key = match auth {
Some(s) if s.starts_with("Bearer ") => &s[7..],
Some(s) => s,
None => {
return Err((
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({
"type": "error",
"error": {
"type": "authentication_error",
"message": "No API key provided. Please set the x-api-key header."
}
})),
))
}
};
if key != expected_key {
return Err((
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({
"type": "error",
"error": {
"type": "authentication_error",
"message": "Invalid API key"
}
})),
));
}
Ok(())
}
pub async fn chat_completions(
State(state): State<AppState>,
headers: HeaderMap,
Json(mut request): Json<ChatCompletionRequest>,
) -> Response {
if let Err(e) = verify_api_key(&headers, &state.api_key).await {
state
.logs
.write()
.await
.add("warn", "Unauthorized request to /v1/chat/completions");
return e.into_response();
}
// 创建请求上下文
let mut ctx = RequestContext::new(request.model.clone()).with_stream(request.stream);
state.logs.write().await.add(
"info",
&format!(
"POST /v1/chat/completions request_id={} model={} stream={}",
ctx.request_id, request.model, request.stream
),
);
// 使用 RequestProcessor 解析模型别名和路由
let provider = state.processor.resolve_and_route(&mut ctx).await;
// 更新请求中的模型名为解析后的模型
if ctx.resolved_model != ctx.original_model {
request.model = ctx.resolved_model.clone();
state.logs.write().await.add(
"info",
&format!(
"[MAPPER] request_id={} alias={} -> model={}",
ctx.request_id, ctx.original_model, ctx.resolved_model
),
);
}
// 应用参数注入
let injection_enabled = *state.injection_enabled.read().await;
if injection_enabled {
let injector = state.processor.injector.read().await;
let mut payload = serde_json::to_value(&request).unwrap_or_default();
let result = injector.inject(&request.model, &mut payload);
if result.has_injections() {
state.logs.write().await.add(
"info",
&format!(
"[INJECT] request_id={} applied_rules={:?} injected_params={:?}",
ctx.request_id, result.applied_rules, result.injected_params
),
);
// 更新请求
if let Ok(updated) = serde_json::from_value(payload) {
request = updated;
}
}
}
// 获取当前默认 provider(用于凭证池选择)
let default_provider = state.default_provider.read().await.clone();
// 记录路由结果
state.logs.write().await.add(
"info",
&format!(
"[ROUTE] request_id={} model={} provider={}",
ctx.request_id, ctx.resolved_model, provider
),
);
// 尝试从凭证池中选择凭证
let credential = match &state.db {
Some(db) => state
.pool_service
.select_credential(db, &default_provider, Some(&request.model))
.ok()
.flatten(),
None => None,
};
// 如果找到凭证池中的凭证,使用它
if let Some(cred) = credential {
state.logs.write().await.add(
"info",
&format!(
"[ROUTE] Using pool credential: type={} name={:?} uuid={}",
cred.provider_type,
cred.name,
&cred.uuid[..8]
),
);
let response = call_provider_openai(&state, &cred, &request).await;
// 记录请求统计
let is_success = response.status().is_success();
let status = if is_success {
crate::telemetry::RequestStatus::Success
} else {
crate::telemetry::RequestStatus::Failed
};
record_request_telemetry(&state, &ctx, status, None);
// 如果成功,记录估算的 Token 使用量
if is_success {
let estimated_input_tokens = request
.messages
.iter()
.map(|m| {
let content_len = match &m.content {
Some(c) => message_content_len(c),
None => 0,
};
content_len / 4
})
.sum::<usize>() as u32;
// 输出 Token 使用估算值(假设平均响应长度)
let estimated_output_tokens = 100u32;
record_token_usage(
&state,
&ctx,
Some(estimated_input_tokens),
Some(estimated_output_tokens),
);
}
return response;
}
// 回退到旧的单凭证模式
state.logs.write().await.add(
"debug",
&format!(
"[ROUTE] No pool credential found for '{}', using legacy mode",
default_provider
),
);
// 检查是否需要刷新 token(无 token 或即将过期)
{
let _guard = state.kiro_refresh_lock.lock().await;
let mut kiro = state.kiro.write().await;
let needs_refresh =
kiro.credentials.access_token.is_none() || kiro.is_token_expiring_soon();
if needs_refresh {
if let Err(e) = kiro.refresh_token().await {
state
.logs
.write()
.await
.add("error", &format!("Token refresh failed: {e}"));
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})),
).into_response();
}
}
}
let kiro = state.kiro.read().await;
match kiro.call_api(&request).await {
Ok(resp) => {
let status = resp.status();
if status.is_success() {
match resp.text().await {
Ok(body) => {
let parsed = parse_cw_response(&body);
let has_tool_calls = !parsed.tool_calls.is_empty();
state.logs.write().await.add(
"info",
&format!(
"Request completed: content_len={}, tool_calls={}",
parsed.content.len(),
parsed.tool_calls.len()
),
);
// 构建消息
let message = if has_tool_calls {
serde_json::json!({
"role": "assistant",
"content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) },
"tool_calls": parsed.tool_calls.iter().map(|tc| {
serde_json::json!({
"id": tc.id,
"type": "function",
"function": {
"name": tc.function.name,
"arguments": tc.function.arguments
}
})
}).collect::<Vec<_>>()
})
} else {
serde_json::json!({
"role": "assistant",
"content": parsed.content
})
};
// 估算 Token 数量(基于字符数,约 4 字符 = 1 token)
let estimated_output_tokens = (parsed.content.len() / 4) as u32;
// 估算输入 Token(基于请求消息)
let estimated_input_tokens = request
.messages
.iter()
.map(|m| {
let content_len = match &m.content {
Some(c) => message_content_len(c),
None => 0,
};
content_len / 4
})
.sum::<usize>()
as u32;
let response = serde_json::json!({
"id": format!("chatcmpl-{}", uuid::Uuid::new_v4()),
"object": "chat.completion",
"created": std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
"model": request.model,
"choices": [{
"index": 0,
"message": message,
"finish_reason": if has_tool_calls { "tool_calls" } else { "stop" }
}],
"usage": {
"prompt_tokens": estimated_input_tokens,
"completion_tokens": estimated_output_tokens,
"total_tokens": estimated_input_tokens + estimated_output_tokens
}
});
// 记录成功请求统计
record_request_telemetry(
&state,
&ctx,
crate::telemetry::RequestStatus::Success,
None,
);
// 记录 Token 使用量
record_token_usage(
&state,
&ctx,
Some(estimated_input_tokens),
Some(estimated_output_tokens),
);
Json(response).into_response()
}
Err(e) => {
// 记录失败请求统计
record_request_telemetry(
&state,
&ctx,
crate::telemetry::RequestStatus::Failed,
Some(e.to_string()),
);
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
} else if status.as_u16() == 403 || status.as_u16() == 402 {
// Token 过期或账户问题,尝试重新加载凭证并刷新
drop(kiro);
let _guard = state.kiro_refresh_lock.lock().await;
let mut kiro = state.kiro.write().await;
state.logs.write().await.add(
"warn",
&format!(
"[AUTH] Got {}, reloading credentials and attempting token refresh...",
status.as_u16()
),
);
// 先重新加载凭证文件(可能用户换了账户)
if let Err(e) = kiro.load_credentials().await {
state.logs.write().await.add(
"error",
&format!("[AUTH] Failed to reload credentials: {e}"),
);
}
match kiro.refresh_token().await {
Ok(_) => {
state
.logs
.write()
.await
.add("info", "[AUTH] Token refreshed successfully after reload");
// 重试请求
drop(kiro);
let kiro = state.kiro.read().await;
match kiro.call_api(&request).await {
Ok(retry_resp) => {
if retry_resp.status().is_success() {
match retry_resp.text().await {
Ok(body) => {
let parsed = parse_cw_response(&body);
let has_tool_calls = !parsed.tool_calls.is_empty();
let message = if has_tool_calls {
serde_json::json!({
"role": "assistant",
"content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) },
"tool_calls": parsed.tool_calls.iter().map(|tc| {
serde_json::json!({
"id": tc.id,
"type": "function",
"function": {
"name": tc.function.name,
"arguments": tc.function.arguments
}
})
}).collect::<Vec<_>>()
})
} else {
serde_json::json!({
"role": "assistant",
"content": parsed.content
})
};
let response = serde_json::json!({
"id": format!("chatcmpl-{}", uuid::Uuid::new_v4()),
"object": "chat.completion",
"created": std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
"model": request.model,
"choices": [{
"index": 0,
"message": message,
"finish_reason": if has_tool_calls { "tool_calls" } else { "stop" }
}],
"usage": {
"prompt_tokens": 0,
"completion_tokens": 0,
"total_tokens": 0
}
});
return Json(response).into_response();
}
Err(e) => return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
).into_response(),
}
}
let body = retry_resp.text().await.unwrap_or_default();
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": format!("Retry failed: {}", body)}})),
).into_response()
}
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response(),
}
}
Err(e) => {
state
.logs
.write()
.await
.add("error", &format!("[AUTH] Token refresh failed: {e}"));
(
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})),
)
.into_response()
}
}
} else {
let body = resp.text().await.unwrap_or_default();
state.logs.write().await.add(
"error",
&format!("Upstream error {}: {}", status, safe_truncate(&body, 200)),
);
(
StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR),
Json(serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}}))
).into_response()
}
}
Err(e) => {
state
.logs
.write()
.await
.add("error", &format!("API call failed: {e}"));
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
}
pub async fn anthropic_messages(
State(state): State<AppState>,
headers: HeaderMap,
Json(mut request): Json<AnthropicMessagesRequest>,
) -> Response {
// 使用 Anthropic 格式的认证验证(优先检查 x-api-key)
if let Err(e) = verify_api_key_anthropic(&headers, &state.api_key).await {
state
.logs
.write()
.await
.add("warn", "Unauthorized request to /v1/messages");
return e.into_response();
}
// 创建请求上下文
let mut ctx = RequestContext::new(request.model.clone()).with_stream(request.stream);
// 详细记录请求信息
let msg_count = request.messages.len();
let has_tools = request.tools.as_ref().map(|t| t.len()).unwrap_or(0);
let has_system = request.system.is_some();
state.logs.write().await.add(
"info",
&format!(
"[REQ] POST /v1/messages request_id={} model={} stream={} messages={} tools={} has_system={}",
ctx.request_id, request.model, request.stream, msg_count, has_tools, has_system
),
);
// 使用 RequestProcessor 解析模型别名和路由
let provider = state.processor.resolve_and_route(&mut ctx).await;
// 更新请求中的模型名为解析后的模型
if ctx.resolved_model != ctx.original_model {
request.model = ctx.resolved_model.clone();
state.logs.write().await.add(
"info",
&format!(
"[MAPPER] request_id={} alias={} -> model={}",
ctx.request_id, ctx.original_model, ctx.resolved_model
),
);
}
// 记录最后一条消息的角色和内容预览
if let Some(last_msg) = request.messages.last() {
let content_preview = match &last_msg.content {
serde_json::Value::String(s) => s.chars().take(100).collect::<String>(),
serde_json::Value::Array(arr) => {
if let Some(first) = arr.first() {
if let Some(text) = first.get("text").and_then(|t| t.as_str()) {
text.chars().take(100).collect::<String>()
} else {
format!("[{} blocks]", arr.len())
}
} else {
"[empty]".to_string()
}
}
_ => "[unknown]".to_string(),
};
state.logs.write().await.add(
"debug",
&format!(
"[REQ] request_id={} last_message: role={} content={}",
ctx.request_id, last_msg.role, content_preview
),
);
}
// 应用参数注入
let injection_enabled = *state.injection_enabled.read().await;
if injection_enabled {
let injector = state.processor.injector.read().await;
let mut payload = serde_json::to_value(&request).unwrap_or_default();
let result = injector.inject(&request.model, &mut payload);
if result.has_injections() {
state.logs.write().await.add(
"info",
&format!(
"[INJECT] request_id={} applied_rules={:?} injected_params={:?}",
ctx.request_id, result.applied_rules, result.injected_params
),
);
// 更新请求
if let Ok(updated) = serde_json::from_value(payload) {
request = updated;
}
}
}
// 获取当前默认 provider(用于凭证池选择)
let default_provider = state.default_provider.read().await.clone();
// 记录路由结果
state.logs.write().await.add(
"info",
&format!(
"[ROUTE] request_id={} model={} provider={}",
ctx.request_id, ctx.resolved_model, provider
),
);
// 尝试从凭证池中选择凭证
let credential = match &state.db {
Some(db) => {
// 根据 default_provider 配置选择凭证
state
.pool_service
.select_credential(db, &default_provider, Some(&request.model))
.ok()
.flatten()
}
None => None,
};
// 如果找到凭证池中的凭证,使用它
if let Some(cred) = credential {
state.logs.write().await.add(
"info",
&format!(
"[ROUTE] Using pool credential: type={} name={:?} uuid={}",
cred.provider_type,
cred.name,
&cred.uuid[..8]
),
);
let response = call_provider_anthropic(&state, &cred, &request).await;
// 记录请求统计
let is_success = response.status().is_success();
let status = if is_success {
crate::telemetry::RequestStatus::Success
} else {
crate::telemetry::RequestStatus::Failed
};
record_request_telemetry(&state, &ctx, status, None);
// 如果成功,记录估算的 Token 使用量
if is_success {
let estimated_input_tokens = request
.messages
.iter()
.map(|m| {
let content_len = match &m.content {
serde_json::Value::String(s) => s.len(),
serde_json::Value::Array(arr) => arr
.iter()
.filter_map(|v| v.get("text").and_then(|t| t.as_str()))
.map(|s| s.len())
.sum(),
_ => 0,
};
content_len / 4
})
.sum::<usize>() as u32;
// 输出 Token 使用估算值
let estimated_output_tokens = 100u32;
record_token_usage(
&state,
&ctx,
Some(estimated_input_tokens),
Some(estimated_output_tokens),
);
}
return response;
}
// 回退到旧的单凭证模式
state.logs.write().await.add(
"debug",
&format!(
"[ROUTE] No pool credential found for '{}', using legacy mode",
default_provider
),
);
// 检查是否需要刷新 token(无 token 或即将过期)
{
let _guard = state.kiro_refresh_lock.lock().await;
let mut kiro = state.kiro.write().await;
let needs_refresh =
kiro.credentials.access_token.is_none() || kiro.is_token_expiring_soon();
if needs_refresh {
state.logs.write().await.add(
"info",
"[AUTH] No access token or token expiring soon, attempting refresh...",
);
if let Err(e) = kiro.refresh_token().await {
state
.logs
.write()
.await
.add("error", &format!("[AUTH] Token refresh failed: {e}"));
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})),
)
.into_response();
}
state
.logs
.write()
.await
.add("info", "[AUTH] Token refreshed successfully");
}
}
// 转换为 OpenAI 格式
let openai_request = convert_anthropic_to_openai(&request);
// 记录转换后的请求信息
state.logs.write().await.add(
"debug",
&format!(
"[CONVERT] OpenAI format: messages={} tools={} stream={}",
openai_request.messages.len(),
openai_request.tools.as_ref().map(|t| t.len()).unwrap_or(0),
openai_request.stream
),
);
let kiro = state.kiro.read().await;
match kiro.call_api(&openai_request).await {
Ok(resp) => {
let status = resp.status();
state
.logs
.write()
.await
.add("info", &format!("[RESP] Upstream status: {status}"));
if status.is_success() {
match resp.bytes().await {
Ok(bytes) => {
// 使用 lossy 转换,避免无效 UTF-8 导致崩溃
let body = String::from_utf8_lossy(&bytes).to_string();
// 记录原始响应长度
state.logs.write().await.add(
"debug",
&format!("[RESP] Raw body length: {} bytes", bytes.len()),
);
// 保存原始响应到文件用于调试
let request_id = uuid::Uuid::new_v4().to_string()[..8].to_string();
state.logs.read().await.log_raw_response(&request_id, &body);
state.logs.write().await.add(
"debug",
&format!("[RESP] Raw response saved to raw_response_{request_id}.txt"),
);
// 记录响应的前200字符用于调试(减少日志量)
let preview: String =
body.chars().filter(|c| !c.is_control()).take(200).collect();
state
.logs
.write()
.await
.add("debug", &format!("[RESP] Body preview: {preview}"));
let parsed = parse_cw_response(&body);
// 详细记录解析结果
state.logs.write().await.add(
"info",
&format!(
"[RESP] Parsed: content_len={}, tool_calls={}, content_preview={}",
parsed.content.len(),
parsed.tool_calls.len(),
parsed.content.chars().take(100).collect::<String>()
),
);
// 记录 tool calls 详情
for (i, tc) in parsed.tool_calls.iter().enumerate() {
state.logs.write().await.add(
"debug",
&format!(
"[RESP] Tool call {}: name={} id={}",
i, tc.function.name, tc.id
),
);
}
// 如果请求流式响应,返回 SSE 格式
if request.stream {
return build_anthropic_stream_response(&request.model, &parsed);
}
// 非流式响应
build_anthropic_response(&request.model, &parsed)
}
Err(e) => {
state
.logs
.write()
.await
.add("error", &format!("[ERROR] Response body read failed: {e}"));
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
} else if status.as_u16() == 403 || status.as_u16() == 402 {
// Token 过期或账户问题,尝试重新加载凭证并刷新
drop(kiro);
let _guard = state.kiro_refresh_lock.lock().await;
let mut kiro = state.kiro.write().await;
state.logs.write().await.add(
"warn",
&format!(
"[AUTH] Got {}, reloading credentials and attempting token refresh...",
status.as_u16()
),
);
// 先重新加载凭证文件(可能用户换了账户)
if let Err(e) = kiro.load_credentials().await {
state.logs.write().await.add(
"error",
&format!("[AUTH] Failed to reload credentials: {e}"),
);
}
match kiro.refresh_token().await {
Ok(_) => {
state.logs.write().await.add(
"info",
"[AUTH] Token refreshed successfully, retrying request...",
);
drop(kiro);
let kiro = state.kiro.read().await;
match kiro.call_api(&openai_request).await {
Ok(retry_resp) => {
let retry_status = retry_resp.status();
state.logs.write().await.add(
"info",
&format!("[RETRY] Response status: {retry_status}"),
);
if retry_resp.status().is_success() {
match retry_resp.bytes().await {
Ok(bytes) => {
let body = String::from_utf8_lossy(&bytes).to_string();
let parsed = parse_cw_response(&body);
state.logs.write().await.add(
"info",
&format!(
"[RETRY] Success: content_len={}, tool_calls={}",
parsed.content.len(), parsed.tool_calls.len()
),
);
if request.stream {
return build_anthropic_stream_response(
&request.model,
&parsed,
);
}
return build_anthropic_response(
&request.model,
&parsed,
);
}
Err(e) => {
state.logs.write().await.add(
"error",
&format!("[RETRY] Body read failed: {e}"),
);
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response();
}
}
}
let body = retry_resp
.bytes()
.await
.map(|b| String::from_utf8_lossy(&b).to_string())
.unwrap_or_default();
state.logs.write().await.add(
"error",
&format!(
"[RETRY] Failed with status {retry_status}: {}",
safe_truncate(&body, 500)
),
);
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": format!("Retry failed: {}", body)}})),
)
.into_response()
}
Err(e) => {
state
.logs
.write()
.await
.add("error", &format!("[RETRY] Request failed: {e}"));
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
}
Err(e) => {
state
.logs
.write()
.await
.add("error", &format!("[AUTH] Token refresh failed: {e}"));
(
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})),
)
.into_response()
}
}
} else {
let body = resp.text().await.unwrap_or_default();
state.logs.write().await.add(
"error",
&format!(
"[ERROR] Upstream error HTTP {}: {}",
status,
safe_truncate(&body, 500)
),
);
(
StatusCode::from_u16(status.as_u16())
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR),
Json(
serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}}),
),
)
.into_response()
}
}
Err(e) => {
// 详细记录网络/连接错误
let error_details = format!("{e:?}");
state
.logs
.write()
.await
.add("error", &format!("[ERROR] Kiro API call failed: {e}"));
state.logs.write().await.add(
"debug",
&format!("[ERROR] Full error details: {error_details}"),
);
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
}
+557
View File
@@ -0,0 +1,557 @@
//! Management API 处理器
//!
//! 提供服务器状态查询、凭证管理、配置管理等功能
use axum::{extract::State, http::StatusCode, response::IntoResponse, Json};
use serde::{Deserialize, Serialize};
use crate::database::dao::provider_pool::ProviderPoolDao;
use crate::server::AppState;
// ============ Types ============
/// 管理 API 状态响应
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ManagementStatusResponse {
/// 服务器是否运行中
pub running: bool,
/// 监听地址
pub host: String,
/// 监听端口
pub port: u16,
/// 处理的请求数
pub requests: u64,
/// 运行时间(秒)
pub uptime_secs: u64,
/// 版本号
pub version: String,
/// TLS 是否启用
pub tls_enabled: bool,
/// 默认 Provider
pub default_provider: String,
}
/// 凭证信息(用于列表显示)
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CredentialInfo {
/// 凭证 ID
pub id: String,
/// Provider 类型
pub provider_type: String,
/// 是否禁用
pub disabled: bool,
/// 是否有效
pub is_valid: bool,
}
/// 凭证列表响应
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CredentialsListResponse {
/// 凭证列表
pub credentials: Vec<CredentialInfo>,
/// 总数
pub total: usize,
}
/// 添加凭证请求
#[derive(Debug, Clone, Deserialize)]
pub struct AddCredentialRequest {
/// Provider 类型
pub provider_type: String,
/// 凭证 ID
pub id: String,
/// API Key(用于 API Key 类型的凭证)
#[serde(default)]
pub api_key: Option<String>,
/// Token 文件路径(用于 OAuth 类型的凭证)
#[serde(default)]
pub token_file: Option<String>,
/// Base URL
#[serde(default)]
pub base_url: Option<String>,
/// 代理 URL
#[serde(default)]
pub proxy_url: Option<String>,
}
/// 添加凭证响应
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AddCredentialResponse {
/// 是否成功
pub success: bool,
/// 消息
pub message: String,
/// 凭证 ID
pub id: Option<String>,
}
/// 配置响应(简化版,不包含敏感信息)
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ManagementConfigResponse {
/// 服务器配置
pub server: ManagementServerConfigInfo,
/// 路由配置
pub routing: ManagementRoutingConfigInfo,
/// 重试配置
pub retry: ManagementRetryConfigInfo,
/// 远程管理配置(不包含 secret_key)
pub remote_management: ManagementRemoteInfo,
}
/// 服务器配置信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ManagementServerConfigInfo {
pub host: String,
pub port: u16,
pub tls_enabled: bool,
}
/// 路由配置信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ManagementRoutingConfigInfo {
pub default_provider: String,
pub rules_count: usize,
}
/// 重试配置信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ManagementRetryConfigInfo {
pub max_retries: u32,
pub base_delay_ms: u64,
pub max_delay_ms: u64,
}
/// 远程管理配置信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ManagementRemoteInfo {
pub allow_remote: bool,
pub has_secret_key: bool,
pub disable_control_panel: bool,
}
/// 更新配置请求
#[derive(Debug, Clone, Deserialize)]
pub struct UpdateConfigRequest {
/// 默认 Provider
#[serde(default)]
pub default_provider: Option<String>,
/// 是否允许远程访问
#[serde(default)]
pub allow_remote: Option<bool>,
}
/// 更新配置响应
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct UpdateConfigResponse {
pub success: bool,
pub message: String,
}
// ============ Handlers ============
/// GET /v0/management/status - 获取服务器状态
pub async fn management_status(State(state): State<AppState>) -> impl IntoResponse {
let default_provider = state.default_provider.read().await.clone();
// 获取请求数量
let requests = state.processor.stats.read().len() as u64;
let response = ManagementStatusResponse {
running: true,
host: "0.0.0.0".to_string(),
port: 8999,
requests,
uptime_secs: 0, // TODO: Track actual uptime
version: env!("CARGO_PKG_VERSION").to_string(),
tls_enabled: false,
default_provider,
};
Json(response)
}
/// GET /v0/management/credentials - 获取凭证列表
pub async fn management_list_credentials(State(state): State<AppState>) -> impl IntoResponse {
let mut credentials = Vec::new();
// 从数据库获取凭证列表
if let Some(ref db) = state.db {
if let Ok(conn) = db.lock() {
if let Ok(pool_credentials) = ProviderPoolDao::get_all(&conn) {
for cred in pool_credentials {
credentials.push(CredentialInfo {
id: cred.uuid.clone(),
provider_type: cred.provider_type.to_string(),
disabled: cred.is_disabled,
is_valid: cred.is_healthy,
});
}
}
}
}
let total = credentials.len();
Json(CredentialsListResponse { credentials, total })
}
/// POST /v0/management/credentials - 添加凭证
pub async fn management_add_credential(
State(state): State<AppState>,
Json(request): Json<AddCredentialRequest>,
) -> impl IntoResponse {
use crate::models::provider_pool_model::{
CredentialData, PoolProviderType, ProviderCredential,
};
// 验证请求
if request.id.is_empty() {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "Credential ID is required".to_string(),
id: None,
}),
);
}
if request.provider_type.is_empty() {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "Provider type is required".to_string(),
id: None,
}),
);
}
// 解析 provider 类型
let provider_type: PoolProviderType = match request.provider_type.parse() {
Ok(pt) => pt,
Err(_) => {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: format!("Invalid provider type: {}", request.provider_type),
id: None,
}),
);
}
};
// 根据 provider 类型创建凭证数据
let credential_data = match provider_type {
PoolProviderType::OpenAI => {
if let Some(api_key) = request.api_key {
CredentialData::OpenAIKey {
api_key,
base_url: request.base_url,
}
} else {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "API key is required for OpenAI provider".to_string(),
id: None,
}),
);
}
}
PoolProviderType::Claude => {
if let Some(api_key) = request.api_key {
CredentialData::ClaudeKey {
api_key,
base_url: request.base_url,
}
} else {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "API key is required for Claude provider".to_string(),
id: None,
}),
);
}
}
PoolProviderType::Vertex => {
if let Some(api_key) = request.api_key {
CredentialData::VertexKey {
api_key,
base_url: request.base_url,
model_aliases: std::collections::HashMap::new(),
}
} else {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "API key is required for Vertex provider".to_string(),
id: None,
}),
);
}
}
PoolProviderType::Kiro => {
if let Some(token_file) = request.token_file {
CredentialData::KiroOAuth {
creds_file_path: token_file,
}
} else {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "Token file is required for Kiro provider".to_string(),
id: None,
}),
);
}
}
PoolProviderType::Gemini => {
if let Some(token_file) = request.token_file {
CredentialData::GeminiOAuth {
creds_file_path: token_file,
project_id: None,
}
} else {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "Token file is required for Gemini provider".to_string(),
id: None,
}),
);
}
}
PoolProviderType::Qwen => {
if let Some(token_file) = request.token_file {
CredentialData::QwenOAuth {
creds_file_path: token_file,
}
} else {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "Token file is required for Qwen provider".to_string(),
id: None,
}),
);
}
}
PoolProviderType::Antigravity => {
if let Some(token_file) = request.token_file {
CredentialData::AntigravityOAuth {
creds_file_path: token_file,
project_id: None,
}
} else {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "Token file is required for Antigravity provider".to_string(),
id: None,
}),
);
}
}
PoolProviderType::GeminiApiKey => {
if let Some(api_key) = request.api_key {
CredentialData::GeminiApiKey {
api_key,
base_url: request.base_url,
excluded_models: Vec::new(),
}
} else {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "API key is required for Gemini API Key provider".to_string(),
id: None,
}),
);
}
}
PoolProviderType::Codex => {
if let Some(token_file) = request.token_file {
CredentialData::CodexOAuth {
creds_file_path: token_file,
}
} else {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "Token file is required for Codex provider".to_string(),
id: None,
}),
);
}
}
PoolProviderType::ClaudeOAuth => {
if let Some(token_file) = request.token_file {
CredentialData::ClaudeOAuth {
creds_file_path: token_file,
}
} else {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "Token file is required for Claude OAuth provider".to_string(),
id: None,
}),
);
}
}
PoolProviderType::IFlow => {
if let Some(token_file) = request.token_file {
// 默认使用 OAuth 类型,Cookie 类型需要通过其他方式添加
CredentialData::IFlowOAuth {
creds_file_path: token_file,
}
} else {
return (
StatusCode::BAD_REQUEST,
Json(AddCredentialResponse {
success: false,
message: "Token file is required for iFlow provider".to_string(),
id: None,
}),
);
}
}
};
// 创建凭证
let mut credential = ProviderCredential::new(provider_type, credential_data);
credential.uuid = request.id.clone();
credential.name = Some(request.id.clone());
// 添加凭证到数据库
if let Some(ref db) = state.db {
if let Ok(conn) = db.lock() {
match ProviderPoolDao::insert(&conn, &credential) {
Ok(_) => {
tracing::info!(
"[MANAGEMENT] Added credential: {} ({})",
request.id,
request.provider_type
);
return (
StatusCode::CREATED,
Json(AddCredentialResponse {
success: true,
message: "Credential added successfully".to_string(),
id: Some(request.id),
}),
);
}
Err(e) => {
tracing::error!("[MANAGEMENT] Failed to add credential: {}", e);
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(AddCredentialResponse {
success: false,
message: format!("Failed to add credential: {}", e),
id: None,
}),
);
}
}
}
}
(
StatusCode::SERVICE_UNAVAILABLE,
Json(AddCredentialResponse {
success: false,
message: "Database not available".to_string(),
id: None,
}),
)
}
/// GET /v0/management/config - 获取配置
pub async fn management_get_config(State(state): State<AppState>) -> impl IntoResponse {
let default_provider = state.default_provider.read().await.clone();
// 获取路由规则数量
let rules_count = state.processor.router.read().await.rules().len();
let response = ManagementConfigResponse {
server: ManagementServerConfigInfo {
host: "0.0.0.0".to_string(),
port: 8999,
tls_enabled: false,
},
routing: ManagementRoutingConfigInfo {
default_provider,
rules_count,
},
retry: ManagementRetryConfigInfo {
max_retries: 3,
base_delay_ms: 1000,
max_delay_ms: 30000,
},
remote_management: ManagementRemoteInfo {
allow_remote: false,
has_secret_key: true,
disable_control_panel: false,
},
};
Json(response)
}
/// PUT /v0/management/config - 更新配置
pub async fn management_update_config(
State(state): State<AppState>,
Json(request): Json<UpdateConfigRequest>,
) -> impl IntoResponse {
let mut updated = false;
// 更新默认 Provider
if let Some(provider) = request.default_provider {
// 验证 provider 类型
if provider.parse::<crate::ProviderType>().is_ok() {
let mut dp = state.default_provider.write().await;
*dp = provider.clone();
tracing::info!("[MANAGEMENT] Updated default_provider to: {}", provider);
updated = true;
} else {
return (
StatusCode::BAD_REQUEST,
Json(UpdateConfigResponse {
success: false,
message: format!("Invalid provider type: {}", provider),
}),
);
}
}
if updated {
(
StatusCode::OK,
Json(UpdateConfigResponse {
success: true,
message: "Configuration updated successfully".to_string(),
}),
)
} else {
(
StatusCode::OK,
Json(UpdateConfigResponse {
success: true,
message: "No changes applied".to_string(),
}),
)
}
}
+13
View File
@@ -0,0 +1,13 @@
//! HTTP 请求处理器模块
//!
//! 将 server 中的各类处理器拆分到独立文件
pub mod api;
pub mod management;
pub mod provider_calls;
pub mod websocket;
pub use api::*;
pub use management::*;
pub use provider_calls::*;
pub use websocket::*;
@@ -0,0 +1,925 @@
//! Provider 调用处理器
//!
//! 根据凭证类型调用不同的 Provider API
use axum::{
body::Body,
http::{header, StatusCode},
response::{IntoResponse, Response},
Json,
};
use futures::stream;
use crate::converter::anthropic_to_openai::convert_anthropic_to_openai;
use crate::converter::openai_to_antigravity::{
convert_antigravity_to_openai_response, convert_openai_to_antigravity_with_context,
};
use crate::models::anthropic::AnthropicMessagesRequest;
use crate::models::openai::ChatCompletionRequest;
use crate::models::provider_pool_model::{CredentialData, ProviderCredential};
use crate::providers::{
AntigravityProvider, ClaudeCustomProvider, GeminiProvider, KiroProvider, OpenAICustomProvider,
QwenProvider, VertexProvider,
};
use crate::server::AppState;
use crate::server_utils::{
build_anthropic_response, build_anthropic_stream_response, parse_cw_response, safe_truncate,
CWParsedResponse,
};
/// 根据凭证调用 Provider (Anthropic 格式)
pub async fn call_provider_anthropic(
state: &AppState,
credential: &ProviderCredential,
request: &AnthropicMessagesRequest,
) -> Response {
match &credential.credential {
CredentialData::KiroOAuth { creds_file_path } => {
// 使用 TokenCacheService 获取有效 token
let db = match &state.db {
Some(db) => db,
None => {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": "Database not available"}})),
)
.into_response();
}
};
// 获取缓存的 token
let token = match state
.token_cache
.get_valid_token(db, &credential.uuid)
.await
{
Ok(t) => t,
Err(e) => {
tracing::warn!("[POOL] Token cache miss, loading from source: {}", e);
// 回退到从源文件加载
let mut kiro = KiroProvider::new();
if let Err(e) = kiro.load_credentials_from_path(creds_file_path).await {
// 记录凭证加载失败
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&format!("Failed to load credentials: {}", e)),
);
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": format!("Failed to load Kiro credentials: {}", e)}})),
)
.into_response();
}
if let Err(e) = kiro.refresh_token().await {
// 记录 Token 刷新失败
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&format!("Token refresh failed: {}", e)),
);
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})),
)
.into_response();
}
kiro.credentials.access_token.unwrap_or_default()
}
};
// 使用获取到的 token 创建 KiroProvider
let mut kiro = KiroProvider::new();
kiro.credentials.access_token = Some(token);
// 从源文件加载其他配置(region, profile_arn 等)
let _ = kiro.load_credentials_from_path(creds_file_path).await;
let openai_request = convert_anthropic_to_openai(request);
let resp = match kiro.call_api(&openai_request).await {
Ok(r) => r,
Err(e) => {
// 记录 API 调用失败
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response();
}
};
let status = resp.status();
if status.is_success() {
match resp.bytes().await {
Ok(bytes) => {
let body = String::from_utf8_lossy(&bytes).to_string();
let parsed = parse_cw_response(&body);
// 记录成功
let _ = state.pool_service.mark_healthy(
db,
&credential.uuid,
Some(&request.model),
);
let _ = state.pool_service.record_usage(db, &credential.uuid);
if request.stream {
build_anthropic_stream_response(&request.model, &parsed)
} else {
build_anthropic_response(&request.model, &parsed)
}
}
Err(e) => {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
} else if status.as_u16() == 401 || status.as_u16() == 403 {
// Token 过期,强制刷新并重试
tracing::info!(
"[POOL] Got {}, forcing token refresh for {}",
status,
&credential.uuid[..8]
);
let new_token = match state
.token_cache
.refresh_and_cache(db, &credential.uuid, true)
.await
{
Ok(t) => t,
Err(e) => {
// 记录 Token 刷新失败
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&format!("Token refresh failed: {}", e)),
);
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})),
)
.into_response();
}
};
// 使用新 token 重试
kiro.credentials.access_token = Some(new_token);
match kiro.call_api(&openai_request).await {
Ok(retry_resp) => {
if retry_resp.status().is_success() {
match retry_resp.bytes().await {
Ok(bytes) => {
let body = String::from_utf8_lossy(&bytes).to_string();
let parsed = parse_cw_response(&body);
// 记录重试成功
let _ = state.pool_service.mark_healthy(
db,
&credential.uuid,
Some(&request.model),
);
let _ = state.pool_service.record_usage(db, &credential.uuid);
if request.stream {
build_anthropic_stream_response(&request.model, &parsed)
} else {
build_anthropic_response(&request.model, &parsed)
}
}
Err(e) => {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
} else {
let body = retry_resp.text().await.unwrap_or_default();
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&format!("Retry failed: {}", body)),
);
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": format!("Retry failed: {}", body)}})),
)
.into_response()
}
}
Err(e) => {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
} else {
let body = resp.text().await.unwrap_or_default();
let _ = state
.pool_service
.mark_unhealthy(db, &credential.uuid, Some(&body));
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": body}})),
)
.into_response()
}
}
CredentialData::GeminiOAuth { .. } => {
// Gemini OAuth 路由暂不支持
(
StatusCode::NOT_IMPLEMENTED,
Json(serde_json::json!({"error": {"message": "Gemini OAuth routing not yet implemented. Use /v1/messages with Gemini models instead."}})),
)
.into_response()
}
CredentialData::QwenOAuth { .. } => {
// Qwen OAuth 路由暂不支持
(
StatusCode::NOT_IMPLEMENTED,
Json(serde_json::json!({"error": {"message": "Qwen OAuth routing not yet implemented. Use /v1/messages with Qwen models instead."}})),
)
.into_response()
}
CredentialData::AntigravityOAuth {
creds_file_path,
project_id,
} => {
let mut antigravity = AntigravityProvider::new();
if let Err(e) = antigravity
.load_credentials_from_path(creds_file_path)
.await
{
// 记录凭证加载失败
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&format!("Failed to load credentials: {}", e)),
);
}
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": format!("Failed to load Antigravity credentials: {}", e)}})),
)
.into_response();
}
// 检查并刷新 token
if antigravity.is_token_expiring_soon() {
if let Err(e) = antigravity.refresh_token().await {
// 记录 Token 刷新失败
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&format!("Token refresh failed: {}", e)),
);
}
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})),
)
.into_response();
}
}
// 设置项目 ID
if let Some(pid) = project_id {
antigravity.project_id = Some(pid.clone());
} else if let Err(e) = antigravity.discover_project().await {
tracing::warn!("[Antigravity] Failed to discover project: {}", e);
}
// 获取 project_id 用于请求
let proj_id = antigravity.project_id.clone().unwrap_or_default();
// 先转换为 OpenAI 格式,再转换为 Antigravity 格式
let openai_request = convert_anthropic_to_openai(request);
let antigravity_request = convert_openai_to_antigravity_with_context(&openai_request, &proj_id);
match antigravity
.generate_content(&request.model, &antigravity_request)
.await
{
Ok(resp) => {
// 转换为 OpenAI 格式,再构建 Anthropic 响应
let content = resp["candidates"][0]["content"]["parts"][0]["text"]
.as_str()
.unwrap_or("");
let parsed = CWParsedResponse {
content: content.to_string(),
tool_calls: Vec::new(),
usage_credits: 0.0,
context_usage_percentage: 0.0,
};
// 记录成功
if let Some(db) = &state.db {
let _ = state.pool_service.mark_healthy(
db,
&credential.uuid,
Some(&request.model),
);
let _ = state.pool_service.record_usage(db, &credential.uuid);
}
if request.stream {
build_anthropic_stream_response(&request.model, &parsed)
} else {
build_anthropic_response(&request.model, &parsed)
}
}
Err(e) => {
// 记录 API 调用失败
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
}
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
}
CredentialData::OpenAIKey { api_key, base_url } => {
let openai = OpenAICustomProvider::with_config(api_key.clone(), base_url.clone());
let openai_request = convert_anthropic_to_openai(request);
match openai.call_api(&openai_request).await {
Ok(resp) => {
if resp.status().is_success() {
match resp.text().await {
Ok(body) => {
if let Ok(openai_resp) =
serde_json::from_str::<serde_json::Value>(&body)
{
let content = openai_resp["choices"][0]["message"]["content"]
.as_str()
.unwrap_or("");
let parsed = CWParsedResponse {
content: content.to_string(),
tool_calls: Vec::new(),
usage_credits: 0.0,
context_usage_percentage: 0.0,
};
// 记录成功
if let Some(db) = &state.db {
let _ = state.pool_service.mark_healthy(
db,
&credential.uuid,
Some(&request.model),
);
let _ =
state.pool_service.record_usage(db, &credential.uuid);
}
if request.stream {
build_anthropic_stream_response(&request.model, &parsed)
} else {
build_anthropic_response(&request.model, &parsed)
}
} else {
// 记录解析失败
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some("Failed to parse OpenAI response"),
);
}
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": "Failed to parse OpenAI response"}})),
)
.into_response()
}
}
Err(e) => {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
}
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
} else {
let body = resp.text().await.unwrap_or_default();
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&body),
);
}
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": body}})),
)
.into_response()
}
}
Err(e) => {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
}
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
}
CredentialData::ClaudeKey { api_key, base_url } => {
// 打印 Claude 代理 URL 用于调试
let actual_base_url = base_url.as_deref().unwrap_or("https://api.anthropic.com");
let claude = ClaudeCustomProvider::with_config(api_key.clone(), base_url.clone());
let request_url = claude.get_base_url();
state.logs.write().await.add(
"info",
&format!(
"[CLAUDE] 使用 Claude API 代理: base_url={} -> {}/v1/messages credential_uuid={}",
actual_base_url,
request_url,
&credential.uuid[..8]
),
);
// 打印请求参数
let request_json = serde_json::to_string(request).unwrap_or_default();
state.logs.write().await.add(
"debug",
&format!(
"[CLAUDE] 请求参数: {}",
&request_json.chars().take(500).collect::<String>()
),
);
match claude.call_api(request).await {
Ok(resp) => {
let status = resp.status();
// 打印响应状态
state.logs.write().await.add(
"info",
&format!(
"[CLAUDE] 响应状态: status={} model={}",
status,
request.model
),
);
match resp.text().await {
Ok(body) => {
if status.is_success() {
// 打印响应内容预览
state.logs.write().await.add(
"debug",
&format!(
"[CLAUDE] 响应内容: {}",
&body.chars().take(500).collect::<String>()
),
);
// 记录成功
if let Some(db) = &state.db {
let _ = state.pool_service.mark_healthy(
db,
&credential.uuid,
Some(&request.model),
);
let _ = state.pool_service.record_usage(db, &credential.uuid);
}
Response::builder()
.status(StatusCode::OK)
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(body))
.unwrap_or_else(|_| {
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": "Failed to build response"}})),
)
.into_response()
})
} else {
state.logs.write().await.add(
"error",
&format!(
"[CLAUDE] 请求失败: status={} body={}",
status,
&body.chars().take(200).collect::<String>()
),
);
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&body),
);
}
(
StatusCode::from_u16(status.as_u16())
.unwrap_or(StatusCode::INTERNAL_SERVER_ERROR),
Json(serde_json::json!({"error": {"message": body}})),
)
.into_response()
}
}
Err(e) => {
state.logs.write().await.add(
"error",
&format!("[CLAUDE] 读取响应失败: {}", e),
);
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
}
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
}
Err(e) => {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
}
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
}
CredentialData::VertexKey { api_key, base_url, .. } => {
// Vertex AI uses Gemini-compatible API, convert Anthropic to OpenAI format first
let openai_request = convert_anthropic_to_openai(request);
let vertex = VertexProvider::with_config(api_key.clone(), base_url.clone());
match vertex.chat_completions(&serde_json::to_value(&openai_request).unwrap_or_default()).await {
Ok(resp) => {
let status = resp.status();
match resp.text().await {
Ok(body) => {
if status.is_success() {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_healthy(db, &credential.uuid, Some(&request.model));
let _ = state.pool_service.record_usage(db, &credential.uuid);
}
Response::builder()
.status(StatusCode::OK)
.header(header::CONTENT_TYPE, "application/json")
.body(Body::from(body))
.unwrap_or_else(|_| {
(StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": "Failed to build response"}}))).into_response()
})
} else {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&body));
}
(StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), Json(serde_json::json!({"error": {"message": body}}))).into_response()
}
}
Err(e) => {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&e.to_string()));
}
(StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}}))).into_response()
}
}
}
Err(e) => {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&e.to_string()));
}
(StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}}))).into_response()
}
}
}
// Gemini API Key credentials - not supported for Anthropic format
CredentialData::GeminiApiKey { .. } => {
(
StatusCode::BAD_REQUEST,
Json(serde_json::json!({"error": {"message": "Gemini API Key credentials do not support Anthropic format"}})),
)
.into_response()
}
// 新增的凭证类型暂不支持 Anthropic 格式
CredentialData::CodexOAuth { .. }
| CredentialData::ClaudeOAuth { .. }
| CredentialData::IFlowOAuth { .. }
| CredentialData::IFlowCookie { .. } => {
(
StatusCode::BAD_REQUEST,
Json(serde_json::json!({"error": {"message": "This credential type does not support Anthropic format yet"}})),
)
.into_response()
}
}
}
/// 根据凭证调用 Provider (OpenAI 格式)
pub async fn call_provider_openai(
state: &AppState,
credential: &ProviderCredential,
request: &ChatCompletionRequest,
) -> Response {
let start_time = std::time::Instant::now();
match &credential.credential {
CredentialData::KiroOAuth { creds_file_path } => {
let mut kiro = KiroProvider::new();
if let Err(e) = kiro.load_credentials_from_path(creds_file_path).await {
// 记录凭证加载失败
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&e.to_string()));
}
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": format!("Failed to load Kiro credentials: {}", e)}})),
)
.into_response();
}
if let Err(e) = kiro.refresh_token().await {
// 记录 Token 刷新失败
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&format!("Token refresh failed: {}", e)));
}
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})),
)
.into_response();
}
match kiro.call_api(request).await {
Ok(resp) => {
let status = resp.status();
if status.is_success() {
// 记录成功
if let Some(db) = &state.db {
let _ = state.pool_service.mark_healthy(db, &credential.uuid, Some(&request.model));
let _ = state.pool_service.record_usage(db, &credential.uuid);
}
match resp.text().await {
Ok(body) => {
let parsed = parse_cw_response(&body);
let has_tool_calls = !parsed.tool_calls.is_empty();
let message = if has_tool_calls {
serde_json::json!({
"role": "assistant",
"content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) },
"tool_calls": parsed.tool_calls.iter().map(|tc| {
serde_json::json!({
"id": tc.id,
"type": "function",
"function": {
"name": tc.function.name,
"arguments": tc.function.arguments
}
})
}).collect::<Vec<_>>()
})
} else {
serde_json::json!({
"role": "assistant",
"content": parsed.content
})
};
Json(serde_json::json!({
"id": format!("chatcmpl-{}", uuid::Uuid::new_v4()),
"object": "chat.completion",
"created": std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
"model": request.model,
"choices": [{
"index": 0,
"message": message,
"finish_reason": if has_tool_calls { "tool_calls" } else { "stop" }
}],
"usage": {
"prompt_tokens": 0,
"completion_tokens": 0,
"total_tokens": 0
}
}))
.into_response()
}
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response(),
}
} else {
// 记录 API 调用失败
let body = resp.text().await.unwrap_or_default();
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&format!("HTTP {}: {}", status, safe_truncate(&body, 100))));
}
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": body}})),
)
.into_response()
}
}
Err(e) => {
// 记录请求错误
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(db, &credential.uuid, Some(&e.to_string()));
}
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response()
}
}
}
CredentialData::GeminiOAuth { .. } => {
(
StatusCode::NOT_IMPLEMENTED,
Json(serde_json::json!({"error": {"message": "Gemini OAuth routing not yet implemented."}})),
)
.into_response()
}
CredentialData::QwenOAuth { .. } => {
(
StatusCode::NOT_IMPLEMENTED,
Json(serde_json::json!({"error": {"message": "Qwen OAuth routing not yet implemented."}})),
)
.into_response()
}
CredentialData::AntigravityOAuth { creds_file_path, project_id } => {
let mut antigravity = AntigravityProvider::new();
if let Err(e) = antigravity.load_credentials_from_path(creds_file_path).await {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": format!("Failed to load Antigravity credentials: {}", e)}})),
)
.into_response();
}
// 检查并刷新 token
if antigravity.is_token_expiring_soon() {
if let Err(e) = antigravity.refresh_token().await {
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})),
)
.into_response();
}
}
// 设置项目 ID
if let Some(pid) = project_id {
antigravity.project_id = Some(pid.clone());
} else if let Err(e) = antigravity.discover_project().await {
tracing::warn!("[Antigravity] Failed to discover project: {}", e);
}
// 获取 project_id 用于请求
let proj_id = antigravity.project_id.clone().unwrap_or_default();
// 转换请求格式
let antigravity_request = convert_openai_to_antigravity_with_context(request, &proj_id);
match antigravity.generate_content(&request.model, &antigravity_request).await {
Ok(resp) => {
let openai_response = convert_antigravity_to_openai_response(&resp, &request.model);
Json(openai_response).into_response()
}
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response(),
}
}
CredentialData::OpenAIKey { api_key, base_url } => {
let openai = OpenAICustomProvider::with_config(api_key.clone(), base_url.clone());
match openai.call_api(request).await {
Ok(resp) => {
if resp.status().is_success() {
match resp.text().await {
Ok(body) => {
if let Ok(json) = serde_json::from_str::<serde_json::Value>(&body) {
Json(json).into_response()
} else {
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": "Invalid JSON response"}})),
)
.into_response()
}
}
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response(),
}
} else {
let body = resp.text().await.unwrap_or_default();
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": body}})),
)
.into_response()
}
}
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response(),
}
}
CredentialData::ClaudeKey { api_key, base_url } => {
// 打印 Claude 代理 URL 用于调试
let actual_base_url = base_url.as_deref().unwrap_or("https://api.anthropic.com");
tracing::info!(
"[CLAUDE] 使用 Claude API 代理: base_url={} credential_uuid={}",
actual_base_url,
&credential.uuid[..8]
);
let claude = ClaudeCustomProvider::with_config(api_key.clone(), base_url.clone());
match claude.call_openai_api(request).await {
Ok(resp) => Json(resp).into_response(),
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response(),
}
}
CredentialData::VertexKey { api_key, base_url, model_aliases } => {
// Resolve model alias if present
let resolved_model = model_aliases.get(&request.model).cloned().unwrap_or_else(|| request.model.clone());
let mut modified_request = request.clone();
modified_request.model = resolved_model;
let vertex = VertexProvider::with_config(api_key.clone(), base_url.clone());
match vertex.chat_completions(&serde_json::to_value(&modified_request).unwrap_or_default()).await {
Ok(resp) => {
if resp.status().is_success() {
match resp.text().await {
Ok(body) => {
if let Ok(json) = serde_json::from_str::<serde_json::Value>(&body) {
Json(json).into_response()
} else {
(StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": "Invalid JSON response"}}))).into_response()
}
}
Err(e) => (StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}}))).into_response(),
}
} else {
let body = resp.text().await.unwrap_or_default();
(StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": body}}))).into_response()
}
}
Err(e) => (StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}}))).into_response(),
}
}
// Gemini API Key credentials - not supported for OpenAI format yet
CredentialData::GeminiApiKey { .. } => {
(
StatusCode::BAD_REQUEST,
Json(serde_json::json!({"error": {"message": "Gemini API Key credentials do not support OpenAI format yet"}})),
)
.into_response()
}
// 新增的凭证类型暂不支持 OpenAI 格式
CredentialData::CodexOAuth { .. }
| CredentialData::ClaudeOAuth { .. }
| CredentialData::IFlowOAuth { .. }
| CredentialData::IFlowCookie { .. } => {
(
StatusCode::BAD_REQUEST,
Json(serde_json::json!({"error": {"message": "This credential type does not support OpenAI format yet"}})),
)
.into_response()
}
}
}
+845
View File
@@ -0,0 +1,845 @@
//! WebSocket 连接处理器
//!
//! 处理 WebSocket 连接的建立、消息收发和 API 请求转发
use axum::{
body::Body,
extract::{
ws::{Message as WsMessage, WebSocket, WebSocketUpgrade},
State,
},
http::HeaderMap,
response::IntoResponse,
};
use futures::{SinkExt, StreamExt as FuturesStreamExt};
use crate::converter::anthropic_to_openai::convert_anthropic_to_openai;
use crate::converter::openai_to_antigravity::{
convert_antigravity_to_openai_response, convert_openai_to_antigravity_with_context,
};
use crate::models::anthropic::AnthropicMessagesRequest;
use crate::models::openai::ChatCompletionRequest;
use crate::models::provider_pool_model::ProviderCredential;
use crate::processor::RequestContext;
use crate::providers::{
AntigravityProvider, ClaudeCustomProvider, KiroProvider, OpenAICustomProvider,
};
use crate::server::AppState;
use crate::server_utils::parse_cw_response;
use crate::websocket::{
WsApiRequest, WsApiResponse, WsEndpoint, WsError, WsMessage as WsProtoMessage,
};
/// WebSocket 升级处理器
pub async fn ws_upgrade_handler(
ws: WebSocketUpgrade,
State(state): State<AppState>,
headers: HeaderMap,
) -> impl IntoResponse {
// 验证 API 密钥
let auth = headers
.get("authorization")
.or_else(|| headers.get("x-api-key"))
.and_then(|v| v.to_str().ok());
let key = match auth {
Some(s) if s.starts_with("Bearer ") => &s[7..],
Some(s) => s,
None => {
return axum::http::Response::builder()
.status(401)
.body(Body::from("No API key provided"))
.unwrap()
.into_response();
}
};
if key != state.api_key {
return axum::http::Response::builder()
.status(401)
.body(Body::from("Invalid API key"))
.unwrap()
.into_response();
}
// 获取客户端信息
let client_info = headers
.get("user-agent")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
ws.on_upgrade(move |socket| handle_websocket(socket, state, client_info))
}
/// 处理 WebSocket 连接
pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: Option<String>) {
let conn_id = uuid::Uuid::new_v4().to_string();
// 注册连接
if let Err(e) = state
.ws_manager
.register(conn_id.clone(), client_info.clone())
{
state.logs.write().await.add(
"error",
&format!("[WS] Failed to register connection: {}", e.message),
);
return;
}
state.logs.write().await.add(
"info",
&format!(
"[WS] New connection: {} (client: {:?})",
&conn_id[..8],
client_info
),
);
let (mut sender, mut receiver) = socket.split();
// 消息处理循环
while let Some(msg) = receiver.next().await {
match msg {
Ok(WsMessage::Text(text)) => {
state.ws_manager.on_message();
state.ws_manager.increment_request_count(&conn_id);
match serde_json::from_str::<WsProtoMessage>(&text) {
Ok(ws_msg) => {
let response = handle_ws_message(&state, &conn_id, ws_msg).await;
if let Some(resp) = response {
let resp_text = serde_json::to_string(&resp).unwrap_or_default();
if sender
.send(WsMessage::Text(resp_text.into()))
.await
.is_err()
{
break;
}
}
}
Err(e) => {
state.ws_manager.on_error();
let error = WsProtoMessage::Error(WsError::invalid_message(format!(
"Failed to parse message: {}",
e
)));
let error_text = serde_json::to_string(&error).unwrap_or_default();
if sender
.send(WsMessage::Text(error_text.into()))
.await
.is_err()
{
break;
}
}
}
}
Ok(WsMessage::Binary(_)) => {
state.ws_manager.on_error();
let error = WsProtoMessage::Error(WsError::invalid_message(
"Binary messages not supported",
));
let error_text = serde_json::to_string(&error).unwrap_or_default();
if sender
.send(WsMessage::Text(error_text.into()))
.await
.is_err()
{
break;
}
}
Ok(WsMessage::Ping(data)) => {
if sender.send(WsMessage::Pong(data)).await.is_err() {
break;
}
}
Ok(WsMessage::Pong(_)) => {
// 收到 pong,连接正常
}
Ok(WsMessage::Close(_)) => {
break;
}
Err(e) => {
state.logs.write().await.add(
"error",
&format!("[WS] Connection {} error: {}", &conn_id[..8], e),
);
break;
}
}
}
// 清理连接
state.ws_manager.unregister(&conn_id);
state.logs.write().await.add(
"info",
&format!("[WS] Connection closed: {}", &conn_id[..8]),
);
}
/// 处理 WebSocket 消息
async fn handle_ws_message(
state: &AppState,
conn_id: &str,
msg: WsProtoMessage,
) -> Option<WsProtoMessage> {
match msg {
WsProtoMessage::Ping { timestamp } => Some(WsProtoMessage::Pong { timestamp }),
WsProtoMessage::Pong { .. } => None,
WsProtoMessage::Request(request) => {
state.logs.write().await.add(
"info",
&format!(
"[WS] Request from {}: id={} endpoint={:?}",
&conn_id[..8],
request.request_id,
request.endpoint
),
);
// 处理 API 请求
let response = handle_ws_api_request(state, &request).await;
Some(response)
}
WsProtoMessage::Response(_)
| WsProtoMessage::StreamChunk(_)
| WsProtoMessage::StreamEnd(_) => Some(WsProtoMessage::Error(WsError::invalid_request(
None,
"Invalid message type from client",
))),
WsProtoMessage::Error(_) => None,
}
}
/// 处理 WebSocket API 请求
async fn handle_ws_api_request(state: &AppState, request: &WsApiRequest) -> WsProtoMessage {
match request.endpoint {
WsEndpoint::Models => {
// 返回模型列表
let models = serde_json::json!({
"object": "list",
"data": [
{"id": "claude-sonnet-4-5", "object": "model", "owned_by": "anthropic"},
{"id": "claude-sonnet-4-5-20250929", "object": "model", "owned_by": "anthropic"},
{"id": "claude-3-7-sonnet-20250219", "object": "model", "owned_by": "anthropic"},
{"id": "gemini-2.5-flash", "object": "model", "owned_by": "google"},
{"id": "gemini-2.5-pro", "object": "model", "owned_by": "google"},
{"id": "qwen3-coder-plus", "object": "model", "owned_by": "alibaba"},
]
});
WsProtoMessage::Response(WsApiResponse {
request_id: request.request_id.clone(),
payload: models,
})
}
WsEndpoint::ChatCompletions => {
// 解析 ChatCompletionRequest
match serde_json::from_value::<ChatCompletionRequest>(request.payload.clone()) {
Ok(chat_request) => {
handle_ws_chat_completions(state, &request.request_id, chat_request).await
}
Err(e) => WsProtoMessage::Error(WsError::invalid_request(
Some(request.request_id.clone()),
format!("Invalid chat completion request: {}", e),
)),
}
}
WsEndpoint::Messages => {
// 解析 AnthropicMessagesRequest
match serde_json::from_value::<AnthropicMessagesRequest>(request.payload.clone()) {
Ok(messages_request) => {
handle_ws_anthropic_messages(state, &request.request_id, messages_request).await
}
Err(e) => WsProtoMessage::Error(WsError::invalid_request(
Some(request.request_id.clone()),
format!("Invalid messages request: {}", e),
)),
}
}
}
}
/// 处理 WebSocket chat completions 请求
async fn handle_ws_chat_completions(
state: &AppState,
request_id: &str,
mut request: ChatCompletionRequest,
) -> WsProtoMessage {
// 创建请求上下文
let mut ctx = RequestContext::new(request.model.clone()).with_stream(request.stream);
// 使用 RequestProcessor 解析模型别名和路由
let _provider = state.processor.resolve_and_route(&mut ctx).await;
// 更新请求中的模型名为解析后的模型
if ctx.resolved_model != ctx.original_model {
request.model = ctx.resolved_model.clone();
}
// 应用参数注入
let injection_enabled = *state.injection_enabled.read().await;
if injection_enabled {
let injector = state.processor.injector.read().await;
let mut payload = serde_json::to_value(&request).unwrap_or_default();
let result = injector.inject(&request.model, &mut payload);
if result.has_injections() {
if let Ok(updated) = serde_json::from_value(payload) {
request = updated;
}
}
}
// 获取默认 provider
let default_provider = state.default_provider.read().await.clone();
// 尝试从凭证池中选择凭证
let credential = match &state.db {
Some(db) => state
.pool_service
.select_credential(db, &default_provider, Some(&request.model))
.ok()
.flatten(),
None => None,
};
// 如果找到凭证,使用它调用 API
if let Some(cred) = credential {
// 简化实现:直接调用 provider 并返回结果
// 实际实现应该复用 call_provider_openai 的逻辑
match call_provider_openai_for_ws(state, &cred, &request).await {
Ok(response) => WsProtoMessage::Response(WsApiResponse {
request_id: request_id.to_string(),
payload: response,
}),
Err(e) => WsProtoMessage::Error(WsError::upstream(Some(request_id.to_string()), e)),
}
} else {
// 回退到 Kiro provider
let kiro = state.kiro.read().await;
match kiro.call_api(&request).await {
Ok(resp) => {
if resp.status().is_success() {
match resp.text().await {
Ok(body) => {
let parsed = parse_cw_response(&body);
let has_tool_calls = !parsed.tool_calls.is_empty();
let message = if has_tool_calls {
serde_json::json!({
"role": "assistant",
"content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) },
"tool_calls": parsed.tool_calls.iter().map(|tc| {
serde_json::json!({
"id": tc.id,
"type": "function",
"function": {
"name": tc.function.name,
"arguments": tc.function.arguments
}
})
}).collect::<Vec<_>>()
})
} else {
serde_json::json!({
"role": "assistant",
"content": parsed.content
})
};
let response = serde_json::json!({
"id": format!("chatcmpl-{}", uuid::Uuid::new_v4()),
"object": "chat.completion",
"created": std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
"model": request.model,
"choices": [{
"index": 0,
"message": message,
"finish_reason": if has_tool_calls { "tool_calls" } else { "stop" }
}],
"usage": {
"prompt_tokens": 0,
"completion_tokens": 0,
"total_tokens": 0
}
});
WsProtoMessage::Response(WsApiResponse {
request_id: request_id.to_string(),
payload: response,
})
}
Err(e) => WsProtoMessage::Error(WsError::internal(
Some(request_id.to_string()),
e.to_string(),
)),
}
} else {
let body = resp.text().await.unwrap_or_default();
WsProtoMessage::Error(WsError::upstream(
Some(request_id.to_string()),
format!("Upstream error: {}", body),
))
}
}
Err(e) => WsProtoMessage::Error(WsError::internal(
Some(request_id.to_string()),
e.to_string(),
)),
}
}
}
/// 处理 WebSocket anthropic messages 请求
async fn handle_ws_anthropic_messages(
state: &AppState,
request_id: &str,
mut request: AnthropicMessagesRequest,
) -> WsProtoMessage {
// 创建请求上下文
let mut ctx = RequestContext::new(request.model.clone()).with_stream(request.stream);
// 使用 RequestProcessor 解析模型别名和路由
let _provider = state.processor.resolve_and_route(&mut ctx).await;
// 更新请求中的模型名为解析后的模型
if ctx.resolved_model != ctx.original_model {
request.model = ctx.resolved_model.clone();
}
// 应用参数注入
let injection_enabled = *state.injection_enabled.read().await;
if injection_enabled {
let injector = state.processor.injector.read().await;
let mut payload = serde_json::to_value(&request).unwrap_or_default();
let result = injector.inject(&request.model, &mut payload);
if result.has_injections() {
if let Ok(updated) = serde_json::from_value(payload) {
request = updated;
}
}
}
// 获取默认 provider
let default_provider = state.default_provider.read().await.clone();
// 尝试从凭证池中选择凭证
let credential = match &state.db {
Some(db) => state
.pool_service
.select_credential(db, &default_provider, Some(&request.model))
.ok()
.flatten(),
None => None,
};
// 如果找到凭证,使用它调用 API
if let Some(cred) = credential {
match call_provider_anthropic_for_ws(state, &cred, &request).await {
Ok(response) => WsProtoMessage::Response(WsApiResponse {
request_id: request_id.to_string(),
payload: response,
}),
Err(e) => WsProtoMessage::Error(WsError::upstream(Some(request_id.to_string()), e)),
}
} else {
// 回退到 Kiro provider
let kiro = state.kiro.read().await;
// 转换为 OpenAI 格式
let openai_request = convert_anthropic_to_openai(&request);
match kiro.call_api(&openai_request).await {
Ok(resp) => {
if resp.status().is_success() {
match resp.text().await {
Ok(body) => {
let parsed = parse_cw_response(&body);
// 转换为 Anthropic 格式响应
let response = serde_json::json!({
"id": format!("msg_{}", uuid::Uuid::new_v4()),
"type": "message",
"role": "assistant",
"content": [{
"type": "text",
"text": parsed.content
}],
"model": request.model,
"stop_reason": "end_turn",
"usage": {
"input_tokens": 0,
"output_tokens": 0
}
});
WsProtoMessage::Response(WsApiResponse {
request_id: request_id.to_string(),
payload: response,
})
}
Err(e) => WsProtoMessage::Error(WsError::internal(
Some(request_id.to_string()),
e.to_string(),
)),
}
} else {
let body = resp.text().await.unwrap_or_default();
WsProtoMessage::Error(WsError::upstream(
Some(request_id.to_string()),
format!("Upstream error: {}", body),
))
}
}
Err(e) => WsProtoMessage::Error(WsError::internal(
Some(request_id.to_string()),
e.to_string(),
)),
}
}
}
/// WebSocket 专用的 OpenAI 格式 Provider 调用
pub async fn call_provider_openai_for_ws(
state: &AppState,
credential: &ProviderCredential,
request: &ChatCompletionRequest,
) -> Result<serde_json::Value, String> {
use crate::models::provider_pool_model::CredentialData;
match &credential.credential {
CredentialData::KiroOAuth { creds_file_path } => {
let mut kiro = KiroProvider::new();
if let Err(e) = kiro.load_credentials_from_path(creds_file_path).await {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&format!("Failed to load credentials: {}", e)),
);
}
return Err(e.to_string());
}
if let Err(e) = kiro.refresh_token().await {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&format!("Token refresh failed: {}", e)),
);
}
return Err(e.to_string());
}
let resp = match kiro.call_api(request).await {
Ok(r) => r,
Err(e) => {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
}
return Err(e.to_string());
}
};
if resp.status().is_success() {
let body = resp.text().await.map_err(|e| e.to_string())?;
let parsed = parse_cw_response(&body);
let has_tool_calls = !parsed.tool_calls.is_empty();
// 记录成功
if let Some(db) = &state.db {
let _ =
state
.pool_service
.mark_healthy(db, &credential.uuid, Some(&request.model));
let _ = state.pool_service.record_usage(db, &credential.uuid);
}
let message = if has_tool_calls {
serde_json::json!({
"role": "assistant",
"content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) },
"tool_calls": parsed.tool_calls.iter().map(|tc| {
serde_json::json!({
"id": tc.id,
"type": "function",
"function": {
"name": tc.function.name,
"arguments": tc.function.arguments
}
})
}).collect::<Vec<_>>()
})
} else {
serde_json::json!({
"role": "assistant",
"content": parsed.content
})
};
Ok(serde_json::json!({
"id": format!("chatcmpl-{}", uuid::Uuid::new_v4()),
"object": "chat.completion",
"created": std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
"model": request.model,
"choices": [{
"index": 0,
"message": message,
"finish_reason": if has_tool_calls { "tool_calls" } else { "stop" }
}],
"usage": {
"prompt_tokens": 0,
"completion_tokens": 0,
"total_tokens": 0
}
}))
} else {
let body = resp.text().await.unwrap_or_default();
if let Some(db) = &state.db {
let _ = state
.pool_service
.mark_unhealthy(db, &credential.uuid, Some(&body));
}
Err(format!("Upstream error: {}", body))
}
}
CredentialData::OpenAIKey { api_key, base_url } => {
let provider = OpenAICustomProvider::with_config(api_key.clone(), base_url.clone());
let resp = match provider.call_api(request).await {
Ok(r) => r,
Err(e) => {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
}
return Err(e.to_string());
}
};
if resp.status().is_success() {
// 记录成功
if let Some(db) = &state.db {
let _ =
state
.pool_service
.mark_healthy(db, &credential.uuid, Some(&request.model));
let _ = state.pool_service.record_usage(db, &credential.uuid);
}
resp.json::<serde_json::Value>()
.await
.map_err(|e| e.to_string())
} else {
let body = resp.text().await.unwrap_or_default();
if let Some(db) = &state.db {
let _ = state
.pool_service
.mark_unhealthy(db, &credential.uuid, Some(&body));
}
Err(format!("Upstream error: {}", body))
}
}
CredentialData::ClaudeKey { api_key, base_url } => {
// 打印 Claude 代理 URL 用于调试
let actual_base_url = base_url.as_deref().unwrap_or("https://api.anthropic.com");
tracing::info!(
"[CLAUDE] 使用 Claude API 代理: base_url={} credential_uuid={}",
actual_base_url,
&credential.uuid[..8]
);
let provider = ClaudeCustomProvider::with_config(api_key.clone(), base_url.clone());
match provider.call_openai_api(request).await {
Ok(result) => {
// 记录成功
if let Some(db) = &state.db {
let _ = state.pool_service.mark_healthy(
db,
&credential.uuid,
Some(&request.model),
);
let _ = state.pool_service.record_usage(db, &credential.uuid);
}
Ok(result)
}
Err(e) => {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
}
Err(e.to_string())
}
}
}
CredentialData::AntigravityOAuth {
creds_file_path,
project_id,
} => {
let mut antigravity = AntigravityProvider::new();
if let Err(e) = antigravity
.load_credentials_from_path(creds_file_path)
.await
{
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&format!("Failed to load credentials: {}", e)),
);
}
return Err(e.to_string());
}
if !antigravity.is_token_valid() {
if let Err(e) = antigravity.refresh_token().await {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&format!("Token refresh failed: {}", e)),
);
}
return Err(e.to_string());
}
}
// 设置项目 ID
if let Some(pid) = project_id {
antigravity.project_id = Some(pid.clone());
}
let proj_id = antigravity.project_id.clone().unwrap_or_default();
let antigravity_request = convert_openai_to_antigravity_with_context(request, &proj_id);
match antigravity
.call_api("generateContent", &antigravity_request)
.await
{
Ok(resp) => {
// 记录成功
if let Some(db) = &state.db {
let _ = state.pool_service.mark_healthy(
db,
&credential.uuid,
Some(&request.model),
);
let _ = state.pool_service.record_usage(db, &credential.uuid);
}
Ok(convert_antigravity_to_openai_response(
&resp,
&request.model,
))
}
Err(e) => {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
}
Err(e.to_string())
}
}
}
// GeminiOAuth 和 QwenOAuth 暂不支持 WebSocket,需要使用 HTTP 端点
_ => Err(
"This credential type is not yet supported via WebSocket. Please use HTTP endpoints."
.to_string(),
),
}
}
/// WebSocket 专用的 Anthropic 格式 Provider 调用
pub async fn call_provider_anthropic_for_ws(
state: &AppState,
credential: &ProviderCredential,
request: &AnthropicMessagesRequest,
) -> Result<serde_json::Value, String> {
use crate::models::provider_pool_model::CredentialData;
match &credential.credential {
CredentialData::ClaudeKey { api_key, base_url } => {
// 打印 Claude 代理 URL 用于调试
let actual_base_url = base_url.as_deref().unwrap_or("https://api.anthropic.com");
tracing::info!(
"[CLAUDE] 使用 Claude API 代理: base_url={} credential_uuid={}",
actual_base_url,
&credential.uuid[..8]
);
let provider = ClaudeCustomProvider::with_config(api_key.clone(), base_url.clone());
let resp = match provider.call_api(request).await {
Ok(r) => r,
Err(e) => {
if let Some(db) = &state.db {
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&e.to_string()),
);
}
return Err(e.to_string());
}
};
if resp.status().is_success() {
// 记录成功
if let Some(db) = &state.db {
let _ =
state
.pool_service
.mark_healthy(db, &credential.uuid, Some(&request.model));
let _ = state.pool_service.record_usage(db, &credential.uuid);
}
resp.json::<serde_json::Value>()
.await
.map_err(|e| e.to_string())
} else {
let body = resp.text().await.unwrap_or_default();
if let Some(db) = &state.db {
let _ = state
.pool_service
.mark_unhealthy(db, &credential.uuid, Some(&body));
}
Err(format!("Upstream error: {}", body))
}
}
_ => {
// 转换为 OpenAI 格式并调用(健康状态更新在 call_provider_openai_for_ws 中处理)
let openai_request = convert_anthropic_to_openai(request);
let result = call_provider_openai_for_ws(state, credential, &openai_request).await?;
// 转换响应为 Anthropic 格式
Ok(serde_json::json!({
"id": format!("msg_{}", uuid::Uuid::new_v4()),
"type": "message",
"role": "assistant",
"content": [{
"type": "text",
"text": result.get("choices")
.and_then(|c| c.get(0))
.and_then(|c| c.get("message"))
.and_then(|m| m.get("content"))
.and_then(|c| c.as_str())
.unwrap_or("")
}],
"model": request.model,
"stop_reason": "end_turn",
"usage": {
"input_tokens": 0,
"output_tokens": 0
}
}))
}
}
}
File diff suppressed because it is too large Load Diff
+656
View File
@@ -0,0 +1,656 @@
//! 服务器工具函数
//!
//! 包含响应解析、字符串处理、响应构建等公共工具函数。
use crate::models::openai::{ContentPart, FunctionCall, MessageContent, ToolCall};
use axum::{
body::Body,
http::{header, StatusCode},
response::{IntoResponse, Response},
Json,
};
use futures::stream;
use std::collections::HashMap;
/// CodeWhisperer 响应解析结果
#[derive(Debug, Default)]
pub struct CWParsedResponse {
pub content: String,
pub tool_calls: Vec<ToolCall>,
pub usage_credits: f64,
pub context_usage_percentage: f64,
}
/// 安全截断字符串到指定字符数,避免 UTF-8 边界问题
pub fn safe_truncate(s: &str, max_chars: usize) -> String {
let chars: Vec<char> = s.chars().collect();
if chars.len() <= max_chars {
s.to_string()
} else {
chars[..max_chars].iter().collect()
}
}
/// 计算 MessageContent 的字符长度
pub fn message_content_len(content: &MessageContent) -> usize {
match content {
MessageContent::Text(s) => s.len(),
MessageContent::Parts(parts) => parts
.iter()
.filter_map(|p| {
if let ContentPart::Text { text } = p {
Some(text.len())
} else {
None
}
})
.sum(),
}
}
/// 解析 CodeWhisperer AWS Event Stream 响应
///
/// AWS Event Stream 是二进制格式,JSON payload 嵌入在二进制头部之间
pub fn parse_cw_response(body: &str) -> CWParsedResponse {
let mut result = CWParsedResponse::default();
// 使用 HashMap 来跟踪多个并发的 tool calls
// key: toolUseId, value: (name, input_accumulated)
let mut tool_map: HashMap<String, (String, String)> = HashMap::new();
// 将字符串转换为字节,因为 AWS Event Stream 包含二进制数据
let bytes = body.as_bytes();
// 搜索所有 JSON 对象的模式
// AWS Event Stream 格式: [binary headers]{"content":"..."}[binary trailer]
let json_patterns: &[&[u8]] = &[
b"{\"content\":",
b"{\"name\":",
b"{\"input\":",
b"{\"stop\":",
b"{\"followupPrompt\":",
b"{\"toolUseId\":",
b"{\"unit\":", // meteringEvent
b"{\"contextUsagePercentage\":", // contextUsageEvent
];
let mut pos = 0;
while pos < bytes.len() {
// 找到下一个 JSON 对象的开始
let mut next_start: Option<usize> = None;
for pattern in json_patterns {
if let Some(idx) = find_subsequence(&bytes[pos..], pattern) {
let abs_pos = pos + idx;
if next_start.is_none_or(|start| abs_pos < start) {
next_start = Some(abs_pos);
}
}
}
let start = match next_start {
Some(s) => s,
None => break,
};
// 从 start 位置提取完整的 JSON 对象
if let Some(json_str) = extract_json_from_bytes(&bytes[start..]) {
if let Ok(value) = serde_json::from_str::<serde_json::Value>(&json_str) {
// 处理 content 事件
if let Some(content) = value.get("content").and_then(|v| v.as_str()) {
// 跳过 followupPrompt
if value.get("followupPrompt").is_none() {
result.content.push_str(content);
}
}
// 处理 tool use 事件 (包含 toolUseId)
else if let Some(tool_use_id) = value.get("toolUseId").and_then(|v| v.as_str()) {
let name = value
.get("name")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let input_chunk = value
.get("input")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let is_stop = value.get("stop").and_then(|v| v.as_bool()).unwrap_or(false);
// 获取或创建 tool entry
let entry = tool_map
.entry(tool_use_id.to_string())
.or_insert_with(|| (String::new(), String::new()));
// 更新 name(如果有)
if !name.is_empty() {
entry.0 = name;
}
// 累积 input
entry.1.push_str(&input_chunk);
// 如果是 stop 事件,完成这个 tool call
if is_stop {
if let Some((name, input)) = tool_map.remove(tool_use_id) {
if !name.is_empty() {
result.tool_calls.push(ToolCall {
id: tool_use_id.to_string(),
call_type: "function".to_string(),
function: FunctionCall {
name,
arguments: input,
},
});
}
}
}
}
// 处理独立的 stop 事件(没有 toolUseId)- 这种情况不应该发生,但以防万一
else if value.get("stop").and_then(|v| v.as_bool()).unwrap_or(false) {
// no-op
}
// 处理 meteringEvent: {"unit":"credit","unitPlural":"credits","usage":0.34}
else if let Some(usage) = value.get("usage").and_then(|v| v.as_f64()) {
result.usage_credits = usage;
}
// 处理 contextUsageEvent: {"contextUsagePercentage":54.36}
else if let Some(ctx_usage) =
value.get("contextUsagePercentage").and_then(|v| v.as_f64())
{
result.context_usage_percentage = ctx_usage;
}
}
pos = start + json_str.len();
} else {
pos = start + 1;
}
}
// 处理未完成的 tool calls(没有收到 stop 事件的)
for (id, (name, input)) in tool_map {
if !name.is_empty() {
result.tool_calls.push(ToolCall {
id,
call_type: "function".to_string(),
function: FunctionCall {
name,
arguments: input,
},
});
}
}
// 解析 bracket 格式的 tool calls: [Called xxx with args: {...}]
parse_bracket_tool_calls(&mut result);
result
}
/// 在字节数组中查找子序列
pub fn find_subsequence(haystack: &[u8], needle: &[u8]) -> Option<usize> {
haystack
.windows(needle.len())
.position(|window| window == needle)
}
/// 从字节数组中提取 JSON 对象字符串
pub fn extract_json_from_bytes(bytes: &[u8]) -> Option<String> {
if bytes.is_empty() || bytes[0] != b'{' {
return None;
}
let mut brace_count = 0;
let mut in_string = false;
let mut escape_next = false;
let mut end_pos = None;
for (i, &b) in bytes.iter().enumerate() {
if escape_next {
escape_next = false;
continue;
}
match b {
b'\\' if in_string => escape_next = true,
b'"' => in_string = !in_string,
b'{' if !in_string => brace_count += 1,
b'}' if !in_string => {
brace_count -= 1;
if brace_count == 0 {
end_pos = Some(i + 1);
break;
}
}
_ => {}
}
}
end_pos.and_then(|end| String::from_utf8(bytes[..end].to_vec()).ok())
}
/// 从字符串中提取完整的 JSON 对象 (保留用于兼容)
#[allow(dead_code)]
pub fn extract_json_object(s: &str) -> Option<&str> {
if !s.starts_with('{') {
return None;
}
let mut brace_count = 0;
let mut in_string = false;
let mut escape_next = false;
for (i, c) in s.char_indices() {
if escape_next {
escape_next = false;
continue;
}
match c {
'\\' if in_string => escape_next = true,
'"' => in_string = !in_string,
'{' if !in_string => brace_count += 1,
'}' if !in_string => {
brace_count -= 1;
if brace_count == 0 {
return Some(&s[..=i]);
}
}
_ => {}
}
}
None
}
/// 解析 bracket 格式的 tool calls
///
/// 格式: [Called xxx with args: {...}]
pub fn parse_bracket_tool_calls(result: &mut CWParsedResponse) {
let re =
regex::Regex::new(r"\[Called\s+(\w+)\s+with\s+args:\s*(\{[^}]*(?:\{[^}]*\}[^}]*)*\})\]")
.ok();
if let Some(re) = re {
let mut to_remove = Vec::new();
for cap in re.captures_iter(&result.content) {
if let (Some(name), Some(args)) = (cap.get(1), cap.get(2)) {
let tool_id = format!(
"call_{}",
&uuid::Uuid::new_v4().to_string().replace('-', "")[..8]
);
result.tool_calls.push(ToolCall {
id: tool_id,
call_type: "function".to_string(),
function: FunctionCall {
name: name.as_str().to_string(),
arguments: args.as_str().to_string(),
},
});
if let Some(full_match) = cap.get(0) {
to_remove.push(full_match.as_str().to_string());
}
}
}
// 从 content 中移除 tool call 文本
for s in to_remove {
result.content = result.content.replace(&s, "");
}
result.content = result.content.trim().to_string();
}
}
/// 构建 Anthropic 非流式响应
pub fn build_anthropic_response(model: &str, parsed: &CWParsedResponse) -> Response {
let has_tool_calls = !parsed.tool_calls.is_empty();
let mut content_array: Vec<serde_json::Value> = Vec::new();
if !parsed.content.is_empty() {
content_array.push(serde_json::json!({
"type": "text",
"text": parsed.content
}));
}
for tc in &parsed.tool_calls {
let input: serde_json::Value =
serde_json::from_str(&tc.function.arguments).unwrap_or(serde_json::json!({}));
content_array.push(serde_json::json!({
"type": "tool_use",
"id": tc.id,
"name": tc.function.name,
"input": input
}));
}
if content_array.is_empty() {
content_array.push(serde_json::json!({"type": "text", "text": ""}));
}
// 估算 output tokens: 基于响应内容长度 (约 4 字符 = 1 token)
let mut output_tokens: u32 = (parsed.content.len() / 4) as u32;
for tc in &parsed.tool_calls {
output_tokens += (tc.function.arguments.len() / 4) as u32;
}
// 从 context_usage_percentage 估算 input tokens
// 假设 100% = 200k tokens (Claude 的上下文窗口)
let input_tokens = ((parsed.context_usage_percentage / 100.0) * 200000.0) as u32;
let response = serde_json::json!({
"id": format!("msg_{}", uuid::Uuid::new_v4()),
"type": "message",
"role": "assistant",
"content": content_array,
"model": model,
"stop_reason": if has_tool_calls { "tool_use" } else { "end_turn" },
"stop_sequence": null,
"usage": {
"input_tokens": input_tokens,
"output_tokens": output_tokens
}
});
Json(response).into_response()
}
/// 构建 Anthropic 流式响应 (SSE)
pub fn build_anthropic_stream_response(model: &str, parsed: &CWParsedResponse) -> Response {
let has_tool_calls = !parsed.tool_calls.is_empty();
let message_id = format!("msg_{}", uuid::Uuid::new_v4());
let model = model.to_string();
let content = parsed.content.clone();
let tool_calls = parsed.tool_calls.clone();
// 估算 output tokens: 基于响应内容长度 (约 4 字符 = 1 token)
let mut output_tokens: u32 = (parsed.content.len() / 4) as u32;
for tc in &parsed.tool_calls {
output_tokens += (tc.function.arguments.len() / 4) as u32;
}
// 从 context_usage_percentage 估算 input tokens
let input_tokens = ((parsed.context_usage_percentage / 100.0) * 200000.0) as u32;
// 构建 SSE 事件流
let mut events: Vec<String> = Vec::new();
// 1. message_start
let message_start = serde_json::json!({
"type": "message_start",
"message": {
"id": message_id,
"type": "message",
"role": "assistant",
"model": model,
"content": [],
"stop_reason": null,
"stop_sequence": null,
"usage": {"input_tokens": input_tokens, "output_tokens": 0}
}
});
events.push(format!("event: message_start\ndata: {message_start}\n\n"));
let mut block_index = 0;
// 2. 文本内容块 - 即使为空也要发送,Claude Code 需要至少一个 content block
// content_block_start
let block_start = serde_json::json!({
"type": "content_block_start",
"index": block_index,
"content_block": {"type": "text", "text": ""}
});
events.push(format!(
"event: content_block_start\ndata: {block_start}\n\n"
));
if !content.is_empty() {
// content_block_delta - 发送完整内容
let block_delta = serde_json::json!({
"type": "content_block_delta",
"index": block_index,
"delta": {"type": "text_delta", "text": content}
});
events.push(format!(
"event: content_block_delta\ndata: {block_delta}\n\n"
));
}
// content_block_stop
let block_stop = serde_json::json!({
"type": "content_block_stop",
"index": block_index
});
events.push(format!("event: content_block_stop\ndata: {block_stop}\n\n"));
block_index += 1;
// 3. Tool use 块
for tc in &tool_calls {
// content_block_start
let block_start = serde_json::json!({
"type": "content_block_start",
"index": block_index,
"content_block": {
"type": "tool_use",
"id": tc.id,
"name": tc.function.name,
"input": {}
}
});
events.push(format!(
"event: content_block_start\ndata: {block_start}\n\n"
));
// content_block_delta - input_json_delta
let partial_json = if tc.function.arguments.is_empty() {
"{}".to_string()
} else {
tc.function.arguments.clone()
};
let block_delta = serde_json::json!({
"type": "content_block_delta",
"index": block_index,
"delta": {
"type": "input_json_delta",
"partial_json": partial_json
}
});
events.push(format!(
"event: content_block_delta\ndata: {block_delta}\n\n"
));
// content_block_stop
let block_stop = serde_json::json!({
"type": "content_block_stop",
"index": block_index
});
events.push(format!("event: content_block_stop\ndata: {block_stop}\n\n"));
block_index += 1;
}
// 4. message_delta
let message_delta = serde_json::json!({
"type": "message_delta",
"delta": {
"stop_reason": if has_tool_calls { "tool_use" } else { "end_turn" },
"stop_sequence": null
},
"usage": {"output_tokens": output_tokens}
});
events.push(format!("event: message_delta\ndata: {message_delta}\n\n"));
// 5. message_stop
let message_stop = serde_json::json!({"type": "message_stop"});
events.push(format!("event: message_stop\ndata: {message_stop}\n\n"));
// 创建 SSE 响应
let body_stream = stream::iter(events.into_iter().map(Ok::<_, std::convert::Infallible>));
let body = Body::from_stream(body_stream);
Response::builder()
.status(StatusCode::OK)
.header(header::CONTENT_TYPE, "text/event-stream")
.header(header::CACHE_CONTROL, "no-cache")
.header(header::CONNECTION, "keep-alive")
.body(body)
.unwrap_or_else(|e| {
tracing::error!("Failed to build SSE response: {}", e);
Response::builder()
.status(StatusCode::INTERNAL_SERVER_ERROR)
.body(Body::empty())
.unwrap_or_default()
})
}
/// 构建 Gemini 原生请求体
///
/// 将用户传入的 Gemini 格式请求转换为 Antigravity 请求格式
pub fn build_gemini_native_request(
request: &serde_json::Value,
model: &str,
project_id: &str,
) -> serde_json::Value {
// 模型名称映射
let actual_model = match model {
"gemini-2.5-computer-use-preview-10-2025" => "rev19-uic3-1p",
"gemini-3-pro-image-preview" => "gemini-3-pro-image",
"gemini-3-pro-preview" => "gemini-3-pro-high",
"gemini-claude-sonnet-4-5" => "claude-sonnet-4-5",
"gemini-claude-sonnet-4-5-thinking" => "claude-sonnet-4-5-thinking",
_ => model,
};
// 是否启用思维链
let enable_thinking = model.ends_with("-thinking")
|| model == "gemini-2.5-pro"
|| model.starts_with("gemini-3-pro-")
|| model == "rev19-uic3-1p"
|| model == "gpt-oss-120b-medium";
// 生成请求 ID 和会话 ID
let request_id = format!("agent-{}", uuid::Uuid::new_v4());
let session_id = {
let uuid = uuid::Uuid::new_v4();
let bytes = uuid.as_bytes();
let n: u64 = u64::from_le_bytes([
bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7],
]) % 9_000_000_000_000_000_000;
format!("-{}", n)
};
// 构建内部请求
let mut inner_request = request.clone();
// 添加会话 ID
inner_request["sessionId"] = serde_json::json!(session_id);
// 确保有 generationConfig
if inner_request.get("generationConfig").is_none() {
inner_request["generationConfig"] = serde_json::json!({
"temperature": 1.0,
"maxOutputTokens": 8096,
"topP": 0.85,
"topK": 50,
"candidateCount": 1,
"stopSequences": [
"<|user|>",
"<|bot|>",
"<|context_request|>",
"<|endoftext|>",
"<|end_of_turn|>"
],
"thinkingConfig": {
"includeThoughts": enable_thinking,
"thinkingBudget": if enable_thinking { 1024 } else { 0 }
}
});
} else {
// 确保有 thinkingConfig
if inner_request["generationConfig"]
.get("thinkingConfig")
.is_none()
{
inner_request["generationConfig"]["thinkingConfig"] = serde_json::json!({
"includeThoughts": enable_thinking,
"thinkingBudget": if enable_thinking { 1024 } else { 0 }
});
}
}
// 删除安全设置(Antigravity 不支持)
if let Some(obj) = inner_request.as_object_mut() {
obj.remove("safetySettings");
}
// 构建完整的 Antigravity 请求体
serde_json::json!({
"project": project_id,
"requestId": request_id,
"request": inner_request,
"model": actual_model,
"userAgent": "antigravity"
})
}
/// 健康检查端点响应
pub async fn health() -> impl IntoResponse {
Json(serde_json::json!({
"status": "healthy",
"version": env!("CARGO_PKG_VERSION")
}))
}
/// 模型列表端点响应
pub async fn models() -> impl IntoResponse {
Json(serde_json::json!({
"object": "list",
"data": [
// Kiro/Claude models
{"id": "claude-sonnet-4-5", "object": "model", "owned_by": "anthropic"},
{"id": "claude-sonnet-4-5-20250929", "object": "model", "owned_by": "anthropic"},
{"id": "claude-3-7-sonnet-20250219", "object": "model", "owned_by": "anthropic"},
{"id": "claude-3-5-sonnet-latest", "object": "model", "owned_by": "anthropic"},
// Gemini models
{"id": "gemini-2.5-flash", "object": "model", "owned_by": "google"},
{"id": "gemini-2.5-flash-lite", "object": "model", "owned_by": "google"},
{"id": "gemini-2.5-pro", "object": "model", "owned_by": "google"},
{"id": "gemini-2.5-pro-preview-06-05", "object": "model", "owned_by": "google"},
{"id": "gemini-3-pro-preview", "object": "model", "owned_by": "google"},
// Qwen models
{"id": "qwen3-coder-plus", "object": "model", "owned_by": "alibaba"},
{"id": "qwen3-coder-flash", "object": "model", "owned_by": "alibaba"}
]
}))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_safe_truncate() {
assert_eq!(safe_truncate("hello", 10), "hello");
assert_eq!(safe_truncate("hello world", 5), "hello");
assert_eq!(safe_truncate("你好世界", 2), "你好");
}
#[test]
fn test_find_subsequence() {
let haystack = b"hello world";
assert_eq!(find_subsequence(haystack, b"world"), Some(6));
assert_eq!(find_subsequence(haystack, b"foo"), None);
}
#[test]
fn test_extract_json_from_bytes() {
let json = b"{\"key\":\"value\"}";
assert_eq!(
extract_json_from_bytes(json),
Some("{\"key\":\"value\"}".to_string())
);
let nested = b"{\"outer\":{\"inner\":\"value\"}}";
assert_eq!(
extract_json_from_bytes(nested),
Some("{\"outer\":{\"inner\":\"value\"}}".to_string())
);
assert_eq!(extract_json_from_bytes(b"not json"), None);
}
}