diff --git a/docs/project-memory/shared-memory/decision-log.md b/docs/project-memory/shared-memory/decision-log.md index d57a6ae4a..b583c8932 100644 --- a/docs/project-memory/shared-memory/decision-log.md +++ b/docs/project-memory/shared-memory/decision-log.md @@ -7768,5 +7768,6 @@ CI 上 `background_agent_runtime_recovers_stale_running_before_pending_task` 在 - 决策:外部生成 worker 每轮主动 `try_join_next` 回收已完成 `JoinHandle`,避免持续有队列任务时只归还 semaphore permit 却让 `JoinSet` 句柄集合无界增长;超时脱管任务在 abort 后等待句柄结束,执行许可保持到 work 真正结束或被取消。 - 决策:`module-ai` 的阶段终态写入与流式增量统一受文本、结构化 JSON、warning 和全局 retained 工作集上限约束;认证投影恢复对过滤后的 refresh session 重新计数,超过 8192 条直接拒绝启动恢复,避免超限快照灌入内存。 +- 复审补充:`spacetime-module` 的 AI procedure 复用同一组任务元数据 / payload / 输出 / 结果引用上限,流式文本聚合超限时回滚事务,terminal task 收口后删除 `ai_text_chunk` 明细,避免真实持久化链路绕过内存边界。 - 决策:备份脚本停库前写入受保护 `.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 a1b22d2cc..1579cc82d 100644 --- a/docs/【后端架构】server-rs与SpacetimeDB数据契约-2026-05-15.md +++ b/docs/【后端架构】server-rs与SpacetimeDB数据契约-2026-05-15.md @@ -330,6 +330,7 @@ Responses 的终态载荷既是工具调用的恢复源,也是正文的恢复 - Rust 结构体:`AiTask` - 源码:`server-rs/crates/spacetime-module/src/ai/tasks.rs` - `module-ai` 的进程内热状态不是持久化真相:文本增量按阶段有序聚合并受单阶段 512 KiB 上限约束;terminal task 立即释放增量明细,内存工作集最多保留 1024 个任务。需要长期查询时必须读取 SpacetimeDB 的 `ai_task` / `ai_task_stage` 投影,不得依赖进程重启后仍存在的内存快照。 +- SpacetimeDB 的 AI 写入 procedure 必须复用同一组任务元数据、payload、文本、结构化输出、warning、失败消息和结果引用上限;流式聚合超过 512 KiB 时在事务内拒绝,terminal task 收口后删除 `ai_text_chunk` 明细,只保留阶段最终快照和结果引用。 ### `ai_task_event` diff --git a/server-rs/crates/module-ai/src/application/service.rs b/server-rs/crates/module-ai/src/application/service.rs index 713f3cd33..6f42f1d69 100644 --- a/server-rs/crates/module-ai/src/application/service.rs +++ b/server-rs/crates/module-ai/src/application/service.rs @@ -28,7 +28,7 @@ impl AiTaskService { validate_task_create_input(&input).map_err(AiTaskServiceError::Field)?; let snapshot = AiTaskSnapshot { - task_id: input.task_id.clone(), + task_id: normalize_required_string(input.task_id).unwrap_or_default(), task_kind: input.task_kind, owner_user_id: normalize_required_string(input.owner_user_id).unwrap_or_default(), request_label: normalize_required_string(input.request_label).unwrap_or_default(), diff --git a/server-rs/crates/module-ai/src/application/store.rs b/server-rs/crates/module-ai/src/application/store.rs index 98df6c189..6150dc397 100644 --- a/server-rs/crates/module-ai/src/application/store.rs +++ b/server-rs/crates/module-ai/src/application/store.rs @@ -5,16 +5,12 @@ use std::{ use crate::{ AiTaskServiceError, AiTaskSnapshot, AiTaskStageStatus, AiTaskStatus, AiTextChunkSnapshot, + MAX_AI_TASK_RETAINED_OUTPUT_BYTES, MAX_AI_TASK_RETAINED_TASKS, MAX_AI_TASK_TEXT_OUTPUT_BYTES, + validate_ai_task_snapshot_memory_limits, }; use super::ensure_task_is_not_terminal; -const MAX_RETAINED_TASKS: usize = 1024; -const MAX_TASK_TEXT_OUTPUT_BYTES: usize = 512 * 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 { inner: Arc>, @@ -42,9 +38,9 @@ impl InMemoryAiTaskStore { return Err(AiTaskServiceError::TaskAlreadyExists); } - validate_task_output_limits(&task)?; + validate_task_memory_limits(&task)?; - let oldest_terminal = if state.tasks.len() >= MAX_RETAINED_TASKS { + let oldest_terminal = if state.tasks.len() >= MAX_AI_TASK_RETAINED_TASKS { let oldest_terminal = state .tasks .values() @@ -70,7 +66,7 @@ impl InMemoryAiTaskStore { .unwrap_or_default(), ) .saturating_add(task_output_bytes(&task)); - if retained_output_bytes > MAX_RETAINED_TASK_OUTPUT_BYTES { + if retained_output_bytes > MAX_AI_TASK_RETAINED_OUTPUT_BYTES { return Err(AiTaskServiceError::Store( "AI 任务仓储输出工作集超过内存上限".to_string(), )); @@ -112,14 +108,14 @@ impl InMemoryAiTaskStore { } (previous_task, task.clone()) }; - if let Err(error) = validate_task_output_limits(&snapshot) { + if let Err(error) = validate_task_memory_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 { + if retained_output_bytes > MAX_AI_TASK_RETAINED_OUTPUT_BYTES { state .tasks .insert(task_id.trim().to_string(), previous_task); @@ -141,7 +137,7 @@ impl InMemoryAiTaskStore { .inner .lock() .map_err(|_| AiTaskServiceError::Store("AI 任务仓储锁已中毒".to_string()))?; - if chunk.delta_text.len() > MAX_TASK_TEXT_OUTPUT_BYTES { + if chunk.delta_text.len() > MAX_AI_TASK_TEXT_OUTPUT_BYTES { return Err(AiTaskServiceError::Store( "AI 任务文本输出超过内存上限".to_string(), )); @@ -180,7 +176,7 @@ impl InMemoryAiTaskStore { } (previous_chunk, aggregated_bytes, aggregated_text) }; - if aggregated_bytes > MAX_TASK_TEXT_OUTPUT_BYTES { + if aggregated_bytes > MAX_AI_TASK_TEXT_OUTPUT_BYTES { rollback_text_chunk(&mut state, &chunk, previous_chunk); return Err(AiTaskServiceError::Store( "AI 任务文本输出超过内存上限".to_string(), @@ -191,7 +187,7 @@ impl InMemoryAiTaskStore { .saturating_sub(previous_stage_output_bytes) .saturating_sub(previous_latest_output_bytes) .saturating_add(aggregated_bytes.saturating_mul(2)); - if projected_retained_output_bytes > MAX_RETAINED_TASK_OUTPUT_BYTES { + if projected_retained_output_bytes > MAX_AI_TASK_RETAINED_OUTPUT_BYTES { rollback_text_chunk(&mut state, &chunk, previous_chunk); return Err(AiTaskServiceError::Store( "AI 任务仓储输出工作集超过内存上限".to_string(), @@ -228,7 +224,7 @@ impl InMemoryAiTaskStore { task.version += 1; task.clone() }; - if let Err(error) = validate_task_output_limits(&snapshot) + if let Err(error) = validate_task_memory_limits(&snapshot) .and_then(|_| validate_retained_output_bytes(&state)) { state.tasks.insert(chunk.task_id.clone(), previous_task); @@ -288,7 +284,7 @@ fn retained_output_bytes(state: &InMemoryAiTaskStoreState) -> usize { fn validate_retained_output_bytes( state: &InMemoryAiTaskStoreState, ) -> Result<(), AiTaskServiceError> { - if retained_output_bytes(state) > MAX_RETAINED_TASK_OUTPUT_BYTES { + if retained_output_bytes(state) > MAX_AI_TASK_RETAINED_OUTPUT_BYTES { return Err(AiTaskServiceError::Store( "AI 任务仓储输出工作集超过内存上限".to_string(), )); @@ -297,6 +293,30 @@ fn validate_retained_output_bytes( } fn task_output_bytes(task: &AiTaskSnapshot) -> usize { + let task_metadata = task + .task_id + .len() + .saturating_add(task.owner_user_id.len()) + .saturating_add(task.request_label.len()) + .saturating_add(task.source_module.len()) + .saturating_add(task.source_entity_id.as_ref().map_or(0, String::len)) + .saturating_add(task.stages.iter().fold(0_usize, |total, stage| { + total + .saturating_add(stage.label.len()) + .saturating_add(stage.detail.len()) + })); + let request_payload = task.request_payload_json.as_ref().map_or(0, String::len); + let failure_message = task.failure_message.as_ref().map_or(0, String::len); + let result_references = task + .result_references + .iter() + .fold(0_usize, |total, reference| { + total + .saturating_add(reference.result_ref_id.len()) + .saturating_add(reference.task_id.len()) + .saturating_add(reference.reference_id.len()) + .saturating_add(reference.label.as_ref().map_or(0, String::len)) + }); let latest_text = task.latest_text_output.as_ref().map_or(0, String::len); let latest_structured = task .latest_structured_payload_json @@ -319,57 +339,16 @@ fn task_output_bytes(task: &AiTaskSnapshot) -> usize { .saturating_add(structured) .saturating_add(warnings) }); - latest_text + task_metadata + .saturating_add(request_payload) + .saturating_add(failure_message) + .saturating_add(result_references) + .saturating_add(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(()) +fn validate_task_memory_limits(task: &AiTaskSnapshot) -> Result<(), AiTaskServiceError> { + validate_ai_task_snapshot_memory_limits(task) + .map_err(|message| AiTaskServiceError::Store(message.to_string())) } diff --git a/server-rs/crates/module-ai/src/domain.rs b/server-rs/crates/module-ai/src/domain.rs index a9931048f..df225816e 100644 --- a/server-rs/crates/module-ai/src/domain.rs +++ b/server-rs/crates/module-ai/src/domain.rs @@ -1,4 +1,5 @@ mod ids; +mod limits; mod stages; mod types; @@ -8,6 +9,16 @@ pub use ids::{ generate_ai_task_stage_id, generate_ai_text_chunk_id, normalize_optional_text, normalize_string_list, }; +pub use limits::{ + MAX_AI_TASK_FAILURE_MESSAGE_BYTES, MAX_AI_TASK_ID_BYTES, MAX_AI_TASK_OWNER_USER_ID_BYTES, + MAX_AI_TASK_REFERENCE_ID_BYTES, MAX_AI_TASK_REFERENCE_LABEL_BYTES, + MAX_AI_TASK_REQUEST_LABEL_BYTES, MAX_AI_TASK_REQUEST_PAYLOAD_BYTES, + MAX_AI_TASK_RESULT_REFERENCES, MAX_AI_TASK_RETAINED_OUTPUT_BYTES, MAX_AI_TASK_RETAINED_TASKS, + MAX_AI_TASK_SOURCE_ENTITY_ID_BYTES, MAX_AI_TASK_SOURCE_MODULE_BYTES, + MAX_AI_TASK_STAGE_DETAIL_BYTES, MAX_AI_TASK_STAGE_LABEL_BYTES, + MAX_AI_TASK_STRUCTURED_OUTPUT_BYTES, MAX_AI_TASK_TEXT_OUTPUT_BYTES, MAX_AI_TASK_WARNING_BYTES, + validate_ai_task_snapshot_memory_limits, +}; pub use types::{ AiResultReferenceKind, AiResultReferenceSnapshot, AiTaskKind, AiTaskSnapshot, AiTaskStageBlueprint, AiTaskStageKind, AiTaskStageSnapshot, AiTaskStageStatus, AiTaskStatus, diff --git a/server-rs/crates/module-ai/src/domain/limits.rs b/server-rs/crates/module-ai/src/domain/limits.rs new file mode 100644 index 000000000..759b4ac6a --- /dev/null +++ b/server-rs/crates/module-ai/src/domain/limits.rs @@ -0,0 +1,115 @@ +use super::types::AiTaskSnapshot; + +pub const MAX_AI_TASK_RETAINED_TASKS: usize = 1024; +pub const MAX_AI_TASK_ID_BYTES: usize = 256; +pub const MAX_AI_TASK_OWNER_USER_ID_BYTES: usize = 256; +pub const MAX_AI_TASK_REQUEST_LABEL_BYTES: usize = 4 * 1024; +pub const MAX_AI_TASK_SOURCE_MODULE_BYTES: usize = 256; +pub const MAX_AI_TASK_SOURCE_ENTITY_ID_BYTES: usize = 512; +pub const MAX_AI_TASK_STAGE_LABEL_BYTES: usize = 4 * 1024; +pub const MAX_AI_TASK_STAGE_DETAIL_BYTES: usize = 8 * 1024; +pub const MAX_AI_TASK_TEXT_OUTPUT_BYTES: usize = 512 * 1024; +pub const MAX_AI_TASK_STRUCTURED_OUTPUT_BYTES: usize = 512 * 1024; +pub const MAX_AI_TASK_WARNING_BYTES: usize = 64 * 1024; +pub const MAX_AI_TASK_REQUEST_PAYLOAD_BYTES: usize = 512 * 1024; +pub const MAX_AI_TASK_FAILURE_MESSAGE_BYTES: usize = 64 * 1024; +pub const MAX_AI_TASK_RESULT_REFERENCES: usize = 64; +pub const MAX_AI_TASK_REFERENCE_ID_BYTES: usize = 512; +pub const MAX_AI_TASK_REFERENCE_LABEL_BYTES: usize = 2 * 1024; +pub const MAX_AI_TASK_RETAINED_OUTPUT_BYTES: usize = 64 * 1024 * 1024; + +pub fn validate_ai_task_snapshot_memory_limits(task: &AiTaskSnapshot) -> Result<(), &'static str> { + if task.task_id.len() > MAX_AI_TASK_ID_BYTES { + return Err("AI 任务 ID 超过内存上限"); + } + if task.owner_user_id.len() > MAX_AI_TASK_OWNER_USER_ID_BYTES { + return Err("AI 任务用户 ID 超过内存上限"); + } + if task.request_label.len() > MAX_AI_TASK_REQUEST_LABEL_BYTES { + return Err("AI 任务请求标签超过内存上限"); + } + if task.source_module.len() > MAX_AI_TASK_SOURCE_MODULE_BYTES { + return Err("AI 任务来源模块超过内存上限"); + } + if task + .source_entity_id + .as_ref() + .is_some_and(|entity_id| entity_id.len() > MAX_AI_TASK_SOURCE_ENTITY_ID_BYTES) + { + return Err("AI 任务来源实体 ID 超过内存上限"); + } + if task.stages.iter().any(|stage| { + stage.label.len() > MAX_AI_TASK_STAGE_LABEL_BYTES + || stage.detail.len() > MAX_AI_TASK_STAGE_DETAIL_BYTES + }) { + return Err("AI 任务阶段元数据超过内存上限"); + } + if task + .request_payload_json + .as_ref() + .is_some_and(|payload| payload.len() > MAX_AI_TASK_REQUEST_PAYLOAD_BYTES) + { + return Err("AI 任务请求 payload 超过内存上限"); + } + if task + .failure_message + .as_ref() + .is_some_and(|message| message.len() > MAX_AI_TASK_FAILURE_MESSAGE_BYTES) + { + return Err("AI 任务失败消息超过内存上限"); + } + if task.stages.iter().any(|stage| { + stage + .text_output + .as_ref() + .is_some_and(|text| text.len() > MAX_AI_TASK_TEXT_OUTPUT_BYTES) + }) || task + .latest_text_output + .as_ref() + .is_some_and(|text| text.len() > MAX_AI_TASK_TEXT_OUTPUT_BYTES) + { + return Err("AI 任务文本输出超过内存上限"); + } + if task.stages.iter().any(|stage| { + stage + .structured_payload_json + .as_ref() + .is_some_and(|payload| payload.len() > MAX_AI_TASK_STRUCTURED_OUTPUT_BYTES) + }) || task + .latest_structured_payload_json + .as_ref() + .is_some_and(|payload| payload.len() > MAX_AI_TASK_STRUCTURED_OUTPUT_BYTES) + { + return Err("AI 任务结构化输出超过内存上限"); + } + if task.stages.iter().any(|stage| { + stage + .warning_messages + .iter() + .fold(0_usize, |total, warning| { + total.saturating_add(warning.len()) + }) + > MAX_AI_TASK_WARNING_BYTES + }) { + return Err("AI 任务 warning 输出超过内存上限"); + } + if task.result_references.len() > MAX_AI_TASK_RESULT_REFERENCES { + return Err("AI 任务结果引用数量超过内存上限"); + } + if task + .result_references + .iter() + .any(|reference| reference.reference_id.len() > MAX_AI_TASK_REFERENCE_ID_BYTES) + { + return Err("AI 任务结果引用 ID 超过内存上限"); + } + if task.result_references.iter().any(|reference| { + reference + .label + .as_ref() + .is_some_and(|label| label.len() > MAX_AI_TASK_REFERENCE_LABEL_BYTES) + }) { + return Err("AI 任务结果引用标签超过内存上限"); + } + Ok(()) +} diff --git a/server-rs/crates/module-ai/src/lib.rs b/server-rs/crates/module-ai/src/lib.rs index e69b5609f..2eba88eb0 100644 --- a/server-rs/crates/module-ai/src/lib.rs +++ b/server-rs/crates/module-ai/src/lib.rs @@ -14,9 +14,17 @@ pub use domain::{ AI_RESULT_REF_ID_PREFIX, AI_TASK_ID_PREFIX, AI_TASK_STAGE_ID_PREFIX, AI_TEXT_CHUNK_ID_PREFIX, AiResultReferenceKind, AiResultReferenceSnapshot, AiTaskKind, AiTaskSnapshot, AiTaskStageBlueprint, AiTaskStageKind, AiTaskStageSnapshot, AiTaskStageStatus, AiTaskStatus, - AiTextChunkSnapshot, INITIAL_AI_TASK_VERSION, generate_ai_result_ref_id, generate_ai_task_id, - generate_ai_task_stage_id, generate_ai_text_chunk_id, normalize_optional_text, - normalize_string_list, + AiTextChunkSnapshot, INITIAL_AI_TASK_VERSION, MAX_AI_TASK_FAILURE_MESSAGE_BYTES, + MAX_AI_TASK_ID_BYTES, MAX_AI_TASK_OWNER_USER_ID_BYTES, MAX_AI_TASK_REFERENCE_ID_BYTES, + MAX_AI_TASK_REFERENCE_LABEL_BYTES, MAX_AI_TASK_REQUEST_LABEL_BYTES, + MAX_AI_TASK_REQUEST_PAYLOAD_BYTES, MAX_AI_TASK_RESULT_REFERENCES, + MAX_AI_TASK_RETAINED_OUTPUT_BYTES, MAX_AI_TASK_RETAINED_TASKS, + MAX_AI_TASK_SOURCE_ENTITY_ID_BYTES, MAX_AI_TASK_SOURCE_MODULE_BYTES, + MAX_AI_TASK_STAGE_DETAIL_BYTES, MAX_AI_TASK_STAGE_LABEL_BYTES, + MAX_AI_TASK_STRUCTURED_OUTPUT_BYTES, MAX_AI_TASK_TEXT_OUTPUT_BYTES, MAX_AI_TASK_WARNING_BYTES, + generate_ai_result_ref_id, generate_ai_task_id, generate_ai_task_stage_id, + generate_ai_text_chunk_id, normalize_optional_text, normalize_string_list, + validate_ai_task_snapshot_memory_limits, }; pub use errors::{AiTaskFieldError, AiTaskServiceError}; pub use events::AiTaskDomainEvent; diff --git a/server-rs/crates/module-ai/src/tests.rs b/server-rs/crates/module-ai/src/tests.rs index 23a17e80b..a1ba8f6cb 100644 --- a/server-rs/crates/module-ai/src/tests.rs +++ b/server-rs/crates/module-ai/src/tests.rs @@ -43,6 +43,32 @@ fn create_task_rejects_duplicate_stage_blueprints() { assert_eq!(error, AiTaskFieldError::DuplicateStageBlueprint); } +#[test] +fn create_task_rejects_oversized_request_payload() { + let service = build_service(); + let mut input = build_create_input(AiTaskKind::StoryGeneration); + input.request_payload_json = Some("x".repeat(MAX_AI_TASK_REQUEST_PAYLOAD_BYTES + 1)); + + let error = service + .create_task(input) + .expect_err("request payload over the memory cap should fail"); + assert!( + matches!(error, AiTaskServiceError::Store(message) if message.contains("请求 payload")) + ); +} + +#[test] +fn create_task_rejects_oversized_request_metadata() { + let service = build_service(); + let mut input = build_create_input(AiTaskKind::StoryGeneration); + input.request_label = "x".repeat(MAX_AI_TASK_REQUEST_LABEL_BYTES + 1); + + let error = service + .create_task(input) + .expect_err("request metadata over the memory cap should fail"); + assert!(matches!(error, AiTaskServiceError::Store(message) if message.contains("请求标签"))); +} + #[test] fn generate_ai_task_stage_id_contains_task_and_stage_slug() { let stage_id = generate_ai_task_stage_id("aitask_demo", AiTaskStageKind::NormalizeResult); @@ -290,7 +316,7 @@ fn complete_stage_rejects_oversized_warning_output_without_mutating_task() { fn complete_stage_enforces_global_retained_output_cap() { let service = build_service(); let structured_payload = "x".repeat(512 * 1024); - for index in 0..64 { + for index in 0..63 { let task = service .create_task(AiTaskCreateInput { task_id: format!("task-structured-cap-{index}"), @@ -306,7 +332,7 @@ fn complete_stage_enforces_global_retained_output_cap() { warning_messages: Vec::new(), completed_at_micros: task.created_at_micros + 1, }) - .expect("64 MiB retained output should remain within the global cap"); + .expect("63 MiB retained output should remain within the global cap"); } let task = service @@ -357,6 +383,47 @@ fn attach_result_reference_appends_binding() { assert_eq!(updated.result_references[0].reference_id, "profile_001"); } +#[test] +fn attach_result_reference_rejects_unbounded_reference_growth() { + let service = build_service(); + let task = service + .create_task(build_create_input(AiTaskKind::CustomWorldGeneration)) + .expect("task should create"); + + for index in 0..MAX_AI_TASK_RESULT_REFERENCES { + service + .attach_result_reference( + &task.task_id, + AiResultReferenceKind::CustomWorldProfile, + format!("profile_{index}"), + None, + task.created_at_micros + index as i64 + 1, + ) + .expect("references within the cap should attach"); + } + + let error = service + .attach_result_reference( + &task.task_id, + AiResultReferenceKind::CustomWorldProfile, + "profile_overflow".to_string(), + None, + task.created_at_micros + MAX_AI_TASK_RESULT_REFERENCES as i64 + 1, + ) + .expect_err("references over the 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"); + assert_eq!( + unchanged.result_references.len(), + MAX_AI_TASK_RESULT_REFERENCES + ); +} + #[test] fn fail_and_cancel_task_move_into_terminal_states() { let service = build_service(); diff --git a/server-rs/crates/spacetime-module/src/ai/snapshots.rs b/server-rs/crates/spacetime-module/src/ai/snapshots.rs index 8ee4c0bec..f76d209e7 100644 --- a/server-rs/crates/spacetime-module/src/ai/snapshots.rs +++ b/server-rs/crates/spacetime-module/src/ai/snapshots.rs @@ -139,17 +139,6 @@ pub(crate) fn build_ai_text_chunk_row_id(snapshot: &AiTextChunkSnapshot) -> Stri ) } -pub(crate) fn build_ai_text_chunk_snapshot_from_row(row: &AiTextChunk) -> AiTextChunkSnapshot { - AiTextChunkSnapshot { - chunk_id: row.chunk_id.clone(), - task_id: row.task_id.clone(), - stage_kind: row.stage_kind, - sequence: row.sequence, - delta_text: row.delta_text.clone(), - created_at_micros: row.created_at.to_micros_since_unix_epoch(), - } -} - pub(crate) fn build_ai_result_reference_row( snapshot: &AiResultReferenceSnapshot, ) -> AiResultReference { diff --git a/server-rs/crates/spacetime-module/src/ai/stages.rs b/server-rs/crates/spacetime-module/src/ai/stages.rs index 8ffea6f22..1fa2cab3f 100644 --- a/server-rs/crates/spacetime-module/src/ai/stages.rs +++ b/server-rs/crates/spacetime-module/src/ai/stages.rs @@ -1,7 +1,7 @@ use crate::*; use module_ai::{ - generate_ai_result_ref_id, generate_ai_text_chunk_id, normalize_optional_text, - normalize_string_list, + MAX_AI_TASK_TEXT_OUTPUT_BYTES, generate_ai_result_ref_id, generate_ai_text_chunk_id, + normalize_optional_text, normalize_string_list, validate_ai_task_snapshot_memory_limits, }; #[spacetimedb::table( @@ -178,6 +178,9 @@ pub(crate) fn append_ai_text_chunk_tx( if input.sequence == 0 { return Err("ai_text_chunk.sequence 必须大于 0".to_string()); } + if input.delta_text.trim().len() > MAX_AI_TASK_TEXT_OUTPUT_BYTES { + return Err("AI 任务文本输出超过内存上限".to_string()); + } let mut snapshot = get_ai_task_snapshot_tx(ctx, &input.task_id)?; ensure_ai_task_can_transition(snapshot.status)?; @@ -200,7 +203,7 @@ pub(crate) fn append_ai_text_chunk_tx( .ai_text_chunk() .insert(build_ai_text_chunk_row(&chunk)); - let aggregated_text = collect_ai_stage_text_output(ctx, &chunk.task_id, chunk.stage_kind); + let aggregated_text = collect_ai_stage_text_output(ctx, &chunk.task_id, chunk.stage_kind)?; snapshot.status = AiTaskStatus::Running; if snapshot.started_at_micros.is_none() { @@ -215,6 +218,7 @@ pub(crate) fn append_ai_text_chunk_tx( snapshot.updated_at_micros = input.created_at_micros; snapshot.version += 1; + validate_ai_task_snapshot_memory_limits(&snapshot).map_err(str::to_string)?; persist_ai_task_snapshot(ctx, &snapshot)?; emit_ai_task_event( ctx, @@ -252,6 +256,7 @@ pub(crate) fn complete_ai_stage_tx( snapshot.updated_at_micros = input.completed_at_micros; snapshot.version += 1; + validate_ai_task_snapshot_memory_limits(&snapshot).map_err(str::to_string)?; persist_ai_task_snapshot(ctx, &snapshot)?; emit_ai_task_event( ctx, @@ -285,26 +290,27 @@ pub(crate) fn attach_ai_result_reference_tx( label: normalize_optional_text(input.label), created_at_micros: input.created_at_micros, }; - ctx.db - .ai_result_reference() - .insert(build_ai_result_reference_row(&reference)); - snapshot.result_references.push(reference); snapshot.updated_at_micros = input.created_at_micros; snapshot.version += 1; - persist_ai_task_snapshot(ctx, &snapshot)?; + validate_ai_task_snapshot_memory_limits(&snapshot).map_err(str::to_string)?; let reference = snapshot .result_references .last() + .cloned() .ok_or_else(|| "ai_result_reference 写入后缺少快照".to_string())?; + ctx.db + .ai_result_reference() + .insert(build_ai_result_reference_row(&reference)); + persist_ai_task_snapshot(ctx, &snapshot)?; emit_ai_task_event( ctx, &snapshot, AiTaskEventKind::ResultReferenceAttached, None, None, - Some(build_ai_result_reference_row_id(reference)), + Some(build_ai_result_reference_row_id(&reference)), input.created_at_micros, ); Ok(snapshot) @@ -333,29 +339,48 @@ pub(crate) fn replace_ai_task_stages( } } +pub(crate) fn delete_ai_text_chunks_for_task(ctx: &ReducerContext, task_id: &str) { + let chunk_row_ids = ctx + .db + .ai_text_chunk() + .by_ai_text_chunk_task_id() + .filter(task_id) + .map(|row| row.text_chunk_row_id.clone()) + .collect::>(); + for row_id in chunk_row_ids { + ctx.db.ai_text_chunk().text_chunk_row_id().delete(&row_id); + } +} + pub(crate) fn collect_ai_stage_text_output( ctx: &ReducerContext, task_id: &str, stage_kind: AiTaskStageKind, -) -> Option { - let mut chunks = ctx +) -> Result, String> { + let mut chunks = Vec::new(); + let mut aggregated_bytes = 0_usize; + for row in ctx .db .ai_text_chunk() .by_ai_text_chunk_task_id() .filter(task_id) .filter(|row| row.task_id == task_id && row.stage_kind == stage_kind) - .map(|row| build_ai_text_chunk_snapshot_from_row(&row)) - .collect::>(); - chunks.sort_by_key(|chunk| chunk.sequence); + { + aggregated_bytes = aggregated_bytes.saturating_add(row.delta_text.len()); + if aggregated_bytes > MAX_AI_TASK_TEXT_OUTPUT_BYTES { + return Err("AI 任务文本输出超过内存上限".to_string()); + } + chunks.push((row.sequence, row.delta_text.clone())); + } + chunks.sort_by_key(|(sequence, _)| *sequence); - let aggregated = chunks - .into_iter() - .map(|chunk| chunk.delta_text) - .collect::>() - .join(""); + let mut aggregated = String::with_capacity(aggregated_bytes); + for (_, delta) in chunks { + aggregated.push_str(&delta); + } if aggregated.trim().is_empty() { - None + Ok(None) } else { - Some(aggregated) + Ok(Some(aggregated)) } } diff --git a/server-rs/crates/spacetime-module/src/ai/tasks.rs b/server-rs/crates/spacetime-module/src/ai/tasks.rs index 66c0909d1..b6c646789 100644 --- a/server-rs/crates/spacetime-module/src/ai/tasks.rs +++ b/server-rs/crates/spacetime-module/src/ai/tasks.rs @@ -1,5 +1,8 @@ use crate::*; -use module_ai::{INITIAL_AI_TASK_VERSION, normalize_optional_text, validate_task_create_input}; +use module_ai::{ + INITIAL_AI_TASK_VERSION, normalize_optional_text, validate_ai_task_snapshot_memory_limits, + validate_task_create_input, +}; #[spacetimedb::table( accessor = ai_task, @@ -133,6 +136,7 @@ fn create_ai_task_tx( } let task_snapshot = build_ai_task_snapshot_from_create_input(&input); + validate_ai_task_snapshot_memory_limits(&task_snapshot).map_err(str::to_string)?; ctx.db.ai_task().insert(build_ai_task_row(&task_snapshot)); replace_ai_task_stages(ctx, &task_snapshot.task_id, &task_snapshot.stages); emit_ai_task_event( @@ -187,7 +191,9 @@ fn complete_ai_task_tx( snapshot.updated_at_micros = input.completed_at_micros; snapshot.version += 1; + validate_ai_task_snapshot_memory_limits(&snapshot).map_err(str::to_string)?; persist_ai_task_snapshot(ctx, &snapshot)?; + delete_ai_text_chunks_for_task(ctx, &snapshot.task_id); emit_ai_task_event( ctx, &snapshot, @@ -218,7 +224,9 @@ fn fail_ai_task_tx( snapshot.updated_at_micros = input.completed_at_micros; snapshot.version += 1; + validate_ai_task_snapshot_memory_limits(&snapshot).map_err(str::to_string)?; persist_ai_task_snapshot(ctx, &snapshot)?; + delete_ai_text_chunks_for_task(ctx, &snapshot.task_id); emit_ai_task_event( ctx, &snapshot, @@ -243,7 +251,9 @@ fn cancel_ai_task_tx( snapshot.updated_at_micros = input.completed_at_micros; snapshot.version += 1; + validate_ai_task_snapshot_memory_limits(&snapshot).map_err(str::to_string)?; persist_ai_task_snapshot(ctx, &snapshot)?; + delete_ai_text_chunks_for_task(ctx, &snapshot.task_id); emit_ai_task_event( ctx, &snapshot,