v0.7.0: Add Antigravity Provider, dynamic routes, pre-commit hooks

Features:
- Add Antigravity Provider for Google internal Gemini 3 Pro API
- Add dynamic route support (/:selector/v1/messages)
- Add protocol converter for Antigravity
- Add protocol selector for smart conversion path

Improvements:
- Add pre-commit hooks for local lint/format checks
- Simplify CI to only build check (lint moved to pre-commit)
- Update frontend with Antigravity support

Breaking Changes:
- Route parameter syntax changed from {param} to :param for Axum 0.7
This commit is contained in:
coso
2025-12-15 12:54:23 +08:00
parent f3272b93a2
commit 2e6de831fa
38 changed files with 2844 additions and 755 deletions
-71
View File
@@ -10,80 +10,9 @@ env:
CARGO_TERM_COLOR: always
jobs:
lint:
name: Lint & Format
runs-on: ubuntu-latest
steps:
- name: Checkout
uses: actions/checkout@v4
- name: Setup Node.js
uses: actions/setup-node@v4
with:
node-version: '20'
cache: 'npm'
- name: Setup Rust
uses: dtolnay/rust-toolchain@stable
with:
components: clippy, rustfmt
- name: Setup Rust cache
uses: Swatinem/rust-cache@v2
with:
workspaces: src-tauri
shared-key: "rust-cache-lint"
- name: Install system dependencies (Linux)
run: |
sudo apt-get update
sudo apt-get install -y libwebkit2gtk-4.1-dev libappindicator3-dev librsvg2-dev patchelf libssl-dev
- name: Install frontend dependencies
run: npm ci
- name: Build frontend
run: npm run build
- name: Check formatting
working-directory: src-tauri
run: cargo fmt --all -- --check
- name: Run Clippy
working-directory: src-tauri
run: cargo clippy --all-targets --all-features -- -D warnings
frontend-lint:
name: Frontend Lint
runs-on: ubuntu-latest
steps:
- name: Checkout
uses: actions/checkout@v4
- name: Setup Node.js
uses: actions/setup-node@v4
with:
node-version: '20'
cache: 'npm'
- name: Install dependencies
run: npm ci
- name: TypeScript check
run: npx tsc --noEmit
- name: ESLint check
run: npm run lint
- name: Prettier check
run: npx prettier --check "src/**/*.{ts,tsx,css}"
build-check:
name: Build Check
runs-on: ubuntu-latest
needs: [lint, frontend-lint]
steps:
- name: Checkout
+54
View File
@@ -0,0 +1,54 @@
#!/bin/sh
echo "🔍 Running pre-commit checks..."
# 前端检查
echo "📦 Checking frontend..."
# TypeScript 检查
echo " → TypeScript check..."
npx tsc --noEmit
if [ $? -ne 0 ]; then
echo "❌ TypeScript check failed!"
exit 1
fi
# ESLint 检查
echo " → ESLint check..."
npm run lint
if [ $? -ne 0 ]; then
echo "❌ ESLint check failed!"
exit 1
fi
# Prettier 检查
echo " → Prettier check..."
npx prettier --check "src/**/*.{ts,tsx,css}"
if [ $? -ne 0 ]; then
echo "❌ Prettier check failed! Run 'npm run format' to fix."
exit 1
fi
# Rust 检查
echo "🦀 Checking Rust..."
# Rust 格式检查
echo " → Rust format check..."
cd src-tauri
cargo fmt --all -- --check
if [ $? -ne 0 ]; then
echo "❌ Rust format check failed! Run 'cargo fmt' in src-tauri to fix."
exit 1
fi
# Rust 编译检查 (不使用严格的 clippy)
echo " → Rust build check..."
cargo check --all-targets
if [ $? -ne 0 ]; then
echo "❌ Rust build check failed!"
exit 1
fi
cd ..
echo "✅ All pre-commit checks passed!"
+19 -2
View File
@@ -1,12 +1,12 @@
{
"name": "proxycast",
"version": "0.6.0",
"version": "0.6.1",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"name": "proxycast",
"version": "0.6.0",
"version": "0.6.1",
"dependencies": {
"@radix-ui/react-dialog": "^1.1.2",
"@radix-ui/react-dropdown-menu": "^2.1.2",
@@ -41,6 +41,7 @@
"eslint-plugin-react-hooks": "^5.0.0",
"eslint-plugin-react-refresh": "^0.4.14",
"globals": "^15.12.0",
"husky": "^9.1.7",
"postcss": "^8.4.47",
"prettier": "^3.3.3",
"tailwindcss": "^3.4.14",
@@ -3982,6 +3983,22 @@
"node": ">= 0.4"
}
},
"node_modules/husky": {
"version": "9.1.7",
"resolved": "https://registry.npmjs.org/husky/-/husky-9.1.7.tgz",
"integrity": "sha512-5gs5ytaNjBrh5Ow3zrvdUUY+0VxIuWVL4i9irt6friV+BqdCfmV11CQTWMiBYWHbXhco+J1kHfTOUkePhCDvMA==",
"dev": true,
"license": "MIT",
"bin": {
"husky": "bin.js"
},
"engines": {
"node": ">=18"
},
"funding": {
"url": "https://github.com/sponsors/typicode"
}
},
"node_modules/ignore": {
"version": "7.0.5",
"resolved": "https://registry.npmjs.org/ignore/-/ignore-7.0.5.tgz",
+4 -2
View File
@@ -1,7 +1,7 @@
{
"name": "proxycast",
"private": true,
"version": "0.6.1",
"version": "0.7.0",
"type": "module",
"scripts": {
"dev": "vite",
@@ -9,7 +9,8 @@
"preview": "vite preview",
"tauri": "tauri",
"lint": "eslint src --max-warnings 0",
"format": "prettier --write \"src/**/*.{ts,tsx,css}\""
"format": "prettier --write \"src/**/*.{ts,tsx,css}\"",
"prepare": "husky"
},
"dependencies": {
"@radix-ui/react-dialog": "^1.1.2",
@@ -45,6 +46,7 @@
"eslint-plugin-react-hooks": "^5.0.0",
"eslint-plugin-react-refresh": "^0.4.14",
"globals": "^15.12.0",
"husky": "^9.1.7",
"postcss": "^8.4.47",
"prettier": "^3.3.3",
"tailwindcss": "^3.4.14",
+1 -1
View File
@@ -3153,7 +3153,7 @@ dependencies = [
[[package]]
name = "proxycast"
version = "0.6.1"
version = "0.7.0"
dependencies = [
"anyhow",
"async-stream",
+1 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "proxycast"
version = "0.6.1"
version = "0.7.0"
description = "AI API Proxy Desktop App"
authors = ["you"]
edition = "2021"
+153 -51
View File
@@ -1,18 +1,18 @@
//! Provider Pool Tauri 命令
use crate::database::DbConnection;
use crate::database::dao::provider_pool::ProviderPoolDao;
use crate::database::DbConnection;
use crate::models::provider_pool_model::{
AddCredentialRequest, CredentialData, CredentialDisplay, HealthCheckResult, OAuthStatus,
ProviderCredential, ProviderPoolOverview, UpdateCredentialRequest,
};
use crate::services::provider_pool_service::ProviderPoolService;
use std::sync::Arc;
use tauri::State;
use chrono::Utc;
use std::fs;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use tauri::State;
use uuid::Uuid;
use chrono::Utc;
pub struct ProviderPoolServiceState(pub Arc<ProviderPoolService>);
@@ -35,8 +35,7 @@ fn get_credentials_dir() -> Result<PathBuf, String> {
// 确保目录存在
if !app_data_dir.exists() {
fs::create_dir_all(&app_data_dir)
.map_err(|e| format!("创建凭证存储目录失败: {}", e))?;
fs::create_dir_all(&app_data_dir).map_err(|e| format!("创建凭证存储目录失败: {}", e))?;
}
Ok(app_data_dir)
@@ -45,7 +44,7 @@ fn get_credentials_dir() -> Result<PathBuf, String> {
/// 复制并重命名 OAuth 凭证文件
fn copy_and_rename_credential_file(
source_path: &str,
provider_type: &str
provider_type: &str,
) -> Result<String, String> {
let expanded_source = expand_tilde(source_path);
let source = Path::new(&expanded_source);
@@ -62,7 +61,8 @@ fn copy_and_rename_credential_file(
.unwrap()
.as_secs();
let new_filename = format!("{}_{}_{}_{}.json",
let new_filename = format!(
"{}_{}_{}_{}.json",
provider_type,
&uuid[..8], // 使用 UUID 前8位
timestamp,
@@ -74,8 +74,7 @@ fn copy_and_rename_credential_file(
let target_path = credentials_dir.join(&new_filename);
// 复制文件
fs::copy(&source, &target_path)
.map_err(|e| format!("复制凭证文件失败: {}", e))?;
fs::copy(&source, &target_path).map_err(|e| format!("复制凭证文件失败: {}", e))?;
// 返回新的文件路径
Ok(target_path.to_string_lossy().to_string())
@@ -160,17 +159,26 @@ pub fn update_provider_pool_credential(
// 清理旧文件
cleanup_credential_file(creds_file_path)?;
copy_and_rename_credential_file(&new_file_path, "kiro")?
},
CredentialData::GeminiOAuth { creds_file_path, .. } => {
}
CredentialData::GeminiOAuth {
creds_file_path, ..
} => {
// 清理旧文件
cleanup_credential_file(creds_file_path)?;
copy_and_rename_credential_file(&new_file_path, "gemini")?
},
}
CredentialData::QwenOAuth { creds_file_path } => {
// 清理旧文件
cleanup_credential_file(creds_file_path)?;
copy_and_rename_credential_file(&new_file_path, "qwen")?
},
}
CredentialData::AntigravityOAuth {
creds_file_path, ..
} => {
// 清理旧文件
cleanup_credential_file(creds_file_path)?;
copy_and_rename_credential_file(&new_file_path, "antigravity")?
}
_ => {
return Err("只有 OAuth 凭证支持重新上传文件".to_string());
}
@@ -183,16 +191,28 @@ pub fn update_provider_pool_credential(
match &mut updated_cred.credential {
CredentialData::KiroOAuth { creds_file_path } => {
*creds_file_path = new_stored_path;
},
CredentialData::GeminiOAuth { creds_file_path, project_id } => {
}
CredentialData::GeminiOAuth {
creds_file_path,
project_id,
} => {
*creds_file_path = new_stored_path;
if let Some(new_pid) = request.new_project_id {
*project_id = Some(new_pid);
}
},
}
CredentialData::QwenOAuth { creds_file_path } => {
*creds_file_path = new_stored_path;
},
}
CredentialData::AntigravityOAuth {
creds_file_path,
project_id,
} => {
*creds_file_path = new_stored_path;
if let Some(new_pid) = request.new_project_id {
*project_id = Some(new_pid);
}
}
_ => {}
}
@@ -251,15 +271,9 @@ pub fn toggle_provider_pool_credential(
uuid: String,
is_disabled: bool,
) -> Result<ProviderCredential, String> {
pool_service.0.update_credential(
&db,
&uuid,
None,
Some(is_disabled),
None,
None,
None,
)
pool_service
.0
.update_credential(&db, &uuid, None, Some(is_disabled), None, None, None)
}
/// 重置凭证计数器
@@ -292,7 +306,11 @@ pub async fn check_provider_pool_credential_health(
tracing::info!("[DEBUG] 开始健康检查 for uuid: {}", uuid);
let result = pool_service.0.check_credential_health(&db, &uuid).await;
match &result {
Ok(health) => tracing::info!("[DEBUG] 健康检查完成: success={}, message={:?}", health.success, health.message),
Ok(health) => tracing::info!(
"[DEBUG] 健康检查完成: success={}, message={:?}",
health.success,
health.message
),
Err(err) => tracing::error!("[DEBUG] 健康检查失败: {}", err),
}
result
@@ -379,6 +397,31 @@ pub fn add_qwen_oauth_credential(
)
}
/// 添加 Antigravity OAuth 凭证(通过文件路径)
#[tauri::command]
pub fn add_antigravity_oauth_credential(
db: State<'_, DbConnection>,
pool_service: State<'_, ProviderPoolServiceState>,
creds_file_path: String,
project_id: Option<String>,
name: Option<String>,
) -> Result<ProviderCredential, String> {
// 复制并重命名文件到应用存储目录
let stored_file_path = copy_and_rename_credential_file(&creds_file_path, "antigravity")?;
pool_service.0.add_credential(
&db,
"antigravity",
CredentialData::AntigravityOAuth {
creds_file_path: stored_file_path,
project_id,
},
name,
Some(true),
None,
)
}
/// 添加 OpenAI API Key 凭证
#[tauri::command]
pub fn add_openai_key_credential(
@@ -456,10 +499,22 @@ pub async fn debug_kiro_credentials() -> Result<String, String> {
match provider.load_credentials().await {
Ok(_) => {
result.push_str("✅ 凭证加载成功!\n");
result.push_str(&format!("📄 认证方式: {:?}\n", provider.credentials.auth_method));
result.push_str(&format!("🔑 有 client_id: {}\n", provider.credentials.client_id.is_some()));
result.push_str(&format!("🔒 有 client_secret: {}\n", provider.credentials.client_secret.is_some()));
result.push_str(&format!("🏷️ 有 clientIdHash: {}\n", provider.credentials.client_id_hash.is_some()));
result.push_str(&format!(
"📄 认证方式: {:?}\n",
provider.credentials.auth_method
));
result.push_str(&format!(
"🔑 有 client_id: {}\n",
provider.credentials.client_id.is_some()
));
result.push_str(&format!(
"🔒 有 client_secret: {}\n",
provider.credentials.client_secret.is_some()
));
result.push_str(&format!(
"🏷️ 有 clientIdHash: {}\n",
provider.credentials.client_id_hash.is_some()
));
if let Some(hash) = &provider.credentials.client_id_hash {
result.push_str(&format!("🔗 clientIdHash: {}\n", hash));
@@ -472,14 +527,20 @@ pub async fn debug_kiro_credentials() -> Result<String, String> {
result.push_str(&format!("🌐 刷新端点: {}\n", refresh_url));
if let Some(client_id) = &provider.credentials.client_id {
result.push_str(&format!("🆔 client_id 前缀: {}...\n", &client_id[..std::cmp::min(20, client_id.len())]));
result.push_str(&format!(
"🆔 client_id 前缀: {}...\n",
&client_id[..std::cmp::min(20, client_id.len())]
));
}
result.push_str("\n🚀 尝试刷新 token...\n");
match provider.refresh_token().await {
Ok(token) => {
result.push_str(&format!("✅ Token 刷新成功! Token 长度: {}\n", token.len()));
result.push_str(&format!("🎫 Token 前缀: {}...\n", &token[..std::cmp::min(50, token.len())]));
result.push_str(&format!(
"🎫 Token 前缀: {}...\n",
&token[..std::cmp::min(50, token.len())]
));
}
Err(e) => {
result.push_str(&format!("❌ Token 刷新失败: {}\n", e));
@@ -498,7 +559,6 @@ pub async fn debug_kiro_credentials() -> Result<String, String> {
#[tauri::command]
pub async fn test_user_credentials() -> Result<String, String> {
use crate::providers::kiro::KiroProvider;
use std::path::PathBuf;
let mut result = String::new();
result.push_str("🧪 测试用户上传的凭证文件...\n\n");
@@ -506,7 +566,9 @@ pub async fn test_user_credentials() -> Result<String, String> {
// 测试用户上传的凭证文件路径
let user_creds_path = dirs::home_dir()
.ok_or("无法获取用户主目录".to_string())?
.join("Library/Application Support/proxycast/credentials/kiro_d8da9d58_1765757992_kiro.json");
.join(
"Library/Application Support/proxycast/credentials/kiro_d8da9d58_1765757992_kiro.json",
);
result.push_str(&format!("📂 用户凭证路径: {}\n", user_creds_path.display()));
@@ -531,8 +593,10 @@ pub async fn test_user_credentials() -> Result<String, String> {
result.push_str("✅ JSON 格式有效\n");
// 检查关键字段
let has_access_token = json.get("accessToken").and_then(|v| v.as_str()).is_some();
let has_refresh_token = json.get("refreshToken").and_then(|v| v.as_str()).is_some();
let has_access_token =
json.get("accessToken").and_then(|v| v.as_str()).is_some();
let has_refresh_token =
json.get("refreshToken").and_then(|v| v.as_str()).is_some();
let auth_method = json.get("authMethod").and_then(|v| v.as_str());
let client_id_hash = json.get("clientIdHash").and_then(|v| v.as_str());
let region = json.get("region").and_then(|v| v.as_str());
@@ -550,7 +614,10 @@ pub async fn test_user_credentials() -> Result<String, String> {
.join(".aws/sso/cache")
.join(format!("{}.json", hash));
result.push_str(&format!("\n🔗 检查 clientIdHash 文件: {}\n", hash_file_path.display()));
result.push_str(&format!(
"\n🔗 检查 clientIdHash 文件: {}\n",
hash_file_path.display()
));
if hash_file_path.exists() {
result.push_str("✅ clientIdHash 文件存在\n");
@@ -559,20 +626,37 @@ pub async fn test_user_credentials() -> Result<String, String> {
Ok(hash_content) => {
match serde_json::from_str::<serde_json::Value>(&hash_content) {
Ok(hash_json) => {
let has_client_id = hash_json.get("clientId").and_then(|v| v.as_str()).is_some();
let has_client_secret = hash_json.get("clientSecret").and_then(|v| v.as_str()).is_some();
let has_client_id = hash_json
.get("clientId")
.and_then(|v| v.as_str())
.is_some();
let has_client_secret = hash_json
.get("clientSecret")
.and_then(|v| v.as_str())
.is_some();
result.push_str(&format!("🆔 hash 文件有 clientId: {}\n", has_client_id));
result.push_str(&format!("🔒 hash 文件有 clientSecret: {}\n", has_client_secret));
result.push_str(&format!(
"🆔 hash 文件有 clientId: {}\n",
has_client_id
));
result.push_str(&format!(
"🔒 hash 文件有 clientSecret: {}\n",
has_client_secret
));
if has_client_id && has_client_secret {
result.push_str("✅ IdC 认证配置完整!\n");
} else {
result.push_str("⚠️ IdC 认证配置不完整,将使用 social 认证\n");
result.push_str(
"⚠️ IdC 认证配置不完整,将使用 social 认证\n",
);
}
}
Err(e) => {
result.push_str(&format!("❌ 无法解析 hash 文件 JSON: {}\n", e));
result.push_str(&format!(
"❌ 无法解析 hash 文件 JSON: {}\n",
e
));
}
}
}
@@ -592,12 +676,24 @@ pub async fn test_user_credentials() -> Result<String, String> {
// 设置凭证路径到用户文件
provider.creds_path = Some(user_creds_path.clone());
match provider.load_credentials_from_path(&user_creds_path.to_string_lossy()).await {
match provider
.load_credentials_from_path(&user_creds_path.to_string_lossy())
.await
{
Ok(_) => {
result.push_str("✅ KiroProvider 加载成功!\n");
result.push_str(&format!("📄 最终认证方式: {:?}\n", provider.credentials.auth_method));
result.push_str(&format!("🔑 最终有 client_id: {}\n", provider.credentials.client_id.is_some()));
result.push_str(&format!("🔒 最终有 client_secret: {}\n", provider.credentials.client_secret.is_some()));
result.push_str(&format!(
"📄 最终认证方式: {:?}\n",
provider.credentials.auth_method
));
result.push_str(&format!(
"🔑 最终有 client_id: {}\n",
provider.credentials.client_id.is_some()
));
result.push_str(&format!(
"🔒 最终有 client_secret: {}\n",
provider.credentials.client_secret.is_some()
));
let detected_method = provider.detect_auth_method();
result.push_str(&format!("🎯 检测到的认证方式: {}\n", detected_method));
@@ -608,8 +704,14 @@ pub async fn test_user_credentials() -> Result<String, String> {
result.push_str("\n🚀 尝试刷新 token...\n");
match provider.refresh_token().await {
Ok(token) => {
result.push_str(&format!("✅ Token 刷新成功! Token 长度: {}\n", token.len()));
result.push_str(&format!("🎫 Token 前缀: {}...\n", &token[..std::cmp::min(50, token.len())]));
result.push_str(&format!(
"✅ Token 刷新成功! Token 长度: {}\n",
token.len()
));
result.push_str(&format!(
"🎫 Token 前缀: {}...\n",
&token[..std::cmp::min(50, token.len())]
));
}
Err(e) => {
result.push_str(&format!("❌ Token 刷新失败: {}\n", e));
+8 -11
View File
@@ -65,17 +65,14 @@ pub async fn get_route_curl_examples(
.map_err(|e| e.to_string())?;
// 查找匹配的路由
let route = routes
.iter()
.find(|r| r.selector == selector)
.or_else(|| {
// 如果是默认路由
if selector == "default" {
None // 返回 None 让下面的代码生成默认示例
} else {
None
}
});
let route = routes.iter().find(|r| r.selector == selector).or_else(|| {
// 如果是默认路由
if selector == "default" {
None // 返回 None 让下面的代码生成默认示例
} else {
None
}
});
let api_key = &config.server.api_key;
+6
View File
@@ -1,10 +1,16 @@
pub mod anthropic_to_openai;
pub mod cw_to_openai;
pub mod openai_to_antigravity;
pub mod openai_to_cw;
pub mod protocol_selector;
#[allow(unused_imports)]
pub use anthropic_to_openai::*;
#[allow(unused_imports)]
pub use cw_to_openai::*;
#[allow(unused_imports)]
pub use openai_to_antigravity::*;
#[allow(unused_imports)]
pub use openai_to_cw::*;
#[allow(unused_imports)]
pub use protocol_selector::*;
@@ -0,0 +1,376 @@
//! OpenAI 格式转换为 Antigravity (Gemini) 格式
use crate::models::openai::*;
use serde::{Deserialize, Serialize};
/// Antigravity/Gemini 内容部分
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GeminiPart {
#[serde(skip_serializing_if = "Option::is_none")]
pub text: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub inline_data: Option<InlineData>,
#[serde(skip_serializing_if = "Option::is_none")]
pub function_call: Option<GeminiFunctionCall>,
#[serde(skip_serializing_if = "Option::is_none")]
pub function_response: Option<GeminiFunctionResponse>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct InlineData {
pub mime_type: String,
pub data: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GeminiFunctionCall {
pub name: String,
pub args: serde_json::Value,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GeminiFunctionResponse {
pub name: String,
pub response: serde_json::Value,
}
/// Antigravity/Gemini 内容
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GeminiContent {
pub role: String,
pub parts: Vec<GeminiPart>,
}
/// Antigravity/Gemini 工具定义
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GeminiTool {
pub function_declarations: Vec<GeminiFunctionDeclaration>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GeminiFunctionDeclaration {
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub parameters: Option<serde_json::Value>,
}
/// Antigravity/Gemini 生成配置
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct GeminiGenerationConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_output_tokens: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_p: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_k: Option<i32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stop_sequences: Option<Vec<String>>,
}
/// Antigravity 请求体
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AntigravityRequestBody {
pub contents: Vec<GeminiContent>,
#[serde(skip_serializing_if = "Option::is_none")]
pub system_instruction: Option<GeminiContent>,
#[serde(skip_serializing_if = "Option::is_none")]
pub generation_config: Option<GeminiGenerationConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<GeminiTool>>,
}
/// 将 OpenAI ChatCompletionRequest 转换为 Antigravity 请求体
pub fn convert_openai_to_antigravity(request: &ChatCompletionRequest) -> serde_json::Value {
let mut contents: Vec<GeminiContent> = Vec::new();
let mut system_instruction: Option<GeminiContent> = None;
// 处理消息
for msg in &request.messages {
match msg.role.as_str() {
"system" => {
let text = msg.get_content_text();
if !text.is_empty() {
system_instruction = Some(GeminiContent {
role: "user".to_string(),
parts: vec![GeminiPart {
text: Some(text),
inline_data: None,
function_call: None,
function_response: None,
}],
});
}
}
"user" => {
let parts = convert_user_content(msg);
if !parts.is_empty() {
contents.push(GeminiContent {
role: "user".to_string(),
parts,
});
}
}
"assistant" => {
let parts = convert_assistant_content(msg);
if !parts.is_empty() {
contents.push(GeminiContent {
role: "model".to_string(),
parts,
});
}
}
"tool" => {
// Tool 响应需要合并到 user 消息
let tool_id = msg.tool_call_id.clone().unwrap_or_default();
let content = msg.get_content_text();
// 尝试解析为 JSON,否则包装为对象
let response_value = serde_json::from_str(&content)
.unwrap_or_else(|_| serde_json::json!({ "result": content }));
contents.push(GeminiContent {
role: "user".to_string(),
parts: vec![GeminiPart {
text: None,
inline_data: None,
function_call: None,
function_response: Some(GeminiFunctionResponse {
name: tool_id,
response: response_value,
}),
}],
});
}
_ => {}
}
}
// 构建生成配置
let generation_config = Some(GeminiGenerationConfig {
temperature: request.temperature,
max_output_tokens: request.max_tokens.map(|t| t as i32),
top_p: None,
top_k: None,
stop_sequences: None,
});
// 转换工具
let tools = request.tools.as_ref().map(|tools| {
vec![GeminiTool {
function_declarations: tools
.iter()
.map(|t| GeminiFunctionDeclaration {
name: t.function.name.clone(),
description: t.function.description.clone(),
parameters: t.function.parameters.clone(),
})
.collect(),
}]
});
let body = AntigravityRequestBody {
contents,
system_instruction,
generation_config,
tools,
};
// 包装为 Antigravity 请求格式
serde_json::json!({
"request": body
})
}
/// 转换用户消息内容
fn convert_user_content(msg: &ChatMessage) -> Vec<GeminiPart> {
let mut parts = Vec::new();
match &msg.content {
Some(MessageContent::Text(text)) => {
parts.push(GeminiPart {
text: Some(text.clone()),
inline_data: None,
function_call: None,
function_response: None,
});
}
Some(MessageContent::Parts(content_parts)) => {
for part in content_parts {
match part {
ContentPart::Text { text } => {
parts.push(GeminiPart {
text: Some(text.clone()),
inline_data: None,
function_call: None,
function_response: None,
});
}
ContentPart::ImageUrl { image_url } => {
// 处理 base64 图片
if let Some((mime, data)) = parse_data_url(&image_url.url) {
parts.push(GeminiPart {
text: None,
inline_data: Some(InlineData {
mime_type: mime,
data,
}),
function_call: None,
function_response: None,
});
}
}
}
}
}
None => {}
}
parts
}
/// 转换助手消息内容
fn convert_assistant_content(msg: &ChatMessage) -> Vec<GeminiPart> {
let mut parts = Vec::new();
// 文本内容
let text = msg.get_content_text();
if !text.is_empty() {
parts.push(GeminiPart {
text: Some(text),
inline_data: None,
function_call: None,
function_response: None,
});
}
// 工具调用
if let Some(tool_calls) = &msg.tool_calls {
for tc in tool_calls {
let args: serde_json::Value =
serde_json::from_str(&tc.function.arguments).unwrap_or(serde_json::json!({}));
parts.push(GeminiPart {
text: None,
inline_data: None,
function_call: Some(GeminiFunctionCall {
name: tc.function.name.clone(),
args,
}),
function_response: None,
});
}
}
parts
}
/// 解析 data URL
fn parse_data_url(url: &str) -> Option<(String, String)> {
if url.starts_with("data:") {
let parts: Vec<&str> = url.splitn(2, ',').collect();
if parts.len() == 2 {
let meta = parts[0].strip_prefix("data:")?;
let mime = meta.split(';').next()?.to_string();
let data = parts[1].to_string();
return Some((mime, data));
}
}
None
}
/// 将 Antigravity 响应转换为 OpenAI 格式
pub fn convert_antigravity_to_openai_response(
antigravity_resp: &serde_json::Value,
model: &str,
) -> serde_json::Value {
let mut choices = Vec::new();
if let Some(candidates) = antigravity_resp
.get("candidates")
.and_then(|c| c.as_array())
{
for (i, candidate) in candidates.iter().enumerate() {
let mut content = String::new();
let mut tool_calls: Vec<serde_json::Value> = Vec::new();
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(text) = part.get("text").and_then(|t| t.as_str()) {
content.push_str(text);
}
if let Some(fc) = part.get("functionCall") {
let call_id =
format!("call_{}", uuid::Uuid::new_v4().to_string()[..8].to_string());
tool_calls.push(serde_json::json!({
"id": call_id,
"type": "function",
"function": {
"name": fc.get("name").and_then(|n| n.as_str()).unwrap_or(""),
"arguments": serde_json::to_string(fc.get("args").unwrap_or(&serde_json::json!({}))).unwrap_or_default()
}
}));
}
}
}
let finish_reason = candidate
.get("finishReason")
.and_then(|r| r.as_str())
.map(|r| match r {
"STOP" => "stop",
"MAX_TOKENS" => "length",
"SAFETY" => "content_filter",
"RECITATION" => "content_filter",
_ => "stop",
})
.unwrap_or("stop");
let mut message = serde_json::json!({
"role": "assistant",
"content": if content.is_empty() { serde_json::Value::Null } else { serde_json::Value::String(content) }
});
if !tool_calls.is_empty() {
message["tool_calls"] = serde_json::json!(tool_calls);
}
choices.push(serde_json::json!({
"index": i,
"message": message,
"finish_reason": finish_reason
}));
}
}
// 构建 usage
let usage = antigravity_resp.get("usageMetadata").map(|u| {
serde_json::json!({
"prompt_tokens": u.get("promptTokenCount").and_then(|t| t.as_i64()).unwrap_or(0),
"completion_tokens": u.get("candidatesTokenCount").and_then(|t| t.as_i64()).unwrap_or(0),
"total_tokens": u.get("totalTokenCount").and_then(|t| t.as_i64()).unwrap_or(0)
})
});
let mut response = serde_json::json!({
"id": format!("chatcmpl-{}", uuid::Uuid::new_v4()),
"object": "chat.completion",
"created": chrono::Utc::now().timestamp(),
"model": model,
"choices": choices
});
if let Some(u) = usage {
response["usage"] = u;
}
response
}
@@ -0,0 +1,220 @@
//! 协议选择器 - 智能选择最优协议转换路径
//!
//! 根据源协议、目标 Provider 和请求特征,选择最优的协议转换路径。
use crate::models::provider_pool_model::PoolProviderType;
/// 协议类型
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Protocol {
/// OpenAI Chat Completions API
OpenAI,
/// Anthropic Messages API (Claude)
Anthropic,
/// CodeWhisperer API (Kiro)
CodeWhisperer,
/// Gemini API (Google)
Gemini,
/// Antigravity API (Google Internal)
Antigravity,
}
impl Protocol {
pub fn as_str(&self) -> &'static str {
match self {
Protocol::OpenAI => "openai",
Protocol::Anthropic => "anthropic",
Protocol::CodeWhisperer => "codewhisperer",
Protocol::Gemini => "gemini",
Protocol::Antigravity => "antigravity",
}
}
}
/// 转换路径
#[derive(Debug, Clone)]
pub struct ConversionPath {
/// 源协议
pub source: Protocol,
/// 目标协议
pub target: Protocol,
/// 是否需要转换
pub needs_conversion: bool,
/// 转换复杂度 (0-10, 越低越好)
pub complexity: u8,
}
/// 协议选择器
pub struct ProtocolSelector;
impl ProtocolSelector {
/// 获取 Provider 的原生协议
pub fn native_protocol(provider: PoolProviderType) -> Protocol {
match provider {
PoolProviderType::Kiro => Protocol::CodeWhisperer,
PoolProviderType::Gemini => Protocol::Gemini,
PoolProviderType::Qwen => Protocol::OpenAI,
PoolProviderType::OpenAI => Protocol::OpenAI,
PoolProviderType::Claude => Protocol::Anthropic,
PoolProviderType::Antigravity => Protocol::Antigravity,
}
}
/// 选择最优转换路径
pub fn select_path(
source_protocol: Protocol,
target_provider: PoolProviderType,
) -> ConversionPath {
let target_protocol = Self::native_protocol(target_provider);
// 如果源和目标协议相同,无需转换
if source_protocol == target_protocol {
return ConversionPath {
source: source_protocol,
target: target_protocol,
needs_conversion: false,
complexity: 0,
};
}
// 计算转换复杂度
let complexity = Self::calculate_complexity(source_protocol, target_protocol);
ConversionPath {
source: source_protocol,
target: target_protocol,
needs_conversion: true,
complexity,
}
}
/// 计算转换复杂度
fn calculate_complexity(source: Protocol, target: Protocol) -> u8 {
match (source, target) {
// OpenAI <-> Anthropic: 中等复杂度
(Protocol::OpenAI, Protocol::Anthropic) => 3,
(Protocol::Anthropic, Protocol::OpenAI) => 3,
// OpenAI <-> CodeWhisperer: 较高复杂度(需要处理历史格式)
(Protocol::OpenAI, Protocol::CodeWhisperer) => 5,
(Protocol::CodeWhisperer, Protocol::OpenAI) => 5,
// OpenAI <-> Gemini/Antigravity: 中等复杂度
(Protocol::OpenAI, Protocol::Gemini) => 4,
(Protocol::OpenAI, Protocol::Antigravity) => 4,
(Protocol::Gemini, Protocol::OpenAI) => 4,
(Protocol::Antigravity, Protocol::OpenAI) => 4,
// Anthropic <-> CodeWhisperer: 较高复杂度
(Protocol::Anthropic, Protocol::CodeWhisperer) => 6,
(Protocol::CodeWhisperer, Protocol::Anthropic) => 6,
// Anthropic <-> Gemini/Antigravity: 中等复杂度
(Protocol::Anthropic, Protocol::Gemini) => 5,
(Protocol::Anthropic, Protocol::Antigravity) => 5,
(Protocol::Gemini, Protocol::Anthropic) => 5,
(Protocol::Antigravity, Protocol::Anthropic) => 5,
// Gemini <-> Antigravity: 低复杂度(格式相似)
(Protocol::Gemini, Protocol::Antigravity) => 1,
(Protocol::Antigravity, Protocol::Gemini) => 1,
// 其他情况
_ => 7,
}
}
/// 检查是否支持直接转换
pub fn supports_direct_conversion(source: Protocol, target: Protocol) -> bool {
matches!(
(source, target),
(Protocol::OpenAI, Protocol::Anthropic)
| (Protocol::Anthropic, Protocol::OpenAI)
| (Protocol::OpenAI, Protocol::CodeWhisperer)
| (Protocol::CodeWhisperer, Protocol::OpenAI)
| (Protocol::OpenAI, Protocol::Gemini)
| (Protocol::OpenAI, Protocol::Antigravity)
| (Protocol::Gemini, Protocol::OpenAI)
| (Protocol::Antigravity, Protocol::OpenAI)
| (Protocol::Gemini, Protocol::Antigravity)
| (Protocol::Antigravity, Protocol::Gemini)
)
}
/// 获取推荐的中间协议(用于不支持直接转换的情况)
pub fn intermediate_protocol(source: Protocol, target: Protocol) -> Option<Protocol> {
// 大多数情况下,OpenAI 是最好的中间协议
if !Self::supports_direct_conversion(source, target) {
if source != Protocol::OpenAI && target != Protocol::OpenAI {
return Some(Protocol::OpenAI);
}
}
None
}
/// 获取 Provider 支持的输入协议列表
pub fn supported_input_protocols(_provider: PoolProviderType) -> Vec<Protocol> {
// 所有 Provider 都支持 OpenAI 和 Anthropic 协议输入
vec![Protocol::OpenAI, Protocol::Anthropic]
}
/// 检查请求是否需要特殊处理
pub fn needs_special_handling(
source: Protocol,
target_provider: PoolProviderType,
has_tools: bool,
has_images: bool,
) -> bool {
// 工具调用在某些转换中需要特殊处理
if has_tools {
match (source, target_provider) {
(Protocol::OpenAI, PoolProviderType::Kiro) => true,
(Protocol::Anthropic, PoolProviderType::Kiro) => true,
_ => false,
}
} else if has_images {
// 图片在某些 Provider 中需要特殊处理
match target_provider {
PoolProviderType::Kiro => true, // Kiro 不支持图片
_ => false,
}
} else {
false
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_native_protocol() {
assert_eq!(
ProtocolSelector::native_protocol(PoolProviderType::Kiro),
Protocol::CodeWhisperer
);
assert_eq!(
ProtocolSelector::native_protocol(PoolProviderType::OpenAI),
Protocol::OpenAI
);
assert_eq!(
ProtocolSelector::native_protocol(PoolProviderType::Claude),
Protocol::Anthropic
);
}
#[test]
fn test_select_path_no_conversion() {
let path = ProtocolSelector::select_path(Protocol::OpenAI, PoolProviderType::OpenAI);
assert!(!path.needs_conversion);
assert_eq!(path.complexity, 0);
}
#[test]
fn test_select_path_with_conversion() {
let path = ProtocolSelector::select_path(Protocol::OpenAI, PoolProviderType::Kiro);
assert!(path.needs_conversion);
assert_eq!(path.target, Protocol::CodeWhisperer);
}
}
+25 -13
View File
@@ -48,7 +48,9 @@ impl ProviderPoolDao {
ORDER BY created_at ASC",
)?;
let rows = stmt.query_map([provider_type.to_string()], |row| Self::row_to_credential(row))?;
let rows = stmt.query_map([provider_type.to_string()], |row| {
Self::row_to_credential(row)
})?;
let mut credentials = Vec::new();
for row in rows {
@@ -245,7 +247,12 @@ impl ProviderPoolDao {
"UPDATE provider_pool_credentials SET
usage_count = ?2, last_used = ?3, updated_at = ?4
WHERE uuid = ?1",
params![uuid, usage_count, last_used.timestamp(), Utc::now().timestamp()],
params![
uuid,
usage_count,
last_used.timestamp(),
Utc::now().timestamp()
],
)?;
Ok(())
}
@@ -298,16 +305,11 @@ impl ProviderPoolDao {
let created_at_ts: i64 = row.get(16)?;
let updated_at_ts: i64 = row.get(17)?;
let provider_type: PoolProviderType = provider_type_str
.parse()
.unwrap_or(PoolProviderType::Kiro);
let provider_type: PoolProviderType =
provider_type_str.parse().unwrap_or(PoolProviderType::Kiro);
let credential: CredentialData = serde_json::from_str(&credential_json).map_err(|e| {
rusqlite::Error::FromSqlConversionFailure(
2,
rusqlite::types::Type::Text,
Box::new(e),
)
rusqlite::Error::FromSqlConversionFailure(2, rusqlite::types::Type::Text, Box::new(e))
})?;
let not_supported_models: Vec<String> = not_supported_models_json
@@ -332,8 +334,14 @@ impl ProviderPoolDao {
last_health_check_time: last_health_check_time_ts
.and_then(|ts| Utc.timestamp_opt(ts, 0).single()),
last_health_check_model,
created_at: Utc.timestamp_opt(created_at_ts, 0).single().unwrap_or_default(),
updated_at: Utc.timestamp_opt(updated_at_ts, 0).single().unwrap_or_default(),
created_at: Utc
.timestamp_opt(created_at_ts, 0)
.single()
.unwrap_or_default(),
updated_at: Utc
.timestamp_opt(updated_at_ts, 0)
.single()
.unwrap_or_default(),
cached_token: None, // 从 get_token_cache 单独获取
})
}
@@ -452,7 +460,11 @@ impl ProviderPoolDao {
}
/// 重置 Token 刷新错误计数
pub fn reset_token_refresh_errors(conn: &Connection, uuid: &str) -> Result<(), rusqlite::Error> {
#[allow(dead_code)]
pub fn reset_token_refresh_errors(
conn: &Connection,
uuid: &str,
) -> Result<(), rusqlite::Error> {
conn.execute(
"UPDATE provider_pool_credentials SET
refresh_error_count = 0,
+13 -1
View File
@@ -5,6 +5,7 @@ mod database;
mod logger;
mod models;
mod providers;
mod router;
mod server;
mod services;
@@ -30,6 +31,7 @@ pub enum ProviderType {
#[serde(rename = "openai")]
OpenAI,
Claude,
Antigravity,
}
impl std::fmt::Display for ProviderType {
@@ -40,6 +42,7 @@ impl std::fmt::Display for ProviderType {
ProviderType::Qwen => write!(f, "qwen"),
ProviderType::OpenAI => write!(f, "openai"),
ProviderType::Claude => write!(f, "claude"),
ProviderType::Antigravity => write!(f, "antigravity"),
}
}
}
@@ -54,6 +57,7 @@ impl std::str::FromStr for ProviderType {
"qwen" => Ok(ProviderType::Qwen),
"openai" => Ok(ProviderType::OpenAI),
"claude" => Ok(ProviderType::Claude),
"antigravity" => Ok(ProviderType::Antigravity),
_ => Err(format!("Invalid provider: {s}")),
}
}
@@ -944,6 +948,10 @@ async fn check_api_compatibility(
("qwen3-coder-plus", "basic"),
("qwen3-coder-plus", "tool_call"),
],
ProviderType::Antigravity => vec![
("gemini-3-pro-preview", "basic"),
("gemini-3-pro-preview", "tool_call"),
],
ProviderType::OpenAI | ProviderType::Claude => vec![],
};
@@ -1343,7 +1351,10 @@ pub fn run() {
logs.write()
.await
.add("info", "[启动] 正在自动启动服务器...");
match s.start(logs.clone(), pool_service, token_cache, Some(db)).await {
match s
.start(logs.clone(), pool_service, token_cache, Some(db))
.await
{
Ok(_) => {
let host = s.config.server.host.clone();
let port = s.config.server.port;
@@ -1470,6 +1481,7 @@ pub fn run() {
commands::provider_pool_cmd::add_kiro_oauth_credential,
commands::provider_pool_cmd::add_gemini_oauth_credential,
commands::provider_pool_cmd::add_qwen_oauth_credential,
commands::provider_pool_cmd::add_antigravity_oauth_credential,
commands::provider_pool_cmd::add_openai_key_credential,
commands::provider_pool_cmd::add_claude_key_credential,
commands::provider_pool_cmd::refresh_pool_credential_token,
+1
View File
@@ -19,5 +19,6 @@ pub use mcp_model::McpServer;
pub use openai::*;
pub use prompt_model::Prompt;
pub use provider_model::Provider;
#[allow(unused_imports)]
pub use provider_pool_model::*;
pub use skill_model::{Skill, SkillMetadata, SkillRepo, SkillState, SkillStates};
+34 -5
View File
@@ -17,6 +17,7 @@ pub enum PoolProviderType {
#[serde(rename = "openai")]
OpenAI,
Claude,
Antigravity,
}
impl std::fmt::Display for PoolProviderType {
@@ -27,6 +28,7 @@ impl std::fmt::Display for PoolProviderType {
PoolProviderType::Qwen => write!(f, "qwen"),
PoolProviderType::OpenAI => write!(f, "openai"),
PoolProviderType::Claude => write!(f, "claude"),
PoolProviderType::Antigravity => write!(f, "antigravity"),
}
}
}
@@ -41,6 +43,7 @@ impl std::str::FromStr for PoolProviderType {
"qwen" => Ok(PoolProviderType::Qwen),
"openai" => Ok(PoolProviderType::OpenAI),
"claude" => Ok(PoolProviderType::Claude),
"antigravity" => Ok(PoolProviderType::Antigravity),
_ => Err(format!("Invalid provider type: {s}")),
}
}
@@ -51,17 +54,18 @@ impl std::str::FromStr for PoolProviderType {
#[serde(tag = "type", rename_all = "snake_case")]
pub enum CredentialData {
/// Kiro OAuth 凭证(文件路径)
KiroOAuth {
creds_file_path: String,
},
KiroOAuth { creds_file_path: String },
/// Gemini OAuth 凭证(文件路径)
GeminiOAuth {
creds_file_path: String,
project_id: Option<String>,
},
/// Qwen OAuth 凭证(文件路径)
QwenOAuth {
QwenOAuth { creds_file_path: String },
/// Antigravity OAuth 凭证(文件路径)- Google 内部 Gemini 3 Pro
AntigravityOAuth {
creds_file_path: String,
project_id: Option<String>,
},
/// OpenAI API Key 凭证
OpenAIKey {
@@ -90,6 +94,11 @@ impl CredentialData {
CredentialData::QwenOAuth { creds_file_path } => {
format!("Qwen OAuth: {}", mask_path(creds_file_path))
}
CredentialData::AntigravityOAuth {
creds_file_path, ..
} => {
format!("Antigravity OAuth: {}", mask_path(creds_file_path))
}
CredentialData::OpenAIKey { api_key, .. } => {
format!("OpenAI: {}", mask_key(api_key))
}
@@ -105,6 +114,7 @@ impl CredentialData {
CredentialData::KiroOAuth { .. } => PoolProviderType::Kiro,
CredentialData::GeminiOAuth { .. } => PoolProviderType::Gemini,
CredentialData::QwenOAuth { .. } => PoolProviderType::Qwen,
CredentialData::AntigravityOAuth { .. } => PoolProviderType::Antigravity,
CredentialData::OpenAIKey { .. } => PoolProviderType::OpenAI,
CredentialData::ClaudeKey { .. } => PoolProviderType::Claude,
}
@@ -370,6 +380,7 @@ pub fn get_default_check_model(provider_type: PoolProviderType) -> &'static str
PoolProviderType::Qwen => "qwen3-coder-flash",
PoolProviderType::OpenAI => "gpt-3.5-turbo",
PoolProviderType::Claude => "claude-3-5-haiku-latest",
PoolProviderType::Antigravity => "gemini-3-pro-preview",
}
}
@@ -383,14 +394,20 @@ pub struct CredentialDisplay {
pub display_credential: String,
pub is_healthy: bool,
pub is_disabled: bool,
pub check_health: bool,
pub check_model_name: Option<String>,
pub not_supported_models: Vec<String>,
pub usage_count: u64,
pub error_count: u32,
pub last_used: Option<String>,
pub last_error_time: Option<String>,
pub last_error_message: Option<String>,
pub last_health_check_time: Option<String>,
pub last_health_check_model: Option<String>,
pub oauth_status: Option<OAuthStatus>,
pub token_cache_status: Option<TokenCacheStatus>,
pub created_at: String,
pub updated_at: String,
}
/// 获取凭证类型字符串
@@ -399,6 +416,7 @@ fn get_credential_type(cred: &CredentialData) -> String {
CredentialData::KiroOAuth { .. } => "kiro_oauth".to_string(),
CredentialData::GeminiOAuth { .. } => "gemini_oauth".to_string(),
CredentialData::QwenOAuth { .. } => "qwen_oauth".to_string(),
CredentialData::AntigravityOAuth { .. } => "antigravity_oauth".to_string(),
CredentialData::OpenAIKey { .. } => "openai_key".to_string(),
CredentialData::ClaudeKey { .. } => "claude_key".to_string(),
}
@@ -408,8 +426,13 @@ fn get_credential_type(cred: &CredentialData) -> String {
pub fn get_oauth_creds_path(cred: &CredentialData) -> Option<String> {
match cred {
CredentialData::KiroOAuth { creds_file_path } => Some(creds_file_path.clone()),
CredentialData::GeminiOAuth { creds_file_path, .. } => Some(creds_file_path.clone()),
CredentialData::GeminiOAuth {
creds_file_path, ..
} => Some(creds_file_path.clone()),
CredentialData::QwenOAuth { creds_file_path } => Some(creds_file_path.clone()),
CredentialData::AntigravityOAuth {
creds_file_path, ..
} => Some(creds_file_path.clone()),
_ => None,
}
}
@@ -435,14 +458,20 @@ impl From<&ProviderCredential> for CredentialDisplay {
display_credential: cred.credential.display_name(),
is_healthy: cred.is_healthy,
is_disabled: cred.is_disabled,
check_health: cred.check_health,
check_model_name: cred.check_model_name.clone(),
not_supported_models: cred.not_supported_models.clone(),
usage_count: cred.usage_count,
error_count: cred.error_count,
last_used: cred.last_used.map(|t| t.to_rfc3339()),
last_error_time: cred.last_error_time.map(|t| t.to_rfc3339()),
last_error_message: cred.last_error_message.clone(),
last_health_check_time: cred.last_health_check_time.map(|t| t.to_rfc3339()),
last_health_check_model: cred.last_health_check_model.clone(),
oauth_status: None, // 需要单独调用获取
token_cache_status,
created_at: cred.created_at.to_rfc3339(),
updated_at: cred.updated_at.to_rfc3339(),
}
}
}
+17 -5
View File
@@ -85,7 +85,7 @@ impl RouteInfo {
let mut examples = Vec::new();
for endpoint in &self.endpoints {
let (model, body) = match endpoint.protocol.as_str() {
let (_model, body) = match endpoint.protocol.as_str() {
"claude" => {
let model = match self.provider_type.as_str() {
"kiro" | "claude" => "claude-sonnet-4-5",
@@ -94,11 +94,17 @@ impl RouteInfo {
"openai" => "gpt-4",
_ => "claude-sonnet-4-5",
};
(model, format!(r#"{{
(
model,
format!(
r#"{{
"model": "{}",
"max_tokens": 1024,
"messages": [{{"role": "user", "content": "Hello!"}}]
}}"#, model))
}}"#,
model
),
)
}
"openai" => {
let model = match self.provider_type.as_str() {
@@ -108,10 +114,16 @@ impl RouteInfo {
"openai" => "gpt-4",
_ => "claude-sonnet-4-5",
};
(model, format!(r#"{{
(
model,
format!(
r#"{{
"model": "{}",
"messages": [{{"role": "user", "content": "Hello!"}}]
}}"#, model))
}}"#,
model
),
)
}
_ => continue,
};
+1
View File
@@ -50,6 +50,7 @@ impl Default for SkillRepo {
}
}
#[allow(dead_code)]
impl SkillRepo {
pub fn new(owner: String, name: String, branch: String) -> Self {
Self {
+478
View File
@@ -0,0 +1,478 @@
//! Antigravity Provider - Google 内部 Gemini 3 Pro 接口
//!
//! 支持 Gemini 3 Pro 等高级模型,通过 Google 内部 API 访问。
use reqwest::Client;
use serde::{Deserialize, Serialize};
use std::error::Error;
use std::path::PathBuf;
use uuid::Uuid;
// Constants
const ANTIGRAVITY_BASE_URL_DAILY: &str = "https://daily-cloudcode-pa.sandbox.googleapis.com";
const ANTIGRAVITY_BASE_URL_AUTOPUSH: &str = "https://autopush-cloudcode-pa.sandbox.googleapis.com";
const ANTIGRAVITY_API_VERSION: &str = "v1internal";
const CREDENTIALS_DIR: &str = ".antigravity";
const CREDENTIALS_FILE: &str = "oauth_creds.json";
// OAuth credentials - 与 Antigravity CLI 相同
const OAUTH_CLIENT_ID: &str =
"1071006060591-tmhssin2h21lcre235vtolojh4g403ep.apps.googleusercontent.com";
const OAUTH_CLIENT_SECRET: &str = "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";
// Token 刷新提前量(秒)
const REFRESH_SKEW: i64 = 3000;
/// Antigravity 支持的模型列表
pub const ANTIGRAVITY_MODELS: &[&str] = &[
"gemini-3-pro-preview",
"gemini-3-pro-image-preview",
"gemini-2.5-computer-use-preview-10-2025",
"gemini-claude-sonnet-4-5",
"gemini-claude-sonnet-4-5-thinking",
];
/// 模型别名映射(用户友好名称 -> 内部名称)
fn alias_to_model_name(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,
}
}
/// 内部模型名称 -> 用户友好名称
#[allow(dead_code)]
fn model_name_to_alias(model: &str) -> &str {
match model {
"rev19-uic3-1p" => "gemini-2.5-computer-use-preview-10-2025",
"gemini-3-pro-image" => "gemini-3-pro-image-preview",
"gemini-3-pro-high" => "gemini-3-pro-preview",
"claude-sonnet-4-5" => "gemini-claude-sonnet-4-5",
"claude-sonnet-4-5-thinking" => "gemini-claude-sonnet-4-5-thinking",
_ => model,
}
}
/// 生成随机请求 ID
fn generate_request_id() -> String {
format!("agent-{}", Uuid::new_v4())
}
/// 生成随机会话 ID
fn generate_session_id() -> String {
// 使用 UUID 的一部分作为随机数
let uuid = Uuid::new_v4();
let bytes = uuid.as_bytes();
let n: u64 = u64::from_le_bytes([
bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7],
]) % 9_000_000_000_000_000_000;
format!("-{}", n)
}
/// 生成随机项目 ID
fn generate_project_id() -> String {
let adjectives = ["useful", "bright", "swift", "calm", "bold"];
let nouns = ["fuze", "wave", "spark", "flow", "core"];
let uuid = Uuid::new_v4();
let bytes = uuid.as_bytes();
let adj = adjectives[(bytes[0] as usize) % adjectives.len()];
let noun = nouns[(bytes[1] as usize) % nouns.len()];
let random_part: String = uuid.to_string()[..5].to_lowercase();
format!("{}-{}-{}", adj, noun, random_part)
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AntigravityCredentials {
pub access_token: Option<String>,
pub refresh_token: Option<String>,
pub token_type: Option<String>,
pub expiry_date: Option<i64>,
pub scope: Option<String>,
}
impl Default for AntigravityCredentials {
fn default() -> Self {
Self {
access_token: None,
refresh_token: None,
token_type: Some("Bearer".to_string()),
expiry_date: None,
scope: None,
}
}
}
/// Antigravity Provider
pub struct AntigravityProvider {
pub credentials: AntigravityCredentials,
pub project_id: Option<String>,
pub client: Client,
pub base_urls: Vec<String>,
pub available_models: Vec<String>,
}
impl Default for AntigravityProvider {
fn default() -> Self {
Self {
credentials: AntigravityCredentials::default(),
project_id: None,
client: Client::builder()
.timeout(std::time::Duration::from_secs(120))
.build()
.unwrap_or_else(|_| Client::new()),
base_urls: vec![
ANTIGRAVITY_BASE_URL_DAILY.to_string(),
ANTIGRAVITY_BASE_URL_AUTOPUSH.to_string(),
],
available_models: ANTIGRAVITY_MODELS.iter().map(|s| s.to_string()).collect(),
}
}
}
impl AntigravityProvider {
pub fn new() -> Self {
Self::default()
}
pub fn default_creds_path() -> PathBuf {
dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(CREDENTIALS_DIR)
.join(CREDENTIALS_FILE)
}
pub async fn load_credentials(&mut self) -> Result<(), Box<dyn Error + Send + Sync>> {
let path = Self::default_creds_path();
if tokio::fs::try_exists(&path).await.unwrap_or(false) {
let content = tokio::fs::read_to_string(&path).await?;
let creds: AntigravityCredentials = serde_json::from_str(&content)?;
self.credentials = creds;
}
Ok(())
}
pub async fn load_credentials_from_path(
&mut self,
path: &str,
) -> Result<(), Box<dyn Error + Send + Sync>> {
let content = tokio::fs::read_to_string(path).await?;
let creds: AntigravityCredentials = serde_json::from_str(&content)?;
self.credentials = creds;
Ok(())
}
pub async fn save_credentials(&self) -> Result<(), Box<dyn Error + Send + Sync>> {
let path = Self::default_creds_path();
if let Some(parent) = path.parent() {
tokio::fs::create_dir_all(parent).await?;
}
let content = serde_json::to_string_pretty(&self.credentials)?;
tokio::fs::write(&path, content).await?;
Ok(())
}
pub fn is_token_valid(&self) -> bool {
if self.credentials.access_token.is_none() {
return false;
}
if let Some(expiry) = self.credentials.expiry_date {
let now = chrono::Utc::now().timestamp_millis();
// Token valid if more than 5 minutes until expiry
return expiry > now + 300_000;
}
true
}
pub fn is_token_expiring_soon(&self) -> bool {
if let Some(expiry) = self.credentials.expiry_date {
let now = chrono::Utc::now().timestamp_millis();
let refresh_skew_ms = REFRESH_SKEW * 1000;
return expiry <= now + refresh_skew_ms;
}
true
}
pub async fn refresh_token(&mut self) -> Result<String, Box<dyn Error + Send + Sync>> {
let refresh_token = self
.credentials
.refresh_token
.as_ref()
.ok_or("No refresh token available")?;
let params = [
("client_id", OAUTH_CLIENT_ID),
("client_secret", OAUTH_CLIENT_SECRET),
("refresh_token", refresh_token.as_str()),
("grant_type", "refresh_token"),
];
let resp = self
.client
.post("https://oauth2.googleapis.com/token")
.form(&params)
.send()
.await?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
return Err(format!("Token refresh failed: {status} - {body}").into());
}
let data: serde_json::Value = resp.json().await?;
let new_token = data["access_token"]
.as_str()
.ok_or("No access token in response")?;
self.credentials.access_token = Some(new_token.to_string());
if let Some(expires_in) = data["expires_in"].as_i64() {
self.credentials.expiry_date =
Some(chrono::Utc::now().timestamp_millis() + expires_in * 1000);
}
// 如果返回了新的 refresh_token,也更新它
if let Some(new_refresh) = data["refresh_token"].as_str() {
self.credentials.refresh_token = Some(new_refresh.to_string());
}
// Save refreshed credentials
self.save_credentials().await?;
Ok(new_token.to_string())
}
/// 调用 Antigravity API
async fn call_api_internal(
&self,
base_url: &str,
method: &str,
body: &serde_json::Value,
) -> Result<serde_json::Value, Box<dyn Error + Send + Sync>> {
let token = self
.credentials
.access_token
.as_ref()
.ok_or("No access token")?;
let url = format!("{}/{ANTIGRAVITY_API_VERSION}:{method}", base_url);
let resp = self
.client
.post(&url)
.header("Authorization", format!("Bearer {token}"))
.header("Content-Type", "application/json")
.header("User-Agent", "antigravity/1.11.5 windows/amd64")
.json(body)
.send()
.await?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
return Err(format!("API call failed: {status} - {body}").into());
}
let data: serde_json::Value = resp.json().await?;
Ok(data)
}
/// 调用 API,支持多环境降级
pub async fn call_api(
&self,
method: &str,
body: &serde_json::Value,
) -> Result<serde_json::Value, Box<dyn Error + Send + Sync>> {
let mut last_error: Option<Box<dyn Error + Send + Sync>> = None;
for base_url in &self.base_urls {
match self.call_api_internal(base_url, method, body).await {
Ok(data) => return Ok(data),
Err(e) => {
tracing::warn!("[Antigravity] Failed on {}: {}", base_url, e);
last_error = Some(e);
}
}
}
Err(last_error.unwrap_or_else(|| "All Antigravity base URLs failed".into()))
}
/// 发现项目 ID
pub async fn discover_project(&mut self) -> Result<String, Box<dyn Error + Send + Sync>> {
if let Some(ref project_id) = self.project_id {
return Ok(project_id.clone());
}
let body = serde_json::json!({
"cloudaicompanionProject": "",
"metadata": {
"ideType": "IDE_UNSPECIFIED",
"platform": "PLATFORM_UNSPECIFIED",
"pluginType": "GEMINI",
"duetProject": ""
}
});
let resp = self.call_api("loadCodeAssist", &body).await?;
if let Some(project) = resp["cloudaicompanionProject"].as_str() {
if !project.is_empty() {
self.project_id = Some(project.to_string());
return Ok(project.to_string());
}
}
// Need to onboard
let onboard_body = serde_json::json!({
"tierId": "free-tier",
"cloudaicompanionProject": "",
"metadata": {
"ideType": "IDE_UNSPECIFIED",
"platform": "PLATFORM_UNSPECIFIED",
"pluginType": "GEMINI",
"duetProject": ""
}
});
let mut lro_resp = self.call_api("onboardUser", &onboard_body).await?;
// Poll until done
for _ in 0..30 {
if lro_resp["done"].as_bool().unwrap_or(false) {
break;
}
tokio::time::sleep(tokio::time::Duration::from_secs(2)).await;
lro_resp = self.call_api("onboardUser", &onboard_body).await?;
}
let project_id = lro_resp["response"]["cloudaicompanionProject"]["id"]
.as_str()
.unwrap_or("")
.to_string();
if project_id.is_empty() {
// 生成一个随机项目 ID 作为后备
let fallback = generate_project_id();
self.project_id = Some(fallback.clone());
return Ok(fallback);
}
self.project_id = Some(project_id.clone());
Ok(project_id)
}
/// 获取可用模型列表
pub async fn fetch_available_models(
&mut self,
) -> Result<Vec<String>, Box<dyn Error + Send + Sync>> {
let body = serde_json::json!({});
match self.call_api("fetchAvailableModels", &body).await {
Ok(resp) => {
if let Some(models) = resp["models"].as_object() {
self.available_models = models
.keys()
.filter_map(|name| {
let alias = model_name_to_alias(name);
if alias.is_empty() {
None
} else {
Some(alias.to_string())
}
})
.collect();
}
}
Err(e) => {
tracing::warn!(
"[Antigravity] Failed to fetch models: {}, using defaults",
e
);
}
}
Ok(self.available_models.clone())
}
/// 生成内容(非流式)
pub async fn generate_content(
&self,
model: &str,
request_body: &serde_json::Value,
) -> Result<serde_json::Value, Box<dyn Error + Send + Sync>> {
let project_id = self.project_id.clone().unwrap_or_else(generate_project_id);
let actual_model = alias_to_model_name(model);
let payload = self.build_antigravity_request(actual_model, &project_id, request_body);
let resp = self.call_api("generateContent", &payload).await?;
// 转换为 Gemini 格式响应
Ok(self.to_gemini_response(&resp))
}
/// 构建 Antigravity 请求
fn build_antigravity_request(
&self,
model: &str,
project_id: &str,
request_body: &serde_json::Value,
) -> serde_json::Value {
let mut payload = request_body.clone();
// 设置基本字段
payload["model"] = serde_json::json!(model);
payload["userAgent"] = serde_json::json!("antigravity");
payload["project"] = serde_json::json!(project_id);
payload["requestId"] = serde_json::json!(generate_request_id());
// 确保 request 对象存在
if payload.get("request").is_none() {
payload["request"] = serde_json::json!({});
}
// 设置会话 ID
payload["request"]["sessionId"] = serde_json::json!(generate_session_id());
// 删除安全设置
if let Some(request) = payload.get_mut("request") {
if let Some(obj) = request.as_object_mut() {
obj.remove("safetySettings");
}
}
payload
}
/// 转换为 Gemini 格式响应
fn to_gemini_response(&self, antigravity_resp: &serde_json::Value) -> serde_json::Value {
let mut response = serde_json::json!({});
if let Some(candidates) = antigravity_resp.get("candidates") {
response["candidates"] = candidates.clone();
}
if let Some(usage) = antigravity_resp.get("usageMetadata") {
response["usageMetadata"] = usage.clone();
}
if let Some(feedback) = antigravity_resp.get("promptFeedback") {
response["promptFeedback"] = feedback.clone();
}
response
}
/// 检查模型是否支持
pub fn supports_model(&self, model: &str) -> bool {
self.available_models.iter().any(|m| m == model)
}
}
+60 -24
View File
@@ -120,8 +120,14 @@ impl KiroProvider {
// 如果有 clientIdHash,尝试加载对应的 client_id 和 client_secret
if let Some(hash) = &merged.client_id_hash {
let hash_file_path = dir.join(format!("{}.json", hash));
tracing::info!("[KIRO] 检查 clientIdHash 文件: {}", hash_file_path.display());
if tokio::fs::try_exists(&hash_file_path).await.unwrap_or(false) {
tracing::info!(
"[KIRO] 检查 clientIdHash 文件: {}",
hash_file_path.display()
);
if tokio::fs::try_exists(&hash_file_path)
.await
.unwrap_or(false)
{
if let Ok(content) = tokio::fs::read_to_string(&hash_file_path).await {
if let Ok(creds) = serde_json::from_str::<KiroCredentials>(&content) {
tracing::info!(
@@ -132,13 +138,23 @@ impl KiroProvider {
);
merge_credentials(&mut merged, &creds);
} else {
tracing::error!("[KIRO] 无法解析 clientIdHash 文件: {}", hash_file_path.display());
tracing::error!(
"[KIRO] 无法解析 clientIdHash 文件: {}",
hash_file_path.display()
);
}
} else {
tracing::error!("[KIRO] 无法读取 clientIdHash 文件: {}", hash_file_path.display());
tracing::error!(
"[KIRO] 无法读取 clientIdHash 文件: {}",
hash_file_path.display()
);
}
} else {
tracing::warn!("[KIRO] clientIdHash {} 指向的文件不存在: {}", hash, hash_file_path.display());
tracing::warn!(
"[KIRO] clientIdHash {} 指向的文件不存在: {}",
hash,
hash_file_path.display()
);
}
} else {
tracing::info!("[KIRO] 没有 clientIdHash 字段");
@@ -181,7 +197,10 @@ impl KiroProvider {
// 加载完成后,智能检测并更新认证方式(如果需要)
let detected_auth_method = self.detect_auth_method();
if self.credentials.auth_method.as_deref().unwrap_or("social") != detected_auth_method {
tracing::info!("[KIRO] 加载后检测到需要调整认证方式为: {}", detected_auth_method);
tracing::info!(
"[KIRO] 加载后检测到需要调整认证方式为: {}",
detected_auth_method
);
self.set_auth_method(&detected_auth_method);
}
@@ -229,7 +248,10 @@ impl KiroProvider {
hash_file_path.display()
);
if tokio::fs::try_exists(&hash_file_path).await.unwrap_or(false) {
if tokio::fs::try_exists(&hash_file_path)
.await
.unwrap_or(false)
{
if let Ok(content) = tokio::fs::read_to_string(&hash_file_path).await {
// 使用 serde_json::Value 来更灵活地解析,因为 hash 文件可能包含额外字段
if let Ok(json_value) = serde_json::from_str::<serde_json::Value>(&content) {
@@ -285,7 +307,8 @@ impl KiroProvider {
let mut entries = tokio::fs::read_dir(dir).await?;
while let Some(entry) = entries.next_entry().await? {
let file_path = entry.path();
if file_path.extension().map(|e| e == "json").unwrap_or(false) && file_path != path
if file_path.extension().map(|e| e == "json").unwrap_or(false)
&& file_path != path
{
if let Ok(content) = tokio::fs::read_to_string(&file_path).await {
if let Ok(creds) = serde_json::from_str::<KiroCredentials>(&content) {
@@ -318,7 +341,10 @@ impl KiroProvider {
// 加载完成后,智能检测并更新认证方式(如果需要)
let detected_auth_method = self.detect_auth_method();
if self.credentials.auth_method.as_deref().unwrap_or("social") != detected_auth_method {
tracing::info!("[KIRO] 从路径加载后检测到需要调整认证方式为: {}", detected_auth_method);
tracing::info!(
"[KIRO] 从路径加载后检测到需要调整认证方式为: {}",
detected_auth_method
);
self.set_auth_method(&detected_auth_method);
}
@@ -354,13 +380,10 @@ impl KiroProvider {
/// 从凭证文件中提取 region 信息的静态方法,供健康检查服务使用
pub fn extract_region_from_creds(creds_content: &str) -> Result<String, String> {
let creds: serde_json::Value = serde_json::from_str(creds_content)
.map_err(|e| format!("解析凭证失败: {}", e))?;
let creds: serde_json::Value =
serde_json::from_str(creds_content).map_err(|e| format!("解析凭证失败: {}", e))?;
let region = creds["region"]
.as_str()
.unwrap_or("us-east-1")
.to_string();
let region = creds["region"].as_str().unwrap_or("us-east-1").to_string();
Ok(region)
}
@@ -444,7 +467,8 @@ impl KiroProvider {
self.validate_refresh_token()?;
tracing::info!("[KIRO] 开始 Token 刷新流程");
tracing::info!("[KIRO] 当前凭证状态: has_client_id={}, has_client_secret={}, auth_method={:?}",
tracing::info!(
"[KIRO] 当前凭证状态: has_client_id={}, has_client_secret={}, auth_method={:?}",
self.credentials.client_id.is_some(),
self.credentials.client_secret.is_some(),
self.credentials.auth_method
@@ -465,16 +489,25 @@ impl KiroProvider {
// 如果检测到的方式与配置中的不同,更新配置
let current_auth = self.credentials.auth_method.as_deref().unwrap_or("social");
if current_auth != detected_auth_method {
tracing::info!("[KIRO] 认证方式从 {} 切换到 {}", current_auth, detected_auth_method);
tracing::info!(
"[KIRO] 认证方式从 {} 切换到 {}",
current_auth,
detected_auth_method
);
self.set_auth_method(&detected_auth_method);
}
let auth_method = detected_auth_method.to_lowercase();
let refresh_url = self.get_refresh_url();
tracing::debug!("[KIRO] refresh_token: auth_method={}, refresh_url={}", auth_method, refresh_url);
tracing::debug!("[KIRO] has_client_id={}, has_client_secret={}",
self.credentials.client_id.is_some(),
tracing::debug!(
"[KIRO] refresh_token: auth_method={}, refresh_url={}",
auth_method,
refresh_url
);
tracing::debug!(
"[KIRO] has_client_id={}, has_client_secret={}",
self.credentials.client_id.is_some(),
self.credentials.client_secret.is_some()
);
@@ -521,11 +554,11 @@ impl KiroProvider {
};
tracing::info!("[KIRO] Token 刷新响应状态: {}", resp.status());
if !resp.status().is_success() {
let status = resp.status();
let body_text = resp.text().await.unwrap_or_default();
tracing::warn!("[KIRO] Token 刷新失败: {} - {}", status, body_text);
// 根据具体的HTTP状态码提供更友好的错误信息
@@ -575,7 +608,10 @@ impl KiroProvider {
pub async fn save_credentials(&self) -> Result<(), Box<dyn Error + Send + Sync>> {
// 使用加载时的路径或默认路径
let path = self.creds_path.clone().unwrap_or_else(Self::default_creds_path);
let path = self
.creds_path
.clone()
.unwrap_or_else(Self::default_creds_path);
// 读取现有文件内容
let mut existing: serde_json::Value = if tokio::fs::try_exists(&path).await.unwrap_or(false)
+3
View File
@@ -1,9 +1,12 @@
pub mod antigravity;
pub mod claude_custom;
pub mod gemini;
pub mod kiro;
pub mod openai_custom;
pub mod qwen;
#[allow(unused_imports)]
pub use antigravity::AntigravityProvider;
#[allow(unused_imports)]
pub use claude_custom::ClaudeCustomProvider;
#[allow(unused_imports)]
+13
View File
@@ -0,0 +1,13 @@
//! 路由系统模块
//!
//! 支持动态路由注册和命名空间路由解析。
//! 路由格式:
//! - `/{provider-name}/v1/messages` - Provider 命名空间路由
//! - `/{selector}/v1/messages` - 凭证选择器路由(向后兼容)
//! - `/v1/messages` - 默认路由
mod provider_router;
mod route_registry;
pub use provider_router::ProviderRouter;
pub use route_registry::{RegisteredRoute, RouteRegistry, RouteType};
+262
View File
@@ -0,0 +1,262 @@
//! Provider 路由解析器
//!
//! 解析请求路径,确定目标 Provider 和协议。
use super::route_registry::{RegisteredRoute, RouteRegistry, RouteType};
use crate::models::provider_pool_model::PoolProviderType;
use std::sync::Arc;
use tokio::sync::RwLock;
/// 路由解析结果
#[derive(Debug, Clone)]
pub struct RouteMatch {
/// 匹配的路由
pub route: RegisteredRoute,
/// 请求的协议 (openai 或 claude)
pub protocol: String,
/// 请求的端点 (messages, chat/completions 等)
pub endpoint: String,
/// 路由选择器(从路径中提取)
pub selector: Option<String>,
}
impl RouteMatch {
/// 是否是 Claude 协议
pub fn is_claude_protocol(&self) -> bool {
self.protocol == "claude" || self.endpoint == "messages"
}
/// 是否是 OpenAI 协议
pub fn is_openai_protocol(&self) -> bool {
self.protocol == "openai" || self.endpoint == "chat/completions"
}
/// 获取 Provider 类型
pub fn provider_type(&self) -> Option<PoolProviderType> {
self.route
.provider_type
.as_ref()
.and_then(|s| s.parse().ok())
}
}
/// Provider 路由器
pub struct ProviderRouter {
/// 路由注册表
registry: Arc<RwLock<RouteRegistry>>,
}
impl ProviderRouter {
/// 创建新的路由器
pub fn new(registry: Arc<RwLock<RouteRegistry>>) -> Self {
Self { registry }
}
/// 解析请求路径
///
/// 支持的路径格式:
/// - `/v1/messages` - 默认路由,Claude 协议
/// - `/v1/chat/completions` - 默认路由,OpenAI 协议
/// - `/{selector}/v1/messages` - 选择器路由,Claude 协议
/// - `/{selector}/v1/chat/completions` - 选择器路由,OpenAI 协议
pub async fn resolve(&self, path: &str) -> Option<RouteMatch> {
let path = path.trim_start_matches('/');
let parts: Vec<&str> = path.split('/').collect();
match parts.as_slice() {
// /v1/messages
["v1", "messages"] => {
let registry = self.registry.read().await;
let route = registry
.enabled_routes()
.into_iter()
.find(|r| r.route_type == RouteType::Default)
.cloned()
.unwrap_or_else(|| RegisteredRoute::default_route("kiro"));
Some(RouteMatch {
route,
protocol: "claude".to_string(),
endpoint: "messages".to_string(),
selector: None,
})
}
// /v1/chat/completions
["v1", "chat", "completions"] => {
let registry = self.registry.read().await;
let route = registry
.enabled_routes()
.into_iter()
.find(|r| r.route_type == RouteType::Default)
.cloned()
.unwrap_or_else(|| RegisteredRoute::default_route("kiro"));
Some(RouteMatch {
route,
protocol: "openai".to_string(),
endpoint: "chat/completions".to_string(),
selector: None,
})
}
// /{selector}/v1/messages
[selector, "v1", "messages"] => {
let registry = self.registry.read().await;
let route = registry
.find_by_selector(selector)
.cloned()
.unwrap_or_else(|| {
// 创建一个临时的选择器路由
RegisteredRoute {
path_pattern: format!("/{}/v1/messages", selector),
route_type: RouteType::CredentialSelector,
provider_type: None,
credential_uuid: None,
credential_name: Some(selector.to_string()),
protocols: vec!["claude".to_string()],
enabled: true,
priority: 50,
}
});
Some(RouteMatch {
route,
protocol: "claude".to_string(),
endpoint: "messages".to_string(),
selector: Some(selector.to_string()),
})
}
// /{selector}/v1/chat/completions
[selector, "v1", "chat", "completions"] => {
let registry = self.registry.read().await;
let route = registry
.find_by_selector(selector)
.cloned()
.unwrap_or_else(|| RegisteredRoute {
path_pattern: format!("/{}/v1/chat/completions", selector),
route_type: RouteType::CredentialSelector,
provider_type: None,
credential_uuid: None,
credential_name: Some(selector.to_string()),
protocols: vec!["openai".to_string()],
enabled: true,
priority: 50,
});
Some(RouteMatch {
route,
protocol: "openai".to_string(),
endpoint: "chat/completions".to_string(),
selector: Some(selector.to_string()),
})
}
_ => None,
}
}
/// 注册凭证路由
pub async fn register_credential(
&self,
provider_type: &str,
credential_uuid: &str,
credential_name: Option<&str>,
) {
let mut registry = self.registry.write().await;
// 注册命名空间路由
let route =
RegisteredRoute::provider_namespace(provider_type, credential_uuid, credential_name);
registry.register(route);
// 注册 UUID 选择器路由
let selector_route = RegisteredRoute::credential_selector(credential_uuid, provider_type);
registry.register(selector_route);
}
/// 注销凭证路由
pub async fn unregister_credential(&self, credential_uuid: &str) {
let mut registry = self.registry.write().await;
registry.unregister(credential_uuid);
}
/// 获取所有注册的路由
pub async fn list_routes(&self) -> Vec<RegisteredRoute> {
let registry = self.registry.read().await;
registry.all_routes().to_vec()
}
/// 生成路由 URL
pub fn generate_url(&self, base_url: &str, route: &RegisteredRoute, protocol: &str) -> String {
let endpoint = if protocol == "claude" {
"messages"
} else {
"chat/completions"
};
if route.route_type == RouteType::Default {
format!("{}/v1/{}", base_url, endpoint)
} else {
let selector = route
.credential_name
.as_ref()
.or(route.credential_uuid.as_ref())
.map(|s| s.as_str())
.unwrap_or("unknown");
format!("{}/{}/v1/{}", base_url, selector, endpoint)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_resolve_default_routes() {
let registry = Arc::new(RwLock::new(RouteRegistry::new()));
let router = ProviderRouter::new(registry);
let match1 = router.resolve("/v1/messages").await.unwrap();
assert_eq!(match1.protocol, "claude");
assert_eq!(match1.endpoint, "messages");
assert!(match1.selector.is_none());
let match2 = router.resolve("/v1/chat/completions").await.unwrap();
assert_eq!(match2.protocol, "openai");
assert_eq!(match2.endpoint, "chat/completions");
assert!(match2.selector.is_none());
}
#[tokio::test]
async fn test_resolve_selector_routes() {
let registry = Arc::new(RwLock::new(RouteRegistry::new()));
let router = ProviderRouter::new(registry);
let match1 = router.resolve("/my-kiro/v1/messages").await.unwrap();
assert_eq!(match1.protocol, "claude");
assert_eq!(match1.selector, Some("my-kiro".to_string()));
let match2 = router
.resolve("/my-kiro/v1/chat/completions")
.await
.unwrap();
assert_eq!(match2.protocol, "openai");
assert_eq!(match2.selector, Some("my-kiro".to_string()));
}
#[tokio::test]
async fn test_register_and_resolve() {
let registry = Arc::new(RwLock::new(RouteRegistry::new()));
let router = ProviderRouter::new(registry);
router
.register_credential("kiro", "uuid-123", Some("my-kiro-account"))
.await;
let match1 = router
.resolve("/my-kiro-account/v1/messages")
.await
.unwrap();
assert_eq!(match1.route.credential_uuid, Some("uuid-123".to_string()));
assert_eq!(match1.route.provider_type, Some("kiro".to_string()));
}
}
+266
View File
@@ -0,0 +1,266 @@
//! 路由注册表
//!
//! 管理所有注册的路由,支持动态添加和查询。
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
/// 路由类型
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum RouteType {
/// Provider 命名空间路由 (如 /claude-kiro-oauth/v1/messages)
ProviderNamespace,
/// 凭证选择器路由 (如 /{uuid}/v1/messages)
CredentialSelector,
/// 默认路由 (如 /v1/messages)
Default,
}
/// 注册的路由信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RegisteredRoute {
/// 路由路径模式
pub path_pattern: String,
/// 路由类型
pub route_type: RouteType,
/// Provider 类型(如果适用)
pub provider_type: Option<String>,
/// 凭证 UUID(如果适用)
pub credential_uuid: Option<String>,
/// 凭证名称(如果适用)
pub credential_name: Option<String>,
/// 支持的协议
pub protocols: Vec<String>,
/// 是否启用
pub enabled: bool,
/// 优先级(数字越小优先级越高)
pub priority: u32,
}
impl RegisteredRoute {
/// 创建 Provider 命名空间路由
pub fn provider_namespace(
provider_type: &str,
credential_uuid: &str,
credential_name: Option<&str>,
) -> Self {
let path_pattern = format!(
"/{}/v1/{{endpoint}}",
Self::generate_route_name(provider_type, credential_name)
);
Self {
path_pattern,
route_type: RouteType::ProviderNamespace,
provider_type: Some(provider_type.to_string()),
credential_uuid: Some(credential_uuid.to_string()),
credential_name: credential_name.map(|s| s.to_string()),
protocols: vec!["openai".to_string(), "claude".to_string()],
enabled: true,
priority: 10,
}
}
/// 创建凭证选择器路由
pub fn credential_selector(credential_uuid: &str, provider_type: &str) -> Self {
Self {
path_pattern: format!("/{}/v1/{{endpoint}}", credential_uuid),
route_type: RouteType::CredentialSelector,
provider_type: Some(provider_type.to_string()),
credential_uuid: Some(credential_uuid.to_string()),
credential_name: None,
protocols: vec!["openai".to_string(), "claude".to_string()],
enabled: true,
priority: 20,
}
}
/// 创建默认路由
pub fn default_route(provider_type: &str) -> Self {
Self {
path_pattern: "/v1/{endpoint}".to_string(),
route_type: RouteType::Default,
provider_type: Some(provider_type.to_string()),
credential_uuid: None,
credential_name: None,
protocols: vec!["openai".to_string(), "claude".to_string()],
enabled: true,
priority: 100,
}
}
/// 生成路由名称
fn generate_route_name(provider_type: &str, credential_name: Option<&str>) -> String {
if let Some(name) = credential_name {
// 将名称转换为 URL 友好格式
name.to_lowercase()
.replace(' ', "-")
.chars()
.filter(|c| c.is_alphanumeric() || *c == '-' || *c == '_')
.collect()
} else {
provider_type.to_lowercase()
}
}
/// 获取路由的显示名称
pub fn display_name(&self) -> String {
if let Some(name) = &self.credential_name {
name.clone()
} else if let Some(provider) = &self.provider_type {
provider.clone()
} else {
"default".to_string()
}
}
}
/// 路由注册表
#[derive(Debug, Default)]
pub struct RouteRegistry {
/// 所有注册的路由
routes: Vec<RegisteredRoute>,
/// 路由名称到索引的映射
name_index: HashMap<String, usize>,
/// UUID 到索引的映射
uuid_index: HashMap<String, usize>,
}
impl RouteRegistry {
/// 创建新的路由注册表
pub fn new() -> Self {
Self::default()
}
/// 注册路由
pub fn register(&mut self, route: RegisteredRoute) {
let index = self.routes.len();
// 更新索引
if let Some(name) = &route.credential_name {
self.name_index.insert(name.to_lowercase(), index);
}
if let Some(uuid) = &route.credential_uuid {
self.uuid_index.insert(uuid.clone(), index);
}
self.routes.push(route);
// 按优先级排序
self.sort_by_priority();
}
/// 注销路由
pub fn unregister(&mut self, credential_uuid: &str) {
if let Some(&index) = self.uuid_index.get(credential_uuid) {
if index < self.routes.len() {
let route = &self.routes[index];
if let Some(name) = &route.credential_name {
self.name_index.remove(&name.to_lowercase());
}
self.uuid_index.remove(credential_uuid);
self.routes.remove(index);
// 重建索引
self.rebuild_indices();
}
}
}
/// 按名称查找路由
pub fn find_by_name(&self, name: &str) -> Option<&RegisteredRoute> {
self.name_index
.get(&name.to_lowercase())
.and_then(|&idx| self.routes.get(idx))
}
/// 按 UUID 查找路由
pub fn find_by_uuid(&self, uuid: &str) -> Option<&RegisteredRoute> {
self.uuid_index
.get(uuid)
.and_then(|&idx| self.routes.get(idx))
}
/// 按选择器查找路由(名称或 UUID)
pub fn find_by_selector(&self, selector: &str) -> Option<&RegisteredRoute> {
self.find_by_name(selector)
.or_else(|| self.find_by_uuid(selector))
}
/// 获取所有路由
pub fn all_routes(&self) -> &[RegisteredRoute] {
&self.routes
}
/// 获取启用的路由
pub fn enabled_routes(&self) -> Vec<&RegisteredRoute> {
self.routes.iter().filter(|r| r.enabled).collect()
}
/// 按 Provider 类型获取路由
pub fn routes_by_provider(&self, provider_type: &str) -> Vec<&RegisteredRoute> {
self.routes
.iter()
.filter(|r| r.provider_type.as_deref() == Some(provider_type))
.collect()
}
/// 清空所有路由
pub fn clear(&mut self) {
self.routes.clear();
self.name_index.clear();
self.uuid_index.clear();
}
/// 按优先级排序
fn sort_by_priority(&mut self) {
self.routes.sort_by_key(|r| r.priority);
self.rebuild_indices();
}
/// 重建索引
fn rebuild_indices(&mut self) {
self.name_index.clear();
self.uuid_index.clear();
for (index, route) in self.routes.iter().enumerate() {
if let Some(name) = &route.credential_name {
self.name_index.insert(name.to_lowercase(), index);
}
if let Some(uuid) = &route.credential_uuid {
self.uuid_index.insert(uuid.clone(), index);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_route_registration() {
let mut registry = RouteRegistry::new();
let route =
RegisteredRoute::provider_namespace("kiro", "uuid-123", Some("my-kiro-account"));
registry.register(route);
assert!(registry.find_by_name("my-kiro-account").is_some());
assert!(registry.find_by_uuid("uuid-123").is_some());
assert!(registry.find_by_selector("my-kiro-account").is_some());
}
#[test]
fn test_route_unregistration() {
let mut registry = RouteRegistry::new();
let route =
RegisteredRoute::provider_namespace("kiro", "uuid-123", Some("my-kiro-account"));
registry.register(route);
registry.unregister("uuid-123");
assert!(registry.find_by_name("my-kiro-account").is_none());
assert!(registry.find_by_uuid("uuid-123").is_none());
}
}
+168 -23
View File
@@ -1,11 +1,15 @@
//! HTTP API 服务器
use crate::config::Config;
use crate::converter::anthropic_to_openai::convert_anthropic_to_openai;
use crate::converter::openai_to_antigravity::{
convert_antigravity_to_openai_response, convert_openai_to_antigravity,
};
use crate::database::DbConnection;
use crate::logger::LogStore;
use crate::models::anthropic::*;
use crate::models::openai::*;
use crate::models::route_model::{RouteInfo, RouteListResponse};
use crate::providers::antigravity::AntigravityProvider;
use crate::providers::claude_custom::ClaudeCustomProvider;
use crate::providers::gemini::GeminiProvider;
use crate::providers::kiro::KiroProvider;
@@ -113,7 +117,19 @@ impl ServerState {
let kiro = self.kiro_provider.clone();
tokio::spawn(async move {
if let Err(e) = run_server(&host, port, &api_key, kiro, logs, rx, pool_service, token_cache, db).await {
if let Err(e) = run_server(
&host,
port,
&api_key,
kiro,
logs,
rx,
pool_service,
token_cache,
db,
)
.await
{
tracing::error!("Server error: {}", e);
}
});
@@ -190,8 +206,14 @@ async fn run_server(
.route("/v1/messages", post(anthropic_messages))
.route("/v1/messages/count_tokens", post(count_tokens))
// 多供应商路由
.route("/{selector}/v1/messages", post(anthropic_messages_with_selector))
.route("/{selector}/v1/chat/completions", post(chat_completions_with_selector))
.route(
"/:selector/v1/messages",
post(anthropic_messages_with_selector),
)
.route(
"/:selector/v1/chat/completions",
post(chat_completions_with_selector),
)
.with_state(state);
let addr: std::net::SocketAddr = format!("{host}:{port}").parse()?;
@@ -1384,11 +1406,10 @@ async fn anthropic_messages_with_selector(
Json(request): Json<AnthropicMessagesRequest>,
) -> Response {
if let Err(e) = verify_api_key(&headers, &state.api_key).await {
state
.logs
.write()
.await
.add("warn", &format!("Unauthorized request to /{}/v1/messages", selector));
state.logs.write().await.add(
"warn",
&format!("Unauthorized request to /{}/v1/messages", selector),
);
return e.into_response();
}
@@ -1412,9 +1433,10 @@ async fn anthropic_messages_with_selector(
Some(cred)
}
// 最后尝试按 provider 类型轮询
else if let Ok(Some(cred)) = state
.pool_service
.select_credential(db, &selector, Some(&request.model))
else if let Ok(Some(cred)) =
state
.pool_service
.select_credential(db, &selector, Some(&request.model))
{
Some(cred)
} else {
@@ -1462,11 +1484,10 @@ async fn chat_completions_with_selector(
Json(request): Json<ChatCompletionRequest>,
) -> Response {
if let Err(e) = verify_api_key(&headers, &state.api_key).await {
state
.logs
.write()
.await
.add("warn", &format!("Unauthorized request to /{}/v1/chat/completions", selector));
state.logs.write().await.add(
"warn",
&format!("Unauthorized request to /{}/v1/chat/completions", selector),
);
return e.into_response();
}
@@ -1485,9 +1506,10 @@ async fn chat_completions_with_selector(
Some(cred)
} else if let Ok(Some(cred)) = state.pool_service.get_by_uuid(db, &selector) {
Some(cred)
} else if let Ok(Some(cred)) = state
.pool_service
.select_credential(db, &selector, Some(&request.model))
} else if let Ok(Some(cred)) =
state
.pool_service
.select_credential(db, &selector, Some(&request.model))
{
Some(cred)
} else {
@@ -1708,7 +1730,11 @@ async fn call_provider_anthropic(
};
// 获取缓存的 token
let token = match state.token_cache.get_valid_token(db, &credential.uuid).await {
let token = match state
.token_cache
.get_valid_token(db, &credential.uuid)
.await
{
Ok(t) => t,
Err(e) => {
tracing::warn!("[POOL] Token cache miss, loading from source: {}", e);
@@ -1770,9 +1796,17 @@ async fn call_provider_anthropic(
}
} else if status.as_u16() == 401 || status.as_u16() == 403 {
// Token 过期,强制刷新并重试
tracing::info!("[POOL] Got {}, forcing token refresh for {}", status, &credential.uuid[..8]);
tracing::info!(
"[POOL] Got {}, forcing token refresh for {}",
status,
&credential.uuid[..8]
);
let new_token = match state.token_cache.refresh_and_cache(db, &credential.uuid, true).await {
let new_token = match state
.token_cache
.refresh_and_cache(db, &credential.uuid, true)
.await
{
Ok(t) => t,
Err(e) => {
return (
@@ -1844,6 +1878,72 @@ async fn call_provider_anthropic(
)
.into_response()
}
CredentialData::AntigravityOAuth {
creds_file_path,
project_id,
} => {
let mut antigravity = AntigravityProvider::new();
if let Err(e) = antigravity
.load_credentials_from_path(creds_file_path)
.await
{
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": format!("Failed to load Antigravity credentials: {}", e)}})),
)
.into_response();
}
// 检查并刷新 token
if antigravity.is_token_expiring_soon() {
if let Err(e) = antigravity.refresh_token().await {
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})),
)
.into_response();
}
}
// 设置项目 ID
if let Some(pid) = project_id {
antigravity.project_id = Some(pid.clone());
} else if let Err(e) = antigravity.discover_project().await {
tracing::warn!("[Antigravity] Failed to discover project: {}", e);
}
// 先转换为 OpenAI 格式,再转换为 Antigravity 格式
let openai_request = convert_anthropic_to_openai(request);
let antigravity_request = convert_openai_to_antigravity(&openai_request);
match antigravity
.generate_content(&request.model, &antigravity_request)
.await
{
Ok(resp) => {
// 转换为 OpenAI 格式,再构建 Anthropic 响应
let content = resp["candidates"][0]["content"]["parts"][0]["text"]
.as_str()
.unwrap_or("");
let parsed = CWParsedResponse {
content: content.to_string(),
tool_calls: Vec::new(),
usage_credits: 0.0,
context_usage_percentage: 0.0,
};
if request.stream {
build_anthropic_stream_response(&request.model, &parsed)
} else {
build_anthropic_response(&request.model, &parsed)
}
}
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response(),
}
}
CredentialData::OpenAIKey { api_key, base_url } => {
let openai = OpenAICustomProvider::with_config(api_key.clone(), base_url.clone());
let openai_request = convert_anthropic_to_openai(request);
@@ -1852,7 +1952,9 @@ async fn call_provider_anthropic(
if resp.status().is_success() {
match resp.text().await {
Ok(body) => {
if let Ok(openai_resp) = serde_json::from_str::<serde_json::Value>(&body) {
if let Ok(openai_resp) =
serde_json::from_str::<serde_json::Value>(&body)
{
let content = openai_resp["choices"][0]["message"]["content"]
.as_str()
.unwrap_or("");
@@ -2053,6 +2155,49 @@ async fn call_provider_openai(
)
.into_response()
}
CredentialData::AntigravityOAuth { creds_file_path, project_id } => {
let mut antigravity = AntigravityProvider::new();
if let Err(e) = antigravity.load_credentials_from_path(creds_file_path).await {
return (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": format!("Failed to load Antigravity credentials: {}", e)}})),
)
.into_response();
}
// 检查并刷新 token
if antigravity.is_token_expiring_soon() {
if let Err(e) = antigravity.refresh_token().await {
return (
StatusCode::UNAUTHORIZED,
Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {}", e)}})),
)
.into_response();
}
}
// 设置项目 ID
if let Some(pid) = project_id {
antigravity.project_id = Some(pid.clone());
} else if let Err(e) = antigravity.discover_project().await {
tracing::warn!("[Antigravity] Failed to discover project: {}", e);
}
// 转换请求格式
let antigravity_request = convert_openai_to_antigravity(request);
match antigravity.generate_content(&request.model, &antigravity_request).await {
Ok(resp) => {
let openai_response = convert_antigravity_to_openai_response(&resp, &request.model);
Json(openai_response).into_response()
}
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": e.to_string()}})),
)
.into_response(),
}
}
CredentialData::OpenAIKey { api_key, base_url } => {
let openai = OpenAICustomProvider::with_config(api_key.clone(), base_url.clone());
match openai.call_api(request).await {
+1
View File
@@ -3,6 +3,7 @@ use serde_json::{json, Value};
use std::path::PathBuf;
/// Get the configuration file path for an app type
#[allow(dead_code)]
pub fn get_app_config_path(app_type: &AppType) -> Option<PathBuf> {
let home = dirs::home_dir()?;
match app_type {
+1
View File
@@ -3,6 +3,7 @@ use serde_json::{json, Map, Value};
use std::path::PathBuf;
/// Get the MCP config file path for an app type
#[allow(dead_code)]
pub fn get_mcp_config_path(app_type: &AppType) -> Option<PathBuf> {
let home = dirs::home_dir()?;
match app_type {
+1
View File
@@ -6,6 +6,7 @@ use std::collections::HashMap;
pub struct PromptService;
#[allow(dead_code)]
impl PromptService {
/// Get all prompts for an app type
pub fn get_all(db: &DbConnection, app_type: &str) -> Result<Vec<Prompt>, String> {
+2
View File
@@ -18,6 +18,7 @@ pub fn get_prompt_file_path(app: &AppType) -> Option<PathBuf> {
}
/// Get the prompt file name for an app
#[allow(dead_code)]
pub fn get_prompt_filename(app: &AppType) -> &'static str {
match app {
AppType::Claude => "CLAUDE.md",
@@ -67,6 +68,7 @@ pub fn write_live_prompt(app: &AppType, content: &str) -> Result<(), String> {
}
/// Delete the live prompt file
#[allow(dead_code)]
pub fn delete_live_prompt(app: &AppType) -> Result<(), String> {
let path = get_prompt_file_path(app).ok_or("Cannot determine prompt file path")?;
+121 -28
View File
@@ -85,7 +85,8 @@ impl ProviderPoolService {
) -> Result<Vec<CredentialDisplay>, String> {
let pt: PoolProviderType = provider_type.parse().map_err(|e: String| e)?;
let conn = db.lock().map_err(|e| e.to_string())?;
let mut credentials = ProviderPoolDao::get_by_type(&conn, &pt).map_err(|e| e.to_string())?;
let mut credentials =
ProviderPoolDao::get_by_type(&conn, &pt).map_err(|e| e.to_string())?;
// 为每个凭证加载 token 缓存
for cred in &mut credentials {
@@ -390,6 +391,13 @@ impl ProviderPoolService {
CredentialData::QwenOAuth { creds_file_path } => {
self.check_qwen_health(creds_file_path, model).await
}
CredentialData::AntigravityOAuth {
creds_file_path,
project_id,
} => {
self.check_antigravity_health(creds_file_path, project_id.as_deref(), model)
.await
}
CredentialData::OpenAIKey { api_key, base_url } => {
self.check_openai_health(api_key, base_url.as_deref(), model)
.await
@@ -411,9 +419,13 @@ impl ProviderPoolService {
format!("{} 认证失败,凭证可能已过期或无效。\n💡 解决方案:\n1. 点击\"刷新\"按钮尝试更新 Token\n2. 如刷新失败,请删除后重新添加此凭证\n3. 检查账户权限是否正常", provider_type)
} else if error.contains("HTTP 429") {
format!("{} 请求频率过高,已被限流。\n💡 解决方案:\n1. 稍等几分钟后再次尝试\n2. 考虑添加更多凭证分散负载", provider_type)
} else if error.contains("HTTP 500") || error.contains("HTTP 502") || error.contains("HTTP 503") {
} else if error.contains("HTTP 500")
|| error.contains("HTTP 502")
|| error.contains("HTTP 503")
{
format!("{} 服务暂时不可用。\n💡 解决方案:\n1. 这通常是服务提供方的临时问题\n2. 请稍后重试\n3. 如问题持续,可尝试其他凭证", provider_type)
} else if error.contains("读取凭证文件失败") || error.contains("解析凭证失败") {
} else if error.contains("读取凭证文件失败") || error.contains("解析凭证失败")
{
format!("凭证文件损坏或不可读。\n💡 解决方案:\n1. 凭证文件可能已损坏\n2. 建议删除此凭证后重新添加\n3. 确保文件权限正确且格式为有效的 JSON")
} else {
// 对于其他未识别的错误,提供通用建议
@@ -424,17 +436,24 @@ impl ProviderPoolService {
// Kiro OAuth 健康检查
async fn check_kiro_health(&self, creds_path: &str, model: &str) -> Result<(), String> {
tracing::debug!("[KIRO HEALTH] 开始健康检查,凭证路径: {}", creds_path);
// 使用 KiroProvider 加载凭证(包括 clientIdHash 文件)
let mut provider = KiroProvider::new();
provider.load_credentials_from_path(creds_path).await
.map_err(|e| self.format_user_friendly_error(&format!("加载凭证失败: {}", e), "Kiro"))?;
provider
.load_credentials_from_path(creds_path)
.await
.map_err(|e| {
self.format_user_friendly_error(&format!("加载凭证失败: {}", e), "Kiro")
})?;
let access_token = provider.credentials.access_token.as_ref()
let access_token = provider
.credentials
.access_token
.as_ref()
.ok_or_else(|| "凭证中缺少 access_token".to_string())?;
let health_check_url = provider.get_health_check_url();
// 获取 modelId 映射
let model_id = match model {
"claude-opus-4-5" | "claude-opus-4-5-20251101" => "claude-opus-4.5",
@@ -444,7 +463,7 @@ impl ProviderPoolService {
"claude-3-7-sonnet-20250219" => "CLAUDE_3_7_SONNET_20250219_V1_0",
_ => "claude-haiku-4.5", // 默认使用 haiku
};
tracing::debug!("[KIRO HEALTH] 健康检查 URL: {}", health_check_url);
tracing::debug!("[KIRO HEALTH] 使用模型: {} -> {}", model, model_id);
@@ -585,6 +604,44 @@ impl ProviderPoolService {
}
}
// Antigravity OAuth 健康检查
async fn check_antigravity_health(
&self,
creds_path: &str,
_project_id: Option<&str>,
_model: &str,
) -> Result<(), String> {
let creds_content =
std::fs::read_to_string(creds_path).map_err(|e| format!("读取凭证文件失败: {}", e))?;
let creds: serde_json::Value =
serde_json::from_str(&creds_content).map_err(|e| format!("解析凭证失败: {}", e))?;
let access_token = creds["access_token"]
.as_str()
.ok_or_else(|| "凭证中缺少 access_token".to_string())?;
// 使用 fetchAvailableModels 作为健康检查
let url =
"https://daily-cloudcode-pa.sandbox.googleapis.com/v1internal:fetchAvailableModels";
let response = self
.client
.post(url)
.bearer_auth(access_token)
.header("User-Agent", "antigravity/1.11.5 windows/amd64")
.json(&serde_json::json!({}))
.timeout(self.health_check_timeout)
.send()
.await
.map_err(|e| format!("请求失败: {}", e))?;
if response.status().is_success() {
Ok(())
} else {
Err(format!("HTTP {}", response.status()))
}
}
// OpenAI API 健康检查
async fn check_openai_health(
&self,
@@ -726,18 +783,24 @@ impl ProviderPoolService {
}
/// 获取 OAuth 凭证状态
pub fn get_oauth_status(&self, creds_path: &str, provider_type: &str) -> Result<OAuthStatus, String> {
let content = std::fs::read_to_string(creds_path)
.map_err(|e| format!("读取凭证文件失败: {}", e))?;
let creds: serde_json::Value = serde_json::from_str(&content)
.map_err(|e| format!("解析凭证文件失败: {}", e))?;
pub fn get_oauth_status(
&self,
creds_path: &str,
provider_type: &str,
) -> Result<OAuthStatus, String> {
let content =
std::fs::read_to_string(creds_path).map_err(|e| format!("读取凭证文件失败: {}", e))?;
let creds: serde_json::Value =
serde_json::from_str(&content).map_err(|e| format!("解析凭证文件失败: {}", e))?;
let has_access_token = creds.get("accessToken")
let has_access_token = creds
.get("accessToken")
.or_else(|| creds.get("access_token"))
.map(|v| v.as_str().is_some())
.unwrap_or(false);
let has_refresh_token = creds.get("refreshToken")
let has_refresh_token = creds
.get("refreshToken")
.or_else(|| creds.get("refresh_token"))
.map(|v| v.as_str().is_some())
.unwrap_or(false);
@@ -745,7 +808,8 @@ impl ProviderPoolService {
// 检查 token 是否有效(根据 expiry_date 判断)
let (is_token_valid, expiry_info) = match provider_type {
"kiro" => {
let expires_at = creds.get("expiresAt")
let expires_at = creds
.get("expiresAt")
.or_else(|| creds.get("expires_at"))
.and_then(|v| v.as_str())
.map(|s| s.to_string());
@@ -780,27 +844,53 @@ impl ProviderPoolService {
/// 刷新 OAuth Token (Kiro)
pub async fn refresh_kiro_token(&self, creds_path: &str) -> Result<String, String> {
let mut provider = crate::providers::kiro::KiroProvider::new();
provider.load_credentials_from_path(creds_path).await
.map_err(|e| self.format_user_friendly_error(&format!("加载凭证失败: {}", e), "Kiro"))?;
provider.refresh_token().await
.map_err(|e| self.format_user_friendly_error(&format!("刷新 Token 失败: {}", e), "Kiro"))
provider
.load_credentials_from_path(creds_path)
.await
.map_err(|e| {
self.format_user_friendly_error(&format!("加载凭证失败: {}", e), "Kiro")
})?;
provider.refresh_token().await.map_err(|e| {
self.format_user_friendly_error(&format!("刷新 Token 失败: {}", e), "Kiro")
})
}
/// 刷新 OAuth Token (Gemini)
pub async fn refresh_gemini_token(&self, creds_path: &str) -> Result<String, String> {
let mut provider = crate::providers::gemini::GeminiProvider::new();
provider.load_credentials_from_path(creds_path).await
provider
.load_credentials_from_path(creds_path)
.await
.map_err(|e| format!("加载凭证失败: {}", e))?;
provider.refresh_token().await
provider
.refresh_token()
.await
.map_err(|e| format!("刷新 Token 失败: {}", e))
}
/// 刷新 OAuth Token (Qwen)
pub async fn refresh_qwen_token(&self, creds_path: &str) -> Result<String, String> {
let mut provider = crate::providers::qwen::QwenProvider::new();
provider.load_credentials_from_path(creds_path).await
provider
.load_credentials_from_path(creds_path)
.await
.map_err(|e| format!("加载凭证失败: {}", e))?;
provider.refresh_token().await
provider
.refresh_token()
.await
.map_err(|e| format!("刷新 Token 失败: {}", e))
}
/// 刷新 OAuth Token (Antigravity)
pub async fn refresh_antigravity_token(&self, creds_path: &str) -> Result<String, String> {
let mut provider = crate::providers::antigravity::AntigravityProvider::new();
provider
.load_credentials_from_path(creds_path)
.await
.map_err(|e| format!("加载凭证失败: {}", e))?;
provider
.refresh_token()
.await
.map_err(|e| format!("刷新 Token 失败: {}", e))
}
@@ -821,12 +911,15 @@ impl ProviderPoolService {
CredentialData::KiroOAuth { creds_file_path } => {
self.refresh_kiro_token(creds_file_path).await
}
CredentialData::GeminiOAuth { creds_file_path, .. } => {
self.refresh_gemini_token(creds_file_path).await
}
CredentialData::GeminiOAuth {
creds_file_path, ..
} => self.refresh_gemini_token(creds_file_path).await,
CredentialData::QwenOAuth { creds_file_path } => {
self.refresh_qwen_token(creds_file_path).await
}
CredentialData::AntigravityOAuth {
creds_file_path, ..
} => self.refresh_antigravity_token(creds_file_path).await,
_ => Err("此凭证类型不支持 Token 刷新".to_string()),
}
}
+65 -6
View File
@@ -43,11 +43,7 @@ impl TokenCacheService {
/// 1. 检查数据库缓存是否有效
/// 2. 如果缓存有效且未过期,直接返回
/// 3. 如果缓存无效或即将过期,执行刷新
pub async fn get_valid_token(
&self,
db: &DbConnection,
uuid: &str,
) -> Result<String, String> {
pub async fn get_valid_token(&self, db: &DbConnection, uuid: &str) -> Result<String, String> {
// 首先检查缓存
let cached = {
let conn = db.lock().map_err(|e| e.to_string())?;
@@ -153,7 +149,11 @@ impl TokenCacheService {
let _ = ProviderPoolDao::record_token_refresh_error(&conn, uuid, &e);
}
tracing::error!("[TOKEN_CACHE] Token refresh failed for {}: {}", &uuid[..8], e);
tracing::error!(
"[TOKEN_CACHE] Token refresh failed for {}: {}",
&uuid[..8],
e
);
Err(e)
}
@@ -172,6 +172,9 @@ impl TokenCacheService {
CredentialData::QwenOAuth { creds_file_path } => {
self.refresh_qwen(creds_file_path).await
}
CredentialData::AntigravityOAuth {
creds_file_path, ..
} => self.refresh_antigravity(creds_file_path).await,
CredentialData::OpenAIKey { api_key, .. } => {
// API Key 不需要刷新,直接返回
Ok(CachedTokenInfo {
@@ -283,6 +286,38 @@ impl TokenCacheService {
})
}
/// 刷新 Antigravity Token
async fn refresh_antigravity(&self, creds_path: &str) -> Result<CachedTokenInfo, String> {
use crate::providers::antigravity::AntigravityProvider;
let mut provider = AntigravityProvider::new();
provider
.load_credentials_from_path(creds_path)
.await
.map_err(|e| format!("加载 Antigravity 凭证失败: {}", e))?;
let token = provider
.refresh_token()
.await
.map_err(|e| format!("刷新 Antigravity Token 失败: {}", e))?;
// Antigravity token 通常 1 小时过期
let expiry_time = provider
.credentials
.expiry_date
.map(|ts| chrono::DateTime::from_timestamp_millis(ts).unwrap_or_default())
.unwrap_or_else(|| Utc::now() + chrono::Duration::minutes(50));
Ok(CachedTokenInfo {
access_token: Some(token),
refresh_token: provider.credentials.refresh_token.clone(),
expiry_time: Some(expiry_time),
last_refresh: Some(Utc::now()),
refresh_error_count: 0,
last_refresh_error: None,
})
}
/// 从源文件加载初始 Token(首次使用时)
pub async fn load_initial_token(
&self,
@@ -388,6 +423,30 @@ impl TokenCacheService {
last_refresh_error: None,
})
}
CredentialData::AntigravityOAuth {
creds_file_path, ..
} => {
let content = tokio::fs::read_to_string(creds_file_path)
.await
.map_err(|e| format!("读取 Antigravity 凭证文件失败: {}", e))?;
let creds: serde_json::Value =
serde_json::from_str(&content).map_err(|e| format!("解析凭证失败: {}", e))?;
let access_token = creds["access_token"].as_str().map(|s| s.to_string());
let refresh_token = creds["refresh_token"].as_str().map(|s| s.to_string());
let expiry_time = creds["expiry_date"]
.as_i64()
.and_then(|ts| chrono::DateTime::from_timestamp(ts, 0));
Ok(CachedTokenInfo {
access_token,
refresh_token,
expiry_time,
last_refresh: None,
refresh_error_count: 0,
last_refresh_error: None,
})
}
CredentialData::OpenAIKey { api_key, .. } => Ok(CachedTokenInfo {
access_token: Some(api_key.clone()),
refresh_token: None,
+1 -1
View File
@@ -1,7 +1,7 @@
{
"$schema": "https://schema.tauri.app/config/2",
"productName": "ProxyCast",
"version": "0.6.1",
"version": "0.7.0",
"identifier": "com.proxycast.app",
"build": {
"beforeDevCommand": "npm run dev",
+2
View File
@@ -76,6 +76,8 @@ export function RoutesTab() {
return "bg-green-100 text-green-700 dark:bg-green-900/30 dark:text-green-400";
case "claude":
return "bg-amber-100 text-amber-700 dark:bg-amber-900/30 dark:text-amber-400";
case "antigravity":
return "bg-cyan-100 text-cyan-700 dark:bg-cyan-900/30 dark:text-cyan-400";
default:
return "bg-gray-100 text-gray-700 dark:bg-gray-900/30 dark:text-gray-400";
}
@@ -14,6 +14,7 @@ const defaultCredsPath: Record<string, string> = {
kiro: "~/.aws/sso/cache/kiro-auth-token.json",
gemini: "~/.gemini/oauth_creds.json",
qwen: "~/.qwen/oauth_creds.json",
antigravity: "~/.antigravity/oauth_creds.json",
};
export function AddCredentialModal({
@@ -35,7 +36,9 @@ export function AddCredentialModal({
const [apiKey, setApiKey] = useState("");
const [baseUrl, setBaseUrl] = useState("");
const isOAuth = ["kiro", "gemini", "qwen"].includes(providerType);
const isOAuth = ["kiro", "gemini", "qwen", "antigravity"].includes(
providerType,
);
const providerLabels: Record<PoolProviderType, string> = {
kiro: "Kiro (AWS)",
@@ -43,6 +46,7 @@ export function AddCredentialModal({
qwen: "Qwen (阿里)",
openai: "OpenAI",
claude: "Claude (Anthropic)",
antigravity: "Antigravity (Gemini 3 Pro)",
};
const handleSelectFile = async () => {
+133 -280
View File
@@ -9,10 +9,6 @@ import {
Clock,
AlertTriangle,
RefreshCw,
Key,
CheckCircle,
XCircle,
Database,
Settings,
} from "lucide-react";
import type { CredentialDisplay } from "@/lib/api/providerPool";
@@ -58,6 +54,7 @@ export function CredentialCard({
kiro_oauth: "OAuth",
gemini_oauth: "OAuth",
qwen_oauth: "OAuth",
antigravity_oauth: "OAuth",
openai_key: "API Key",
claude_key: "API Key",
};
@@ -70,332 +67,188 @@ export function CredentialCard({
return (
<div
className={`rounded-xl border p-5 transition-all hover:shadow-md ${
className={`rounded-xl border p-4 transition-all hover:shadow-md ${
credential.is_disabled
? "border-gray-200 bg-gray-50/80 opacity-60 dark:border-gray-700 dark:bg-gray-900/60"
? "border-gray-200 bg-gray-50/80 opacity-70 dark:border-gray-700 dark:bg-gray-900/60"
: isHealthy
? "border-green-200 bg-gradient-to-br from-green-50/80 to-green-100/40 dark:border-green-800 dark:bg-gradient-to-br dark:from-green-950/40 dark:to-green-900/20 shadow-green-500/5"
: "border-red-200 bg-gradient-to-br from-red-50/80 to-red-100/40 dark:border-red-800 dark:bg-gradient-to-br dark:from-red-950/40 dark:to-red-900/20 shadow-red-500/5"
? "border-green-200 bg-gradient-to-r from-green-50/80 to-white dark:border-green-800 dark:bg-gradient-to-r dark:from-green-950/40 dark:to-transparent"
: "border-red-200 bg-gradient-to-r from-red-50/80 to-white dark:border-red-800 dark:bg-gradient-to-r dark:from-red-950/40 dark:to-transparent"
}`}
>
{/* Header */}
<div className="flex items-start justify-between mb-3">
<div className="flex items-center gap-4">
{/* Status Icon */}
<div
className={`shrink-0 rounded-full p-2.5 ${
credential.is_disabled
? "bg-gray-100 dark:bg-gray-800"
: isHealthy
? "bg-green-100 dark:bg-green-900/30"
: "bg-red-100 dark:bg-red-900/30"
}`}
>
{credential.is_disabled ? (
<PowerOff className="h-5 w-5 text-gray-400" />
) : isHealthy ? (
<Heart className="h-5 w-5 text-green-600 dark:text-green-400" />
) : (
<HeartOff className="h-5 w-5 text-red-600 dark:text-red-400" />
)}
</div>
{/* Main Info */}
<div className="flex-1 min-w-0">
<div className="flex items-center gap-3 mb-2">
<div className="flex items-center gap-2 mb-1">
<h4 className="font-semibold text-base truncate">
{credential.name || `凭证 #${credential.uuid.slice(0, 12)}`}
{credential.name || `凭证 #${credential.uuid.slice(0, 8)}`}
</h4>
<span className="rounded-full bg-muted px-2.5 py-1 text-xs font-medium whitespace-nowrap">
<span className="rounded-full bg-muted px-2 py-0.5 text-xs font-medium">
{getCredentialTypeLabel(credential.credential_type)}
</span>
</div>
<p className="text-xs text-muted-foreground font-mono">
{credential.uuid.slice(0, 24)}...
<p className="text-xs text-muted-foreground font-mono truncate">
{credential.uuid}
</p>
</div>
<div className="flex items-center gap-2">
<div
className={`rounded-full p-2 ${
{/* Stats */}
<div className="hidden sm:flex items-center gap-6 shrink-0">
<div className="flex items-center gap-2">
<Activity className="h-4 w-4 text-blue-500" />
<div className="text-center">
<div className="text-xs text-muted-foreground">使用次数</div>
<div className="font-semibold">{credential.usage_count}</div>
</div>
</div>
<div className="flex items-center gap-2">
<AlertTriangle
className={`h-4 w-4 ${hasError ? "text-yellow-500" : "text-green-500"}`}
/>
<div className="text-center">
<div className="text-xs text-muted-foreground">错误次数</div>
<div className="font-semibold">{credential.error_count}</div>
</div>
</div>
<div className="flex items-center gap-2 text-muted-foreground">
<Clock className="h-4 w-4" />
<div>
<div className="text-xs">最后使用</div>
<div className="text-xs font-medium">
{formatDate(credential.last_used)}
</div>
</div>
</div>
</div>
{/* Health Check Info */}
{credential.last_health_check_time && (
<div className="hidden lg:block shrink-0 text-xs text-muted-foreground border-l pl-4">
<div>检查: {formatDate(credential.last_health_check_time)}</div>
{credential.last_health_check_model && (
<div className="text-primary">
({credential.last_health_check_model})
</div>
)}
</div>
)}
{/* Actions */}
<div className="flex items-center gap-1.5 shrink-0">
<button
onClick={onToggle}
className={`rounded-lg p-2 text-xs font-medium transition-colors ${
credential.is_disabled
? "bg-gray-100 dark:bg-gray-800"
: isHealthy
? "bg-green-100 dark:bg-green-900/30"
: "bg-red-100 dark:bg-red-900/30"
? "bg-green-100 text-green-700 hover:bg-green-200 dark:bg-green-900/30 dark:text-green-400"
: "bg-gray-100 text-gray-700 hover:bg-gray-200 dark:bg-gray-800 dark:text-gray-300"
}`}
title={credential.is_disabled ? "启用" : "禁用"}
>
{credential.is_disabled ? (
<PowerOff className="h-4 w-4 text-gray-400" />
) : isHealthy ? (
<Heart className="h-4 w-4 text-green-600 dark:text-green-400" />
<Power className="h-4 w-4" />
) : (
<HeartOff className="h-4 w-4 text-red-600 dark:text-red-400" />
<PowerOff className="h-4 w-4" />
)}
</div>
</div>
</div>
</button>
{/* Stats */}
<div className="bg-white/50 dark:bg-black/20 rounded-lg p-3 mb-3">
<div className="grid grid-cols-2 gap-3">
<div className="flex items-center gap-2">
<div className="rounded-full bg-blue-100 dark:bg-blue-900/30 p-1.5">
<Activity className="h-3 w-3 text-blue-600 dark:text-blue-400" />
</div>
<div>
<div className="text-xs text-muted-foreground">使用次数</div>
<div className="font-semibold text-sm">
{credential.usage_count}
</div>
</div>
</div>
<div className="flex items-center gap-2">
<div
className={`rounded-full p-1.5 ${
hasError
? "bg-yellow-100 dark:bg-yellow-900/30"
: "bg-green-100 dark:bg-green-900/30"
}`}
>
<AlertTriangle
className={`h-3 w-3 ${
hasError
? "text-yellow-600 dark:text-yellow-400"
: "text-green-600 dark:text-green-400"
}`}
/>
</div>
<div>
<div className="text-xs text-muted-foreground">错误次数</div>
<div className="font-semibold text-sm">
{credential.error_count}
</div>
</div>
</div>
<div className="col-span-2 flex items-center gap-2 pt-2 border-t border-border/50">
<Clock className="h-3 w-3 text-muted-foreground" />
<span className="text-xs text-muted-foreground">
最后使用: {formatDate(credential.last_used)}
</span>
</div>
</div>
</div>
<button
onClick={onEdit}
className="rounded-lg bg-blue-100 p-2 text-blue-700 hover:bg-blue-200 dark:bg-blue-900/30 dark:text-blue-400 transition-colors"
title="编辑"
>
<Settings className="h-4 w-4" />
</button>
{/* Health Check Info */}
{credential.last_health_check_time && (
<div className="mt-2 text-xs text-muted-foreground">
<span>
检查: {formatDate(credential.last_health_check_time)}
{credential.last_health_check_model &&
` (${credential.last_health_check_model})`}
</span>
</div>
)}
{/* OAuth Status */}
{isOAuth && credential.oauth_status && (
<div className="mb-3 rounded-lg border border-blue-200 dark:border-blue-800 bg-blue-50/50 dark:bg-blue-950/30 p-3">
<div className="flex items-center gap-2 mb-3">
<div className="rounded-full bg-blue-100 dark:bg-blue-900/30 p-1.5">
<Key className="h-3 w-3 text-blue-600 dark:text-blue-400" />
</div>
<span className="font-semibold text-sm">OAuth 状态</span>
</div>
<div className="space-y-2">
<div className="flex items-center justify-between">
<span className="text-xs text-muted-foreground">
Access Token
</span>
<div className="flex items-center gap-1">
{credential.oauth_status.has_access_token ? (
<CheckCircle className="h-4 w-4 text-green-500" />
) : (
<XCircle className="h-4 w-4 text-red-500" />
)}
<span className="text-xs font-medium">
{credential.oauth_status.has_access_token ? "有效" : "缺失"}
</span>
</div>
</div>
<div className="flex items-center justify-between">
<span className="text-xs text-muted-foreground">Token 状态</span>
<div className="flex items-center gap-1">
{credential.oauth_status.is_token_valid ? (
<CheckCircle className="h-4 w-4 text-green-500" />
) : (
<XCircle className="h-4 w-4 text-yellow-500" />
)}
<span className="text-xs font-medium">
{credential.oauth_status.is_token_valid ? "有效" : "需刷新"}
</span>
</div>
</div>
{credential.oauth_status.expiry_info && (
<div className="pt-2 border-t border-blue-200/50 dark:border-blue-800/50">
<span className="text-xs text-muted-foreground">
过期时间: {credential.oauth_status.expiry_info}
</span>
</div>
)}
</div>
</div>
)}
{/* Token Cache Status */}
{isOAuth && credential.token_cache_status && (
<div className="mb-3 rounded-lg border border-purple-200 dark:border-purple-800 bg-purple-50/50 dark:bg-purple-950/30 p-3">
<div className="flex items-center gap-2 mb-3">
<div className="rounded-full bg-purple-100 dark:bg-purple-900/30 p-1.5">
<Database className="h-3 w-3 text-purple-600 dark:text-purple-400" />
</div>
<span className="font-semibold text-sm">Token 缓存</span>
</div>
<div className="space-y-2">
<div className="flex items-center justify-between">
<span className="text-xs text-muted-foreground">缓存状态</span>
<div className="flex items-center gap-1">
{credential.token_cache_status.has_cached_token ? (
<CheckCircle className="h-4 w-4 text-green-500" />
) : (
<XCircle className="h-4 w-4 text-gray-400" />
)}
<span className="text-xs font-medium">
{credential.token_cache_status.has_cached_token
? "已缓存"
: "未缓存"}
</span>
</div>
</div>
<div className="flex items-center justify-between">
<span className="text-xs text-muted-foreground">有效性</span>
<div className="flex items-center gap-1">
{credential.token_cache_status.is_valid ? (
credential.token_cache_status.is_expiring_soon ? (
<AlertTriangle className="h-4 w-4 text-yellow-500" />
) : (
<CheckCircle className="h-4 w-4 text-green-500" />
)
) : (
<XCircle className="h-4 w-4 text-red-500" />
)}
<span className="text-xs font-medium">
{credential.token_cache_status.is_valid
? credential.token_cache_status.is_expiring_soon
? "即将过期"
: "有效"
: "已过期"}
</span>
</div>
</div>
{(credential.token_cache_status.last_refresh ||
credential.token_cache_status.expiry_time) && (
<div className="pt-2 border-t border-purple-200/50 dark:border-purple-800/50 space-y-1">
{credential.token_cache_status.last_refresh && (
<div className="text-xs text-muted-foreground">
最后刷新:{" "}
{formatDate(credential.token_cache_status.last_refresh)}
</div>
)}
{credential.token_cache_status.expiry_time && (
<div className="text-xs text-muted-foreground">
过期时间:{" "}
{formatDate(credential.token_cache_status.expiry_time)}
</div>
)}
</div>
)}
{credential.token_cache_status.refresh_error_count > 0 && (
<div className="pt-2 border-t border-red-200/50 dark:border-red-800/50">
<div className="flex items-center gap-2 text-red-600 dark:text-red-400">
<AlertTriangle className="h-3 w-3" />
<span className="text-xs font-medium">
刷新失败 {credential.token_cache_status.refresh_error_count}{" "}
次
</span>
</div>
{credential.token_cache_status.last_refresh_error && (
<div className="mt-1 text-xs text-red-600 dark:text-red-400 truncate">
{credential.token_cache_status.last_refresh_error.slice(
0,
60,
)}
{credential.token_cache_status.last_refresh_error.length >
60 && "..."}
</div>
)}
</div>
)}
</div>
</div>
)}
{/* Error Message */}
{credential.last_error_message && (
<div className="mt-2 rounded bg-red-100 p-2 text-xs text-red-700 dark:bg-red-900/30 dark:text-red-300">
{credential.last_error_message.slice(0, 100)}
{credential.last_error_message.length > 100 && "..."}
</div>
)}
{/* Actions */}
<div className="flex items-center gap-2 pt-4 border-t border-border/30">
<button
onClick={onToggle}
className={`flex items-center gap-1 rounded-lg px-3 py-2 text-xs font-medium transition-colors ${
credential.is_disabled
? "bg-green-100 text-green-700 hover:bg-green-200 dark:bg-green-900/30 dark:text-green-400 dark:hover:bg-green-800/40"
: "bg-gray-100 text-gray-700 hover:bg-gray-200 dark:bg-gray-800 dark:text-gray-300 dark:hover:bg-gray-700"
}`}
title={credential.is_disabled ? "启用凭证" : "禁用凭证"}
>
{credential.is_disabled ? (
<>
<Power className="h-3 w-3" />
启用
</>
) : (
<>
<PowerOff className="h-3 w-3" />
禁用
</>
)}
</button>
<button
onClick={onEdit}
className="flex items-center gap-1 rounded-lg bg-blue-100 px-3 py-2 text-xs font-medium text-blue-700 hover:bg-blue-200 dark:bg-blue-900/30 dark:text-blue-400 dark:hover:bg-blue-800/40 transition-colors"
title="编辑凭证配置"
>
<Settings className="h-3 w-3" />
编辑
</button>
<div className="flex items-center gap-1">
<button
onClick={onCheckHealth}
disabled={checkingHealth}
className="flex items-center gap-1 rounded-lg bg-emerald-100 px-3 py-2 text-xs font-medium text-emerald-700 hover:bg-emerald-200 disabled:opacity-50 dark:bg-emerald-900/30 dark:text-emerald-400 dark:hover:bg-emerald-800/40 transition-colors"
title="执行健康检测"
className="rounded-lg bg-emerald-100 p-2 text-emerald-700 hover:bg-emerald-200 disabled:opacity-50 dark:bg-emerald-900/30 dark:text-emerald-400 transition-colors"
title="检测"
>
<Activity
className={`h-3 w-3 ${checkingHealth ? "animate-pulse" : ""}`}
className={`h-4 w-4 ${checkingHealth ? "animate-pulse" : ""}`}
/>
检测
</button>
{isOAuth && onRefreshToken && (
<button
onClick={onRefreshToken}
disabled={refreshingToken}
className="flex items-center gap-1 rounded-lg bg-purple-100 px-3 py-2 text-xs font-medium text-purple-700 hover:bg-purple-200 disabled:opacity-50 dark:bg-purple-900/30 dark:text-purple-400 dark:hover:bg-purple-800/40 transition-colors"
title="刷新 OAuth Token"
className="rounded-lg bg-purple-100 p-2 text-purple-700 hover:bg-purple-200 disabled:opacity-50 dark:bg-purple-900/30 dark:text-purple-400 transition-colors"
title="刷新 Token"
>
<RefreshCw
className={`h-3 w-3 ${refreshingToken ? "animate-spin" : ""}`}
className={`h-4 w-4 ${refreshingToken ? "animate-spin" : ""}`}
/>
刷新
</button>
)}
<button
onClick={onReset}
className="flex items-center gap-1 rounded-lg bg-orange-100 px-3 py-2 text-xs font-medium text-orange-700 hover:bg-orange-200 dark:bg-orange-900/30 dark:text-orange-400 dark:hover:bg-orange-800/40 transition-colors"
title="重置统计计数器"
className="rounded-lg bg-orange-100 p-2 text-orange-700 hover:bg-orange-200 dark:bg-orange-900/30 dark:text-orange-400 transition-colors"
title="重置"
>
<RotateCcw className="h-3 w-3" />
重置
<RotateCcw className="h-4 w-4" />
</button>
<button
onClick={onDelete}
disabled={deleting}
className="flex items-center gap-1 rounded-lg bg-red-100 px-3 py-2 text-xs font-medium text-red-700 hover:bg-red-200 disabled:opacity-50 dark:bg-red-900/30 dark:text-red-400 dark:hover:bg-red-800/40 transition-colors"
title="删除凭证"
className="rounded-lg bg-red-100 p-2 text-red-700 hover:bg-red-200 disabled:opacity-50 dark:bg-red-900/30 dark:text-red-400 transition-colors"
title="删除"
>
<Trash2 className="h-3 w-3" />
删除
<Trash2 className="h-4 w-4" />
</button>
</div>
</div>
{/* Mobile Stats - shown on small screens */}
<div className="sm:hidden mt-3 pt-3 border-t border-border/30">
<div className="flex items-center justify-between text-xs">
<div className="flex items-center gap-4">
<span className="flex items-center gap-1">
<Activity className="h-3 w-3 text-blue-500" />
使用: {credential.usage_count}
</span>
<span className="flex items-center gap-1">
<AlertTriangle
className={`h-3 w-3 ${hasError ? "text-yellow-500" : "text-green-500"}`}
/>
错误: {credential.error_count}
</span>
</div>
<span className="text-muted-foreground">
<Clock className="h-3 w-3 inline mr-1" />
{formatDate(credential.last_used)}
</span>
</div>
</div>
{/* Error Message */}
{credential.last_error_message && (
<div className="mt-3 rounded-lg bg-red-100 p-2 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>
)}
</div>
);
}
@@ -1,9 +1,18 @@
import { useState, useEffect } from "react";
import { X, Eye, EyeOff, Settings, FolderOpen, Upload, CheckCircle } from "lucide-react";
import {
X,
Eye,
EyeOff,
Settings,
Upload,
CheckCircle,
Ban,
} from "lucide-react";
import { open } from "@tauri-apps/plugin-dialog";
import {
CredentialDisplay,
UpdateCredentialRequest,
PoolProviderType,
} from "@/lib/api/providerPool";
interface EditCredentialModalProps {
@@ -13,6 +22,37 @@ interface EditCredentialModalProps {
onEdit: (uuid: string, request: UpdateCredentialRequest) => Promise<void>;
}
// 各 Provider 支持的模型列表 (参考 AIClient-2-API/src/provider-models.js)
const providerModels: Record<PoolProviderType, string[]> = {
kiro: [
"claude-opus-4-5",
"claude-opus-4-5-20251101",
"claude-haiku-4-5",
"claude-sonnet-4-5",
"claude-sonnet-4-5-20250929",
"claude-sonnet-4-20250514",
"claude-3-7-sonnet-20250219",
],
gemini: [
"gemini-2.5-flash",
"gemini-2.5-flash-lite",
"gemini-2.5-pro",
"gemini-2.5-pro-preview-06-05",
"gemini-2.5-flash-preview-09-2025",
"gemini-3-pro-preview",
],
qwen: ["qwen3-coder-plus", "qwen3-coder-flash"],
antigravity: [
"gemini-3-pro-preview",
"gemini-3-pro-image-preview",
"gemini-2.5-computer-use-preview-10-2025",
"gemini-claude-sonnet-4-5",
"gemini-claude-sonnet-4-5-thinking",
],
openai: [], // 自定义 API,无预设模型
claude: [], // 自定义 API,无预设模型
};
export function EditCredentialModal({
credential,
isOpen,
@@ -22,7 +62,7 @@ export function EditCredentialModal({
const [name, setName] = useState("");
const [checkHealth, setCheckHealth] = useState(true);
const [checkModelName, setCheckModelName] = useState("");
const [notSupportedModelsText, setNotSupportedModelsText] = useState("");
const [notSupportedModels, setNotSupportedModels] = useState<string[]>([]);
const [loading, setLoading] = useState(false);
const [error, setError] = useState<string | null>(null);
const [showCredentialDetails, setShowCredentialDetails] = useState(false);
@@ -37,9 +77,7 @@ export function EditCredentialModal({
setName(credential.name || "");
setCheckHealth(credential.check_health);
setCheckModelName(credential.check_model_name || "");
setNotSupportedModelsText(
(credential.not_supported_models || []).join(", "),
);
setNotSupportedModels(credential.not_supported_models || []);
setNewCredFilePath("");
setNewProjectId("");
setError(null);
@@ -52,6 +90,18 @@ export function EditCredentialModal({
const isOAuth = credential.credential_type.includes("oauth");
// 获取当前 provider 类型
const getProviderType = (): PoolProviderType => {
if (credential.credential_type.includes("kiro")) return "kiro";
if (credential.credential_type.includes("gemini")) return "gemini";
if (credential.credential_type.includes("qwen")) return "qwen";
if (credential.credential_type.includes("openai")) return "openai";
if (credential.credential_type.includes("claude")) return "claude";
return "kiro";
};
const currentProviderModels = providerModels[getProviderType()] || [];
const handleSelectNewFile = async () => {
try {
const selected = await open({
@@ -68,7 +118,6 @@ export function EditCredentialModal({
const getMaskedCredentialInfo = () => {
if (isOAuth) {
// OAuth 凭证显示文件路径(部分遮罩)
const path = credential.display_credential;
const parts = path.split("/");
if (parts.length > 1) {
@@ -78,30 +127,27 @@ export function EditCredentialModal({
}
return `***${path.slice(-12)}`;
} else {
// API Key 显示遮罩
return credential.display_credential;
}
};
const toggleModelSupport = (model: string) => {
setNotSupportedModels((prev) =>
prev.includes(model) ? prev.filter((m) => m !== model) : [...prev, model],
);
};
const handleSubmit = async () => {
setLoading(true);
setError(null);
try {
// 解析不支持的模型列表
const parsedNotSupportedModels = notSupportedModelsText
.split(",")
.map((model) => model.trim())
.filter((model) => model.length > 0);
const updateRequest: UpdateCredentialRequest = {
name: name.trim() || undefined,
check_health: checkHealth,
check_model_name: checkModelName.trim() || undefined,
not_supported_models:
parsedNotSupportedModels.length > 0
? parsedNotSupportedModels
: undefined,
// 始终传递 not_supported_models,即使为空数组(用于清除选择)
not_supported_models: notSupportedModels,
new_creds_file_path: newCredFilePath.trim() || undefined,
new_project_id: newProjectId.trim() || undefined,
};
@@ -117,9 +163,9 @@ export function EditCredentialModal({
return (
<div className="fixed inset-0 z-50 flex items-center justify-center bg-black/50 p-4">
<div className="w-full max-w-2xl h-[80vh] rounded-lg bg-background shadow-xl flex flex-col">
<div className="w-full max-w-2xl max-h-[85vh] rounded-lg bg-background shadow-xl flex flex-col">
{/* Header */}
<div className="flex items-center justify-between border-b pb-4 px-6 pt-6">
<div className="flex items-center justify-between border-b pb-4 px-6 pt-6 shrink-0">
<h3 className="text-lg font-semibold flex items-center gap-2">
<Settings className="h-5 w-5" />
编辑凭证
@@ -131,204 +177,189 @@ export function EditCredentialModal({
{/* Content - Scrollable */}
<div className="flex-1 overflow-y-auto px-6 py-4">
<div className="space-y-4">
{/* 凭证信息(只读) */}
<div className="rounded-lg bg-muted/50 p-3">
<div className="flex items-center justify-between mb-2">
<label className="text-sm font-medium">凭证信息</label>
<button
type="button"
onClick={() => setShowCredentialDetails(!showCredentialDetails)}
className="flex items-center gap-1 text-xs text-muted-foreground hover:text-foreground"
>
{showCredentialDetails ? (
<>
<EyeOff className="h-3 w-3" />
隐藏
</>
) : (
<>
<Eye className="h-3 w-3" />
显示
</>
)}
</button>
</div>
<div className="space-y-2 text-sm">
<div className="flex items-center justify-between">
<span className="text-muted-foreground">类型:</span>
<span className="font-mono">{credential.credential_type}</span>
</div>
<div className="flex items-center justify-between">
<span className="text-muted-foreground">UUID:</span>
<span className="font-mono">
{credential.uuid.slice(0, 24)}...
</span>
</div>
<div className="flex items-center justify-between">
<span className="text-muted-foreground">
{isOAuth ? "文件路径:" : "API Key:"}
</span>
<span className="font-mono">
{showCredentialDetails
? credential.display_credential
: getMaskedCredentialInfo()}
</span>
</div>
</div>
<p className="mt-2 text-xs text-muted-foreground">
🔒 敏感信息(API Key、文件路径)无法修改,如需更改请删除后重新添加
</p>
</div>
{/* 可编辑字段 */}
<div>
<label className="mb-1 block text-sm font-medium">名称</label>
<input
type="text"
value={name}
onChange={(e) => setName(e.target.value)}
placeholder="给这个凭证起个名字..."
className="w-full rounded-lg border bg-background px-3 py-2 text-sm"
/>
</div>
{/* 健康检查设置 */}
<div>
<label className="mb-2 flex items-center gap-2 text-sm font-medium">
<input
type="checkbox"
checked={checkHealth}
onChange={(e) => setCheckHealth(e.target.checked)}
className="rounded"
/>
启用自动健康检查
</label>
{checkHealth && (
<div className="ml-6">
<label className="mb-1 block text-xs font-medium text-muted-foreground">
检查模型(可选)
<div className="space-y-5">
{/* 名称 + 健康检查 */}
<div className="grid grid-cols-2 gap-4">
<div>
<label className="mb-1 block text-sm font-medium">
名称 (选填)
</label>
<input
type="text"
value={checkModelName}
onChange={(e) => setCheckModelName(e.target.value)}
placeholder="留空使用默认模型..."
value={name}
onChange={(e) => setName(e.target.value)}
placeholder="给这个凭证起个名字..."
className="w-full rounded-lg border bg-background px-3 py-2 text-sm"
/>
</div>
)}
</div>
{/* 不支持的模型列表 */}
<div>
<label className="mb-1 block text-sm font-medium">
不支持的模型
</label>
<textarea
value={notSupportedModelsText}
onChange={(e) => setNotSupportedModelsText(e.target.value)}
placeholder="用逗号分隔多个模型,例如: model-1, model-2"
className="w-full rounded-lg border bg-background px-3 py-2 text-sm"
rows={3}
/>
<p className="mt-1 text-xs text-muted-foreground">
这些模型将不会路由到此凭证
</p>
</div>
{/* OAuth 文件重新上传 */}
{isOAuth && (
<div className="rounded-lg border border-amber-200 dark:border-amber-800 bg-amber-50/50 dark:bg-amber-950/30 p-4">
<div className="flex items-center gap-2 mb-3">
<div className="rounded-full bg-amber-100 dark:bg-amber-900/30 p-1.5">
<Upload className="h-3 w-3 text-amber-600 dark:text-amber-400" />
</div>
<span className="font-semibold text-sm">重新上传凭证文件</span>
</div>
<p className="text-xs text-muted-foreground mb-3">
选择新的凭证文件来替换当前文件。新文件将被复制到应用存储目录。
</p>
<div className="space-y-3">
<div>
<label className="mb-1 block text-xs font-medium text-muted-foreground">
新凭证文件
</label>
<div className="flex items-center gap-2">
<input
type="text"
value={newCredFilePath}
onChange={(e) => setNewCredFilePath(e.target.value)}
placeholder="选择新的凭证文件..."
className="flex-1 rounded-lg border bg-background px-3 py-2 text-sm"
readOnly
/>
<button
type="button"
onClick={handleSelectNewFile}
className="flex items-center gap-1 rounded-lg bg-blue-100 px-3 py-2 text-xs font-medium text-blue-700 hover:bg-blue-200 dark:bg-blue-900/30 dark:text-blue-400 dark:hover:bg-blue-800/40 transition-colors"
>
<FolderOpen className="h-3 w-3" />
选择文件
</button>
</div>
</div>
{credential.credential_type === "gemini_oauth" && (
<div>
<label className="mb-1 block text-xs font-medium text-muted-foreground">
项目ID(可选)
</label>
<input
type="text"
value={newProjectId}
onChange={(e) => setNewProjectId(e.target.value)}
placeholder="留空保持当前项目ID..."
className="w-full rounded-lg border bg-background px-3 py-2 text-sm"
/>
</div>
)}
{newCredFilePath && (
<div className="text-xs text-green-600 dark:text-green-400 flex items-center gap-1">
<CheckCircle className="h-3 w-3" />
文件已选择,保存后将替换当前凭证文件
</div>
)}
<div>
<label className="mb-1 block text-sm font-medium">
健康检查
</label>
<select
value={checkHealth ? "enabled" : "disabled"}
onChange={(e) => setCheckHealth(e.target.value === "enabled")}
className="w-full rounded-lg border bg-background px-3 py-2 text-sm"
>
<option value="enabled">启用</option>
<option value="disabled">禁用</option>
</select>
</div>
</div>
)}
{/* 统计信息(只读) */}
<div className="rounded-lg bg-muted/50 p-3">
<label className="mb-2 block text-sm font-medium">使用统计</label>
<div className="grid grid-cols-2 gap-3 text-sm">
<div className="flex justify-between">
<span className="text-muted-foreground">使用次数:</span>
<span className="font-mono">{credential.usage_count}</span>
{/* 检查模型名称 */}
<div>
<label className="mb-1 block text-sm font-medium">
检查模型名称 (选填)
</label>
<input
type="text"
value={checkModelName}
onChange={(e) => setCheckModelName(e.target.value)}
placeholder="用于健康检查的模型名称..."
className="w-full rounded-lg border bg-background px-3 py-2 text-sm"
/>
</div>
{/* OAuth凭据文件路径 */}
{isOAuth && (
<div>
<label className="mb-1 block text-sm font-medium">
OAuth凭据文件路径
</label>
<div className="flex items-center gap-2">
<input
type="text"
value={
showCredentialDetails
? credential.display_credential
: getMaskedCredentialInfo()
}
readOnly
className="flex-1 rounded-lg border bg-muted/50 px-3 py-2 text-sm text-muted-foreground"
/>
<button
type="button"
onClick={() =>
setShowCredentialDetails(!showCredentialDetails)
}
className="rounded-lg border p-2 hover:bg-muted"
title={showCredentialDetails ? "隐藏" : "显示"}
>
{showCredentialDetails ? (
<EyeOff className="h-4 w-4" />
) : (
<Eye className="h-4 w-4" />
)}
</button>
<button
type="button"
onClick={handleSelectNewFile}
className="rounded-lg border p-2 hover:bg-muted"
title="上传新文件"
>
<Upload className="h-4 w-4" />
</button>
</div>
{newCredFilePath && (
<div className="mt-2 text-xs text-green-600 dark:text-green-400 flex items-center gap-1">
<CheckCircle className="h-3 w-3" />
新文件已选择: {newCredFilePath.split("/").pop()}
</div>
)}
</div>
<div className="flex justify-between">
<span className="text-muted-foreground">错误次数:</span>
<span className="font-mono">{credential.error_count}</span>
</div>
<div className="col-span-2 flex justify-between">
<span className="text-muted-foreground">最后使用:</span>
<span className="text-xs">
{credential.last_used || "从未"}
)}
{/* Gemini Project ID */}
{credential.credential_type === "gemini_oauth" &&
newCredFilePath && (
<div>
<label className="mb-1 block text-sm font-medium">
项目ID(可选)
</label>
<input
type="text"
value={newProjectId}
onChange={(e) => setNewProjectId(e.target.value)}
placeholder="留空保持当前项目ID..."
className="w-full rounded-lg border bg-background px-3 py-2 text-sm"
/>
</div>
)}
{/* 不支持的模型 - Checkbox Grid */}
<div>
<div className="flex items-center gap-2 mb-3">
<Ban className="h-4 w-4 text-muted-foreground" />
<label className="text-sm font-medium">不支持的模型</label>
<span className="text-xs text-muted-foreground">
选择此提供商不支持的模型,系统会自动排除这些模型
</span>
</div>
<div className="grid grid-cols-2 sm:grid-cols-3 gap-2">
{currentProviderModels.map((model) => (
<label
key={model}
className={`flex items-center gap-2 rounded-lg border px-3 py-2 cursor-pointer transition-colors ${
notSupportedModels.includes(model)
? "border-red-300 bg-red-50 dark:border-red-800 dark:bg-red-950/30"
: "border-border hover:bg-muted/50"
}`}
>
<input
type="checkbox"
checked={notSupportedModels.includes(model)}
onChange={() => toggleModelSupport(model)}
className="rounded border-gray-300"
/>
<span className="text-sm truncate">{model}</span>
</label>
))}
</div>
</div>
</div>
{/* Error */}
{error && (
<div className="rounded-lg border border-red-500 bg-red-50 p-3 text-sm text-red-700 dark:bg-red-950/30">
{error}
{/* 使用统计(只读) */}
<div className="rounded-lg bg-muted/50 p-4">
<label className="mb-3 block text-sm font-medium">使用统计</label>
<div className="grid grid-cols-3 gap-4 text-sm">
<div>
<span className="text-muted-foreground block text-xs">
使用次数
</span>
<span className="font-semibold">
{credential.usage_count}
</span>
</div>
<div>
<span className="text-muted-foreground block text-xs">
错误次数
</span>
<span className="font-semibold">
{credential.error_count}
</span>
</div>
<div>
<span className="text-muted-foreground block text-xs">
最后使用
</span>
<span className="text-xs">
{credential.last_used || "从未"}
</span>
</div>
</div>
</div>
)}
{/* Error */}
{error && (
<div className="rounded-lg border border-red-500 bg-red-50 p-3 text-sm text-red-700 dark:bg-red-950/30">
{error}
</div>
)}
</div>
</div>
{/* Footer */}
<div className="border-t px-6 py-4 flex justify-end gap-2">
<div className="border-t px-6 py-4 flex justify-end gap-2 shrink-0">
<button
onClick={onClose}
className="rounded-lg border px-4 py-2 text-sm hover:bg-muted"
+56 -26
View File
@@ -1,10 +1,24 @@
import { useState, useEffect } from "react";
import { AlertTriangle, X, RotateCcw, Trash2, Settings, CheckCircle2 } from "lucide-react";
import {
AlertTriangle,
X,
RotateCcw,
Trash2,
Settings,
CheckCircle2,
} from "lucide-react";
export interface ErrorInfo {
id: string;
message: string;
type: "delete" | "toggle" | "reset" | "health_check" | "refresh_token" | "general" | "success";
type:
| "delete"
| "toggle"
| "reset"
| "health_check"
| "refresh_token"
| "general"
| "success";
uuid?: string; // 相关凭证的UUID(如果有的话)
}
@@ -19,47 +33,51 @@ const ErrorTypeConfig = {
icon: Trash2,
color: "text-red-600 dark:text-red-400",
bgColor: "bg-red-50 dark:bg-red-950/30",
borderColor: "border-red-200 dark:border-red-800"
borderColor: "border-red-200 dark:border-red-800",
},
toggle: {
icon: Settings,
color: "text-blue-600 dark:text-blue-400",
bgColor: "bg-blue-50 dark:bg-blue-950/30",
borderColor: "border-blue-200 dark:border-blue-800"
borderColor: "border-blue-200 dark:border-blue-800",
},
reset: {
icon: RotateCcw,
color: "text-orange-600 dark:text-orange-400",
bgColor: "bg-orange-50 dark:bg-orange-950/30",
borderColor: "border-orange-200 dark:border-orange-800"
borderColor: "border-orange-200 dark:border-orange-800",
},
health_check: {
icon: AlertTriangle,
color: "text-yellow-600 dark:text-yellow-400",
bgColor: "bg-yellow-50 dark:bg-yellow-950/30",
borderColor: "border-yellow-200 dark:border-yellow-800"
borderColor: "border-yellow-200 dark:border-yellow-800",
},
refresh_token: {
icon: RotateCcw,
color: "text-purple-600 dark:text-purple-400",
bgColor: "bg-purple-50 dark:bg-purple-950/30",
borderColor: "border-purple-200 dark:border-purple-800"
borderColor: "border-purple-200 dark:border-purple-800",
},
general: {
icon: AlertTriangle,
color: "text-gray-600 dark:text-gray-400",
bgColor: "bg-gray-50 dark:bg-gray-950/30",
borderColor: "border-gray-200 dark:border-gray-800"
borderColor: "border-gray-200 dark:border-gray-800",
},
success: {
icon: CheckCircle2,
color: "text-green-600 dark:text-green-400",
bgColor: "bg-green-50 dark:bg-green-950/30",
borderColor: "border-green-200 dark:border-green-800"
}
borderColor: "border-green-200 dark:border-green-800",
},
};
function ErrorItem({ error, onDismiss, onRetry }: {
function ErrorItem({
error,
onDismiss,
onRetry,
}: {
error: ErrorInfo;
onDismiss: (id: string) => void;
onRetry?: (error: ErrorInfo) => void;
@@ -68,7 +86,9 @@ function ErrorItem({ error, onDismiss, onRetry }: {
const IconComponent = config.icon;
return (
<div className={`rounded-lg border p-4 ${config.bgColor} ${config.borderColor}`}>
<div
className={`rounded-lg border p-4 ${config.bgColor} ${config.borderColor}`}
>
<div className="flex items-start gap-3">
<IconComponent className={`h-5 w-5 mt-0.5 ${config.color}`} />
<div className="flex-1 min-w-0">
@@ -99,12 +119,16 @@ function ErrorItem({ error, onDismiss, onRetry }: {
);
}
export function ErrorDisplay({ errors, onDismiss, onRetry }: ErrorDisplayProps) {
export function ErrorDisplay({
errors,
onDismiss,
onRetry,
}: ErrorDisplayProps) {
// 自动关闭通知
useEffect(() => {
const timers: ReturnType<typeof setTimeout>[] = [];
errors.forEach(error => {
errors.forEach((error) => {
// 成功消息 3 秒后自动关闭,其他类型 15 秒后自动关闭
if (error.type === "success") {
const timer = setTimeout(() => {
@@ -120,7 +144,7 @@ export function ErrorDisplay({ errors, onDismiss, onRetry }: ErrorDisplayProps)
});
return () => {
timers.forEach(timer => clearTimeout(timer));
timers.forEach((timer) => clearTimeout(timer));
};
}, [errors, onDismiss]);
@@ -131,7 +155,7 @@ export function ErrorDisplay({ errors, onDismiss, onRetry }: ErrorDisplayProps)
return (
<div className="fixed top-4 right-4 z-50 w-96 max-w-full">
<div className="space-y-3 max-h-96 overflow-y-auto">
{errors.map(error => (
{errors.map((error) => (
<ErrorItem
key={error.id}
error={error}
@@ -149,20 +173,26 @@ export function ErrorDisplay({ errors, onDismiss, onRetry }: ErrorDisplayProps)
export function useErrorDisplay() {
const [errors, setErrors] = useState<ErrorInfo[]>([]);
const showError = (message: string, type: ErrorInfo["type"] = "general", uuid?: string) => {
const showError = (
message: string,
type: ErrorInfo["type"] = "general",
uuid?: string,
) => {
// 检查是否已经存在相同的错误消息(基于 message, type, uuid 的组合)
setErrors(prev => {
const isDuplicate = prev.some(existing =>
existing.message === message &&
existing.type === type &&
existing.uuid === uuid
setErrors((prev) => {
const isDuplicate = prev.some(
(existing) =>
existing.message === message &&
existing.type === type &&
existing.uuid === uuid,
);
if (isDuplicate) {
return prev; // 如果重复,不添加新的错误
}
const id = Date.now().toString() + Math.random().toString(36).substr(2, 9);
const id =
Date.now().toString() + Math.random().toString(36).substr(2, 9);
const error: ErrorInfo = { id, message, type, uuid };
return [...prev, error];
});
@@ -171,11 +201,11 @@ export function useErrorDisplay() {
const showSuccess = (message: string, uuid?: string) => {
const id = Date.now().toString() + Math.random().toString(36).substr(2, 9);
const info: ErrorInfo = { id, message, type: "success", uuid };
setErrors(prev => [...prev, info]);
setErrors((prev) => [...prev, info]);
};
const dismissError = (id: string) => {
setErrors(prev => prev.filter(error => error.id !== id));
setErrors((prev) => prev.filter((error) => error.id !== id));
};
const clearErrors = () => {
@@ -189,4 +219,4 @@ export function useErrorDisplay() {
dismissError,
clearErrors,
};
}
}
@@ -27,6 +27,7 @@ const allProviderTypes: PoolProviderType[] = [
"kiro",
"gemini",
"qwen",
"antigravity",
"openai",
"claude",
];
@@ -35,6 +36,7 @@ const providerLabels: Record<PoolProviderType, string> = {
kiro: "Kiro (AWS)",
gemini: "Gemini (Google)",
qwen: "Qwen (阿里)",
antigravity: "Antigravity (Gemini 3 Pro)",
openai: "OpenAI",
claude: "Claude (Anthropic)",
};
@@ -92,7 +94,11 @@ export const ProviderPoolPage = forwardRef<ProviderPoolPageRef>(
try {
await toggleCredential(credential.uuid, !credential.is_disabled);
} catch (e) {
showError(e instanceof Error ? e.message : String(e), "toggle", credential.uuid);
showError(
e instanceof Error ? e.message : String(e),
"toggle",
credential.uuid,
);
}
};
@@ -113,7 +119,11 @@ export const ProviderPoolPage = forwardRef<ProviderPoolPageRef>(
showError(result.message || "健康检查未通过", "health_check", uuid);
}
} catch (e) {
showError(e instanceof Error ? e.message : String(e), "health_check", uuid);
showError(
e instanceof Error ? e.message : String(e),
"health_check",
uuid,
);
}
};
@@ -138,7 +148,11 @@ export const ProviderPoolPage = forwardRef<ProviderPoolPageRef>(
await refreshCredentialToken(uuid);
showSuccess("Token 刷新成功!", uuid);
} catch (e) {
showError(e instanceof Error ? e.message : String(e), "refresh_token", uuid);
showError(
e instanceof Error ? e.message : String(e),
"refresh_token",
uuid,
);
}
};
@@ -302,7 +316,7 @@ export const ProviderPoolPage = forwardRef<ProviderPoolPageRef>(
</button>
</div>
) : (
<div className="grid grid-cols-1 gap-4 md:grid-cols-2 lg:grid-cols-3">
<div className="flex flex-col gap-4">
{currentCredentials.map((credential) => (
<CredentialCard
key={credential.uuid}
+26 -1
View File
@@ -1,7 +1,13 @@
import { invoke } from "@tauri-apps/api/core";
// Provider types supported by the pool
export type PoolProviderType = "kiro" | "gemini" | "qwen" | "openai" | "claude";
export type PoolProviderType =
| "kiro"
| "gemini"
| "qwen"
| "antigravity"
| "openai"
| "claude";
// Credential data types
export interface KiroOAuthCredential {
@@ -20,6 +26,12 @@ export interface QwenOAuthCredential {
creds_file_path: string;
}
export interface AntigravityOAuthCredential {
type: "antigravity_oauth";
creds_file_path: string;
project_id?: string;
}
export interface OpenAIKeyCredential {
type: "openai_key";
api_key: string;
@@ -36,6 +48,7 @@ export type CredentialData =
| KiroOAuthCredential
| GeminiOAuthCredential
| QwenOAuthCredential
| AntigravityOAuthCredential
| OpenAIKeyCredential
| ClaudeKeyCredential;
@@ -259,6 +272,18 @@ export const providerPoolApi = {
return invoke("add_claude_key_credential", { apiKey, baseUrl, name });
},
async addAntigravityOAuth(
credsFilePath: string,
projectId?: string,
name?: string,
): Promise<ProviderCredential> {
return invoke("add_antigravity_oauth_credential", {
credsFilePath,
projectId,
name,
});
},
// OAuth token management
async refreshCredentialToken(uuid: string): Promise<string> {
return invoke("refresh_pool_credential_token", { uuid });