Merge pull request #66 from lvhxg6/feature-antigravity-banana

Feature antigravity banana
This commit is contained in:
coso
2026-01-04 12:05:17 +08:00
committed by GitHub
31 changed files with 10835 additions and 198 deletions
+6411
View File
File diff suppressed because it is too large Load Diff
+369
View File
@@ -0,0 +1,369 @@
#!/usr/bin/env python3
"""
OpenAI 兼容图像生成 API 测试脚本
使用 OpenAI Python SDK 测试 Antigravity 图像生成 API。
使用方法:
# 安装依赖
pip install openai
# 运行测试(需要先启动 API Server)
python scripts/test_image_api.py
# 指定自定义 API 地址和密钥
python scripts/test_image_api.py --base-url http://localhost:8999 --api-key your-key
环境变量:
PROXYCAST_BASE_URL: API 服务器地址(默认: http://localhost:8999)
PROXYCAST_API_KEY: API 密钥(默认: pc_LXZbIv3o78WpHuQwqgmwC0U4G0cY5UtQ)
"""
import argparse
import base64
import os
import sys
from datetime import datetime
try:
from openai import OpenAI
except ImportError:
print("错误: 请先安装 openai 库")
print("运行: pip install openai")
sys.exit(1)
def test_image_generation_url(client: OpenAI, prompt: str) -> bool:
"""
测试 URL 响应格式的图像生成
Args:
client: OpenAI 客户端
prompt: 图像生成提示词
Returns:
测试是否通过
"""
print("\n" + "=" * 60)
print("测试 1: URL 响应格式")
print("=" * 60)
print(f"提示词: {prompt}")
try:
response = client.images.generate(
model="dall-e-3", # 会被映射到 gemini-3-pro-image
prompt=prompt,
n=1,
size="1024x1024",
response_format="url"
)
# 验证响应结构
print(f"\n响应时间戳: {response.created}")
print(f"生成图片数量: {len(response.data)}")
if len(response.data) == 0:
print("❌ 错误: 没有生成图片")
return False
image = response.data[0]
# 验证 URL 格式
if image.url:
print(f"URL 长度: {len(image.url)} 字符")
if image.url.startswith("data:image/"):
print("✅ URL 格式正确 (data URL)")
else:
print(f"⚠️ URL 格式: {image.url[:50]}...")
else:
print("❌ 错误: URL 为空")
return False
# 验证 revised_prompt
if image.revised_prompt:
print(f"修订提示词: {image.revised_prompt[:100]}...")
else:
print("ℹ️ 没有修订提示词")
print("\n✅ 测试 1 通过")
return True
except Exception as e:
print(f"\n❌ 测试 1 失败: {e}")
return False
def test_image_generation_b64(client: OpenAI, prompt: str) -> bool:
"""
测试 b64_json 响应格式的图像生成
Args:
client: OpenAI 客户端
prompt: 图像生成提示词
Returns:
测试是否通过
"""
print("\n" + "=" * 60)
print("测试 2: b64_json 响应格式")
print("=" * 60)
print(f"提示词: {prompt}")
try:
response = client.images.generate(
model="gemini-3-pro-image-preview", # 直接使用 Gemini 模型名
prompt=prompt,
n=1,
response_format="b64_json"
)
# 验证响应结构
print(f"\n响应时间戳: {response.created}")
print(f"生成图片数量: {len(response.data)}")
if len(response.data) == 0:
print("❌ 错误: 没有生成图片")
return False
image = response.data[0]
# 验证 b64_json 格式
if image.b64_json:
print(f"Base64 数据长度: {len(image.b64_json)} 字符")
# 尝试解码验证
try:
decoded = base64.b64decode(image.b64_json)
print(f"解码后大小: {len(decoded)} 字节")
# 检查图片魔数
if decoded[:8] == b'\x89PNG\r\n\x1a\n':
print("✅ 图片格式: PNG")
elif decoded[:2] == b'\xff\xd8':
print("✅ 图片格式: JPEG")
elif decoded[:4] == b'GIF8':
print("✅ 图片格式: GIF")
elif decoded[:4] == b'RIFF':
print("✅ 图片格式: WebP")
else:
print(f"⚠️ 未知图片格式: {decoded[:8].hex()}")
except Exception as e:
print(f"⚠️ Base64 解码失败: {e}")
else:
print("❌ 错误: b64_json 为空")
return False
# 验证 revised_prompt
if image.revised_prompt:
print(f"修订提示词: {image.revised_prompt[:100]}...")
else:
print("ℹ️ 没有修订提示词")
print("\n✅ 测试 2 通过")
return True
except Exception as e:
print(f"\n❌ 测试 2 失败: {e}")
return False
def test_error_handling(client: OpenAI) -> bool:
"""
测试错误处理
Args:
client: OpenAI 客户端
Returns:
测试是否通过
"""
print("\n" + "=" * 60)
print("测试 3: 错误处理")
print("=" * 60)
try:
# 测试空提示词
print("测试空提示词...")
try:
response = client.images.generate(
model="dall-e-3",
prompt="", # 空提示词
n=1
)
print("❌ 错误: 应该拒绝空提示词")
return False
except Exception as e:
error_msg = str(e).lower()
if "prompt" in error_msg or "empty" in error_msg or "required" in error_msg:
print(f"✅ 正确拒绝空提示词: {e}")
else:
print(f"⚠️ 收到错误但消息不明确: {e}")
print("\n✅ 测试 3 通过")
return True
except Exception as e:
print(f"\n❌ 测试 3 失败: {e}")
return False
def test_response_structure(client: OpenAI, prompt: str) -> bool:
"""
测试响应结构符合 OpenAI 规范
Args:
client: OpenAI 客户端
prompt: 图像生成提示词
Returns:
测试是否通过
"""
print("\n" + "=" * 60)
print("测试 4: 响应结构验证")
print("=" * 60)
print(f"提示词: {prompt}")
try:
response = client.images.generate(
model="dall-e-3",
prompt=prompt,
n=1,
response_format="url"
)
# 验证 created 字段
if response.created:
print(f"✅ created 字段存在: {response.created}")
# 验证是有效的 Unix 时间戳
if response.created > 0:
dt = datetime.fromtimestamp(response.created)
print(f" 时间: {dt}")
else:
print("❌ created 不是有效时间戳")
return False
else:
print("❌ created 字段缺失")
return False
# 验证 data 字段
if response.data is not None:
print(f"✅ data 字段存在: {len(response.data)} 项")
if len(response.data) > 0:
print("✅ data 数组非空")
else:
print("❌ data 数组为空")
return False
else:
print("❌ data 字段缺失")
return False
# 验证每个图片项
for i, image in enumerate(response.data):
print(f"\n图片 {i + 1}:")
has_url = image.url is not None
has_b64 = image.b64_json is not None
if has_url:
print(f" ✅ url 字段存在")
if has_b64:
print(f" ✅ b64_json 字段存在")
if not has_url and not has_b64:
print(f" ❌ 缺少 url 和 b64_json")
return False
if image.revised_prompt:
print(f" ✅ revised_prompt 字段存在")
else:
print(f" ℹ️ revised_prompt 字段为空")
print("\n✅ 测试 4 通过")
return True
except Exception as e:
print(f"\n❌ 测试 4 失败: {e}")
return False
def main():
parser = argparse.ArgumentParser(
description="测试 OpenAI 兼容图像生成 API"
)
parser.add_argument(
"--base-url",
default=os.environ.get("PROXYCAST_BASE_URL", "http://localhost:8999"),
help="API 服务器地址"
)
parser.add_argument(
"--api-key",
default=os.environ.get("PROXYCAST_API_KEY", "pc_LXZbIv3o78WpHuQwqgmwC0U4G0cY5UtQ"),
help="API 密钥"
)
parser.add_argument(
"--prompt",
default="A cute fluffy cat sitting on a windowsill, looking at the sunset",
help="测试用的图像生成提示词"
)
parser.add_argument(
"--skip-generation",
action="store_true",
help="跳过实际图像生成测试(仅测试错误处理)"
)
args = parser.parse_args()
print("=" * 60)
print("OpenAI 兼容图像生成 API 测试")
print("=" * 60)
print(f"API 地址: {args.base_url}")
print(f"API 密钥: {args.api_key[:8]}...")
print(f"测试提示词: {args.prompt[:50]}...")
# 创建 OpenAI 客户端
client = OpenAI(
base_url=f"{args.base_url}/v1",
api_key=args.api_key
)
results = []
if not args.skip_generation:
# 测试 1: URL 响应格式
results.append(("URL 响应格式", test_image_generation_url(client, args.prompt)))
# 测试 2: b64_json 响应格式
results.append(("b64_json 响应格式", test_image_generation_b64(client, args.prompt)))
# 测试 4: 响应结构验证
results.append(("响应结构验证", test_response_structure(client, args.prompt)))
# 测试 3: 错误处理
results.append(("错误处理", test_error_handling(client)))
# 打印总结
print("\n" + "=" * 60)
print("测试总结")
print("=" * 60)
passed = 0
failed = 0
for name, result in results:
status = "✅ 通过" if result else "❌ 失败"
print(f" {name}: {status}")
if result:
passed += 1
else:
failed += 1
print(f"\n总计: {passed} 通过, {failed} 失败")
if failed > 0:
sys.exit(1)
else:
print("\n🎉 所有测试通过!")
sys.exit(0)
if __name__ == "__main__":
main()
@@ -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"
+5
View File
@@ -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;
}
+21 -1
View File
@@ -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);
+47 -16
View File
@@ -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)
}
+29 -2
View File
@@ -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());
}
+686 -24
View File
@@ -152,8 +152,9 @@ pub struct AntigravityRequestInner {
pub system_instruction: Option<GeminiContent>,
#[serde(skip_serializing_if = "Option::is_none")]
pub generation_config: Option<GeminiGenerationConfig>,
/// 工具定义 - 支持 Gemini 格式(function_declarations)和 Claude 格式(custom + input_schema)
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<GeminiTool>>,
pub tools: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_config: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
@@ -246,6 +247,17 @@ 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"
}
/// 检查是否是 Claude 模型(通过 Antigravity 访问)
/// Claude 模型需要使用不同的工具格式(custom + input_schema)
fn is_claude_model(model: &str) -> bool {
model.starts_with("claude-") || model.contains("claude")
}
// ============================================================================
// 主转换函数
// ============================================================================
@@ -524,6 +536,24 @@ 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
);
eprintln!(
"[ANTIGRAVITY] 图片生成模型 {} 已启用 IMAGE 响应模态",
actual_model
);
} else {
eprintln!(
"[ANTIGRAVITY] 模型 {} 不是图片生成模型,不启用 IMAGE 响应模态",
actual_model
);
}
// 处理 reasoning_effort(思维链配置)
if supports_thinking {
if let Some(ref effort) = request.reasoning_effort {
@@ -559,30 +589,49 @@ pub fn convert_openai_to_antigravity_with_context(
}
// 转换工具定义
let tools: Option<Vec<GeminiTool>> = request.tools.as_ref().and_then(|tools| {
let mut function_declarations: Vec<GeminiFunctionDeclaration> = Vec::new();
// 注意:Antigravity API 统一使用 Gemini 格式(function_declarations)
// Claude 模型在 Antigravity 内部会自动转换
let tools: Option<serde_json::Value> = request.tools.as_ref().and_then(|tools| {
let is_claude = is_claude_model(actual_model);
// Gemini 模型使用 function_declarations + parametersJsonSchema
// Claude 模型使用 function_declarations + inputSchema(注意字段名不同)
let mut function_declarations: Vec<serde_json::Value> = Vec::new();
for t in tools {
match t {
Tool::Function { function } => {
// 转换 parameters -> parametersJsonSchema
let params_schema = function.parameters.as_ref().map(|p| {
let mut schema = clean_parameters(Some(p.clone())).unwrap_or_default();
// 确保有 type 和 properties
if schema.get("type").is_none() {
schema["type"] = serde_json::json!("object");
}
if schema.get("properties").is_none() {
schema["properties"] = serde_json::json!({});
}
schema
});
let params_schema = function
.parameters
.as_ref()
.map(|p| {
let mut schema = clean_parameters(Some(p.clone())).unwrap_or_default();
// 确保有 type 和 properties
if schema.get("type").is_none() {
schema["type"] = serde_json::json!("object");
}
if schema.get("properties").is_none() {
schema["properties"] = serde_json::json!({});
}
schema
})
.unwrap_or_else(|| serde_json::json!({"type": "object", "properties": {}}));
function_declarations.push(GeminiFunctionDeclaration {
name: function.name.clone(),
description: function.description.clone(),
parameters_json_schema: params_schema,
});
if is_claude {
// Claude 模型使用 inputSchema 字段名
function_declarations.push(serde_json::json!({
"name": function.name,
"description": function.description.clone().unwrap_or_default(),
"inputSchema": params_schema
}));
} else {
// Gemini 模型使用 parametersJsonSchema 字段名
function_declarations.push(serde_json::json!({
"name": function.name,
"description": function.description.clone(),
"parametersJsonSchema": params_schema
}));
}
}
Tool::WebSearch | Tool::WebSearch20250305 => {
// web_search 工具不转换
@@ -593,10 +642,9 @@ pub fn convert_openai_to_antigravity_with_context(
if function_declarations.is_empty() {
None
} else {
Some(vec![GeminiTool {
function_declarations: Some(function_declarations),
google_search: None,
}])
Some(serde_json::json!([{
"functionDeclarations": function_declarations
}]))
}
});
@@ -951,3 +999,617 @@ pub fn convert_antigravity_to_openai_response(
response
}
// ============================================================================
// 图像生成 API 转换函数
// ============================================================================
use crate::models::openai::{ImageData, ImageGenerationRequest, ImageGenerationResponse};
/// 图像生成模型名称映射
///
/// 注意:Antigravity API 使用内部模型名称 `gemini-3-pro-image`,
/// 而不是用户友好名称 `gemini-3-pro-image-preview`。
fn image_model_mapping(model: &str) -> &str {
match model {
// OpenAI 兼容模型名 -> Antigravity 内部名称
"dall-e-3" | "dall-e-2" => "gemini-3-pro-image",
// 用户友好名称 -> 内部名称
"gemini-3-pro-image-preview" => "gemini-3-pro-image",
_ => model,
}
}
/// 将 OpenAI 图像生成请求转换为 Antigravity 格式
///
/// # 参数
/// - `request`: OpenAI 图像生成请求
/// - `project_id`: Antigravity 项目 ID
///
/// # 返回
/// Antigravity 格式的请求 JSON
pub fn convert_image_request_to_antigravity(
request: &ImageGenerationRequest,
project_id: &str,
) -> serde_json::Value {
// 模型映射
let actual_model = image_model_mapping(&request.model);
// 构建 Gemini 内容结构
let contents = vec![serde_json::json!({
"role": "user",
"parts": [{"text": request.prompt}]
})];
// 构建生成配置
let generation_config = serde_json::json!({
"temperature": 1.0,
"maxOutputTokens": 8096,
"responseModalities": ["TEXT", "IMAGE"],
"candidateCount": request.n
});
// 构建安全设置
let safety_settings = default_safety_settings();
// 构建完整请求
serde_json::json!({
"project": project_id,
"requestId": format!("img-{}", Uuid::new_v4()),
"request": {
"contents": contents,
"generationConfig": generation_config,
"safetySettings": safety_settings
},
"model": actual_model,
"userAgent": "antigravity"
})
}
/// 将 Antigravity 图像响应转换为 OpenAI 格式
///
/// # 参数
/// - `antigravity_resp`: Antigravity 响应 JSON
/// - `response_format`: 响应格式 ("url" 或 "b64_json")
///
/// # 返回
/// OpenAI 格式的图像生成响应,或错误信息
pub fn convert_antigravity_image_response(
antigravity_resp: &serde_json::Value,
response_format: &str,
) -> Result<ImageGenerationResponse, String> {
let resp = antigravity_resp.get("response").unwrap_or(antigravity_resp);
let mut images = Vec::new();
let mut revised_prompt: Option<String> = None;
if let Some(candidates) = resp.get("candidates").and_then(|c| c.as_array()) {
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 {
// 提取文本作为 revised_prompt
if let Some(text) = part.get("text").and_then(|t| t.as_str()) {
if !text.is_empty() {
revised_prompt = Some(text.to_string());
}
}
// 提取图像数据
if let Some(inline_data) =
part.get("inlineData").or_else(|| part.get("inline_data"))
{
if let (Some(data), Some(mime_type)) = (
inline_data.get("data").and_then(|d| d.as_str()),
inline_data
.get("mimeType")
.or_else(|| inline_data.get("mime_type"))
.and_then(|m| m.as_str()),
) {
let image_data = if response_format == "b64_json" {
ImageData {
b64_json: Some(data.to_string()),
url: None,
revised_prompt: revised_prompt.clone(),
}
} else {
// 构建 data URL
let data_url = format!("data:{};base64,{}", mime_type, data);
ImageData {
b64_json: None,
url: Some(data_url),
revised_prompt: revised_prompt.clone(),
}
};
images.push(image_data);
}
}
}
}
}
}
if images.is_empty() {
return Err("No image generated".to_string());
}
Ok(ImageGenerationResponse {
created: chrono::Utc::now().timestamp(),
data: images,
})
}
// ============================================================================
// 图像生成 API 测试
// ============================================================================
#[cfg(test)]
mod image_tests {
use super::*;
#[test]
fn test_image_model_mapping() {
// OpenAI 兼容模型名映射到内部名称
assert_eq!(image_model_mapping("dall-e-3"), "gemini-3-pro-image");
assert_eq!(image_model_mapping("dall-e-2"), "gemini-3-pro-image");
// 用户友好名称映射到内部名称
assert_eq!(
image_model_mapping("gemini-3-pro-image-preview"),
"gemini-3-pro-image"
);
// 内部名称保持不变
assert_eq!(
image_model_mapping("gemini-3-pro-image"),
"gemini-3-pro-image"
);
// 其他模型名称透传
assert_eq!(image_model_mapping("other-model"), "other-model");
}
#[test]
fn test_convert_image_request_basic() {
let request = ImageGenerationRequest {
prompt: "A cute cat".to_string(),
model: "gemini-3-pro-image-preview".to_string(),
n: 1,
size: None,
response_format: "url".to_string(),
quality: None,
style: None,
user: None,
};
let result = convert_image_request_to_antigravity(&request, "test-project");
// 验证基本结构
assert_eq!(result["project"], "test-project");
// 模型名应该映射为内部名称
assert_eq!(result["model"], "gemini-3-pro-image");
assert!(result["requestId"].as_str().unwrap().starts_with("img-"));
// 验证内容
let contents = result["request"]["contents"].as_array().unwrap();
assert_eq!(contents.len(), 1);
assert_eq!(contents[0]["role"], "user");
assert_eq!(contents[0]["parts"][0]["text"], "A cute cat");
// 验证生成配置
let gen_config = &result["request"]["generationConfig"];
let modalities = gen_config["responseModalities"].as_array().unwrap();
assert!(modalities.contains(&serde_json::json!("TEXT")));
assert!(modalities.contains(&serde_json::json!("IMAGE")));
assert_eq!(gen_config["candidateCount"], 1);
}
#[test]
fn test_convert_image_request_with_n() {
let request = ImageGenerationRequest {
prompt: "A beautiful sunset".to_string(),
model: "dall-e-3".to_string(),
n: 3,
size: Some("1024x1024".to_string()),
response_format: "b64_json".to_string(),
quality: Some("hd".to_string()),
style: Some("vivid".to_string()),
user: None,
};
let result = convert_image_request_to_antigravity(&request, "project-123");
// dall-e-3 应该映射为内部名称
assert_eq!(result["model"], "gemini-3-pro-image");
assert_eq!(result["request"]["generationConfig"]["candidateCount"], 3);
}
#[test]
fn test_convert_antigravity_image_response_b64_json() {
let antigravity_resp = serde_json::json!({
"response": {
"candidates": [{
"content": {
"parts": [
{"text": "Here is your image"},
{
"inlineData": {
"mimeType": "image/png",
"data": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="
}
}
]
}
}]
}
});
let result = convert_antigravity_image_response(&antigravity_resp, "b64_json").unwrap();
assert!(result.created > 0);
assert_eq!(result.data.len(), 1);
assert!(result.data[0].b64_json.is_some());
assert!(result.data[0].url.is_none());
assert_eq!(
result.data[0].revised_prompt,
Some("Here is your image".to_string())
);
}
#[test]
fn test_convert_antigravity_image_response_url() {
let antigravity_resp = serde_json::json!({
"response": {
"candidates": [{
"content": {
"parts": [{
"inlineData": {
"mimeType": "image/jpeg",
"data": "base64data"
}
}]
}
}]
}
});
let result = convert_antigravity_image_response(&antigravity_resp, "url").unwrap();
assert!(result.created > 0);
assert_eq!(result.data.len(), 1);
assert!(result.data[0].b64_json.is_none());
assert!(result.data[0].url.is_some());
assert_eq!(
result.data[0].url,
Some("data:image/jpeg;base64,base64data".to_string())
);
}
#[test]
fn test_convert_antigravity_image_response_no_image() {
let antigravity_resp = serde_json::json!({
"response": {
"candidates": [{
"content": {
"parts": [{"text": "Sorry, I cannot generate that image"}]
}
}]
}
});
let result = convert_antigravity_image_response(&antigravity_resp, "url");
assert!(result.is_err());
assert_eq!(result.unwrap_err(), "No image generated");
}
#[test]
fn test_convert_antigravity_image_response_snake_case() {
// 测试 snake_case 字段名兼容性
let antigravity_resp = serde_json::json!({
"response": {
"candidates": [{
"content": {
"parts": [{
"inline_data": {
"mime_type": "image/png",
"data": "testdata"
}
}]
}
}]
}
});
let result = convert_antigravity_image_response(&antigravity_resp, "b64_json").unwrap();
assert_eq!(result.data.len(), 1);
assert_eq!(result.data[0].b64_json, Some("testdata".to_string()));
}
}
// ============================================================================
// 图像生成 API 属性测试
// ============================================================================
#[cfg(test)]
mod image_property_tests {
use super::*;
use proptest::prelude::*;
// 生成随机提示词
fn arb_prompt() -> impl Strategy<Value = String> {
"[a-zA-Z0-9 .,!?]{1,200}".prop_map(|s| s)
}
// 生成随机模型名称
fn arb_image_model() -> impl Strategy<Value = String> {
prop_oneof![
Just("dall-e-3".to_string()),
Just("dall-e-2".to_string()),
Just("gemini-3-pro-image-preview".to_string()),
Just("gemini-3-pro-image".to_string()),
Just("other-model".to_string()),
]
}
// 生成随机 n 值
fn arb_n() -> impl Strategy<Value = u32> {
1u32..5u32
}
// 生成随机响应格式
fn arb_response_format() -> impl Strategy<Value = String> {
prop_oneof![Just("url".to_string()), Just("b64_json".to_string()),]
}
// 生成随机图像请求
fn arb_image_request() -> impl Strategy<Value = ImageGenerationRequest> {
(
arb_prompt(),
arb_image_model(),
arb_n(),
arb_response_format(),
)
.prop_map(
|(prompt, model, n, response_format)| ImageGenerationRequest {
prompt,
model,
n,
size: None,
response_format,
quality: None,
style: None,
user: None,
},
)
}
// 生成随机 Base64 数据
fn arb_base64_data() -> impl Strategy<Value = String> {
"[a-zA-Z0-9+/]{10,100}={0,2}".prop_map(|s| s)
}
// 生成随机 MIME 类型
fn arb_mime_type() -> impl Strategy<Value = String> {
prop_oneof![
Just("image/png".to_string()),
Just("image/jpeg".to_string()),
Just("image/gif".to_string()),
Just("image/webp".to_string()),
]
}
proptest! {
/// Property 1: Request Conversion Correctness
///
/// *For any* valid OpenAI image generation request with a non-empty prompt,
/// the converted Antigravity request SHALL:
/// - Contain the prompt text in `request.contents[0].parts[0].text`
/// - Have `responseModalities` set to `["TEXT", "IMAGE"]`
/// - Use the correct mapped model name
/// - Include the specified `n` value in `candidateCount`
///
/// **Feature: antigravity-image-api, Property 1: Request Conversion Correctness**
/// **Validates: Requirements 1.2, 1.3, 1.4, 2.1, 2.2**
#[test]
fn prop_request_conversion_correctness(request in arb_image_request()) {
let result = convert_image_request_to_antigravity(&request, "test-project");
// 验证 prompt 正确传递
let contents = result["request"]["contents"].as_array().unwrap();
prop_assert_eq!(contents.len(), 1);
prop_assert_eq!(contents[0]["parts"][0]["text"].as_str().unwrap(), request.prompt.as_str());
// 验证 responseModalities 设置正确
let modalities = result["request"]["generationConfig"]["responseModalities"]
.as_array()
.unwrap();
prop_assert!(modalities.contains(&serde_json::json!("TEXT")));
prop_assert!(modalities.contains(&serde_json::json!("IMAGE")));
// 验证模型映射正确
let expected_model = image_model_mapping(&request.model);
prop_assert_eq!(result["model"].as_str().unwrap(), expected_model);
// 验证 n 值正确传递
prop_assert_eq!(
result["request"]["generationConfig"]["candidateCount"].as_u64().unwrap(),
request.n as u64
);
}
/// Property 2: Response Format Correctness
///
/// *For any* Antigravity response containing image data:
/// - WHEN response_format is "b64_json", the output SHALL have `b64_json` field set and `url` field null
/// - WHEN response_format is "url", the output SHALL have `url` field as a valid data URL and `b64_json` field null
///
/// **Feature: antigravity-image-api, Property 2: Response Format Correctness**
/// **Validates: Requirements 1.6, 1.7, 3.2, 3.3**
#[test]
fn prop_response_format_correctness(
base64_data in arb_base64_data(),
mime_type in arb_mime_type(),
response_format in arb_response_format()
) {
let antigravity_resp = serde_json::json!({
"response": {
"candidates": [{
"content": {
"parts": [{
"inlineData": {
"mimeType": mime_type.clone(),
"data": base64_data.clone()
}
}]
}
}]
}
});
let result = convert_antigravity_image_response(&antigravity_resp, &response_format).unwrap();
prop_assert!(result.data.len() >= 1);
if response_format == "b64_json" {
// b64_json 格式
prop_assert!(result.data[0].b64_json.is_some());
prop_assert!(result.data[0].url.is_none());
prop_assert_eq!(result.data[0].b64_json.as_ref().unwrap(), &base64_data);
} else {
// url 格式
prop_assert!(result.data[0].b64_json.is_none());
prop_assert!(result.data[0].url.is_some());
// 验证 data URL 格式
let url = result.data[0].url.as_ref().unwrap();
let expected_url = format!("data:{};base64,{}", mime_type, base64_data);
prop_assert_eq!(url, &expected_url);
}
}
/// Property 3: OpenAI Response Structure Compliance
///
/// *For any* successful image generation response:
/// - The response SHALL have a `created` field with a valid Unix timestamp (positive integer)
/// - The response SHALL have a `data` field that is a non-empty array
/// - Each item in `data` SHALL have either `url` or `b64_json` field (not both)
///
/// **Feature: antigravity-image-api, Property 3: OpenAI Response Structure Compliance**
/// **Validates: Requirements 3.4, 3.5, 5.3, 5.4**
#[test]
fn prop_openai_response_structure(
base64_data in arb_base64_data(),
mime_type in arb_mime_type(),
response_format in arb_response_format()
) {
let antigravity_resp = serde_json::json!({
"response": {
"candidates": [{
"content": {
"parts": [{
"inlineData": {
"mimeType": mime_type,
"data": base64_data
}
}]
}
}]
}
});
let result = convert_antigravity_image_response(&antigravity_resp, &response_format).unwrap();
// 验证 created 是有效时间戳
prop_assert!(result.created > 0);
// 验证 data 是非空数组
prop_assert!(!result.data.is_empty());
// 验证每个 item 只有 url 或 b64_json 之一
for item in &result.data {
let has_url = item.url.is_some();
let has_b64 = item.b64_json.is_some();
prop_assert!(has_url != has_b64, "Each item should have exactly one of url or b64_json");
}
}
/// Property 4: Model Name Mapping Correctness
///
/// *For any* model name in the request:
/// - "dall-e-3" SHALL map to "gemini-3-pro-image"
/// - "dall-e-2" SHALL map to "gemini-3-pro-image"
/// - "gemini-3-pro-image-preview" SHALL map to "gemini-3-pro-image"
/// - Other model names SHALL pass through unchanged
///
/// **Feature: antigravity-image-api, Property 4: Model Name Mapping Correctness**
/// **Validates: Requirements 2.3, 2.4**
#[test]
fn prop_model_name_mapping(model in arb_image_model()) {
let mapped = image_model_mapping(&model);
match model.as_str() {
"dall-e-3" | "dall-e-2" | "gemini-3-pro-image-preview" => {
prop_assert_eq!(mapped, "gemini-3-pro-image");
}
other => {
prop_assert_eq!(mapped, other);
}
}
}
/// Property 5: Image Data Extraction
///
/// *For any* Antigravity response with `inlineData` containing `data` and `mimeType`:
/// - The converter SHALL successfully extract the Base64 data
/// - The converter SHALL successfully extract the MIME type
/// - If text is present alongside the image, it SHALL be included in `revised_prompt`
///
/// **Feature: antigravity-image-api, Property 5: Image Data Extraction**
/// **Validates: Requirements 3.1, 3.6**
#[test]
fn prop_image_data_extraction(
base64_data in arb_base64_data(),
mime_type in arb_mime_type(),
text in proptest::option::of("[a-zA-Z0-9 ]{0,100}")
) {
let mut parts = vec![];
// 可选的文本部分
if let Some(ref t) = text {
if !t.is_empty() {
parts.push(serde_json::json!({"text": t}));
}
}
// 图像部分
parts.push(serde_json::json!({
"inlineData": {
"mimeType": mime_type.clone(),
"data": base64_data.clone()
}
}));
let antigravity_resp = serde_json::json!({
"response": {
"candidates": [{
"content": {
"parts": parts
}
}]
}
});
let result = convert_antigravity_image_response(&antigravity_resp, "b64_json").unwrap();
// 验证 Base64 数据提取
prop_assert_eq!(result.data[0].b64_json.as_ref().unwrap(), &base64_data);
// 验证 revised_prompt
if let Some(ref t) = text {
if !t.is_empty() {
prop_assert_eq!(result.data[0].revised_prompt.as_ref().unwrap(), t);
}
}
}
}
}
+2
View File
@@ -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,
+79
View File
@@ -189,3 +189,82 @@ pub struct ChatCompletionChunk {
pub model: String,
pub choices: Vec<StreamChoice>,
}
// ============================================================================
// 图像生成 API 数据模型
// ============================================================================
/// OpenAI 图像生成请求
///
/// 兼容 OpenAI Images API,支持通过 Antigravity 生成图像。
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageGenerationRequest {
/// 图像生成提示词
pub prompt: String,
/// 模型名称 (默认: gemini-3-pro-image-preview)
#[serde(default = "default_image_model")]
pub model: String,
/// 生成图像数量 (默认: 1)
#[serde(default = "default_n")]
pub n: u32,
/// 图像尺寸 (可选,Antigravity 可能忽略)
#[serde(skip_serializing_if = "Option::is_none")]
pub size: Option<String>,
/// 响应格式: "url" 或 "b64_json" (默认: "url")
#[serde(default = "default_response_format")]
pub response_format: String,
/// 图像质量 (可选,Antigravity 可能忽略)
#[serde(skip_serializing_if = "Option::is_none")]
pub quality: Option<String>,
/// 图像风格 (可选,Antigravity 可能忽略)
#[serde(skip_serializing_if = "Option::is_none")]
pub style: Option<String>,
/// 用户标识 (可选)
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
}
fn default_image_model() -> String {
"gemini-3-pro-image-preview".to_string()
}
fn default_n() -> u32 {
1
}
fn default_response_format() -> String {
"url".to_string()
}
/// OpenAI 图像生成响应
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageGenerationResponse {
/// 创建时间戳 (Unix epoch seconds)
pub created: i64,
/// 生成的图像数组
pub data: Vec<ImageData>,
}
/// 单个图像数据
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageData {
/// Base64 编码的图像数据 (当 response_format="b64_json")
#[serde(skip_serializing_if = "Option::is_none")]
pub b64_json: Option<String>,
/// 图像 URL (当 response_format="url",返回 data URL)
#[serde(skip_serializing_if = "Option::is_none")]
pub url: Option<String>,
/// 修订后的提示词 (如果 Antigravity 返回了文本)
#[serde(skip_serializing_if = "Option::is_none")]
pub revised_prompt: Option<String>,
}
+4 -1
View File
@@ -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);
}
+20 -4
View File
@@ -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()),
+697 -8
View File
@@ -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(&params)
.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,23 @@ impl StreamingProvider for AntigravityProvider {
match result {
Ok(resp) => {
let status = resp.status();
eprintln!("[ANTIGRAVITY_STREAM] HTTP 响应状态: {}", status);
tracing::info!("[ANTIGRAVITY_STREAM] HTTP 响应状态: {}", status);
if status.is_success() {
tracing::info!("[ANTIGRAVITY_STREAM] 流式响应开始: status={}", status);
eprintln!("[ANTIGRAVITY_STREAM] ✓ 流式响应成功建立");
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] 请求失败 ({}): {} - {}",
eprintln!(
"[ANTIGRAVITY_STREAM] ✗ 请求失败\n Base URL: {}\n Status: {}\n Body: {}",
base_url,
status,
&body[..body.len().min(500)]
);
tracing::error!(
"[ANTIGRAVITY_STREAM] ✗ 请求失败\n Base URL: {}\n Status: {}\n Body: {}",
base_url,
status,
body
@@ -1666,12 +2032,21 @@ impl StreamingProvider for AntigravityProvider {
}
}
Err(e) => {
tracing::warn!("[ANTIGRAVITY_STREAM] 连接失败 ({}): {}", base_url, e);
eprintln!(
"[ANTIGRAVITY_STREAM] ✗ 连接失败\n Base URL: {}\n Error: {}",
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 +2064,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());
}
}
+90 -6
View File
@@ -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();
@@ -0,0 +1,352 @@
//! 图像生成 API 处理器
//!
//! 实现 OpenAI 兼容的 `/v1/images/generations` 端点,
//! 通过 Antigravity Provider 调用 Gemini 图像生成模型。
//!
//! # 功能
//! - 接收 OpenAI 格式的图像生成请求
//! - 转换为 Antigravity/Gemini 格式
//! - 调用 Antigravity Provider
//! - 返回 OpenAI 格式的响应
//!
//! # 需求覆盖
//! - 需求 1.1: 实现 `/v1/images/generations` 端点
//! - 需求 4.1: 验证请求参数
//! - 需求 4.2: 获取 Antigravity 凭证
//! - 需求 4.3: 调用 Antigravity Provider
//! - 需求 4.4: 转换响应格式
use axum::{
extract::State,
http::{HeaderMap, StatusCode},
response::{IntoResponse, Response},
Json,
};
use crate::converter::openai_to_antigravity::{
convert_antigravity_image_response, convert_image_request_to_antigravity,
};
use crate::models::openai::ImageGenerationRequest;
use crate::models::provider_pool_model::CredentialData;
use crate::providers::AntigravityProvider;
use crate::server::handlers::verify_api_key;
use crate::server::AppState;
/// 处理图像生成请求
///
/// # 端点
/// `POST /v1/images/generations`
///
/// # 请求格式
/// ```json
/// {
/// "prompt": "A cute cat",
/// "model": "dall-e-3",
/// "n": 1,
/// "size": "1024x1024",
/// "response_format": "url"
/// }
/// ```
///
/// # 响应格式
/// ```json
/// {
/// "created": 1234567890,
/// "data": [
/// {
/// "url": "data:image/png;base64,...",
/// "revised_prompt": "A cute fluffy cat"
/// }
/// ]
/// }
/// ```
pub async fn handle_image_generation(
State(state): State<AppState>,
headers: HeaderMap,
Json(request): Json<ImageGenerationRequest>,
) -> Response {
// 验证 API Key
if let Err(e) = verify_api_key(&headers, &state.api_key).await {
return e.into_response();
}
// 验证请求参数
if request.prompt.trim().is_empty() {
return (
StatusCode::BAD_REQUEST,
Json(serde_json::json!({
"error": {
"message": "prompt is required and cannot be empty",
"type": "invalid_request_error",
"code": "invalid_prompt"
}
})),
)
.into_response();
}
// 记录请求日志
// 安全截取 prompt,避免 UTF-8 字符边界问题
let prompt_preview: String = request.prompt.chars().take(50).collect();
let prompt_display = if request.prompt.chars().count() > 50 {
format!("{}...", prompt_preview)
} else {
request.prompt.clone()
};
state.logs.write().await.add(
"info",
&format!(
"[IMAGE] 收到图像生成请求: model={}, prompt={}, n={}, response_format={}",
request.model, prompt_display, request.n, request.response_format
),
);
// 获取 Antigravity 凭证
let db = match &state.db {
Some(db) => db,
None => {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": "Database not available",
"type": "server_error"
}
})),
)
.into_response();
}
};
// 从凭证池获取 Antigravity 凭证
let credential = match state
.pool_service
.select_credential(db, "antigravity", None)
{
Ok(Some(cred)) => cred,
Ok(None) => {
state
.logs
.write()
.await
.add("error", "[IMAGE] 没有可用的 Antigravity 凭证");
return (
StatusCode::SERVICE_UNAVAILABLE,
Json(serde_json::json!({
"error": {
"message": "No Antigravity credentials available for image generation",
"type": "server_error",
"code": "no_credentials"
}
})),
)
.into_response();
}
Err(e) => {
state
.logs
.write()
.await
.add("error", &format!("[IMAGE] 获取凭证失败: {}", e));
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": format!("Failed to get credentials: {}", e),
"type": "server_error"
}
})),
)
.into_response();
}
};
// 提取 Antigravity 凭证信息
let (creds_file_path, project_id) = match &credential.credential {
CredentialData::AntigravityOAuth {
creds_file_path,
project_id,
} => (creds_file_path.clone(), project_id.clone()),
_ => {
state
.logs
.write()
.await
.add("error", "[IMAGE] 选中的凭证不是 Antigravity 类型");
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": "Selected credential is not Antigravity type",
"type": "server_error"
}
})),
)
.into_response();
}
};
// 创建 Antigravity Provider
let mut antigravity = AntigravityProvider::new();
if let Err(e) = antigravity
.load_credentials_from_path(&creds_file_path)
.await
{
let _ = state.pool_service.mark_unhealthy(
db,
&credential.uuid,
Some(&format!("Failed to load credentials: {}", e)),
);
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": format!("Failed to load Antigravity credentials: {}", e),
"type": "server_error"
}
})),
)
.into_response();
}
// 验证并刷新 Token
let validation_result = antigravity.validate_token();
if validation_result.needs_refresh() {
tracing::info!("[IMAGE] Token 需要刷新,开始刷新...");
if let Err(refresh_error) = antigravity.refresh_token_with_retry(3).await {
tracing::error!("[IMAGE] Token 刷新失败: {:?}", refresh_error);
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,
"type": "authentication_error"
}
})),
)
.into_response();
}
}
// 设置项目 ID
if let Some(pid) = project_id {
antigravity.project_id = Some(pid);
} else if let Err(e) = antigravity.discover_project().await {
tracing::warn!("[IMAGE] Failed to discover project: {}", e);
}
let proj_id = antigravity.project_id.clone().unwrap_or_default();
// 转换请求为 Antigravity 格式
let antigravity_request = convert_image_request_to_antigravity(&request, &proj_id);
state.logs.write().await.add(
"debug",
&format!(
"[IMAGE] Antigravity 请求: model={}",
antigravity_request["model"].as_str().unwrap_or("unknown")
),
);
// 调用 Antigravity API - 直接使用 call_api 而不是 generate_content
// 因为 generate_content 内部的 to_gemini_response 会丢失嵌套在 response 字段下的数据
let model = antigravity_request["model"]
.as_str()
.unwrap_or("gemini-3-pro-image-preview");
eprintln!("[IMAGE] 调用 Antigravity API: model={}", model);
eprintln!(
"[IMAGE] 请求内容: {}",
serde_json::to_string_pretty(&antigravity_request).unwrap_or_default()
);
match antigravity
.call_api("generateContent", &antigravity_request)
.await
{
Ok(resp) => {
// 调试:打印原始响应
eprintln!(
"[IMAGE] Antigravity 原始响应: {}",
serde_json::to_string_pretty(&resp).unwrap_or_default()
);
state.logs.write().await.add(
"debug",
&format!(
"[IMAGE] Antigravity 原始响应: {}",
serde_json::to_string(&resp).unwrap_or_default()
),
);
// 转换响应为 OpenAI 格式
match convert_antigravity_image_response(&resp, &request.response_format) {
Ok(image_response) => {
// 记录成功
let _ = state
.pool_service
.mark_healthy(db, &credential.uuid, Some(model));
let _ = state.pool_service.record_usage(db, &credential.uuid);
state.logs.write().await.add(
"info",
&format!("[IMAGE] 图像生成成功: {} 张图片", image_response.data.len()),
);
(StatusCode::OK, Json(image_response)).into_response()
}
Err(e) => {
state
.logs
.write()
.await
.add("error", &format!("[IMAGE] 响应转换失败: {}", e));
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": e,
"type": "server_error",
"code": "image_generation_failed"
}
})),
)
.into_response()
}
}
}
Err(e) => {
let _ = state
.pool_service
.mark_unhealthy(db, &credential.uuid, Some(&e.to_string()));
state
.logs
.write()
.await
.add("error", &format!("[IMAGE] Antigravity API 调用失败: {}", e));
(
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({
"error": {
"message": format!("Image generation failed: {}", e),
"type": "server_error",
"code": "api_error"
}
})),
)
.into_response()
}
}
}
+2
View File
@@ -4,6 +4,7 @@
pub mod api;
pub mod credentials_api;
pub mod image_handler;
pub mod kiro_credential;
pub mod management;
pub mod provider_calls;
@@ -11,6 +12,7 @@ pub mod websocket;
pub use api::*;
pub use credentials_api::*;
pub use image_handler::*;
pub use kiro_credential::*;
pub use management::*;
pub use provider_calls::*;
+696 -23
View File
@@ -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![Generated Image]({})", 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![Generated Image]({})", 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()))
}
+29 -8
View File
@@ -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());
}
}
+80 -27
View File
@@ -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()
);
}
// 更新模型映射器
@@ -851,6 +871,11 @@ async fn run_server(
.route("/v1/chat/completions", post(handlers::chat_completions))
.route("/v1/messages", post(handlers::anthropic_messages))
.route("/v1/messages/count_tokens", post(count_tokens))
// 图像生成 API 路由
.route(
"/v1/images/generations",
post(handlers::handle_image_generation),
)
// Gemini 原生协议路由
.route("/v1/gemini/*path", post(gemini_generate_content))
// WebSocket 路由
@@ -1033,18 +1058,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();
}
}
}
+308
View File
@@ -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}`);
+4
View File
@@ -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);
}
});
});
});
+28 -27
View File
@@ -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({
+3 -1
View File
@@ -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,
});
}
+35
View File
@@ -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 {
/** 浏览器是否可用 */