4234a21eed
- module-runtime:删除写死的初始目录(高质量→gpt-6-astra、快速→gpt-5.6-luna),新增 from_upstream_models:id 用模型名 slug、alias/modelId 用上游原名、默认项取排序后第一项;字段与校验规则不变 - spacetime-module:read_agc_model_catalog 缺行返回 AGC_MODEL_CATALOG_NOT_INITIALIZED,不再返回内置目录 - api-server:新增启动期 ensure_agc_model_catalog_initialized,未初始化时从分组定价列表 GET /api/pricing?group=taonier 生成目录并按存量 revision 写回;拉不到只记录 error、不写替代目录,下次启动重试 - api-server:目录未初始化时 AGC 目录/对话接口与后台目录接口返回 503,不回落任何内置模型名;后台 PUT 同样要求已初始化 - api-server:上游请求 10s 超时、1 MiB 流式上限、禁止重定向、不带凭据;目标校验只校验地址,不再绑定已下线的固定模型哨兵;写回冲突后重读校验既有目录
447 lines
18 KiB
Rust
447 lines
18 KiB
Rust
use crate::{
|
|
admin::AuthenticatedAdmin, api_response::json_success_body, http_error::AppError,
|
|
request_context::RequestContext, state::AppState,
|
|
};
|
|
use axum::{
|
|
Json,
|
|
extract::{Extension, State},
|
|
http::StatusCode,
|
|
};
|
|
use module_runtime::{
|
|
AGC_MODEL_CATALOG_CONFLICT, AGC_MODEL_CATALOG_NOT_INITIALIZED, AgcModelCatalog,
|
|
};
|
|
use shared_contracts::admin::{AdminAgcModel, AdminAgcModelCatalog};
|
|
use spacetime_client::SpacetimeClientError;
|
|
use std::time::Duration;
|
|
use tracing::warn;
|
|
|
|
/// 目录未初始化时对外统一的失败文案:目录只能来自上游同步或后台保存。
|
|
pub(crate) const AGC_MODEL_CATALOG_NOT_INITIALIZED_MESSAGE: &str =
|
|
"模型目录未初始化,服务端正在尝试从上游同步,请稍后重试";
|
|
/// 上游模型列表请求超时与响应大小上限;越界按同步失败处理。
|
|
///
|
|
/// 同步发生在启动期、且在开始对外服务之前,超时必须足够短:上游挂起时不能让
|
|
/// 每个 API/All 实例都延迟三十秒才可用。单次失败只记录 error,下次启动会重试。
|
|
const AGC_MODEL_LIST_REQUEST_TIMEOUT: Duration = Duration::from_secs(10);
|
|
const AGC_MODEL_LIST_MAX_BYTES: usize = 1024 * 1024;
|
|
|
|
/// 只读 `revision`:存量目录内容不合法时,覆盖写入仍需对齐乐观锁版本。
|
|
#[derive(serde::Deserialize)]
|
|
struct StoredCatalogRevision {
|
|
revision: u64,
|
|
}
|
|
|
|
pub(crate) async fn load_catalog(state: &AppState) -> Result<AgcModelCatalog, AppError> {
|
|
let stored = read_stored_catalog(state).await.map_err(|_| {
|
|
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE).with_message("模型目录暂不可用")
|
|
})?;
|
|
let Some(json) = stored else {
|
|
return Err(uninitialized_error());
|
|
};
|
|
parse_catalog(&json).map_err(|_| uninitialized_error())
|
|
}
|
|
|
|
fn uninitialized_error() -> AppError {
|
|
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE)
|
|
.with_message(AGC_MODEL_CATALOG_NOT_INITIALIZED_MESSAGE)
|
|
}
|
|
|
|
/// 解析并校验目录内容;解析或校验失败都按“未初始化”处理,由启动期重新同步。
|
|
fn parse_catalog(json: &str) -> Result<AgcModelCatalog, String> {
|
|
let catalog: AgcModelCatalog =
|
|
serde_json::from_str(json).map_err(|_| "模型目录格式无效".to_string())?;
|
|
catalog.validate()?;
|
|
Ok(catalog)
|
|
}
|
|
|
|
async fn read_stored_catalog(state: &AppState) -> Result<Option<String>, SpacetimeClientError> {
|
|
match state.spacetime_client().read_agc_model_catalog().await {
|
|
Ok(json) => Ok(Some(json)),
|
|
Err(SpacetimeClientError::Procedure(message))
|
|
if message == AGC_MODEL_CATALOG_NOT_INITIALIZED =>
|
|
{
|
|
Ok(None)
|
|
}
|
|
Err(error) => Err(error),
|
|
}
|
|
}
|
|
|
|
/// 启动期确保目录已初始化:未初始化时从上游模型列表重建,失败只返回错误由调用方记录。
|
|
///
|
|
/// 目录已可用时不做任何写入;只有缺失、结构与当前定义不符或校验不通过才重建,因此上游模型
|
|
/// 变化不会自动覆盖后台维护过的目录。
|
|
pub(crate) async fn ensure_agc_model_catalog_initialized(state: &AppState) -> Result<(), String> {
|
|
let stored = read_stored_catalog(state)
|
|
.await
|
|
.map_err(|error| format!("读取 AGC 模型目录失败:{error}"))?;
|
|
let revision = match stored.as_deref() {
|
|
Some(json) => match parse_catalog(json) {
|
|
Ok(_) => return Ok(()),
|
|
Err(message) => {
|
|
warn!(
|
|
error = %message,
|
|
"AGC 模型目录内容与当前定义不符,按未初始化处理并从上游重建"
|
|
);
|
|
stored_catalog_revision(json)
|
|
}
|
|
},
|
|
None => Some(0),
|
|
};
|
|
let revision = revision.ok_or_else(|| {
|
|
"存量 AGC 模型目录缺少可解析的 revision,需要先清理该行再重启".to_string()
|
|
})?;
|
|
|
|
let models = fetch_upstream_model_names(state).await?;
|
|
let catalog = AgcModelCatalog::from_upstream_models(models, revision)?;
|
|
let payload =
|
|
serde_json::to_string(&catalog).map_err(|_| "AGC 模型目录序列化失败".to_string())?;
|
|
match state
|
|
.spacetime_client()
|
|
.save_agc_model_catalog(payload)
|
|
.await
|
|
{
|
|
Ok(saved) => {
|
|
let saved: AgcModelCatalog = serde_json::from_str(&saved)
|
|
.map_err(|_| "AGC 模型目录写回结果格式无效".to_string())?;
|
|
tracing::info!(
|
|
revision = saved.revision,
|
|
model_count = saved.models.len(),
|
|
"已按上游模型列表初始化 AGC 模型目录"
|
|
);
|
|
Ok(())
|
|
}
|
|
// 多实例同时启动时只有一个写入成功:接受既有目录,但仍要确认它可用,
|
|
// 否则会静默地把「每次启动都冲突、目录一直不可用」变成没有任何线索的黑洞。
|
|
Err(SpacetimeClientError::Procedure(message)) if message == AGC_MODEL_CATALOG_CONFLICT => {
|
|
let stored = read_stored_catalog(state)
|
|
.await
|
|
.map_err(|error| format!("写入冲突后重读 AGC 模型目录失败:{error}"))?;
|
|
if stored
|
|
.as_deref()
|
|
.map(parse_catalog)
|
|
.is_some_and(|result| result.is_ok())
|
|
{
|
|
warn!("AGC 模型目录写入冲突:已接受其它实例写入的目录");
|
|
Ok(())
|
|
} else {
|
|
Err("AGC 模型目录写入冲突后仍不可用,需要人工检查该行内容与 revision".to_string())
|
|
}
|
|
}
|
|
Err(error) => Err(format!("写入 AGC 模型目录失败:{error}")),
|
|
}
|
|
}
|
|
|
|
fn stored_catalog_revision(json: &str) -> Option<u64> {
|
|
serde_json::from_str::<StoredCatalogRevision>(json)
|
|
.ok()
|
|
.map(|stored| stored.revision)
|
|
}
|
|
|
|
/// 上游在售模型列表:`GET {Router 控制面}/api/pricing?group=taonier` 的 `data[].model_name`。
|
|
///
|
|
/// 取“该分组可见的在售模型”,而不是管理面模型注册表:注册表里会残留已下线、没有路由绑定的
|
|
/// 条目(例如已从上游移除的 `gpt-6-astra`/`gpt-6-luna`),而定价列表就是 AGC 账号实际能调用的集合。
|
|
/// 该端点是公开只读接口,不需要管理凭据。
|
|
async fn fetch_upstream_model_names(state: &AppState) -> Result<Vec<String>, String> {
|
|
crate::external_api_keys::ensure_llm_router_url_allowed(state)?;
|
|
let origin =
|
|
crate::external_api_keys::router_control_origin(&state.config.llm_router_base_url)?;
|
|
let url = format!(
|
|
"{origin}/api/pricing?group={}",
|
|
crate::external_api_keys::LLM_ROUTER_TOKEN_GROUP
|
|
);
|
|
let client = reqwest::Client::builder()
|
|
.timeout(AGC_MODEL_LIST_REQUEST_TIMEOUT)
|
|
.redirect(reqwest::redirect::Policy::none())
|
|
.build()
|
|
.map_err(|error| format!("构建 LLM Router 客户端失败:{error}"))?;
|
|
let response = client
|
|
.get(url)
|
|
.send()
|
|
.await
|
|
.map_err(|error| format!("请求上游模型列表失败:{error}"))?;
|
|
let status = response.status();
|
|
if !status.is_success() {
|
|
return Err(format!("上游模型列表返回 HTTP {status}"));
|
|
}
|
|
let bytes = read_bounded_json_body(response).await?;
|
|
let payload: serde_json::Value =
|
|
serde_json::from_slice(&bytes).map_err(|_| "上游模型列表格式无效".to_string())?;
|
|
parse_upstream_model_names(&payload)
|
|
}
|
|
|
|
async fn read_bounded_json_body(mut response: reqwest::Response) -> Result<Vec<u8>, String> {
|
|
// 先按 Content-Length 快速拒绝,再流式累加做兜底:不信任上游声明的长度,
|
|
// 逐块累计超阈值立即中断,避免 `bytes()` 一次性分配任意大小响应撑爆内存。
|
|
if response
|
|
.content_length()
|
|
.is_some_and(|length| length > AGC_MODEL_LIST_MAX_BYTES as u64)
|
|
{
|
|
return Err("上游模型列表响应超过大小上限".to_string());
|
|
}
|
|
let mut bytes = Vec::new();
|
|
while let Some(chunk) = response
|
|
.chunk()
|
|
.await
|
|
.map_err(|error| format!("读取上游模型列表失败:{error}"))?
|
|
{
|
|
if bytes.len().saturating_add(chunk.len()) > AGC_MODEL_LIST_MAX_BYTES {
|
|
return Err("上游模型列表响应超过大小上限".to_string());
|
|
}
|
|
bytes.extend_from_slice(chunk.as_ref());
|
|
}
|
|
Ok(bytes)
|
|
}
|
|
|
|
fn parse_upstream_model_names(payload: &serde_json::Value) -> Result<Vec<String>, String> {
|
|
let data = payload
|
|
.get("data")
|
|
.and_then(serde_json::Value::as_array)
|
|
.ok_or_else(|| "上游模型列表缺少 data 数组".to_string())?;
|
|
let models = data
|
|
.iter()
|
|
.filter_map(|entry| entry.get("model_name").and_then(serde_json::Value::as_str))
|
|
.map(str::to_string)
|
|
.collect::<Vec<_>>();
|
|
if models.iter().all(|model| model.trim().is_empty()) {
|
|
return Err("上游模型列表为空".to_string());
|
|
}
|
|
Ok(models)
|
|
}
|
|
|
|
pub async fn admin_get_agc_models(
|
|
State(state): State<AppState>,
|
|
Extension(context): Extension<RequestContext>,
|
|
Extension(_admin): Extension<AuthenticatedAdmin>,
|
|
) -> Result<Json<serde_json::Value>, AppError> {
|
|
Ok(json_success_body(
|
|
Some(&context),
|
|
catalog_dto(load_catalog(&state).await?),
|
|
))
|
|
}
|
|
|
|
pub async fn admin_save_agc_models(
|
|
State(state): State<AppState>,
|
|
Extension(context): Extension<RequestContext>,
|
|
Extension(_admin): Extension<AuthenticatedAdmin>,
|
|
Json(payload): Json<AdminAgcModelCatalog>,
|
|
) -> Result<Json<serde_json::Value>, AppError> {
|
|
// 目录只来自上游同步:未初始化时后台写入同样失败关闭,避免出现第二条绕过同步的写入口。
|
|
load_catalog(&state).await?;
|
|
let catalog = AgcModelCatalog {
|
|
revision: payload.revision,
|
|
default_model_id: payload.default_model_id,
|
|
models: payload
|
|
.models
|
|
.into_iter()
|
|
.map(|m| module_runtime::AgcModel {
|
|
id: m.id,
|
|
alias: m.alias,
|
|
model_id: m.model_id,
|
|
enabled: m.enabled,
|
|
})
|
|
.collect(),
|
|
};
|
|
catalog
|
|
.validate()
|
|
.map_err(|message| AppError::from_status(StatusCode::BAD_REQUEST).with_message(message))?;
|
|
let payload = serde_json::to_string(&catalog)
|
|
.map_err(|_| AppError::from_status(StatusCode::INTERNAL_SERVER_ERROR))?;
|
|
let saved = state
|
|
.spacetime_client()
|
|
.save_agc_model_catalog(payload)
|
|
.await
|
|
.map_err(|error| {
|
|
if matches!(&error, SpacetimeClientError::Procedure(message) if message == AGC_MODEL_CATALOG_CONFLICT) {
|
|
AppError::from_status(StatusCode::CONFLICT).with_message("模型目录已被更新,请重新读取")
|
|
} else {
|
|
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE).with_message("保存模型目录失败,请稍后重试")
|
|
}
|
|
})?;
|
|
let catalog: AgcModelCatalog = serde_json::from_str(&saved)
|
|
.map_err(|_| AppError::from_status(StatusCode::INTERNAL_SERVER_ERROR))?;
|
|
Ok(json_success_body(Some(&context), catalog_dto(catalog)))
|
|
}
|
|
|
|
fn catalog_dto(catalog: AgcModelCatalog) -> AdminAgcModelCatalog {
|
|
AdminAgcModelCatalog {
|
|
revision: catalog.revision,
|
|
default_model_id: catalog.default_model_id,
|
|
models: catalog
|
|
.models
|
|
.into_iter()
|
|
.map(|m| AdminAgcModel {
|
|
id: m.id,
|
|
alias: m.alias,
|
|
model_id: m.model_id,
|
|
enabled: m.enabled,
|
|
})
|
|
.collect(),
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use serde_json::json;
|
|
|
|
#[test]
|
|
fn upstream_model_names_come_from_pricing_data_array() {
|
|
let payload = json!({
|
|
"auto_groups": ["default"],
|
|
"data": [
|
|
{"model_name": "glm-5.3", "model_ratio": 1.0},
|
|
{"model_name": "deepseek-flash", "model_ratio": 0.075},
|
|
{"model_ratio": 1.0}
|
|
]
|
|
});
|
|
assert_eq!(
|
|
parse_upstream_model_names(&payload).unwrap(),
|
|
vec!["glm-5.3".to_string(), "deepseek-flash".to_string()]
|
|
);
|
|
|
|
assert_eq!(
|
|
parse_upstream_model_names(&json!({"data": []})).unwrap_err(),
|
|
"上游模型列表为空"
|
|
);
|
|
assert_eq!(
|
|
parse_upstream_model_names(&json!({"data": [{"model_name": " "}]})).unwrap_err(),
|
|
"上游模型列表为空"
|
|
);
|
|
assert!(parse_upstream_model_names(&json!({"object": "list"})).is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn stored_catalog_revision_reads_row_revision() {
|
|
assert_eq!(
|
|
stored_catalog_revision(
|
|
r#"{"revision":4,"defaultModelId":"quality","models":[{"id":"quality","alias":"高质量","modelId":"gpt-6-astra","enabled":true}]}"#
|
|
),
|
|
Some(4)
|
|
);
|
|
assert_eq!(stored_catalog_revision("not json"), None);
|
|
assert_eq!(stored_catalog_revision(r#"{"models":[]}"#), None);
|
|
}
|
|
|
|
#[test]
|
|
fn catalog_parsing_marks_unusable_content_as_uninitialized() {
|
|
// 后台保存过的目录结构必须能直接解析。
|
|
let catalog = AgcModelCatalog::from_upstream_models(
|
|
vec!["deepseek-v4-pro".to_string(), "glm-5.3".to_string()],
|
|
4,
|
|
)
|
|
.unwrap();
|
|
assert_eq!(
|
|
parse_catalog(&serde_json::to_string(&catalog).unwrap()).unwrap(),
|
|
catalog
|
|
);
|
|
|
|
// 结构或内容不合法(例如被外部工具改过)都按未初始化处理,由启动期重新同步。
|
|
assert!(
|
|
parse_catalog(r#"{"revision":1,"defaultModel":"deepseek-v4-pro","models":[]}"#)
|
|
.is_err()
|
|
);
|
|
assert!(parse_catalog(
|
|
r#"{"revision":1,"defaultModelId":"quality","models":[{"id":"quality","alias":"高质量","modelId":"gpt-6-astra","enabled":false}]}"#
|
|
)
|
|
.is_err());
|
|
}
|
|
|
|
struct MockModelListServer {
|
|
base_url: String,
|
|
captured: std::sync::Arc<std::sync::Mutex<Option<String>>>,
|
|
_handle: std::thread::JoinHandle<()>,
|
|
}
|
|
|
|
fn spawn_mock_model_list_server(status_line: &str, body: &str) -> MockModelListServer {
|
|
use std::io::{Read, Write};
|
|
|
|
let listener = std::net::TcpListener::bind("127.0.0.1:0").expect("mock listener binds");
|
|
let address = listener.local_addr().expect("mock address");
|
|
let captured = std::sync::Arc::new(std::sync::Mutex::new(None));
|
|
let captured_for_thread = std::sync::Arc::clone(&captured);
|
|
let response = format!(
|
|
"HTTP/1.1 {status_line}\r\ncontent-type: application/json; charset=utf-8\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
|
|
body.len()
|
|
);
|
|
let handle = std::thread::spawn(move || {
|
|
let (mut stream, _) = listener.accept().expect("mock accept");
|
|
let mut buffer = [0u8; 8192];
|
|
let read = stream.read(&mut buffer).unwrap_or_default();
|
|
*captured_for_thread.lock().expect("captured lock") =
|
|
Some(String::from_utf8_lossy(&buffer[..read]).to_string());
|
|
let _ = stream.write_all(response.as_bytes());
|
|
let _ = stream.flush();
|
|
});
|
|
MockModelListServer {
|
|
base_url: format!("http://{address}/v1"),
|
|
captured,
|
|
_handle: handle,
|
|
}
|
|
}
|
|
|
|
fn model_list_state(base_url: &str) -> AppState {
|
|
AppState::new(crate::config::AppConfig {
|
|
llm_router_base_url: base_url.to_string(),
|
|
..crate::config::AppConfig::default()
|
|
})
|
|
.expect("state should build")
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn fetch_upstream_model_names_reads_group_pricing_without_credentials() {
|
|
let server = spawn_mock_model_list_server(
|
|
"200 OK",
|
|
&json!({"data": [{"model_name": "glm-5.3"}, {"model_name": "deepseek-flash"}]})
|
|
.to_string(),
|
|
);
|
|
let state = model_list_state(&server.base_url);
|
|
|
|
assert_eq!(
|
|
fetch_upstream_model_names(&state).await.unwrap(),
|
|
vec!["glm-5.3".to_string(), "deepseek-flash".to_string()]
|
|
);
|
|
|
|
let request = server
|
|
.captured
|
|
.lock()
|
|
.expect("captured lock")
|
|
.clone()
|
|
.expect("mock server should capture request");
|
|
// 控制面路径由 base_url 推导(去掉 /v1),并显式带 AGC 账号所在分组。
|
|
assert!(
|
|
request.starts_with("GET /api/pricing?group=taonier HTTP/1.1"),
|
|
"{request}"
|
|
);
|
|
// 定价列表是公开只读接口:不得把任何凭据发过去。
|
|
assert!(
|
|
!request.to_ascii_lowercase().contains("authorization:"),
|
|
"{request}"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn fetch_upstream_model_names_fails_closed_when_upstream_unavailable_or_empty() {
|
|
let unauthorized = spawn_mock_model_list_server("401 Unauthorized", "{}");
|
|
let state = model_list_state(&unauthorized.base_url);
|
|
assert_eq!(
|
|
fetch_upstream_model_names(&state).await.unwrap_err(),
|
|
"上游模型列表返回 HTTP 401 Unauthorized"
|
|
);
|
|
|
|
let empty = spawn_mock_model_list_server("200 OK", &json!({"data": []}).to_string());
|
|
let state = model_list_state(&empty.base_url);
|
|
assert_eq!(
|
|
fetch_upstream_model_names(&state).await.unwrap_err(),
|
|
"上游模型列表为空"
|
|
);
|
|
|
|
let failing = spawn_mock_model_list_server("500 Internal Server Error", "{}");
|
|
let state = model_list_state(&failing.base_url);
|
|
assert_eq!(
|
|
fetch_upstream_model_names(&state).await.unwrap_err(),
|
|
"上游模型列表返回 HTTP 500 Internal Server Error"
|
|
);
|
|
}
|
|
}
|