From a02a3187580dd0b86b1c64e150698740d9cf91e6 Mon Sep 17 00:00:00 2001 From: kdletters Date: Thu, 27 Aug 2026 18:09:44 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E7=94=9F=E4=BA=A7=E5=86=85?= =?UTF-8?q?=E5=AD=98=E5=B7=A5=E4=BD=9C=E9=9B=86=E5=9B=9E=E6=94=B6=E4=B8=8E?= =?UTF-8?q?=20OOM=20=E6=81=A2=E5=A4=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 回收持续队列中的已完成 worker 句柄并收紧脱管许可生命周期 统一 AI 任务终态写入的文本、结构化输出和 warning 内存上限 限制认证投影恢复的 retained refresh session 数量 为备份停库增加 marker 与 systemd OOM 兜底恢复 补充回归测试并同步后端、运维和项目决策文档 --- .../genarrative-database-backup.service | 5 + .../shared-memory/decision-log.md | 7 + ...】server-rs与SpacetimeDB数据契约-2026-05-15.md | 2 + ...发运维】本地开发验证与生产运维-2026-05-15.md | 2 + scripts/check-database-backup-to-oss.mjs | 26 ++ scripts/check-production-ops-guardrails.mjs | 15 +- scripts/database-backup-to-oss.mjs | 42 +++- .../src/external_generation_worker.rs | 38 ++- .../crates/module-ai/src/application/store.rs | 224 ++++++++++++++---- server-rs/crates/module-ai/src/tests.rs | 144 +++++++++++ server-rs/crates/module-auth/src/lib.rs | 42 ++++ 11 files changed, 498 insertions(+), 49 deletions(-) diff --git a/deploy/systemd/genarrative-database-backup.service b/deploy/systemd/genarrative-database-backup.service index 4ccc00922..8a7d95535 100644 --- a/deploy/systemd/genarrative-database-backup.service +++ b/deploy/systemd/genarrative-database-backup.service @@ -13,10 +13,15 @@ ExecStart=/usr/bin/node -- /opt/genarrative/current/scripts/database-backup-to-o # 备份脚本必须受独立内存上限保护,不能因目录扫描异常拖垮整台 release 主机。 Environment=NODE_OPTIONS=--max-old-space-size=768 +Environment=GENARRATIVE_DATABASE_BACKUP_STOP_MARKER=/var/lib/genarrative/database-backups/.spacetimedb-stopped MemoryHigh=768M MemoryMax=1G OOMPolicy=stop +# 主进程可能在停库后被 MemoryMax/OOMPolicy 强制终止,JS finally 无法执行; +# 仅当备份脚本留下停库 marker 且本次 service 非正常成功时,由 systemd 兜底恢复全部依赖服务。 +ExecStopPost=/bin/sh -c 'if [ "${SERVICE_RESULT}" != "success" ] && [ -f "${GENARRATIVE_DATABASE_BACKUP_STOP_MARKER}" ]; then systemctl start spacetimedb.service; systemctl restart genarrative-api.service; systemctl restart genarrative-external-generation-worker@1.service; systemctl restart genarrative-external-generation-controller.service; if systemctl is-active --quiet spacetimedb.service && systemctl is-active --quiet genarrative-api.service && systemctl is-active --quiet genarrative-external-generation-worker@1.service && systemctl is-active --quiet genarrative-external-generation-controller.service; then rm -f "${GENARRATIVE_DATABASE_BACKUP_STOP_MARKER}"; fi; fi' + # 备份需要停止 / 启动 spacetimedb.service,并读取 /stdb、写入 /var/lib/genarrative/database-backups。 # 停止 SpacetimeDB 会连带停止 Requires 它的 API / worker / controller,冷备份后必须显式拉起。 PrivateTmp=true diff --git a/docs/project-memory/shared-memory/decision-log.md b/docs/project-memory/shared-memory/decision-log.md index 614376650..d57a6ae4a 100644 --- a/docs/project-memory/shared-memory/decision-log.md +++ b/docs/project-memory/shared-memory/decision-log.md @@ -7763,3 +7763,10 @@ CI 上 `background_agent_runtime_recovers_stale_running_before_pending_task` 在 - 决策:`.codex/skills/` 下的 SpacetimeDB 指导收敛为单一 `.codex/skills/genarrative-spacetimedb/SKILL.md`。官方 `spacetimedb` 插件负责通用 concepts、Rust server、CLI、TypeScript client 和 MCP 知识;项目 skill 只保留 Genarrative 的架构边界、schema/migration 门禁、目标 server 安全规则、运行时排障和验证路径。 - 路由:涉及 SpacetimeDB 的任务统一先读取项目适配 skill,再按需读取 `spacetimedb:concepts`、`spacetimedb:rust-server`、`spacetimedb:cli`、`spacetimedb:typescript-client` 或 `spacetimedb:mcp`。插件通用示例不得覆盖项目禁止 `maincloud`、禁止人工 `spacetime --root-dir`、显式 server 和后端分层等约束。 - 安装:团队环境缺少插件时使用 `codex plugin marketplace add clockworklabs/SpacetimeDB --sparse .agents --sparse codex-plugin` 和 `codex plugin add spacetimedb\@spacetimedb-plugins`;个人配置、缓存和凭据不进入仓库。 + +## 2026-08-27 release 内存增长修复与备份 OOM 恢复兜底 + +- 决策:外部生成 worker 每轮主动 `try_join_next` 回收已完成 `JoinHandle`,避免持续有队列任务时只归还 semaphore permit 却让 `JoinSet` 句柄集合无界增长;超时脱管任务在 abort 后等待句柄结束,执行许可保持到 work 真正结束或被取消。 +- 决策:`module-ai` 的阶段终态写入与流式增量统一受文本、结构化 JSON、warning 和全局 retained 工作集上限约束;认证投影恢复对过滤后的 refresh session 重新计数,超过 8192 条直接拒绝启动恢复,避免超限快照灌入内存。 +- 决策:备份脚本停库前写入受保护 `.spacetimedb-stopped` marker,正常恢复完成后清理;systemd 备份 service 通过 `MemoryHigh/MemoryMax/OOMPolicy` 和 `ExecStopPost` 在 Node OOM kill、无法执行 JS finally 时兜底拉起 SpacetimeDB、API、worker、controller,恢复不完整则保留 marker。 +- 验证:worker/module-auth/module-ai 定向 Rust 测试、database-backup/production-ops/encoding 门禁和 `git diff --check` 必须在提交前通过;release 现场需按 archive-full 重新 provision 并核验旧 files-history drop-in 已删除、四个服务 active。 diff --git a/docs/【后端架构】server-rs与SpacetimeDB数据契约-2026-05-15.md b/docs/【后端架构】server-rs与SpacetimeDB数据契约-2026-05-15.md index 2f10ad243..a1b22d2cc 100644 --- a/docs/【后端架构】server-rs与SpacetimeDB数据契约-2026-05-15.md +++ b/docs/【后端架构】server-rs与SpacetimeDB数据契约-2026-05-15.md @@ -411,6 +411,8 @@ Responses 的终态载荷既是工具调用的恢复源,也是正文的恢复 ### `auth_store_projection_meta` +启动投影恢复会对过滤后的 retained refresh session 重新计数;超过 8192 条时直接失败关闭并继续重试,不得把超限快照一次性灌入内存。 + - Rust 结构体:`AuthStoreProjectionMeta` - 源码:`server-rs/crates/spacetime-module/src/auth/tables.rs` diff --git a/docs/【开发运维】本地开发验证与生产运维-2026-05-15.md b/docs/【开发运维】本地开发验证与生产运维-2026-05-15.md index 133a00c5c..a344a4425 100644 --- a/docs/【开发运维】本地开发验证与生产运维-2026-05-15.md +++ b/docs/【开发运维】本地开发验证与生产运维-2026-05-15.md @@ -401,6 +401,8 @@ UI 相关修改要重点验证: ### SpacetimeDB 数据目录 OSS 备份 +脚本停库前会在固定 work-dir 写入 `.spacetimedb-stopped` marker;正常 finally 恢复 SpacetimeDB 及 `--restart-service-after` 指定的 API / worker / controller 后才清理 marker。若 Node 因 `MemoryMax` / OOM 被强制终止,systemd `ExecStopPost` 会根据仍存在的 marker 兜底恢复这些服务;恢复未全部成功时保留 marker 供后续重试。 + 数据库备份不放进 `spacetime-module` reducer / procedure:备份属于文件系统与 OSS 外部副作用,必须由运维脚本在 SpacetimeDB 宿主外执行。当前统一脚本为 `scripts/database-backup-to-oss.mjs`(npm 命令 `npm run database:backup:oss`)。默认 `--storage-format archive --mode full` 保持原有全量压缩包冷备行为;`--storage-format files` 不生成 tar.gz,而是把目录树映射成逐文件 CAS 对象与 catalog,full 重跑只上传新增或内容变化的文件,history 只处理已被最新 snapshot 完全覆盖的历史 commitlog 与旧 snapshot。`Genarrative-Server-Provision` 的 `DATABASE_BACKUP_PROFILE` 默认是 `archive-full`,继续安装每天 `03:20` 左右执行的全量冷备主 service;当前 release 只允许 `archive-full`,避免 `files-history` 在大目录上构造全量 catalog 导致 Node 内存峰值;development 才可以显式选择 `files-history`,且指定 work-dir 必须已经有与本机 database/bucket 匹配且已发布的 full baseline state: ```bash diff --git a/scripts/check-database-backup-to-oss.mjs b/scripts/check-database-backup-to-oss.mjs index 57bc09368..689295ca6 100644 --- a/scripts/check-database-backup-to-oss.mjs +++ b/scripts/check-database-backup-to-oss.mjs @@ -50,6 +50,7 @@ async function main() { assertDeferredArchiveDiscoveryIsBoundedAndDeterministic(); assertCanonicalQueryAndAuthorizationIncludeMultipartParameters(); assertInsufficientSpaceStopsBeforeServiceChanges(); + assertStopFailureRetainsRecoveryMarker(); assertArchiveFailureStillRestoresDependentServices(); await assertMultipartUploadRetriesAndVerifiesRemoteLength(); await assertUploadBandwidthLimiterSharesBudgetAndPropagatesErrors(); @@ -748,6 +749,27 @@ function assertInsufficientSpaceStopsBeforeServiceChanges() { assertFileMissing(fixture.tarLog, '空间不足时不能调用 tar。'); } +function assertStopFailureRetainsRecoveryMarker() { + const fixture = createFixture('stop-failure-marker'); + writeExecutable( + path.join(fixture.binDir, 'systemctl'), + `#!/usr/bin/env bash +printf 'systemctl %s\\n' "$*" >> "${fixture.systemctlLog}" +if [ "$1" = stop ]; then + exit 9 +fi +exit 0 +`, + ); + const result = runBackup(fixture, ['--stop-service', 'spacetimedb.service']); + + assertStatus(result, 1, '停止服务失败时备份必须失败。'); + assertTrue( + existsSync(path.join(fixture.workDir, '.spacetimedb-stopped')), + '停止服务命令失败时必须保留 marker,供 systemd ExecStopPost 兜底恢复。', + ); +} + function assertArchiveFailureStillRestoresDependentServices() { const fixture = createFixture('tar-failure'); const result = runBackup(fixture, [ @@ -776,6 +798,10 @@ function assertArchiveFailureStillRestoresDependentServices() { for (const command of expectedCommands) { assertIncludes(systemctlLog, command, `tar 失败后必须执行: ${command}`); } + assertFileMissing( + path.join(fixture.workDir, '.spacetimedb-stopped'), + '正常执行 finally 恢复全部服务后必须清理停库 marker。', + ); } async function assertMultipartUploadRetriesAndVerifiesRemoteLength() { diff --git a/scripts/check-production-ops-guardrails.mjs b/scripts/check-production-ops-guardrails.mjs index 251dcfd30..bacd4b949 100644 --- a/scripts/check-production-ops-guardrails.mjs +++ b/scripts/check-production-ops-guardrails.mjs @@ -827,12 +827,25 @@ const checks = [ reason: '备份 Node 进程必须设置独立 heap 上限,避免目录扫描异常拖垮 release 主机。', }, + { + file: 'deploy/systemd/genarrative-database-backup.service', + includes: + 'Environment=GENARRATIVE_DATABASE_BACKUP_STOP_MARKER=/var/lib/genarrative/database-backups/.spacetimedb-stopped', + reason: + '备份停库 marker 必须固定在受保护的 release work-dir,供 OOM 后 systemd 兜底恢复服务。', + }, { file: 'deploy/systemd/genarrative-database-backup.service', includes: 'MemoryMax=1G', reason: '备份 service 必须设置 systemd 内存硬上限,避免异常进程消耗整机内存。', }, + { + file: 'deploy/systemd/genarrative-database-backup.service', + includes: 'ExecStopPost=/bin/sh -c', + reason: + '备份主进程被 OOM kill 后必须由 systemd 兜底恢复停掉的 SpacetimeDB、API、worker 和 controller。', + }, { file: 'deploy/systemd/genarrative-database-backup.service', excludes: '--storage-format files', @@ -965,7 +978,7 @@ const checks = [ { file: 'scripts/database-backup-to-oss.mjs', includes: - 'restoreServicesAfterBackup({stopService, serviceStopped, restartServicesAfter})', + 'restoreServicesAfterBackup({stopService, serviceStopped, restartServicesAfter, stopMarkerPath})', reason: '生产冷备份打包失败时也必须恢复 SpacetimeDB 及依赖服务。', }, { diff --git a/scripts/database-backup-to-oss.mjs b/scripts/database-backup-to-oss.mjs index 86057f07a..f3788301a 100644 --- a/scripts/database-backup-to-oss.mjs +++ b/scripts/database-backup-to-oss.mjs @@ -35,6 +35,7 @@ const DEFAULT_LOCAL_DATA_DIR = resolve(REPO_ROOT, 'server-rs/.spacetimedb/local/ const DEFAULT_LOCAL_WORK_DIR = resolve(REPO_ROOT, 'server-rs/.data/database-backups'); const DEFAULT_PRODUCTION_DATA_DIR = '/stdb'; const DEFAULT_PRODUCTION_WORK_DIR = '/var/lib/genarrative/database-backups'; +const DEFAULT_DATABASE_BACKUP_STOP_MARKER = join(DEFAULT_PRODUCTION_WORK_DIR, '.spacetimedb-stopped'); const DEFAULT_SPACE_SAFETY_RATIO = 1.1; const DEFAULT_EXTRA_FREE_BYTES = 512 * 1024 * 1024; const OSS_ALGORITHM = 'OSS4-HMAC-SHA256'; @@ -890,11 +891,37 @@ function collectRestartServicesAfterBackup({args, env}) { return [...new Set(serviceNames.filter(Boolean))]; } -function stopServiceIfNeeded(serviceName) { +function databaseBackupStopMarkerPath(workDir) { + return resolvePath(firstNonEmpty( + process.env.GENARRATIVE_DATABASE_BACKUP_STOP_MARKER, + workDir === DEFAULT_PRODUCTION_WORK_DIR + ? DEFAULT_DATABASE_BACKUP_STOP_MARKER + : join(workDir, '.spacetimedb-stopped'), + )); +} + +function writeDatabaseBackupStopMarker(markerPath, serviceName) { + atomicWriteJson(markerPath, { + serviceName, + pid: process.pid, + stoppedAt: new Date().toISOString(), + }); +} + +function clearDatabaseBackupStopMarker(markerPath) { + if (markerPath) { + rmSync(markerPath, {force: true}); + } +} + +function stopServiceIfNeeded(serviceName, stopMarkerPath) { if (!serviceName) { return false; } console.log(`[database-backup] 停止服务以获取冷备份: ${serviceName}`); + writeDatabaseBackupStopMarker(stopMarkerPath, serviceName); + // stop 命令失败时仍保留 marker:systemd 的 ExecStopPost 需要它判断是否要 + // 兜底恢复,不能因为当前进程还能捕获异常就抹掉上一次停库证据。 runCommand('systemctl', ['stop', serviceName], {stdio: 'inherit'}); return true; } @@ -925,7 +952,7 @@ function restartServicesAfterBackup(serviceNames) { } } -function restoreServicesAfterBackup({stopService, serviceStopped, restartServicesAfter}) { +function restoreServicesAfterBackup({stopService, serviceStopped, restartServicesAfter, stopMarkerPath}) { const errors = []; try { startServiceIfNeeded(stopService, serviceStopped); @@ -940,6 +967,7 @@ function restoreServicesAfterBackup({stopService, serviceStopped, restartService if (errors.length > 0) { throw new AggregateError(errors, `恢复冷备份相关服务失败: ${errors.map((error) => error.message).join('; ')}`); } + clearDatabaseBackupStopMarker(stopMarkerPath); } function createArchive({dataDir, workDir, fileName}) { @@ -3341,12 +3369,13 @@ async function main() { } const stopService = args.stopService || firstNonEmpty(env.GENARRATIVE_DATABASE_BACKUP_STOP_SERVICE); const restartServicesAfter = collectRestartServicesAfterBackup({args, env}); + const stopMarkerPath = databaseBackupStopMarkerPath(workDir); let serviceStopped = false; let backupError = null; let restoreError = null; try { if (args.mode === 'full' && !args.dryRun) { - serviceStopped = stopServiceIfNeeded(stopService); + serviceStopped = stopServiceIfNeeded(stopService, stopMarkerPath); } await runDirectFilesBackup({ mode: args.mode, @@ -3365,7 +3394,7 @@ async function main() { } finally { try { if (serviceStopped) { - restoreServicesAfterBackup({stopService, serviceStopped, restartServicesAfter}); + restoreServicesAfterBackup({stopService, serviceStopped, restartServicesAfter, stopMarkerPath}); } else if (!backupError && args.mode === 'full' && !args.dryRun) { restartServicesAfterBackup(restartServicesAfter); } @@ -3419,16 +3448,17 @@ async function main() { let restoreError = null; const stopService = args.stopService || firstNonEmpty(env.GENARRATIVE_DATABASE_BACKUP_STOP_SERVICE); const restartServicesAfter = collectRestartServicesAfterBackup({args, env}); + const stopMarkerPath = databaseBackupStopMarkerPath(workDir); try { assertSufficientWorkDirSpace({dataDir, workDir, args, env}); - serviceStopped = stopServiceIfNeeded(stopService); + serviceStopped = stopServiceIfNeeded(stopService, stopMarkerPath); archivePath = createArchive({dataDir, workDir, fileName}); } catch (error) { backupError = error; } finally { try { if (serviceStopped) { - restoreServicesAfterBackup({stopService, serviceStopped, restartServicesAfter}); + restoreServicesAfterBackup({stopService, serviceStopped, restartServicesAfter, stopMarkerPath}); } else if (!backupError) { restartServicesAfterBackup(restartServicesAfter); } diff --git a/server-rs/crates/api-server/src/external_generation_worker.rs b/server-rs/crates/api-server/src/external_generation_worker.rs index 76223aefc..fa80d9b32 100644 --- a/server-rs/crates/api-server/src/external_generation_worker.rs +++ b/server-rs/crates/api-server/src/external_generation_worker.rs @@ -118,6 +118,9 @@ pub(crate) async fn run_external_generation_worker(state: AppState) -> Result<() ); loop { + // 持续有队列任务时不会进入等待分支,因此必须在每轮主动回收已完成的 + // JoinHandle;否则 permit 虽已归还,JoinSet 仍会保留每个历史任务的句柄。 + reap_finished_external_generation_worker_tasks(&mut tasks); ensure_external_generation_queue_wake_subscription(&state, &mut queue_wake).await; while work_slots.available_permits() == 0 { @@ -274,6 +277,14 @@ async fn await_worker_task(tasks: &mut JoinSet<()>) { } } +fn reap_finished_external_generation_worker_tasks(tasks: &mut JoinSet<()>) { + while let Some(result) = tasks.try_join_next() { + if let Err(error) = result { + error!(error = %error, "external generation worker 子任务 panic"); + } + } +} + async fn await_one_task_or_queue_wake_or_sleep_or_shutdown( tasks: &mut JoinSet<()>, queue_wake: &mut Option, @@ -387,7 +398,7 @@ async fn process_external_generation_job( job_id = %job.job_id, job_kind = %job.job_kind, timeout_seconds = job_timeout.as_secs(), - "external generation worker 任务超过执行预算,停止续租并释放 worker 槽位,在途执行交由租约仲裁" + "external generation worker 任务超过执行预算,停止续租并保留 worker 槽位,在途执行交由租约仲裁" ); detach_external_generation_work_until_lease_expiry( work_handle, @@ -458,6 +469,9 @@ fn detach_external_generation_work_until_lease_expiry( ), Err(_) => { work_handle.abort(); + // 仅调用 abort 不会从 JoinHandle/JoinSet 中消费完成结果;等待被取消 + // 的 handle,确保 permit 与任务句柄在同一生命周期内一起释放。 + let _ = work_handle.await; warn!( job_id = %job_id, job_kind = %job_kind, @@ -2031,6 +2045,28 @@ mod tests { ); } + #[tokio::test] + async fn worker_reaps_completed_tasks_while_queue_remains_busy() { + let mut tasks = JoinSet::new(); + let completed = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)); + const TASK_COUNT: usize = 128; + for _ in 0..TASK_COUNT { + let completed = completed.clone(); + tasks.spawn(async move { + completed.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + }); + } + + while completed.load(std::sync::atomic::Ordering::SeqCst) < TASK_COUNT { + tokio::task::yield_now().await; + } + assert_eq!(tasks.len(), TASK_COUNT); + + reap_finished_external_generation_worker_tasks(&mut tasks); + + assert!(tasks.is_empty(), "已完成任务的 JoinHandle 应在每轮被回收"); + } + #[tokio::test] async fn worker_detached_work_is_aborted_after_lease_arbitration_window() { let connection = std::sync::Arc::new(tokio::sync::Semaphore::new(1)); diff --git a/server-rs/crates/module-ai/src/application/store.rs b/server-rs/crates/module-ai/src/application/store.rs index 4ff835d1a..98df6c189 100644 --- a/server-rs/crates/module-ai/src/application/store.rs +++ b/server-rs/crates/module-ai/src/application/store.rs @@ -11,7 +11,9 @@ use super::ensure_task_is_not_terminal; const MAX_RETAINED_TASKS: usize = 1024; const MAX_TASK_TEXT_OUTPUT_BYTES: usize = 512 * 1024; -const MAX_RETAINED_TASK_TEXT_BYTES: usize = 64 * 1024 * 1024; +const MAX_TASK_STRUCTURED_OUTPUT_BYTES: usize = 512 * 1024; +const MAX_TASK_WARNING_BYTES: usize = 64 * 1024; +const MAX_RETAINED_TASK_OUTPUT_BYTES: usize = 64 * 1024 * 1024; #[derive(Clone, Debug, Default)] pub struct InMemoryAiTaskStore { @@ -40,18 +42,41 @@ impl InMemoryAiTaskStore { return Err(AiTaskServiceError::TaskAlreadyExists); } - if state.tasks.len() >= MAX_RETAINED_TASKS { + validate_task_output_limits(&task)?; + + let oldest_terminal = if state.tasks.len() >= MAX_RETAINED_TASKS { let oldest_terminal = state .tasks .values() .filter(|value| value.status.is_terminal()) .min_by_key(|value| value.completed_at_micros.or(Some(value.updated_at_micros))) .map(|value| value.task_id.clone()); - let Some(task_id) = oldest_terminal else { + if oldest_terminal.is_none() { return Err(AiTaskServiceError::Store( "AI 任务仓储已达到内存容量上限".to_string(), )); - }; + } + oldest_terminal + } else { + None + }; + + let retained_output_bytes = retained_output_bytes(&state) + .saturating_sub( + oldest_terminal + .as_deref() + .and_then(|task_id| state.tasks.get(task_id)) + .map(task_output_bytes) + .unwrap_or_default(), + ) + .saturating_add(task_output_bytes(&task)); + if retained_output_bytes > MAX_RETAINED_TASK_OUTPUT_BYTES { + return Err(AiTaskServiceError::Store( + "AI 任务仓储输出工作集超过内存上限".to_string(), + )); + } + + if let Some(task_id) = oldest_terminal { state.tasks.remove(&task_id); state.text_chunks.remove(&task_id); } @@ -75,12 +100,33 @@ impl InMemoryAiTaskStore { .inner .lock() .map_err(|_| AiTaskServiceError::Store("AI 任务仓储锁已中毒".to_string()))?; - let task = state - .tasks - .get_mut(task_id.trim()) - .ok_or(AiTaskServiceError::TaskNotFound)?; - apply(task)?; - let snapshot = task.clone(); + let (previous_task, snapshot) = { + let task = state + .tasks + .get_mut(task_id.trim()) + .ok_or(AiTaskServiceError::TaskNotFound)?; + let previous_task = task.clone(); + if let Err(error) = apply(task) { + *task = previous_task; + return Err(error); + } + (previous_task, task.clone()) + }; + if let Err(error) = validate_task_output_limits(&snapshot) { + state + .tasks + .insert(task_id.trim().to_string(), previous_task); + return Err(error); + } + let retained_output_bytes = retained_output_bytes(&state); + if retained_output_bytes > MAX_RETAINED_TASK_OUTPUT_BYTES { + state + .tasks + .insert(task_id.trim().to_string(), previous_task); + return Err(AiTaskServiceError::Store( + "AI 任务仓储输出工作集超过内存上限".to_string(), + )); + } if snapshot.status.is_terminal() { state.text_chunks.remove(task_id.trim()); } @@ -100,7 +146,7 @@ impl InMemoryAiTaskStore { "AI 任务文本输出超过内存上限".to_string(), )); } - let (previous_stage_output_bytes, previous_latest_output_bytes) = { + let (previous_stage_output_bytes, previous_latest_output_bytes, previous_task) = { let task = state .tasks .get(&chunk.task_id) @@ -114,6 +160,7 @@ impl InMemoryAiTaskStore { ( stage.text_output.as_ref().map_or(0, String::len), task.latest_text_output.as_ref().map_or(0, String::len), + task.clone(), ) }; @@ -140,14 +187,14 @@ impl InMemoryAiTaskStore { )); } - let retained_text_bytes = retained_text_bytes(&state) + let projected_retained_output_bytes = retained_output_bytes(&state) .saturating_sub(previous_stage_output_bytes) .saturating_sub(previous_latest_output_bytes) .saturating_add(aggregated_bytes.saturating_mul(2)); - if retained_text_bytes > MAX_RETAINED_TASK_TEXT_BYTES { + if projected_retained_output_bytes > MAX_RETAINED_TASK_OUTPUT_BYTES { rollback_text_chunk(&mut state, &chunk, previous_chunk); return Err(AiTaskServiceError::Store( - "AI 任务仓储文本工作集超过内存上限".to_string(), + "AI 任务仓储输出工作集超过内存上限".to_string(), )); } @@ -157,28 +204,38 @@ impl InMemoryAiTaskStore { Some(aggregated_text) }; - let task = state - .tasks - .get_mut(&chunk.task_id) - .ok_or(AiTaskServiceError::TaskNotFound)?; - ensure_task_is_not_terminal(task.status)?; - let stage = task - .stages - .iter_mut() - .find(|stage| stage.stage_kind == chunk.stage_kind) - .ok_or(AiTaskServiceError::StageNotFound)?; - if stage.status == AiTaskStageStatus::Pending { - stage.status = AiTaskStageStatus::Running; - stage.started_at_micros = Some(chunk.created_at_micros); + let snapshot = { + let task = state + .tasks + .get_mut(&chunk.task_id) + .ok_or(AiTaskServiceError::TaskNotFound)?; + ensure_task_is_not_terminal(task.status)?; + let stage = task + .stages + .iter_mut() + .find(|stage| stage.stage_kind == chunk.stage_kind) + .ok_or(AiTaskServiceError::StageNotFound)?; + if stage.status == AiTaskStageStatus::Pending { + stage.status = AiTaskStageStatus::Running; + stage.started_at_micros = Some(chunk.created_at_micros); + } + task.status = AiTaskStatus::Running; + task.started_at_micros + .get_or_insert(chunk.created_at_micros); + stage.text_output = normalized_output.clone(); + task.latest_text_output = normalized_output; + task.updated_at_micros = chunk.created_at_micros; + task.version += 1; + task.clone() + }; + if let Err(error) = validate_task_output_limits(&snapshot) + .and_then(|_| validate_retained_output_bytes(&state)) + { + state.tasks.insert(chunk.task_id.clone(), previous_task); + rollback_text_chunk(&mut state, &chunk, previous_chunk); + return Err(error); } - task.status = AiTaskStatus::Running; - task.started_at_micros - .get_or_insert(chunk.created_at_micros); - stage.text_output = normalized_output.clone(); - task.latest_text_output = normalized_output; - task.updated_at_micros = chunk.created_at_micros; - task.version += 1; - Ok(task.clone()) + Ok(snapshot) } pub(super) fn get_task(&self, task_id: &str) -> Result { @@ -212,13 +269,9 @@ fn rollback_text_chunk( } } -fn retained_text_bytes(state: &InMemoryAiTaskStoreState) -> usize { +fn retained_output_bytes(state: &InMemoryAiTaskStoreState) -> usize { let snapshot_bytes = state.tasks.values().fold(0_usize, |total, task| { - let latest = task.latest_text_output.as_ref().map_or(0, String::len); - let stages = task.stages.iter().fold(0_usize, |stage_total, stage| { - stage_total.saturating_add(stage.text_output.as_ref().map_or(0, String::len)) - }); - total.saturating_add(latest).saturating_add(stages) + total.saturating_add(task_output_bytes(task)) }); state .text_chunks @@ -231,3 +284,92 @@ fn retained_text_bytes(state: &InMemoryAiTaskStoreState) -> usize { }) }) } + +fn validate_retained_output_bytes( + state: &InMemoryAiTaskStoreState, +) -> Result<(), AiTaskServiceError> { + if retained_output_bytes(state) > MAX_RETAINED_TASK_OUTPUT_BYTES { + return Err(AiTaskServiceError::Store( + "AI 任务仓储输出工作集超过内存上限".to_string(), + )); + } + Ok(()) +} + +fn task_output_bytes(task: &AiTaskSnapshot) -> usize { + let latest_text = task.latest_text_output.as_ref().map_or(0, String::len); + let latest_structured = task + .latest_structured_payload_json + .as_ref() + .map_or(0, String::len); + let stage_bytes = task.stages.iter().fold(0_usize, |total, stage| { + let text = stage.text_output.as_ref().map_or(0, String::len); + let structured = stage + .structured_payload_json + .as_ref() + .map_or(0, String::len); + let warnings = stage + .warning_messages + .iter() + .fold(0_usize, |warning_total, warning| { + warning_total.saturating_add(warning.len()) + }); + total + .saturating_add(text) + .saturating_add(structured) + .saturating_add(warnings) + }); + latest_text + .saturating_add(latest_structured) + .saturating_add(stage_bytes) +} + +fn validate_task_output_limits(task: &AiTaskSnapshot) -> Result<(), AiTaskServiceError> { + if task.stages.iter().any(|stage| { + stage + .text_output + .as_ref() + .is_some_and(|text| text.len() > MAX_TASK_TEXT_OUTPUT_BYTES) + }) { + return Err(AiTaskServiceError::Store( + "AI 任务文本输出超过内存上限".to_string(), + )); + } + if task + .latest_text_output + .as_ref() + .is_some_and(|text| text.len() > MAX_TASK_TEXT_OUTPUT_BYTES) + { + return Err(AiTaskServiceError::Store( + "AI 任务文本输出超过内存上限".to_string(), + )); + } + if task.stages.iter().any(|stage| { + stage + .structured_payload_json + .as_ref() + .is_some_and(|payload| payload.len() > MAX_TASK_STRUCTURED_OUTPUT_BYTES) + }) || task + .latest_structured_payload_json + .as_ref() + .is_some_and(|payload| payload.len() > MAX_TASK_STRUCTURED_OUTPUT_BYTES) + { + return Err(AiTaskServiceError::Store( + "AI 任务结构化输出超过内存上限".to_string(), + )); + } + if task.stages.iter().any(|stage| { + stage + .warning_messages + .iter() + .fold(0_usize, |total, warning| { + total.saturating_add(warning.len()) + }) + > MAX_TASK_WARNING_BYTES + }) { + return Err(AiTaskServiceError::Store( + "AI 任务 warning 输出超过内存上限".to_string(), + )); + } + Ok(()) +} diff --git a/server-rs/crates/module-ai/src/tests.rs b/server-rs/crates/module-ai/src/tests.rs index 513f27a66..23a17e80b 100644 --- a/server-rs/crates/module-ai/src/tests.rs +++ b/server-rs/crates/module-ai/src/tests.rs @@ -188,6 +188,150 @@ fn complete_stage_updates_latest_outputs() { assert_eq!(stage.warning_messages, vec!["使用了 fallback 选项池"]); } +#[test] +fn complete_stage_rejects_oversized_text_output_without_mutating_task() { + let service = build_service(); + let task = service + .create_task(build_create_input(AiTaskKind::StoryGeneration)) + .expect("task should create"); + let oversized = "x".repeat(512 * 1024 + 1); + + let error = service + .complete_stage(AiStageCompletionInput { + task_id: task.task_id.clone(), + stage_kind: AiTaskStageKind::NormalizeResult, + text_output: Some(oversized), + structured_payload_json: None, + warning_messages: Vec::new(), + completed_at_micros: task.created_at_micros + 1, + }) + .expect_err("text output over the per-stage cap should fail"); + assert!(matches!(error, AiTaskServiceError::Store(message) if message.contains("文本输出"))); + + let unchanged = service + .get_task(&task.task_id) + .expect("task should remain readable"); + let stage = unchanged + .stages + .iter() + .find(|stage| stage.stage_kind == AiTaskStageKind::NormalizeResult) + .expect("normalize stage should exist"); + assert_eq!(stage.status, AiTaskStageStatus::Pending); + assert!(stage.text_output.is_none()); + assert!(unchanged.latest_text_output.is_none()); +} + +#[test] +fn complete_stage_rejects_oversized_structured_output_without_mutating_task() { + let service = build_service(); + let task = service + .create_task(build_create_input(AiTaskKind::StoryGeneration)) + .expect("task should create"); + let oversized = "x".repeat(512 * 1024 + 1); + + let error = service + .complete_stage(AiStageCompletionInput { + task_id: task.task_id.clone(), + stage_kind: AiTaskStageKind::NormalizeResult, + text_output: None, + structured_payload_json: Some(oversized), + warning_messages: Vec::new(), + completed_at_micros: task.created_at_micros + 1, + }) + .expect_err("structured output over the per-stage cap should fail"); + assert!(matches!(error, AiTaskServiceError::Store(message) if message.contains("结构化输出"))); + + let unchanged = service + .get_task(&task.task_id) + .expect("task should remain readable"); + let stage = unchanged + .stages + .iter() + .find(|stage| stage.stage_kind == AiTaskStageKind::NormalizeResult) + .expect("normalize stage should exist"); + assert_eq!(stage.status, AiTaskStageStatus::Pending); + assert!(stage.structured_payload_json.is_none()); + assert!(unchanged.latest_structured_payload_json.is_none()); +} + +#[test] +fn complete_stage_rejects_oversized_warning_output_without_mutating_task() { + let service = build_service(); + let task = service + .create_task(build_create_input(AiTaskKind::StoryGeneration)) + .expect("task should create"); + let oversized_warning = "w".repeat(64 * 1024 + 1); + + let error = service + .complete_stage(AiStageCompletionInput { + task_id: task.task_id.clone(), + stage_kind: AiTaskStageKind::NormalizeResult, + text_output: None, + structured_payload_json: None, + warning_messages: vec![oversized_warning], + completed_at_micros: task.created_at_micros + 1, + }) + .expect_err("warning output over the per-stage cap should fail"); + assert!(matches!(error, AiTaskServiceError::Store(message) if message.contains("warning"))); + + let unchanged = service + .get_task(&task.task_id) + .expect("task should remain readable"); + let stage = unchanged + .stages + .iter() + .find(|stage| stage.stage_kind == AiTaskStageKind::NormalizeResult) + .expect("normalize stage should exist"); + assert_eq!(stage.status, AiTaskStageStatus::Pending); + assert!(stage.warning_messages.is_empty()); +} + +#[test] +fn complete_stage_enforces_global_retained_output_cap() { + let service = build_service(); + let structured_payload = "x".repeat(512 * 1024); + for index in 0..64 { + let task = service + .create_task(AiTaskCreateInput { + task_id: format!("task-structured-cap-{index}"), + ..build_create_input(AiTaskKind::StoryGeneration) + }) + .expect("task should create"); + service + .complete_stage(AiStageCompletionInput { + task_id: task.task_id, + stage_kind: AiTaskStageKind::NormalizeResult, + text_output: None, + structured_payload_json: Some(structured_payload.clone()), + warning_messages: Vec::new(), + completed_at_micros: task.created_at_micros + 1, + }) + .expect("64 MiB retained output should remain within the global cap"); + } + + let task = service + .create_task(AiTaskCreateInput { + task_id: "task-structured-cap-overflow".to_string(), + ..build_create_input(AiTaskKind::StoryGeneration) + }) + .expect("the overflow candidate task itself should create"); + let error = service + .complete_stage(AiStageCompletionInput { + task_id: task.task_id.clone(), + stage_kind: AiTaskStageKind::NormalizeResult, + text_output: None, + structured_payload_json: Some(structured_payload), + warning_messages: Vec::new(), + completed_at_micros: task.created_at_micros + 1, + }) + .expect_err("global retained output cap should reject the overflow"); + assert!(matches!(error, AiTaskServiceError::Store(message) if message.contains("工作集"))); + let unchanged = service + .get_task(&task.task_id) + .expect("overflow task should remain readable"); + assert!(unchanged.latest_structured_payload_json.is_none()); +} + #[test] fn attach_result_reference_appends_binding() { let service = build_service(); diff --git a/server-rs/crates/module-auth/src/lib.rs b/server-rs/crates/module-auth/src/lib.rs index 58309b45a..14d080f7b 100644 --- a/server-rs/crates/module-auth/src/lib.rs +++ b/server-rs/crates/module-auth/src/lib.rs @@ -1078,6 +1078,7 @@ impl InMemoryAuthStoreState { } let now = OffsetDateTime::now_utc(); + let mut retained_refresh_session_count = 0_usize; for session in view.refresh_sessions { if !existing_user_ids.contains(&session.user_id) { continue; @@ -1089,6 +1090,12 @@ impl InMemoryAuthStoreState { ) { continue; } + retained_refresh_session_count += 1; + if retained_refresh_session_count > MAX_REFRESH_SESSIONS { + return Err(format!( + "认证投影中的 refresh session 数量超过内存上限(最多 {MAX_REFRESH_SESSIONS} 条)" + )); + } let client_info = serde_json::from_str::(&session.client_info_json) .map_err(|error| format!("解析 refresh session 客户端信息失败:{error}"))?; @@ -4239,6 +4246,41 @@ mod tests { assert_eq!(stale_error, RefreshSessionError::SessionNotFound); } + #[test] + fn projection_restore_rejects_too_many_retained_refresh_sessions() { + let client_info_json = + serde_json::to_string(&build_client_info()).expect("client info should serialize"); + let refresh_sessions = (0..=MAX_REFRESH_SESSIONS) + .map(|index| AuthStoreProjectionRefreshSession { + session_id: format!("session-{index}"), + user_id: "user_projection_cap".to_string(), + refresh_token_hash: format!("hash-{index}"), + issued_by_provider: "password".to_string(), + client_info_json: client_info_json.clone(), + expires_at: "2999-01-01T00:00:00Z".to_string(), + revoked_at: None, + created_at: "2026-01-01T00:00:00Z".to_string(), + updated_at: "2026-01-01T00:00:00Z".to_string(), + last_seen_at: "2026-01-01T00:00:00Z".to_string(), + }) + .collect(); + + let error = InMemoryAuthStore::from_projection_view(AuthStoreProjectionView { + updated_at_micros: 1, + users: vec![projection_user( + "user_projection_cap", + "projection_cap", + None, + )], + identities: vec![], + refresh_sessions, + }) + .expect_err("projection restore must enforce the refresh session cap"); + + assert!(error.contains("refresh session")); + assert!(error.contains(&MAX_REFRESH_SESSIONS.to_string())); + } + #[tokio::test] async fn wechat_login_hits_existing_user_by_union_id_before_openid() { let store = build_store();