Files
Genarrative/server-rs/crates/api-server/src/agc_models.rs
T
suzmii 4234a21eed 服务端:AGC 模型目录缺配置时改为启动期从上游同步,保留原目录格式
- 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 流式上限、禁止重定向、不带凭据;目标校验只校验地址,不再绑定已下线的固定模型哨兵;写回冲突后重读校验既有目录
2026-09-24 16:18:06 +08:00

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"
);
}
}