"""异步入库集成验证 在真实内存 Qdrant + FakeOllama 环境下验证异步入库全链路闭环: POST 202 拿 task_id → wait_done 终态 → 任务查询 → 文档详情 → 检索命中 → 删除; 并覆盖两条边界路径: - 失败路径:分类阶段抛错 → failed + error.stage/partial_summary,Redis 失败镜像 TTL=7 天 - Redis 降级路径:无 Redis / Redis 全异常时任务照常完成、查询正常 """ from collections.abc import Callable from typing import Any import pytest from fastapi.testclient import TestClient from qdrant_client import AsyncQdrantClient from app.api.v1 import document as document_module from app.api.v1 import search as search_module from app.config import Settings from app.core.chunker import Chunker from app.core.classifier import Classifier from app.core.ingest_tasks import IngestTaskManager, IngestTaskStatus from app.core.ingestion import Ingester from app.core.query_parser import QueryParser from app.core.retriever import Retriever from app.core.sparse import SparseEncoder from app.core.summarizer import Summarizer from app.main import app from app.models.document import DocumentInput, IngestionResult from app.models.knowledge import load_taxonomy from app.services.qdrant import QdrantService from tests.test_e2e_integration import ( FAKE_CATEGORY, FAKE_L1_SUMMARY, DeterministicEmbedding, FakeCache, FakeOllama, ) class ClassifyFailOllama(FakeOllama): """分类阶段抛错的 FakeOllama(总结等其余行为与正常版一致)""" async def generate(self, prompt: str, json_mode: bool = False) -> str: if "你是知识库分类助手" in prompt: raise RuntimeError("模拟分类模型不可用") return await super().generate(prompt, json_mode=json_mode) class RecordingRedis: """记录 set_json 调用的假 Redis(get_json 恒未命中)""" def __init__(self) -> None: self.set_calls: list[tuple[str, dict[str, Any], int | None]] = [] async def get_json(self, key: str) -> dict[str, Any] | None: return None async def set_json(self, key: str, value: dict[str, Any], ttl: int | None = None) -> bool: self.set_calls.append((key, value, ttl)) return True class BrokenRedis: """全部方法抛异常的假 Redis,验证任务管理器的容错降级""" async def get_json(self, key: str) -> dict[str, Any] | None: raise ConnectionError("模拟 Redis 不可用") async def set_json(self, key: str, value: dict[str, Any], ttl: int | None = None) -> bool: raise ConnectionError("模拟 Redis 不可用") class RecordingIngester: """包装真实 Ingester,记录 progress_cb 回调的阶段序列(接口与 Ingester 一致)""" def __init__(self, ingester: Ingester) -> None: self._ingester = ingester self.stages: list[str] = [] async def ingest( self, doc: DocumentInput, progress_cb: Callable[[str], None] | None = None ) -> IngestionResult: def _cb(stage: str) -> None: self.stages.append(stage) if progress_cb is not None: progress_cb(stage) return await self._ingester.ingest(doc, progress_cb=_cb) async def _make_env(ollama: FakeOllama) -> tuple[QdrantService, Ingester, Retriever]: """构建集成环境:真实组件 + 内存 Qdrant + 指定 FakeOllama + 确定性向量""" qdrant = QdrantService(client=AsyncQdrantClient(location=":memory:")) await qdrant.ensure_collections() taxonomy = load_taxonomy() embedding = DeterministicEmbedding() ingester = Ingester( summarizer=Summarizer(ollama=ollama), # type: ignore[arg-type] classifier=Classifier(ollama=ollama, taxonomy=taxonomy), # type: ignore[arg-type] chunker=Chunker(), embedding=embedding, # type: ignore[arg-type] sparse=SparseEncoder(), qdrant=qdrant, ) retriever = Retriever( qdrant=qdrant, query_parser=QueryParser(ollama=ollama, taxonomy=taxonomy, cache=FakeCache()), # type: ignore[arg-type] embedding=embedding, # type: ignore[arg-type] sparse_encoder=SparseEncoder(), ) return qdrant, ingester, retriever def _structured_doc() -> DocumentInput: """带 Markdown 标题结构的中文长文档(>500 字符,走完整三级总结)""" paragraph = "这是章节正文内容,包含足够多的信息量,用于测试切分与向量化流程。" * 20 text = f"# 安装指南\n{paragraph}\n\n## 环境准备\n{paragraph}\n\n## 安装步骤\n{paragraph}" return DocumentInput(text=text, title="安装文档") def _patch_app( monkeypatch: pytest.MonkeyPatch, manager: IngestTaskManager, qdrant: QdrantService, retriever: Retriever ) -> None: """替换模块级单例:任务管理器 / Qdrant / Retriever,lifespan 建集合改空操作""" async def _noop_ensure_collections(self: QdrantService) -> None: return None monkeypatch.setattr(QdrantService, "ensure_collections", _noop_ensure_collections) monkeypatch.setattr(document_module, "_task_manager", manager) monkeypatch.setattr(document_module, "_qdrant", qdrant) monkeypatch.setattr(search_module, "_retriever", retriever) monkeypatch.setattr(search_module, "get_cache", lambda: FakeCache()) class TestIngestAsyncFullLoop: """全链路闭环:异步入库 → 任务查询 → 文档管理 → 检索 → 删除""" async def test_full_loop(self, monkeypatch: pytest.MonkeyPatch, admin_headers: dict[str, str]): qdrant, ingester, retriever = await _make_env(FakeOllama()) recording = RecordingIngester(ingester) manager = IngestTaskManager(recording, None, Settings()) # type: ignore[arg-type] _patch_app(monkeypatch, manager, qdrant, retriever) doc = _structured_doc() with TestClient(app) as client: # 1. 提交入库:202 + task_id(变更类端点需 admin 认证头) resp_post = client.post( "/api/v1/documents", json={"text": doc.text, "title": doc.title}, headers=admin_headers ) assert resp_post.status_code == 202 body_post = resp_post.json() assert body_post["code"] == 0 assert body_post["data"]["status"] == "pending" task_id = body_post["data"]["task_id"] assert task_id # 2. 等待终态:done + 结果完整 + 阶段序列完整 final = await manager.wait_done(task_id) assert final["status"] == IngestTaskStatus.DONE result = final["result"] assert result["document_id"] assert result["category"] == FAKE_CATEGORY assert result["chunks_count"] >= 1 assert recording.stages == ["summarizing", "classifying", "embedding", "writing"] document_id = result["document_id"] # 3. 任务状态查询:code=0、done、含 result resp_task = client.get(f"/api/v1/documents/tasks/{task_id}") body_task = resp_task.json() assert body_task["code"] == 0 assert body_task["data"]["status"] == "done" assert body_task["data"]["result"]["document_id"] == document_id # 4. 文档详情:入库完成后即可管理 resp_detail = client.get(f"/api/v1/documents/{document_id}") body_detail = resp_detail.json() assert body_detail["code"] == 0 # 5. 检索:相关 query 命中该文档 resp_search = client.post("/api/v1/search", json={"query": "安装步骤有哪些注意事项?"}) body_search = resp_search.json() assert body_search["code"] == 0 hits = body_search["data"]["hits"] assert hits assert any(hit["doc_id"] == document_id for hit in hits) # 6. 删除:四层集合中该文档全部清除 resp_delete = client.delete(f"/api/v1/documents/{document_id}", headers=admin_headers) body_delete = resp_delete.json() assert body_delete["code"] == 0 assert body_delete["data"]["deleted_total"] > 0 class TestIngestAsyncFailure: """失败路径:分类阶段抛错 → failed 状态、错误透传与 Redis 失败镜像 TTL""" async def test_classify_failure(self, monkeypatch: pytest.MonkeyPatch): qdrant, ingester, retriever = await _make_env(ClassifyFailOllama()) redis = RecordingRedis() manager = IngestTaskManager(ingester, redis, Settings()) # type: ignore[arg-type] _patch_app(monkeypatch, manager, qdrant, retriever) # 直接经 manager 提交(与 API 提交同路径),在测试事件循环内等待终态与镜像写完 task_id = await manager.submit(_structured_doc()) final = await manager.wait_done(task_id) # 终态 failed:阶段与已产出总结透传 assert final["status"] == IngestTaskStatus.FAILED error = final["error"] assert error["stage"] == "classify" assert error["partial_summary"] assert error["partial_summary"]["l1_summary"] == FAKE_L1_SUMMARY # Redis 镜像:failed 快照使用失败 TTL(7 天),其余快照用 done TTL(24h) failed_mirrors = [c for c in redis.set_calls if c[1]["status"] == IngestTaskStatus.FAILED] assert failed_mirrors assert all(ttl == 604800 for _, _, ttl in failed_mirrors) non_failed = [c for c in redis.set_calls if c[1]["status"] != IngestTaskStatus.FAILED] assert non_failed assert all(ttl == 86400 for _, _, ttl in non_failed) # API 查询:GET tasks/{task_id} 返回同样的失败信息 with TestClient(app) as client: resp_task = client.get(f"/api/v1/documents/tasks/{task_id}") body_task = resp_task.json() assert body_task["code"] == 0 assert body_task["data"]["status"] == "failed" assert body_task["data"]["error"]["stage"] == "classify" assert body_task["data"]["error"]["partial_summary"] assert body_task["data"]["error"]["partial_summary"]["l1_summary"] == FAKE_L1_SUMMARY class TestIngestAsyncRedisDegraded: """Redis 降级路径:无 Redis 或 Redis 全异常时任务照常完成""" async def test_redis_none(self): """redis=None:纯内存模式,提交→done→查询全流程正常""" _, ingester, _ = await _make_env(FakeOllama()) manager = IngestTaskManager(ingester, None, Settings()) task_id = await manager.submit(_structured_doc()) final = await manager.wait_done(task_id) assert final["status"] == IngestTaskStatus.DONE assert final["result"]["document_id"] record = await manager.get(task_id) assert record is not None assert record["status"] == IngestTaskStatus.DONE async def test_redis_broken(self): """Redis 读写全抛异常:镜像/读取容错降级,任务流程与查询不受影响""" _, ingester, _ = await _make_env(FakeOllama()) manager = IngestTaskManager(ingester, BrokenRedis(), Settings()) # type: ignore[arg-type] task_id = await manager.submit(_structured_doc()) final = await manager.wait_done(task_id) assert final["status"] == IngestTaskStatus.DONE assert final["result"]["document_id"] record = await manager.get(task_id) assert record is not None assert record["status"] == IngestTaskStatus.DONE