Files
Genarrative/server-rs/crates/module-ai/src/tests.rs
T
kdletters a7337c67a1
Project CI / Repository checks (push) Successful in 3m24s
Project CI / Frontend tests (push) Successful in 4m40s
Project CI / Backend tests (push) Successful in 10m5s
Project CI / Native shell tests (push) Successful in 15m8s
修复生产发布内存持续增长 (#203)
修复 release 内存持续增长问题。

本次范围:
- 收口备份扫描与历史维护的内存峰值。
- 限制外部生成 worker 脱管任务的实际并发。
- 为 API 内存态和历史数据增加有界留存。

验收:
- 定向测试、cargo 检查和生产运维门禁通过。
- 备份不再触发全局 OOM。
- worker/API/SpacetimeDB 内存曲线在空闲期停止单调增长。

Reviewed-on: http://192.168.35.82/git/GenarrativeAI/Genarrative/pulls/203
Co-authored-by: kdletters <kdletters@qq.com>
Co-committed-by: kdletters <kdletters@qq.com>
2026-08-27 21:46:56 +08:00

478 lines
17 KiB
Rust

use super::*;
fn build_service() -> AiTaskService {
AiTaskService::new(InMemoryAiTaskStore::default())
}
fn build_create_input(task_kind: AiTaskKind) -> AiTaskCreateInput {
AiTaskCreateInput {
task_id: generate_ai_task_id(1_713_680_000_000_000),
task_kind,
owner_user_id: "user_001".to_string(),
request_label: "首轮故事生成".to_string(),
source_module: "story".to_string(),
source_entity_id: Some("storysess_001".to_string()),
request_payload_json: Some("{\"scene\":\"camp\"}".to_string()),
stages: task_kind.default_stage_blueprints(),
created_at_micros: 1_713_680_000_000_000,
}
}
#[test]
fn default_stage_blueprints_match_story_baseline() {
let stages = AiTaskKind::StoryGeneration.default_stage_blueprints();
assert_eq!(stages.len(), 4);
assert_eq!(stages[0].stage_kind, AiTaskStageKind::PreparePrompt);
assert_eq!(stages[1].stage_kind, AiTaskStageKind::RequestModel);
assert_eq!(stages[2].stage_kind, AiTaskStageKind::RepairResponse);
assert_eq!(stages[3].stage_kind, AiTaskStageKind::NormalizeResult);
}
#[test]
fn create_task_rejects_duplicate_stage_blueprints() {
let mut input = build_create_input(AiTaskKind::StoryGeneration);
input.stages.push(AiTaskStageBlueprint {
stage_kind: AiTaskStageKind::PreparePrompt,
label: "重复阶段".to_string(),
detail: "重复阶段".to_string(),
order: 99,
});
let error = validate_task_create_input(&input).expect_err("duplicate stages should fail");
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);
assert_eq!(stage_id, "aistage_aitask_demo_normalize_result");
}
#[test]
fn generate_ai_task_id_uses_task_prefix_for_new_ids() {
assert!(generate_ai_task_id(1_713_680_000_000_000).starts_with("task-"));
}
#[test]
fn create_and_start_task_updates_status() {
let service = build_service();
let created = service
.create_task(build_create_input(AiTaskKind::QuestIntent))
.expect("task should create");
let started = service
.start_task(&created.task_id, created.created_at_micros + 1)
.expect("task should start");
assert_eq!(created.status, AiTaskStatus::Pending);
assert_eq!(started.status, AiTaskStatus::Running);
assert_eq!(
started.started_at_micros,
Some(created.created_at_micros + 1)
);
assert_eq!(started.version, INITIAL_AI_TASK_VERSION + 1);
}
#[test]
fn append_text_chunk_aggregates_stream_output_by_stage() {
let service = build_service();
let task = service
.create_task(build_create_input(AiTaskKind::CharacterChat))
.expect("task should create");
service
.start_stage(
&task.task_id,
AiTaskStageKind::RequestModel,
task.created_at_micros + 10,
)
.expect("stage should start");
let (after_first, _) = service
.append_text_chunk(
&task.task_id,
AiTaskStageKind::RequestModel,
1,
"你".to_string(),
task.created_at_micros + 20,
)
.expect("first chunk should append");
let (after_second, second_chunk) = service
.append_text_chunk(
&task.task_id,
AiTaskStageKind::RequestModel,
2,
"好。".to_string(),
task.created_at_micros + 30,
)
.expect("second chunk should append");
assert_eq!(after_first.latest_text_output.as_deref(), Some("你"));
assert_eq!(after_second.latest_text_output.as_deref(), Some("你好。"));
assert_eq!(second_chunk.sequence, 2);
}
#[test]
fn append_text_chunk_rejects_output_over_stage_memory_limit_without_mutating_task() {
let service = build_service();
let task = service
.create_task(build_create_input(AiTaskKind::CharacterChat))
.expect("task should create");
let max_output = "a".repeat(512 * 1024);
let (updated, _) = service
.append_text_chunk(
&task.task_id,
AiTaskStageKind::RequestModel,
1,
max_output.clone(),
task.created_at_micros + 1,
)
.expect("the stage limit itself should be accepted");
assert_eq!(
updated.latest_text_output.as_deref().map(str::len),
Some(max_output.len())
);
let error = service
.append_text_chunk(
&task.task_id,
AiTaskStageKind::RequestModel,
2,
"b".to_string(),
task.created_at_micros + 2,
)
.expect_err("output beyond the stage limit should fail");
assert!(matches!(error, AiTaskServiceError::Store(_)));
let after_rejection = service
.get_task(&task.task_id)
.expect("task should remain readable");
assert_eq!(
after_rejection.latest_text_output.as_deref().map(str::len),
Some(max_output.len())
);
}
#[test]
fn complete_stage_updates_latest_outputs() {
let service = build_service();
let task = service
.create_task(build_create_input(AiTaskKind::StoryGeneration))
.expect("task should create");
let completed = service
.complete_stage(AiStageCompletionInput {
task_id: task.task_id.clone(),
stage_kind: AiTaskStageKind::NormalizeResult,
text_output: Some("营地前的篝火重新亮了起来。".to_string()),
structured_payload_json: Some("{\"choices\":3}".to_string()),
warning_messages: vec!["使用了 fallback 选项池".to_string()],
completed_at_micros: task.created_at_micros + 50,
})
.expect("stage should complete");
let stage = completed
.stages
.iter()
.find(|stage| stage.stage_kind == AiTaskStageKind::NormalizeResult)
.expect("normalize stage should exist");
assert_eq!(stage.status, AiTaskStageStatus::Completed);
assert_eq!(
completed.latest_text_output.as_deref(),
Some("营地前的篝火重新亮了起来。")
);
assert_eq!(
completed.latest_structured_payload_json.as_deref(),
Some("{\"choices\":3}")
);
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..63 {
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("63 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();
let task = service
.create_task(build_create_input(AiTaskKind::CustomWorldGeneration))
.expect("task should create");
let updated = service
.attach_result_reference(
&task.task_id,
AiResultReferenceKind::CustomWorldProfile,
"profile_001".to_string(),
Some("主世界档案".to_string()),
task.created_at_micros + 10,
)
.expect("reference should attach");
assert_eq!(updated.result_references.len(), 1);
assert_eq!(
updated.result_references[0].reference_kind,
AiResultReferenceKind::CustomWorldProfile
);
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();
let first = service
.create_task(build_create_input(AiTaskKind::NpcChat))
.expect("task should create");
let failed = service
.fail_task(
&first.task_id,
"上游模型超时".to_string(),
first.created_at_micros + 10,
)
.expect("task should fail");
assert_eq!(failed.status, AiTaskStatus::Failed);
assert_eq!(failed.failure_message.as_deref(), Some("上游模型超时"));
let second = service
.create_task(AiTaskCreateInput {
task_id: generate_ai_task_id(1_713_680_000_000_999),
..build_create_input(AiTaskKind::RuntimeItemIntent)
})
.expect("second task should create");
let cancelled = service
.cancel_task(&second.task_id, second.created_at_micros + 20)
.expect("task should cancel");
assert_eq!(cancelled.status, AiTaskStatus::Cancelled);
assert_eq!(
cancelled.completed_at_micros,
Some(second.created_at_micros + 20)
);
}
#[test]
fn complete_task_marks_terminal_success() {
let service = build_service();
let task = service
.create_task(build_create_input(AiTaskKind::QuestIntent))
.expect("task should create");
let completed = service
.complete_task(&task.task_id, task.created_at_micros + 100)
.expect("task should complete");
assert_eq!(completed.status, AiTaskStatus::Completed);
assert_eq!(
completed.completed_at_micros,
Some(task.created_at_micros + 100)
);
}