From 4234a21eed98249917951229d4edb68dc715d85f Mon Sep 17 00:00:00 2001 From: Suzumiya Date: Thu, 24 Sep 2026 16:14:35 +0800 Subject: [PATCH] =?UTF-8?q?=E6=9C=8D=E5=8A=A1=E7=AB=AF=EF=BC=9AAGC=20?= =?UTF-8?q?=E6=A8=A1=E5=9E=8B=E7=9B=AE=E5=BD=95=E7=BC=BA=E9=85=8D=E7=BD=AE?= =?UTF-8?q?=E6=97=B6=E6=94=B9=E4=B8=BA=E5=90=AF=E5=8A=A8=E6=9C=9F=E4=BB=8E?= =?UTF-8?q?=E4=B8=8A=E6=B8=B8=E5=90=8C=E6=AD=A5=EF=BC=8C=E4=BF=9D=E7=95=99?= =?UTF-8?q?=E5=8E=9F=E7=9B=AE=E5=BD=95=E6=A0=BC=E5=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 流式上限、禁止重定向、不带凭据;目标校验只校验地址,不再绑定已下线的固定模型哨兵;写回冲突后重读校验既有目录 --- server-rs/crates/api-server/src/agc_models.rs | 377 +++++++++++++++++- .../api-server/src/external_api_keys.rs | 29 +- server-rs/crates/api-server/src/llm/mod.rs | 133 ++++-- server-rs/crates/api-server/src/main.rs | 14 + .../crates/module-runtime/src/agc_models.rs | 274 +++++++++++-- .../crates/spacetime-module/src/agc_models.rs | 10 +- 6 files changed, 749 insertions(+), 88 deletions(-) diff --git a/server-rs/crates/api-server/src/agc_models.rs b/server-rs/crates/api-server/src/agc_models.rs index dbd83c0cc..f0634b0df 100644 --- a/server-rs/crates/api-server/src/agc_models.rs +++ b/server-rs/crates/api-server/src/agc_models.rs @@ -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 { - 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 { + let catalog: AgcModelCatalog = + serde_json::from_str(json).map_err(|_| "模型目录格式无效".to_string())?; + catalog.validate()?; Ok(catalog) } +async fn read_stored_catalog(state: &AppState) -> Result, 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 { + serde_json::from_str::(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, 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, 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, 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::>(); + 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, Extension(context): Extension, @@ -44,6 +226,8 @@ pub async fn admin_save_agc_models( Extension(_admin): Extension, Json(payload): Json, ) -> Result, 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>>, + _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" + ); + } +} diff --git a/server-rs/crates/api-server/src/external_api_keys.rs b/server-rs/crates/api-server/src/external_api_keys.rs index 3e70d8e5f..cef612c13 100644 --- a/server-rs/crates/api-server/src/external_api_keys.rs +++ b/server-rs/crates/api-server/src/external_api_keys.rs @@ -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 { +pub(crate) fn router_control_origin(base_url: &str) -> Result { 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 { 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 字节的 diff --git a/server-rs/crates/api-server/src/llm/mod.rs b/server-rs/crates/api-server/src/llm/mod.rs index 32ee19402..025823e3f 100644 --- a/server-rs/crates/api-server/src/llm/mod.rs +++ b/server-rs/crates/api-server/src/llm/mod.rs @@ -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( diff --git a/server-rs/crates/api-server/src/main.rs b/server-rs/crates/api-server/src/main.rs index 241551c49..8515721f5 100644 --- a/server-rs/crates/api-server/src/main.rs +++ b/server-rs/crates/api-server/src/main.rs @@ -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) } diff --git a/server-rs/crates/module-runtime/src/agc_models.rs b/server-rs/crates/module-runtime/src/agc_models.rs index 74c9f96ec..216dc1aaf 100644 --- a/server-rs/crates/module-runtime/src/agc_models.rs +++ b/server-rs/crates/module-runtime/src/agc_models.rs @@ -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, } -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, + revision: u64, + ) -> Result { + let mut model_names = models + .into_iter() + .map(|model| model.trim().to_string()) + .filter(|model| !model.is_empty()) + .collect::>(); + 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 { + 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::(); + 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 { + 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::>(); + 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::>(); + assert_eq!( + AgcModelCatalog::from_upstream_models(too_many, 0).unwrap_err(), + format!("模型列表必须包含 1 至 {AGC_MODEL_CATALOG_MAX_MODELS} 项") + ); } } diff --git a/server-rs/crates/spacetime-module/src/agc_models.rs b/server-rs/crates/spacetime-module/src/agc_models.rs index 5fba815ab..03bbb3536 100644 --- a/server-rs/crates/spacetime-module/src/agc_models.rs +++ b/server-rs/crates/spacetime-module/src/agc_models.rs @@ -16,16 +16,14 @@ pub fn read_agc_model_catalog(ctx: &mut ProcedureContext) -> Result