From bd9748fe18b63a078f9ecd1f2b4796fbb4a0e6c0 Mon Sep 17 00:00:00 2001 From: Chiron <598621670@qq.com> Date: Mon, 12 Jan 2026 11:40:29 +0800 Subject: [PATCH 1/2] =?UTF-8?q?feat:=20=E5=8A=A8=E6=80=81=E6=A3=80?= =?UTF-8?q?=E6=B5=8B=20IP=20=E5=9C=B0=E5=9D=80=E5=8F=98=E5=8C=96=E5=B9=B6?= =?UTF-8?q?=E8=87=AA=E5=8A=A8=E6=9B=B4=E6=96=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 主要改进: 1. 服务器启动时检查配置的 IP 是否有效,无效时自动切换到当前局域网 IP 2. 前端定期刷新网络信息,实时检测 IP 变化 3. 路由端点和 API 测试地址动态更新 4. 优先选择 192.168.x.x 或 10.x.x.x 的局域网 IP 修复 Antigravity API: 1. 修正 API 端点为 cloudcode-pa.googleapis.com 2. 透传上游错误状态码 新增功能: 1. 添加 model_service 和 model_cmd 模块 --- src-tauri/src/app/commands/api_test.rs | 4 +- src-tauri/src/app/commands/server.rs | 5 +- src-tauri/src/app/runner.rs | 13 +- src-tauri/src/app/setup.rs | 6 +- src-tauri/src/commands/mod.rs | 1 + src-tauri/src/commands/model_cmd.rs | 146 ++++++ src-tauri/src/commands/route_cmd.rs | 43 +- .../src/converter/openai_to_antigravity.rs | 18 +- src-tauri/src/database/dao/provider_pool.rs | 52 +- src-tauri/src/database/schema.rs | 6 + src-tauri/src/models/provider_pool_model.rs | 12 + src-tauri/src/providers/README.md | 8 + src-tauri/src/providers/antigravity.rs | 175 ++++++- src-tauri/src/providers/mod.rs | 2 + .../src/server/handlers/provider_calls.rs | 63 +-- src-tauri/src/server/mod.rs | 121 ++++- src-tauri/src/server_utils.rs | 86 +++- .../src/services/api_key_provider_service.rs | 2 + src-tauri/src/services/mod.rs | 1 + src-tauri/src/services/model_service.rs | 467 ++++++++++++++++++ src/components/api-server/ApiServerPage.tsx | 146 +++++- src/lib/api/providerPool.ts | 34 ++ 22 files changed, 1290 insertions(+), 121 deletions(-) create mode 100644 src-tauri/src/commands/model_cmd.rs create mode 100644 src-tauri/src/services/model_service.rs diff --git a/src-tauri/src/app/commands/api_test.rs b/src-tauri/src/app/commands/api_test.rs index 4e009c805..be2f39352 100644 --- a/src-tauri/src/app/commands/api_test.rs +++ b/src-tauri/src/app/commands/api_test.rs @@ -310,7 +310,9 @@ pub async fn test_api( auth: bool, ) -> Result { let s = state.read().await; - let base_url = format!("http://{}:{}", s.config.server.host, s.config.server.port); + // 使用 status() 获取实际监听的地址(可能与配置不同) + let status = s.status(); + let base_url = format!("http://{}:{}", status.host, status.port); let api_key = s .running_api_key .as_ref() diff --git a/src-tauri/src/app/commands/server.rs b/src-tauri/src/app/commands/server.rs index 2bae835fc..ea7779c0d 100644 --- a/src-tauri/src/app/commands/server.rs +++ b/src-tauri/src/app/commands/server.rs @@ -28,11 +28,14 @@ pub async fn start_server( ) .await .map_err(|e| e.to_string())?; + + // 使用 status() 获取实际使用的地址(可能已经自动切换到有效的 IP) + let status = s.status(); logs.write().await.add( "info", &format!( "Server started on {}:{}", - s.config.server.host, s.config.server.port + status.host, status.port ), ); Ok("Server started".to_string()) diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index e158010eb..81ea2e0f6 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -522,8 +522,10 @@ pub fn run() { .await { Ok(_) => { - let host = s.config.server.host.clone(); - let port = s.config.server.port; + // 使用 status() 获取实际使用的地址(可能已经自动切换到有效的 IP) + let status = s.status(); + let host = status.host; + let port = status.port; logs.write() .await .add("info", &format!("[启动] 服务器已启动: {host}:{port}")); @@ -1158,6 +1160,13 @@ pub fn run() { commands::model_registry_cmd::get_models_by_tier, commands::model_registry_cmd::get_provider_alias_config, commands::model_registry_cmd::get_all_alias_configs, + // Model Management commands (动态模型列表) + commands::model_cmd::get_credential_models, + commands::model_cmd::refresh_credential_models, + commands::model_cmd::get_all_models_by_provider, + commands::model_cmd::get_all_available_models, + commands::model_cmd::refresh_all_credential_models, + commands::model_cmd::get_default_models_for_provider, // Terminal commands commands::terminal_cmd::terminal_create_session, commands::terminal_cmd::terminal_write, diff --git a/src-tauri/src/app/setup.rs b/src-tauri/src/app/setup.rs index 0053d9b4a..53db57a0c 100644 --- a/src-tauri/src/app/setup.rs +++ b/src-tauri/src/app/setup.rs @@ -187,8 +187,10 @@ async fn start_server_async( .await { Ok(_) => { - let host = s.config.server.host.clone(); - let port = s.config.server.port; + // 获取服务器实际使用的地址(可能已经自动切换到有效的 IP) + let status = s.status(); + let host = status.host; + let port = status.port; logs.write() .await .add("info", &format!("[启动] 服务器已启动: {host}:{port}")); diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index c90da7ead..519b9f71f 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -10,6 +10,7 @@ pub mod injection_cmd; pub mod kiro_local; pub mod machine_id_cmd; pub mod mcp_cmd; +pub mod model_cmd; pub mod model_registry_cmd; pub mod models_cmd; pub mod music_cmd; diff --git a/src-tauri/src/commands/model_cmd.rs b/src-tauri/src/commands/model_cmd.rs new file mode 100644 index 000000000..0f822264a --- /dev/null +++ b/src-tauri/src/commands/model_cmd.rs @@ -0,0 +1,146 @@ +//! 模型管理相关命令 + +use crate::database::DbConnection; +use crate::database::dao::provider_pool::ProviderPoolDao; +use crate::services::model_service::ModelService; +use std::collections::HashMap; +use tauri::State; + +/// 获取凭证支持的模型列表(从数据库缓存) +#[tauri::command] +pub fn get_credential_models( + db: State<'_, DbConnection>, + credential_uuid: String, +) -> Result, String> { + tracing::info!("[GET_CREDENTIAL_MODELS] 获取凭证模型列表: {}", credential_uuid); + + let model_service = ModelService::new(); + model_service.get_credential_models(&db, &credential_uuid) +} + +/// 刷新凭证的模型列表(从 Provider API 重新获取) +#[tauri::command] +pub async fn refresh_credential_models( + db: State<'_, DbConnection>, + credential_uuid: String, +) -> Result, String> { + tracing::info!("[REFRESH_CREDENTIAL_MODELS] ========== 开始刷新凭证模型列表 =========="); + tracing::info!("[REFRESH_CREDENTIAL_MODELS] credential_uuid: {}", credential_uuid); + + let model_service = ModelService::new(); + + // 从数据库获取凭证信息 + let credential = { + let conn = db.lock().map_err(|e| e.to_string())?; + ProviderPoolDao::get_by_uuid(&conn, &credential_uuid) + .map_err(|e| e.to_string())? + .ok_or_else(|| format!("凭证不存在: {}", credential_uuid))? + }; + + tracing::info!( + "[REFRESH_CREDENTIAL_MODELS] 凭证信息: provider_type={}, name={:?}", + credential.provider_type, + credential.name + ); + + // 从 Provider API 获取模型列表 + tracing::info!("[REFRESH_CREDENTIAL_MODELS] 开始从 Provider API 获取模型列表..."); + let models = model_service.fetch_models_for_credential(&credential).await?; + + tracing::info!( + "[REFRESH_CREDENTIAL_MODELS] 成功获取 {} 个模型: {:?}", + models.len(), + models + ); + + // 更新到数据库 + tracing::info!("[REFRESH_CREDENTIAL_MODELS] 更新模型列表到数据库..."); + model_service.update_credential_models(&db, &credential_uuid, models.clone())?; + + tracing::info!("[REFRESH_CREDENTIAL_MODELS] ========== 刷新完成 =========="); + + Ok(models) +} + +/// 获取所有凭证的模型列表(按 Provider 类型分组) +#[tauri::command] +pub fn get_all_models_by_provider( + db: State<'_, DbConnection>, +) -> Result>, String> { + tracing::info!("[GET_ALL_MODELS_BY_PROVIDER] 获取所有 Provider 的模型列表"); + + let model_service = ModelService::new(); + model_service.get_all_models_by_provider(&db) +} + +/// 获取所有可用的模型列表(合并所有健康凭证的模型) +#[tauri::command] +pub fn get_all_available_models( + db: State<'_, DbConnection>, +) -> Result, String> { + tracing::info!("[GET_ALL_AVAILABLE_MODELS] 获取所有可用模型"); + + let model_service = ModelService::new(); + model_service.get_all_available_models(&db) +} + +/// 批量刷新所有凭证的模型列表 +#[tauri::command] +pub async fn refresh_all_credential_models( + db: State<'_, DbConnection>, +) -> Result, String>>, String> { + tracing::info!("[REFRESH_ALL_CREDENTIAL_MODELS] 批量刷新所有凭证的模型列表"); + + let model_service = ModelService::new(); + + // 获取所有凭证 + let credentials = { + let conn = db.lock().map_err(|e| e.to_string())?; + ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())? + }; + + let mut results = HashMap::new(); + + for credential in credentials { + if credential.is_disabled { + tracing::debug!("[REFRESH_ALL] 跳过已禁用的凭证: {}", credential.uuid); + continue; + } + + tracing::info!("[REFRESH_ALL] 刷新凭证: {} ({})", credential.uuid, credential.provider_type); + + // 尝试获取模型列表 + let result = match model_service.fetch_models_for_credential(&credential).await { + Ok(models) => { + // 更新到数据库 + if let Err(e) = model_service.update_credential_models(&db, &credential.uuid, models.clone()) { + tracing::error!("[REFRESH_ALL] 更新数据库失败: {}", e); + Err(format!("更新数据库失败: {}", e)) + } else { + tracing::info!("[REFRESH_ALL] 成功刷新 {} 个模型", models.len()); + Ok(models) + } + } + Err(e) => { + tracing::warn!("[REFRESH_ALL] 获取模型列表失败: {}", e); + Err(e) + } + }; + + results.insert(credential.uuid.clone(), result); + } + + Ok(results) +} + +/// 获取 Provider 的默认模型列表 +#[tauri::command] +pub fn get_default_models_for_provider( + provider_type: String, +) -> Result, String> { + let pt: crate::models::provider_pool_model::PoolProviderType = + provider_type.parse().map_err(|e: String| e)?; + + let model_service = ModelService::new(); + Ok(model_service.get_default_models_for_provider(&pt)) +} diff --git a/src-tauri/src/commands/route_cmd.rs b/src-tauri/src/commands/route_cmd.rs index 9df74a994..9842feeca 100644 --- a/src-tauri/src/commands/route_cmd.rs +++ b/src-tauri/src/commands/route_cmd.rs @@ -5,6 +5,45 @@ use crate::config; use crate::database::DbConnection; use crate::models::route_model::{RouteInfo, RouteListResponse}; +/// 获取有效的服务器地址 +/// 如果配置的 IP 不在当前网卡列表中,自动替换为当前的局域网 IP +fn get_valid_base_url(config: &config::Config) -> String { + let configured_host = &config.server.host; + let port = config.server.port; + + // 特殊地址不需要检查 + if configured_host == "127.0.0.1" || configured_host == "localhost" { + return format!("http://{}:{}", configured_host, port); + } + + // 0.0.0.0 或其他 IP 需要检查 + if let Ok(network_info) = crate::commands::network_cmd::get_network_info() { + let host = if configured_host == "0.0.0.0" { + // 0.0.0.0 替换为局域网 IP + network_info.all_ips.iter() + .find(|ip| ip.starts_with("192.168.") || ip.starts_with("10.")) + .or_else(|| network_info.lan_ip.as_ref()) + .or_else(|| network_info.all_ips.first()) + .cloned() + .unwrap_or_else(|| "localhost".to_string()) + } else if network_info.all_ips.contains(configured_host) { + // IP 在当前网卡列表中,使用配置的 IP + configured_host.clone() + } else { + // IP 不在当前网卡列表中,替换为局域网 IP + network_info.all_ips.iter() + .find(|ip| ip.starts_with("192.168.") || ip.starts_with("10.")) + .or_else(|| network_info.lan_ip.as_ref()) + .or_else(|| network_info.all_ips.first()) + .cloned() + .unwrap_or_else(|| "localhost".to_string()) + }; + format!("http://{}:{}", host, port) + } else { + format!("http://{}:{}", configured_host, port) + } +} + /// 获取所有可用的路由端点 #[tauri::command] pub async fn get_available_routes( @@ -13,7 +52,7 @@ pub async fn get_available_routes( ) -> Result { // 获取配置中的服务器地址和默认 Provider let config = config::load_config().unwrap_or_default(); - let base_url = format!("http://{}:{}", config.server.host, config.server.port); + let base_url = get_valid_base_url(&config); let default_provider = config.default_provider.clone(); let routes = pool_service @@ -58,7 +97,7 @@ pub async fn get_route_curl_examples( pool_service: tauri::State<'_, ProviderPoolServiceState>, ) -> Result, String> { let config = config::load_config().unwrap_or_default(); - let base_url = format!("http://{}:{}", config.server.host, config.server.port); + let base_url = get_valid_base_url(&config); let default_provider = config.default_provider.clone(); let routes = pool_service diff --git a/src-tauri/src/converter/openai_to_antigravity.rs b/src-tauri/src/converter/openai_to_antigravity.rs index c62988420..6b4d517c1 100644 --- a/src-tauri/src/converter/openai_to_antigravity.rs +++ b/src-tauri/src/converter/openai_to_antigravity.rs @@ -273,8 +273,17 @@ pub fn convert_openai_to_antigravity_with_context( request: &ChatCompletionRequest, project_id: &str, ) -> serde_json::Value { + eprintln!("========== [CONVERT] OpenAI -> Antigravity 转换开始 =========="); + eprintln!("[CONVERT] 原始模型: {}", request.model); + eprintln!("[CONVERT] 项目ID: {}", project_id); + eprintln!("[CONVERT] 消息数量: {}", request.messages.len()); + eprintln!("[CONVERT] 流式: {}", request.stream); + let actual_model = model_mapping(&request.model); + eprintln!("[CONVERT] 映射后模型: {}", actual_model); + let supports_thinking = model_supports_thinking(actual_model); + eprintln!("[CONVERT] 支持思维链: {}", supports_thinking); let mut contents: Vec = Vec::new(); let mut system_instruction: Option = None; @@ -656,13 +665,18 @@ pub fn convert_openai_to_antigravity_with_context( }; // 构建完整的 Antigravity 请求体 - serde_json::json!({ + let result = serde_json::json!({ "project": project_id, "requestId": generate_request_id(), "request": inner, "model": actual_model, "userAgent": "antigravity" - }) + }); + + eprintln!("[CONVERT] 转换后的请求体: {}", serde_json::to_string_pretty(&result).unwrap_or_default()); + eprintln!("========== [CONVERT] OpenAI -> Antigravity 转换完成 =========="); + + result } // ============================================================================ diff --git a/src-tauri/src/database/dao/provider_pool.rs b/src-tauri/src/database/dao/provider_pool.rs index 58972348b..84a892f55 100644 --- a/src-tauri/src/database/dao/provider_pool.rs +++ b/src-tauri/src/database/dao/provider_pool.rs @@ -16,7 +16,7 @@ impl ProviderPoolDao { pub fn get_all(conn: &Connection) -> Result, rusqlite::Error> { let mut stmt = conn.prepare( "SELECT uuid, provider_type, credential_data, name, is_healthy, is_disabled, - check_health, check_model_name, not_supported_models, usage_count, error_count, + check_health, check_model_name, not_supported_models, supported_models, usage_count, error_count, last_used, last_error_time, last_error_message, last_health_check_time, last_health_check_model, created_at, updated_at, source, proxy_url FROM provider_pool_credentials @@ -39,7 +39,7 @@ impl ProviderPoolDao { ) -> Result, rusqlite::Error> { let mut stmt = conn.prepare( "SELECT uuid, provider_type, credential_data, name, is_healthy, is_disabled, - check_health, check_model_name, not_supported_models, usage_count, error_count, + check_health, check_model_name, not_supported_models, supported_models, usage_count, error_count, last_used, last_error_time, last_error_message, last_health_check_time, last_health_check_model, created_at, updated_at, source, proxy_url FROM provider_pool_credentials @@ -65,7 +65,7 @@ impl ProviderPoolDao { ) -> Result, rusqlite::Error> { let mut stmt = conn.prepare( "SELECT uuid, provider_type, credential_data, name, is_healthy, is_disabled, - check_health, check_model_name, not_supported_models, usage_count, error_count, + check_health, check_model_name, not_supported_models, supported_models, usage_count, error_count, last_used, last_error_time, last_error_message, last_health_check_time, last_health_check_model, created_at, updated_at, source, proxy_url FROM provider_pool_credentials @@ -87,7 +87,7 @@ impl ProviderPoolDao { ) -> Result, rusqlite::Error> { let mut stmt = conn.prepare( "SELECT uuid, provider_type, credential_data, name, is_healthy, is_disabled, - check_health, check_model_name, not_supported_models, usage_count, error_count, + check_health, check_model_name, not_supported_models, supported_models, usage_count, error_count, last_used, last_error_time, last_error_message, last_health_check_time, last_health_check_model, created_at, updated_at, source, proxy_url FROM provider_pool_credentials @@ -120,6 +120,8 @@ impl ProviderPoolDao { serde_json::to_string(&cred.credential).unwrap_or_else(|_| "{}".to_string()); let not_supported_models_json = serde_json::to_string(&cred.not_supported_models).unwrap_or_else(|_| "[]".to_string()); + let supported_models_json = + serde_json::to_string(&cred.supported_models).unwrap_or_else(|_| "[]".to_string()); let source_str = match cred.source { CredentialSource::Manual => "manual", CredentialSource::Imported => "imported", @@ -129,10 +131,10 @@ impl ProviderPoolDao { conn.execute( "INSERT INTO provider_pool_credentials (uuid, provider_type, credential_data, name, is_healthy, is_disabled, - check_health, check_model_name, not_supported_models, usage_count, error_count, + check_health, check_model_name, not_supported_models, supported_models, usage_count, error_count, last_used, last_error_time, last_error_message, last_health_check_time, last_health_check_model, created_at, updated_at, source, proxy_url) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20)", + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20, ?21)", params![ cred.uuid, cred.provider_type.to_string(), @@ -143,6 +145,7 @@ impl ProviderPoolDao { cred.check_health, cred.check_model_name, not_supported_models_json, + supported_models_json, cred.usage_count, cred.error_count, cred.last_used.map(|t| t.timestamp()), @@ -165,14 +168,16 @@ impl ProviderPoolDao { serde_json::to_string(&cred.credential).unwrap_or_else(|_| "{}".to_string()); let not_supported_models_json = serde_json::to_string(&cred.not_supported_models).unwrap_or_else(|_| "[]".to_string()); + let supported_models_json = + serde_json::to_string(&cred.supported_models).unwrap_or_else(|_| "[]".to_string()); conn.execute( "UPDATE provider_pool_credentials SET provider_type = ?2, credential_data = ?3, name = ?4, is_healthy = ?5, is_disabled = ?6, check_health = ?7, check_model_name = ?8, - not_supported_models = ?9, usage_count = ?10, error_count = ?11, - last_used = ?12, last_error_time = ?13, last_error_message = ?14, - last_health_check_time = ?15, last_health_check_model = ?16, updated_at = ?17, proxy_url = ?18 + not_supported_models = ?9, supported_models = ?10, usage_count = ?11, error_count = ?12, + last_used = ?13, last_error_time = ?14, last_error_message = ?15, + last_health_check_time = ?16, last_health_check_model = ?17, updated_at = ?18, proxy_url = ?19 WHERE uuid = ?1", params![ cred.uuid, @@ -184,6 +189,7 @@ impl ProviderPoolDao { cred.check_health, cred.check_model_name, not_supported_models_json, + supported_models_json, cred.usage_count, cred.error_count, cred.last_used.map(|t| t.timestamp()), @@ -297,17 +303,18 @@ impl ProviderPoolDao { let check_health: bool = row.get(6)?; let check_model_name: Option = row.get(7)?; let not_supported_models_json: Option = row.get(8)?; - let usage_count: u64 = row.get::<_, i64>(9)? as u64; - let error_count: u32 = row.get::<_, i32>(10)? as u32; - let last_used_ts: Option = row.get(11)?; - let last_error_time_ts: Option = row.get(12)?; - let last_error_message: Option = row.get(13)?; - let last_health_check_time_ts: Option = row.get(14)?; - let last_health_check_model: Option = row.get(15)?; - let created_at_ts: i64 = row.get(16)?; - let updated_at_ts: i64 = row.get(17)?; - let source_str: Option = row.get(18).ok(); - let proxy_url: Option = row.get(19).ok(); + let supported_models_json: Option = row.get(9)?; + let usage_count: u64 = row.get::<_, i64>(10)? as u64; + let error_count: u32 = row.get::<_, i32>(11)? as u32; + let last_used_ts: Option = row.get(12)?; + let last_error_time_ts: Option = row.get(13)?; + let last_error_message: Option = row.get(14)?; + let last_health_check_time_ts: Option = row.get(15)?; + let last_health_check_model: Option = row.get(16)?; + let created_at_ts: i64 = row.get(17)?; + let updated_at_ts: i64 = row.get(18)?; + let source_str: Option = row.get(19).ok(); + let proxy_url: Option = row.get(20).ok(); let provider_type: PoolProviderType = provider_type_str.parse().unwrap_or(PoolProviderType::Kiro); @@ -320,6 +327,10 @@ impl ProviderPoolDao { .and_then(|s| serde_json::from_str(&s).ok()) .unwrap_or_default(); + let supported_models: Vec = supported_models_json + .and_then(|s| serde_json::from_str(&s).ok()) + .unwrap_or_default(); + let source = match source_str.as_deref() { Some("imported") => CredentialSource::Imported, Some("private") => CredentialSource::Private, @@ -336,6 +347,7 @@ impl ProviderPoolDao { check_health, check_model_name, not_supported_models, + supported_models, usage_count, error_count, last_used: last_used_ts.and_then(|ts| Utc.timestamp_opt(ts, 0).single()), diff --git a/src-tauri/src/database/schema.rs b/src-tauri/src/database/schema.rs index 820482f84..4a6a0b216 100644 --- a/src-tauri/src/database/schema.rs +++ b/src-tauri/src/database/schema.rs @@ -219,6 +219,12 @@ pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> { [], ); + // Migration: 添加支持的模型列表字段 + let _ = conn.execute( + "ALTER TABLE provider_pool_credentials ADD COLUMN supported_models TEXT", + [], + ); + // Migration: 添加代理URL字段 - 使用重建表结构的方式 migrate_add_proxy_url_column(conn)?; diff --git a/src-tauri/src/models/provider_pool_model.rs b/src-tauri/src/models/provider_pool_model.rs index 2615f4657..a0ef773bc 100644 --- a/src-tauri/src/models/provider_pool_model.rs +++ b/src-tauri/src/models/provider_pool_model.rs @@ -213,6 +213,9 @@ pub struct ProviderCredential { /// 不支持的模型列表(黑名单) #[serde(default)] pub not_supported_models: Vec, + /// 支持的模型列表(从 /v1/models 接口获取) + #[serde(default)] + pub supported_models: Vec, /// 使用次数 #[serde(default)] pub usage_count: u64, @@ -261,6 +264,7 @@ impl ProviderCredential { check_health: true, check_model_name: None, not_supported_models: Vec::new(), + supported_models: Vec::new(), usage_count: 0, error_count: 0, last_used: None, @@ -534,6 +538,7 @@ pub struct CredentialDisplay { pub check_health: bool, pub check_model_name: Option, pub not_supported_models: Vec, + pub supported_models: Vec, pub usage_count: u64, pub error_count: u32, pub last_used: Option, @@ -639,6 +644,7 @@ impl From<&ProviderCredential> for CredentialDisplay { check_health: cred.check_health, check_model_name: cred.check_model_name.clone(), not_supported_models: cred.not_supported_models.clone(), + supported_models: cred.supported_models.clone(), usage_count: cred.usage_count, error_count: cred.error_count, last_used: cred.last_used.map(|t| t.to_rfc3339()), @@ -769,6 +775,7 @@ mod tests { check_health: true, check_model_name: None, not_supported_models: vec!["claude-opus".to_string()], + supported_models: vec![], usage_count: 0, error_count: 0, last_used: None, @@ -803,6 +810,7 @@ mod tests { check_health: true, check_model_name: None, not_supported_models: vec![], + supported_models: vec![], usage_count: 0, error_count: 0, last_used: None, @@ -839,6 +847,7 @@ mod tests { check_health: true, check_model_name: None, not_supported_models: vec![], + supported_models: vec![], usage_count: 0, error_count: 0, last_used: None, @@ -879,6 +888,7 @@ mod tests { check_health: true, check_model_name: None, not_supported_models: vec![], + supported_models: vec![], usage_count: 0, error_count: 0, last_used: None, @@ -916,6 +926,7 @@ mod tests { check_health: true, check_model_name: None, not_supported_models: vec!["gemini-3-pro".to_string()], + supported_models: vec![], usage_count: 0, error_count: 0, last_used: None, @@ -954,6 +965,7 @@ mod tests { check_health: true, check_model_name: None, not_supported_models: vec![], + supported_models: vec![], usage_count: 0, error_count: 0, last_used: None, diff --git a/src-tauri/src/providers/README.md b/src-tauri/src/providers/README.md index c0a78a2a2..496ffae21 100644 --- a/src-tauri/src/providers/README.md +++ b/src-tauri/src/providers/README.md @@ -1,3 +1,11 @@ + # providers diff --git a/src-tauri/src/providers/antigravity.rs b/src-tauri/src/providers/antigravity.rs index 26a8a1a1c..850b8e254 100644 --- a/src-tauri/src/providers/antigravity.rs +++ b/src-tauri/src/providers/antigravity.rs @@ -14,7 +14,80 @@ use std::sync::Arc; use tokio::sync::oneshot; use uuid::Uuid; +// ============================================================================ +// Antigravity API 错误类型 +// ============================================================================ + +/// Antigravity API 错误 +/// +/// 携带 HTTP 状态码,便于调用方透传给客户端 +#[derive(Debug, Clone)] +pub struct AntigravityApiError { + /// HTTP 状态码 + pub status_code: u16, + /// 错误消息 + pub message: String, + /// 原始响应体(如果有) + pub body: Option, +} + +impl AntigravityApiError { + /// 创建新的 API 错误 + pub fn new(status_code: u16, message: impl Into) -> Self { + Self { + status_code, + message: message.into(), + body: None, + } + } + + /// 创建带响应体的 API 错误 + pub fn with_body(status_code: u16, message: impl Into, body: impl Into) -> Self { + Self { + status_code, + message: message.into(), + body: Some(body.into()), + } + } + + /// 是否是可重试的错误(429 或 5xx) + pub fn is_retryable(&self) -> bool { + self.status_code == 429 || (self.status_code >= 500 && self.status_code < 600) + } + + /// 是否是权限错误(401 或 403) + pub fn is_auth_error(&self) -> bool { + self.status_code == 401 || self.status_code == 403 + } + + /// 是否是配额耗尽错误(429) + pub fn is_rate_limit(&self) -> bool { + self.status_code == 429 + } + + /// 获取用户友好的错误消息 + pub fn user_message(&self) -> String { + match self.status_code { + 401 => format!("认证失败,请重新登录: {}", self.message), + 403 => format!("权限不足: {}", self.message), + 429 => format!("请求过于频繁,请稍后重试: {}", self.message), + 500..=599 => format!("服务器错误 ({}): {}", self.status_code, self.message), + _ => format!("API 错误 ({}): {}", self.status_code, self.message), + } + } +} + +impl std::fmt::Display for AntigravityApiError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "HTTP {} - {}", self.status_code, self.message) + } +} + +impl std::error::Error for AntigravityApiError {} + // Constants +// 正确的 Cloud Code API 端点(参考 Antigravity-Manager) +const ANTIGRAVITY_BASE_URL_PROD: &str = "https://cloudcode-pa.googleapis.com"; 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"; @@ -287,9 +360,11 @@ impl Default for AntigravityProvider { .timeout(std::time::Duration::from_secs(120)) .build() .unwrap_or_else(|_| Client::new()), + // 只使用生产环境和 daily 环境(参考 Antigravity-Manager) + // 沙盒环境(autopush)需要特殊许可证,不适合普通用户 base_urls: vec![ + ANTIGRAVITY_BASE_URL_PROD.to_string(), ANTIGRAVITY_BASE_URL_DAILY.to_string(), - ANTIGRAVITY_BASE_URL_AUTOPUSH.to_string(), ], available_models: ANTIGRAVITY_MODELS_FALLBACK .iter() @@ -744,21 +819,30 @@ impl AntigravityProvider { Ok(new_token.to_string()) } - /// 调用 Antigravity API + /// 调用 Antigravity API(内部方法) + /// + /// 返回 `AntigravityApiError` 以便调用方获取 HTTP 状态码 async fn call_api_internal( &self, base_url: &str, method: &str, body: &serde_json::Value, - ) -> Result> { + ) -> Result { let token = self .credentials .access_token .as_ref() - .ok_or("No access token")?; + .ok_or_else(|| AntigravityApiError::new(401, "No access token"))?; let url = format!("{}/{ANTIGRAVITY_API_VERSION}:{method}", base_url); + // 打印详细的请求信息 + eprintln!("========== [ANTIGRAVITY_API] 请求详情 =========="); + eprintln!("[ANTIGRAVITY_API] URL: {}", url); + eprintln!("[ANTIGRAVITY_API] Method: {}", method); + eprintln!("[ANTIGRAVITY_API] Token (前20字符): {}...", &token[..token.len().min(20)]); + eprintln!("[ANTIGRAVITY_API] 请求体: {}", serde_json::to_string_pretty(body).unwrap_or_default()); + let resp = self .client .post(&url) @@ -767,37 +851,78 @@ impl AntigravityProvider { .header("User-Agent", "antigravity/1.11.5 windows/amd64") .json(body) .send() - .await?; + .await + .map_err(|e| { + eprintln!("[ANTIGRAVITY_API] 网络错误: {}", e); + AntigravityApiError::new(503, format!("Network error: {}", e)) + })?; - 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 status = resp.status(); + let status_code = status.as_u16(); + eprintln!("[ANTIGRAVITY_API] 响应状态码: {}", status); + + if !status.is_success() { + let body_text = resp.text().await.unwrap_or_default(); + eprintln!("[ANTIGRAVITY_API] 错误响应体: {}", body_text); + eprintln!("========== [ANTIGRAVITY_API] 请求失败 =========="); + return Err(AntigravityApiError::with_body( + status_code, + format!("API call failed: {}", status), + body_text, + )); } - let data: serde_json::Value = resp.json().await?; + let response_text = resp.text().await.map_err(|e| { + AntigravityApiError::new(500, format!("Failed to read response: {}", e)) + })?; + eprintln!("[ANTIGRAVITY_API] 响应体: {}", response_text); + + let data: serde_json::Value = serde_json::from_str(&response_text) + .map_err(|e| AntigravityApiError::new(500, format!("Failed to parse response: {}", e)))?; + + eprintln!("========== [ANTIGRAVITY_API] 请求成功 =========="); Ok(data) } /// 调用 API,支持多环境降级 + /// + /// 返回 `AntigravityApiError` 以便调用方获取 HTTP 状态码并透传给客户端。 + /// 只在特定错误(429 配额耗尽、5xx 服务器错误)时才降级到备用端点, + /// 403 权限错误等不应该降级,因为换端点也没用。 pub async fn call_api( &self, method: &str, body: &serde_json::Value, - ) -> Result> { - let mut last_error: Option> = None; + ) -> Result { + let mut last_error: Option = None; - for base_url in &self.base_urls { + for (idx, base_url) in self.base_urls.iter().enumerate() { 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); + // 使用 AntigravityApiError 的方法判断是否可重试 + let should_fallback = e.is_retryable(); + + if should_fallback && idx + 1 < self.base_urls.len() { + tracing::warn!( + "[Antigravity] {} 返回可重试错误 (HTTP {}), 尝试下一个端点", + base_url, e.status_code + ); + last_error = Some(e); + continue; + } + + // 403、401 等权限错误直接返回,不降级 + tracing::warn!( + "[Antigravity] {} 失败 (HTTP {}): {}", + base_url, e.status_code, e.message + ); + return Err(e); } } } - Err(last_error.unwrap_or_else(|| "All Antigravity base URLs failed".into())) + Err(last_error.unwrap_or_else(|| AntigravityApiError::new(503, "All Antigravity base URLs failed"))) } /// 发现项目 ID @@ -898,20 +1023,34 @@ impl AntigravityProvider { } /// 生成内容(非流式) + /// + /// 返回 `AntigravityApiError` 以便调用方获取 HTTP 状态码并透传给客户端。 pub async fn generate_content( &self, model: &str, request_body: &serde_json::Value, - ) -> Result> { + ) -> Result { + eprintln!("========== [ANTIGRAVITY_GENERATE] 开始生成内容 =========="); + eprintln!("[ANTIGRAVITY_GENERATE] 模型: {}", model); + eprintln!("[ANTIGRAVITY_GENERATE] 请求体: {}", serde_json::to_string_pretty(request_body).unwrap_or_default()); + let project_id = self.project_id.clone().unwrap_or_else(generate_project_id); + eprintln!("[ANTIGRAVITY_GENERATE] 项目ID: {}", project_id); + let actual_model = alias_to_model_name(model); + eprintln!("[ANTIGRAVITY_GENERATE] 实际模型名: {}", actual_model); let payload = self.build_antigravity_request(&actual_model, &project_id, request_body); + eprintln!("[ANTIGRAVITY_GENERATE] 构建的 payload: {}", serde_json::to_string_pretty(&payload).unwrap_or_default()); + eprintln!("[ANTIGRAVITY_GENERATE] 调用 call_api..."); let resp = self.call_api("generateContent", &payload).await?; + eprintln!("[ANTIGRAVITY_GENERATE] call_api 返回成功"); // 转换为 Gemini 格式响应 - Ok(self.to_gemini_response(&resp)) + let result = self.to_gemini_response(&resp); + eprintln!("========== [ANTIGRAVITY_GENERATE] 生成内容完成 =========="); + Ok(result) } /// 构建 Antigravity 请求 diff --git a/src-tauri/src/providers/mod.rs b/src-tauri/src/providers/mod.rs index 3a3225c93..0dd530d81 100644 --- a/src-tauri/src/providers/mod.rs +++ b/src-tauri/src/providers/mod.rs @@ -18,6 +18,8 @@ mod tests; #[allow(unused_imports)] pub use traits::{CredentialProvider, ProviderResult, TokenManager}; +#[allow(unused_imports)] +pub use antigravity::AntigravityApiError; #[allow(unused_imports)] pub use antigravity::AntigravityProvider; #[allow(unused_imports)] diff --git a/src-tauri/src/server/handlers/provider_calls.rs b/src-tauri/src/server/handlers/provider_calls.rs index a7dfe1149..97c5a39b8 100644 --- a/src-tauri/src/server/handlers/provider_calls.rs +++ b/src-tauri/src/server/handlers/provider_calls.rs @@ -59,13 +59,13 @@ use crate::models::anthropic::AnthropicMessagesRequest; use crate::models::openai::ChatCompletionRequest; use crate::models::provider_pool_model::{CredentialData, ProviderCredential}; use crate::providers::{ - AntigravityProvider, ClaudeCustomProvider, IFlowProvider, KiroProvider, OpenAICustomProvider, - VertexProvider, + AntigravityApiError, AntigravityProvider, ClaudeCustomProvider, IFlowProvider, KiroProvider, + OpenAICustomProvider, VertexProvider, }; use crate::server::AppState; use crate::server_utils::{ - build_anthropic_response, build_anthropic_stream_response, parse_cw_response, safe_truncate, - CWParsedResponse, + build_anthropic_response, build_anthropic_stream_response, build_error_response, + build_error_response_with_status, parse_cw_response, safe_truncate, CWParsedResponse, }; use crate::stream::{PipelineConfig, StreamPipeline}; use crate::streaming::traits::StreamingProvider; @@ -439,20 +439,18 @@ pub async fn call_provider_anthropic( build_anthropic_response(&request.model, &parsed) } } - Err(e) => { + Err(api_err) => { // 记录 API 调用失败 if let Some(db) = &state.db { let _ = state.pool_service.mark_unhealthy( db, &credential.uuid, - Some(&e.to_string()), + Some(&api_err.message), ); } - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response() + + // 直接使用 AntigravityApiError 的状态码构建响应 + build_error_response_with_status(api_err.status_code, &api_err.to_string()) } } } @@ -1441,13 +1439,10 @@ pub async fn call_provider_openai( .into_response() }); } - Err(e) => { - tracing::error!("[ANTIGRAVITY_STREAM] 图片生成失败: {}", e); - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(); + Err(api_err) => { + tracing::error!("[ANTIGRAVITY_STREAM] 图片生成失败 (HTTP {}): {}", api_err.status_code, api_err.message); + // 直接使用 AntigravityApiError 的状态码构建响应 + return build_error_response_with_status(api_err.status_code, &api_err.to_string()); } } } @@ -1540,31 +1535,41 @@ pub async fn call_provider_openai( .into_response() }); } - Err(e) => { - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(); + Err(provider_err) => { + // call_api_stream 返回 ProviderError,使用字符串解析状态码 + return build_error_response(&provider_err.to_string()); } } } // 非流式请求处理 + eprintln!("[ANTIGRAVITY_OPENAI] ========== 开始处理非流式请求 =========="); + eprintln!("[ANTIGRAVITY_OPENAI] 模型: {}", request.model); + // 获取 project_id 用于请求 let proj_id = antigravity.project_id.clone().unwrap_or_default(); + eprintln!("[ANTIGRAVITY_OPENAI] 项目ID: {}", proj_id); + // 转换请求格式 + eprintln!("[ANTIGRAVITY_OPENAI] 开始转换请求格式..."); let antigravity_request = convert_openai_to_antigravity_with_context(request, &proj_id); + eprintln!("[ANTIGRAVITY_OPENAI] 请求格式转换完成"); + + eprintln!("[ANTIGRAVITY_OPENAI] 调用 generate_content..."); match antigravity.generate_content(&request.model, &antigravity_request).await { Ok(resp) => { + eprintln!("[ANTIGRAVITY_OPENAI] generate_content 返回成功"); let openai_response = convert_antigravity_to_openai_response(&resp, &request.model); + eprintln!("[ANTIGRAVITY_OPENAI] ========== 非流式请求处理完成 =========="); Json(openai_response).into_response() } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .into_response(), + Err(api_err) => { + eprintln!("[ANTIGRAVITY_OPENAI] generate_content 失败 (HTTP {}): {}", api_err.status_code, api_err.message); + eprintln!("[ANTIGRAVITY_OPENAI] ========== 非流式请求处理失败 =========="); + + // 直接使用 AntigravityApiError 的状态码构建响应 + build_error_response_with_status(api_err.status_code, &api_err.to_string()) + } } } CredentialData::OpenAIKey { api_key, base_url } => { diff --git a/src-tauri/src/server/mod.rs b/src-tauri/src/server/mod.rs index 27635ce1b..ccf09925b 100644 --- a/src-tauri/src/server/mod.rs +++ b/src-tauri/src/server/mod.rs @@ -25,8 +25,8 @@ use crate::providers::kiro::KiroProvider; use crate::providers::openai_custom::OpenAICustomProvider; use crate::providers::qwen::QwenProvider; use crate::server_utils::{ - build_anthropic_response, build_anthropic_stream_response, build_gemini_native_request, health, - models, parse_cw_response, + build_anthropic_response, build_anthropic_stream_response, build_error_response, + build_error_response_with_status, build_gemini_native_request, health, models, parse_cw_response, }; use crate::services::kiro_event_service::KiroEventService; use crate::services::provider_pool_service::ProviderPoolService; @@ -171,6 +171,8 @@ pub struct ServerState { /// 服务器运行时使用的 API key(启动时从配置复制) /// 用于 test_api 命令,确保测试使用的 API key 和服务器一致 pub running_api_key: Option, + /// 服务器实际监听的 host(可能与配置不同,因为会自动切换到有效的 IP) + pub running_host: Option, } impl ServerState { @@ -196,13 +198,15 @@ impl ServerState { router_ref: None, shutdown_tx: None, running_api_key: None, + running_host: None, } } pub fn status(&self) -> ServerStatus { ServerStatus { running: self.running, - host: self.config.server.host.clone(), + // 使用实际运行的 host,如果没有则使用配置的 host + host: self.running_host.clone().unwrap_or_else(|| self.config.server.host.clone()), port: self.config.server.port, requests: self.requests, uptime_secs: self.start_time.map(|t| t.elapsed().as_secs()).unwrap_or(0), @@ -277,7 +281,48 @@ impl ServerState { let (tx, rx) = oneshot::channel(); self.shutdown_tx = Some(tx); - let host = self.config.server.host.clone(); + // 检查配置的 host 是否有效(在当前网卡列表中或是特殊地址) + let host = { + let configured_host = &self.config.server.host; + + // 特殊地址不需要检查 + if configured_host == "0.0.0.0" || configured_host == "127.0.0.1" || configured_host == "localhost" { + configured_host.clone() + } else { + // 检查 IP 是否在当前网卡列表中 + match crate::commands::network_cmd::get_network_info() { + Ok(network_info) => { + if network_info.all_ips.contains(configured_host) { + configured_host.clone() + } else { + // IP 不在当前网卡列表中,使用当前的局域网 IP + // 优先选择 192.168.x.x 或 10.x.x.x 开头的 IP(真正的局域网 IP) + let preferred_ip = network_info.all_ips.iter() + .find(|ip| ip.starts_with("192.168.") || ip.starts_with("10.")); + + let new_ip = preferred_ip + .or_else(|| network_info.lan_ip.as_ref()) + .or_else(|| network_info.all_ips.first()) + .cloned() + .unwrap_or_else(|| "127.0.0.1".to_string()); + + tracing::warn!( + "[SERVER] 配置的 IP {} 不在当前网卡列表中,自动切换到 {}", + configured_host, new_ip + ); + eprintln!( + "[SERVER] 警告:配置的 IP {} 不在当前网卡列表中,自动切换到 {}", + configured_host, new_ip + ); + new_ip + } + } + Err(_) => { + configured_host.clone() + } + } + } + }; let port = self.config.server.port; let api_key = self.config.server.api_key.clone(); let api_key_for_state = api_key.clone(); // 用于保存到 running_api_key @@ -346,6 +391,9 @@ impl ServerState { // 保存 router_ref 以便后续动态更新 self.router_ref = Some(processor.router.clone()); + // 保存实际使用的 host(在移动到 spawn 之前克隆) + let running_host = host.clone(); + tokio::spawn(async move { if let Err(e) = run_server( &host, @@ -379,6 +427,8 @@ impl ServerState { self.start_time = Some(std::time::Instant::now()); // 保存服务器运行时使用的 API key,用于 test_api 命令 self.running_api_key = Some(api_key_for_state); + // 保存服务器实际监听的 host(可能与配置不同) + self.running_host = Some(running_host); Ok(()) } @@ -389,6 +439,7 @@ impl ServerState { self.running = false; self.start_time = None; self.running_api_key = None; + self.running_host = None; self.router_ref = None; } } @@ -1275,22 +1326,15 @@ async fn gemini_generate_content( // 直接返回 Gemini 格式响应 Json(resp).into_response() } - Err(e) => { + Err(api_err) => { state .logs .write() .await - .add("error", &format!("[GEMINI] 请求失败: {}", e)); + .add("error", &format!("[GEMINI] 请求失败 (HTTP {}): {}", api_err.status_code, api_err.message)); - ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({ - "error": { - "message": e.to_string() - } - })), - ) - .into_response() + // 直接使用 AntigravityApiError 的状态码构建响应 + build_error_response_with_status(api_err.status_code, &api_err.to_string()) } } } @@ -1308,10 +1352,49 @@ async fn gemini_generate_content( /// 列出所有可用路由 async fn list_routes(State(state): State) -> impl IntoResponse { + // 处理 base_url:检查 IP 是否有效(在当前网卡列表中或是特殊地址) + let display_base_url = { + // 从 base_url 中提取 host 部分 + let url_parts: Vec<&str> = state.base_url.split("://").collect(); + let host_port = if url_parts.len() > 1 { url_parts[1] } else { &state.base_url }; + let host = host_port.split(':').next().unwrap_or("localhost"); + + // 检查是否需要替换 IP + let should_replace = if host == "0.0.0.0" || host == "127.0.0.1" || host == "localhost" { + // 0.0.0.0 需要替换为局域网 IP,127.0.0.1 和 localhost 保持不变 + host == "0.0.0.0" + } else { + // 检查 IP 是否在当前网卡列表中 + if let Ok(network_info) = crate::commands::network_cmd::get_network_info() { + !network_info.all_ips.contains(&host.to_string()) + } else { + false + } + }; + + if should_replace { + // 获取局域网 IP 进行替换 + // 优先选择 192.168.x.x 或 10.x.x.x 开头的 IP(真正的局域网 IP) + if let Ok(network_info) = crate::commands::network_cmd::get_network_info() { + let new_ip = network_info.all_ips.iter() + .find(|ip| ip.starts_with("192.168.") || ip.starts_with("10.")) + .or_else(|| network_info.lan_ip.as_ref()) + .or_else(|| network_info.all_ips.first()) + .cloned() + .unwrap_or_else(|| "localhost".to_string()); + state.base_url.replace(host, &new_ip) + } else { + state.base_url.replace(host, "localhost") + } + } else { + state.base_url.clone() + } + }; + let routes = match &state.db { Some(db) => state .pool_service - .get_available_routes(db, &state.base_url) + .get_available_routes(db, &display_base_url) .unwrap_or_default(), None => Vec::new(), }; @@ -1328,12 +1411,12 @@ async fn list_routes(State(state): State) -> impl IntoResponse { crate::models::route_model::RouteEndpoint { path: "/v1/messages".to_string(), protocol: "claude".to_string(), - url: format!("{}/v1/messages", state.base_url), + url: format!("{}/v1/messages", display_base_url), }, crate::models::route_model::RouteEndpoint { path: "/v1/chat/completions".to_string(), protocol: "openai".to_string(), - url: format!("{}/v1/chat/completions", state.base_url), + url: format!("{}/v1/chat/completions", display_base_url), }, ], tags: vec!["默认".to_string()], @@ -1342,7 +1425,7 @@ async fn list_routes(State(state): State) -> impl IntoResponse { all_routes.extend(routes); let response = RouteListResponse { - base_url: state.base_url.clone(), + base_url: display_base_url, default_provider, routes: all_routes, }; diff --git a/src-tauri/src/server_utils.rs b/src-tauri/src/server_utils.rs index 05c4bf362..60edde852 100644 --- a/src-tauri/src/server_utils.rs +++ b/src-tauri/src/server_utils.rs @@ -12,6 +12,90 @@ use axum::{ use futures::stream; use std::collections::HashMap; +/// 从错误信息中解析 HTTP 状态码 +/// +/// 用于将上游 API 返回的错误状态码透传给客户端,而不是统一返回 500。 +/// 支持解析常见的 HTTP 状态码:429、403、401、404、400、503、502、500。 +/// +/// # 参数 +/// - `error_message`: 错误信息字符串,通常包含状态码(如 "API call failed: 429 - ...") +/// +/// # 返回 +/// 解析出的 HTTP 状态码,如果无法解析则返回 500 INTERNAL_SERVER_ERROR +/// +/// # 示例 +/// ``` +/// let status = parse_error_status_code("API call failed: 429 Too Many Requests"); +/// assert_eq!(status, StatusCode::TOO_MANY_REQUESTS); +/// ``` +pub fn parse_error_status_code(error_message: &str) -> StatusCode { + if error_message.contains("429") { + StatusCode::TOO_MANY_REQUESTS + } else if error_message.contains("403") { + StatusCode::FORBIDDEN + } else if error_message.contains("401") { + StatusCode::UNAUTHORIZED + } else if error_message.contains("404") { + StatusCode::NOT_FOUND + } else if error_message.contains("400") { + StatusCode::BAD_REQUEST + } else if error_message.contains("503") { + StatusCode::SERVICE_UNAVAILABLE + } else if error_message.contains("502") { + StatusCode::BAD_GATEWAY + } else if error_message.contains("500") { + StatusCode::INTERNAL_SERVER_ERROR + } else { + StatusCode::INTERNAL_SERVER_ERROR + } +} + +/// 构建错误响应 +/// +/// 从错误信息中解析状态码并构建标准的 JSON 错误响应。 +/// +/// # 参数 +/// - `error_message`: 错误信息字符串 +/// +/// # 返回 +/// 包含正确状态码的 HTTP 响应 +pub fn build_error_response(error_message: &str) -> Response { + let status_code = parse_error_status_code(error_message); + ( + status_code, + Json(serde_json::json!({ + "error": { + "message": error_message + } + })), + ) + .into_response() +} + +/// 从 HTTP 状态码构建错误响应 +/// +/// 直接使用状态码构建响应,无需解析字符串。 +/// 适用于已知状态码的场景(如 AntigravityApiError)。 +/// +/// # 参数 +/// - `status_code`: HTTP 状态码(u16) +/// - `error_message`: 错误信息字符串 +/// +/// # 返回 +/// 包含指定状态码的 HTTP 响应 +pub fn build_error_response_with_status(status_code: u16, error_message: &str) -> Response { + let status = StatusCode::from_u16(status_code).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); + ( + status, + Json(serde_json::json!({ + "error": { + "message": error_message + } + })), + ) + .into_response() +} + /// CodeWhisperer 响应解析结果 #[derive(Debug, Default)] pub struct CWParsedResponse { @@ -620,7 +704,7 @@ pub async fn health() -> impl IntoResponse { })) } -/// 模型列表端点响应 +/// 模型列表端点响应(静态列表,用于不指定凭证的情况) pub async fn models() -> impl IntoResponse { Json(serde_json::json!({ "object": "list", diff --git a/src-tauri/src/services/api_key_provider_service.rs b/src-tauri/src/services/api_key_provider_service.rs index 90d9a4290..d67834b71 100644 --- a/src-tauri/src/services/api_key_provider_service.rs +++ b/src-tauri/src/services/api_key_provider_service.rs @@ -1085,6 +1085,7 @@ impl ApiKeyProviderService { check_health: false, check_model_name: None, not_supported_models: Vec::new(), + supported_models: Vec::new(), usage_count: 0, error_count: 0, last_used: None, @@ -1142,6 +1143,7 @@ impl ApiKeyProviderService { check_health: false, // 降级凭证不参与健康检查 check_model_name: None, not_supported_models: Vec::new(), + supported_models: Vec::new(), usage_count: 0, error_count: 0, last_used: None, diff --git a/src-tauri/src/services/mod.rs b/src-tauri/src/services/mod.rs index 8df69d58d..526c0b7e0 100644 --- a/src-tauri/src/services/mod.rs +++ b/src-tauri/src/services/mod.rs @@ -7,6 +7,7 @@ pub mod machine_id_service; pub mod mcp_service; pub mod mcp_sync; pub mod model_registry_service; +pub mod model_service; pub mod prompt_service; pub mod prompt_sync; pub mod provider_pool_service; diff --git a/src-tauri/src/services/model_service.rs b/src-tauri/src/services/model_service.rs new file mode 100644 index 000000000..9a864fc3f --- /dev/null +++ b/src-tauri/src/services/model_service.rs @@ -0,0 +1,467 @@ +//! 模型管理服务 +//! +//! 提供统一的模型获取、缓存和查询接口,支持从不同 Provider 获取模型列表。 + +use crate::database::dao::provider_pool::ProviderPoolDao; +use crate::database::DbConnection; +use crate::models::provider_pool_model::{CredentialData, PoolProviderType, ProviderCredential}; +use reqwest::Client; +use serde::{Deserialize, Serialize}; +use std::collections::HashMap; +use std::time::Duration; + +/// 模型信息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ModelInfo { + /// 模型 ID + pub id: String, + /// 模型对象类型(通常是 "model") + pub object: String, + /// 拥有者(如 "anthropic", "google", "openai") + pub owned_by: String, + /// 创建时间(可选) + #[serde(skip_serializing_if = "Option::is_none")] + pub created: Option, +} + +/// /v1/models 接口的响应格式 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ModelsResponse { + pub object: String, + pub data: Vec, +} + +/// 模型服务 +pub struct ModelService { + /// HTTP 客户端 + client: Client, + /// 请求超时时间 + timeout: Duration, +} + +impl Default for ModelService { + fn default() -> Self { + Self::new() + } +} + +impl ModelService { + /// 创建新的模型服务实例 + pub fn new() -> Self { + Self { + client: Client::builder() + .timeout(Duration::from_secs(10)) + .build() + .unwrap_or_default(), + timeout: Duration::from_secs(10), + } + } + + /// 从凭证获取支持的模型列表 + /// + /// 根据凭证类型调用相应的 /v1/models 接口 + pub async fn fetch_models_for_credential( + &self, + credential: &ProviderCredential, + ) -> Result, String> { + tracing::info!( + "[MODEL_SERVICE] 获取凭证模型列表: uuid={}, provider_type={}", + credential.uuid, + credential.provider_type + ); + + match &credential.credential { + // Antigravity 使用固定的模型列表(从配置文件读取) + CredentialData::AntigravityOAuth { .. } => { + // Antigravity 不提供标准的 /v1/models 接口 + // 直接返回预定义的模型列表 + tracing::info!("[MODEL_SERVICE] Antigravity 使用预定义模型列表"); + Ok(self.get_default_models_for_provider(&credential.provider_type)) + } + + // OAuth 凭证:由于需要处理 Token 刷新等复杂逻辑,暂时使用默认模型列表 + // TODO: 未来可以通过 ProviderPoolService 来获取动态模型列表 + CredentialData::KiroOAuth { .. } + | CredentialData::GeminiOAuth { .. } + | CredentialData::QwenOAuth { .. } + | CredentialData::CodexOAuth { .. } + | CredentialData::ClaudeOAuth { .. } + | CredentialData::IFlowOAuth { .. } + | CredentialData::IFlowCookie { .. } => { + tracing::info!("[MODEL_SERVICE] OAuth 凭证使用默认模型列表"); + Ok(self.get_default_models_for_provider(&credential.provider_type)) + } + + // API Key 类型凭证:直接调用 Provider 的 API + CredentialData::OpenAIKey { base_url, api_key } => { + tracing::info!("[MODEL_SERVICE] 使用 OpenAI API Key"); + self.fetch_models_openai(base_url.as_deref(), api_key).await + } + CredentialData::ClaudeKey { base_url, api_key } => { + tracing::info!("[MODEL_SERVICE] 使用 Claude API Key"); + self.fetch_models_claude(base_url.as_deref(), api_key).await + } + CredentialData::AnthropicKey { base_url, api_key } => { + tracing::info!("[MODEL_SERVICE] 使用 Anthropic API Key"); + self.fetch_models_anthropic(base_url.as_deref(), api_key) + .await + } + CredentialData::GeminiApiKey { + api_key, base_url, .. + } => { + tracing::info!("[MODEL_SERVICE] 使用 Gemini API Key"); + self.fetch_models_gemini(base_url.as_deref(), api_key).await + } + CredentialData::VertexKey { .. } => { + tracing::info!("[MODEL_SERVICE] Vertex AI 使用固定模型列表"); + // Vertex AI 使用固定的模型列表 + Ok(self.get_default_models_for_provider(&credential.provider_type)) + } + } + } + + /// 获取 OpenAI 兼容 API 的模型列表 + async fn fetch_models_openai( + &self, + base_url: Option<&str>, + api_key: &str, + ) -> Result, String> { + let url = format!("{}/v1/models", base_url.unwrap_or("https://api.openai.com")); + + tracing::info!("[MODEL_SERVICE] 请求 OpenAI API 获取模型列表: url={}", url); + + let response = self + .client + .get(&url) + .header("Authorization", format!("Bearer {}", api_key)) + .timeout(self.timeout) + .send() + .await + .map_err(|e| { + tracing::error!("[MODEL_SERVICE] OpenAI 请求失败: {}", e); + format!("请求失败: {}", e) + })?; + + let status = response.status(); + tracing::info!("[MODEL_SERVICE] OpenAI 响应状态码: {}", status); + + if !status.is_success() { + let error_body = response.text().await.unwrap_or_default(); + tracing::error!("[MODEL_SERVICE] OpenAI HTTP 错误: status={}, body={}", status, error_body); + return Err(format!("HTTP 错误: {}", status)); + } + + let response_text = response.text().await.map_err(|e| { + tracing::error!("[MODEL_SERVICE] 读取 OpenAI 响应体失败: {}", e); + format!("读取响应体失败: {}", e) + })?; + + tracing::debug!("[MODEL_SERVICE] OpenAI 响应体: {}", response_text); + + let models_response: ModelsResponse = serde_json::from_str(&response_text).map_err(|e| { + tracing::error!("[MODEL_SERVICE] 解析 OpenAI 响应失败: {}, 响应内容: {}", e, response_text); + format!("解析响应失败: {}", e) + })?; + + let model_ids: Vec = models_response.data.into_iter().map(|m| m.id).collect(); + + tracing::info!("[MODEL_SERVICE] OpenAI 成功获取 {} 个模型", model_ids.len()); + + Ok(model_ids) + } + + /// 获取 Claude API 的模型列表 + async fn fetch_models_claude( + &self, + base_url: Option<&str>, + api_key: &str, + ) -> Result, String> { + // Claude API 使用 OpenAI 兼容格式 + self.fetch_models_openai(base_url, api_key).await + } + + /// 获取 Anthropic API 的模型列表 + async fn fetch_models_anthropic( + &self, + base_url: Option<&str>, + api_key: &str, + ) -> Result, String> { + // Anthropic API 使用 OpenAI 兼容格式 + let url = format!( + "{}/v1/models", + base_url.unwrap_or("https://api.anthropic.com") + ); + + tracing::info!("[MODEL_SERVICE] 请求 Anthropic API 获取模型列表: url={}", url); + + let response = self + .client + .get(&url) + .header("x-api-key", api_key) + .header("anthropic-version", "2023-06-01") + .timeout(self.timeout) + .send() + .await + .map_err(|e| { + tracing::error!("[MODEL_SERVICE] Anthropic 请求失败: {}", e); + format!("请求失败: {}", e) + })?; + + let status = response.status(); + tracing::info!("[MODEL_SERVICE] Anthropic 响应状态码: {}", status); + + if !status.is_success() { + let error_body = response.text().await.unwrap_or_default(); + tracing::error!("[MODEL_SERVICE] Anthropic HTTP 错误: status={}, body={}", status, error_body); + return Err(format!("HTTP 错误: {}", status)); + } + + let response_text = response.text().await.map_err(|e| { + tracing::error!("[MODEL_SERVICE] 读取 Anthropic 响应体失败: {}", e); + format!("读取响应体失败: {}", e) + })?; + + tracing::debug!("[MODEL_SERVICE] Anthropic 响应体: {}", response_text); + + let models_response: ModelsResponse = serde_json::from_str(&response_text).map_err(|e| { + tracing::error!("[MODEL_SERVICE] 解析 Anthropic 响应失败: {}, 响应内容: {}", e, response_text); + format!("解析响应失败: {}", e) + })?; + + let model_ids: Vec = models_response.data.into_iter().map(|m| m.id).collect(); + + tracing::info!("[MODEL_SERVICE] Anthropic 成功获取 {} 个模型", model_ids.len()); + + Ok(model_ids) + } + + /// 获取 Gemini API 的模型列表 + async fn fetch_models_gemini( + &self, + base_url: Option<&str>, + api_key: &str, + ) -> Result, String> { + let url = format!( + "{}/v1/models?key={}", + base_url.unwrap_or("https://generativelanguage.googleapis.com"), + api_key + ); + + tracing::info!("[MODEL_SERVICE] 请求 Gemini API 获取模型列表: url={}", url); + + let response = self + .client + .get(&url) + .timeout(self.timeout) + .send() + .await + .map_err(|e| { + tracing::error!("[MODEL_SERVICE] Gemini 请求失败: {}", e); + format!("请求失败: {}", e) + })?; + + let status = response.status(); + tracing::info!("[MODEL_SERVICE] Gemini 响应状态码: {}", status); + + if !status.is_success() { + let error_body = response.text().await.unwrap_or_default(); + tracing::error!("[MODEL_SERVICE] Gemini HTTP 错误: status={}, body={}", status, error_body); + return Err(format!("HTTP 错误: {}", status)); + } + + let response_text = response.text().await.map_err(|e| { + tracing::error!("[MODEL_SERVICE] 读取 Gemini 响应体失败: {}", e); + format!("读取响应体失败: {}", e) + })?; + + tracing::debug!("[MODEL_SERVICE] Gemini 响应体: {}", response_text); + + // Gemini API 返回格式不同,需要特殊处理 + let response_json: serde_json::Value = serde_json::from_str(&response_text).map_err(|e| { + tracing::error!("[MODEL_SERVICE] 解析 Gemini 响应失败: {}, 响应内容: {}", e, response_text); + format!("解析响应失败: {}", e) + })?; + + let models = response_json + .get("models") + .and_then(|m| m.as_array()) + .ok_or_else(|| { + tracing::error!("[MODEL_SERVICE] Gemini 响应格式错误: 缺少 models 字段"); + "响应格式错误".to_string() + })?; + + let model_ids: Vec = models + .iter() + .filter_map(|m| m.get("name").and_then(|n| n.as_str())) + .map(|name| { + // Gemini API 返回的是 "models/gemini-pro",需要提取模型名 + name.strip_prefix("models/").unwrap_or(name).to_string() + }) + .collect(); + + tracing::info!( + "[MODEL_SERVICE] Gemini 成功获取 {} 个模型: {:?}", + model_ids.len(), + model_ids + ); + + Ok(model_ids) + } + + /// 获取 Provider 的默认模型列表(用于无法动态获取的情况) + pub fn get_default_models_for_provider(&self, provider_type: &PoolProviderType) -> Vec { + match provider_type { + PoolProviderType::Kiro => vec![ + "claude-sonnet-4-5".to_string(), + "claude-sonnet-4-5-20250929".to_string(), + "claude-3-7-sonnet-20250219".to_string(), + "claude-3-5-sonnet-latest".to_string(), + "claude-haiku-4-5".to_string(), + ], + PoolProviderType::Gemini => vec![ + "gemini-2.5-flash".to_string(), + "gemini-2.5-flash-lite".to_string(), + "gemini-2.5-pro".to_string(), + "gemini-2.5-pro-preview-06-05".to_string(), + ], + PoolProviderType::Qwen => vec![ + "qwen3-coder-plus".to_string(), + "qwen3-coder-flash".to_string(), + ], + PoolProviderType::Antigravity => vec![ + "gemini-2.5-computer-use-preview-10-2025".to_string(), + "gemini-3-pro-image-preview".to_string(), + "gemini-3-pro-preview".to_string(), + "gemini-3-flash-preview".to_string(), + "gemini-2.5-flash-preview".to_string(), + "gemini-claude-sonnet-4-5".to_string(), + "gemini-claude-sonnet-4-5-thinking".to_string(), + "gemini-claude-opus-4-5-thinking".to_string(), + ], + PoolProviderType::OpenAI => vec![ + "gpt-4o".to_string(), + "gpt-4o-mini".to_string(), + "gpt-3.5-turbo".to_string(), + ], + PoolProviderType::Claude | PoolProviderType::Anthropic => vec![ + "claude-sonnet-4-5-20250929".to_string(), + "claude-3-5-sonnet-20241022".to_string(), + "claude-3-5-haiku-20241022".to_string(), + ], + PoolProviderType::GeminiApiKey => vec![ + "gemini-2.5-flash".to_string(), + "gemini-2.5-pro".to_string(), + ], + _ => vec![], + } + } + + /// 更新凭证的支持模型列表到数据库 + pub fn update_credential_models( + &self, + db: &DbConnection, + credential_uuid: &str, + models: Vec, + ) -> Result<(), String> { + let conn = db.lock().map_err(|e| e.to_string())?; + + // 序列化模型列表为 JSON + let models_json = serde_json::to_string(&models).map_err(|e| e.to_string())?; + + conn.execute( + "UPDATE provider_pool_credentials SET supported_models = ?1, updated_at = ?2 WHERE uuid = ?3", + rusqlite::params![models_json, chrono::Utc::now().timestamp(), credential_uuid], + ) + .map_err(|e| e.to_string())?; + + Ok(()) + } + + /// 获取凭证的支持模型列表(从数据库) + pub fn get_credential_models( + &self, + db: &DbConnection, + credential_uuid: &str, + ) -> Result, String> { + let conn = db.lock().map_err(|e| e.to_string())?; + + let mut stmt = conn + .prepare("SELECT supported_models FROM provider_pool_credentials WHERE uuid = ?1") + .map_err(|e| e.to_string())?; + + let models_json: Option = stmt + .query_row([credential_uuid], |row| row.get(0)) + .ok(); + + match models_json { + Some(json) => serde_json::from_str(&json).map_err(|e| e.to_string()), + None => Ok(vec![]), + } + } + + /// 获取所有凭证的模型列表(按 Provider 类型分组) + pub fn get_all_models_by_provider( + &self, + db: &DbConnection, + ) -> Result>, String> { + let conn = db.lock().map_err(|e| e.to_string())?; + let credentials = ProviderPoolDao::get_all(&conn).map_err(|e| e.to_string())?; + + let mut models_by_provider: HashMap> = HashMap::new(); + + for cred in credentials { + if cred.is_disabled || !cred.is_healthy { + continue; + } + + let models = self.get_credential_models(db, &cred.uuid)?; + let provider_key = cred.provider_type.to_string(); + + models_by_provider + .entry(provider_key) + .or_default() + .extend(models); + } + + // 去重 + for models in models_by_provider.values_mut() { + models.sort(); + models.dedup(); + } + + Ok(models_by_provider) + } + + /// 获取可用的所有模型列表(合并所有健康凭证的模型) + pub fn get_all_available_models(&self, db: &DbConnection) -> Result, String> { + let models_by_provider = self.get_all_models_by_provider(db)?; + + let mut all_models: Vec = models_by_provider + .into_values() + .flatten() + .collect(); + + all_models.sort(); + all_models.dedup(); + + Ok(all_models) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_get_default_models_for_provider() { + let service = ModelService::new(); + + let kiro_models = service.get_default_models_for_provider(&PoolProviderType::Kiro); + assert!(!kiro_models.is_empty()); + assert!(kiro_models.contains(&"claude-sonnet-4-5".to_string())); + + let gemini_models = service.get_default_models_for_provider(&PoolProviderType::Gemini); + assert!(!gemini_models.is_empty()); + assert!(gemini_models.contains(&"gemini-2.5-flash".to_string())); + } +} diff --git a/src/components/api-server/ApiServerPage.tsx b/src/components/api-server/ApiServerPage.tsx index dadac7102..e351f2bf4 100644 --- a/src/components/api-server/ApiServerPage.tsx +++ b/src/components/api-server/ApiServerPage.tsx @@ -109,6 +109,34 @@ export function ApiServerPage() { } }; + const loadNetworkInfo = async () => { + try { + const info = await getNetworkInfo(); + setNetworkInfo(info); + + // 如果配置的 host 不在当前网卡列表中(且不是 127.0.0.1 或 0.0.0.0), + // 自动更新为当前的局域网 IP + if (config && editHost) { + const isValidHost = + editHost === "127.0.0.1" || + editHost === "0.0.0.0" || + info.all_ips.includes(editHost); + + if (!isValidHost && info.all_ips.length > 0) { + // 选择第一个局域网 IP(通常是 192.168.x.x 或 10.x.x.x) + const lanIp = info.all_ips.find(ip => + ip.startsWith("192.168.") || ip.startsWith("10.") + ) || info.all_ips[0]; + + console.log(`配置的 IP ${editHost} 不在当前网卡列表中,自动更新为 ${lanIp}`); + setEditHost(lanIp); + } + } + } catch (e) { + console.error("Failed to get network info:", e); + } + }; + useEffect(() => { fetchStatus(); fetchConfig(); @@ -116,19 +144,32 @@ export function ApiServerPage() { loadNetworkInfo(); const statusInterval = setInterval(fetchStatus, 3000); - return () => clearInterval(statusInterval); + // 定期刷新网络信息,以便检测 IP 变化 + const networkInterval = setInterval(loadNetworkInfo, 5000); + return () => { + clearInterval(statusInterval); + clearInterval(networkInterval); + }; }, []); - const loadNetworkInfo = async () => { - try { - const info = await getNetworkInfo(); - setNetworkInfo(info); - } catch (e) { - console.error("Failed to get network info:", e); + // 当 config 和 editHost 加载完成后,检查并更新网络信息 + useEffect(() => { + if (config && editHost && networkInfo) { + const isValidHost = + editHost === "127.0.0.1" || + editHost === "0.0.0.0" || + networkInfo.all_ips.includes(editHost); + + if (!isValidHost && networkInfo.all_ips.length > 0) { + const lanIp = networkInfo.all_ips.find(ip => + ip.startsWith("192.168.") || ip.startsWith("10.") + ) || networkInfo.all_ips[0]; + + console.log(`配置的 IP ${editHost} 不在当前网卡列表中,自动更新为 ${lanIp}`); + setEditHost(lanIp); + } } - }; - - const loadDefaultProvider = async () => { + }, [config, networkInfo]); const loadDefaultProvider = async () => { try { const dp = await getDefaultProvider(); setDefaultProviderState(dp); @@ -143,6 +184,8 @@ export function ApiServerPage() { try { await reloadCredentials(); await startServer(); + // 等待服务器完全启动 + await new Promise((resolve) => setTimeout(resolve, 500)); await fetchStatus(); setMessage({ type: "success", text: "服务已启动" }); } catch (e: unknown) { @@ -408,14 +451,13 @@ export function ApiServerPage() { const handleSetDefaultProvider = async (providerId: string) => { try { - await setDefaultProvider(providerId); + // 先更新 UI 状态,提供即时反馈 setDefaultProviderState(providerId); + + // 异步调用后端 + await setDefaultProvider(providerId); - // 获取最新的凭证池数据 - const freshOverview = await providerPoolApi.getOverview(); - setPoolOverview(freshOverview); - - // 获取该 Provider 的凭证信息 + // 获取该 Provider 的凭证信息(用于显示消息) const provider = availableProviders.find((p) => p.id === providerId); const label = providerLabels[providerId] || providerId; @@ -434,25 +476,51 @@ export function ApiServerPage() { } else { setProviderSwitchMsg(`已切换到 ${label}`); } + + // 在后台异步刷新凭证池数据,不阻塞 UI + providerPoolApi.getOverview().then(setPoolOverview).catch(console.error); } catch (e: unknown) { const errMsg = e instanceof Error ? e.message : String(e); setProviderSwitchMsg(`切换失败: ${errMsg}`); + // 切换失败时恢复原来的状态 + loadDefaultProvider(); } }; // 根据监听地址智能选择测试 URL // - 127.0.0.1: 使用 127.0.0.1(仅本机) - // - 0.0.0.0: 使用 127.0.0.1(本机访问所有接口) + // - 0.0.0.0: 使用当前局域网 IP(优先 192.168.x.x 或 10.x.x.x) // - 局域网 IP: 使用该 IP(允许局域网测试) const getTestUrl = (host: string, port: number) => { if (host === "0.0.0.0") { - return `http://127.0.0.1:${port}`; + // 0.0.0.0 时,使用当前局域网 IP 以便局域网设备访问 + const lanIp = networkInfo?.all_ips.find(ip => + ip.startsWith("192.168.") || ip.startsWith("10.") + ) || networkInfo?.all_ips[0] || "127.0.0.1"; + return `http://${lanIp}:${port}`; } return `http://${host}:${port}`; }; // 使用 editHost 而不是 status.host,这样可以实时反映用户的选择 - const currentHost = status?.running ? status.host : editHost; + // 同时检查配置的 IP 是否仍然有效(在当前网卡列表中) + const getValidHost = () => { + const host = status?.running ? status.host : editHost; + // 如果是特殊地址,直接返回 + if (host === "127.0.0.1" || host === "0.0.0.0") { + return host; + } + // 检查配置的 IP 是否在当前网卡列表中 + if (networkInfo?.all_ips && !networkInfo.all_ips.includes(host)) { + // IP 已失效,返回当前有效的局域网 IP + return networkInfo.all_ips.find(ip => + ip.startsWith("192.168.") || ip.startsWith("10.") + ) || networkInfo.all_ips[0] || host; + } + return host; + }; + + const currentHost = getValidHost(); const currentPort = status?.running ? status.port : parseInt(editPort) || 8999; @@ -599,10 +667,7 @@ export function ApiServerPage() { }, })); - // 测试成功后立即刷新凭证池数据,更新使用次数 - if (result.success) { - await loadPoolOverview(); - } + return result.success; } catch (e: unknown) { const errMsg = e instanceof Error ? e.message : String(e); setTestResults((prev) => ({ @@ -613,12 +678,21 @@ export function ApiServerPage() { response: `请求失败: ${errMsg}`, }, })); + return false; } }; const runAllTests = async () => { + let hasSuccess = false; for (const endpoint of testEndpoints) { - await runTest(endpoint); + const success = await runTest(endpoint); + if (success) { + hasSuccess = true; + } + } + // 所有测试完成后,如果有成功的测试,刷新一次凭证池数据 + if (hasSuccess) { + await loadPoolOverview(); } }; @@ -811,6 +885,30 @@ export function ApiServerPage() { ))} + + {/* 当前值不在预定义选项中时,显示为自定义选项 */} + {editHost && + editHost !== "127.0.0.1" && + editHost !== "0.0.0.0" && + !(networkInfo?.all_ips.includes(editHost) ?? false) && ( + + + + + + + {editHost} + + (自定义) + + + + + )} diff --git a/src/lib/api/providerPool.ts b/src/lib/api/providerPool.ts index 57fe23eca..0121f90ff 100644 --- a/src/lib/api/providerPool.ts +++ b/src/lib/api/providerPool.ts @@ -580,6 +580,40 @@ export const providerPoolApi = { async getAllCredentialHealth(): Promise { return safeInvoke("get_all_credential_health"); }, + + // ============ 模型管理 ============ + + // 获取凭证支持的模型列表(从数据库缓存) + async getCredentialModels(credentialUuid: string): Promise { + return safeInvoke("get_credential_models", { credentialUuid }); + }, + + // 刷新凭证的模型列表(从 Provider API 重新获取) + async refreshCredentialModels(credentialUuid: string): Promise { + return safeInvoke("refresh_credential_models", { credentialUuid }); + }, + + // 获取所有凭证的模型列表(按 Provider 类型分组) + async getAllModelsByProvider(): Promise> { + return safeInvoke("get_all_models_by_provider"); + }, + + // 获取所有可用的模型列表(合并所有健康凭证的模型) + async getAllAvailableModels(): Promise { + return safeInvoke("get_all_available_models"); + }, + + // 批量刷新所有凭证的模型列表 + async refreshAllCredentialModels(): Promise< + Record + > { + return safeInvoke("refresh_all_credential_models"); + }, + + // 获取 Provider 的默认模型列表 + async getDefaultModelsForProvider(providerType: string): Promise { + return safeInvoke("get_default_models_for_provider", { providerType }); + }, }; // Migration result From 8564c9007e5295d469babf929d05ac0b99ce79fa Mon Sep 17 00:00:00 2001 From: Chiron <598621670@qq.com> Date: Mon, 12 Jan 2026 13:36:27 +0800 Subject: [PATCH 2/2] =?UTF-8?q?feat:=20=E4=BF=AE=E5=A4=8D=E6=99=BA?= =?UTF-8?q?=E8=B0=B1=20API=20=E8=B7=AF=E5=BE=84=E9=97=AE=E9=A2=98=E5=B9=B6?= =?UTF-8?q?=E6=94=AF=E6=8C=81=E8=87=AA=E5=AE=9A=E4=B9=89=E6=A8=A1=E5=9E=8B?= =?UTF-8?q?=E5=88=97=E8=A1=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 修复: 1. 修复 OpenAI Custom Provider 的 URL 路径拼接问题 - 支持检测 /v1, /v2, /v3, /v4 等版本号路径 - 避免重复拼接导致 /v4/v1/chat/completions 错误 新功能: 1. API Key Provider 支持自定义模型列表 - 数据库添加 custom_models 字段 - 前端 Provider 配置表单添加自定义模型输入 - 用于不支持 /models 接口的 Provider(如智谱) 2. API Server 测试页面使用自定义模型 - 自动使用当前 Provider 的自定义模型进行测试 - 为多个自定义模型生成独立测试端点 --- pnpm-lock.yaml | 66 +++++++++++++++++++ .../src/commands/api_key_provider_cmd.rs | 6 ++ .../src/database/dao/api_key_provider.rs | 56 ++++++++++++---- src-tauri/src/database/schema.rs | 7 ++ src-tauri/src/database/system_providers.rs | 1 + src-tauri/src/providers/README.md | 8 --- src-tauri/src/providers/openai_custom.rs | 32 +++++++-- .../src/services/api_key_provider_service.rs | 5 ++ src/components/api-server/ApiServerPage.tsx | 40 ++++++++++- .../api-key/ProviderConfigForm.tsx | 30 +++++++++ src/lib/api/apiKeyProvider.ts | 4 ++ 11 files changed, 229 insertions(+), 26 deletions(-) diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index ea8c1db8f..f54b81db2 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -71,6 +71,9 @@ importers: '@tauri-apps/plugin-shell': specifier: ^2.0.0 version: 2.3.3 + '@tonejs/midi': + specifier: ^2.0.28 + version: 2.0.28 '@types/lodash-es': specifier: ^4.17.12 version: 4.17.12 @@ -164,6 +167,9 @@ importers: tailwind-merge: specifier: ^2.6.0 version: 2.6.0 + tone: + specifier: ^15.1.22 + version: 15.1.22 devDependencies: '@babel/plugin-transform-react-jsx-source': specifier: ^7.27.1 @@ -1320,56 +1326,67 @@ packages: resolution: {integrity: sha512-EHMUcDwhtdRGlXZsGSIuXSYwD5kOT9NVnx9sqzYiwAc91wfYOE1g1djOEDseZJKKqtHAHGwnGPQu3kytmfaXLQ==} cpu: [arm] os: [linux] + libc: [glibc] '@rollup/rollup-linux-arm-musleabihf@4.54.0': resolution: {integrity: sha512-+pBrqEjaakN2ySv5RVrj/qLytYhPKEUwk+e3SFU5jTLHIcAtqh2rLrd/OkbNuHJpsBgxsD8ccJt5ga/SeG0JmA==} cpu: [arm] os: [linux] + libc: [musl] '@rollup/rollup-linux-arm64-gnu@4.54.0': resolution: {integrity: sha512-NSqc7rE9wuUaRBsBp5ckQ5CVz5aIRKCwsoa6WMF7G01sX3/qHUw/z4pv+D+ahL1EIKy6Enpcnz1RY8pf7bjwng==} cpu: [arm64] os: [linux] + libc: [glibc] '@rollup/rollup-linux-arm64-musl@4.54.0': resolution: {integrity: sha512-gr5vDbg3Bakga5kbdpqx81m2n9IX8M6gIMlQQIXiLTNeQW6CucvuInJ91EuCJ/JYvc+rcLLsDFcfAD1K7fMofg==} cpu: [arm64] os: [linux] + libc: [musl] '@rollup/rollup-linux-loong64-gnu@4.54.0': resolution: {integrity: sha512-gsrtB1NA3ZYj2vq0Rzkylo9ylCtW/PhpLEivlgWe0bpgtX5+9j9EZa0wtZiCjgu6zmSeZWyI/e2YRX1URozpIw==} cpu: [loong64] os: [linux] + libc: [glibc] '@rollup/rollup-linux-ppc64-gnu@4.54.0': resolution: {integrity: sha512-y3qNOfTBStmFNq+t4s7Tmc9hW2ENtPg8FeUD/VShI7rKxNW7O4fFeaYbMsd3tpFlIg1Q8IapFgy7Q9i2BqeBvA==} cpu: [ppc64] os: [linux] + libc: [glibc] '@rollup/rollup-linux-riscv64-gnu@4.54.0': resolution: {integrity: sha512-89sepv7h2lIVPsFma8iwmccN7Yjjtgz0Rj/Ou6fEqg3HDhpCa+Et+YSufy27i6b0Wav69Qv4WBNl3Rs6pwhebQ==} cpu: [riscv64] os: [linux] + libc: [glibc] '@rollup/rollup-linux-riscv64-musl@4.54.0': resolution: {integrity: sha512-ZcU77ieh0M2Q8Ur7D5X7KvK+UxbXeDHwiOt/CPSBTI1fBmeDMivW0dPkdqkT4rOgDjrDDBUed9x4EgraIKoR2A==} cpu: [riscv64] os: [linux] + libc: [musl] '@rollup/rollup-linux-s390x-gnu@4.54.0': resolution: {integrity: sha512-2AdWy5RdDF5+4YfG/YesGDDtbyJlC9LHmL6rZw6FurBJ5n4vFGupsOBGfwMRjBYH7qRQowT8D/U4LoSvVwOhSQ==} cpu: [s390x] os: [linux] + libc: [glibc] '@rollup/rollup-linux-x64-gnu@4.54.0': resolution: {integrity: sha512-WGt5J8Ij/rvyqpFexxk3ffKqqbLf9AqrTBbWDk7ApGUzaIs6V+s2s84kAxklFwmMF/vBNGrVdYgbblCOFFezMQ==} cpu: [x64] os: [linux] + libc: [glibc] '@rollup/rollup-linux-x64-musl@4.54.0': resolution: {integrity: sha512-JzQmb38ATzHjxlPHuTH6tE7ojnMKM2kYNzt44LO/jJi8BpceEC8QuXYA908n8r3CNuG/B3BV8VR3Hi1rYtmPiw==} cpu: [x64] os: [linux] + libc: [musl] '@rollup/rollup-openharmony-arm64@4.54.0': resolution: {integrity: sha512-huT3fd0iC7jigGh7n3q/+lfPcXxBi+om/Rs3yiFxjvSxbSB6aohDFXbWvlspaqjeOh+hx7DDHS+5Es5qRkWkZg==} @@ -1493,30 +1510,35 @@ packages: engines: {node: '>= 10'} cpu: [arm64] os: [linux] + libc: [glibc] '@tauri-apps/cli-linux-arm64-musl@2.9.6': resolution: {integrity: sha512-02TKUndpodXBCR0oP//6dZWGYcc22Upf2eP27NvC6z0DIqvkBBFziQUcvi2n6SrwTRL0yGgQjkm9K5NIn8s6jw==} engines: {node: '>= 10'} cpu: [arm64] os: [linux] + libc: [musl] '@tauri-apps/cli-linux-riscv64-gnu@2.9.6': resolution: {integrity: sha512-fmp1hnulbqzl1GkXl4aTX9fV+ubHw2LqlLH1PE3BxZ11EQk+l/TmiEongjnxF0ie4kV8DQfDNJ1KGiIdWe1GvQ==} engines: {node: '>= 10'} cpu: [riscv64] os: [linux] + libc: [glibc] '@tauri-apps/cli-linux-x64-gnu@2.9.6': resolution: {integrity: sha512-vY0le8ad2KaV1PJr+jCd8fUF9VOjwwQP/uBuTJvhvKTloEwxYA/kAjKK9OpIslGA9m/zcnSo74czI6bBrm2sYA==} engines: {node: '>= 10'} cpu: [x64] os: [linux] + libc: [glibc] '@tauri-apps/cli-linux-x64-musl@2.9.6': resolution: {integrity: sha512-TOEuB8YCFZTWVDzsO2yW0+zGcoMiPPwcUgdnW1ODnmgfwccpnihDRoks+ABT1e3fHb1ol8QQWsHSCovb3o2ENQ==} engines: {node: '>= 10'} cpu: [x64] os: [linux] + libc: [musl] '@tauri-apps/cli-win32-arm64-msvc@2.9.6': resolution: {integrity: sha512-ujmDGMRc4qRLAnj8nNG26Rlz9klJ0I0jmZs2BPpmNNf0gM/rcVHhqbEkAaHPTBVIrtUdf7bGvQAD2pyIiUrBHQ==} @@ -1553,6 +1575,9 @@ packages: '@tauri-apps/plugin-shell@2.3.3': resolution: {integrity: sha512-Xod+pRcFxmOWFWEnqH5yZcA7qwAMuaaDkMR1Sply+F8VfBj++CGnj2xf5UoialmjZ2Cvd8qrvSCbU+7GgNVsKQ==} + '@tonejs/midi@2.0.28': + resolution: {integrity: sha512-RII6YpInPsOZ5t3Si/20QKpNqB1lZ2OCFJSOzJxz38YdY/3zqDr3uaml4JuCWkdixuPqP1/TBnXzhQ39csyoVg==} + '@tootallnate/once@2.0.0': resolution: {integrity: sha512-XCuKFP5PS55gnMVu3dty8KPatLqUoy/ZYzDzAGCQ8JNFCkLXzmI7vNHCR+XpbZaMWQK/vQubr7PkYq8g470J/A==} engines: {node: '>= 10'} @@ -1833,6 +1858,9 @@ packages: resolution: {integrity: sha512-ik3ZgC9dY/lYVVM++OISsaYDeg1tb0VtP5uL3ouh1koGOaUMDPpbFIei4JkFimWUFPn90sbMNMXQAIVOlnYKJA==} engines: {node: '>=10'} + array-flatten@3.0.0: + resolution: {integrity: sha512-zPMVc3ZYlGLNk4mpK1NzP2wg0ml9t7fUgDsayR5Y5rSzxQilzR9FGu/EH2jQOcKSAeAfWeylyW8juy3OkWRvNA==} + assertion-error@2.0.1: resolution: {integrity: sha512-Izi8RQcffqCeNVgFigKli1ssklIbpHnCYc6AknXGYoB6grJqyeby7jv12JUQgmTAnIDnbck1uxksT4dzN3PWBA==} engines: {node: '>=12'} @@ -1840,6 +1868,10 @@ packages: asynckit@0.4.0: resolution: {integrity: sha512-Oei9OH4tRh0YqU3GxhX79dM/mwVgvbZJaSNaRk+bshkj0S5cfHcgYakreBjrHwatXKbz+IoIdYLxrKim2MjW0Q==} + automation-events@7.1.14: + resolution: {integrity: sha512-33doW0iTYXR2gSNBEKosQfcUZw1j3PCo3le41Wp3LRKStYXjWejTSjkd38Tm6b5AF0k+0IHmgDO0hfVROyDoUQ==} + engines: {node: '>=18.2.0'} + autoprefixer@10.4.23: resolution: {integrity: sha512-YYTXSFulfwytnjAPlw8QHncHJmlvFKtczb8InXaAx9Q0LbfDnfEYDE55omerIJKihhmU61Ft+cAOSzQVaBUmeA==} engines: {node: ^10 || ^12 || >=14} @@ -2993,6 +3025,9 @@ packages: resolution: {integrity: sha512-PXwfBhYu0hBCPw8Dn0E+WDYb7af3dSLVWKi3HGv84IdF4TyFoC0ysxFd0Goxw7nSv4T/PzEJQxsYsEiFCKo2BA==} engines: {node: '>=8.6'} + midi-file@1.2.4: + resolution: {integrity: sha512-B5SnBC6i2bwJIXTY9MElIydJwAmnKx+r5eJ1jknTLetzLflEl0GWveuBB6ACrQpecSRkOB6fhTx1PwXk2BVxnA==} + mime-db@1.52.0: resolution: {integrity: sha512-sPU4uV7dYlvtWJxwwxHD0PuihVNiE7TyAbQ5SWxDCB9mUYvOgroQOwYQQOKPJ8CIbE+1ETVlOoK1UC2nU3gYvg==} engines: {node: '>= 0.6'} @@ -3499,6 +3534,9 @@ packages: stackback@0.0.2: resolution: {integrity: sha512-1XMJE5fQo1jGH6Y/7ebnwPOBEkIEnT4QF32d5R1+VXdXveM0IBMJt8zfaxX1P3QhVwrYe+576+jkANtSS2mBbw==} + standardized-audio-context@25.3.77: + resolution: {integrity: sha512-Ki9zNz6pKcC5Pi+QPjPyVsD9GwJIJWgryji0XL9cAJXMGyn+dPOf6Qik1AHei0+UNVcc4BOCa0hWLBzlwqsW/A==} + std-env@3.10.0: resolution: {integrity: sha512-5GS12FdOZNliM5mAOxFRg7Ir0pWz8MdpYm6AY6VPkGpbA7ZzmbzNcBJQ0GPvvyWgcY7QAhCgf9Uy89I03faLkg==} @@ -3603,6 +3641,9 @@ packages: resolution: {integrity: sha512-65P7iz6X5yEr1cwcgvQxbbIw7Uk3gOy5dIdtZ4rDveLqhrdJP+Li/Hx6tyK0NEb+2GCyneCMJiGqrADCSNk8sQ==} engines: {node: '>=8.0'} + tone@15.1.22: + resolution: {integrity: sha512-TCScAGD4sLsama5DjvTUXlLDXSqPealhL64nsdV1hhr6frPWve0DeSo63AKnSJwgfg55fhvxj0iPPRwPN5o0ag==} + tough-cookie@4.1.4: resolution: {integrity: sha512-Loo5UUvLD9ScZ6jh8beX1T6sO1w2/MpCRpEP7V280GKMVUQ0Jzar2U3UJPsrdbziLEMMhu3Ujnq//rhiFuIeag==} engines: {node: '>=6'} @@ -5111,6 +5152,11 @@ snapshots: dependencies: '@tauri-apps/api': 2.9.1 + '@tonejs/midi@2.0.28': + dependencies: + array-flatten: 3.0.0 + midi-file: 1.2.4 + '@tootallnate/once@2.0.0': optional: true @@ -5439,11 +5485,18 @@ snapshots: dependencies: tslib: 2.8.1 + array-flatten@3.0.0: {} + assertion-error@2.0.1: {} asynckit@0.4.0: optional: true + automation-events@7.1.14: + dependencies: + '@babel/runtime': 7.28.4 + tslib: 2.8.1 + autoprefixer@10.4.23(postcss@8.5.6): dependencies: browserslist: 4.28.1 @@ -7006,6 +7059,8 @@ snapshots: braces: 3.0.3 picomatch: 2.3.1 + midi-file@1.2.4: {} + mime-db@1.52.0: optional: true @@ -7543,6 +7598,12 @@ snapshots: stackback@0.0.2: {} + standardized-audio-context@25.3.77: + dependencies: + '@babel/runtime': 7.28.4 + automation-events: 7.1.14 + tslib: 2.8.1 + std-env@3.10.0: {} string-width@4.2.3: @@ -7684,6 +7745,11 @@ snapshots: dependencies: is-number: 7.0.0 + tone@15.1.22: + dependencies: + standardized-audio-context: 25.3.77 + tslib: 2.8.1 + tough-cookie@4.1.4: dependencies: psl: 1.15.0 diff --git a/src-tauri/src/commands/api_key_provider_cmd.rs b/src-tauri/src/commands/api_key_provider_cmd.rs index 7dc76c124..ced6faba7 100644 --- a/src-tauri/src/commands/api_key_provider_cmd.rs +++ b/src-tauri/src/commands/api_key_provider_cmd.rs @@ -45,6 +45,8 @@ pub struct UpdateProviderRequest { pub project: Option, pub location: Option, pub region: Option, + /// 自定义模型列表 + pub custom_models: Option>, } /// 添加 API Key 请求 @@ -71,6 +73,8 @@ pub struct ProviderDisplay { pub project: Option, pub location: Option, pub region: Option, + /// 自定义模型列表 + pub custom_models: Vec, pub api_key_count: usize, pub created_at: String, pub updated_at: String, @@ -130,6 +134,7 @@ fn provider_to_display(provider: &ApiKeyProvider, api_key_count: usize) -> Provi project: provider.project.clone(), location: provider.location.clone(), region: provider.region.clone(), + custom_models: provider.custom_models.clone(), api_key_count, created_at: provider.created_at.to_rfc3339(), updated_at: provider.updated_at.to_rfc3339(), @@ -247,6 +252,7 @@ pub fn update_api_key_provider( request.project, request.location, request.region, + request.custom_models, )?; // 获取 API Key 数量 diff --git a/src-tauri/src/database/dao/api_key_provider.rs b/src-tauri/src/database/dao/api_key_provider.rs index 2b673fddc..c2d476db0 100644 --- a/src-tauri/src/database/dao/api_key_provider.rs +++ b/src-tauri/src/database/dao/api_key_provider.rs @@ -126,6 +126,10 @@ pub struct ApiKeyProvider { pub project: Option, pub location: Option, pub region: Option, + /// 自定义模型列表(JSON 数组格式存储) + /// 用于不支持 /models 接口的 Provider(如智谱) + #[serde(default)] + pub custom_models: Vec, pub created_at: DateTime, pub updated_at: DateTime, } @@ -166,7 +170,7 @@ impl ApiKeyProviderDao { pub fn get_all_providers(conn: &Connection) -> Result, rusqlite::Error> { let mut stmt = conn.prepare( "SELECT id, name, type, api_host, is_system, group_name, enabled, sort_order, - api_version, project, location, region, created_at, updated_at + api_version, project, location, region, custom_models, created_at, updated_at FROM api_key_providers ORDER BY sort_order ASC, created_at ASC", )?; @@ -186,7 +190,7 @@ impl ApiKeyProviderDao { ) -> Result, rusqlite::Error> { let mut stmt = conn.prepare( "SELECT id, name, type, api_host, is_system, group_name, enabled, sort_order, - api_version, project, location, region, created_at, updated_at + api_version, project, location, region, custom_models, created_at, updated_at FROM api_key_providers WHERE id = ?1", )?; @@ -206,7 +210,7 @@ impl ApiKeyProviderDao { ) -> Result, rusqlite::Error> { let mut stmt = conn.prepare( "SELECT id, name, type, api_host, is_system, group_name, enabled, sort_order, - api_version, project, location, region, created_at, updated_at + api_version, project, location, region, custom_models, created_at, updated_at FROM api_key_providers WHERE group_name = ?1 ORDER BY sort_order ASC, created_at ASC", @@ -225,11 +229,17 @@ impl ApiKeyProviderDao { conn: &Connection, provider: &ApiKeyProvider, ) -> Result<(), rusqlite::Error> { + let custom_models_json = if provider.custom_models.is_empty() { + None + } else { + Some(serde_json::to_string(&provider.custom_models).unwrap_or_default()) + }; + conn.execute( "INSERT INTO api_key_providers (id, name, type, api_host, is_system, group_name, enabled, sort_order, - api_version, project, location, region, created_at, updated_at) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)", + api_version, project, location, region, custom_models, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15)", params![ provider.id, provider.name, @@ -243,6 +253,7 @@ impl ApiKeyProviderDao { provider.project, provider.location, provider.region, + custom_models_json, provider.created_at.to_rfc3339(), provider.updated_at.to_rfc3339(), ], @@ -255,11 +266,17 @@ impl ApiKeyProviderDao { conn: &Connection, provider: &ApiKeyProvider, ) -> Result<(), rusqlite::Error> { + let custom_models_json = if provider.custom_models.is_empty() { + None + } else { + Some(serde_json::to_string(&provider.custom_models).unwrap_or_default()) + }; + conn.execute( "UPDATE api_key_providers SET name = ?2, type = ?3, api_host = ?4, is_system = ?5, group_name = ?6, enabled = ?7, sort_order = ?8, api_version = ?9, project = ?10, - location = ?11, region = ?12, updated_at = ?13 + location = ?11, region = ?12, custom_models = ?13, updated_at = ?14 WHERE id = ?1", params![ provider.id, @@ -274,6 +291,7 @@ impl ApiKeyProviderDao { provider.project, provider.location, provider.region, + custom_models_json, provider.updated_at.to_rfc3339(), ], )?; @@ -311,8 +329,9 @@ impl ApiKeyProviderDao { let project: Option = row.get(9)?; let location: Option = row.get(10)?; let region: Option = row.get(11)?; - let created_at_str: String = row.get(12)?; - let updated_at_str: String = row.get(13)?; + let custom_models_json: Option = row.get(12)?; + let created_at_str: String = row.get(13)?; + let updated_at_str: String = row.get(14)?; let provider_type: ApiProviderType = type_str.parse().unwrap_or(ApiProviderType::Openai); let group: ProviderGroup = group_str.parse().unwrap_or(ProviderGroup::Custom); @@ -324,6 +343,11 @@ impl ApiKeyProviderDao { .map(|dt| dt.with_timezone(&Utc)) .unwrap_or_else(|_| Utc::now()); + // 解析自定义模型列表 + let custom_models: Vec = custom_models_json + .and_then(|json| serde_json::from_str(&json).ok()) + .unwrap_or_default(); + Ok(ApiKeyProvider { id, name, @@ -337,6 +361,7 @@ impl ApiKeyProviderDao { project, location, region, + custom_models, created_at, updated_at, }) @@ -398,7 +423,7 @@ impl ApiKeyProviderDao { k.usage_count, k.error_count, k.last_used_at, k.created_at, p.id, p.name, p.type, p.api_host, p.is_system, p.group_name, p.enabled, p.sort_order, p.api_version, p.project, p.location, p.region, - p.created_at, p.updated_at + p.custom_models, p.created_at, p.updated_at FROM api_keys k JOIN api_key_providers p ON k.provider_id = p.id WHERE p.type = ?1 AND k.enabled = 1 AND p.enabled = 1 @@ -431,8 +456,9 @@ impl ApiKeyProviderDao { }; // 解析 Provider - let provider_created_at_str: String = row.get(21)?; - let provider_updated_at_str: String = row.get(22)?; + let custom_models_json: Option = row.get(21)?; + let provider_created_at_str: String = row.get(22)?; + let provider_updated_at_str: String = row.get(23)?; let provider_created_at = DateTime::parse_from_rfc3339(&provider_created_at_str) .map(|dt| dt.with_timezone(&Utc)) .unwrap_or_else(|_| Utc::now()); @@ -440,6 +466,11 @@ impl ApiKeyProviderDao { .map(|dt| dt.with_timezone(&Utc)) .unwrap_or_else(|_| Utc::now()); + // 解析自定义模型列表 + let custom_models: Vec = custom_models_json + .and_then(|json| serde_json::from_str(&json).ok()) + .unwrap_or_default(); + let provider = ApiKeyProvider { id: row.get(9)?, name: row.get(10)?, @@ -459,6 +490,7 @@ impl ApiKeyProviderDao { project: row.get(18)?, location: row.get(19)?, region: row.get(20)?, + custom_models, created_at: provider_created_at, updated_at: provider_updated_at, }; @@ -666,7 +698,7 @@ impl ApiKeyProviderDao { ) -> Result, rusqlite::Error> { let mut stmt = conn.prepare( "SELECT id, name, type, api_host, is_system, group_name, enabled, sort_order, - api_version, project, location, region, created_at, updated_at + api_version, project, location, region, custom_models, created_at, updated_at FROM api_key_providers WHERE enabled = 1 ORDER BY sort_order ASC, created_at ASC", diff --git a/src-tauri/src/database/schema.rs b/src-tauri/src/database/schema.rs index 4a6a0b216..2b4ca1f96 100644 --- a/src-tauri/src/database/schema.rs +++ b/src-tauri/src/database/schema.rs @@ -17,12 +17,19 @@ pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> { project TEXT, location TEXT, region TEXT, + custom_models TEXT, created_at TEXT NOT NULL, updated_at TEXT NOT NULL )", [], )?; + // Migration: 添加 custom_models 列(如果不存在) + let _ = conn.execute( + "ALTER TABLE api_key_providers ADD COLUMN custom_models TEXT", + [], + ); + // 创建 api_key_providers 索引 conn.execute( "CREATE INDEX IF NOT EXISTS idx_api_key_providers_group ON api_key_providers(group_name)", diff --git a/src-tauri/src/database/system_providers.rs b/src-tauri/src/database/system_providers.rs index 3606a82d5..646b89233 100644 --- a/src-tauri/src/database/system_providers.rs +++ b/src-tauri/src/database/system_providers.rs @@ -626,6 +626,7 @@ pub fn to_api_key_provider(def: &SystemProviderDef) -> ApiKeyProvider { project: None, location: None, region: None, + custom_models: Vec::new(), created_at: now, updated_at: now, } diff --git a/src-tauri/src/providers/README.md b/src-tauri/src/providers/README.md index 496ffae21..c0a78a2a2 100644 --- a/src-tauri/src/providers/README.md +++ b/src-tauri/src/providers/README.md @@ -1,11 +1,3 @@ - # providers diff --git a/src-tauri/src/providers/openai_custom.rs b/src-tauri/src/providers/openai_custom.rs index 461bc51e4..99b8f4e54 100644 --- a/src-tauri/src/providers/openai_custom.rs +++ b/src-tauri/src/providers/openai_custom.rs @@ -54,16 +54,34 @@ impl OpenAICustomProvider { } /// 构建完整的 API URL - /// 智能处理用户输入的 base_url,无论是否带 /v1 都能正确工作 + /// 智能处理用户输入的 base_url,支持多种 API 版本格式 + /// + /// 支持的格式: + /// - `https://api.openai.com` -> `https://api.openai.com/v1/chat/completions` + /// - `https://api.openai.com/v1` -> `https://api.openai.com/v1/chat/completions` + /// - `https://open.bigmodel.cn/api/paas/v4` -> `https://open.bigmodel.cn/api/paas/v4/chat/completions` + /// - `https://api.deepseek.com/v1` -> `https://api.deepseek.com/v1/chat/completions` fn build_url(&self, endpoint: &str) -> String { let base = self.get_base_url(); let base = base.trim_end_matches('/'); - // 如果用户输入了带 /v1 的 URL,直接拼接 endpoint - // 否则拼接 /v1/endpoint - if base.ends_with("/v1") { + // 检查是否已经包含版本号路径(/v1, /v2, /v3, /v4 等) + // 使用正则匹配 /v 后跟数字的模式 + let has_version = base + .rsplit('/') + .next() + .map(|last_segment| { + last_segment.starts_with('v') + && last_segment.len() >= 2 + && last_segment[1..].chars().all(|c| c.is_ascii_digit()) + }) + .unwrap_or(false); + + if has_version { + // 已有版本号,直接拼接 endpoint format!("{}/{}", base, endpoint) } else { + // 没有版本号,添加 /v1 format!("{}/v1/{}", base, endpoint) } } @@ -104,6 +122,9 @@ impl OpenAICustomProvider { .ok_or("OpenAI API key not configured")?; let url = self.build_url("chat/completions"); + + eprintln!("[OPENAI_CUSTOM] chat_completions URL: {}", url); + eprintln!("[OPENAI_CUSTOM] chat_completions base_url: {}", self.get_base_url()); let resp = self .client @@ -125,6 +146,8 @@ impl OpenAICustomProvider { .ok_or("OpenAI API key not configured")?; let url = self.build_url("models"); + + eprintln!("[OPENAI_CUSTOM] list_models URL: {}", url); let resp = self .client @@ -136,6 +159,7 @@ impl OpenAICustomProvider { if !resp.status().is_success() { let status = resp.status(); let body = resp.text().await.unwrap_or_default(); + eprintln!("[OPENAI_CUSTOM] list_models 失败: {} - {}", status, body); return Err(format!("Failed to list models: {status} - {body}").into()); } diff --git a/src-tauri/src/services/api_key_provider_service.rs b/src-tauri/src/services/api_key_provider_service.rs index d67834b71..5465d4580 100644 --- a/src-tauri/src/services/api_key_provider_service.rs +++ b/src-tauri/src/services/api_key_provider_service.rs @@ -242,6 +242,7 @@ impl ApiKeyProviderService { project, location, region, + custom_models: Vec::new(), created_at: now, updated_at: now, }; @@ -265,6 +266,7 @@ impl ApiKeyProviderService { project: Option, location: Option, region: Option, + custom_models: Option>, ) -> Result { let conn = db.lock().map_err(|e| e.to_string())?; let mut provider = ApiKeyProviderDao::get_provider_by_id(&conn, id) @@ -296,6 +298,9 @@ impl ApiKeyProviderService { if let Some(r) = region { provider.region = if r.is_empty() { None } else { Some(r) }; } + if let Some(models) = custom_models { + provider.custom_models = models; + } provider.updated_at = Utc::now(); ApiKeyProviderDao::update_provider(&conn, &provider).map_err(|e| e.to_string())?; diff --git a/src/components/api-server/ApiServerPage.tsx b/src/components/api-server/ApiServerPage.tsx index e351f2bf4..d40242945 100644 --- a/src/components/api-server/ApiServerPage.tsx +++ b/src/components/api-server/ApiServerPage.tsx @@ -527,8 +527,27 @@ export function ApiServerPage() { const serverUrl = getTestUrl(currentHost, currentPort); const apiKey = config?.server.api_key ?? ""; + // 获取当前选中 Provider 的自定义模型列表 + const getCurrentProviderCustomModels = (): string[] => { + // 先从 API Key Provider 中查找 + const apiKeyProvider = apiKeyProviders.find( + (p) => p.id === defaultProvider && p.enabled + ); + if (apiKeyProvider?.custom_models && apiKeyProvider.custom_models.length > 0) { + return apiKeyProvider.custom_models; + } + return []; + }; + // 根据 Provider 类型获取测试模型 const getTestModel = (provider: string): string => { + // 优先使用自定义模型列表中的第一个模型 + const customModels = getCurrentProviderCustomModels(); + if (customModels.length > 0) { + return customModels[0]; + } + + // 否则使用默认模型 switch (provider) { case "antigravity": return "gemini-3-pro-preview"; @@ -542,6 +561,8 @@ export function ApiServerPage() { return "claude-sonnet-4-20250514"; case "deepseek": return "deepseek-chat"; + case "zhipu": + return "glm-4"; case "kiro": default: return "claude-opus-4-5-20251101"; @@ -549,6 +570,7 @@ export function ApiServerPage() { }; const testModel = getTestModel(defaultProvider); + const customModels = getCurrentProviderCustomModels(); // 根据 Provider 类型获取 Gemini 测试模型列表 const getGeminiTestModels = (provider: string): string[] => { @@ -593,7 +615,7 @@ export function ApiServerPage() { }, { id: "chat", - name: "OpenAI Chat", + name: `OpenAI Chat (${testModel})`, method: "POST", path: "/v1/chat/completions", needsAuth: true, @@ -602,9 +624,23 @@ export function ApiServerPage() { messages: [{ role: "user", content: "Say hi in one word" }], }), }, + // 为自定义模型列表中的其他模型生成测试端点 + ...(customModels.length > 1 + ? customModels.slice(1).map((model, index) => ({ + id: `custom-model-${index}`, + name: `OpenAI Chat (${model})`, + method: "POST", + path: "/v1/chat/completions", + needsAuth: true, + body: JSON.stringify({ + model: model, + messages: [{ role: "user", content: "Say hi in one word" }], + }), + })) + : []), { id: "anthropic", - name: "Anthropic Messages", + name: `Anthropic Messages (${testModel})`, method: "POST", path: "/v1/messages", needsAuth: true, diff --git a/src/components/provider-pool/api-key/ProviderConfigForm.tsx b/src/components/provider-pool/api-key/ProviderConfigForm.tsx index 22c46b4fa..44a96de65 100644 --- a/src/components/provider-pool/api-key/ProviderConfigForm.tsx +++ b/src/components/provider-pool/api-key/ProviderConfigForm.tsx @@ -86,6 +86,7 @@ interface FormState { project: string; location: string; region: string; + customModels: string; } // ============================================================================ @@ -125,6 +126,7 @@ export const ProviderConfigForm: React.FC = ({ project: provider.project || "", location: provider.location || "", region: provider.region || "", + customModels: (provider.custom_models || []).join(", "), }); // 保存状态 @@ -143,6 +145,7 @@ export const ProviderConfigForm: React.FC = ({ project: provider.project || "", location: provider.location || "", region: provider.region || "", + customModels: (provider.custom_models || []).join(", "), }); setSaveError(null); }, [ @@ -152,6 +155,7 @@ export const ProviderConfigForm: React.FC = ({ provider.project, provider.location, provider.region, + provider.custom_models, ]); // 保存配置 @@ -163,12 +167,19 @@ export const ProviderConfigForm: React.FC = ({ setSaveError(null); try { + // 解析自定义模型列表(逗号分隔) + const customModels = state.customModels + .split(",") + .map((m) => m.trim()) + .filter((m) => m.length > 0); + const request: UpdateProviderRequest = { api_host: state.apiHost || undefined, api_version: state.apiVersion || undefined, project: state.project || undefined, location: state.location || undefined, region: state.region || undefined, + custom_models: customModels.length > 0 ? customModels : undefined, }; await onUpdate(provider.id, request); @@ -330,6 +341,25 @@ export const ProviderConfigForm: React.FC = ({ )} + {/* 自定义模型列表 */} +
+ + handleFieldChange("customModels", e.target.value)} + placeholder="glm-4, glm-4-flash, glm-4.7" + disabled={loading || isSaving} + data-testid="custom-models-input" + /> +

+ 该 Provider 支持的模型列表,用逗号分隔。用于不支持 /models 接口的 Provider(如智谱) +

+
+ {/* 保存状态指示 */}
{isSaving ? ( diff --git a/src/lib/api/apiKeyProvider.ts b/src/lib/api/apiKeyProvider.ts index 3cea6a91f..7622f0d80 100644 --- a/src/lib/api/apiKeyProvider.ts +++ b/src/lib/api/apiKeyProvider.ts @@ -38,6 +38,8 @@ export interface UpdateProviderRequest { project?: string; location?: string; region?: string; + /** 自定义模型列表 */ + custom_models?: string[]; } /** @@ -69,6 +71,8 @@ export interface ProviderDisplay { project?: string; location?: string; region?: string; + /** 自定义模型列表 */ + custom_models?: string[]; api_key_count: number; created_at: string; updated_at: string;