"""IngestTaskManager 单元测试(FakeIngester/FakeRedis,不真实联网)""" import asyncio from collections.abc import Callable from datetime import datetime from typing import Any from app.config import Settings from app.core.ingest_tasks import REDIS_KEY_PREFIX, IngestTaskManager, IngestTaskStatus from app.core.ingestion import IngestionError from app.models.document import DocumentInput, DocumentSummary, IngestionResult, SummaryLevel def make_summary() -> DocumentSummary: """构造固定的三级总结""" return DocumentSummary( l1_summary="一句话总结", l2_outline=None, l3_content_outline="内容大纲", level=SummaryLevel.L3 ) def make_result(doc_id: str = "doc-1") -> IngestionResult: """构造固定的入库结果""" return IngestionResult( document_id=doc_id, summary=make_summary(), category="tech", collection="四层集合", chunks_count=2, tags=["t"], category_confidence=0.9, ) class FakeIngester: """假入库器:上报固定阶段序列;可配置抛错或用闸门阻塞以验证并发限流""" def __init__(self, error: Exception | None = None) -> None: self.error = error self.stages: list[str] = [] self.calls: list[DocumentInput] = [] self.gate: asyncio.Event | None = None self.started = asyncio.Event() async def ingest(self, doc: DocumentInput, progress_cb: Callable[[str], None] | None = None) -> IngestionResult: self.calls.append(doc) self.started.set() if progress_cb is not None: for stage in ("summarizing", "classifying", "embedding", "writing"): progress_cb(stage) self.stages.append(stage) if self.gate is not None: await self.gate.wait() if self.error is not None: raise self.error return make_result() class FakeRedis: """内存版 Redis:记录全部写入(含 TTL),可配置写入抛错模拟故障降级""" def __init__(self, fail_writes: bool = False) -> None: self.fail_writes = fail_writes self.writes: list[tuple[str, dict[str, Any], int | None]] = [] self.store: dict[str, dict[str, Any]] = {} async def set_json(self, key: str, value: dict[str, Any], ttl: int | None = None) -> bool: self.writes.append((key, value, ttl)) if self.fail_writes: raise RuntimeError("redis down") self.store[key] = value return True async def get_json(self, key: str) -> dict[str, Any] | None: return self.store.get(key) async def test_submit_returns_immediately_and_completes() -> None: """submit 立即返回;任务后台跑完为 done,结果完整,阶段序列齐全,Redis 镜像同步""" ingester = FakeIngester() redis = FakeRedis() manager = IngestTaskManager(ingester, redis, Settings()) task_id = await manager.submit(DocumentInput(text="正文", title="标题")) assert isinstance(task_id, str) and len(task_id) == 32 # submit 后立即可查:状态在合法集合内(pending 或已进入某阶段) record = await manager.get(task_id) assert record is not None assert record["status"] in set(IngestTaskStatus) final = await manager.wait_done(task_id, timeout=5) assert final["status"] == IngestTaskStatus.DONE assert final["result"] == make_result().model_dump(mode="json") assert final["error"] is None assert ingester.stages == ["summarizing", "classifying", "embedding", "writing"] # 时间字段为 ISO8601 字符串 datetime.fromisoformat(final["created_at"]) datetime.fromisoformat(final["updated_at"]) # Redis 镜像已写入终态 mirrored = await redis.get_json(f"{REDIS_KEY_PREFIX}{task_id}") assert mirrored is not None assert mirrored["status"] == IngestTaskStatus.DONE async def test_ingestion_error_marks_failed_with_stage_and_partial_summary() -> None: """IngestionError:任务 failed,error.stage 透传,partial_summary 保留""" summary = DocumentSummary(l1_summary="L1", l2_outline=None, l3_content_outline="L3", level=SummaryLevel.L3) ingester = FakeIngester(error=IngestionError("classify", "分类判定失败: boom", summary=summary)) manager = IngestTaskManager(ingester, FakeRedis(), Settings()) task_id = await manager.submit(DocumentInput(text="正文")) final = await manager.wait_done(task_id, timeout=5) assert final["status"] == IngestTaskStatus.FAILED assert final["result"] is None assert final["error"]["stage"] == "classify" assert "boom" in final["error"]["message"] assert final["error"]["partial_summary"] == summary.model_dump(mode="json") async def test_unexpected_error_marks_failed_with_unknown_stage() -> None: """普通异常:任务 failed,error.stage 为 unknown""" ingester = FakeIngester(error=RuntimeError("炸了")) manager = IngestTaskManager(ingester, None, Settings()) task_id = await manager.submit(DocumentInput(text="正文")) final = await manager.wait_done(task_id, timeout=5) assert final["status"] == IngestTaskStatus.FAILED assert final["error"]["stage"] == "unknown" assert "炸了" in final["error"]["message"] async def test_concurrency_limited_by_semaphore() -> None: """并发上限 1 时两个任务串行:第一个放行前第二个不得进入执行""" ingester = FakeIngester() ingester.gate = asyncio.Event() # 首次调用阻塞,验证第二个任务在排队 manager = IngestTaskManager(ingester, None, Settings(ingest_max_concurrency=1)) task1 = await manager.submit(DocumentInput(text="a", title="t1")) task2 = await manager.submit(DocumentInput(text="b", title="t2")) await asyncio.wait_for(ingester.started.wait(), timeout=1) await asyncio.sleep(0.05) # 给第二个任务调度机会 assert len(ingester.calls) == 1 queued = await manager.get(task2) assert queued is not None assert queued["status"] == IngestTaskStatus.PENDING ingester.gate.set() final1 = await manager.wait_done(task1, timeout=5) final2 = await manager.wait_done(task2, timeout=5) assert final1["status"] == IngestTaskStatus.DONE assert final2["status"] == IngestTaskStatus.DONE assert [doc.title for doc in ingester.calls] == ["t1", "t2"] async def test_redis_write_failure_does_not_break_task() -> None: """Redis 写失败仅告警:任务仍正常跑完为 done""" redis = FakeRedis(fail_writes=True) manager = IngestTaskManager(FakeIngester(), redis, Settings()) task_id = await manager.submit(DocumentInput(text="正文")) final = await manager.wait_done(task_id, timeout=5) assert final["status"] == IngestTaskStatus.DONE assert final["result"] is not None assert redis.writes # 确实尝试过写 Redis async def test_redis_mirror_ttl_done_and_failed() -> None: """Redis 镜像 TTL:进行中与 done 用 ttl_done,failed 用 ttl_failed""" settings = Settings(ingest_task_ttl_done=86400, ingest_task_ttl_failed=604800) redis = FakeRedis() manager = IngestTaskManager(FakeIngester(), redis, settings) done_id = await manager.submit(DocumentInput(text="x")) await manager.wait_done(done_id, timeout=5) done_writes = [(v, ttl) for key, v, ttl in redis.writes if key.endswith(done_id)] assert done_writes assert all(ttl == 86400 for _, ttl in done_writes) assert done_writes[-1][0]["status"] == IngestTaskStatus.DONE redis2 = FakeRedis() failing_manager = IngestTaskManager(FakeIngester(error=RuntimeError("boom")), redis2, settings) failed_id = await failing_manager.submit(DocumentInput(text="y")) await failing_manager.wait_done(failed_id, timeout=5) failed_writes = [(v, ttl) for key, v, ttl in redis2.writes if key.endswith(failed_id)] assert failed_writes assert failed_writes[-1][0]["status"] == IngestTaskStatus.FAILED assert failed_writes[-1][1] == 604800 async def test_get_falls_back_to_redis_then_none() -> None: """get:内存 miss 时回查 Redis;两者都没有返回 None""" redis = FakeRedis() manager = IngestTaskManager(FakeIngester(), redis, Settings()) assert await manager.get("不存在") is None redis.store[f"{REDIS_KEY_PREFIX}abc"] = {"task_id": "abc", "status": "done"} record = await manager.get("abc") assert record is not None assert record["status"] == "done"