mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
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:
co-authored by
factory-droid[bot]
parent
2af7093b07
commit
6923f5062a
Generated
+1
-1
@@ -3367,7 +3367,7 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "proxycast"
|
||||
version = "0.15.0"
|
||||
version = "0.15.1"
|
||||
dependencies = [
|
||||
"anyhow",
|
||||
"async-stream",
|
||||
|
||||
@@ -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"
|
||||
@@ -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![],
|
||||
};
|
||||
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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(),
|
||||
}),
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user