修复生产内存工作集回收与 OOM 恢复
Project CI / Frontend tests (pull_request) Successful in 2m47s
Project CI / Backend tests (pull_request) Successful in 6m38s
Project CI / Repository checks (pull_request) Successful in 2m45s
Project CI / Native shell tests (pull_request) Successful in 15m30s

回收持续队列中的已完成 worker 句柄并收紧脱管许可生命周期

统一 AI 任务终态写入的文本、结构化输出和 warning 内存上限

限制认证投影恢复的 retained refresh session 数量

为备份停库增加 marker 与 systemd OOM 兜底恢复

补充回归测试并同步后端、运维和项目决策文档
This commit is contained in:
2026-08-27 18:09:44 +08:00
parent 810bb1d12b
commit a02a318758
11 changed files with 498 additions and 49 deletions
@@ -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
@@ -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。
@@ -411,6 +411,8 @@ Responses 的终态载荷既是工具调用的恢复源,也是正文的恢复
### `auth_store_projection_meta`
启动投影恢复会对过滤后的 retained refresh session 重新计数;超过 8192 条时直接失败关闭并继续重试,不得把超限快照一次性灌入内存。
- Rust 结构体:`AuthStoreProjectionMeta`
- 源码:`server-rs/crates/spacetime-module/src/auth/tables.rs`
@@ -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
+26
View File
@@ -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() {
+14 -1
View File
@@ -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 及依赖服务。',
},
{
+36 -6
View File
@@ -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 命令失败时仍保留 markersystemd 的 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);
}
@@ -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<ExternalGenerationQueueWakeSubscription>,
@@ -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));
@@ -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<AiTaskSnapshot, AiTaskServiceError> {
@@ -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(())
}
+144
View File
@@ -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();
+42
View File
@@ -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::<RefreshSessionClientInfo>(&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();