Merge pull request #108 from Chiron-Brahm/main

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