服务端: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 流式上限、禁止重定向、不带凭据;目标校验只校验地址,不再绑定已下线的固定模型哨兵;写回冲突后重读校验既有目录
This commit is contained in:
@@ -7,26 +7,208 @@ use axum::{
|
||||
extract::{Extension, State},
|
||||
http::StatusCode,
|
||||
};
|
||||
use module_runtime::AgcModelCatalog;
|
||||
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 json = state
|
||||
.spacetime_client()
|
||||
.read_agc_model_catalog()
|
||||
.await
|
||||
.map_err(|_| {
|
||||
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE).with_message("模型目录暂不可用")
|
||||
})?;
|
||||
let catalog: AgcModelCatalog = serde_json::from_str(&json).map_err(|_| {
|
||||
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE).with_message("模型目录格式无效")
|
||||
})?;
|
||||
catalog.validate().map_err(|message| {
|
||||
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE).with_message(message)
|
||||
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>,
|
||||
@@ -44,6 +226,8 @@ pub async fn admin_save_agc_models(
|
||||
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,
|
||||
@@ -68,7 +252,7 @@ pub async fn admin_save_agc_models(
|
||||
.save_agc_model_catalog(payload)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
if matches!(error, spacetime_client::SpacetimeClientError::Procedure(ref message) if message == module_runtime::AGC_MODEL_CATALOG_CONFLICT) {
|
||||
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("保存模型目录失败,请稍后重试")
|
||||
@@ -95,3 +279,168 @@ fn catalog_dto(catalog: AgcModelCatalog) -> AdminAgcModelCatalog {
|
||||
.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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -49,7 +49,7 @@ const EXTERNAL_API_KEY_SCOPES: [&str; 4] = [
|
||||
const LLM_ROUTER_TOKEN_IDENTIFIER: &str = "agc_auto_generate";
|
||||
/// Router 用户(账号)与它名下固定 Token / API Key 都归属同一分组 `taonier`。
|
||||
const LLM_ROUTER_USER_GROUP: &str = "taonier";
|
||||
const LLM_ROUTER_TOKEN_GROUP: &str = "taonier";
|
||||
pub(crate) const LLM_ROUTER_TOKEN_GROUP: &str = "taonier";
|
||||
const LLM_ROUTER_API_KEY_SCOPES: [&str; 1] = ["llm:responses"];
|
||||
const LLM_ROUTER_SUBSCRIPTION_PLAN_ID: i64 = 1;
|
||||
const LLM_ROUTER_SUBSCRIPTION_RENEWAL_THRESHOLD_SECONDS: i64 = 24 * 60 * 60;
|
||||
@@ -1624,7 +1624,7 @@ async fn ensure_router_token_contract(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn router_control_origin(base_url: &str) -> Result<String, String> {
|
||||
pub(crate) fn router_control_origin(base_url: &str) -> Result<String, String> {
|
||||
let mut url = reqwest::Url::parse(base_url.trim_end_matches('/'))
|
||||
.map_err(|error| format!("LLM Router 地址无效:{error}"))?;
|
||||
let is_loopback = url.host_str().is_some_and(|host| {
|
||||
@@ -1643,7 +1643,11 @@ fn router_control_origin(base_url: &str) -> Result<String, String> {
|
||||
Ok(url.to_string().trim_end_matches('/').to_string())
|
||||
}
|
||||
|
||||
fn ensure_llm_router_target_allowed(state: &AppState) -> Result<(), String> {
|
||||
/// 只校验 LLM Router 目标地址是否允许(官方路由 / loopback、scheme),不校验固定模型。
|
||||
///
|
||||
/// 与具体模型无关的调用(例如按分组定价列表同步 AGC 模型目录)用这个入口,
|
||||
/// 避免被“必须使用官方固定模型”的哨兵常量挡住。
|
||||
pub(crate) fn ensure_llm_router_url_allowed(state: &AppState) -> Result<(), String> {
|
||||
let base_url = state.config.llm_router_base_url.trim_end_matches('/');
|
||||
let url =
|
||||
reqwest::Url::parse(base_url).map_err(|error| format!("LLM Router 地址无效:{error}"))?;
|
||||
@@ -1664,9 +1668,6 @@ fn ensure_llm_router_target_allowed(state: &AppState) -> Result<(), String> {
|
||||
if base_url != OFFICIAL_LLM_ROUTER_BASE_URL {
|
||||
return Err("生产环境 LLM Router 必须使用官方固定路由".to_string());
|
||||
}
|
||||
if state.config.llm_router_model.trim() != OFFICIAL_LLM_ROUTER_MODEL {
|
||||
return Err("生产环境 LLM Router 必须使用官方固定模型".to_string());
|
||||
}
|
||||
if url.scheme() != "https" {
|
||||
return Err("生产环境 LLM Router 只允许 HTTPS 地址".to_string());
|
||||
}
|
||||
@@ -1674,9 +1675,6 @@ fn ensure_llm_router_target_allowed(state: &AppState) -> Result<(), String> {
|
||||
}
|
||||
|
||||
if base_url == OFFICIAL_LLM_ROUTER_BASE_URL {
|
||||
if state.config.llm_router_model.trim() != OFFICIAL_LLM_ROUTER_MODEL {
|
||||
return Err("LLM Router 必须使用官方固定模型".to_string());
|
||||
}
|
||||
if url.scheme() != "https" {
|
||||
return Err("官方 LLM Router 只允许 HTTPS 地址".to_string());
|
||||
}
|
||||
@@ -1698,6 +1696,19 @@ fn ensure_llm_router_target_allowed(state: &AppState) -> Result<(), String> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn ensure_llm_router_target_allowed(state: &AppState) -> Result<(), String> {
|
||||
ensure_llm_router_url_allowed(state)?;
|
||||
if state.config.llm_router_model.trim() != OFFICIAL_LLM_ROUTER_MODEL {
|
||||
if state.config.is_production() {
|
||||
return Err("生产环境 LLM Router 必须使用官方固定模型".to_string());
|
||||
}
|
||||
if state.config.llm_router_base_url.trim_end_matches('/') == OFFICIAL_LLM_ROUTER_BASE_URL {
|
||||
return Err("LLM Router 必须使用官方固定模型".to_string());
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn router_username_for_owner(owner_user_id: &str) -> String {
|
||||
// New API 的 User.Username 校验上限是 20 个字符。保留可读前缀后只
|
||||
// 能放 11 个字符;使用完整 owner id 做 SHA-256,再编码成 8 字节的
|
||||
|
||||
@@ -37,18 +37,24 @@ mod model_catalog_tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn public_catalog_only_exposes_alias_and_stable_id() {
|
||||
let mut catalog = module_runtime::AgcModelCatalog::default();
|
||||
catalog.revision = 7;
|
||||
fn public_catalog_exposes_stable_id_and_upstream_alias() {
|
||||
let mut catalog = module_runtime::AgcModelCatalog::from_upstream_models(
|
||||
vec!["gpt-5.6-sol".to_string(), "gpt-5.6-terra".to_string()],
|
||||
7,
|
||||
)
|
||||
.expect("catalog should build");
|
||||
catalog.models[1].enabled = false;
|
||||
let payload = serde_json::to_value(public_model_catalog(catalog)).unwrap();
|
||||
// 客户端拿到稳定标识 + 别名(别名就是上游原始模型名),实际模型名不下发。
|
||||
assert_eq!(
|
||||
payload["models"],
|
||||
json!([{"id": "quality", "displayName": "高质量"}])
|
||||
json!([{"id": "gpt-5-6-sol", "displayName": "gpt-5.6-sol"}])
|
||||
);
|
||||
assert_eq!(payload["defaultModelId"], "quality");
|
||||
assert_eq!(payload["defaultModelId"], "gpt-5-6-sol");
|
||||
assert_eq!(payload["revision"], json!(7));
|
||||
assert!(!payload.to_string().contains("gpt-"));
|
||||
assert!(payload.get("defaultModel").is_none());
|
||||
assert!(payload["models"][0].get("enabled").is_none());
|
||||
assert!(payload["models"][0].get("modelId").is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -194,6 +200,7 @@ fn public_model_catalog(catalog: module_runtime::AgcModelCatalog) -> LlmModelsRe
|
||||
.filter(|model| model.enabled)
|
||||
.map(|model| LlmModelSummary {
|
||||
id: model.id,
|
||||
// 初始目录里别名就是上游原始模型名(不再填“高质量/快速”这类人工别名)。
|
||||
display_name: model.alias,
|
||||
})
|
||||
.collect(),
|
||||
@@ -201,6 +208,29 @@ fn public_model_catalog(catalog: module_runtime::AgcModelCatalog) -> LlmModelsRe
|
||||
}
|
||||
}
|
||||
|
||||
/// 测试用目录:两项。上游模型名带 `.`,标识是它的 slug —— 既验证「客户端只回传目录标识」,
|
||||
/// 也验证标识 → 实际模型名的映射;默认项是排序后的第一项,`TEST_AGC_MODEL_ID` 不是默认项。
|
||||
#[cfg(test)]
|
||||
pub(crate) const TEST_AGC_MODEL_ID: &str = "test-router-model";
|
||||
#[cfg(test)]
|
||||
pub(crate) const TEST_AGC_MODEL_MODEL_ID: &str = "test-router.model";
|
||||
#[cfg(test)]
|
||||
pub(crate) const TEST_AGC_MODEL_DEFAULT_ID: &str = "test-router-default";
|
||||
#[cfg(test)]
|
||||
pub(crate) const TEST_AGC_MODEL_DEFAULT_MODEL_ID: &str = "test-router.default";
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn test_agc_model_catalog() -> module_runtime::AgcModelCatalog {
|
||||
module_runtime::AgcModelCatalog::from_upstream_models(
|
||||
vec![
|
||||
TEST_AGC_MODEL_DEFAULT_MODEL_ID.to_string(),
|
||||
TEST_AGC_MODEL_MODEL_ID.to_string(),
|
||||
],
|
||||
0,
|
||||
)
|
||||
.expect("test catalog should build")
|
||||
}
|
||||
|
||||
async fn load_llm_catalog(
|
||||
state: &AppState,
|
||||
owner: &str,
|
||||
@@ -211,7 +241,7 @@ async fn load_llm_catalog(
|
||||
.expect("fixture lock")
|
||||
.contains_key(owner)
|
||||
{
|
||||
return Ok(module_runtime::AgcModelCatalog::default());
|
||||
return Ok(test_agc_model_catalog());
|
||||
}
|
||||
let _ = owner;
|
||||
crate::agc_models::load_catalog(state).await
|
||||
@@ -283,9 +313,8 @@ pub async fn proxy_llm_responses(
|
||||
] {
|
||||
object.remove(field);
|
||||
}
|
||||
// The AGC client may select a model from the server-provided Router
|
||||
// directory. Older callers without the reserved marker remain pinned to
|
||||
// the official default model.
|
||||
// AGC 客户端可以在服务端目录内选择模型;`model` 就是上游原始模型名。
|
||||
// 老客户端存的历史稳定标识与目录外模型一律拒绝,不回退其它模型。
|
||||
let agc_client = headers
|
||||
.get("x-genarrative-client")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
@@ -293,16 +322,13 @@ pub async fn proxy_llm_responses(
|
||||
let catalog = load_llm_catalog(&state, authenticated.claims().user_id())
|
||||
.await
|
||||
.map_err(|error| llm_error_response(&request_context, error))?;
|
||||
let selected_id = if agc_client {
|
||||
requested_model
|
||||
.as_deref()
|
||||
.filter(|id| *id != "platform-default")
|
||||
let requested_model = if agc_client {
|
||||
requested_model.as_deref()
|
||||
} else {
|
||||
None
|
||||
}
|
||||
.unwrap_or(&catalog.default_model_id);
|
||||
};
|
||||
let selected_model = catalog
|
||||
.resolve(selected_id)
|
||||
.resolve_requested(requested_model)
|
||||
.map_err(|message| {
|
||||
llm_error_response(
|
||||
&request_context,
|
||||
@@ -847,7 +873,7 @@ async fn resolve_llm_router_client(
|
||||
let catalog = load_llm_catalog(state, owner_user_id)
|
||||
.await
|
||||
.map_err(|_| "模型目录暂不可用".to_string())?;
|
||||
let model = catalog.resolve(&catalog.default_model_id)?;
|
||||
let model = catalog.resolve_requested(None)?;
|
||||
let config = platform_llm::LlmConfig::new(
|
||||
platform_llm::LlmProvider::OpenAiCompatible,
|
||||
base_url.to_string(),
|
||||
@@ -1304,11 +1330,14 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn llm_responses_proxy_forces_official_model_and_keeps_router_key_server_side() {
|
||||
async fn llm_responses_without_agc_marker_uses_catalog_default_and_keeps_router_key_server_side()
|
||||
{
|
||||
let (server_url, captured_request) = spawn_capturing_mock_server(MockResponse {
|
||||
status_line: "200 OK",
|
||||
content_type: "application/json; charset=utf-8",
|
||||
body: r#"{"id":"resp_proxy_01","model":"gpt-6-astra","output":[]}"#.to_string(),
|
||||
body: format!(
|
||||
r#"{{"id":"resp_proxy_01","model":"{TEST_AGC_MODEL_DEFAULT_MODEL_ID}","output":[]}}"#
|
||||
),
|
||||
extra_headers: Vec::new(),
|
||||
});
|
||||
let (state, user_id) = seed_authenticated_state(AppConfig {
|
||||
@@ -1373,12 +1402,64 @@ mod tests {
|
||||
.expect("upstream request body");
|
||||
let upstream_payload: Value =
|
||||
serde_json::from_str(upstream_body).expect("upstream body should be json");
|
||||
assert_eq!(upstream_payload["model"], "gpt-6-astra");
|
||||
assert_eq!(upstream_payload["model"], TEST_AGC_MODEL_DEFAULT_MODEL_ID);
|
||||
assert_ne!(upstream_payload["model"], "client-must-not-control");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn llm_responses_rejects_upstream_names_and_unknown_catalog_ids() {
|
||||
async fn llm_responses_forwards_catalog_model_selected_by_agc_client() {
|
||||
let (server_url, captured_request) = spawn_capturing_mock_server(MockResponse {
|
||||
status_line: "200 OK",
|
||||
content_type: "application/json; charset=utf-8",
|
||||
body: format!(
|
||||
r#"{{"id":"resp_proxy_02","model":"{TEST_AGC_MODEL_MODEL_ID}","output":[]}}"#
|
||||
),
|
||||
extra_headers: Vec::new(),
|
||||
});
|
||||
let (state, user_id) = seed_authenticated_state(AppConfig {
|
||||
llm_router_base_url: server_url.clone(),
|
||||
llm_router_api_key_encryption_secret: Some("fixture-encryption-secret".to_string()),
|
||||
..AppConfig::default()
|
||||
})
|
||||
.await;
|
||||
install_test_provisioned_router_credential(&user_id, server_url, "fixture-router-key");
|
||||
let token = issue_access_token(&state, &user_id);
|
||||
let app = build_router(state);
|
||||
|
||||
let response = app
|
||||
.oneshot(
|
||||
Request::builder()
|
||||
.method("POST")
|
||||
.uri("/api/llm/responses")
|
||||
.header("authorization", format!("Bearer {token}"))
|
||||
.header("x-genarrative-client", "agc")
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from(
|
||||
json!({"model": TEST_AGC_MODEL_ID, "input": "hello"}).to_string(),
|
||||
))
|
||||
.expect("request should build"),
|
||||
)
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
|
||||
let upstream_request = captured_request
|
||||
.lock()
|
||||
.expect("captured request lock")
|
||||
.clone()
|
||||
.expect("mock server should capture upstream request");
|
||||
let (_, upstream_body) = upstream_request
|
||||
.split_once("\r\n\r\n")
|
||||
.expect("upstream request body");
|
||||
let upstream_payload: Value =
|
||||
serde_json::from_str(upstream_body).expect("upstream body should be json");
|
||||
// 客户端只能回传目录标识,服务端映射成上游实际模型名;默认项不参与。
|
||||
assert_eq!(upstream_payload["model"], TEST_AGC_MODEL_MODEL_ID);
|
||||
assert_ne!(upstream_payload["model"], TEST_AGC_MODEL_DEFAULT_MODEL_ID);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn llm_responses_rejects_models_outside_catalog() {
|
||||
let (state, user_id) = seed_authenticated_state(AppConfig::default()).await;
|
||||
install_test_provisioned_router_credential(
|
||||
&user_id,
|
||||
@@ -1387,7 +1468,13 @@ mod tests {
|
||||
);
|
||||
let token = issue_access_token(&state, &user_id);
|
||||
let app = build_router(state);
|
||||
for model in ["gpt-6-astra", "unlisted"] {
|
||||
// 历史稳定标识、目录外名称、以及「直接拿上游实际模型名当标识」都必须拒绝。
|
||||
for model in [
|
||||
"quality",
|
||||
"gpt-6-astra",
|
||||
"unlisted",
|
||||
TEST_AGC_MODEL_MODEL_ID,
|
||||
] {
|
||||
let response = app
|
||||
.clone()
|
||||
.oneshot(
|
||||
|
||||
@@ -500,6 +500,10 @@ fn should_initialize_editor_generation_pricing_for_startup(process_role: Process
|
||||
process_role.runs_http()
|
||||
}
|
||||
|
||||
fn should_initialize_agc_model_catalog_for_startup(process_role: ProcessRole) -> bool {
|
||||
process_role.runs_http()
|
||||
}
|
||||
|
||||
async fn run_http_role(config: AppConfig) -> Result<(), io::Error> {
|
||||
let bind_address = config.bind_socket_addr();
|
||||
let listen_backlog = config.listen_backlog;
|
||||
@@ -764,6 +768,16 @@ async fn try_restore_app_state_for_startup(
|
||||
))
|
||||
})?;
|
||||
}
|
||||
// AGC 模型目录只来自上游同步或后台保存;这里同步失败不阻塞启动,由下一次启动重试,
|
||||
// 未初始化期间 AGC 相关接口失败关闭。
|
||||
if should_initialize_agc_model_catalog_for_startup(process_role) {
|
||||
if let Err(error) = crate::agc_models::ensure_agc_model_catalog_initialized(&state).await {
|
||||
error!(
|
||||
error = %error,
|
||||
"AGC 模型目录未初始化:本次启动未从上游同步到模型列表,AGC 目录与对话接口将失败关闭,下次启动会重试"
|
||||
);
|
||||
}
|
||||
}
|
||||
Ok(state)
|
||||
}
|
||||
|
||||
|
||||
@@ -1,9 +1,18 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashSet;
|
||||
|
||||
/// 目录 revision 乐观锁冲突。
|
||||
pub const AGC_MODEL_CATALOG_CONFLICT: &str = "AGC_MODEL_CATALOG_CONFLICT";
|
||||
/// 目录尚未初始化:SpacetimeDB 缺行,或存量内容与当前定义不符。
|
||||
pub const AGC_MODEL_CATALOG_NOT_INITIALIZED: &str = "AGC_MODEL_CATALOG_NOT_INITIALIZED";
|
||||
/// 客户端未显式选择模型时使用的占位标识。
|
||||
pub const AGC_MODEL_PLATFORM_DEFAULT: &str = "platform-default";
|
||||
/// 模型标识的长度上限,与客户端 `select_game_creator_model` 的校验保持一致。
|
||||
pub const AGC_MODEL_ID_MAX_BYTES: usize = 64;
|
||||
/// 目录项数上限,与后台「AGC 模型」页的新增上限保持一致。
|
||||
pub const AGC_MODEL_CATALOG_MAX_MODELS: usize = 32;
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct AgcModel {
|
||||
pub id: String,
|
||||
@@ -12,7 +21,7 @@ pub struct AgcModel {
|
||||
pub enabled: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "camelCase", deny_unknown_fields)]
|
||||
pub struct AgcModelCatalog {
|
||||
pub revision: u64,
|
||||
@@ -20,40 +29,62 @@ pub struct AgcModelCatalog {
|
||||
pub models: Vec<AgcModel>,
|
||||
}
|
||||
|
||||
impl Default for AgcModelCatalog {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
revision: 0,
|
||||
default_model_id: "quality".into(),
|
||||
models: vec![
|
||||
AgcModel {
|
||||
id: "quality".into(),
|
||||
alias: "高质量".into(),
|
||||
model_id: "gpt-6-astra".into(),
|
||||
enabled: true,
|
||||
},
|
||||
AgcModel {
|
||||
id: "fast".into(),
|
||||
alias: "快速".into(),
|
||||
model_id: "gpt-5.6-luna".into(),
|
||||
enabled: true,
|
||||
},
|
||||
],
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl AgcModelCatalog {
|
||||
/// 按上游模型列表生成目录:`modelId` 是上游原始模型名,`alias` 也直接用原名
|
||||
/// (不再填「高质量/快速」这类人工别名),`id` 是模型名的稳定 slug。
|
||||
///
|
||||
/// 上游返回顺序不稳定,所以先按原始模型名排序再生成,重复同步得到一致的目录与默认项。
|
||||
pub fn from_upstream_models(
|
||||
models: impl IntoIterator<Item = String>,
|
||||
revision: u64,
|
||||
) -> Result<Self, String> {
|
||||
let mut model_names = models
|
||||
.into_iter()
|
||||
.map(|model| model.trim().to_string())
|
||||
.filter(|model| !model.is_empty())
|
||||
.collect::<Vec<_>>();
|
||||
model_names.sort();
|
||||
model_names.dedup();
|
||||
if model_names.is_empty() {
|
||||
return Err("上游模型列表为空".into());
|
||||
}
|
||||
|
||||
let mut used_ids = HashSet::new();
|
||||
let mut entries = Vec::with_capacity(model_names.len());
|
||||
for model_id in model_names {
|
||||
let id = unique_model_id(&model_id, &mut used_ids);
|
||||
entries.push(AgcModel {
|
||||
id,
|
||||
alias: model_id.clone(),
|
||||
model_id,
|
||||
enabled: true,
|
||||
});
|
||||
}
|
||||
let default_model_id = entries
|
||||
.first()
|
||||
.map(|entry| entry.id.clone())
|
||||
.ok_or_else(|| "上游模型列表为空".to_string())?;
|
||||
let catalog = Self {
|
||||
revision,
|
||||
default_model_id,
|
||||
models: entries,
|
||||
};
|
||||
catalog.validate()?;
|
||||
Ok(catalog)
|
||||
}
|
||||
|
||||
pub fn validate(&self) -> Result<(), String> {
|
||||
if self.models.is_empty() || self.models.len() > 32 {
|
||||
return Err("模型列表必须包含 1 至 32 项".into());
|
||||
if self.models.is_empty() || self.models.len() > AGC_MODEL_CATALOG_MAX_MODELS {
|
||||
return Err(format!(
|
||||
"模型列表必须包含 1 至 {AGC_MODEL_CATALOG_MAX_MODELS} 项"
|
||||
));
|
||||
}
|
||||
let mut ids = HashSet::new();
|
||||
let mut aliases = HashSet::new();
|
||||
for model in &self.models {
|
||||
if model.id.is_empty()
|
||||
|| model.id == "platform-default"
|
||||
|| model.id.len() > 64
|
||||
|| model.id == AGC_MODEL_PLATFORM_DEFAULT
|
||||
|| model.id.len() > AGC_MODEL_ID_MAX_BYTES
|
||||
|| !model
|
||||
.id
|
||||
.bytes()
|
||||
@@ -88,31 +119,202 @@ impl AgcModelCatalog {
|
||||
.map(|m| m.model_id.as_str())
|
||||
.ok_or_else(|| "所选模型不可用,请刷新模型列表".into())
|
||||
}
|
||||
|
||||
/// 请求侧解析:显式选择的标识按目录校验,未选择或占位标识使用默认项。
|
||||
pub fn resolve_requested(&self, requested: Option<&str>) -> Result<&str, String> {
|
||||
let requested = requested
|
||||
.map(str::trim)
|
||||
.filter(|id| !id.is_empty() && *id != AGC_MODEL_PLATFORM_DEFAULT);
|
||||
match requested {
|
||||
Some(id) => self.resolve(id),
|
||||
None => self.resolve(&self.default_model_id),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 由上游模型名生成稳定标识:只保留小写字母、数字、连字符与下划线,其余字符折叠成 `-`。
|
||||
fn agc_model_id_from_name(model_name: &str) -> String {
|
||||
let mut id = String::new();
|
||||
let mut separator_pending = false;
|
||||
for value in model_name.chars() {
|
||||
let lowered = value.to_ascii_lowercase();
|
||||
if lowered.is_ascii_alphanumeric() || lowered == '_' {
|
||||
if separator_pending && !id.is_empty() {
|
||||
id.push('-');
|
||||
}
|
||||
separator_pending = false;
|
||||
id.push(lowered);
|
||||
} else {
|
||||
separator_pending = true;
|
||||
}
|
||||
}
|
||||
id
|
||||
}
|
||||
|
||||
/// 生成在本次目录内唯一的标识:同名 slug 追加 `-2`/`-3`,并保证不超过长度上限。
|
||||
fn unique_model_id(model_name: &str, used_ids: &mut HashSet<String>) -> String {
|
||||
let slug = agc_model_id_from_name(model_name);
|
||||
let slug = if slug.is_empty() {
|
||||
"model".to_string()
|
||||
} else {
|
||||
slug
|
||||
};
|
||||
// 预留后缀空间(`-` 加最多两位序号)后截断,保证候选标识仍在长度上限内。
|
||||
let base = slug
|
||||
.char_indices()
|
||||
.take_while(|(index, _)| *index < AGC_MODEL_ID_MAX_BYTES - 3)
|
||||
.map(|(_, value)| value)
|
||||
.collect::<String>();
|
||||
let base = base.trim_end_matches('-').to_string();
|
||||
let base = if base.is_empty() {
|
||||
"model".to_string()
|
||||
} else {
|
||||
base
|
||||
};
|
||||
|
||||
let mut candidate = base.clone();
|
||||
let mut suffix = 2;
|
||||
while !used_ids.insert(candidate.clone()) {
|
||||
candidate = format!("{base}-{suffix}");
|
||||
suffix += 1;
|
||||
}
|
||||
candidate
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn upstream(models: &[&str]) -> Vec<String> {
|
||||
models.iter().map(|model| (*model).to_string()).collect()
|
||||
}
|
||||
|
||||
fn model(id: &str, alias: &str, model_id: &str) -> AgcModel {
|
||||
AgcModel {
|
||||
id: id.into(),
|
||||
alias: alias.into(),
|
||||
model_id: model_id.into(),
|
||||
enabled: true,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn catalog_builds_from_upstream_models_with_stable_ids() {
|
||||
let catalog = AgcModelCatalog::from_upstream_models(
|
||||
upstream(&[
|
||||
" qwen3.8-flash ",
|
||||
"glm-5.3",
|
||||
"qwen3.8-flash",
|
||||
"deepseek-v4-pro",
|
||||
"",
|
||||
"vendor/model.v1:latest",
|
||||
]),
|
||||
3,
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(catalog.revision, 3);
|
||||
// 默认项是排序后第一项,与上游返回顺序无关。
|
||||
assert_eq!(catalog.default_model_id, "deepseek-v4-pro");
|
||||
assert_eq!(
|
||||
catalog.models,
|
||||
vec![
|
||||
model("deepseek-v4-pro", "deepseek-v4-pro", "deepseek-v4-pro"),
|
||||
model("glm-5-3", "glm-5.3", "glm-5.3"),
|
||||
model("qwen3-8-flash", "qwen3.8-flash", "qwen3.8-flash"),
|
||||
model(
|
||||
"vendor-model-v1-latest",
|
||||
"vendor/model.v1:latest",
|
||||
"vendor/model.v1:latest"
|
||||
),
|
||||
]
|
||||
);
|
||||
assert!(catalog.validate().is_ok());
|
||||
// 同一模型集合重复生成结果一致。
|
||||
assert_eq!(
|
||||
AgcModelCatalog::from_upstream_models(
|
||||
upstream(&[
|
||||
"vendor/model.v1:latest",
|
||||
"deepseek-v4-pro",
|
||||
"glm-5.3",
|
||||
"qwen3.8-flash",
|
||||
]),
|
||||
3
|
||||
)
|
||||
.unwrap(),
|
||||
catalog
|
||||
);
|
||||
assert!(AgcModelCatalog::from_upstream_models(upstream(&["", " "]), 0).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn catalog_keeps_ids_unique_and_within_client_contract() {
|
||||
// 不同模型名折叠成同一个 slug 时按排序追加序号,且标识始终符合客户端校验。
|
||||
let catalog = AgcModelCatalog::from_upstream_models(
|
||||
upstream(&["GLM-5.3", "glm/5.3", "glm_5.3", "模型名"]),
|
||||
0,
|
||||
)
|
||||
.unwrap();
|
||||
let ids = catalog
|
||||
.models
|
||||
.iter()
|
||||
.map(|entry| entry.id.as_str())
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(ids, vec!["glm-5-3", "glm-5-3-2", "glm_5-3", "model"]);
|
||||
for entry in &catalog.models {
|
||||
assert!(entry.id.len() <= AGC_MODEL_ID_MAX_BYTES);
|
||||
assert!(
|
||||
entry
|
||||
.id
|
||||
.bytes()
|
||||
.all(|c| c.is_ascii_alphanumeric() || c == b'-' || c == b'_')
|
||||
);
|
||||
}
|
||||
assert!(catalog.validate().is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn catalog_maps_only_enabled_ids() {
|
||||
let mut catalog = AgcModelCatalog::default();
|
||||
let mut catalog =
|
||||
AgcModelCatalog::from_upstream_models(upstream(&["model-a", "model-b"]), 0).unwrap();
|
||||
assert!(catalog.validate().is_ok());
|
||||
assert_eq!(catalog.resolve("quality").unwrap(), "gpt-6-astra");
|
||||
assert!(catalog.resolve("gpt-6-astra").is_err());
|
||||
assert!(catalog.resolve("unknown").is_err());
|
||||
assert_eq!(catalog.resolve("model-a").unwrap(), "model-a");
|
||||
// 客户端不能直接指定实际模型名,只能回传目录标识。
|
||||
assert!(catalog.resolve("model-c").is_err());
|
||||
assert_eq!(catalog.resolve_requested(None).unwrap(), "model-a");
|
||||
assert_eq!(
|
||||
catalog
|
||||
.resolve_requested(Some(AGC_MODEL_PLATFORM_DEFAULT))
|
||||
.unwrap(),
|
||||
"model-a"
|
||||
);
|
||||
catalog.models[0].enabled = false;
|
||||
assert!(catalog.resolve("quality").is_err());
|
||||
assert!(catalog.resolve("model-a").is_err());
|
||||
assert!(catalog.validate().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn catalog_rejects_duplicate_aliases_and_ids() {
|
||||
let mut catalog = AgcModelCatalog::default();
|
||||
let mut catalog =
|
||||
AgcModelCatalog::from_upstream_models(upstream(&["model-a", "model-b"]), 0).unwrap();
|
||||
catalog.models[1].alias = catalog.models[0].alias.clone();
|
||||
assert!(catalog.validate().is_err());
|
||||
catalog.models[1].alias = "快速".into();
|
||||
catalog.models[1].alias = "model-b".into();
|
||||
catalog.models[1].id = catalog.models[0].id.clone();
|
||||
assert!(catalog.validate().is_err());
|
||||
catalog.models[1].id = "model-b".into();
|
||||
catalog.models[1].id = AGC_MODEL_PLATFORM_DEFAULT.into();
|
||||
assert!(catalog.validate().is_err());
|
||||
catalog.models[1].id = "model-b".into();
|
||||
catalog.models[1].model_id = "".into();
|
||||
assert!(catalog.validate().is_err());
|
||||
|
||||
let too_many = (0..AGC_MODEL_CATALOG_MAX_MODELS + 1)
|
||||
.map(|index| format!("model-{index}"))
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
AgcModelCatalog::from_upstream_models(too_many, 0).unwrap_err(),
|
||||
format!("模型列表必须包含 1 至 {AGC_MODEL_CATALOG_MAX_MODELS} 项")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -16,16 +16,14 @@ pub fn read_agc_model_catalog(ctx: &mut ProcedureContext) -> Result<String, Stri
|
||||
crate::editor_project_storage::require_editor_generation_runtime_service_identity(
|
||||
tx, caller,
|
||||
)?;
|
||||
Ok(tx
|
||||
.db
|
||||
// 目录只有一份事实来源:上游同步或后台保存。缺行时不得返回任何内置目录,
|
||||
// 由 api-server 按“未初始化”处理并触发启动期同步。
|
||||
tx.db
|
||||
.agc_model_catalog()
|
||||
.id()
|
||||
.find(0)
|
||||
.map(|row| row.catalog_json)
|
||||
.unwrap_or_else(|| {
|
||||
serde_json::to_string(&module_runtime::AgcModelCatalog::default())
|
||||
.expect("default catalog")
|
||||
}))
|
||||
.ok_or_else(|| module_runtime::AGC_MODEL_CATALOG_NOT_INITIALIZED.to_string())
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user