mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
feat: 增强 Antigravity Token 管理和 Provider Pool 健康监控机制
主要更新: - 新增 Antigravity Token 验证和自动刷新机制 - 实现 Provider Pool 凭证健康状态监控 - 添加多个新 Gemini 模型支持(gemini-3-pro-image-preview, gemini-3-flash-preview 等) - 优化前端 Agent Chat 和 API Server 页面交互 - 增加凭证错误处理和重新授权提示功能
This commit is contained in:
Generated
+6411
File diff suppressed because it is too large
Load Diff
@@ -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 eac52c1d42823583ca0588a698ae96e4360b6cd26ff9f36c289a3ca06fc1a276 # shrinks to secret_key = "__O__Ew02R17oSvs94e--h"
|
||||
@@ -686,6 +686,11 @@ impl NativeAgentState {
|
||||
self.agent.read().is_some()
|
||||
}
|
||||
|
||||
/// 获取当前 Agent 的 provider 类型
|
||||
pub fn get_provider_type(&self) -> Option<ProviderType> {
|
||||
self.agent.read().as_ref().map(|a| a.provider_type)
|
||||
}
|
||||
|
||||
pub fn reset(&self) {
|
||||
*self.agent.write() = None;
|
||||
}
|
||||
|
||||
@@ -159,10 +159,21 @@ impl OpenAIProtocol {
|
||||
let mut parser = OpenAISSEParser::new();
|
||||
let mut final_usage = None;
|
||||
|
||||
eprintln!("[OpenAIProtocol] 开始处理 SSE 流...");
|
||||
|
||||
while let Some(chunk) = stream.next().await {
|
||||
match chunk {
|
||||
Ok(bytes) => {
|
||||
let text = String::from_utf8_lossy(&bytes);
|
||||
eprintln!(
|
||||
"[OpenAIProtocol] 收到 chunk: {} bytes, 内容: {}",
|
||||
bytes.len(),
|
||||
if text.len() > 200 {
|
||||
format!("{}...", &text[..200])
|
||||
} else {
|
||||
text.to_string()
|
||||
}
|
||||
);
|
||||
buffer.push_str(&text);
|
||||
|
||||
// 处理完整的 SSE 事件(以 \n\n 分隔)
|
||||
@@ -287,6 +298,11 @@ impl Protocol for OpenAIProtocol {
|
||||
|
||||
let url = format!("{}{}", base_url, self.endpoint());
|
||||
|
||||
eprintln!(
|
||||
"[OpenAIProtocol] 发送请求到: {} model={} stream={}",
|
||||
url, model, request.stream
|
||||
);
|
||||
|
||||
let response = client
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bearer {}", api_key))
|
||||
@@ -294,9 +310,13 @@ impl Protocol for OpenAIProtocol {
|
||||
.json(&request)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("请求失败: {}", e))?;
|
||||
.map_err(|e| {
|
||||
eprintln!("[OpenAIProtocol] 请求发送失败: {}", e);
|
||||
format!("请求失败: {}", e)
|
||||
})?;
|
||||
|
||||
let status = response.status();
|
||||
eprintln!("[OpenAIProtocol] 响应状态: {}", status);
|
||||
if !status.is_success() {
|
||||
let body = response.text().await.unwrap_or_default();
|
||||
error!("[OpenAIProtocol] 请求失败: {} - {}", status, body);
|
||||
|
||||
@@ -147,34 +147,65 @@ pub async fn native_agent_chat_stream(
|
||||
session_id: Option<String>,
|
||||
model: Option<String>,
|
||||
images: Option<Vec<ImageInputParam>>,
|
||||
provider: Option<String>,
|
||||
) -> Result<(), String> {
|
||||
tracing::info!(
|
||||
"[NativeAgent] 发送流式消息: message_len={}, model={:?}, event={}, session={:?}",
|
||||
"[NativeAgent] 发送流式消息: message_len={}, model={:?}, provider={:?}, event={}, session={:?}",
|
||||
message.len(),
|
||||
model,
|
||||
provider,
|
||||
event_name,
|
||||
session_id
|
||||
);
|
||||
|
||||
// 如果 Agent 未初始化,自动初始化
|
||||
if !agent_state.is_initialized() {
|
||||
let (port, api_key, running, default_provider) = {
|
||||
let state = app_state.read().await;
|
||||
(
|
||||
state.config.server.port,
|
||||
state.running_api_key.clone(),
|
||||
state.running,
|
||||
state.config.routing.default_provider.clone(),
|
||||
)
|
||||
};
|
||||
// 获取配置信息
|
||||
let (port, api_key, running, default_provider) = {
|
||||
let state = app_state.read().await;
|
||||
(
|
||||
state.config.server.port,
|
||||
state.running_api_key.clone(),
|
||||
state.running,
|
||||
state.config.routing.default_provider.clone(),
|
||||
)
|
||||
};
|
||||
|
||||
if !running {
|
||||
return Err("ProxyCast API Server 未运行".to_string());
|
||||
if !running {
|
||||
return Err("ProxyCast API Server 未运行".to_string());
|
||||
}
|
||||
|
||||
let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?;
|
||||
|
||||
// 使用前端传递的 provider,如果没有则使用默认值
|
||||
let provider_str = provider.unwrap_or(default_provider);
|
||||
let provider_type = ProviderType::from_str(&provider_str);
|
||||
|
||||
tracing::info!(
|
||||
"[NativeAgent] 使用 provider: {:?} (原始值: {})",
|
||||
provider_type,
|
||||
provider_str
|
||||
);
|
||||
|
||||
// 如果 Agent 未初始化,或者 provider 发生变化,重新初始化
|
||||
let need_reinit = if !agent_state.is_initialized() {
|
||||
tracing::info!("[NativeAgent] Agent 未初始化,需要初始化");
|
||||
true
|
||||
} else if let Some(current_provider) = agent_state.get_provider_type() {
|
||||
if current_provider != provider_type {
|
||||
tracing::info!(
|
||||
"[NativeAgent] Provider 发生变化: {:?} -> {:?},需要重新初始化",
|
||||
current_provider,
|
||||
provider_type
|
||||
);
|
||||
true
|
||||
} else {
|
||||
false
|
||||
}
|
||||
} else {
|
||||
true
|
||||
};
|
||||
|
||||
let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?;
|
||||
if need_reinit {
|
||||
let base_url = format!("http://127.0.0.1:{}", port);
|
||||
let provider_type = ProviderType::from_str(&default_provider);
|
||||
agent_state.init(base_url, api_key, provider_type)?;
|
||||
}
|
||||
|
||||
|
||||
@@ -4125,3 +4125,24 @@ mod playwright_tests {
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// 获取单个凭证的健康状态
|
||||
/// Requirements: 4.4
|
||||
#[tauri::command]
|
||||
pub async fn get_credential_health(
|
||||
db: State<'_, DbConnection>,
|
||||
pool_service: State<'_, ProviderPoolServiceState>,
|
||||
uuid: String,
|
||||
) -> Result<Option<crate::services::provider_pool_service::CredentialHealthInfo>, String> {
|
||||
pool_service.0.get_credential_health(&db, &uuid)
|
||||
}
|
||||
|
||||
/// 获取所有凭证的健康状态
|
||||
/// Requirements: 4.4
|
||||
#[tauri::command]
|
||||
pub async fn get_all_credential_health(
|
||||
db: State<'_, DbConnection>,
|
||||
pool_service: State<'_, ProviderPoolServiceState>,
|
||||
) -> Result<Vec<crate::services::provider_pool_service::CredentialHealthInfo>, String> {
|
||||
pool_service.0.get_all_credential_health(&db)
|
||||
}
|
||||
|
||||
@@ -583,13 +583,35 @@ impl Default for RoutingConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
default_provider: default_provider(),
|
||||
rules: Vec::new(),
|
||||
rules: default_routing_rules(),
|
||||
model_aliases: HashMap::new(),
|
||||
exclusions: HashMap::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 默认路由规则
|
||||
///
|
||||
/// 为常见的模型模式提供默认路由:
|
||||
/// - `gemini-*` → Antigravity (Antigravity 支持 Gemini 系列模型)
|
||||
/// - `claude-*` → Kiro (默认使用 Kiro 处理 Claude 模型)
|
||||
fn default_routing_rules() -> Vec<RoutingRuleConfig> {
|
||||
vec![
|
||||
// Gemini 模型路由到 Antigravity
|
||||
RoutingRuleConfig {
|
||||
pattern: "gemini-*".to_string(),
|
||||
provider: "antigravity".to_string(),
|
||||
priority: 10,
|
||||
},
|
||||
// Claude 模型路由到 Kiro
|
||||
RoutingRuleConfig {
|
||||
pattern: "claude-*".to_string(),
|
||||
provider: "kiro".to_string(),
|
||||
priority: 10,
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
/// 路由规则配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
|
||||
pub struct RoutingRuleConfig {
|
||||
@@ -935,7 +957,12 @@ mod unit_tests {
|
||||
fn test_routing_config_default() {
|
||||
let config = RoutingConfig::default();
|
||||
assert_eq!(config.default_provider, "kiro");
|
||||
assert!(config.rules.is_empty());
|
||||
// 默认包含 gemini-* 和 claude-* 的路由规则
|
||||
assert_eq!(config.rules.len(), 2);
|
||||
assert_eq!(config.rules[0].pattern, "gemini-*");
|
||||
assert_eq!(config.rules[0].provider, "antigravity");
|
||||
assert_eq!(config.rules[1].pattern, "claude-*");
|
||||
assert_eq!(config.rules[1].provider, "kiro");
|
||||
assert!(config.model_aliases.is_empty());
|
||||
assert!(config.exclusions.is_empty());
|
||||
}
|
||||
|
||||
@@ -246,6 +246,11 @@ fn is_enable_thinking(model: &str) -> bool {
|
||||
|| model == "gpt-oss-120b-medium"
|
||||
}
|
||||
|
||||
/// 检查是否是图片生成模型
|
||||
fn is_image_generation_model(model: &str) -> bool {
|
||||
model == "gemini-3-pro-image" || model == "gemini-3-pro-image-preview"
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// 主转换函数
|
||||
// ============================================================================
|
||||
@@ -524,6 +529,15 @@ pub fn convert_openai_to_antigravity_with_context(
|
||||
response_modalities: None,
|
||||
};
|
||||
|
||||
// 为图片生成模型设置 response_modalities
|
||||
if is_image_generation_model(actual_model) {
|
||||
generation_config.response_modalities = Some(vec!["TEXT".to_string(), "IMAGE".to_string()]);
|
||||
tracing::info!(
|
||||
"[ANTIGRAVITY] 图片生成模型 {} 已启用 IMAGE 响应模态",
|
||||
actual_model
|
||||
);
|
||||
}
|
||||
|
||||
// 处理 reasoning_effort(思维链配置)
|
||||
if supports_thinking {
|
||||
if let Some(ref effort) = request.reasoning_effort {
|
||||
|
||||
@@ -2166,6 +2166,8 @@ pub fn run() {
|
||||
commands::provider_pool_cmd::start_gemini_oauth_login,
|
||||
commands::provider_pool_cmd::exchange_gemini_code,
|
||||
commands::provider_pool_cmd::get_kiro_credential_fingerprint,
|
||||
commands::provider_pool_cmd::get_credential_health,
|
||||
commands::provider_pool_cmd::get_all_credential_health,
|
||||
// Kiro Builder ID 登录命令
|
||||
commands::provider_pool_cmd::start_kiro_builder_id_login,
|
||||
commands::provider_pool_cmd::poll_kiro_builder_id_auth,
|
||||
|
||||
@@ -307,13 +307,16 @@ impl ProviderCredential {
|
||||
|
||||
// Antigravity 凭证只支持特定的模型
|
||||
if let CredentialData::AntigravityOAuth { .. } = &self.credential {
|
||||
// Antigravity 支持的模型列表
|
||||
// Antigravity 支持的模型列表(与 antigravity.rs 中的 ANTIGRAVITY_MODELS 保持同步)
|
||||
const ANTIGRAVITY_SUPPORTED_MODELS: &[&str] = &[
|
||||
"gemini-3-pro-preview",
|
||||
"gemini-3-pro-image-preview",
|
||||
"gemini-3-flash-preview",
|
||||
"gemini-2.5-flash",
|
||||
"gemini-2.5-computer-use-preview-10-2025",
|
||||
"gemini-claude-sonnet-4-5",
|
||||
"gemini-claude-sonnet-4-5-thinking",
|
||||
"gemini-claude-opus-4-5-thinking",
|
||||
];
|
||||
return ANTIGRAVITY_SUPPORTED_MODELS.contains(&model);
|
||||
}
|
||||
|
||||
@@ -93,9 +93,8 @@ impl RequestProcessor {
|
||||
|
||||
/// 使用默认配置创建请求处理器
|
||||
pub fn with_defaults(pool_service: Arc<ProviderPoolService>) -> Self {
|
||||
use crate::ProviderType;
|
||||
Self {
|
||||
router: Arc::new(RwLock::new(Router::new(ProviderType::Kiro))),
|
||||
router: Arc::new(RwLock::new(Self::create_router_with_defaults())),
|
||||
mapper: Arc::new(RwLock::new(ModelMapper::new())),
|
||||
injector: Arc::new(RwLock::new(Injector::new())),
|
||||
retrier: Arc::new(Retrier::with_defaults()),
|
||||
@@ -109,6 +108,24 @@ impl RequestProcessor {
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建带默认路由规则的路由器
|
||||
fn create_router_with_defaults() -> Router {
|
||||
use crate::router::RoutingRule;
|
||||
use crate::ProviderType;
|
||||
|
||||
let mut router = Router::new(ProviderType::Kiro);
|
||||
|
||||
// 添加默认路由规则:gemini-* → Antigravity
|
||||
router.add_rule(RoutingRule::new("gemini-*", ProviderType::Antigravity, 10));
|
||||
|
||||
// 添加默认路由规则:claude-* → Kiro
|
||||
router.add_rule(RoutingRule::new("claude-*", ProviderType::Kiro, 10));
|
||||
|
||||
tracing::info!("[ROUTER] 初始化默认路由规则: gemini-* → Antigravity, claude-* → Kiro");
|
||||
|
||||
router
|
||||
}
|
||||
|
||||
/// 使用共享的统计和 Token 追踪器创建请求处理器
|
||||
///
|
||||
/// 这允许 RequestProcessor 与 TelemetryState 共享同一个 StatsAggregator 和 TokenTracker,
|
||||
@@ -118,9 +135,8 @@ impl RequestProcessor {
|
||||
stats: Arc<ParkingLotRwLock<StatsAggregator>>,
|
||||
tokens: Arc<ParkingLotRwLock<TokenTracker>>,
|
||||
) -> Self {
|
||||
use crate::ProviderType;
|
||||
Self {
|
||||
router: Arc::new(RwLock::new(Router::new(ProviderType::Kiro))),
|
||||
router: Arc::new(RwLock::new(Self::create_router_with_defaults())),
|
||||
mapper: Arc::new(RwLock::new(ModelMapper::new())),
|
||||
injector: Arc::new(RwLock::new(Injector::new())),
|
||||
retrier: Arc::new(Retrier::with_defaults()),
|
||||
|
||||
@@ -38,6 +38,98 @@ const OAUTH_SCOPES: &[&str] = &[
|
||||
// Token 刷新提前量(秒)
|
||||
const REFRESH_SKEW: i64 = 3000;
|
||||
|
||||
// Token 即将过期的阈值(秒)- 10 分钟
|
||||
const TOKEN_EXPIRING_SOON_THRESHOLD: i64 = 600;
|
||||
|
||||
/// Token 验证结果
|
||||
/// Requirements: 1.1, 1.2, 1.3, 1.4
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum TokenValidationResult {
|
||||
/// Token 有效,包含剩余有效时间(秒)
|
||||
Valid { expires_in_secs: i64 },
|
||||
/// Token 即将过期(少于 10 分钟),需要主动刷新
|
||||
ExpiringSoon { expires_in_secs: i64 },
|
||||
/// Token 已过期
|
||||
Expired,
|
||||
/// Token 无效(缺失、为空或格式错误)
|
||||
Invalid { reason: String },
|
||||
}
|
||||
|
||||
impl TokenValidationResult {
|
||||
/// 是否需要刷新 Token
|
||||
pub fn needs_refresh(&self) -> bool {
|
||||
matches!(
|
||||
self,
|
||||
TokenValidationResult::ExpiringSoon { .. }
|
||||
| TokenValidationResult::Expired
|
||||
| TokenValidationResult::Invalid { .. }
|
||||
)
|
||||
}
|
||||
|
||||
/// 是否可以使用(有效或即将过期但仍可用)
|
||||
pub fn is_usable(&self) -> bool {
|
||||
matches!(
|
||||
self,
|
||||
TokenValidationResult::Valid { .. } | TokenValidationResult::ExpiringSoon { .. }
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/// Token 刷新错误类型
|
||||
/// Requirements: 2.1
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub enum TokenRefreshError {
|
||||
/// OAuth invalid_grant 错误 - 需要用户重新授权
|
||||
InvalidGrant { message: String },
|
||||
/// 网络错误 - 可以重试
|
||||
NetworkError { message: String },
|
||||
/// 服务器错误 (5xx) - 可以重试
|
||||
ServerError { message: String },
|
||||
/// 未知错误
|
||||
Unknown { message: String },
|
||||
}
|
||||
|
||||
impl TokenRefreshError {
|
||||
/// 是否需要用户重新授权
|
||||
pub fn requires_reauth(&self) -> bool {
|
||||
matches!(self, TokenRefreshError::InvalidGrant { .. })
|
||||
}
|
||||
|
||||
/// 是否可以重试
|
||||
pub fn is_retryable(&self) -> bool {
|
||||
matches!(
|
||||
self,
|
||||
TokenRefreshError::NetworkError { .. } | TokenRefreshError::ServerError { .. }
|
||||
)
|
||||
}
|
||||
|
||||
/// 获取用户友好的错误消息
|
||||
pub fn user_message(&self) -> String {
|
||||
match self {
|
||||
TokenRefreshError::InvalidGrant { .. } => {
|
||||
"Antigravity 授权已过期,请重新登录授权".to_string()
|
||||
}
|
||||
TokenRefreshError::NetworkError { message } => {
|
||||
format!("网络连接失败: {}", message)
|
||||
}
|
||||
TokenRefreshError::ServerError { message } => {
|
||||
format!("Google 服务暂时不可用: {}", message)
|
||||
}
|
||||
TokenRefreshError::Unknown { message } => {
|
||||
format!("Token 刷新失败: {}", message)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for TokenRefreshError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
write!(f, "{}", self.user_message())
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for TokenRefreshError {}
|
||||
|
||||
/// Antigravity 支持的模型列表
|
||||
pub const ANTIGRAVITY_MODELS: &[&str] = &[
|
||||
"gemini-3-pro-preview",
|
||||
@@ -330,6 +422,252 @@ impl AntigravityProvider {
|
||||
true
|
||||
}
|
||||
|
||||
/// 验证 Token 状态(支持多种时间格式)
|
||||
/// Requirements: 1.1, 1.2, 1.3, 1.4
|
||||
pub fn validate_token(&self) -> TokenValidationResult {
|
||||
// 检查 access_token 是否存在且非空
|
||||
match &self.credentials.access_token {
|
||||
None => {
|
||||
return TokenValidationResult::Invalid {
|
||||
reason: "access_token 缺失".to_string(),
|
||||
};
|
||||
}
|
||||
Some(token) if token.trim().is_empty() => {
|
||||
return TokenValidationResult::Invalid {
|
||||
reason: "access_token 为空".to_string(),
|
||||
};
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
// 检查是否被禁用
|
||||
if self.credentials.enable == Some(false) {
|
||||
return TokenValidationResult::Invalid {
|
||||
reason: "凭证已被禁用".to_string(),
|
||||
};
|
||||
}
|
||||
|
||||
// 检查 refresh_token 是否存在(用于后续刷新)
|
||||
if self.credentials.refresh_token.is_none() {
|
||||
return TokenValidationResult::Invalid {
|
||||
reason: "refresh_token 缺失,无法刷新".to_string(),
|
||||
};
|
||||
}
|
||||
|
||||
let now = chrono::Utc::now();
|
||||
let now_millis = now.timestamp_millis();
|
||||
|
||||
// 尝试解析过期时间(支持多种格式)
|
||||
let expires_in_secs: Option<i64> = {
|
||||
// 优先检查 RFC3339 格式
|
||||
if let Some(expire_str) = &self.credentials.expire {
|
||||
if let Ok(expires) = chrono::DateTime::parse_from_rfc3339(expire_str) {
|
||||
Some((expires.timestamp_millis() - now_millis) / 1000)
|
||||
} else {
|
||||
// RFC3339 解析失败,尝试其他格式
|
||||
None
|
||||
}
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
.or_else(|| {
|
||||
// 兼容毫秒时间戳格式
|
||||
self.credentials
|
||||
.expiry_date
|
||||
.map(|expiry| (expiry - now_millis) / 1000)
|
||||
})
|
||||
.or_else(|| {
|
||||
// 兼容 timestamp + expires_in 格式
|
||||
match (self.credentials.timestamp, self.credentials.expires_in) {
|
||||
(Some(timestamp), Some(expires_in)) => {
|
||||
let expiry = timestamp + (expires_in * 1000);
|
||||
Some((expiry - now_millis) / 1000)
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
});
|
||||
|
||||
match expires_in_secs {
|
||||
Some(secs) if secs <= 0 => TokenValidationResult::Expired,
|
||||
Some(secs) if secs <= TOKEN_EXPIRING_SOON_THRESHOLD => {
|
||||
TokenValidationResult::ExpiringSoon {
|
||||
expires_in_secs: secs,
|
||||
}
|
||||
}
|
||||
Some(secs) => TokenValidationResult::Valid {
|
||||
expires_in_secs: secs,
|
||||
},
|
||||
None => {
|
||||
// 无法解析过期时间,视为已过期(Requirements: 1.4)
|
||||
TokenValidationResult::Invalid {
|
||||
reason: "无法解析过期时间格式".to_string(),
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 分类 Token 刷新错误
|
||||
/// Requirements: 2.1
|
||||
fn classify_refresh_error(status: u16, body: &str) -> TokenRefreshError {
|
||||
// 检查是否是 invalid_grant 错误
|
||||
if status == 400 && body.contains("invalid_grant") {
|
||||
return TokenRefreshError::InvalidGrant {
|
||||
message: "Refresh token 已失效或被撤销".to_string(),
|
||||
};
|
||||
}
|
||||
|
||||
// 服务器错误 (5xx)
|
||||
if status >= 500 {
|
||||
return TokenRefreshError::ServerError {
|
||||
message: format!("HTTP {}: {}", status, body),
|
||||
};
|
||||
}
|
||||
|
||||
// 其他客户端错误
|
||||
if status >= 400 {
|
||||
return TokenRefreshError::Unknown {
|
||||
message: format!("HTTP {}: {}", status, body),
|
||||
};
|
||||
}
|
||||
|
||||
TokenRefreshError::Unknown {
|
||||
message: body.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// 带重试的 Token 刷新
|
||||
/// Requirements: 2.2, 2.3
|
||||
pub async fn refresh_token_with_retry(
|
||||
&mut self,
|
||||
max_retries: u32,
|
||||
) -> Result<String, TokenRefreshError> {
|
||||
let refresh_token = self
|
||||
.credentials
|
||||
.refresh_token
|
||||
.as_ref()
|
||||
.ok_or_else(|| TokenRefreshError::InvalidGrant {
|
||||
message: "No refresh token available".to_string(),
|
||||
})?
|
||||
.clone();
|
||||
|
||||
let params = [
|
||||
("client_id", OAUTH_CLIENT_ID),
|
||||
("client_secret", OAUTH_CLIENT_SECRET),
|
||||
("refresh_token", refresh_token.as_str()),
|
||||
("grant_type", "refresh_token"),
|
||||
];
|
||||
|
||||
let mut last_error: Option<TokenRefreshError> = None;
|
||||
let mut retry_count = 0;
|
||||
|
||||
while retry_count <= max_retries {
|
||||
if retry_count > 0 {
|
||||
// 指数退避:100ms, 200ms, 400ms, ...
|
||||
let delay_ms = 100 * (1 << (retry_count - 1));
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(delay_ms)).await;
|
||||
tracing::info!(
|
||||
"[Antigravity] Token 刷新重试 {}/{}, 延迟 {}ms",
|
||||
retry_count,
|
||||
max_retries,
|
||||
delay_ms
|
||||
);
|
||||
}
|
||||
|
||||
let result = self
|
||||
.client
|
||||
.post("https://oauth2.googleapis.com/token")
|
||||
.form(¶ms)
|
||||
.send()
|
||||
.await;
|
||||
|
||||
match result {
|
||||
Ok(resp) => {
|
||||
let status = resp.status();
|
||||
if status.is_success() {
|
||||
// 成功,解析响应
|
||||
match resp.json::<serde_json::Value>().await {
|
||||
Ok(data) => {
|
||||
let new_token = data["access_token"].as_str().ok_or_else(|| {
|
||||
TokenRefreshError::Unknown {
|
||||
message: "响应中缺少 access_token".to_string(),
|
||||
}
|
||||
})?;
|
||||
|
||||
self.credentials.access_token = Some(new_token.to_string());
|
||||
|
||||
// 更新过期时间
|
||||
if let Some(expires_in) = data["expires_in"].as_i64() {
|
||||
let now = chrono::Utc::now();
|
||||
let expires_at = now + chrono::Duration::seconds(expires_in);
|
||||
self.credentials.expire = Some(expires_at.to_rfc3339());
|
||||
self.credentials.expiry_date =
|
||||
Some(expires_at.timestamp_millis());
|
||||
self.credentials.expires_in = Some(expires_in);
|
||||
self.credentials.timestamp = Some(now.timestamp_millis());
|
||||
}
|
||||
|
||||
// 更新 refresh_token(如果返回了新的)
|
||||
if let Some(new_refresh) = data["refresh_token"].as_str() {
|
||||
self.credentials.refresh_token = Some(new_refresh.to_string());
|
||||
}
|
||||
|
||||
self.credentials.last_refresh =
|
||||
Some(chrono::Utc::now().to_rfc3339());
|
||||
|
||||
// 保存凭证
|
||||
if let Err(e) = self.save_credentials().await {
|
||||
tracing::warn!("[Antigravity] 保存凭证失败: {}", e);
|
||||
}
|
||||
|
||||
return Ok(new_token.to_string());
|
||||
}
|
||||
Err(e) => {
|
||||
last_error = Some(TokenRefreshError::Unknown {
|
||||
message: format!("解析响应失败: {}", e),
|
||||
});
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// 请求失败
|
||||
let status_code = status.as_u16();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
let error = Self::classify_refresh_error(status_code, &body);
|
||||
|
||||
// invalid_grant 不重试
|
||||
if error.requires_reauth() {
|
||||
return Err(error);
|
||||
}
|
||||
|
||||
// 可重试的错误
|
||||
if error.is_retryable() {
|
||||
last_error = Some(error);
|
||||
retry_count += 1;
|
||||
continue;
|
||||
}
|
||||
|
||||
return Err(error);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
// 网络错误,可重试
|
||||
last_error = Some(TokenRefreshError::NetworkError {
|
||||
message: e.to_string(),
|
||||
});
|
||||
retry_count += 1;
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
retry_count += 1;
|
||||
}
|
||||
|
||||
// 所有重试都失败
|
||||
Err(last_error.unwrap_or_else(|| TokenRefreshError::Unknown {
|
||||
message: "Token 刷新失败,已达到最大重试次数".to_string(),
|
||||
}))
|
||||
}
|
||||
|
||||
pub async fn refresh_token(&mut self) -> Result<String, Box<dyn Error + Send + Sync>> {
|
||||
let refresh_token = self
|
||||
.credentials
|
||||
@@ -1605,20 +1943,31 @@ impl StreamingProvider for AntigravityProvider {
|
||||
&self,
|
||||
request: &ChatCompletionRequest,
|
||||
) -> Result<StreamResponse, ProviderError> {
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] ========== call_api_stream 开始 ==========");
|
||||
|
||||
let token = self
|
||||
.credentials
|
||||
.access_token
|
||||
.as_ref()
|
||||
.ok_or_else(|| ProviderError::AuthenticationError("No access token".to_string()))?;
|
||||
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] Token 长度: {} 字符", token.len());
|
||||
|
||||
let project_id = self.project_id.clone().unwrap_or_else(generate_project_id);
|
||||
let actual_model = alias_to_model_name(&request.model);
|
||||
|
||||
tracing::info!(
|
||||
"[ANTIGRAVITY_STREAM] project_id={}, request.model={}, actual_model={}",
|
||||
project_id,
|
||||
request.model,
|
||||
actual_model
|
||||
);
|
||||
|
||||
// 使用统一的转换函数构建请求体
|
||||
let payload = convert_openai_to_antigravity_with_context(request, &project_id);
|
||||
|
||||
tracing::debug!(
|
||||
"[ANTIGRAVITY_STREAM] 请求体: {}",
|
||||
tracing::info!(
|
||||
"[ANTIGRAVITY_STREAM] 请求体 (完整): {}",
|
||||
serde_json::to_string_pretty(&payload).unwrap_or_default()
|
||||
);
|
||||
|
||||
@@ -1631,8 +1980,15 @@ impl StreamingProvider for AntigravityProvider {
|
||||
base_url
|
||||
);
|
||||
|
||||
eprintln!("[ANTIGRAVITY_STREAM] ========== 发起 HTTP 请求 ==========");
|
||||
eprintln!("[ANTIGRAVITY_STREAM] URL: {}", url);
|
||||
eprintln!("[ANTIGRAVITY_STREAM] Model: {}", actual_model);
|
||||
eprintln!(
|
||||
"[ANTIGRAVITY_STREAM] Token 前20字符: {}...",
|
||||
&token[..20.min(token.len())]
|
||||
);
|
||||
tracing::info!(
|
||||
"[ANTIGRAVITY_STREAM] 发起流式请求: url={} model={}",
|
||||
"[ANTIGRAVITY_STREAM] ========== 发起 HTTP 请求 ==========\n URL: {}\n Model: {}\n Method: POST",
|
||||
url,
|
||||
actual_model
|
||||
);
|
||||
@@ -1640,7 +1996,7 @@ impl StreamingProvider for AntigravityProvider {
|
||||
let result = self
|
||||
.client
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bearer {token}"))
|
||||
.header("Authorization", format!("Bearer {}", token))
|
||||
.header("Content-Type", "application/json")
|
||||
.header("Accept", "text/event-stream")
|
||||
.header("User-Agent", "antigravity/1.11.5 windows/amd64")
|
||||
@@ -1651,13 +2007,15 @@ impl StreamingProvider for AntigravityProvider {
|
||||
match result {
|
||||
Ok(resp) => {
|
||||
let status = resp.status();
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] HTTP 响应状态: {}", status);
|
||||
|
||||
if status.is_success() {
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] 流式响应开始: status={}", status);
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] ✓ 流式响应成功建立,返回流");
|
||||
return Ok(reqwest_stream_to_stream_response(resp));
|
||||
} else {
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
tracing::warn!(
|
||||
"[ANTIGRAVITY_STREAM] 请求失败 ({}): {} - {}",
|
||||
tracing::error!(
|
||||
"[ANTIGRAVITY_STREAM] ✗ 请求失败\n Base URL: {}\n Status: {}\n Body: {}",
|
||||
base_url,
|
||||
status,
|
||||
body
|
||||
@@ -1666,12 +2024,17 @@ impl StreamingProvider for AntigravityProvider {
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::warn!("[ANTIGRAVITY_STREAM] 连接失败 ({}): {}", base_url, e);
|
||||
tracing::error!(
|
||||
"[ANTIGRAVITY_STREAM] ✗ 连接失败\n Base URL: {}\n Error: {}",
|
||||
base_url,
|
||||
e
|
||||
);
|
||||
last_error = Some(ProviderError::from_reqwest_error(&e));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
tracing::error!("[ANTIGRAVITY_STREAM] 所有 base URL 都失败了");
|
||||
Err(last_error.unwrap_or_else(|| {
|
||||
ProviderError::NetworkError("All Antigravity base URLs failed".to_string())
|
||||
}))
|
||||
@@ -1689,3 +2052,317 @@ impl StreamingProvider for AntigravityProvider {
|
||||
StreamFormat::GeminiStream
|
||||
}
|
||||
}
|
||||
|
||||
// ==================== 测试模块 ====================
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use proptest::prelude::*;
|
||||
|
||||
// 辅助函数:检查是否为 Valid 状态
|
||||
fn is_valid(result: &TokenValidationResult) -> bool {
|
||||
matches!(result, TokenValidationResult::Valid { .. })
|
||||
}
|
||||
|
||||
// 辅助函数:检查是否为 ExpiringSoon 状态
|
||||
fn is_expiring_soon(result: &TokenValidationResult) -> bool {
|
||||
matches!(result, TokenValidationResult::ExpiringSoon { .. })
|
||||
}
|
||||
|
||||
// 辅助函数:检查是否为 Expired 状态
|
||||
fn is_expired(result: &TokenValidationResult) -> bool {
|
||||
matches!(result, TokenValidationResult::Expired)
|
||||
}
|
||||
|
||||
// 辅助函数:检查是否为 Invalid 状态
|
||||
fn is_invalid(result: &TokenValidationResult) -> bool {
|
||||
matches!(result, TokenValidationResult::Invalid { .. })
|
||||
}
|
||||
|
||||
// ==================== Property 1: Token 过期时间解析正确性 ====================
|
||||
// Feature: antigravity-token-refresh, Property 1: Token 过期时间解析正确性
|
||||
// Validates: Requirements 1.1, 1.3
|
||||
|
||||
proptest! {
|
||||
#![proptest_config(ProptestConfig::with_cases(100))]
|
||||
|
||||
/// Property 1: 对于任何有效的过期时间(RFC3339 格式),validate_token() 应正确判断状态
|
||||
#[test]
|
||||
fn prop_validate_token_rfc3339_format(
|
||||
expires_in_secs in -3600i64..7200i64, // -1小时到2小时
|
||||
) {
|
||||
let now = chrono::Utc::now();
|
||||
let expires_at = now + chrono::Duration::seconds(expires_in_secs);
|
||||
|
||||
let mut provider = AntigravityProvider::new();
|
||||
provider.credentials.access_token = Some("test_token".to_string());
|
||||
provider.credentials.refresh_token = Some("test_refresh".to_string());
|
||||
provider.credentials.expire = Some(expires_at.to_rfc3339());
|
||||
|
||||
let result = provider.validate_token();
|
||||
|
||||
if expires_in_secs <= 0 {
|
||||
prop_assert!(is_expired(&result), "Expected Expired for expires_in_secs={}", expires_in_secs);
|
||||
} else if expires_in_secs <= TOKEN_EXPIRING_SOON_THRESHOLD {
|
||||
prop_assert!(is_expiring_soon(&result), "Expected ExpiringSoon for expires_in_secs={}", expires_in_secs);
|
||||
} else {
|
||||
prop_assert!(is_valid(&result), "Expected Valid for expires_in_secs={}", expires_in_secs);
|
||||
}
|
||||
}
|
||||
|
||||
/// Property 1: 对于任何有效的过期时间(毫秒时间戳格式),validate_token() 应正确判断状态
|
||||
#[test]
|
||||
fn prop_validate_token_timestamp_format(
|
||||
expires_in_secs in -3600i64..7200i64,
|
||||
) {
|
||||
let now = chrono::Utc::now();
|
||||
let expires_at = now + chrono::Duration::seconds(expires_in_secs);
|
||||
|
||||
let mut provider = AntigravityProvider::new();
|
||||
provider.credentials.access_token = Some("test_token".to_string());
|
||||
provider.credentials.refresh_token = Some("test_refresh".to_string());
|
||||
provider.credentials.expiry_date = Some(expires_at.timestamp_millis());
|
||||
|
||||
let result = provider.validate_token();
|
||||
|
||||
if expires_in_secs <= 0 {
|
||||
prop_assert!(is_expired(&result), "Expected Expired for expires_in_secs={}", expires_in_secs);
|
||||
} else if expires_in_secs <= TOKEN_EXPIRING_SOON_THRESHOLD {
|
||||
prop_assert!(is_expiring_soon(&result), "Expected ExpiringSoon for expires_in_secs={}", expires_in_secs);
|
||||
} else {
|
||||
prop_assert!(is_valid(&result), "Expected Valid for expires_in_secs={}", expires_in_secs);
|
||||
}
|
||||
}
|
||||
|
||||
/// Property 1: 对于任何有效的过期时间(timestamp + expires_in 格式),validate_token() 应正确判断状态
|
||||
#[test]
|
||||
fn prop_validate_token_expires_in_format(
|
||||
expires_in_secs in 1i64..7200i64, // 只测试正数,因为这个格式不支持负数
|
||||
) {
|
||||
let now = chrono::Utc::now();
|
||||
|
||||
let mut provider = AntigravityProvider::new();
|
||||
provider.credentials.access_token = Some("test_token".to_string());
|
||||
provider.credentials.refresh_token = Some("test_refresh".to_string());
|
||||
provider.credentials.timestamp = Some(now.timestamp_millis());
|
||||
provider.credentials.expires_in = Some(expires_in_secs);
|
||||
|
||||
let result = provider.validate_token();
|
||||
|
||||
// 由于时间精度问题,允许 1 秒的误差
|
||||
if expires_in_secs <= 1 {
|
||||
prop_assert!(is_expired(&result) || is_expiring_soon(&result), "Expected Expired or ExpiringSoon for expires_in_secs={}", expires_in_secs);
|
||||
} else if expires_in_secs <= TOKEN_EXPIRING_SOON_THRESHOLD {
|
||||
prop_assert!(is_expiring_soon(&result), "Expected ExpiringSoon for expires_in_secs={}", expires_in_secs);
|
||||
} else {
|
||||
prop_assert!(is_valid(&result), "Expected Valid for expires_in_secs={}", expires_in_secs);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ==================== Property 2: 空 Token 检测 ====================
|
||||
// Feature: antigravity-token-refresh, Property 2: 空 Token 检测
|
||||
// Validates: Requirements 1.2
|
||||
|
||||
proptest! {
|
||||
#![proptest_config(ProptestConfig::with_cases(100))]
|
||||
|
||||
/// Property 2: 对于任何空或仅包含空白字符的 token,validate_token() 应返回 Invalid
|
||||
#[test]
|
||||
fn prop_validate_token_empty_detection(
|
||||
whitespace in "[ \t\n\r]*", // 生成各种空白字符组合
|
||||
) {
|
||||
let mut provider = AntigravityProvider::new();
|
||||
provider.credentials.access_token = Some(whitespace);
|
||||
provider.credentials.refresh_token = Some("test_refresh".to_string());
|
||||
|
||||
let result = provider.validate_token();
|
||||
|
||||
prop_assert!(is_invalid(&result), "Expected Invalid for empty/whitespace token");
|
||||
}
|
||||
}
|
||||
|
||||
/// Property 2: None token 应返回 Invalid
|
||||
#[test]
|
||||
fn test_validate_token_none() {
|
||||
let mut provider = AntigravityProvider::new();
|
||||
provider.credentials.access_token = None;
|
||||
provider.credentials.refresh_token = Some("test_refresh".to_string());
|
||||
|
||||
let result = provider.validate_token();
|
||||
|
||||
assert!(
|
||||
matches!(result, TokenValidationResult::Invalid { reason } if reason.contains("缺失"))
|
||||
);
|
||||
}
|
||||
|
||||
/// Property 2: 缺少 refresh_token 应返回 Invalid
|
||||
#[test]
|
||||
fn test_validate_token_no_refresh_token() {
|
||||
let mut provider = AntigravityProvider::new();
|
||||
provider.credentials.access_token = Some("test_token".to_string());
|
||||
provider.credentials.refresh_token = None;
|
||||
|
||||
let result = provider.validate_token();
|
||||
|
||||
assert!(
|
||||
matches!(result, TokenValidationResult::Invalid { reason } if reason.contains("refresh_token"))
|
||||
);
|
||||
}
|
||||
|
||||
/// Property 2: 禁用的凭证应返回 Invalid
|
||||
#[test]
|
||||
fn test_validate_token_disabled() {
|
||||
let mut provider = AntigravityProvider::new();
|
||||
provider.credentials.access_token = Some("test_token".to_string());
|
||||
provider.credentials.refresh_token = Some("test_refresh".to_string());
|
||||
provider.credentials.enable = Some(false);
|
||||
|
||||
let result = provider.validate_token();
|
||||
|
||||
assert!(
|
||||
matches!(result, TokenValidationResult::Invalid { reason } if reason.contains("禁用"))
|
||||
);
|
||||
}
|
||||
|
||||
// ==================== TokenRefreshError 测试 ====================
|
||||
|
||||
#[test]
|
||||
fn test_classify_refresh_error_invalid_grant() {
|
||||
let error =
|
||||
AntigravityProvider::classify_refresh_error(400, r#"{"error": "invalid_grant"}"#);
|
||||
assert!(matches!(error, TokenRefreshError::InvalidGrant { .. }));
|
||||
assert!(error.requires_reauth());
|
||||
assert!(!error.is_retryable());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_classify_refresh_error_server_error() {
|
||||
let error = AntigravityProvider::classify_refresh_error(500, "Internal Server Error");
|
||||
assert!(matches!(error, TokenRefreshError::ServerError { .. }));
|
||||
assert!(!error.requires_reauth());
|
||||
assert!(error.is_retryable());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_classify_refresh_error_unknown() {
|
||||
let error = AntigravityProvider::classify_refresh_error(403, "Forbidden");
|
||||
assert!(matches!(error, TokenRefreshError::Unknown { .. }));
|
||||
assert!(!error.requires_reauth());
|
||||
assert!(!error.is_retryable());
|
||||
}
|
||||
|
||||
// ==================== TokenValidationResult 方法测试 ====================
|
||||
|
||||
#[test]
|
||||
fn test_token_validation_result_needs_refresh() {
|
||||
assert!(!TokenValidationResult::Valid {
|
||||
expires_in_secs: 3600
|
||||
}
|
||||
.needs_refresh());
|
||||
assert!(TokenValidationResult::ExpiringSoon {
|
||||
expires_in_secs: 300
|
||||
}
|
||||
.needs_refresh());
|
||||
assert!(TokenValidationResult::Expired.needs_refresh());
|
||||
assert!(TokenValidationResult::Invalid {
|
||||
reason: "test".to_string()
|
||||
}
|
||||
.needs_refresh());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_token_validation_result_is_usable() {
|
||||
assert!(TokenValidationResult::Valid {
|
||||
expires_in_secs: 3600
|
||||
}
|
||||
.is_usable());
|
||||
assert!(TokenValidationResult::ExpiringSoon {
|
||||
expires_in_secs: 300
|
||||
}
|
||||
.is_usable());
|
||||
assert!(!TokenValidationResult::Expired.is_usable());
|
||||
assert!(!TokenValidationResult::Invalid {
|
||||
reason: "test".to_string()
|
||||
}
|
||||
.is_usable());
|
||||
}
|
||||
|
||||
// ==================== Property 5: 重试次数限制 ====================
|
||||
// Feature: antigravity-token-refresh, Property 5: 重试次数限制
|
||||
// Validates: Requirements 2.2
|
||||
|
||||
proptest! {
|
||||
#![proptest_config(ProptestConfig::with_cases(100))]
|
||||
|
||||
/// Property 5: 对于任何 HTTP 状态码,错误分类应正确识别可重试错误
|
||||
#[test]
|
||||
fn prop_classify_error_retryable(
|
||||
status in 100u16..600u16,
|
||||
) {
|
||||
let error = AntigravityProvider::classify_refresh_error(status, "test error");
|
||||
|
||||
// 5xx 错误应该是可重试的
|
||||
if status >= 500 {
|
||||
prop_assert!(error.is_retryable(), "5xx errors should be retryable");
|
||||
}
|
||||
|
||||
// 400 + invalid_grant 不应该重试
|
||||
if status == 400 {
|
||||
let invalid_grant_error = AntigravityProvider::classify_refresh_error(400, "invalid_grant");
|
||||
prop_assert!(!invalid_grant_error.is_retryable(), "invalid_grant should not be retryable");
|
||||
prop_assert!(invalid_grant_error.requires_reauth(), "invalid_grant should require reauth");
|
||||
}
|
||||
}
|
||||
|
||||
/// Property 5: 对于任何错误类型,user_message 应返回非空字符串
|
||||
#[test]
|
||||
fn prop_error_user_message_not_empty(
|
||||
status in 100u16..600u16,
|
||||
body in ".*",
|
||||
) {
|
||||
let error = AntigravityProvider::classify_refresh_error(status, &body);
|
||||
let message = error.user_message();
|
||||
prop_assert!(!message.is_empty(), "User message should not be empty");
|
||||
}
|
||||
}
|
||||
|
||||
/// 测试 TokenRefreshError 的 Display 实现
|
||||
#[test]
|
||||
fn test_token_refresh_error_display() {
|
||||
let errors = vec![
|
||||
TokenRefreshError::InvalidGrant {
|
||||
message: "test".to_string(),
|
||||
},
|
||||
TokenRefreshError::NetworkError {
|
||||
message: "test".to_string(),
|
||||
},
|
||||
TokenRefreshError::ServerError {
|
||||
message: "test".to_string(),
|
||||
},
|
||||
TokenRefreshError::Unknown {
|
||||
message: "test".to_string(),
|
||||
},
|
||||
];
|
||||
|
||||
for error in errors {
|
||||
let display = format!("{}", error);
|
||||
assert!(!display.is_empty());
|
||||
}
|
||||
}
|
||||
|
||||
/// 测试缺少 refresh_token 时 refresh_token_with_retry 应返回 InvalidGrant 错误
|
||||
#[tokio::test]
|
||||
async fn test_refresh_token_with_retry_no_refresh_token() {
|
||||
let mut provider = AntigravityProvider::new();
|
||||
provider.credentials.access_token = Some("test_token".to_string());
|
||||
provider.credentials.refresh_token = None;
|
||||
|
||||
let result = provider.refresh_token_with_retry(3).await;
|
||||
|
||||
assert!(result.is_err());
|
||||
let error = result.unwrap_err();
|
||||
assert!(error.requires_reauth());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -622,7 +622,15 @@ pub async fn chat_completions(
|
||||
headers: HeaderMap,
|
||||
Json(mut request): Json<ChatCompletionRequest>,
|
||||
) -> Response {
|
||||
// ========== 详细日志:请求入口 ==========
|
||||
eprintln!("\n========== [CHAT_COMPLETIONS] 收到请求 ==========");
|
||||
eprintln!("[CHAT_COMPLETIONS] URL: /v1/chat/completions");
|
||||
eprintln!("[CHAT_COMPLETIONS] 模型: {}", request.model);
|
||||
eprintln!("[CHAT_COMPLETIONS] 流式: {}", request.stream);
|
||||
eprintln!("[CHAT_COMPLETIONS] 消息数量: {}", request.messages.len());
|
||||
|
||||
if let Err(e) = verify_api_key(&headers, &state.api_key).await {
|
||||
eprintln!("[CHAT_COMPLETIONS] 认证失败!");
|
||||
state
|
||||
.logs
|
||||
.write()
|
||||
@@ -630,9 +638,11 @@ pub async fn chat_completions(
|
||||
.add("warn", "Unauthorized request to /v1/chat/completions");
|
||||
return e.into_response();
|
||||
}
|
||||
eprintln!("[CHAT_COMPLETIONS] 认证成功");
|
||||
|
||||
// 创建请求上下文
|
||||
let mut ctx = RequestContext::new(request.model.clone()).with_stream(request.stream);
|
||||
eprintln!("[CHAT_COMPLETIONS] 请求ID: {}", ctx.request_id);
|
||||
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
@@ -643,11 +653,20 @@ pub async fn chat_completions(
|
||||
);
|
||||
|
||||
// 使用 RequestProcessor 解析模型别名和路由
|
||||
eprintln!("[CHAT_COMPLETIONS] 开始路由解析...");
|
||||
let provider = state.processor.resolve_and_route(&mut ctx).await;
|
||||
eprintln!(
|
||||
"[CHAT_COMPLETIONS] 路由结果: provider={:?}, resolved_model={}",
|
||||
provider, ctx.resolved_model
|
||||
);
|
||||
|
||||
// 更新请求中的模型名为解析后的模型
|
||||
if ctx.resolved_model != ctx.original_model {
|
||||
request.model = ctx.resolved_model.clone();
|
||||
eprintln!(
|
||||
"[CHAT_COMPLETIONS] 模型别名解析: {} -> {}",
|
||||
ctx.original_model, ctx.resolved_model
|
||||
);
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
@@ -681,6 +700,10 @@ pub async fn chat_completions(
|
||||
// 根据客户端类型选择 Provider
|
||||
// **Validates: Requirements 3.1, 3.3, 3.4**
|
||||
let (selected_provider, client_type) = select_provider_for_client(&headers, &state).await;
|
||||
eprintln!(
|
||||
"[CHAT_COMPLETIONS] 客户端类型: {}, 选择的Provider: {}",
|
||||
client_type, selected_provider
|
||||
);
|
||||
|
||||
// 记录客户端检测和 Provider 选择结果
|
||||
state.logs.write().await.add(
|
||||
@@ -701,17 +724,73 @@ pub async fn chat_completions(
|
||||
);
|
||||
|
||||
// 尝试从凭证池中选择凭证
|
||||
// 优先使用路由规则选择的 provider,如果找不到再回退到 selected_provider
|
||||
eprintln!("[CHAT_COMPLETIONS] 开始选择凭证...");
|
||||
let credential = match &state.db {
|
||||
Some(db) => state
|
||||
.pool_service
|
||||
.select_credential(db, &selected_provider, Some(&request.model))
|
||||
.ok()
|
||||
.flatten(),
|
||||
None => None,
|
||||
Some(db) => {
|
||||
// 首先尝试使用路由规则选择的 provider
|
||||
let provider_str = provider.to_string();
|
||||
eprintln!(
|
||||
"[CHAT_COMPLETIONS] 尝试从凭证池选择: provider={}, model={}",
|
||||
provider_str, request.model
|
||||
);
|
||||
let cred = state
|
||||
.pool_service
|
||||
.select_credential(db, &provider_str, Some(&request.model))
|
||||
.ok()
|
||||
.flatten();
|
||||
|
||||
if cred.is_some() {
|
||||
eprintln!("[CHAT_COMPLETIONS] 找到凭证: provider={}", provider_str);
|
||||
} else {
|
||||
eprintln!("[CHAT_COMPLETIONS] 未找到凭证: provider={}", provider_str);
|
||||
}
|
||||
|
||||
// 如果路由规则的 provider 没有找到凭证,回退到 selected_provider
|
||||
if cred.is_none() && provider_str != selected_provider {
|
||||
eprintln!(
|
||||
"[CHAT_COMPLETIONS] 回退到 selected_provider: {}",
|
||||
selected_provider
|
||||
);
|
||||
state.logs.write().await.add(
|
||||
"debug",
|
||||
&format!(
|
||||
"[ROUTE] No credential found for routed provider '{}', trying selected_provider '{}'",
|
||||
provider_str, selected_provider
|
||||
),
|
||||
);
|
||||
let fallback_cred = state
|
||||
.pool_service
|
||||
.select_credential(db, &selected_provider, Some(&request.model))
|
||||
.ok()
|
||||
.flatten();
|
||||
if fallback_cred.is_some() {
|
||||
eprintln!(
|
||||
"[CHAT_COMPLETIONS] 回退凭证找到: provider={}",
|
||||
selected_provider
|
||||
);
|
||||
} else {
|
||||
eprintln!("[CHAT_COMPLETIONS] 回退凭证也未找到!");
|
||||
}
|
||||
fallback_cred
|
||||
} else {
|
||||
cred
|
||||
}
|
||||
}
|
||||
None => {
|
||||
eprintln!("[CHAT_COMPLETIONS] 数据库未初始化!");
|
||||
None
|
||||
}
|
||||
};
|
||||
|
||||
// 如果找到凭证池中的凭证,使用它
|
||||
if let Some(cred) = credential {
|
||||
eprintln!(
|
||||
"[CHAT_COMPLETIONS] 使用凭证: type={}, name={:?}, uuid={}",
|
||||
cred.provider_type,
|
||||
cred.name,
|
||||
&cred.uuid[..8.min(cred.uuid.len())]
|
||||
);
|
||||
state.logs.write().await.add(
|
||||
"info",
|
||||
&format!(
|
||||
@@ -764,7 +843,12 @@ pub async fn chat_completions(
|
||||
}
|
||||
}
|
||||
|
||||
eprintln!("[CHAT_COMPLETIONS] 调用 Provider: {}", cred.provider_type);
|
||||
let response = call_provider_openai(&state, &cred, &request, flow_id.as_deref()).await;
|
||||
eprintln!(
|
||||
"[CHAT_COMPLETIONS] Provider 响应状态: {}",
|
||||
response.status()
|
||||
);
|
||||
|
||||
// 记录请求统计
|
||||
let is_success = response.status().is_success();
|
||||
|
||||
@@ -350,24 +350,53 @@ pub async fn call_provider_anthropic(
|
||||
)
|
||||
.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)),
|
||||
);
|
||||
|
||||
// 使用新的 validate_token() 方法检查 Token 状态
|
||||
let validation_result = antigravity.validate_token();
|
||||
tracing::info!("[Antigravity] Token 验证结果: {:?}", validation_result);
|
||||
|
||||
// 根据验证结果决定是否刷新
|
||||
if validation_result.needs_refresh() {
|
||||
tracing::info!("[Antigravity] Token 需要刷新,开始刷新...");
|
||||
match antigravity.refresh_token_with_retry(3).await {
|
||||
Ok(new_token) => {
|
||||
tracing::info!("[Antigravity] Token 刷新成功,新 token 长度: {}", new_token.len());
|
||||
// 刷新成功,标记为健康
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_healthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
None,
|
||||
);
|
||||
}
|
||||
}
|
||||
Err(refresh_error) => {
|
||||
tracing::error!("[Antigravity] Token 刷新失败: {:?}", refresh_error);
|
||||
// 使用新的 mark_unhealthy_with_details 方法
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_unhealthy_with_details(
|
||||
db,
|
||||
&credential.uuid,
|
||||
&refresh_error,
|
||||
);
|
||||
}
|
||||
|
||||
// 根据错误类型返回不同的状态码和消息
|
||||
let (status, message) = if refresh_error.requires_reauth() {
|
||||
(StatusCode::UNAUTHORIZED, refresh_error.user_message())
|
||||
} else {
|
||||
(StatusCode::INTERNAL_SERVER_ERROR, refresh_error.user_message())
|
||||
};
|
||||
|
||||
return (
|
||||
status,
|
||||
Json(serde_json::json!({"error": {"message": message}})),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
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());
|
||||
@@ -1071,30 +1100,321 @@ pub async fn call_provider_openai(
|
||||
.into_response()
|
||||
}
|
||||
CredentialData::AntigravityOAuth { creds_file_path, project_id } => {
|
||||
eprintln!("\n========== [ANTIGRAVITY] 开始处理 Antigravity 请求 ==========");
|
||||
eprintln!("[ANTIGRAVITY] 凭证文件: {}", creds_file_path);
|
||||
eprintln!("[ANTIGRAVITY] 项目ID: {:?}", project_id);
|
||||
eprintln!("[ANTIGRAVITY] 模型: {}", request.model);
|
||||
eprintln!("[ANTIGRAVITY] 流式: {}", request.stream);
|
||||
|
||||
let mut antigravity = AntigravityProvider::new();
|
||||
if let Err(e) = antigravity.load_credentials_from_path(creds_file_path).await {
|
||||
eprintln!("[ANTIGRAVITY] 加载凭证失败: {}", e);
|
||||
// 记录凭证加载失败
|
||||
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 {
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})),
|
||||
)
|
||||
.into_response();
|
||||
eprintln!("[ANTIGRAVITY] 凭证加载成功");
|
||||
|
||||
// 使用新的 validate_token() 方法检查 Token 状态
|
||||
let validation_result = antigravity.validate_token();
|
||||
eprintln!("[ANTIGRAVITY] Token 验证结果: {:?}", validation_result);
|
||||
eprintln!("[ANTIGRAVITY] needs_refresh() = {}", validation_result.needs_refresh());
|
||||
tracing::info!("[Antigravity] Token 验证结果: {:?}", validation_result);
|
||||
|
||||
// 根据验证结果决定是否刷新
|
||||
if validation_result.needs_refresh() {
|
||||
eprintln!("[ANTIGRAVITY] Token 需要刷新,开始刷新...");
|
||||
tracing::info!("[Antigravity] Token 需要刷新,开始刷新...");
|
||||
match antigravity.refresh_token_with_retry(3).await {
|
||||
Ok(new_token) => {
|
||||
eprintln!("[ANTIGRAVITY] Token 刷新成功,新 token 长度: {}", new_token.len());
|
||||
tracing::info!("[Antigravity] Token 刷新成功,新 token 长度: {}", new_token.len());
|
||||
// 刷新成功,标记为健康
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_healthy(
|
||||
db,
|
||||
&credential.uuid,
|
||||
None,
|
||||
);
|
||||
}
|
||||
}
|
||||
Err(refresh_error) => {
|
||||
eprintln!("[ANTIGRAVITY] Token 刷新失败: {:?}", refresh_error);
|
||||
tracing::error!("[Antigravity] Token 刷新失败: {:?}", refresh_error);
|
||||
// 使用新的 mark_unhealthy_with_details 方法
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_unhealthy_with_details(
|
||||
db,
|
||||
&credential.uuid,
|
||||
&refresh_error,
|
||||
);
|
||||
}
|
||||
|
||||
// 根据错误类型返回不同的状态码和消息
|
||||
let (status, message) = if refresh_error.requires_reauth() {
|
||||
(StatusCode::UNAUTHORIZED, refresh_error.user_message())
|
||||
} else {
|
||||
(StatusCode::INTERNAL_SERVER_ERROR, refresh_error.user_message())
|
||||
};
|
||||
|
||||
return (
|
||||
status,
|
||||
Json(serde_json::json!({"error": {"message": message}})),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
}
|
||||
} else {
|
||||
eprintln!("[ANTIGRAVITY] Token 不需要刷新,继续使用现有 Token");
|
||||
}
|
||||
|
||||
// 设置项目 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);
|
||||
}
|
||||
|
||||
tracing::info!("[ANTIGRAVITY] request.stream = {}, model = {}, project_id = {:?}",
|
||||
request.stream, request.model, antigravity.project_id);
|
||||
|
||||
// 检查是否为流式请求
|
||||
if request.stream {
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] ========== 开始处理流式请求 ==========");
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] model={}, has_token={}",
|
||||
request.model, antigravity.credentials.access_token.is_some());
|
||||
|
||||
// 检查是否是图片生成模型
|
||||
// 注意:gemini-3-pro-image-preview 是支持图片理解的模型,不是图片生成模型
|
||||
// 只有明确的图片生成模型才需要走非流式路径
|
||||
let is_image_generation_model = request.model == "imagen"
|
||||
|| request.model.starts_with("imagen-")
|
||||
|| request.model.contains("image-generation");
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] is_image_generation_model={}", is_image_generation_model);
|
||||
|
||||
// 对于图片生成模型,使用非流式请求然后模拟流式返回
|
||||
if is_image_generation_model {
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] 图片生成模型,使用非流式请求");
|
||||
|
||||
// 获取 project_id 用于请求
|
||||
let proj_id = antigravity.project_id.clone().unwrap_or_default();
|
||||
// 转换请求格式 - 这已经是完整的 Antigravity 请求格式
|
||||
let antigravity_request = convert_openai_to_antigravity_with_context(request, &proj_id);
|
||||
|
||||
// 直接调用 call_api,因为 antigravity_request 已经是完整格式
|
||||
match antigravity.call_api("generateContent", &antigravity_request).await {
|
||||
Ok(resp) => {
|
||||
// 保存原始响应到文件用于调试
|
||||
let resp_str = serde_json::to_string_pretty(&resp).unwrap_or_default();
|
||||
let debug_dir = dirs::home_dir()
|
||||
.map(|h| h.join(".proxycast/logs"))
|
||||
.unwrap_or_else(|| std::path::PathBuf::from("/tmp"));
|
||||
let _ = std::fs::create_dir_all(&debug_dir);
|
||||
let debug_file = debug_dir.join("antigravity_image_response.json");
|
||||
let _ = std::fs::write(&debug_file, &resp_str);
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] 原始响应已保存到: {:?}, 大小: {} bytes", debug_file, resp_str.len());
|
||||
eprintln!("[ANTIGRAVITY_STREAM] 原始响应已保存到: {:?}, 大小: {} bytes", debug_file, resp_str.len());
|
||||
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] 图片生成完成,转换为流式响应");
|
||||
|
||||
// 将非流式响应转换为 OpenAI 格式
|
||||
let openai_response = convert_antigravity_to_openai_response(&resp, &request.model);
|
||||
|
||||
// 保存转换后的响应到文件
|
||||
let openai_str = serde_json::to_string_pretty(&openai_response).unwrap_or_default();
|
||||
let openai_debug_file = debug_dir.join("antigravity_image_openai_response.json");
|
||||
let _ = std::fs::write(&openai_debug_file, &openai_str);
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] OpenAI 响应已保存到: {:?}, 大小: {} bytes", openai_debug_file, openai_str.len());
|
||||
eprintln!("[ANTIGRAVITY_STREAM] OpenAI 响应已保存到: {:?}, 大小: {} bytes", openai_debug_file, openai_str.len());
|
||||
|
||||
// 将非流式响应转换为流式 SSE 格式
|
||||
let model = request.model.clone();
|
||||
let chunk_id = format!("chatcmpl-{}", uuid::Uuid::new_v4());
|
||||
let created = chrono::Utc::now().timestamp();
|
||||
|
||||
// 提取内容
|
||||
let content = openai_response
|
||||
.get("choices")
|
||||
.and_then(|c| c.as_array())
|
||||
.and_then(|arr| arr.first())
|
||||
.and_then(|choice| choice.get("message"))
|
||||
.and_then(|msg| msg.get("content"))
|
||||
.and_then(|c| c.as_str())
|
||||
.unwrap_or("");
|
||||
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] 图片内容长度: {} 字符", content.len());
|
||||
eprintln!("[ANTIGRAVITY_STREAM] 图片内容长度: {} 字符", content.len());
|
||||
|
||||
// 构建 SSE 事件
|
||||
let mut sse_events = String::new();
|
||||
|
||||
// 发送内容 chunk
|
||||
if !content.is_empty() {
|
||||
let chunk_response = serde_json::json!({
|
||||
"id": chunk_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"content": content
|
||||
},
|
||||
"finish_reason": null
|
||||
}]
|
||||
});
|
||||
sse_events.push_str(&format!("data: {}\n\n", chunk_response.to_string()));
|
||||
}
|
||||
|
||||
// 发送结束 chunk
|
||||
let done_response = serde_json::json!({
|
||||
"id": chunk_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"delta": {},
|
||||
"finish_reason": "stop"
|
||||
}]
|
||||
});
|
||||
sse_events.push_str(&format!("data: {}\n\n", done_response.to_string()));
|
||||
sse_events.push_str("data: [DONE]\n\n");
|
||||
|
||||
return 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::from(sse_events))
|
||||
.unwrap_or_else(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": "Failed to build streaming response"}})),
|
||||
)
|
||||
.into_response()
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!("[ANTIGRAVITY_STREAM] 图片生成失败: {}", e);
|
||||
return (
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": e.to_string()}})),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
match antigravity.call_api_stream(request).await {
|
||||
Ok(stream_response) => {
|
||||
eprintln!("[ANTIGRAVITY_STREAM] ✓ 流式响应已建立");
|
||||
tracing::info!("[ANTIGRAVITY_STREAM] ✓ 流式响应已建立");
|
||||
|
||||
let model = request.model.clone();
|
||||
|
||||
// Antigravity 返回的是分片的 JSON,需要累积所有数据后解析
|
||||
// 使用 channel 来收集所有数据,然后一次性返回
|
||||
let (tx, rx) = tokio::sync::oneshot::channel::<Result<String, String>>();
|
||||
|
||||
// 在后台任务中收集所有数据
|
||||
let model_clone = model.clone();
|
||||
tokio::spawn(async move {
|
||||
use futures::StreamExt;
|
||||
let mut stream = stream_response;
|
||||
let mut all_data = String::new();
|
||||
let mut chunk_count = 0u32;
|
||||
|
||||
while let Some(result) = stream.next().await {
|
||||
chunk_count += 1;
|
||||
match result {
|
||||
Ok(bytes) => {
|
||||
let text = String::from_utf8_lossy(&bytes);
|
||||
all_data.push_str(&text);
|
||||
|
||||
if chunk_count <= 3 {
|
||||
eprintln!("[ANTIGRAVITY_STREAM] 收集 chunk #{}: {} bytes", chunk_count, bytes.len());
|
||||
} else if chunk_count % 200 == 0 {
|
||||
eprintln!("[ANTIGRAVITY_STREAM] 已收集 {} 个 chunk, 总大小: {} bytes", chunk_count, all_data.len());
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
eprintln!("[ANTIGRAVITY_STREAM] chunk #{} 错误: {}", chunk_count, e);
|
||||
let _ = tx.send(Err(e.to_string()));
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
eprintln!("[ANTIGRAVITY_STREAM] 流结束,共收集 {} 个 chunk, 总大小: {} bytes", chunk_count, all_data.len());
|
||||
|
||||
// 尝试解析累积的 JSON 数据
|
||||
// Antigravity 返回格式: { "response": { "candidates": [...] } }
|
||||
let result = parse_antigravity_accumulated_response(&all_data, &model_clone);
|
||||
let _ = tx.send(result);
|
||||
});
|
||||
|
||||
// 等待数据收集完成,然后构建 SSE 响应
|
||||
let sse_stream = async_stream::stream! {
|
||||
match rx.await {
|
||||
Ok(Ok(sse_content)) => {
|
||||
// 返回累积的 SSE 事件
|
||||
yield Ok::<_, std::io::Error>(axum::body::Bytes::from(sse_content));
|
||||
}
|
||||
Ok(Err(e)) => {
|
||||
eprintln!("[ANTIGRAVITY_STREAM] 解析错误: {}", e);
|
||||
let error_event = format!(
|
||||
"data: {{\"error\": {{\"message\": \"{}\"}}}}\n\ndata: [DONE]\n\n",
|
||||
e.replace("\"", "\\\"")
|
||||
);
|
||||
yield Ok(axum::body::Bytes::from(error_event));
|
||||
}
|
||||
Err(_) => {
|
||||
eprintln!("[ANTIGRAVITY_STREAM] channel 接收错误");
|
||||
let error_event = "data: {\"error\": {\"message\": \"Internal error\"}}\n\ndata: [DONE]\n\n";
|
||||
yield Ok(axum::body::Bytes::from(error_event.to_string()));
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
return Response::builder()
|
||||
.status(StatusCode::OK)
|
||||
.header(header::CONTENT_TYPE, "text/event-stream")
|
||||
.header(header::CACHE_CONTROL, "no-cache")
|
||||
.header(header::CONNECTION, "keep-alive")
|
||||
.header("X-Accel-Buffering", "no")
|
||||
.body(Body::from_stream(sse_stream))
|
||||
.unwrap_or_else(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(
|
||||
serde_json::json!({"error": {"message": "Failed to build streaming response"}}),
|
||||
),
|
||||
)
|
||||
.into_response()
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
return (
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
Json(serde_json::json!({"error": {"message": e.to_string()}})),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 非流式请求处理
|
||||
// 获取 project_id 用于请求
|
||||
let proj_id = antigravity.project_id.clone().unwrap_or_default();
|
||||
// 转换请求格式
|
||||
@@ -2263,3 +2583,356 @@ pub async fn handle_kiro_stream(
|
||||
.into_response()
|
||||
})
|
||||
}
|
||||
|
||||
/// 解析 Antigravity 累积的流式响应数据
|
||||
///
|
||||
/// Antigravity 返回的流式数据是分片的 JSON,格式如下:
|
||||
/// ```json
|
||||
/// {
|
||||
/// "response": {
|
||||
/// "candidates": [{
|
||||
/// "content": {
|
||||
/// "role": "model",
|
||||
/// "parts": [
|
||||
/// { "text": "..." },
|
||||
/// { "inlineData": { "mimeType": "image/jpeg", "data": "base64..." } }
|
||||
/// ]
|
||||
/// }
|
||||
/// }]
|
||||
/// }
|
||||
/// }
|
||||
/// ```
|
||||
fn parse_antigravity_accumulated_response(data: &str, model: &str) -> Result<String, String> {
|
||||
eprintln!(
|
||||
"[ANTIGRAVITY_PARSE] 开始解析累积数据,大小: {} bytes",
|
||||
data.len()
|
||||
);
|
||||
|
||||
// 保存原始数据到文件用于调试
|
||||
let debug_dir = dirs::home_dir()
|
||||
.map(|h| h.join(".proxycast/logs"))
|
||||
.unwrap_or_else(|| std::path::PathBuf::from("/tmp"));
|
||||
let _ = std::fs::create_dir_all(&debug_dir);
|
||||
let debug_file = debug_dir.join("antigravity_stream_raw.txt");
|
||||
let _ = std::fs::write(&debug_file, data);
|
||||
eprintln!("[ANTIGRAVITY_PARSE] 原始数据已保存到: {:?}", debug_file);
|
||||
|
||||
// 打印数据的前1000字符用于调试
|
||||
eprintln!(
|
||||
"[ANTIGRAVITY_PARSE] 数据前1000字符:\n{}",
|
||||
&data[..data.len().min(1000)]
|
||||
);
|
||||
|
||||
// 尝试解析 JSON
|
||||
// Antigravity 流式响应可能是多个 JSON 对象,每个对象一行
|
||||
// 或者是一个大的 JSON 对象
|
||||
|
||||
// 首先尝试直接解析为单个 JSON
|
||||
if let Ok(json) = serde_json::from_str::<serde_json::Value>(data) {
|
||||
eprintln!("[ANTIGRAVITY_PARSE] 单个 JSON 解析成功");
|
||||
return parse_antigravity_json(&json, model);
|
||||
}
|
||||
|
||||
// 如果失败,尝试按行解析,找到包含 candidates 的 JSON
|
||||
eprintln!("[ANTIGRAVITY_PARSE] 单个 JSON 解析失败,尝试按行解析");
|
||||
|
||||
let mut all_text = String::new();
|
||||
let mut all_images: Vec<(String, String)> = Vec::new(); // (mime_type, data)
|
||||
let mut found_any = false;
|
||||
|
||||
for line in data.lines() {
|
||||
let line = line.trim();
|
||||
if line.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
// 尝试解析每一行
|
||||
if let Ok(json) = serde_json::from_str::<serde_json::Value>(line) {
|
||||
if let Some((text, images)) = extract_content_from_json(&json) {
|
||||
all_text.push_str(&text);
|
||||
all_images.extend(images);
|
||||
found_any = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if found_any {
|
||||
eprintln!(
|
||||
"[ANTIGRAVITY_PARSE] 按行解析成功,文本长度: {}, 图片数: {}",
|
||||
all_text.len(),
|
||||
all_images.len()
|
||||
);
|
||||
return build_sse_response(&all_text, &all_images, model);
|
||||
}
|
||||
|
||||
// 如果还是失败,尝试找到 JSON 对象的边界
|
||||
eprintln!("[ANTIGRAVITY_PARSE] 按行解析失败,尝试查找 JSON 边界");
|
||||
|
||||
// 查找所有 { 开头的位置,尝试解析
|
||||
let mut start = 0;
|
||||
while let Some(pos) = data[start..].find('{') {
|
||||
let json_start = start + pos;
|
||||
// 尝试从这个位置解析 JSON
|
||||
if let Ok(json) = serde_json::from_str::<serde_json::Value>(&data[json_start..]) {
|
||||
eprintln!("[ANTIGRAVITY_PARSE] 在位置 {} 找到有效 JSON", json_start);
|
||||
return parse_antigravity_json(&json, model);
|
||||
}
|
||||
start = json_start + 1;
|
||||
if start >= data.len() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
Err(format!("无法解析响应数据,请查看 {:?}", debug_file))
|
||||
}
|
||||
|
||||
/// 从 JSON 中提取内容
|
||||
fn extract_content_from_json(json: &serde_json::Value) -> Option<(String, Vec<(String, String)>)> {
|
||||
// 尝试多种路径
|
||||
let candidates = json
|
||||
.get("response")
|
||||
.and_then(|r| r.get("candidates"))
|
||||
.or_else(|| json.get("candidates"))
|
||||
.and_then(|c| c.as_array())?;
|
||||
|
||||
if candidates.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let mut text = String::new();
|
||||
let mut images = Vec::new();
|
||||
|
||||
for candidate in candidates {
|
||||
if let Some(parts) = candidate
|
||||
.get("content")
|
||||
.and_then(|c| c.get("parts"))
|
||||
.and_then(|p| p.as_array())
|
||||
{
|
||||
for part in parts {
|
||||
if let Some(t) = part.get("text").and_then(|t| t.as_str()) {
|
||||
text.push_str(t);
|
||||
}
|
||||
if let Some(inline_data) =
|
||||
part.get("inlineData").or_else(|| part.get("inline_data"))
|
||||
{
|
||||
if let Some(data) = inline_data.get("data").and_then(|d| d.as_str()) {
|
||||
let mime = inline_data
|
||||
.get("mimeType")
|
||||
.or_else(|| inline_data.get("mime_type"))
|
||||
.and_then(|m| m.as_str())
|
||||
.unwrap_or("image/png");
|
||||
images.push((mime.to_string(), data.to_string()));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if text.is_empty() && images.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some((text, images))
|
||||
}
|
||||
}
|
||||
|
||||
/// 解析 Antigravity JSON 响应
|
||||
fn parse_antigravity_json(json: &serde_json::Value, model: &str) -> Result<String, String> {
|
||||
eprintln!(
|
||||
"[ANTIGRAVITY_PARSE] 解析 JSON,顶层类型: {}",
|
||||
if json.is_object() {
|
||||
"object"
|
||||
} else if json.is_array() {
|
||||
"array"
|
||||
} else {
|
||||
"other"
|
||||
}
|
||||
);
|
||||
|
||||
if let Some(obj) = json.as_object() {
|
||||
eprintln!(
|
||||
"[ANTIGRAVITY_PARSE] 顶层 keys: {:?}",
|
||||
obj.keys().collect::<Vec<_>>()
|
||||
);
|
||||
}
|
||||
|
||||
if let Some((text, images)) = extract_content_from_json(json) {
|
||||
return build_sse_response(&text, &images, model);
|
||||
}
|
||||
|
||||
// 如果是数组,尝试处理每个元素
|
||||
if let Some(arr) = json.as_array() {
|
||||
eprintln!("[ANTIGRAVITY_PARSE] 顶层是数组,长度: {}", arr.len());
|
||||
let mut all_text = String::new();
|
||||
let mut all_images = Vec::new();
|
||||
|
||||
for item in arr {
|
||||
if let Some((text, images)) = extract_content_from_json(item) {
|
||||
all_text.push_str(&text);
|
||||
all_images.extend(images);
|
||||
}
|
||||
}
|
||||
|
||||
if !all_text.is_empty() || !all_images.is_empty() {
|
||||
return build_sse_response(&all_text, &all_images, model);
|
||||
}
|
||||
}
|
||||
|
||||
Err("响应中没有 candidates".to_string())
|
||||
}
|
||||
|
||||
/// 构建 SSE 响应
|
||||
fn build_sse_response(
|
||||
text: &str,
|
||||
images: &[(String, String)],
|
||||
model: &str,
|
||||
) -> Result<String, String> {
|
||||
let mut content = text.to_string();
|
||||
|
||||
// 添加图片
|
||||
for (mime, data) in images {
|
||||
let image_url = format!("data:{};base64,{}", mime, data);
|
||||
content.push_str(&format!("\n\n", image_url));
|
||||
}
|
||||
|
||||
eprintln!("[ANTIGRAVITY_PARSE] 构建 SSE,内容长度: {}", content.len());
|
||||
|
||||
let chunk_id = format!("chatcmpl-{}", uuid::Uuid::new_v4());
|
||||
let created = chrono::Utc::now().timestamp();
|
||||
|
||||
let mut sse_output = String::new();
|
||||
|
||||
if !content.is_empty() {
|
||||
let content_chunk = serde_json::json!({
|
||||
"id": &chunk_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"delta": { "content": content },
|
||||
"finish_reason": serde_json::Value::Null
|
||||
}]
|
||||
});
|
||||
sse_output.push_str(&format!("data: {}\n\n", content_chunk.to_string()));
|
||||
}
|
||||
|
||||
let done_chunk = serde_json::json!({
|
||||
"id": &chunk_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"delta": {},
|
||||
"finish_reason": "stop"
|
||||
}]
|
||||
});
|
||||
sse_output.push_str(&format!("data: {}\n\n", done_chunk.to_string()));
|
||||
sse_output.push_str("data: [DONE]\n\n");
|
||||
|
||||
Ok(sse_output)
|
||||
}
|
||||
|
||||
/// 将 Gemini 流式响应 chunk 转换为 OpenAI SSE 格式
|
||||
///
|
||||
/// Gemini 流式响应格式:
|
||||
/// ```json
|
||||
/// {"candidates":[{"content":{"parts":[{"text":"Hello"}],"role":"model"},"finishReason":"STOP"}]}
|
||||
/// ```
|
||||
///
|
||||
/// OpenAI SSE 格式:
|
||||
/// ```
|
||||
/// data: {"id":"chatcmpl-xxx","object":"chat.completion.chunk","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]}
|
||||
/// ```
|
||||
fn convert_gemini_chunk_to_openai_sse(json: &serde_json::Value, model: &str) -> Option<String> {
|
||||
// 检查是否有 candidates
|
||||
let candidates = json.get("candidates")?.as_array()?;
|
||||
if candidates.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let candidate = &candidates[0];
|
||||
|
||||
// 提取文本内容
|
||||
let mut content_delta: Option<String> = None;
|
||||
let mut has_image = false;
|
||||
let mut image_data: Option<String> = None;
|
||||
|
||||
if let Some(content) = candidate.get("content") {
|
||||
if let Some(parts) = content.get("parts").and_then(|p| p.as_array()) {
|
||||
for part in parts {
|
||||
// 处理文本
|
||||
if let Some(text) = part.get("text").and_then(|t| t.as_str()) {
|
||||
content_delta = Some(text.to_string());
|
||||
}
|
||||
|
||||
// 处理图片(inlineData)
|
||||
if let Some(inline_data) =
|
||||
part.get("inlineData").or_else(|| part.get("inline_data"))
|
||||
{
|
||||
if let Some(data) = inline_data.get("data").and_then(|d| d.as_str()) {
|
||||
let mime_type = inline_data
|
||||
.get("mimeType")
|
||||
.or_else(|| inline_data.get("mime_type"))
|
||||
.and_then(|m| m.as_str())
|
||||
.unwrap_or("image/png");
|
||||
|
||||
// 将图片作为 markdown 格式的 data URL
|
||||
let image_url = format!("data:{};base64,{}", mime_type, data);
|
||||
image_data = Some(format!("\n\n", image_url));
|
||||
has_image = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 检查 finish_reason
|
||||
let finish_reason = candidate
|
||||
.get("finishReason")
|
||||
.and_then(|f| f.as_str())
|
||||
.map(|r| match r {
|
||||
"STOP" => "stop",
|
||||
"MAX_TOKENS" => "length",
|
||||
"SAFETY" => "content_filter",
|
||||
"RECITATION" => "content_filter",
|
||||
_ => "stop",
|
||||
});
|
||||
|
||||
// 如果没有内容变化且没有 finish_reason,跳过
|
||||
if content_delta.is_none() && !has_image && finish_reason.is_none() {
|
||||
return None;
|
||||
}
|
||||
|
||||
// 合并文本和图片内容
|
||||
let final_content = match (content_delta, image_data) {
|
||||
(Some(text), Some(img)) => Some(format!("{}{}", text, img)),
|
||||
(Some(text), None) => Some(text),
|
||||
(None, Some(img)) => Some(img),
|
||||
(None, None) => None,
|
||||
};
|
||||
|
||||
// 构建 OpenAI 格式的 delta
|
||||
let mut delta = serde_json::json!({});
|
||||
if let Some(content) = final_content {
|
||||
delta["content"] = serde_json::Value::String(content);
|
||||
}
|
||||
|
||||
// 构建完整的 SSE 事件
|
||||
let chunk_id = format!("chatcmpl-{}", uuid::Uuid::new_v4());
|
||||
let created = chrono::Utc::now().timestamp();
|
||||
|
||||
let response = serde_json::json!({
|
||||
"id": chunk_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created,
|
||||
"model": model,
|
||||
"choices": [{
|
||||
"index": 0,
|
||||
"delta": delta,
|
||||
"finish_reason": finish_reason
|
||||
}]
|
||||
});
|
||||
|
||||
Some(format!("data: {}\n\n", response.to_string()))
|
||||
}
|
||||
|
||||
@@ -867,16 +867,37 @@ pub async fn call_provider_openai_for_ws(
|
||||
}
|
||||
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)),
|
||||
|
||||
// 使用新的 validate_token() 方法检查 Token 状态
|
||||
let validation_result = antigravity.validate_token();
|
||||
tracing::info!("[Antigravity WS] Token 验证结果: {:?}", validation_result);
|
||||
|
||||
// 根据验证结果决定是否刷新
|
||||
if validation_result.needs_refresh() {
|
||||
tracing::info!("[Antigravity WS] Token 需要刷新,开始刷新...");
|
||||
match antigravity.refresh_token_with_retry(3).await {
|
||||
Ok(new_token) => {
|
||||
tracing::info!(
|
||||
"[Antigravity WS] Token 刷新成功,新 token 长度: {}",
|
||||
new_token.len()
|
||||
);
|
||||
// 刷新成功,标记为健康
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_healthy(db, &credential.uuid, None);
|
||||
}
|
||||
}
|
||||
Err(refresh_error) => {
|
||||
tracing::error!("[Antigravity WS] Token 刷新失败: {:?}", refresh_error);
|
||||
// 使用新的 mark_unhealthy_with_details 方法
|
||||
if let Some(db) = &state.db {
|
||||
let _ = state.pool_service.mark_unhealthy_with_details(
|
||||
db,
|
||||
&credential.uuid,
|
||||
&refresh_error,
|
||||
);
|
||||
}
|
||||
return Err(refresh_error.user_message());
|
||||
}
|
||||
return Err(e.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+75
-27
@@ -560,23 +560,43 @@ async fn update_processor_config(processor: &RequestProcessor, config: &Config)
|
||||
{
|
||||
let mut router = processor.router.write().await;
|
||||
router.clear_rules();
|
||||
for rule in &config.routing.rules {
|
||||
// 解析 provider 字符串为 ProviderType
|
||||
if let Ok(provider_type) = rule.provider.parse::<crate::ProviderType>() {
|
||||
router.add_rule(crate::router::RoutingRule {
|
||||
pattern: rule.pattern.clone(),
|
||||
target_provider: provider_type,
|
||||
priority: rule.priority,
|
||||
enabled: true,
|
||||
});
|
||||
} else {
|
||||
tracing::warn!("[HOT_RELOAD] 无法解析 provider: {}", rule.provider);
|
||||
|
||||
// 如果配置文件中有路由规则,使用配置文件的规则
|
||||
// 否则使用默认规则
|
||||
if !config.routing.rules.is_empty() {
|
||||
for rule in &config.routing.rules {
|
||||
// 解析 provider 字符串为 ProviderType
|
||||
if let Ok(provider_type) = rule.provider.parse::<crate::ProviderType>() {
|
||||
router.add_rule(crate::router::RoutingRule {
|
||||
pattern: rule.pattern.clone(),
|
||||
target_provider: provider_type,
|
||||
priority: rule.priority,
|
||||
enabled: true,
|
||||
});
|
||||
} else {
|
||||
tracing::warn!("[HOT_RELOAD] 无法解析 provider: {}", rule.provider);
|
||||
}
|
||||
}
|
||||
tracing::debug!(
|
||||
"[HOT_RELOAD] 路由规则已更新: {} 条规则(来自配置文件)",
|
||||
config.routing.rules.len()
|
||||
);
|
||||
} else {
|
||||
// 使用默认路由规则
|
||||
router.add_rule(crate::router::RoutingRule::new(
|
||||
"gemini-*",
|
||||
crate::ProviderType::Antigravity,
|
||||
10,
|
||||
));
|
||||
router.add_rule(crate::router::RoutingRule::new(
|
||||
"claude-*",
|
||||
crate::ProviderType::Kiro,
|
||||
10,
|
||||
));
|
||||
tracing::debug!(
|
||||
"[HOT_RELOAD] 路由规则已更新: 使用默认规则 (gemini-* → Antigravity, claude-* → Kiro)"
|
||||
);
|
||||
}
|
||||
tracing::debug!(
|
||||
"[HOT_RELOAD] 路由规则已更新: {} 条规则",
|
||||
config.routing.rules.len()
|
||||
);
|
||||
}
|
||||
|
||||
// 更新模型映射器
|
||||
@@ -1033,18 +1053,46 @@ async fn gemini_generate_content(
|
||||
.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 刷新失败: {}", e)
|
||||
}
|
||||
})),
|
||||
)
|
||||
.into_response();
|
||||
// 使用新的 validate_token() 方法检查 Token 状态
|
||||
let validation_result = antigravity.validate_token();
|
||||
tracing::info!(
|
||||
"[Antigravity Gemini] Token 验证结果: {:?}",
|
||||
validation_result
|
||||
);
|
||||
|
||||
// 根据验证结果决定是否刷新
|
||||
if validation_result.needs_refresh() {
|
||||
tracing::info!("[Antigravity Gemini] Token 需要刷新,开始刷新...");
|
||||
match antigravity.refresh_token_with_retry(3).await {
|
||||
Ok(new_token) => {
|
||||
tracing::info!(
|
||||
"[Antigravity Gemini] Token 刷新成功,新 token 长度: {}",
|
||||
new_token.len()
|
||||
);
|
||||
}
|
||||
Err(refresh_error) => {
|
||||
tracing::error!("[Antigravity Gemini] Token 刷新失败: {:?}", refresh_error);
|
||||
|
||||
// 根据错误类型返回不同的状态码和消息
|
||||
let (status, message) = if refresh_error.requires_reauth() {
|
||||
(StatusCode::UNAUTHORIZED, refresh_error.user_message())
|
||||
} else {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
refresh_error.user_message(),
|
||||
)
|
||||
};
|
||||
|
||||
return (
|
||||
status,
|
||||
Json(serde_json::json!({
|
||||
"error": {
|
||||
"message": message
|
||||
}
|
||||
})),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -636,6 +636,12 @@ pub async fn models() -> impl IntoResponse {
|
||||
{"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"},
|
||||
{"id": "gemini-3-pro-image-preview", "object": "model", "owned_by": "google"},
|
||||
{"id": "gemini-3-flash-preview", "object": "model", "owned_by": "google"},
|
||||
{"id": "gemini-2.5-computer-use-preview-10-2025", "object": "model", "owned_by": "google"},
|
||||
{"id": "gemini-claude-sonnet-4-5", "object": "model", "owned_by": "google"},
|
||||
{"id": "gemini-claude-sonnet-4-5-thinking", "object": "model", "owned_by": "google"},
|
||||
{"id": "gemini-claude-opus-4-5-thinking", "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"}
|
||||
@@ -958,4 +964,306 @@ mod property_tests {
|
||||
prop_assert_eq!(input_tokens, expected_input);
|
||||
}
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// Property 1: 模型列表结构和归属正确性
|
||||
// **Validates: Requirements 1.2, 1.3**
|
||||
// ========================================================================
|
||||
|
||||
/// 获取模型列表数据用于测试
|
||||
fn get_model_list_data() -> Vec<serde_json::Value> {
|
||||
vec![
|
||||
// Kiro/Claude models
|
||||
serde_json::json!({"id": "claude-sonnet-4-5", "object": "model", "owned_by": "anthropic"}),
|
||||
serde_json::json!({"id": "claude-sonnet-4-5-20250929", "object": "model", "owned_by": "anthropic"}),
|
||||
serde_json::json!({"id": "claude-3-7-sonnet-20250219", "object": "model", "owned_by": "anthropic"}),
|
||||
serde_json::json!({"id": "claude-3-5-sonnet-latest", "object": "model", "owned_by": "anthropic"}),
|
||||
// Gemini models
|
||||
serde_json::json!({"id": "gemini-2.5-flash", "object": "model", "owned_by": "google"}),
|
||||
serde_json::json!({"id": "gemini-2.5-flash-lite", "object": "model", "owned_by": "google"}),
|
||||
serde_json::json!({"id": "gemini-2.5-pro", "object": "model", "owned_by": "google"}),
|
||||
serde_json::json!({"id": "gemini-2.5-pro-preview-06-05", "object": "model", "owned_by": "google"}),
|
||||
serde_json::json!({"id": "gemini-3-pro-preview", "object": "model", "owned_by": "google"}),
|
||||
serde_json::json!({"id": "gemini-3-pro-image-preview", "object": "model", "owned_by": "google"}),
|
||||
serde_json::json!({"id": "gemini-3-flash-preview", "object": "model", "owned_by": "google"}),
|
||||
serde_json::json!({"id": "gemini-2.5-computer-use-preview-10-2025", "object": "model", "owned_by": "google"}),
|
||||
serde_json::json!({"id": "gemini-claude-sonnet-4-5", "object": "model", "owned_by": "google"}),
|
||||
serde_json::json!({"id": "gemini-claude-sonnet-4-5-thinking", "object": "model", "owned_by": "google"}),
|
||||
serde_json::json!({"id": "gemini-claude-opus-4-5-thinking", "object": "model", "owned_by": "google"}),
|
||||
// Qwen models
|
||||
serde_json::json!({"id": "qwen3-coder-plus", "object": "model", "owned_by": "alibaba"}),
|
||||
serde_json::json!({"id": "qwen3-coder-flash", "object": "model", "owned_by": "alibaba"}),
|
||||
]
|
||||
}
|
||||
|
||||
/// Property 1: 模型列表结构和归属正确性
|
||||
///
|
||||
/// *对于任意* 模型列表中的模型,应该包含:
|
||||
/// - id 字段 (非空字符串)
|
||||
/// - object 字段 (值为 "model")
|
||||
/// - owned_by 字段 (与模型类型匹配: gemini-* -> google, claude-* -> anthropic, qwen* -> alibaba)
|
||||
///
|
||||
/// **Validates: Requirements 1.2, 1.3**
|
||||
#[test]
|
||||
fn prop_model_list_structure_and_ownership() {
|
||||
let models = get_model_list_data();
|
||||
|
||||
for model in &models {
|
||||
// 验证 id 字段存在且非空
|
||||
let id = model.get("id").and_then(|v| v.as_str());
|
||||
assert!(id.is_some(), "Model should have id field");
|
||||
assert!(!id.unwrap().is_empty(), "Model id should not be empty");
|
||||
|
||||
// 验证 object 字段为 "model"
|
||||
let object = model.get("object").and_then(|v| v.as_str());
|
||||
assert_eq!(object, Some("model"), "Model object should be 'model'");
|
||||
|
||||
// 验证 owned_by 字段与模型类型匹配
|
||||
let owned_by = model.get("owned_by").and_then(|v| v.as_str());
|
||||
assert!(owned_by.is_some(), "Model should have owned_by field");
|
||||
|
||||
let model_id = id.unwrap();
|
||||
let owner = owned_by.unwrap();
|
||||
|
||||
if model_id.starts_with("gemini-") {
|
||||
assert_eq!(owner, "google", "Gemini models should be owned by google");
|
||||
} else if model_id.starts_with("claude-") {
|
||||
assert_eq!(
|
||||
owner, "anthropic",
|
||||
"Claude models should be owned by anthropic"
|
||||
);
|
||||
} else if model_id.starts_with("qwen") {
|
||||
assert_eq!(owner, "alibaba", "Qwen models should be owned by alibaba");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 验证所有 Antigravity 支持的模型都在列表中
|
||||
#[test]
|
||||
fn test_antigravity_models_present() {
|
||||
let models = get_model_list_data();
|
||||
let model_ids: Vec<&str> = models
|
||||
.iter()
|
||||
.filter_map(|m| m.get("id").and_then(|v| v.as_str()))
|
||||
.collect();
|
||||
|
||||
// 验证所有 Antigravity 支持的模型都存在
|
||||
let required_models = [
|
||||
"gemini-3-pro-preview",
|
||||
"gemini-3-pro-image-preview",
|
||||
"gemini-3-flash-preview",
|
||||
"gemini-2.5-computer-use-preview-10-2025",
|
||||
"gemini-claude-sonnet-4-5",
|
||||
"gemini-claude-sonnet-4-5-thinking",
|
||||
"gemini-claude-opus-4-5-thinking",
|
||||
];
|
||||
|
||||
for required in &required_models {
|
||||
assert!(
|
||||
model_ids.contains(required),
|
||||
"Model {} should be in the list",
|
||||
required
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// Property 2: 模型名称映射正确性
|
||||
// **Validates: Requirements 3.1, 3.2**
|
||||
// ========================================================================
|
||||
|
||||
/// 获取模型名称映射的预期结果
|
||||
fn get_expected_model_mapping(model: &str) -> &str {
|
||||
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,
|
||||
}
|
||||
}
|
||||
|
||||
/// Property 2: 模型名称映射正确性
|
||||
///
|
||||
/// *对于任意* 已知映射表中的模型名称,build_gemini_native_request 应该返回正确的内部模型名称。
|
||||
/// *对于任意* 不在映射表中的模型名称,应该原样返回。
|
||||
///
|
||||
/// **Validates: Requirements 3.1, 3.2**
|
||||
#[test]
|
||||
fn prop_model_name_mapping_correctness() {
|
||||
let test_request = serde_json::json!({
|
||||
"contents": [{"role": "user", "parts": [{"text": "test"}]}]
|
||||
});
|
||||
let project_id = "test-project";
|
||||
|
||||
// 测试已知映射
|
||||
let known_mappings = [
|
||||
("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",
|
||||
),
|
||||
];
|
||||
|
||||
for (input, expected) in &known_mappings {
|
||||
let result = build_gemini_native_request(&test_request, input, project_id);
|
||||
let actual_model = result.get("model").and_then(|v| v.as_str()).unwrap();
|
||||
assert_eq!(
|
||||
actual_model, *expected,
|
||||
"Model {} should map to {}",
|
||||
input, expected
|
||||
);
|
||||
}
|
||||
|
||||
// 测试未知模型名称应该原样返回
|
||||
let unknown_models = ["gemini-2.0-flash", "gemini-2.5-flash", "custom-model"];
|
||||
for model in &unknown_models {
|
||||
let result = build_gemini_native_request(&test_request, model, project_id);
|
||||
let actual_model = result.get("model").and_then(|v| v.as_str()).unwrap();
|
||||
assert_eq!(
|
||||
actual_model, *model,
|
||||
"Unknown model {} should be returned unchanged",
|
||||
model
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// Property 3: 思维链启用逻辑正确性
|
||||
// **Validates: Requirements 4.1, 4.2**
|
||||
// ========================================================================
|
||||
|
||||
/// 判断模型是否应该启用思维链
|
||||
fn should_enable_thinking(model: &str) -> bool {
|
||||
model.ends_with("-thinking")
|
||||
|| model == "gemini-2.5-pro"
|
||||
|| model.starts_with("gemini-3-pro-")
|
||||
|| model == "rev19-uic3-1p"
|
||||
|| model == "gpt-oss-120b-medium"
|
||||
}
|
||||
|
||||
/// Property 3: 思维链启用逻辑正确性
|
||||
///
|
||||
/// *对于任意* 模型名称,思维链启用状态应该根据以下规则正确判断:
|
||||
/// - 以 "-thinking" 结尾的模型启用思维链
|
||||
/// - "gemini-2.5-pro" 启用思维链
|
||||
/// - 以 "gemini-3-pro-" 开头的模型启用思维链
|
||||
/// - "rev19-uic3-1p" 启用思维链
|
||||
/// - "gpt-oss-120b-medium" 启用思维链
|
||||
///
|
||||
/// **Validates: Requirements 4.1, 4.2**
|
||||
#[test]
|
||||
fn prop_thinking_mode_enablement_logic() {
|
||||
// 应该启用思维链的模型
|
||||
let thinking_enabled_models = [
|
||||
"gemini-claude-sonnet-4-5-thinking",
|
||||
"gemini-claude-opus-4-5-thinking",
|
||||
"custom-model-thinking",
|
||||
"gemini-2.5-pro",
|
||||
"gemini-3-pro-preview",
|
||||
"gemini-3-pro-image-preview",
|
||||
"gemini-3-pro-high",
|
||||
"rev19-uic3-1p",
|
||||
"gpt-oss-120b-medium",
|
||||
];
|
||||
|
||||
for model in &thinking_enabled_models {
|
||||
assert!(
|
||||
should_enable_thinking(model),
|
||||
"Model {} should have thinking enabled",
|
||||
model
|
||||
);
|
||||
}
|
||||
|
||||
// 不应该启用思维链的模型
|
||||
let thinking_disabled_models = [
|
||||
"gemini-2.0-flash",
|
||||
"gemini-2.5-flash",
|
||||
"gemini-claude-sonnet-4-5",
|
||||
"claude-sonnet-4-5",
|
||||
"custom-model",
|
||||
];
|
||||
|
||||
for model in &thinking_disabled_models {
|
||||
assert!(
|
||||
!should_enable_thinking(model),
|
||||
"Model {} should have thinking disabled",
|
||||
model
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// ========================================================================
|
||||
// Property 4: 思维链配置值正确性
|
||||
// **Validates: Requirements 4.3, 4.4**
|
||||
// ========================================================================
|
||||
|
||||
/// Property 4: 思维链配置值正确性
|
||||
///
|
||||
/// *对于任意* Gemini 原生请求:
|
||||
/// - 当思维链启用时,includeThoughts 应该为 true,thinkingBudget 应该为 1024
|
||||
/// - 当思维链禁用时,includeThoughts 应该为 false,thinkingBudget 应该为 0
|
||||
///
|
||||
/// **Validates: Requirements 4.3, 4.4**
|
||||
#[test]
|
||||
fn prop_thinking_configuration_values() {
|
||||
let test_request = serde_json::json!({
|
||||
"contents": [{"role": "user", "parts": [{"text": "test"}]}]
|
||||
});
|
||||
let project_id = "test-project";
|
||||
|
||||
// 测试启用思维链的模型
|
||||
let thinking_enabled_models = [
|
||||
"gemini-3-pro-preview",
|
||||
"gemini-2.5-pro",
|
||||
"gemini-claude-sonnet-4-5-thinking",
|
||||
];
|
||||
|
||||
for model in &thinking_enabled_models {
|
||||
let result = build_gemini_native_request(&test_request, model, project_id);
|
||||
let thinking_config = &result["request"]["generationConfig"]["thinkingConfig"];
|
||||
|
||||
assert_eq!(
|
||||
thinking_config["includeThoughts"].as_bool(),
|
||||
Some(true),
|
||||
"Model {} should have includeThoughts=true",
|
||||
model
|
||||
);
|
||||
assert_eq!(
|
||||
thinking_config["thinkingBudget"].as_i64(),
|
||||
Some(1024),
|
||||
"Model {} should have thinkingBudget=1024",
|
||||
model
|
||||
);
|
||||
}
|
||||
|
||||
// 测试禁用思维链的模型
|
||||
let thinking_disabled_models = [
|
||||
"gemini-2.0-flash",
|
||||
"gemini-2.5-flash",
|
||||
"gemini-claude-sonnet-4-5",
|
||||
];
|
||||
|
||||
for model in &thinking_disabled_models {
|
||||
let result = build_gemini_native_request(&test_request, model, project_id);
|
||||
let thinking_config = &result["request"]["generationConfig"]["thinkingConfig"];
|
||||
|
||||
assert_eq!(
|
||||
thinking_config["includeThoughts"].as_bool(),
|
||||
Some(false),
|
||||
"Model {} should have includeThoughts=false",
|
||||
model
|
||||
);
|
||||
assert_eq!(
|
||||
thinking_config["thinkingBudget"].as_i64(),
|
||||
Some(0),
|
||||
"Model {} should have thinkingBudget=0",
|
||||
model
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,13 +12,49 @@ use crate::models::provider_pool_model::{
|
||||
ProviderPoolOverview,
|
||||
};
|
||||
use crate::models::route_model::RouteInfo;
|
||||
use crate::providers::antigravity::TokenRefreshError;
|
||||
use crate::providers::kiro::KiroProvider;
|
||||
use chrono::Utc;
|
||||
use reqwest::Client;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::atomic::AtomicUsize;
|
||||
use std::time::Duration;
|
||||
|
||||
/// 凭证健康信息
|
||||
/// Requirements: 3.1, 3.2
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct CredentialHealthInfo {
|
||||
/// 凭证 UUID
|
||||
pub uuid: String,
|
||||
/// 凭证名称
|
||||
pub name: Option<String>,
|
||||
/// Provider 类型
|
||||
pub provider_type: String,
|
||||
/// 是否健康
|
||||
pub is_healthy: bool,
|
||||
/// 最后错误信息
|
||||
pub last_error: Option<String>,
|
||||
/// 最后错误时间(RFC3339 格式)
|
||||
pub last_error_time: Option<String>,
|
||||
/// 错误次数
|
||||
pub failure_count: u32,
|
||||
/// 是否需要重新授权
|
||||
pub requires_reauth: bool,
|
||||
}
|
||||
|
||||
/// 凭证选择错误
|
||||
/// Requirements: 3.4
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub enum SelectionError {
|
||||
/// 没有凭证
|
||||
NoCredentials,
|
||||
/// 所有凭证都不健康
|
||||
AllUnhealthy { details: Vec<CredentialHealthInfo> },
|
||||
/// 模型不支持
|
||||
ModelNotSupported { model: String },
|
||||
}
|
||||
|
||||
/// 凭证池管理服务
|
||||
pub struct ProviderPoolService {
|
||||
/// HTTP 客户端(用于健康检测)
|
||||
@@ -372,6 +408,215 @@ impl ProviderPoolService {
|
||||
ProviderPoolDao::reset_health_by_type(&conn, &pt).map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 获取凭证健康状态
|
||||
/// Requirements: 3.2
|
||||
pub fn get_credential_health(
|
||||
&self,
|
||||
db: &DbConnection,
|
||||
uuid: &str,
|
||||
) -> Result<Option<CredentialHealthInfo>, String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
let cred = ProviderPoolDao::get_by_uuid(&conn, uuid).map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(cred.map(|c| CredentialHealthInfo {
|
||||
uuid: c.uuid.clone(),
|
||||
name: c.name.clone(),
|
||||
provider_type: c.provider_type.to_string(),
|
||||
is_healthy: c.is_healthy,
|
||||
last_error: c.last_error_message.clone(),
|
||||
last_error_time: c.last_error_time.map(|t| t.to_rfc3339()),
|
||||
failure_count: c.error_count,
|
||||
requires_reauth: c
|
||||
.last_error_message
|
||||
.as_ref()
|
||||
.map(|e| e.contains("invalid_grant") || e.contains("重新授权"))
|
||||
.unwrap_or(false),
|
||||
}))
|
||||
}
|
||||
|
||||
/// 获取所有凭证的健康状态
|
||||
/// Requirements: 3.2
|
||||
pub fn get_all_credential_health(
|
||||
&self,
|
||||
db: &DbConnection,
|
||||
) -> Result<Vec<CredentialHealthInfo>, String> {
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
let credentials = ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())?;
|
||||
|
||||
Ok(credentials
|
||||
.into_iter()
|
||||
.map(|c| CredentialHealthInfo {
|
||||
uuid: c.uuid.clone(),
|
||||
name: c.name.clone(),
|
||||
provider_type: c.provider_type.to_string(),
|
||||
is_healthy: c.is_healthy,
|
||||
last_error: c.last_error_message.clone(),
|
||||
last_error_time: c.last_error_time.map(|t| t.to_rfc3339()),
|
||||
failure_count: c.error_count,
|
||||
requires_reauth: c
|
||||
.last_error_message
|
||||
.as_ref()
|
||||
.map(|e| e.contains("invalid_grant") || e.contains("重新授权"))
|
||||
.unwrap_or(false),
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
/// 标记凭证为不健康(带详细错误信息)
|
||||
/// Requirements: 3.1, 3.2
|
||||
pub fn mark_unhealthy_with_details(
|
||||
&self,
|
||||
db: &DbConnection,
|
||||
uuid: &str,
|
||||
error: &TokenRefreshError,
|
||||
) -> Result<(), String> {
|
||||
let error_message = error.user_message();
|
||||
let requires_reauth = error.requires_reauth();
|
||||
|
||||
let conn = db.lock().map_err(|e| e.to_string())?;
|
||||
let cred = ProviderPoolDao::get_by_uuid(&conn, uuid)
|
||||
.map_err(|e| e.to_string())?
|
||||
.ok_or_else(|| format!("Credential not found: {}", uuid))?;
|
||||
|
||||
let new_error_count = cred.error_count + 1;
|
||||
// 如果需要重新授权,直接标记为不健康
|
||||
let is_healthy = if requires_reauth {
|
||||
false
|
||||
} else {
|
||||
new_error_count < self.max_error_count
|
||||
};
|
||||
|
||||
let error_msg = if requires_reauth {
|
||||
format!("[需要重新授权] {}", error_message)
|
||||
} else {
|
||||
error_message
|
||||
};
|
||||
|
||||
ProviderPoolDao::update_health_status(
|
||||
&conn,
|
||||
uuid,
|
||||
is_healthy,
|
||||
new_error_count,
|
||||
Some(Utc::now()),
|
||||
Some(&error_msg),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.map_err(|e| e.to_string())
|
||||
}
|
||||
|
||||
/// 选择一个健康的凭证
|
||||
/// Requirements: 2.4, 3.3, 3.4
|
||||
pub fn select_healthy_credential(
|
||||
&self,
|
||||
db: &DbConnection,
|
||||
provider_type: &str,
|
||||
model: Option<&str>,
|
||||
) -> Result<ProviderCredential, SelectionError> {
|
||||
let pt: PoolProviderType = provider_type
|
||||
.parse()
|
||||
.map_err(|_| SelectionError::NoCredentials)?;
|
||||
let conn = db.lock().map_err(|_| SelectionError::NoCredentials)?;
|
||||
let credentials =
|
||||
ProviderPoolDao::get_by_type(&conn, &pt).map_err(|_| SelectionError::NoCredentials)?;
|
||||
drop(conn);
|
||||
|
||||
if credentials.is_empty() {
|
||||
return Err(SelectionError::NoCredentials);
|
||||
}
|
||||
|
||||
// 过滤可用的凭证(健康且未禁用)
|
||||
let mut available: Vec<_> = credentials
|
||||
.iter()
|
||||
.filter(|c| c.is_available() && c.is_healthy)
|
||||
.collect();
|
||||
|
||||
// 如果指定了模型,进一步过滤支持该模型的凭证
|
||||
if let Some(m) = model {
|
||||
available.retain(|c| c.supports_model(m));
|
||||
if available.is_empty() {
|
||||
// 检查是否有凭证支持该模型但不健康
|
||||
let unhealthy_supporting: Vec<_> = credentials
|
||||
.iter()
|
||||
.filter(|c| c.supports_model(m) && !c.is_healthy)
|
||||
.collect();
|
||||
|
||||
if !unhealthy_supporting.is_empty() {
|
||||
// 返回不健康凭证的详细信息
|
||||
let details: Vec<CredentialHealthInfo> = unhealthy_supporting
|
||||
.into_iter()
|
||||
.map(|c| CredentialHealthInfo {
|
||||
uuid: c.uuid.clone(),
|
||||
name: c.name.clone(),
|
||||
provider_type: c.provider_type.to_string(),
|
||||
is_healthy: c.is_healthy,
|
||||
last_error: c.last_error_message.clone(),
|
||||
last_error_time: c.last_error_time.map(|t| t.to_rfc3339()),
|
||||
failure_count: c.error_count,
|
||||
requires_reauth: c
|
||||
.last_error_message
|
||||
.as_ref()
|
||||
.map(|e| e.contains("invalid_grant") || e.contains("重新授权"))
|
||||
.unwrap_or(false),
|
||||
})
|
||||
.collect();
|
||||
return Err(SelectionError::AllUnhealthy { details });
|
||||
}
|
||||
|
||||
return Err(SelectionError::ModelNotSupported {
|
||||
model: m.to_string(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
if available.is_empty() {
|
||||
// 所有凭证都不健康
|
||||
let details: Vec<CredentialHealthInfo> = credentials
|
||||
.iter()
|
||||
.filter(|c| !c.is_healthy)
|
||||
.map(|c| CredentialHealthInfo {
|
||||
uuid: c.uuid.clone(),
|
||||
name: c.name.clone(),
|
||||
provider_type: c.provider_type.to_string(),
|
||||
is_healthy: c.is_healthy,
|
||||
last_error: c.last_error_message.clone(),
|
||||
last_error_time: c.last_error_time.map(|t| t.to_rfc3339()),
|
||||
failure_count: c.error_count,
|
||||
requires_reauth: c
|
||||
.last_error_message
|
||||
.as_ref()
|
||||
.map(|e| e.contains("invalid_grant") || e.contains("重新授权"))
|
||||
.unwrap_or(false),
|
||||
})
|
||||
.collect();
|
||||
return Err(SelectionError::AllUnhealthy { details });
|
||||
}
|
||||
|
||||
// 使用轮询策略选择凭证
|
||||
let key = format!("{}:{}", provider_type, model.unwrap_or("*"));
|
||||
let index = {
|
||||
let indices = self.round_robin_index.read().unwrap();
|
||||
indices
|
||||
.get(&key)
|
||||
.map(|i| i.load(std::sync::atomic::Ordering::Relaxed))
|
||||
.unwrap_or(0)
|
||||
};
|
||||
|
||||
let selected_index = index % available.len();
|
||||
let selected = available[selected_index].clone();
|
||||
|
||||
// 更新轮询索引
|
||||
{
|
||||
let mut indices = self.round_robin_index.write().unwrap();
|
||||
indices
|
||||
.entry(key)
|
||||
.or_insert_with(|| AtomicUsize::new(0))
|
||||
.store(index + 1, std::sync::atomic::Ordering::Relaxed);
|
||||
}
|
||||
|
||||
Ok(selected)
|
||||
}
|
||||
|
||||
/// 执行单个凭证的健康检查
|
||||
///
|
||||
/// 如果遇到 401 错误,会自动尝试刷新 token 后重试
|
||||
@@ -1702,3 +1947,140 @@ pub struct MigrationResult {
|
||||
/// 错误信息列表
|
||||
pub errors: Vec<String>,
|
||||
}
|
||||
|
||||
// ==================== 测试模块 ====================
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
// ==================== Property 3: 不健康凭证排除 ====================
|
||||
// Feature: antigravity-token-refresh, Property 3: 不健康凭证排除
|
||||
// Validates: Requirements 2.4, 3.3
|
||||
|
||||
#[test]
|
||||
fn test_credential_health_info_creation() {
|
||||
let info = CredentialHealthInfo {
|
||||
uuid: "test-uuid".to_string(),
|
||||
name: Some("Test Credential".to_string()),
|
||||
provider_type: "antigravity".to_string(),
|
||||
is_healthy: false,
|
||||
last_error: Some("Token refresh failed".to_string()),
|
||||
last_error_time: Some("2024-01-01T00:00:00Z".to_string()),
|
||||
failure_count: 3,
|
||||
requires_reauth: true,
|
||||
};
|
||||
|
||||
assert_eq!(info.uuid, "test-uuid");
|
||||
assert!(!info.is_healthy);
|
||||
assert!(info.requires_reauth);
|
||||
assert_eq!(info.failure_count, 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_selection_error_no_credentials() {
|
||||
let error = SelectionError::NoCredentials;
|
||||
// 验证可以序列化
|
||||
let json = serde_json::to_string(&error).unwrap();
|
||||
assert!(json.contains("NoCredentials"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_selection_error_all_unhealthy() {
|
||||
let details = vec![CredentialHealthInfo {
|
||||
uuid: "test-uuid".to_string(),
|
||||
name: Some("Test".to_string()),
|
||||
provider_type: "antigravity".to_string(),
|
||||
is_healthy: false,
|
||||
last_error: Some("invalid_grant".to_string()),
|
||||
last_error_time: None,
|
||||
failure_count: 1,
|
||||
requires_reauth: true,
|
||||
}];
|
||||
|
||||
let error = SelectionError::AllUnhealthy { details };
|
||||
let json = serde_json::to_string(&error).unwrap();
|
||||
assert!(json.contains("AllUnhealthy"));
|
||||
assert!(json.contains("invalid_grant"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_selection_error_model_not_supported() {
|
||||
let error = SelectionError::ModelNotSupported {
|
||||
model: "gpt-5".to_string(),
|
||||
};
|
||||
let json = serde_json::to_string(&error).unwrap();
|
||||
assert!(json.contains("ModelNotSupported"));
|
||||
assert!(json.contains("gpt-5"));
|
||||
}
|
||||
|
||||
// ==================== Property 4: 健康状态记录完整性 ====================
|
||||
// Feature: antigravity-token-refresh, Property 4: 健康状态记录完整性
|
||||
// Validates: Requirements 3.2
|
||||
|
||||
#[test]
|
||||
fn test_credential_health_info_requires_reauth_detection() {
|
||||
// 测试 invalid_grant 检测
|
||||
let info_with_invalid_grant = CredentialHealthInfo {
|
||||
uuid: "test".to_string(),
|
||||
name: None,
|
||||
provider_type: "antigravity".to_string(),
|
||||
is_healthy: false,
|
||||
last_error: Some("Token refresh failed: invalid_grant".to_string()),
|
||||
last_error_time: Some(chrono::Utc::now().to_rfc3339()),
|
||||
failure_count: 1,
|
||||
requires_reauth: true,
|
||||
};
|
||||
assert!(info_with_invalid_grant.requires_reauth);
|
||||
|
||||
// 测试重新授权检测
|
||||
let info_with_reauth = CredentialHealthInfo {
|
||||
uuid: "test".to_string(),
|
||||
name: None,
|
||||
provider_type: "antigravity".to_string(),
|
||||
is_healthy: false,
|
||||
last_error: Some("[需要重新授权] Token 已过期".to_string()),
|
||||
last_error_time: Some(chrono::Utc::now().to_rfc3339()),
|
||||
failure_count: 1,
|
||||
requires_reauth: true,
|
||||
};
|
||||
assert!(info_with_reauth.requires_reauth);
|
||||
|
||||
// 测试普通错误不需要重新授权
|
||||
let info_normal_error = CredentialHealthInfo {
|
||||
uuid: "test".to_string(),
|
||||
name: None,
|
||||
provider_type: "antigravity".to_string(),
|
||||
is_healthy: false,
|
||||
last_error: Some("Network error".to_string()),
|
||||
last_error_time: Some(chrono::Utc::now().to_rfc3339()),
|
||||
failure_count: 1,
|
||||
requires_reauth: false,
|
||||
};
|
||||
assert!(!info_normal_error.requires_reauth);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_credential_health_info_serialization() {
|
||||
let info = CredentialHealthInfo {
|
||||
uuid: "test-uuid".to_string(),
|
||||
name: Some("Test".to_string()),
|
||||
provider_type: "antigravity".to_string(),
|
||||
is_healthy: true,
|
||||
last_error: None,
|
||||
last_error_time: None,
|
||||
failure_count: 0,
|
||||
requires_reauth: false,
|
||||
};
|
||||
|
||||
// 测试序列化
|
||||
let json = serde_json::to_string(&info).unwrap();
|
||||
assert!(json.contains("test-uuid"));
|
||||
assert!(json.contains("antigravity"));
|
||||
|
||||
// 测试反序列化
|
||||
let deserialized: CredentialHealthInfo = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(deserialized.uuid, info.uuid);
|
||||
assert_eq!(deserialized.is_healthy, info.is_healthy);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -144,7 +144,40 @@ const MarkdownContainer = styled.div`
|
||||
|
||||
img {
|
||||
max-width: 100%;
|
||||
max-height: 512px;
|
||||
border-radius: 8px;
|
||||
object-fit: contain;
|
||||
cursor: pointer;
|
||||
transition: transform 0.2s ease;
|
||||
|
||||
&:hover {
|
||||
transform: scale(1.02);
|
||||
}
|
||||
}
|
||||
`;
|
||||
|
||||
// 图片容器样式
|
||||
const ImageContainer = styled.div`
|
||||
margin: 1em 0;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 8px;
|
||||
`;
|
||||
|
||||
const GeneratedImage = styled.img`
|
||||
max-width: 100%;
|
||||
max-height: 512px;
|
||||
border-radius: 8px;
|
||||
object-fit: contain;
|
||||
cursor: pointer;
|
||||
border: 1px solid hsl(var(--border));
|
||||
transition:
|
||||
transform 0.2s ease,
|
||||
box-shadow 0.2s ease;
|
||||
|
||||
&:hover {
|
||||
transform: scale(1.02);
|
||||
box-shadow: 0 4px 12px rgba(0, 0, 0, 0.15);
|
||||
}
|
||||
`;
|
||||
|
||||
@@ -200,59 +233,186 @@ export const MarkdownRenderer: React.FC<MarkdownRendererProps> = memo(
|
||||
setTimeout(() => setCopied(null), 2000);
|
||||
};
|
||||
|
||||
// 预处理内容:检测并提取 base64 图片
|
||||
const processedContent = React.useMemo(() => {
|
||||
// 匹配 markdown 图片语法中的 base64 data URL
|
||||
const base64ImageRegex =
|
||||
/!\[([^\]]*)\]\((data:image\/[^;]+;base64,[^)]+)\)/g;
|
||||
let result = content;
|
||||
const images: { alt: string; src: string; placeholder: string }[] = [];
|
||||
|
||||
let match;
|
||||
let index = 0;
|
||||
while ((match = base64ImageRegex.exec(content)) !== null) {
|
||||
const placeholder = `__BASE64_IMAGE_${index}__`;
|
||||
images.push({
|
||||
alt: match[1] || "Generated Image",
|
||||
src: match[2],
|
||||
placeholder,
|
||||
});
|
||||
result = result.replace(match[0], placeholder);
|
||||
index++;
|
||||
}
|
||||
|
||||
return { text: result, images };
|
||||
}, [content]);
|
||||
|
||||
// 渲染 base64 图片
|
||||
const renderBase64Images = () => {
|
||||
if (processedContent.images.length === 0) return null;
|
||||
|
||||
return processedContent.images.map((img, idx) => {
|
||||
const handleImageClick = () => {
|
||||
const newWindow = window.open();
|
||||
if (newWindow) {
|
||||
newWindow.document.write(`
|
||||
<html>
|
||||
<head>
|
||||
<title>${img.alt}</title>
|
||||
<style>
|
||||
body {
|
||||
margin: 0;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
min-height: 100vh;
|
||||
background: #1a1a1a;
|
||||
}
|
||||
img {
|
||||
max-width: 100%;
|
||||
max-height: 100vh;
|
||||
object-fit: contain;
|
||||
}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<img src="${img.src}" alt="${img.alt}" />
|
||||
</body>
|
||||
</html>
|
||||
`);
|
||||
newWindow.document.close();
|
||||
}
|
||||
};
|
||||
|
||||
return (
|
||||
<ImageContainer key={`base64-img-${idx}`}>
|
||||
<GeneratedImage
|
||||
src={img.src}
|
||||
alt={img.alt}
|
||||
onClick={handleImageClick}
|
||||
title="点击查看大图"
|
||||
onError={(e) => {
|
||||
console.error("[MarkdownRenderer] 图片加载失败:", img.alt);
|
||||
(e.target as HTMLImageElement).style.display = "none";
|
||||
}}
|
||||
onLoad={() => {
|
||||
console.log("[MarkdownRenderer] 图片加载成功:", img.alt);
|
||||
}}
|
||||
/>
|
||||
<span
|
||||
style={{
|
||||
fontSize: "12px",
|
||||
color: "hsl(var(--muted-foreground))",
|
||||
textAlign: "center",
|
||||
}}
|
||||
>
|
||||
🖼️ AI 生成图片 - 点击查看大图
|
||||
</span>
|
||||
</ImageContainer>
|
||||
);
|
||||
});
|
||||
};
|
||||
|
||||
// 检查处理后的文本是否只包含占位符
|
||||
const hasOnlyPlaceholders = React.useMemo(() => {
|
||||
const trimmed = processedContent.text.trim();
|
||||
return /^(__BASE64_IMAGE_\d+__\s*)+$/.test(trimmed) || trimmed === "";
|
||||
}, [processedContent.text]);
|
||||
|
||||
return (
|
||||
<MarkdownContainer>
|
||||
<ReactMarkdown
|
||||
remarkPlugins={[remarkGfm, remarkMath]}
|
||||
rehypePlugins={[rehypeRaw, rehypeKatex]}
|
||||
components={{
|
||||
code({ inline, className, children, ...props }: any) {
|
||||
const match = /language-(\w+)/.exec(className || "");
|
||||
const codeContent = String(children).replace(/\n$/, "");
|
||||
const language = match ? match[1] : "text";
|
||||
{/* 先渲染 base64 图片 */}
|
||||
{renderBase64Images()}
|
||||
|
||||
{/* 如果还有其他内容,渲染 markdown */}
|
||||
{!hasOnlyPlaceholders && processedContent.text.trim() && (
|
||||
<ReactMarkdown
|
||||
remarkPlugins={[remarkGfm, remarkMath]}
|
||||
rehypePlugins={[rehypeRaw, rehypeKatex]}
|
||||
components={{
|
||||
code({ inline, className, children, ...props }: any) {
|
||||
const match = /language-(\w+)/.exec(className || "");
|
||||
const codeContent = String(children).replace(/\n$/, "");
|
||||
const language = match ? match[1] : "text";
|
||||
|
||||
// Inline code
|
||||
if (inline) {
|
||||
return (
|
||||
<code className={className} {...props}>
|
||||
{children}
|
||||
</code>
|
||||
);
|
||||
}
|
||||
|
||||
// Block code
|
||||
const isCopied = copied === codeContent;
|
||||
|
||||
// Inline code
|
||||
if (inline) {
|
||||
return (
|
||||
<code className={className} {...props}>
|
||||
{children}
|
||||
</code>
|
||||
<CodeBlockContainer>
|
||||
<CodeHeader>
|
||||
<span>{language}</span>
|
||||
<CopyButton onClick={() => handleCopy(codeContent)}>
|
||||
{isCopied ? <Check size={14} /> : <Copy size={14} />}
|
||||
{isCopied ? "Copied" : "Copy"}
|
||||
</CopyButton>
|
||||
</CodeHeader>
|
||||
<SyntaxHighlighter
|
||||
style={oneDark}
|
||||
language={language}
|
||||
PreTag="div"
|
||||
customStyle={{
|
||||
margin: 0,
|
||||
padding: "16px",
|
||||
background: "transparent",
|
||||
fontSize: "13px",
|
||||
}}
|
||||
{...props}
|
||||
>
|
||||
{codeContent}
|
||||
</SyntaxHighlighter>
|
||||
</CodeBlockContainer>
|
||||
);
|
||||
}
|
||||
},
|
||||
// 普通图片渲染(非 base64)
|
||||
img({ src, alt, ...props }: any) {
|
||||
// base64 图片已经在上面单独处理了,这里只处理普通 URL 图片
|
||||
if (src?.startsWith("data:")) {
|
||||
return null; // 跳过 base64 图片,已在上面处理
|
||||
}
|
||||
|
||||
// Block code
|
||||
const isCopied = copied === codeContent;
|
||||
const handleImageClick = () => {
|
||||
if (src) {
|
||||
window.open(src, "_blank");
|
||||
}
|
||||
};
|
||||
|
||||
return (
|
||||
<CodeBlockContainer>
|
||||
<CodeHeader>
|
||||
<span>{language}</span>
|
||||
<CopyButton onClick={() => handleCopy(codeContent)}>
|
||||
{isCopied ? <Check size={14} /> : <Copy size={14} />}
|
||||
{isCopied ? "Copied" : "Copy"}
|
||||
</CopyButton>
|
||||
</CodeHeader>
|
||||
<SyntaxHighlighter
|
||||
style={oneDark}
|
||||
language={language}
|
||||
PreTag="div"
|
||||
customStyle={{
|
||||
margin: 0,
|
||||
padding: "16px",
|
||||
background: "transparent",
|
||||
fontSize: "13px",
|
||||
}}
|
||||
{...props}
|
||||
>
|
||||
{codeContent}
|
||||
</SyntaxHighlighter>
|
||||
</CodeBlockContainer>
|
||||
);
|
||||
},
|
||||
}}
|
||||
>
|
||||
{content}
|
||||
</ReactMarkdown>
|
||||
return (
|
||||
<ImageContainer>
|
||||
<GeneratedImage
|
||||
src={src}
|
||||
alt={alt || "Image"}
|
||||
onClick={handleImageClick}
|
||||
title="点击查看大图"
|
||||
{...props}
|
||||
/>
|
||||
</ImageContainer>
|
||||
);
|
||||
},
|
||||
}}
|
||||
>
|
||||
{processedContent.text}
|
||||
</ReactMarkdown>
|
||||
)}
|
||||
</MarkdownContainer>
|
||||
);
|
||||
},
|
||||
|
||||
@@ -472,6 +472,7 @@ export function useAgentChat() {
|
||||
activeSessionId, // 传递 sessionId 以保持上下文
|
||||
model || undefined,
|
||||
imagesToSend,
|
||||
providerType, // 传递用户选择的 provider
|
||||
);
|
||||
} catch (error) {
|
||||
toast.error(`发送失败: ${error}`);
|
||||
|
||||
@@ -99,6 +99,10 @@ export const PROVIDER_CONFIG: Record<
|
||||
antigravity: {
|
||||
label: "Antigravity",
|
||||
models: [
|
||||
"gemini-3-pro-preview",
|
||||
"gemini-3-pro-image-preview",
|
||||
"gemini-3-flash-preview",
|
||||
"gemini-2.5-computer-use-preview-10-2025",
|
||||
"gemini-claude-sonnet-4-5",
|
||||
"gemini-claude-sonnet-4-5-thinking",
|
||||
"gemini-claude-opus-4-5-thinking",
|
||||
|
||||
@@ -0,0 +1,167 @@
|
||||
/**
|
||||
* API Server Page 测试
|
||||
*
|
||||
* 测试 Antigravity 模型支持功能
|
||||
*
|
||||
* **Feature: antigravity-model-support**
|
||||
*/
|
||||
|
||||
import { describe, expect, test } from "vitest";
|
||||
|
||||
// ============================================================================
|
||||
// 从 ApiServerPage.tsx 提取的测试函数
|
||||
// ============================================================================
|
||||
|
||||
/**
|
||||
* 根据 Provider 类型获取 Gemini 测试模型列表
|
||||
*/
|
||||
function getGeminiTestModels(provider: string): string[] {
|
||||
switch (provider) {
|
||||
case "antigravity":
|
||||
return [
|
||||
"gemini-3-pro-preview",
|
||||
"gemini-3-pro-image-preview",
|
||||
"gemini-3-flash-preview",
|
||||
"gemini-claude-sonnet-4-5",
|
||||
];
|
||||
case "gemini":
|
||||
return ["gemini-2.0-flash", "gemini-2.5-flash", "gemini-2.5-pro"];
|
||||
default:
|
||||
return ["gemini-2.0-flash"];
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 根据 Provider 类型获取测试模型
|
||||
*/
|
||||
function getTestModel(provider: string): string {
|
||||
switch (provider) {
|
||||
case "antigravity":
|
||||
return "gemini-3-pro-preview";
|
||||
case "gemini":
|
||||
return "gemini-2.0-flash";
|
||||
case "qwen":
|
||||
return "qwen-max";
|
||||
case "openai":
|
||||
return "gpt-4o";
|
||||
case "claude":
|
||||
return "claude-sonnet-4-20250514";
|
||||
case "kiro":
|
||||
default:
|
||||
return "claude-opus-4-5-20251101";
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Property 测试: Antigravity 模型支持
|
||||
// **Validates: Requirements 2.1**
|
||||
// ============================================================================
|
||||
|
||||
describe("Antigravity Model Support", () => {
|
||||
/**
|
||||
* Property: Antigravity Provider 测试模型列表
|
||||
*
|
||||
* *对于* Antigravity provider,getGeminiTestModels 应该返回正确的模型列表:
|
||||
* - gemini-3-pro-preview
|
||||
* - gemini-3-pro-image-preview
|
||||
* - gemini-3-flash-preview
|
||||
* - gemini-claude-sonnet-4-5
|
||||
*
|
||||
* **Validates: Requirements 2.1**
|
||||
*/
|
||||
describe("getGeminiTestModels", () => {
|
||||
test("antigravity provider 应返回 4 个 Gemini 模型", () => {
|
||||
const models = getGeminiTestModels("antigravity");
|
||||
|
||||
expect(models).toHaveLength(4);
|
||||
expect(models).toContain("gemini-3-pro-preview");
|
||||
expect(models).toContain("gemini-3-pro-image-preview");
|
||||
expect(models).toContain("gemini-3-flash-preview");
|
||||
expect(models).toContain("gemini-claude-sonnet-4-5");
|
||||
});
|
||||
|
||||
test("gemini provider 应返回 3 个 Gemini 模型", () => {
|
||||
const models = getGeminiTestModels("gemini");
|
||||
|
||||
expect(models).toHaveLength(3);
|
||||
expect(models).toContain("gemini-2.0-flash");
|
||||
expect(models).toContain("gemini-2.5-flash");
|
||||
expect(models).toContain("gemini-2.5-pro");
|
||||
});
|
||||
|
||||
test("其他 provider 应返回默认模型列表", () => {
|
||||
const providers = ["kiro", "openai", "claude", "qwen", "unknown"];
|
||||
|
||||
for (const provider of providers) {
|
||||
const models = getGeminiTestModels(provider);
|
||||
expect(models).toEqual(["gemini-2.0-flash"]);
|
||||
}
|
||||
});
|
||||
});
|
||||
|
||||
/**
|
||||
* Property: Provider 默认测试模型
|
||||
*
|
||||
* *对于任意* provider,getTestModel 应该返回该 provider 的默认测试模型
|
||||
*
|
||||
* **Validates: Requirements 2.1**
|
||||
*/
|
||||
describe("getTestModel", () => {
|
||||
test("antigravity provider 应返回 gemini-3-pro-preview", () => {
|
||||
expect(getTestModel("antigravity")).toBe("gemini-3-pro-preview");
|
||||
});
|
||||
|
||||
test("gemini provider 应返回 gemini-2.0-flash", () => {
|
||||
expect(getTestModel("gemini")).toBe("gemini-2.0-flash");
|
||||
});
|
||||
|
||||
test("kiro provider 应返回 claude-opus-4-5-20251101", () => {
|
||||
expect(getTestModel("kiro")).toBe("claude-opus-4-5-20251101");
|
||||
});
|
||||
|
||||
test("openai provider 应返回 gpt-4o", () => {
|
||||
expect(getTestModel("openai")).toBe("gpt-4o");
|
||||
});
|
||||
|
||||
test("claude provider 应返回 claude-sonnet-4-20250514", () => {
|
||||
expect(getTestModel("claude")).toBe("claude-sonnet-4-20250514");
|
||||
});
|
||||
|
||||
test("qwen provider 应返回 qwen-max", () => {
|
||||
expect(getTestModel("qwen")).toBe("qwen-max");
|
||||
});
|
||||
|
||||
test("未知 provider 应返回默认模型", () => {
|
||||
expect(getTestModel("unknown")).toBe("claude-opus-4-5-20251101");
|
||||
});
|
||||
});
|
||||
|
||||
/**
|
||||
* Property: Gemini 测试端点显示条件
|
||||
*
|
||||
* *对于* antigravity 或 gemini provider,应该显示 Gemini 测试端点
|
||||
*
|
||||
* **Validates: Requirements 2.1**
|
||||
*/
|
||||
describe("showGeminiTest", () => {
|
||||
const shouldShowGeminiTest = (provider: string): boolean => {
|
||||
return provider === "antigravity" || provider === "gemini";
|
||||
};
|
||||
|
||||
test("antigravity provider 应显示 Gemini 测试端点", () => {
|
||||
expect(shouldShowGeminiTest("antigravity")).toBe(true);
|
||||
});
|
||||
|
||||
test("gemini provider 应显示 Gemini 测试端点", () => {
|
||||
expect(shouldShowGeminiTest("gemini")).toBe(true);
|
||||
});
|
||||
|
||||
test("其他 provider 不应显示 Gemini 测试端点", () => {
|
||||
const providers = ["kiro", "openai", "claude", "qwen"];
|
||||
|
||||
for (const provider of providers) {
|
||||
expect(shouldShowGeminiTest(provider)).toBe(false);
|
||||
}
|
||||
});
|
||||
});
|
||||
});
|
||||
@@ -263,19 +263,24 @@ export function ApiServerPage() {
|
||||
|
||||
const testModel = getTestModel(defaultProvider);
|
||||
|
||||
// 根据 Provider 类型获取 Gemini 测试模型
|
||||
const getGeminiTestModel = (provider: string): string => {
|
||||
// 根据 Provider 类型获取 Gemini 测试模型列表
|
||||
const getGeminiTestModels = (provider: string): string[] => {
|
||||
switch (provider) {
|
||||
case "antigravity":
|
||||
return "gemini-3-pro-preview";
|
||||
return [
|
||||
"gemini-3-pro-preview",
|
||||
"gemini-3-pro-image-preview",
|
||||
"gemini-3-flash-preview",
|
||||
"gemini-claude-sonnet-4-5",
|
||||
];
|
||||
case "gemini":
|
||||
return "gemini-2.0-flash";
|
||||
return ["gemini-2.0-flash", "gemini-2.5-flash", "gemini-2.5-pro"];
|
||||
default:
|
||||
return "gemini-2.0-flash";
|
||||
return ["gemini-2.0-flash"];
|
||||
}
|
||||
};
|
||||
|
||||
const geminiTestModel = getGeminiTestModel(defaultProvider);
|
||||
const geminiTestModels = getGeminiTestModels(defaultProvider);
|
||||
|
||||
// 是否显示 Gemini 测试端点
|
||||
const showGeminiTest =
|
||||
@@ -329,28 +334,24 @@ export function ApiServerPage() {
|
||||
},
|
||||
// Gemini 原生协议测试(仅在 Antigravity 或 Gemini Provider 时显示)
|
||||
...(showGeminiTest
|
||||
? [
|
||||
{
|
||||
id: "gemini",
|
||||
name: "Gemini Generate",
|
||||
method: "POST",
|
||||
path: `/v1/gemini/${geminiTestModel}:generateContent`,
|
||||
needsAuth: true,
|
||||
body: JSON.stringify({
|
||||
contents: [
|
||||
{
|
||||
role: "user",
|
||||
parts: [
|
||||
{ text: "What is 2+2? Answer with just the number." },
|
||||
],
|
||||
},
|
||||
],
|
||||
generationConfig: {
|
||||
maxOutputTokens: 100,
|
||||
? geminiTestModels.map((model, index) => ({
|
||||
id: `gemini-${index}`,
|
||||
name: `Gemini ${model}`,
|
||||
method: "POST",
|
||||
path: `/v1/gemini/${model}:generateContent`,
|
||||
needsAuth: true,
|
||||
body: JSON.stringify({
|
||||
contents: [
|
||||
{
|
||||
role: "user",
|
||||
parts: [{ text: "What is 2+2? Answer with just the number." }],
|
||||
},
|
||||
}),
|
||||
},
|
||||
]
|
||||
],
|
||||
generationConfig: {
|
||||
maxOutputTokens: 100,
|
||||
},
|
||||
}),
|
||||
}))
|
||||
: []),
|
||||
];
|
||||
|
||||
|
||||
@@ -672,9 +672,51 @@ export function CredentialCard({
|
||||
|
||||
{/* Error Message */}
|
||||
{credential.last_error_message && (
|
||||
<div className="mx-4 mb-3 rounded-lg bg-red-100 p-3 text-xs text-red-700 dark:bg-red-900/30 dark:text-red-300">
|
||||
{credential.last_error_message.slice(0, 150)}
|
||||
{credential.last_error_message.length > 150 && "..."}
|
||||
<div
|
||||
className={`mx-4 mb-3 rounded-lg p-3 text-xs ${
|
||||
credential.last_error_message.includes("invalid_grant") ||
|
||||
credential.last_error_message.includes("重新授权") ||
|
||||
credential.last_error_message.includes("凭证已过期")
|
||||
? "bg-amber-100 dark:bg-amber-900/30 border border-amber-300 dark:border-amber-700"
|
||||
: "bg-red-100 dark:bg-red-900/30"
|
||||
}`}
|
||||
>
|
||||
<div
|
||||
className={`${
|
||||
credential.last_error_message.includes("invalid_grant") ||
|
||||
credential.last_error_message.includes("重新授权") ||
|
||||
credential.last_error_message.includes("凭证已过期")
|
||||
? "text-amber-700 dark:text-amber-300"
|
||||
: "text-red-700 dark:text-red-300"
|
||||
}`}
|
||||
>
|
||||
{credential.last_error_message.slice(0, 150)}
|
||||
{credential.last_error_message.length > 150 && "..."}
|
||||
</div>
|
||||
{/* 重新授权提示 */}
|
||||
{(credential.last_error_message.includes("invalid_grant") ||
|
||||
credential.last_error_message.includes("重新授权") ||
|
||||
credential.last_error_message.includes("凭证已过期")) && (
|
||||
<div className="mt-2 pt-2 border-t border-amber-300 dark:border-amber-700">
|
||||
<div className="flex items-center justify-between">
|
||||
<span className="text-amber-600 dark:text-amber-400 font-medium">
|
||||
💡 需要重新授权
|
||||
</span>
|
||||
{onRefreshToken && (
|
||||
<button
|
||||
onClick={onRefreshToken}
|
||||
disabled={refreshingToken}
|
||||
className="px-3 py-1 text-xs font-medium bg-amber-600 text-white rounded hover:bg-amber-700 disabled:opacity-50 transition-colors"
|
||||
>
|
||||
{refreshingToken ? "刷新中..." : "尝试刷新"}
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
<p className="mt-1 text-amber-600/80 dark:text-amber-400/80">
|
||||
请删除此凭证并重新添加,或尝试刷新 Token
|
||||
</p>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ import {
|
||||
Trash2,
|
||||
Settings,
|
||||
CheckCircle2,
|
||||
KeyRound,
|
||||
} from "lucide-react";
|
||||
|
||||
export interface ErrorInfo {
|
||||
@@ -20,7 +21,8 @@ export interface ErrorInfo {
|
||||
| "migrate"
|
||||
| "config"
|
||||
| "general"
|
||||
| "success";
|
||||
| "success"
|
||||
| "reauth"; // 需要重新授权
|
||||
uuid?: string; // 相关凭证的UUID(如果有的话)
|
||||
}
|
||||
|
||||
@@ -85,6 +87,12 @@ const ErrorTypeConfig = {
|
||||
bgColor: "bg-green-50 dark:bg-green-950/30",
|
||||
borderColor: "border-green-200 dark:border-green-800",
|
||||
},
|
||||
reauth: {
|
||||
icon: KeyRound,
|
||||
color: "text-amber-600 dark:text-amber-400",
|
||||
bgColor: "bg-amber-50 dark:bg-amber-950/30",
|
||||
borderColor: "border-amber-200 dark:border-amber-800",
|
||||
},
|
||||
};
|
||||
|
||||
function ErrorItem({
|
||||
|
||||
@@ -303,7 +303,7 @@ export async function sendAgentMessage(
|
||||
* // 处理文本增量
|
||||
* }
|
||||
* });
|
||||
* await sendAgentMessageStream(message, eventName, sessionId);
|
||||
* await sendAgentMessageStream(message, eventName, sessionId, model, undefined, provider);
|
||||
* ```
|
||||
*/
|
||||
export async function sendAgentMessageStream(
|
||||
@@ -312,6 +312,7 @@ export async function sendAgentMessageStream(
|
||||
sessionId?: string,
|
||||
model?: string,
|
||||
images?: ImageInput[],
|
||||
provider?: string,
|
||||
): Promise<void> {
|
||||
return await invoke("native_agent_chat_stream", {
|
||||
message,
|
||||
@@ -319,6 +320,7 @@ export async function sendAgentMessageStream(
|
||||
sessionId,
|
||||
model,
|
||||
images,
|
||||
provider,
|
||||
});
|
||||
}
|
||||
|
||||
|
||||
@@ -563,6 +563,20 @@ export const providerPoolApi = {
|
||||
async migratePrivateConfig(config: unknown): Promise<MigrationResult> {
|
||||
return invoke("migrate_private_config_to_pool", { config });
|
||||
},
|
||||
|
||||
// 获取单个凭证的健康状态
|
||||
// Requirements: 4.4
|
||||
async getCredentialHealth(
|
||||
uuid: string,
|
||||
): Promise<CredentialHealthInfo | null> {
|
||||
return invoke("get_credential_health", { uuid });
|
||||
},
|
||||
|
||||
// 获取所有凭证的健康状态
|
||||
// Requirements: 4.4
|
||||
async getAllCredentialHealth(): Promise<CredentialHealthInfo[]> {
|
||||
return invoke("get_all_credential_health");
|
||||
},
|
||||
};
|
||||
|
||||
// Migration result
|
||||
@@ -616,6 +630,27 @@ export interface KiroFingerprintInfo {
|
||||
auth_method: string;
|
||||
}
|
||||
|
||||
// 凭证健康状态信息
|
||||
// Requirements: 4.4
|
||||
export interface CredentialHealthInfo {
|
||||
/** 凭证 UUID */
|
||||
uuid: string;
|
||||
/** 凭证名称 */
|
||||
name?: string;
|
||||
/** Provider 类型 */
|
||||
provider_type: string;
|
||||
/** 是否健康 */
|
||||
is_healthy: boolean;
|
||||
/** 最后错误信息 */
|
||||
last_error?: string;
|
||||
/** 最后错误时间(RFC3339 格式) */
|
||||
last_error_time?: string;
|
||||
/** 错误次数 */
|
||||
failure_count: number;
|
||||
/** 是否需要重新授权 */
|
||||
requires_reauth: boolean;
|
||||
}
|
||||
|
||||
// Playwright 状态
|
||||
export interface PlaywrightStatus {
|
||||
/** 浏览器是否可用 */
|
||||
|
||||
Reference in New Issue
Block a user