mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
Generated
+66
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 数量
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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()),
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 请求
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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,
|
||||
};
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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()));
|
||||
}
|
||||
}
|
||||
@@ -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 ? (
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user