diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 00d521481..c6e925bfa 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -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 diff --git a/.husky/pre-commit b/.husky/pre-commit new file mode 100755 index 000000000..c1fc2d618 --- /dev/null +++ b/.husky/pre-commit @@ -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!" diff --git a/package-lock.json b/package-lock.json index a6be94df8..27f852c39 100644 --- a/package-lock.json +++ b/package-lock.json @@ -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", diff --git a/package.json b/package.json index 202416934..bd91bb0f3 100644 --- a/package.json +++ b/package.json @@ -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", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index b91db6209..4df53c5d0 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -3153,7 +3153,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.6.1" +version = "0.7.0" dependencies = [ "anyhow", "async-stream", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index de31d548e..dda7fceec 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -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" diff --git a/src-tauri/src/commands/provider_pool_cmd.rs b/src-tauri/src/commands/provider_pool_cmd.rs index b2bf6d267..7a887b2dd 100644 --- a/src-tauri/src/commands/provider_pool_cmd.rs +++ b/src-tauri/src/commands/provider_pool_cmd.rs @@ -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); @@ -35,8 +35,7 @@ fn get_credentials_dir() -> Result { // 确保目录存在 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 { /// 复制并重命名 OAuth 凭证文件 fn copy_and_rename_credential_file( source_path: &str, - provider_type: &str + provider_type: &str, ) -> Result { 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 { - 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, + name: Option, +) -> Result { + // 复制并重命名文件到应用存储目录 + 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 { 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 { 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 { #[tauri::command] pub async fn test_user_credentials() -> Result { 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 { // 测试用户上传的凭证文件路径 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 { 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 { .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 { Ok(hash_content) => { match serde_json::from_str::(&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 { // 设置凭证路径到用户文件 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 { 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)); diff --git a/src-tauri/src/commands/route_cmd.rs b/src-tauri/src/commands/route_cmd.rs index a4a72140e..4d80992ef 100644 --- a/src-tauri/src/commands/route_cmd.rs +++ b/src-tauri/src/commands/route_cmd.rs @@ -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; diff --git a/src-tauri/src/converter/mod.rs b/src-tauri/src/converter/mod.rs index cc9e44361..0457c28b3 100644 --- a/src-tauri/src/converter/mod.rs +++ b/src-tauri/src/converter/mod.rs @@ -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::*; diff --git a/src-tauri/src/converter/openai_to_antigravity.rs b/src-tauri/src/converter/openai_to_antigravity.rs new file mode 100644 index 000000000..26a7ec977 --- /dev/null +++ b/src-tauri/src/converter/openai_to_antigravity.rs @@ -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, + #[serde(skip_serializing_if = "Option::is_none")] + pub inline_data: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub function_call: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub function_response: Option, +} + +#[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, +} + +/// Antigravity/Gemini 工具定义 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct GeminiTool { + pub function_declarations: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct GeminiFunctionDeclaration { + pub name: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub description: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub parameters: Option, +} + +/// Antigravity/Gemini 生成配置 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct GeminiGenerationConfig { + #[serde(skip_serializing_if = "Option::is_none")] + pub temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_output_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub top_p: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub top_k: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub stop_sequences: Option>, +} + +/// Antigravity 请求体 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct AntigravityRequestBody { + pub contents: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub system_instruction: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub generation_config: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub tools: Option>, +} + +/// 将 OpenAI ChatCompletionRequest 转换为 Antigravity 请求体 +pub fn convert_openai_to_antigravity(request: &ChatCompletionRequest) -> serde_json::Value { + let mut contents: Vec = Vec::new(); + let mut system_instruction: Option = 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 { + 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 { + 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 = 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 +} diff --git a/src-tauri/src/converter/protocol_selector.rs b/src-tauri/src/converter/protocol_selector.rs new file mode 100644 index 000000000..d6fd2b6a9 --- /dev/null +++ b/src-tauri/src/converter/protocol_selector.rs @@ -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 { + // 大多数情况下,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 { + // 所有 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); + } +} diff --git a/src-tauri/src/database/dao/provider_pool.rs b/src-tauri/src/database/dao/provider_pool.rs index 5b8fa9064..b8094d50f 100644 --- a/src-tauri/src/database/dao/provider_pool.rs +++ b/src-tauri/src/database/dao/provider_pool.rs @@ -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 = 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, diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 1feb33255..f32d9c7f7 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -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, diff --git a/src-tauri/src/models/mod.rs b/src-tauri/src/models/mod.rs index 4264ac3e0..90fb4f251 100644 --- a/src-tauri/src/models/mod.rs +++ b/src-tauri/src/models/mod.rs @@ -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}; diff --git a/src-tauri/src/models/provider_pool_model.rs b/src-tauri/src/models/provider_pool_model.rs index 37be2fcbb..6a687347d 100644 --- a/src-tauri/src/models/provider_pool_model.rs +++ b/src-tauri/src/models/provider_pool_model.rs @@ -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, }, /// Qwen OAuth 凭证(文件路径) - QwenOAuth { + QwenOAuth { creds_file_path: String }, + /// Antigravity OAuth 凭证(文件路径)- Google 内部 Gemini 3 Pro + AntigravityOAuth { creds_file_path: String, + project_id: Option, }, /// 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, + pub not_supported_models: Vec, pub usage_count: u64, pub error_count: u32, pub last_used: Option, + pub last_error_time: Option, pub last_error_message: Option, pub last_health_check_time: Option, pub last_health_check_model: Option, pub oauth_status: Option, pub token_cache_status: Option, + 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 { 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(), } } } diff --git a/src-tauri/src/models/route_model.rs b/src-tauri/src/models/route_model.rs index f6363105b..528f7fb7d 100644 --- a/src-tauri/src/models/route_model.rs +++ b/src-tauri/src/models/route_model.rs @@ -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, }; diff --git a/src-tauri/src/models/skill_model.rs b/src-tauri/src/models/skill_model.rs index a36139421..9c9e70114 100644 --- a/src-tauri/src/models/skill_model.rs +++ b/src-tauri/src/models/skill_model.rs @@ -50,6 +50,7 @@ impl Default for SkillRepo { } } +#[allow(dead_code)] impl SkillRepo { pub fn new(owner: String, name: String, branch: String) -> Self { Self { diff --git a/src-tauri/src/providers/antigravity.rs b/src-tauri/src/providers/antigravity.rs new file mode 100644 index 000000000..169cb63a0 --- /dev/null +++ b/src-tauri/src/providers/antigravity.rs @@ -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, + pub refresh_token: Option, + pub token_type: Option, + pub expiry_date: Option, + pub scope: Option, +} + +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, + pub client: Client, + pub base_urls: Vec, + pub available_models: Vec, +} + +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> { + 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> { + 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> { + 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> { + 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(¶ms) + .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> { + 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> { + let mut last_error: Option> = 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> { + 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, Box> { + 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> { + 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) + } +} diff --git a/src-tauri/src/providers/kiro.rs b/src-tauri/src/providers/kiro.rs index b6f7e44e3..6b2e93ea6 100644 --- a/src-tauri/src/providers/kiro.rs +++ b/src-tauri/src/providers/kiro.rs @@ -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::(&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::(&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::(&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 { - 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> { // 使用加载时的路径或默认路径 - 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) diff --git a/src-tauri/src/providers/mod.rs b/src-tauri/src/providers/mod.rs index 441186c95..324db9bcb 100644 --- a/src-tauri/src/providers/mod.rs +++ b/src-tauri/src/providers/mod.rs @@ -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)] diff --git a/src-tauri/src/router/mod.rs b/src-tauri/src/router/mod.rs new file mode 100644 index 000000000..4dd306a78 --- /dev/null +++ b/src-tauri/src/router/mod.rs @@ -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}; diff --git a/src-tauri/src/router/provider_router.rs b/src-tauri/src/router/provider_router.rs new file mode 100644 index 000000000..c302eb951 --- /dev/null +++ b/src-tauri/src/router/provider_router.rs @@ -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, +} + +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 { + self.route + .provider_type + .as_ref() + .and_then(|s| s.parse().ok()) + } +} + +/// Provider 路由器 +pub struct ProviderRouter { + /// 路由注册表 + registry: Arc>, +} + +impl ProviderRouter { + /// 创建新的路由器 + pub fn new(registry: Arc>) -> 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 { + 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 { + 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())); + } +} diff --git a/src-tauri/src/router/route_registry.rs b/src-tauri/src/router/route_registry.rs new file mode 100644 index 000000000..482300c07 --- /dev/null +++ b/src-tauri/src/router/route_registry.rs @@ -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, + /// 凭证 UUID(如果适用) + pub credential_uuid: Option, + /// 凭证名称(如果适用) + pub credential_name: Option, + /// 支持的协议 + pub protocols: Vec, + /// 是否启用 + 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, + /// 路由名称到索引的映射 + name_index: HashMap, + /// UUID 到索引的映射 + uuid_index: HashMap, +} + +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()); + } +} diff --git a/src-tauri/src/server.rs b/src-tauri/src/server.rs index ab1b04322..4f0fa7d66 100644 --- a/src-tauri/src/server.rs +++ b/src-tauri/src/server.rs @@ -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, ) -> 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, ) -> 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::(&body) { + if let Ok(openai_resp) = + serde_json::from_str::(&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 { diff --git a/src-tauri/src/services/live_sync.rs b/src-tauri/src/services/live_sync.rs index 9ee89e871..34c3bf263 100644 --- a/src-tauri/src/services/live_sync.rs +++ b/src-tauri/src/services/live_sync.rs @@ -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 { let home = dirs::home_dir()?; match app_type { diff --git a/src-tauri/src/services/mcp_sync.rs b/src-tauri/src/services/mcp_sync.rs index 8136b09ea..ccc3b0715 100644 --- a/src-tauri/src/services/mcp_sync.rs +++ b/src-tauri/src/services/mcp_sync.rs @@ -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 { let home = dirs::home_dir()?; match app_type { diff --git a/src-tauri/src/services/prompt_service.rs b/src-tauri/src/services/prompt_service.rs index 0990ecad1..84acf8573 100644 --- a/src-tauri/src/services/prompt_service.rs +++ b/src-tauri/src/services/prompt_service.rs @@ -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, String> { diff --git a/src-tauri/src/services/prompt_sync.rs b/src-tauri/src/services/prompt_sync.rs index d71ebdac4..debf97402 100644 --- a/src-tauri/src/services/prompt_sync.rs +++ b/src-tauri/src/services/prompt_sync.rs @@ -18,6 +18,7 @@ pub fn get_prompt_file_path(app: &AppType) -> Option { } /// 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")?; diff --git a/src-tauri/src/services/provider_pool_service.rs b/src-tauri/src/services/provider_pool_service.rs index daa684218..6eceaaf86 100644 --- a/src-tauri/src/services/provider_pool_service.rs +++ b/src-tauri/src/services/provider_pool_service.rs @@ -85,7 +85,8 @@ impl ProviderPoolService { ) -> Result, 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 { - 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 { + 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 { 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 { 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 { 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 { + 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()), } } diff --git a/src-tauri/src/services/token_cache_service.rs b/src-tauri/src/services/token_cache_service.rs index b0e9d86dc..57a5fcb50 100644 --- a/src-tauri/src/services/token_cache_service.rs +++ b/src-tauri/src/services/token_cache_service.rs @@ -43,11 +43,7 @@ impl TokenCacheService { /// 1. 检查数据库缓存是否有效 /// 2. 如果缓存有效且未过期,直接返回 /// 3. 如果缓存无效或即将过期,执行刷新 - pub async fn get_valid_token( - &self, - db: &DbConnection, - uuid: &str, - ) -> Result { + pub async fn get_valid_token(&self, db: &DbConnection, uuid: &str) -> Result { // 首先检查缓存 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 { + 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, diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index d092963ed..11696d6cc 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -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", diff --git a/src/components/api-server/RoutesTab.tsx b/src/components/api-server/RoutesTab.tsx index 1c3d3ad0b..6cb104ec9 100644 --- a/src/components/api-server/RoutesTab.tsx +++ b/src/components/api-server/RoutesTab.tsx @@ -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"; } diff --git a/src/components/provider-pool/AddCredentialModal.tsx b/src/components/provider-pool/AddCredentialModal.tsx index 37406e9c3..89861c3af 100644 --- a/src/components/provider-pool/AddCredentialModal.tsx +++ b/src/components/provider-pool/AddCredentialModal.tsx @@ -14,6 +14,7 @@ const defaultCredsPath: Record = { 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 = { 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 () => { diff --git a/src/components/provider-pool/CredentialCard.tsx b/src/components/provider-pool/CredentialCard.tsx index e447a1c50..abc66f557 100644 --- a/src/components/provider-pool/CredentialCard.tsx +++ b/src/components/provider-pool/CredentialCard.tsx @@ -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 (
- {/* Header */} -
+
+ {/* Status Icon */} +
+ {credential.is_disabled ? ( + + ) : isHealthy ? ( + + ) : ( + + )} +
+ + {/* Main Info */}
-
+

- {credential.name || `凭证 #${credential.uuid.slice(0, 12)}`} + {credential.name || `凭证 #${credential.uuid.slice(0, 8)}`}

- + {getCredentialTypeLabel(credential.credential_type)}
-

- {credential.uuid.slice(0, 24)}... +

+ {credential.uuid}

-
-
+
+ +
+
使用次数
+
{credential.usage_count}
+
+
+
+ +
+
错误次数
+
{credential.error_count}
+
+
+
+ +
+
最后使用
+
+ {formatDate(credential.last_used)} +
+
+
+
+ + {/* Health Check Info */} + {credential.last_health_check_time && ( +
+
检查: {formatDate(credential.last_health_check_time)}
+ {credential.last_health_check_model && ( +
+ ({credential.last_health_check_model}) +
+ )} +
+ )} + + {/* Actions */} +
+
-
-
+ - {/* Stats */} -
-
-
-
- -
-
-
使用次数
-
- {credential.usage_count} -
-
-
-
-
- -
-
-
错误次数
-
- {credential.error_count} -
-
-
-
- - - 最后使用: {formatDate(credential.last_used)} - -
-
-
+ - {/* Health Check Info */} - {credential.last_health_check_time && ( -
- - 检查: {formatDate(credential.last_health_check_time)} - {credential.last_health_check_model && - ` (${credential.last_health_check_model})`} - -
- )} - - {/* OAuth Status */} - {isOAuth && credential.oauth_status && ( -
-
-
- -
- OAuth 状态 -
-
-
- - Access Token - -
- {credential.oauth_status.has_access_token ? ( - - ) : ( - - )} - - {credential.oauth_status.has_access_token ? "有效" : "缺失"} - -
-
-
- Token 状态 -
- {credential.oauth_status.is_token_valid ? ( - - ) : ( - - )} - - {credential.oauth_status.is_token_valid ? "有效" : "需刷新"} - -
-
- {credential.oauth_status.expiry_info && ( -
- - 过期时间: {credential.oauth_status.expiry_info} - -
- )} -
-
- )} - - {/* Token Cache Status */} - {isOAuth && credential.token_cache_status && ( -
-
-
- -
- Token 缓存 -
-
-
- 缓存状态 -
- {credential.token_cache_status.has_cached_token ? ( - - ) : ( - - )} - - {credential.token_cache_status.has_cached_token - ? "已缓存" - : "未缓存"} - -
-
-
- 有效性 -
- {credential.token_cache_status.is_valid ? ( - credential.token_cache_status.is_expiring_soon ? ( - - ) : ( - - ) - ) : ( - - )} - - {credential.token_cache_status.is_valid - ? credential.token_cache_status.is_expiring_soon - ? "即将过期" - : "有效" - : "已过期"} - -
-
- {(credential.token_cache_status.last_refresh || - credential.token_cache_status.expiry_time) && ( -
- {credential.token_cache_status.last_refresh && ( -
- 最后刷新:{" "} - {formatDate(credential.token_cache_status.last_refresh)} -
- )} - {credential.token_cache_status.expiry_time && ( -
- 过期时间:{" "} - {formatDate(credential.token_cache_status.expiry_time)} -
- )} -
- )} - {credential.token_cache_status.refresh_error_count > 0 && ( -
-
- - - 刷新失败 {credential.token_cache_status.refresh_error_count}{" "} - 次 - -
- {credential.token_cache_status.last_refresh_error && ( -
- {credential.token_cache_status.last_refresh_error.slice( - 0, - 60, - )} - {credential.token_cache_status.last_refresh_error.length > - 60 && "..."} -
- )} -
- )} -
-
- )} - - {/* Error Message */} - {credential.last_error_message && ( -
- {credential.last_error_message.slice(0, 100)} - {credential.last_error_message.length > 100 && "..."} -
- )} - - {/* Actions */} -
- - - - -
{isOAuth && onRefreshToken && ( )}
+ + {/* Mobile Stats - shown on small screens */} +
+
+
+ + + 使用: {credential.usage_count} + + + + 错误: {credential.error_count} + +
+ + + {formatDate(credential.last_used)} + +
+
+ + {/* Error Message */} + {credential.last_error_message && ( +
+ {credential.last_error_message.slice(0, 150)} + {credential.last_error_message.length > 150 && "..."} +
+ )}
); } diff --git a/src/components/provider-pool/EditCredentialModal.tsx b/src/components/provider-pool/EditCredentialModal.tsx index b562e3e86..7047f7aa4 100644 --- a/src/components/provider-pool/EditCredentialModal.tsx +++ b/src/components/provider-pool/EditCredentialModal.tsx @@ -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; } +// 各 Provider 支持的模型列表 (参考 AIClient-2-API/src/provider-models.js) +const providerModels: Record = { + 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([]); const [loading, setLoading] = useState(false); const [error, setError] = useState(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 (
-
+
{/* Header */} -
+

编辑凭证 @@ -131,204 +177,189 @@ export function EditCredentialModal({ {/* Content - Scrollable */}
-
- {/* 凭证信息(只读) */} -
-
- - -
-
-
- 类型: - {credential.credential_type} -
-
- UUID: - - {credential.uuid.slice(0, 24)}... - -
-
- - {isOAuth ? "文件路径:" : "API Key:"} - - - {showCredentialDetails - ? credential.display_credential - : getMaskedCredentialInfo()} - -
-
-

- 🔒 敏感信息(API Key、文件路径)无法修改,如需更改请删除后重新添加 -

-
- - {/* 可编辑字段 */} -
- - setName(e.target.value)} - placeholder="给这个凭证起个名字..." - className="w-full rounded-lg border bg-background px-3 py-2 text-sm" - /> -
- - {/* 健康检查设置 */} -
- - {checkHealth && ( -
-