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