"""IngestTaskManager 单元测试(FakeIngester/FakeRedis,不真实联网)""" import asyncio import hashlib from collections.abc import Callable from datetime import datetime from typing import Any from app.config import Settings from app.core.ingest_tasks import ( DEDUP_KEY_PREFIX, 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" def _text_hash(text: str) -> str: return hashlib.sha256(text.encode("utf-8")).hexdigest() async def test_dedup_hit_reuses_old_doc_id_and_skips_pipeline() -> None: """去重命中:相同 text 第二次提交直接 done,复用旧 doc_id,deduplicated=True,不调 ingester""" ingester = FakeIngester() redis = FakeRedis() manager = IngestTaskManager(ingester, redis, Settings()) # 首次提交:正常跑流水线 text = "重复内容" task1 = await manager.submit(DocumentInput(text=text, title="t1")) final1 = await manager.wait_done(task1, timeout=5) assert final1["status"] == IngestTaskStatus.DONE assert final1["result"]["document_id"] == "doc-1" assert final1["result"]["deduplicated"] is False assert len(ingester.calls) == 1 # dedup 记录已写入 Redis stored = await redis.get_json(f"{DEDUP_KEY_PREFIX}{_text_hash(text)}") assert stored is not None assert stored["document_id"] == "doc-1" # 第二次提交相同 text:直接 done,不重跑流水线 task2 = await manager.submit(DocumentInput(text=text, title="t2")) # wait_done 会先 await 所有 fire-and-forget 镜像任务再返回,确保 Redis 已写 record2 = await manager.wait_done(task2, timeout=5) assert record2["status"] == IngestTaskStatus.DONE assert record2["result"]["document_id"] == "doc-1" assert record2["result"]["deduplicated"] is True # ingester 没被再次调用 assert len(ingester.calls) == 1 # 镜像已写入 Redis mirrored = await redis.get_json(f"{REDIS_KEY_PREFIX}{task2}") assert mirrored is not None assert mirrored["status"] == IngestTaskStatus.DONE assert mirrored["result"]["deduplicated"] is True async def test_dedup_miss_when_text_differs() -> None: """去重未命中:不同 text 走原 _run,完成后写入对应 dedup key""" ingester = FakeIngester() redis = FakeRedis() manager = IngestTaskManager(ingester, redis, Settings()) task1 = await manager.submit(DocumentInput(text="内容A", title="t1")) await manager.wait_done(task1, timeout=5) task2 = await manager.submit(DocumentInput(text="内容B", title="t2")) await manager.wait_done(task2, timeout=5) assert await redis.get_json(f"{DEDUP_KEY_PREFIX}{_text_hash('内容A')}") is not None assert await redis.get_json(f"{DEDUP_KEY_PREFIX}{_text_hash('内容B')}") is not None assert len(ingester.calls) == 2 async def test_dedup_skipped_when_redis_unavailable() -> None: """Redis 不可用:跳过去重,相同 text 仍走完整流水线""" ingester = FakeIngester() manager = IngestTaskManager(ingester, None, Settings()) task1 = await manager.submit(DocumentInput(text="内容", title="t1")) final1 = await manager.wait_done(task1, timeout=5) assert final1["result"]["deduplicated"] is False task2 = await manager.submit(DocumentInput(text="内容", title="t2")) final2 = await manager.wait_done(task2, timeout=5) assert final2["status"] == IngestTaskStatus.DONE assert final2["result"]["deduplicated"] is False assert len(ingester.calls) == 2 async def test_dedup_lookup_failure_falls_back_to_normal_pipeline() -> None: """Redis get_json 抛错:去重查询降级为未命中,走原流水线""" class _ExplodingRedis(FakeRedis): async def get_json(self, key: str) -> dict[str, Any] | None: raise RuntimeError("redis down") ingester = FakeIngester() redis = _ExplodingRedis() manager = IngestTaskManager(ingester, redis, Settings()) task = await manager.submit(DocumentInput(text="x", title="t")) final = await manager.wait_done(task, timeout=5) assert final["status"] == IngestTaskStatus.DONE assert len(ingester.calls) == 1