Files
QMDSearch/tests/test_ingest_async_integration.py
T
kplam 51dc8dc4f6 Initial commit: QMDSearch 分层信息检索服务
- FastAPI + Qdrant + Redis + Ollama 技术栈
- L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合)
- 文档三级总结与 2.5 级回退
- query 解析路由与分类
- /admin 管理页面
2026-07-29 21:24:40 +08:00

263 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""异步入库集成验证
在真实内存 Qdrant + FakeOllama 环境下验证异步入库全链路闭环:
POST 202 拿 task_id → wait_done 终态 → 任务查询 → 文档详情 → 检索命中 → 删除;
并覆盖两条边界路径:
- 失败路径:分类阶段抛错 → failed + error.stage/partial_summaryRedis 失败镜像 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 调用的假 Redisget_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 / Retrieverlifespan 建集合改空操作"""
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):
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
resp_post = client.post("/api/v1/documents", json={"text": doc.text, "title": doc.title})
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}")
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