feat: add provider environment variable sync and UI improvements

- Add updateProviderEnvVars function to sync Claude Code env vars
- Improve API Server page with provider selection and env var updates
- Update codex provider with better error handling
- Fix provider pool service credential lookup
- Format UI components with Prettier
- Fix test files for custom_models field

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
coso
2026-01-12 16:19:22 +08:00
co-authored by Claude Opus 4.5
parent a52ec0123e
commit 9789e7d8a8
23 changed files with 670 additions and 472 deletions
+2 -2
View File
@@ -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"
}
+2 -5
View File
@@ -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())
}
+1
View File
@@ -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,
+43 -33
View File
@@ -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<Vec<String>, 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<Vec<String>, 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<HashMap<String, Vec<String>>, 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<Vec<String>, String> {
pub fn get_all_available_models(db: State<'_, DbConnection>) -> Result<Vec<String>, 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<HashMap<String, Result<Vec<String>, 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<Vec<String>, String> {
let pt: crate::models::provider_pool_model::PoolProviderType =
pub fn get_default_models_for_provider(provider_type: String) -> Result<Vec<String>, 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))
}
+8 -4
View File
@@ -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())
@@ -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
}
@@ -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,
+37 -16
View File
@@ -42,7 +42,11 @@ impl AntigravityApiError {
}
/// 创建带响应体的 API 错误
pub fn with_body(status_code: u16, message: impl Into<String>, body: impl Into<String>) -> Self {
pub fn with_body(
status_code: u16,
message: impl Into<String>,
body: impl Into<String>,
) -> 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<serde_json::Value, AntigravityApiError> {
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?;
+53 -8
View File
@@ -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<RwLock<Option<oneshot::Sender<()>>>> =
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::<Result<CodexOAuthResult, String>>();
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())
}
}
};
+7 -4
View File
@@ -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
@@ -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())
}
+36 -20
View File
@@ -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<AppState>) -> 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<AppState>) -> 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())
+1 -1
View File
@@ -88,7 +88,7 @@ fn get_shell_config_path() -> Result<PathBuf, Box<dyn std::error::Error + Send +
/// 将环境变量写入 shell 配置文件
/// 使用标记块管理,避免重复添加
pub(crate) fn write_env_to_shell_config(
pub fn write_env_to_shell_config(
env_vars: &[(String, String)],
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let config_path = get_shell_config_path()?;
+57 -30
View File
@@ -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<String> = 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<String> = 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<String> = stmt
.query_row([credential_uuid], |row| row.get(0))
.ok();
let models_json: Option<String> = 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<Vec<String>, String> {
let models_by_provider = self.get_all_models_by_provider(db)?;
let mut all_models: Vec<String> = models_by_provider
.into_values()
.flatten()
.collect();
let mut all_models: Vec<String> = models_by_provider.into_values().flatten().collect();
all_models.sort();
all_models.dedup();
+65 -55
View File
@@ -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::<String>()
))
}
}
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::<String>()
))
}
}
@@ -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,
};
+305 -266
View File
@@ -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<NetworkInfo | null>(null);
// 模型库状态
const [allModels, setAllModels] = useState<EnhancedModelMetadata[]>([]);
const [testModel, setTestModel] = useState<string>("");
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() {
</Select.ItemText>
</Select.Item>
))}
{/* 当前值不在预定义选项中时,显示为自定义选项 */}
{editHost &&
editHost !== "127.0.0.1" &&
editHost !== "0.0.0.0" &&
!(networkInfo?.all_ips.includes(editHost) ?? false) && (
<Select.Item
key={editHost}
value={editHost}
className="relative flex cursor-pointer select-none items-center rounded-sm px-8 py-2.5 text-sm outline-none transition-colors hover:bg-accent hover:text-accent-foreground focus:bg-accent focus:text-accent-foreground data-[disabled]:pointer-events-none data-[disabled]:opacity-50"
>
<Select.ItemIndicator className="absolute left-2 flex h-3.5 w-3.5 items-center justify-center">
<Check className="h-4 w-4" />
</Select.ItemIndicator>
<Select.ItemText>
<span className="flex items-center gap-2">
<span className="font-mono">{editHost}</span>
<span className="text-xs text-muted-foreground">
(自定义)
</span>
</span>
</Select.ItemText>
</Select.Item>
)}
</Select.Viewport>
</Select.Content>
</Select.Portal>
@@ -1161,7 +1151,7 @@ export function ApiServerPage() {
<h3 className="font-semibold">API 测试</h3>
<button
onClick={runAllTests}
disabled={!status?.running}
disabled={!status?.running || testEndpoints.length === 0}
className="flex items-center gap-2 rounded-lg bg-primary px-3 py-1.5 text-sm font-medium text-primary-foreground hover:bg-primary/90 disabled:opacity-50"
>
<Play className="h-4 w-4" />
@@ -1169,6 +1159,55 @@ export function ApiServerPage() {
</button>
</div>
{/* 模型选择器 */}
<div className="mb-4 flex items-center gap-3">
<span className="text-sm text-muted-foreground">测试模型:</span>
<Select.Root value={testModel} onValueChange={setTestModel}>
<Select.Trigger className="inline-flex min-w-[300px] items-center justify-between gap-2 rounded-md border border-input bg-background px-3 py-1.5 text-sm shadow-sm transition-colors hover:bg-accent hover:text-accent-foreground focus:outline-none focus:ring-2 focus:ring-ring disabled:cursor-not-allowed disabled:opacity-50">
<Select.Value placeholder="选择模型..." />
<Select.Icon>
<ChevronDown className="h-4 w-4 opacity-50" />
</Select.Icon>
</Select.Trigger>
<Select.Portal>
<Select.Content className="relative z-50 max-h-[300px] min-w-[300px] overflow-hidden rounded-md border border-border bg-white dark:bg-gray-900 text-foreground shadow-lg animate-in fade-in-80 data-[side=bottom]:slide-in-from-top-2 data-[side=top]:slide-in-from-bottom-2">
<Select.Viewport className="p-1 max-h-[280px] overflow-y-auto">
{allModels.map((model) => (
<Select.Item
key={model.id}
value={model.id}
className="relative flex cursor-pointer select-none items-center rounded-sm px-8 py-2 text-sm outline-none transition-colors hover:bg-accent hover:text-accent-foreground focus:bg-accent focus:text-accent-foreground data-[disabled]:pointer-events-none data-[disabled]:opacity-50"
>
<Select.ItemIndicator className="absolute left-2 flex h-3.5 w-3.5 items-center justify-center">
<Check className="h-4 w-4" />
</Select.ItemIndicator>
<Select.ItemText>
<span className="flex items-center gap-2">
<span>{model.display_name}</span>
<span className="text-xs opacity-60">
{model.provider_name}
</span>
<span
className={`text-xs px-1.5 py-0.5 rounded ${
model.tier === "max"
? "bg-purple-100 text-purple-700 dark:bg-purple-900/30 dark:text-purple-400"
: model.tier === "pro"
? "bg-blue-100 text-blue-700 dark:bg-blue-900/30 dark:text-blue-400"
: "bg-gray-100 text-gray-700 dark:bg-gray-900/30 dark:text-gray-400"
}`}
>
{model.tier}
</span>
</span>
</Select.ItemText>
</Select.Item>
))}
</Select.Viewport>
</Select.Content>
</Select.Portal>
</Select.Root>
</div>
<div className="space-y-3">
{testEndpoints.map((endpoint) => {
const result = testResults[endpoint.id];
@@ -356,7 +356,8 @@ export const ProviderConfigForm: React.FC<ProviderConfigFormProps> = ({
data-testid="custom-models-input"
/>
<p className="text-xs text-muted-foreground">
该 Provider 支持的模型列表,用逗号分隔。用于不支持 /models 接口的 Provider(如智谱)
该 Provider 支持的模型列表,用逗号分隔。用于不支持 /models 接口的
Provider(如智谱)
</p>
</div>
+2 -4
View File
@@ -7,10 +7,8 @@ import React from "react";
import { Check } from "lucide-react";
import { cn } from "@/lib/utils";
interface CheckboxProps extends Omit<
React.ButtonHTMLAttributes<HTMLButtonElement>,
"onChange"
> {
interface CheckboxProps
extends Omit<React.ButtonHTMLAttributes<HTMLButtonElement>, "onChange"> {
checked?: boolean;
onCheckedChange?: (checked: boolean) => void;
}
+2 -4
View File
@@ -1,10 +1,8 @@
import React from "react";
import { cn } from "@/lib/utils";
interface SwitchProps extends Omit<
React.ButtonHTMLAttributes<HTMLButtonElement>,
"onChange"
> {
interface SwitchProps
extends Omit<React.ButtonHTMLAttributes<HTMLButtonElement>, "onChange"> {
checked?: boolean;
onCheckedChange?: (checked: boolean) => void;
}
+22
View File
@@ -209,6 +209,28 @@ export async function setDefaultProvider(provider: string): Promise<string> {
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<void> {
return safeInvoke("update_provider_env_vars", {
providerType,
apiHost,
apiKey: apiKey || null,
});
}
export async function refreshKiroToken(): Promise<string> {
return safeInvoke("refresh_kiro_token");
}
+2 -3
View File
@@ -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();
+5 -4
View File
@@ -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: <T>(bound: BoundValue<T>) => T | undefined;
renderChild: (childId: ComponentId) => React.ReactNode;
renderChildren: (children: ChildrenDef) => React.ReactNode[];