diff --git a/package.json b/package.json index 7e0786dd2..196f9a248 100644 --- a/package.json +++ b/package.json @@ -9,7 +9,7 @@ }, "homepage": "https://github.com/aiclientproxy/proxycast", "scripts": { - "dev": "vite", + "dev": "npx vite", "build": "tsc && vite build", "preview": "vite preview", "tauri": "tauri", @@ -105,7 +105,7 @@ "tailwindcss": "^3.4.14", "tsx": "^4.21.0", "typescript": "^5.6.3", - "vite": "^5.4.10", + "vite": "^5.4.21", "vite-plugin-svgr": "^4.5.0", "vitest": "^4.0.16" } diff --git a/src-tauri/src/app/commands/server.rs b/src-tauri/src/app/commands/server.rs index ea7779c0d..95ea6f19b 100644 --- a/src-tauri/src/app/commands/server.rs +++ b/src-tauri/src/app/commands/server.rs @@ -28,15 +28,12 @@ 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 {}:{}", - status.host, status.port - ), + &format!("Server started on {}:{}", 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 81ea2e0f6..d8fc1fa82 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -635,6 +635,7 @@ pub fn run() { app_commands::set_default_provider, app_commands::get_endpoint_providers, app_commands::set_endpoint_provider, + app_commands::update_provider_env_vars, // Unified OAuth commands (new) commands::oauth_cmd::get_oauth_credentials, commands::oauth_cmd::reload_oauth_credentials, diff --git a/src-tauri/src/commands/model_cmd.rs b/src-tauri/src/commands/model_cmd.rs index 0f822264a..541b600dc 100644 --- a/src-tauri/src/commands/model_cmd.rs +++ b/src-tauri/src/commands/model_cmd.rs @@ -1,7 +1,7 @@ //! 模型管理相关命令 -use crate::database::DbConnection; use crate::database::dao::provider_pool::ProviderPoolDao; +use crate::database::DbConnection; use crate::services::model_service::ModelService; use std::collections::HashMap; use tauri::State; @@ -12,8 +12,11 @@ pub fn get_credential_models( db: State<'_, DbConnection>, credential_uuid: String, ) -> Result, String> { - tracing::info!("[GET_CREDENTIAL_MODELS] 获取凭证模型列表: {}", credential_uuid); - + tracing::info!( + "[GET_CREDENTIAL_MODELS] 获取凭证模型列表: {}", + credential_uuid + ); + let model_service = ModelService::new(); model_service.get_credential_models(&db, &credential_uuid) } @@ -25,10 +28,13 @@ pub async fn refresh_credential_models( credential_uuid: String, ) -> Result, String> { tracing::info!("[REFRESH_CREDENTIAL_MODELS] ========== 开始刷新凭证模型列表 =========="); - tracing::info!("[REFRESH_CREDENTIAL_MODELS] credential_uuid: {}", credential_uuid); - + 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())?; @@ -36,29 +42,31 @@ pub async fn refresh_credential_models( .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?; - + 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) } @@ -68,18 +76,16 @@ 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> { +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) } @@ -90,30 +96,36 @@ 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); - + + 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()) { + if let Err(e) = + model_service.update_credential_models(&db, &credential.uuid, models.clone()) + { tracing::error!("[REFRESH_ALL] 更新数据库失败: {}", e); Err(format!("更新数据库失败: {}", e)) } else { @@ -126,21 +138,19 @@ pub async fn refresh_all_credential_models( 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 = +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 9842feeca..8166b535a 100644 --- a/src-tauri/src/commands/route_cmd.rs +++ b/src-tauri/src/commands/route_cmd.rs @@ -10,17 +10,19 @@ use crate::models::route_model::{RouteInfo, RouteListResponse}; 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() + 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()) @@ -31,7 +33,9 @@ fn get_valid_base_url(config: &config::Config) -> String { configured_host.clone() } else { // IP 不在当前网卡列表中,替换为局域网 IP - network_info.all_ips.iter() + 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()) diff --git a/src-tauri/src/converter/openai_to_antigravity.rs b/src-tauri/src/converter/openai_to_antigravity.rs index 6b4d517c1..2225aab26 100644 --- a/src-tauri/src/converter/openai_to_antigravity.rs +++ b/src-tauri/src/converter/openai_to_antigravity.rs @@ -278,10 +278,10 @@ pub fn convert_openai_to_antigravity_with_context( 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); @@ -672,10 +672,13 @@ pub fn convert_openai_to_antigravity_with_context( "model": actual_model, "userAgent": "antigravity" }); - - eprintln!("[CONVERT] 转换后的请求体: {}", serde_json::to_string_pretty(&result).unwrap_or_default()); + + eprintln!( + "[CONVERT] 转换后的请求体: {}", + serde_json::to_string_pretty(&result).unwrap_or_default() + ); eprintln!("========== [CONVERT] OpenAI -> Antigravity 转换完成 =========="); - + result } diff --git a/src-tauri/src/database/dao/api_key_provider.rs b/src-tauri/src/database/dao/api_key_provider.rs index c2d476db0..92263e7c8 100644 --- a/src-tauri/src/database/dao/api_key_provider.rs +++ b/src-tauri/src/database/dao/api_key_provider.rs @@ -234,7 +234,7 @@ impl ApiKeyProviderDao { } 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, @@ -271,7 +271,7 @@ impl ApiKeyProviderDao { } 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, diff --git a/src-tauri/src/providers/antigravity.rs b/src-tauri/src/providers/antigravity.rs index 850b8e254..ebfa56599 100644 --- a/src-tauri/src/providers/antigravity.rs +++ b/src-tauri/src/providers/antigravity.rs @@ -42,7 +42,11 @@ impl AntigravityApiError { } /// 创建带响应体的 API 错误 - pub fn with_body(status_code: u16, message: impl Into, body: impl Into) -> Self { + pub fn with_body( + status_code: u16, + message: impl Into, + body: impl Into, + ) -> Self { Self { status_code, message: message.into(), @@ -840,8 +844,14 @@ impl AntigravityProvider { 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()); + 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 @@ -876,10 +886,11 @@ impl AntigravityProvider { 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)))?; - + + 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) } @@ -902,27 +913,31 @@ impl AntigravityProvider { Err(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 + base_url, + e.status_code ); last_error = Some(e); continue; } - + // 403、401 等权限错误直接返回,不降级 tracing::warn!( "[Antigravity] {} 失败 (HTTP {}): {}", - base_url, e.status_code, e.message + base_url, + e.status_code, + e.message ); return Err(e); } } } - Err(last_error.unwrap_or_else(|| AntigravityApiError::new(503, "All Antigravity base URLs failed"))) + Err(last_error + .unwrap_or_else(|| AntigravityApiError::new(503, "All Antigravity base URLs failed"))) } /// 发现项目 ID @@ -1032,16 +1047,22 @@ impl AntigravityProvider { ) -> Result { eprintln!("========== [ANTIGRAVITY_GENERATE] 开始生成内容 =========="); eprintln!("[ANTIGRAVITY_GENERATE] 模型: {}", model); - eprintln!("[ANTIGRAVITY_GENERATE] 请求体: {}", serde_json::to_string_pretty(request_body).unwrap_or_default()); - + 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] 构建的 payload: {}", + serde_json::to_string_pretty(&payload).unwrap_or_default() + ); eprintln!("[ANTIGRAVITY_GENERATE] 调用 call_api..."); let resp = self.call_api("generateContent", &payload).await?; diff --git a/src-tauri/src/providers/codex.rs b/src-tauri/src/providers/codex.rs index ae80db018..f858eb136 100644 --- a/src-tauri/src/providers/codex.rs +++ b/src-tauri/src/providers/codex.rs @@ -597,7 +597,8 @@ impl CodexProvider { ("client_id", OPENAI_CLIENT_ID), ("response_type", "code"), ("redirect_uri", &self.get_redirect_uri()), - ("scope", "openid email profile offline_access"), + // 必须包含 api.responses.write 才能使用 responses API + ("scope", "openid email profile offline_access api.responses.write api.responses.read api.model.request"), ("state", state), ("code_challenge", &pkce_codes.code_challenge), ("code_challenge_method", "S256"), @@ -750,7 +751,8 @@ impl CodexProvider { ("client_id", OPENAI_CLIENT_ID), ("grant_type", "refresh_token"), ("refresh_token", refresh_token.as_str()), - ("scope", "openid profile email"), + // 必须包含 api.responses.write 才能使用 responses API + ("scope", "openid email profile offline_access api.responses.write api.responses.read api.model.request"), ]; let resp = self @@ -1746,10 +1748,27 @@ mod tests { // OAuth 登录功能(参考 Antigravity 实现) // ============================================================================ +use once_cell::sync::Lazy; use std::sync::Arc; -use tokio::sync::oneshot; +use tokio::sync::{oneshot, RwLock}; use uuid::Uuid; +/// 全局 Codex OAuth 服务器状态 +/// 用于在重新打开授权对话框时关闭之前的服务器 +static CODEX_OAUTH_SERVER_SHUTDOWN: Lazy>>> = + Lazy::new(|| RwLock::new(None)); + +/// 停止之前运行的 Codex OAuth 服务器(如果有) +pub async fn stop_codex_oauth_server() { + let mut guard = CODEX_OAUTH_SERVER_SHUTDOWN.write().await; + if let Some(shutdown_tx) = guard.take() { + tracing::info!("[Codex OAuth] 关闭之前的 OAuth 服务器"); + let _ = shutdown_tx.send(()); + // 给服务器一些时间来关闭 + tokio::time::sleep(std::time::Duration::from_millis(100)).await; + } +} + /// OAuth 登录成功后的凭证信息 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CodexOAuthResult { @@ -1777,7 +1796,8 @@ pub fn generate_codex_auth_url(state: &str, code_challenge: &str) -> String { ("client_id", OPENAI_CLIENT_ID), ("response_type", "code"), ("redirect_uri", redirect_uri.as_str()), - ("scope", "openid email profile offline_access"), + // 必须包含 api.responses.write 才能使用 responses API + ("scope", "openid email profile offline_access api.responses.write api.responses.read api.model.request"), ("state", state), ("code_challenge", code_challenge), ("code_challenge_method", "S256"), @@ -1890,6 +1910,9 @@ pub async fn start_codex_oauth_server_and_get_url() -> Result< use std::collections::HashMap; use tokio::net::TcpListener; + // 首先停止之前可能运行的 OAuth 服务器 + stop_codex_oauth_server().await; + let client = Client::builder() .timeout(std::time::Duration::from_secs(30)) .build()?; @@ -1907,6 +1930,15 @@ pub async fn start_codex_oauth_server_and_get_url() -> Result< let (tx, rx) = oneshot::channel::>(); let tx = Arc::new(tokio::sync::Mutex::new(Some(tx))); + // 创建 shutdown channel + let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>(); + + // 保存 shutdown sender 到全局状态 + { + let mut guard = CODEX_OAUTH_SERVER_SHUTDOWN.write().await; + *guard = Some(shutdown_tx); + } + // 使用固定端口 1455(OpenAI OAuth 要求) let port = OPENAI_OAUTH_CALLBACK_PORT; let listener = TcpListener::bind(format!("127.0.0.1:{}", port)).await.map_err(|e| { @@ -2095,8 +2127,11 @@ pub async fn start_codex_oauth_server_and_get_url() -> Result< }), ); - // 启动服务器 - let server = axum::serve(listener, app); + // 启动服务器(支持优雅关闭) + let server = axum::serve(listener, app).with_graceful_shutdown(async move { + let _ = shutdown_rx.await; + tracing::info!("[Codex OAuth] 服务器收到关闭信号"); + }); // 创建等待 future let wait_future = async move { @@ -2119,8 +2154,18 @@ pub async fn start_codex_oauth_server_and_get_url() -> Result< }); match timeout.await { - Ok(result) => result, - Err(_) => Err("OAuth 登录超时(5分钟)".into()), + Ok(result) => { + // 成功或失败后都清理全局状态 + let mut guard = CODEX_OAUTH_SERVER_SHUTDOWN.write().await; + *guard = None; + result + } + Err(_) => { + // 超时后也清理全局状态 + let mut guard = CODEX_OAUTH_SERVER_SHUTDOWN.write().await; + *guard = None; + Err("OAuth 登录超时(5分钟)".into()) + } } }; diff --git a/src-tauri/src/providers/openai_custom.rs b/src-tauri/src/providers/openai_custom.rs index a6e2794c8..f667008a9 100644 --- a/src-tauri/src/providers/openai_custom.rs +++ b/src-tauri/src/providers/openai_custom.rs @@ -69,7 +69,7 @@ impl OpenAICustomProvider { /// 构建完整的 API URL /// 智能处理用户输入的 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` @@ -136,9 +136,12 @@ 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()); + eprintln!( + "[OPENAI_CUSTOM] chat_completions base_url: {}", + self.get_base_url() + ); let resp = self .client @@ -160,7 +163,7 @@ 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 diff --git a/src-tauri/src/server/handlers/provider_calls.rs b/src-tauri/src/server/handlers/provider_calls.rs index 97c5a39b8..ad258988c 100644 --- a/src-tauri/src/server/handlers/provider_calls.rs +++ b/src-tauri/src/server/handlers/provider_calls.rs @@ -448,7 +448,7 @@ pub async fn call_provider_anthropic( Some(&api_err.message), ); } - + // 直接使用 AntigravityApiError 的状态码构建响应 build_error_response_with_status(api_err.status_code, &api_err.to_string()) } @@ -1545,16 +1545,16 @@ pub async fn call_provider_openai( // 非流式请求处理 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) => { @@ -1566,7 +1566,7 @@ pub async fn call_provider_openai( 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()) } diff --git a/src-tauri/src/server/mod.rs b/src-tauri/src/server/mod.rs index ccf09925b..2a5ad8ea4 100644 --- a/src-tauri/src/server/mod.rs +++ b/src-tauri/src/server/mod.rs @@ -26,7 +26,8 @@ use crate::providers::openai_custom::OpenAICustomProvider; use crate::providers::qwen::QwenProvider; use crate::server_utils::{ build_anthropic_response, build_anthropic_stream_response, build_error_response, - build_error_response_with_status, build_gemini_native_request, health, models, parse_cw_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; @@ -206,7 +207,10 @@ impl ServerState { ServerStatus { running: self.running, // 使用实际运行的 host,如果没有则使用配置的 host - host: self.running_host.clone().unwrap_or_else(|| self.config.server.host.clone()), + 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), @@ -284,9 +288,12 @@ impl ServerState { // 检查配置的 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" { + if configured_host == "0.0.0.0" + || configured_host == "127.0.0.1" + || configured_host == "localhost" + { configured_host.clone() } else { // 检查 IP 是否在当前网卡列表中 @@ -297,18 +304,21 @@ impl ServerState { } else { // IP 不在当前网卡列表中,使用当前的局域网 IP // 优先选择 192.168.x.x 或 10.x.x.x 开头的 IP(真正的局域网 IP) - let preferred_ip = network_info.all_ips.iter() + 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 + configured_host, + new_ip ); eprintln!( "[SERVER] 警告:配置的 IP {} 不在当前网卡列表中,自动切换到 {}", @@ -317,9 +327,7 @@ impl ServerState { new_ip } } - Err(_) => { - configured_host.clone() - } + Err(_) => configured_host.clone(), } } }; @@ -1327,11 +1335,13 @@ async fn gemini_generate_content( Json(resp).into_response() } Err(api_err) => { - state - .logs - .write() - .await - .add("error", &format!("[GEMINI] 请求失败 (HTTP {}): {}", api_err.status_code, api_err.message)); + state.logs.write().await.add( + "error", + &format!( + "[GEMINI] 请求失败 (HTTP {}): {}", + api_err.status_code, api_err.message + ), + ); // 直接使用 AntigravityApiError 的状态码构建响应 build_error_response_with_status(api_err.status_code, &api_err.to_string()) @@ -1356,9 +1366,13 @@ async fn list_routes(State(state): State) -> impl IntoResponse { 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_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 保持不变 @@ -1371,12 +1385,14 @@ async fn list_routes(State(state): State) -> impl IntoResponse { 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() + 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()) diff --git a/src-tauri/src/services/live_sync.rs b/src-tauri/src/services/live_sync.rs index 365c8dec2..1c5da08a2 100644 --- a/src-tauri/src/services/live_sync.rs +++ b/src-tauri/src/services/live_sync.rs @@ -88,7 +88,7 @@ fn get_shell_config_path() -> Result Result<(), Box> { let config_path = get_shell_config_path()?; diff --git a/src-tauri/src/services/model_service.rs b/src-tauri/src/services/model_service.rs index 9a864fc3f..9b9c354c6 100644 --- a/src-tauri/src/services/model_service.rs +++ b/src-tauri/src/services/model_service.rs @@ -147,7 +147,11 @@ impl ModelService { if !status.is_success() { let error_body = response.text().await.unwrap_or_default(); - tracing::error!("[MODEL_SERVICE] OpenAI HTTP 错误: status={}, body={}", status, error_body); + tracing::error!( + "[MODEL_SERVICE] OpenAI HTTP 错误: status={}, body={}", + status, + error_body + ); return Err(format!("HTTP 错误: {}", status)); } @@ -158,13 +162,18 @@ impl ModelService { 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 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) @@ -192,7 +201,10 @@ impl ModelService { base_url.unwrap_or("https://api.anthropic.com") ); - tracing::info!("[MODEL_SERVICE] 请求 Anthropic API 获取模型列表: url={}", url); + tracing::info!( + "[MODEL_SERVICE] 请求 Anthropic API 获取模型列表: url={}", + url + ); let response = self .client @@ -212,7 +224,11 @@ impl ModelService { if !status.is_success() { let error_body = response.text().await.unwrap_or_default(); - tracing::error!("[MODEL_SERVICE] Anthropic HTTP 错误: status={}, body={}", status, error_body); + tracing::error!( + "[MODEL_SERVICE] Anthropic HTTP 错误: status={}, body={}", + status, + error_body + ); return Err(format!("HTTP 错误: {}", status)); } @@ -223,14 +239,22 @@ impl ModelService { 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 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()); + + tracing::info!( + "[MODEL_SERVICE] Anthropic 成功获取 {} 个模型", + model_ids.len() + ); Ok(model_ids) } @@ -265,7 +289,11 @@ impl ModelService { if !status.is_success() { let error_body = response.text().await.unwrap_or_default(); - tracing::error!("[MODEL_SERVICE] Gemini HTTP 错误: status={}, body={}", status, error_body); + tracing::error!( + "[MODEL_SERVICE] Gemini HTTP 错误: status={}, body={}", + status, + error_body + ); return Err(format!("HTTP 错误: {}", status)); } @@ -277,10 +305,15 @@ impl ModelService { 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 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") @@ -348,10 +381,9 @@ impl ModelService { "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(), - ], + PoolProviderType::GeminiApiKey => { + vec!["gemini-2.5-flash".to_string(), "gemini-2.5-pro".to_string()] + } _ => vec![], } } @@ -389,9 +421,7 @@ impl ModelService { .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(); + 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()), @@ -436,10 +466,7 @@ impl ModelService { 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(); + let mut all_models: Vec = models_by_provider.into_values().flatten().collect(); all_models.sort(); all_models.dedup(); diff --git a/src-tauri/src/services/provider_pool_service.rs b/src-tauri/src/services/provider_pool_service.rs index c79576998..97bea58d2 100644 --- a/src-tauri/src/services/provider_pool_service.rs +++ b/src-tauri/src/services/provider_pool_service.rs @@ -1424,65 +1424,75 @@ impl ProviderPoolService { .filter(|s| !s.is_empty()) }); - match base_url { - Some(base) => { - // 使用自定义 base_url (如 Yunyi),与 CodexProvider 的 URL/headers 行为保持一致 - let url = CodexProvider::build_responses_url(base); + // 检查是否使用 API Key 模式(如果有 api_key 且没有 refresh_token/access_token) + let is_api_key_mode = provider + .credentials + .api_key + .as_deref() + .map(|s| !s.trim().is_empty()) + .unwrap_or(false) + && provider.credentials.refresh_token.is_none(); - // Codex/Yunyi 使用 responses API 格式;云驿等代理要求 stream 必须为 true - let request_body = serde_json::json!({ - "model": model, - "input": [{ - "type": "message", - "role": "user", - "content": [{"type": "input_text", "text": "Say OK"}] - }], - "max_output_tokens": 10, - "stream": true - }); + // API Key 模式使用 chat/completions API,OAuth 模式使用 responses API + if is_api_key_mode && base_url.is_none() { + // API Key 直连 OpenAI:使用 chat/completions API + return self.check_openai_health(&token, None, model).await; + } - tracing::debug!( - "[HEALTH_CHECK] Codex responses API URL: {}, model: {}", - url, - model - ); + // OAuth 模式或有自定义 base_url:使用 responses API + let url = match base_url { + Some(base) => CodexProvider::build_responses_url(base), + None => "https://api.openai.com/v1/responses".to_string(), + }; - let response = self - .client - .post(&url) - .bearer_auth(&token) - .header("Content-Type", "application/json") - .header("Accept", "text/event-stream") - .header("Openai-Beta", "responses=experimental") - .header("Originator", "codex_cli_rs") - .header("Session_id", uuid::Uuid::new_v4().to_string()) - .header("Conversation_id", uuid::Uuid::new_v4().to_string()) - .header( - "User-Agent", - "codex_cli_rs/0.77.0 (ProxyCast health check; Mac OS; arm64)", - ) - .json(&request_body) - .timeout(self.health_check_timeout) - .send() - .await - .map_err(|e| format!("请求失败: {}", e))?; + // Codex/Yunyi 使用 responses API 格式;云驿等代理要求 stream 必须为 true + let request_body = serde_json::json!({ + "model": model, + "input": [{ + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "Say OK"}] + }], + "max_output_tokens": 10, + "stream": true + }); - if response.status().is_success() { - Ok(()) - } else { - let status = response.status(); - let body = response.text().await.unwrap_or_default(); - Err(format!( - "HTTP {} - {}", - status, - body.chars().take(200).collect::() - )) - } - } - None => { - // 没有自定义 base_url,使用 OpenAI 官方 chat/completions API - self.check_openai_health(&token, None, model).await - } + tracing::debug!( + "[HEALTH_CHECK] Codex responses API URL: {}, model: {}", + url, + model + ); + + let response = self + .client + .post(&url) + .bearer_auth(&token) + .header("Content-Type", "application/json") + .header("Accept", "text/event-stream") + .header("Openai-Beta", "responses=experimental") + .header("Originator", "codex_cli_rs") + .header("Session_id", uuid::Uuid::new_v4().to_string()) + .header("Conversation_id", uuid::Uuid::new_v4().to_string()) + .header( + "User-Agent", + "codex_cli_rs/0.77.0 (ProxyCast health check; Mac OS; arm64)", + ) + .json(&request_body) + .timeout(self.health_check_timeout) + .send() + .await + .map_err(|e| format!("请求失败: {}", e))?; + + if response.status().is_success() { + Ok(()) + } else { + let status = response.status(); + let body = response.text().await.unwrap_or_default(); + Err(format!( + "HTTP {} - {}", + status, + body.chars().take(200).collect::() + )) } } diff --git a/src-tauri/tests/api_key_provider_tests.rs b/src-tauri/tests/api_key_provider_tests.rs index d54cc8e4d..d27560ce0 100644 --- a/src-tauri/tests/api_key_provider_tests.rs +++ b/src-tauri/tests/api_key_provider_tests.rs @@ -102,6 +102,7 @@ impl TestContext { project: None, location: None, region: None, + custom_models: vec![], created_at: now, updated_at: now, }; @@ -709,6 +710,7 @@ mod unit_tests { None, None, None, + None, ) .expect("Failed to update provider"); @@ -830,6 +832,7 @@ mod unit_tests { project: None, location: None, region: None, + custom_models: vec![], created_at: now, updated_at: now, }; diff --git a/src/components/api-server/ApiServerPage.tsx b/src/components/api-server/ApiServerPage.tsx index d40242945..9bf6bc369 100644 --- a/src/components/api-server/ApiServerPage.tsx +++ b/src/components/api-server/ApiServerPage.tsx @@ -1,4 +1,4 @@ -import { useState, useEffect } from "react"; +import { useState, useEffect, useMemo } from "react"; import { Play, Copy, @@ -24,6 +24,7 @@ import { TestResult, getDefaultProvider, setDefaultProvider, + updateProviderEnvVars, getNetworkInfo, NetworkInfo, } from "@/hooks/useTauri"; @@ -32,6 +33,11 @@ import { apiKeyProviderApi, ProviderWithKeysDisplay, } from "@/lib/api/apiKeyProvider"; +import { + getModelRegistry, + getModelsForProvider, +} from "@/lib/api/modelRegistry"; +import type { EnhancedModelMetadata } from "@/lib/types/modelRegistry"; interface TestState { endpoint: string; @@ -43,6 +49,64 @@ interface TestState { type TabId = "server" | "routes" | "logs"; +// Provider 到 API 类型的映射 +type ApiType = "openai" | "anthropic" | "gemini"; +const getProviderApiType = (provider: string): ApiType => { + // 转换为小写以便匹配 + const p = provider.toLowerCase(); + + // OpenAI 兼容类型 + if ( + p === "codex" || + p === "openai" || + p === "openai-response" || + p === "azure_openai" || + p === "azure-openai" || + p === "qwen" || + p === "iflow" + ) { + return "openai"; + } + + // Anthropic 类型 + if ( + p === "anthropic" || + p === "claude" || + p === "claude_oauth" || + p === "kiro" + ) { + return "anthropic"; + } + + // Gemini 类型 + if ( + p === "gemini" || + p === "gemini_api_key" || + p === "antigravity" || + p === "vertex" || + p === "vertexai" + ) { + return "gemini"; + } + + // 默认返回 openai + return "openai"; +}; + +// 根据 API 类型获取对应的模型 provider_id 列表 +const getModelProviderIds = (apiType: ApiType): string[] => { + switch (apiType) { + case "gemini": + return ["google"]; + case "anthropic": + return ["anthropic"]; + case "openai": + return ["openai", "azure", "deepseek", "alibaba"]; + default: + return []; + } +}; + // 可用的 Provider 信息(合并 OAuth 凭证池和 API Key Provider) interface AvailableProvider { id: string; @@ -78,6 +142,11 @@ export function ApiServerPage() { // 网络信息 const [networkInfo, setNetworkInfo] = useState(null); + // 模型库状态 + const [allModels, setAllModels] = useState([]); + const [testModel, setTestModel] = useState(""); + const [_modelsLoading, setModelsLoading] = useState(false); + // 自动清除消息 useEffect(() => { if (message) { @@ -109,34 +178,6 @@ 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(); @@ -144,32 +185,19 @@ export function ApiServerPage() { loadNetworkInfo(); const statusInterval = setInterval(fetchStatus, 3000); - // 定期刷新网络信息,以便检测 IP 变化 - const networkInterval = setInterval(loadNetworkInfo, 5000); - return () => { - clearInterval(statusInterval); - clearInterval(networkInterval); - }; + return () => clearInterval(statusInterval); }, []); - // 当 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 loadNetworkInfo = async () => { + try { + const info = await getNetworkInfo(); + setNetworkInfo(info); + } catch (e) { + console.error("Failed to get network info:", e); } - }, [config, networkInfo]); const loadDefaultProvider = async () => { + }; + + const loadDefaultProvider = async () => { try { const dp = await getDefaultProvider(); setDefaultProviderState(dp); @@ -178,14 +206,73 @@ export function ApiServerPage() { } }; + // 加载模型库 + const loadModels = async ( + provider?: string, + providers?: ProviderWithKeysDisplay[], + resetSelection: boolean = false, + ) => { + setModelsLoading(true); + try { + let models: EnhancedModelMetadata[]; + + if (provider) { + // 首先检查是否是自定义 API Key Provider + // 如果是,使用其 type 字段来确定 API 类型 + let effectiveProvider = provider; + if (providers) { + const customProvider = providers.find((p) => p.id === provider); + if (customProvider) { + // 使用自定义 Provider 的 type 字段 + effectiveProvider = customProvider.type; + } + } + + // 根据 Provider 的 API 类型过滤模型 + const apiType = getProviderApiType(effectiveProvider); + const providerIds = getModelProviderIds(apiType); + + if (providerIds.length > 0) { + // 获取所有匹配 provider_id 的模型 + const modelPromises = providerIds.map((id) => + getModelsForProvider(id), + ); + const modelArrays = await Promise.all(modelPromises); + models = modelArrays.flat(); + } else { + // 未知类型,显示所有模型 + models = await getModelRegistry(); + } + } else { + models = await getModelRegistry(); + } + + setAllModels(models); + + // 只在需要重置选择时,或当前选择的模型不在新列表中时,才重置测试模型选择 + if (resetSelection || !models.find((m) => m.id === testModel)) { + if (models.length > 0) { + const defaultModel = + models.find((m) => m.tier === "pro") || models[0]; + if (defaultModel) { + setTestModel(defaultModel.id); + } + } else { + setTestModel(""); + } + } + } catch (e) { + console.error("Failed to load models:", e); + } + setModelsLoading(false); + }; + const handleStart = async () => { setLoading(true); setError(null); try { await reloadCredentials(); await startServer(); - // 等待服务器完全启动 - await new Promise((resolve) => setTimeout(resolve, 500)); await fetchStatus(); setMessage({ type: "success", text: "服务已启动" }); } catch (e: unknown) { @@ -298,6 +385,23 @@ export function ApiServerPage() { null, ); + // 当 defaultProvider 变化时,重新加载模型并重置选择 + useEffect(() => { + loadModels(defaultProvider, apiKeyProviders, true); + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [defaultProvider]); + + // 当 apiKeyProviders 首次加载时,重新加载模型(不重置选择) + // 这是为了确保自定义 Provider 的 type 能被正确识别 + const [apiKeyProvidersLoaded, setApiKeyProvidersLoaded] = useState(false); + useEffect(() => { + if (apiKeyProviders.length > 0 && !apiKeyProvidersLoaded) { + setApiKeyProvidersLoaded(true); + loadModels(defaultProvider, apiKeyProviders, true); + } + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [apiKeyProviders]); + // 加载凭证池概览 const loadPoolOverview = async () => { try { @@ -451,13 +555,24 @@ export function ApiServerPage() { const handleSetDefaultProvider = async (providerId: string) => { try { - // 先更新 UI 状态,提供即时反馈 - setDefaultProviderState(providerId); - - // 异步调用后端 await setDefaultProvider(providerId); + setDefaultProviderState(providerId); - // 获取该 Provider 的凭证信息(用于显示消息) + // 获取最新的凭证池数据 + const freshOverview = await providerPoolApi.getOverview(); + setPoolOverview(freshOverview); + + // 如果是 API Key Provider,更新对应的环境变量 + const apiKeyProvider = apiKeyProviders.find((p) => p.id === providerId); + if (apiKeyProvider && apiKeyProvider.api_host) { + // 根据 provider.type 更新对应的环境变量 + await updateProviderEnvVars( + apiKeyProvider.type, + apiKeyProvider.api_host, + ); + } + + // 获取该 Provider 的凭证信息 const provider = availableProviders.find((p) => p.id === providerId); const label = providerLabels[providerId] || providerId; @@ -476,207 +591,112 @@ 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: 使用当前局域网 IP(优先 192.168.x.x 或 10.x.x.x) + // - 0.0.0.0: 使用 127.0.0.1(本机访问所有接口) // - 局域网 IP: 使用该 IP(允许局域网测试) const getTestUrl = (host: string, port: number) => { if (host === "0.0.0.0") { - // 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://127.0.0.1:${port}`; } return `http://${host}:${port}`; }; // 使用 editHost 而不是 status.host,这样可以实时反映用户的选择 - // 同时检查配置的 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 currentHost = status?.running ? status.host : editHost; const currentPort = status?.running ? status.port : parseInt(editPort) || 8999; 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 + // 动态生成测试端点 + const testEndpoints = useMemo(() => { + if (!testModel) return []; + + // 首先检查是否是自定义 API Key Provider + // 如果是,使用其 type 字段来确定 API 类型 + let effectiveProvider = defaultProvider; + const customProvider = apiKeyProviders.find( + (p) => p.id === defaultProvider, ); - 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]; + if (customProvider) { + effectiveProvider = customProvider.type; } - // 否则使用默认模型 - switch (provider) { - case "antigravity": - return "gemini-3-pro-preview"; - case "gemini": - return "gemini-2.0-flash"; - case "qwen": - return "qwen-max"; + const apiType = getProviderApiType(effectiveProvider); + + switch (apiType) { case "openai": - return "gpt-4o"; - case "claude": - return "claude-sonnet-4-20250514"; - case "deepseek": - return "deepseek-chat"; - case "zhipu": - return "glm-4"; - case "kiro": - default: - return "claude-opus-4-5-20251101"; - } - }; - - const testModel = getTestModel(defaultProvider); - const customModels = getCurrentProviderCustomModels(); - - // 根据 Provider 类型获取 Gemini 测试模型列表 - const getGeminiTestModels = (provider: string): string[] => { - switch (provider) { - case "antigravity": return [ - "gemini-3-pro-preview", - "gemini-3-pro-image-preview", - "gemini-3-flash-preview", - "gemini-claude-sonnet-4-5", - ]; - case "gemini": - return ["gemini-2.0-flash", "gemini-2.5-flash", "gemini-2.5-pro"]; - default: - return ["gemini-2.0-flash"]; - } - }; - - const geminiTestModels = getGeminiTestModels(defaultProvider); - - // 是否显示 Gemini 测试端点 - const showGeminiTest = - defaultProvider === "antigravity" || defaultProvider === "gemini"; - - // Test endpoints - const testEndpoints = [ - { - id: "health", - name: "健康检查", - method: "GET", - path: "/health", - needsAuth: false, - body: null, - }, - { - id: "models", - name: "模型列表", - method: "GET", - path: "/v1/models", - needsAuth: true, - body: null, - }, - { - id: "chat", - name: `OpenAI Chat (${testModel})`, - method: "POST", - path: "/v1/chat/completions", - needsAuth: true, - body: JSON.stringify({ - model: testModel, - 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 (${testModel})`, - method: "POST", - path: "/v1/messages", - needsAuth: true, - body: JSON.stringify({ - model: testModel, - max_tokens: 100, - messages: [ { - role: "user", - content: "What is 1+1? Answer with just the number.", + id: "chat", + name: "OpenAI Chat", + method: "POST", + path: "/v1/chat/completions", + needsAuth: true, + body: JSON.stringify({ + model: testModel, + messages: [{ role: "user", content: "Say hi in one word" }], + }), }, - ], - }), - }, - // Gemini 原生协议测试(仅在 Antigravity 或 Gemini Provider 时显示) - ...(showGeminiTest - ? geminiTestModels.map((model, index) => ({ - id: `gemini-${index}`, - name: `Gemini ${model}`, - method: "POST", - path: `/v1/gemini/${model}:generateContent`, - needsAuth: true, - body: JSON.stringify({ - contents: [ - { - role: "user", - parts: [{ text: "What is 2+2? Answer with just the number." }], + ]; + + case "anthropic": + return [ + { + id: "anthropic", + name: "Anthropic Messages", + method: "POST", + path: "/v1/messages", + needsAuth: true, + body: JSON.stringify({ + model: testModel, + max_tokens: 100, + messages: [ + { + role: "user", + content: "What is 1+1? Answer with just the number.", + }, + ], + }), + }, + ]; + + case "gemini": + return [ + { + id: "gemini", + name: `Gemini ${testModel}`, + method: "POST", + path: `/v1/gemini/${testModel}:generateContent`, + needsAuth: true, + body: JSON.stringify({ + contents: [ + { + role: "user", + parts: [ + { text: "What is 2+2? Answer with just the number." }, + ], + }, + ], + generationConfig: { + maxOutputTokens: 100, }, - ], - generationConfig: { - maxOutputTokens: 100, - }, - }), - })) - : []), - ]; + }), + }, + ]; + + default: + return []; + } + }, [defaultProvider, testModel, apiKeyProviders]); const runTest = async (endpoint: (typeof testEndpoints)[0]) => { setTestResults((prev) => ({ @@ -703,7 +723,10 @@ export function ApiServerPage() { }, })); - return result.success; + // 测试成功后立即刷新凭证池数据,更新使用次数 + if (result.success) { + await loadPoolOverview(); + } } catch (e: unknown) { const errMsg = e instanceof Error ? e.message : String(e); setTestResults((prev) => ({ @@ -714,21 +737,12 @@ export function ApiServerPage() { response: `请求失败: ${errMsg}`, }, })); - return false; } }; const runAllTests = async () => { - let hasSuccess = false; for (const endpoint of testEndpoints) { - const success = await runTest(endpoint); - if (success) { - hasSuccess = true; - } - } - // 所有测试完成后,如果有成功的测试,刷新一次凭证池数据 - if (hasSuccess) { - await loadPoolOverview(); + await runTest(endpoint); } }; @@ -921,30 +935,6 @@ export function ApiServerPage() { ))} - - {/* 当前值不在预定义选项中时,显示为自定义选项 */} - {editHost && - editHost !== "127.0.0.1" && - editHost !== "0.0.0.0" && - !(networkInfo?.all_ips.includes(editHost) ?? false) && ( - - - - - - - {editHost} - - (自定义) - - - - - )} @@ -1161,7 +1151,7 @@ export function ApiServerPage() {

API 测试

+ {/* 模型选择器 */} +
+ 测试模型: + + + + + + + + + + + {allModels.map((model) => ( + + + + + + + {model.display_name} + + {model.provider_name} + + + {model.tier} + + + + + ))} + + + + +
+
{testEndpoints.map((endpoint) => { const result = testResults[endpoint.id]; diff --git a/src/components/provider-pool/api-key/ProviderConfigForm.tsx b/src/components/provider-pool/api-key/ProviderConfigForm.tsx index 44a96de65..c6c52671f 100644 --- a/src/components/provider-pool/api-key/ProviderConfigForm.tsx +++ b/src/components/provider-pool/api-key/ProviderConfigForm.tsx @@ -356,7 +356,8 @@ export const ProviderConfigForm: React.FC = ({ data-testid="custom-models-input" />

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

diff --git a/src/components/ui/checkbox.tsx b/src/components/ui/checkbox.tsx index 1207fb99e..d8fd71d21 100644 --- a/src/components/ui/checkbox.tsx +++ b/src/components/ui/checkbox.tsx @@ -7,10 +7,8 @@ import React from "react"; import { Check } from "lucide-react"; import { cn } from "@/lib/utils"; -interface CheckboxProps extends Omit< - React.ButtonHTMLAttributes, - "onChange" -> { +interface CheckboxProps + extends Omit, "onChange"> { checked?: boolean; onCheckedChange?: (checked: boolean) => void; } diff --git a/src/components/ui/switch.tsx b/src/components/ui/switch.tsx index c973fab51..6bc20548d 100644 --- a/src/components/ui/switch.tsx +++ b/src/components/ui/switch.tsx @@ -1,10 +1,8 @@ import React from "react"; import { cn } from "@/lib/utils"; -interface SwitchProps extends Omit< - React.ButtonHTMLAttributes, - "onChange" -> { +interface SwitchProps + extends Omit, "onChange"> { checked?: boolean; onCheckedChange?: (checked: boolean) => void; } diff --git a/src/hooks/useTauri.ts b/src/hooks/useTauri.ts index fa7d1e74c..c9822accb 100644 --- a/src/hooks/useTauri.ts +++ b/src/hooks/useTauri.ts @@ -209,6 +209,28 @@ export async function setDefaultProvider(provider: string): Promise { return safeInvoke("set_default_provider", { provider }); } +/** + * 更新 Provider 的环境变量 + * + * 当用户在 API Server 页面选择一个 API Key Provider 时调用 + * 会更新 ~/.claude/settings.json 和 shell 配置文件中的环境变量 + * + * @param providerType Provider 类型(如 "anthropic", "openai", "gemini") + * @param apiHost Provider 的 API Host + * @param apiKey 可选的 API Key + */ +export async function updateProviderEnvVars( + providerType: string, + apiHost: string, + apiKey?: string, +): Promise { + return safeInvoke("update_provider_env_vars", { + providerType, + apiHost, + apiKey: apiKey || null, + }); +} + export async function refreshKiroToken(): Promise { return safeInvoke("refresh_kiro_token"); } diff --git a/src/lib/notificationService.ts b/src/lib/notificationService.ts index 27e9704ab..da000d96d 100644 --- a/src/lib/notificationService.ts +++ b/src/lib/notificationService.ts @@ -207,9 +207,8 @@ class NotificationService { private playSound(type?: NotificationType): void { // 使用 Web Audio API 播放简单的提示音 try { - const audioContext = new ( - window.AudioContext || (window as any).webkitAudioContext - )(); + const audioContext = new (window.AudioContext || + (window as any).webkitAudioContext)(); const oscillator = audioContext.createOscillator(); const gainNode = audioContext.createGain(); diff --git a/src/lib/plugin-ui/PluginUIRenderer.tsx b/src/lib/plugin-ui/PluginUIRenderer.tsx index 45b4001d9..5800610af 100644 --- a/src/lib/plugin-ui/PluginUIRenderer.tsx +++ b/src/lib/plugin-ui/PluginUIRenderer.tsx @@ -201,10 +201,11 @@ function getSurfaceStyles(surface: SurfaceState): React.CSSProperties { /** * 单个组件渲染器 */ -interface ComponentRendererInternalProps extends Omit< - ComponentRendererProps, - "resolveValue" | "renderChild" | "renderChildren" -> { +interface ComponentRendererInternalProps + extends Omit< + ComponentRendererProps, + "resolveValue" | "renderChild" | "renderChildren" + > { resolveValue: (bound: BoundValue) => T | undefined; renderChild: (childId: ComponentId) => React.ReactNode; renderChildren: (children: ChildrenDef) => React.ReactNode[];