mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
feat: 动态检测 IP 地址变化并自动更新
主要改进: 1. 服务器启动时检查配置的 IP 是否有效,无效时自动切换到当前局域网 IP 2. 前端定期刷新网络信息,实时检测 IP 变化 3. 路由端点和 API 测试地址动态更新 4. 优先选择 192.168.x.x 或 10.x.x.x 的局域网 IP 修复 Antigravity API: 1. 修正 API 端点为 cloudcode-pa.googleapis.com 2. 透传上游错误状态码 新增功能: 1. 添加 model_service 和 model_cmd 模块
This commit is contained in:
@@ -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}"));
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
|
||||
@@ -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()),
|
||||
|
||||
@@ -219,6 +219,12 @@ pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> {
|
||||
[],
|
||||
);
|
||||
|
||||
// Migration: 添加支持的模型列表字段
|
||||
let _ = conn.execute(
|
||||
"ALTER TABLE provider_pool_credentials ADD COLUMN supported_models TEXT",
|
||||
[],
|
||||
);
|
||||
|
||||
// Migration: 添加代理URL字段 - 使用重建表结构的方式
|
||||
migrate_add_proxy_url_column(conn)?;
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -1,3 +1,11 @@
|
||||
<!--
|
||||
* @Author: Chiron 598621670@qq.com
|
||||
* @Date: 2026-01-06 17:34:03
|
||||
* @LastEditors: Chiron 598621670@qq.com
|
||||
* @LastEditTime: 2026-01-11 01:09:49
|
||||
* @FilePath: /proxycast/src-tauri/src/providers/README.md
|
||||
* @Description: 这是默认设置,请设置`customMade`, 打开koroFileHeader查看配置 进行设置: https://github.com/OBKoro1/koro1FileHeader/wiki/%E9%85%8D%E7%BD%AE
|
||||
-->
|
||||
# providers
|
||||
|
||||
<!-- 一旦我所属的文件夹有所变化,请更新我 -->
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -1085,6 +1085,7 @@ impl ApiKeyProviderService {
|
||||
check_health: false,
|
||||
check_model_name: None,
|
||||
not_supported_models: Vec::new(),
|
||||
supported_models: Vec::new(),
|
||||
usage_count: 0,
|
||||
error_count: 0,
|
||||
last_used: None,
|
||||
@@ -1142,6 +1143,7 @@ impl ApiKeyProviderService {
|
||||
check_health: false, // 降级凭证不参与健康检查
|
||||
check_model_name: None,
|
||||
not_supported_models: Vec::new(),
|
||||
supported_models: Vec::new(),
|
||||
usage_count: 0,
|
||||
error_count: 0,
|
||||
last_used: None,
|
||||
|
||||
@@ -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,25 +476,51 @@ export function ApiServerPage() {
|
||||
} else {
|
||||
setProviderSwitchMsg(`已切换到 ${label}`);
|
||||
}
|
||||
|
||||
// 在后台异步刷新凭证池数据,不阻塞 UI
|
||||
providerPoolApi.getOverview().then(setPoolOverview).catch(console.error);
|
||||
} catch (e: unknown) {
|
||||
const errMsg = e instanceof Error ? e.message : String(e);
|
||||
setProviderSwitchMsg(`切换失败: ${errMsg}`);
|
||||
// 切换失败时恢复原来的状态
|
||||
loadDefaultProvider();
|
||||
}
|
||||
};
|
||||
|
||||
// 根据监听地址智能选择测试 URL
|
||||
// - 127.0.0.1: 使用 127.0.0.1(仅本机)
|
||||
// - 0.0.0.0: 使用 127.0.0.1(本机访问所有接口)
|
||||
// - 0.0.0.0: 使用当前局域网 IP(优先 192.168.x.x 或 10.x.x.x)
|
||||
// - 局域网 IP: 使用该 IP(允许局域网测试)
|
||||
const getTestUrl = (host: string, port: number) => {
|
||||
if (host === "0.0.0.0") {
|
||||
return `http://127.0.0.1:${port}`;
|
||||
// 0.0.0.0 时,使用当前局域网 IP 以便局域网设备访问
|
||||
const lanIp = networkInfo?.all_ips.find(ip =>
|
||||
ip.startsWith("192.168.") || ip.startsWith("10.")
|
||||
) || networkInfo?.all_ips[0] || "127.0.0.1";
|
||||
return `http://${lanIp}:${port}`;
|
||||
}
|
||||
return `http://${host}:${port}`;
|
||||
};
|
||||
|
||||
// 使用 editHost 而不是 status.host,这样可以实时反映用户的选择
|
||||
const currentHost = status?.running ? status.host : editHost;
|
||||
// 同时检查配置的 IP 是否仍然有效(在当前网卡列表中)
|
||||
const getValidHost = () => {
|
||||
const host = status?.running ? status.host : editHost;
|
||||
// 如果是特殊地址,直接返回
|
||||
if (host === "127.0.0.1" || host === "0.0.0.0") {
|
||||
return host;
|
||||
}
|
||||
// 检查配置的 IP 是否在当前网卡列表中
|
||||
if (networkInfo?.all_ips && !networkInfo.all_ips.includes(host)) {
|
||||
// IP 已失效,返回当前有效的局域网 IP
|
||||
return networkInfo.all_ips.find(ip =>
|
||||
ip.startsWith("192.168.") || ip.startsWith("10.")
|
||||
) || networkInfo.all_ips[0] || host;
|
||||
}
|
||||
return host;
|
||||
};
|
||||
|
||||
const currentHost = getValidHost();
|
||||
const currentPort = status?.running
|
||||
? status.port
|
||||
: parseInt(editPort) || 8999;
|
||||
@@ -599,10 +667,7 @@ export function ApiServerPage() {
|
||||
},
|
||||
}));
|
||||
|
||||
// 测试成功后立即刷新凭证池数据,更新使用次数
|
||||
if (result.success) {
|
||||
await loadPoolOverview();
|
||||
}
|
||||
return result.success;
|
||||
} catch (e: unknown) {
|
||||
const errMsg = e instanceof Error ? e.message : String(e);
|
||||
setTestResults((prev) => ({
|
||||
@@ -613,12 +678,21 @@ export function ApiServerPage() {
|
||||
response: `请求失败: ${errMsg}`,
|
||||
},
|
||||
}));
|
||||
return false;
|
||||
}
|
||||
};
|
||||
|
||||
const runAllTests = async () => {
|
||||
let hasSuccess = false;
|
||||
for (const endpoint of testEndpoints) {
|
||||
await runTest(endpoint);
|
||||
const success = await runTest(endpoint);
|
||||
if (success) {
|
||||
hasSuccess = true;
|
||||
}
|
||||
}
|
||||
// 所有测试完成后,如果有成功的测试,刷新一次凭证池数据
|
||||
if (hasSuccess) {
|
||||
await loadPoolOverview();
|
||||
}
|
||||
};
|
||||
|
||||
@@ -811,6 +885,30 @@ export function ApiServerPage() {
|
||||
</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>
|
||||
|
||||
@@ -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