Files
group_fqcd_jr/tests/unit/service/test_knowledge_ingest_service.py
T

421 lines
17 KiB
Python
Raw Normal View History

"""知识入库服务契约测试。
入库链路:文件落存储 → 元数据落库 → 切分 chunk → 逐块写 `fin_knowledge_meta`
→ 投递 `knowledge.vector_sync_requested`(Milvus 写入由 Outbox Worker 负责,本层不连 Milvus)。
固定口径(Task 5 派发说明裁定 1):`ingest` 返回 **`list[int]`**(一份文档切出多块即是多行知识),
集合名只能取白名单映射、调用方不得随意指定。
"""
import json
import re
from typing import Any
import pytest
from app.core.errors import ValidationAgentError
from app.service.document_parser import ParsedChunk
from app.service.knowledge_ingest_service import (
TYPE_TO_COLLECTION,
KnowledgeIngestService,
)
_TEXT_QUERY = re.compile(r"[A-Za-z_]+")
class FakeStorage:
def __init__(self) -> None:
self.saved: list[dict[str, Any]] = []
async def save(self, *, key: str, content: bytes, content_type: str) -> str:
self.saved.append({"key": key, "content": content, "content_type": content_type})
return key
class FakeParser:
def __init__(self, chunks: list[ParsedChunk]) -> None:
self._chunks = chunks
self.calls: list[dict[str, Any]] = []
def parse(self, *, filename: str, content: bytes) -> list[ParsedChunk]:
self.calls.append({"filename": filename, "content": content})
return list(self._chunks)
class FakeSession:
"""记录每条带参数的语句;`SELECT LAST_INSERT_ID()` 按预设序列返回。
`previous_ids` 是"同一 source_file + 集合上已经 active 的行"(导入侧幂等的替身输入):
`_supersede_previous_version` 用 `scalars()` 读它们,本替身直接把它当查询结果返回。
"""
def __init__(self, ids: list[Any] | None = None, previous_ids: list[int] | None = None) -> None:
self._ids = list(ids if ids is not None else [])
self.previous_ids = list(previous_ids if previous_ids is not None else [])
self.statements: list[tuple[str, dict[str, Any]]] = []
self.commits = 0
async def execute(self, statement: Any, params: Any = None) -> None:
# SQLAlchemy 构造式语句(UPDATE ... WHERE id IN (...))的绑定值不在 `params` 里,
# 而是编译进语句本身;不取出来就看不见"到底把哪些行改成了什么状态"。
values = dict(params or {}) or _compiled_params(statement)
self.statements.append((str(statement), values))
return None
async def scalars(self, statement: Any, params: Any = None) -> Any:
self.statements.append((str(statement), dict(params or {}) or _compiled_params(statement)))
return list(self.previous_ids)
async def scalar(self, statement: Any, params: Any = None) -> Any:
self.statements.append((str(statement), dict(params or {})))
return self._ids.pop(0) if self._ids else None
async def commit(self) -> None: # pragma: no cover - 事务归调用方
self.commits += 1
def _compiled_params(statement: Any) -> dict[str, Any]:
"""取构造式语句编译后的绑定参数(纯 SQL 字符串没有可编译对象)。"""
compile_ = getattr(statement, "compile", None)
if compile_ is None:
return {}
return dict(compile_().params)
class RecordingEmbedder:
async def embed(self, endpoints: Any, text: str, *, max_attempts: int = 2) -> Any:
raise AssertionError("入库服务不得自己调用 embedding(向量由 Outbox Worker 负责)")
class RecordingResolver:
async def resolve(self, *, agent_type: str, task_type: str) -> list[Any]:
raise AssertionError("入库服务不得自己解析模型端点(交给 Outbox Worker)")
def _chunks() -> list[ParsedChunk]:
return [
ParsedChunk(text="第一块正文", heading_path=("申购",), index=0),
ParsedChunk(text="第二块正文", heading_path=("申购", "确认"), index=1),
]
def _service(
*,
ids: list[Any] | None = None,
chunks: list[ParsedChunk] | None = None,
created_by: int | None = 9003,
previous_ids: list[int] | None = None,
) -> tuple[KnowledgeIngestService, FakeSession, FakeStorage, FakeParser]:
"""默认给 2 个 chunk 配 2 个自增 id(`LAST_INSERT_ID()` 的替身)。"""
session = FakeSession(ids if ids is not None else [11, 12], previous_ids)
storage = FakeStorage()
parser = FakeParser(chunks if chunks is not None else _chunks())
service = KnowledgeIngestService(
session=session, # type: ignore[arg-type]
parser=parser, # type: ignore[arg-type]
storage=storage,
embedder=RecordingEmbedder(),
endpoint_resolver=RecordingResolver(),
created_by=created_by,
)
return service, session, storage, parser
def _by_table(session: FakeSession, table: str) -> list[dict[str, Any]]:
prefix = f"INSERT INTO {table}"
return [params for sql, params in session.statements if sql.startswith(prefix)]
async def test_ingest_returns_one_id_per_chunk() -> None:
"""裁定 1:返回值是 list[int],一份文档多块 → 多行知识、多条事件。"""
service, session, _, _ = _service(ids=[11, 12])
ids = await service.ingest(filename="faq.md", content=b"# \xe7\x94\xb3\xe8\xb4\xad",
knowledge_type="faq")
assert ids == [11, 12]
assert len(_by_table(session, "fin_knowledge_meta")) == 2
assert len(_by_table(session, "domain_event_outbox")) == 2
async def test_ingest_writes_allowlisted_collection_and_chunk_tags() -> None:
service, session, _, _ = _service(ids=[11, 12])
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq",
collection="fin_faq_collection")
rows = _by_table(session, "fin_knowledge_meta")
assert {row["milvus_collection"] for row in rows} == {"fin_faq_collection"}
first = json.loads(rows[0]["tags"])
assert first["heading_path"] == ["申购"]
assert first["chunk_index"] == 0
assert json.loads(rows[1]["tags"])["chunk_index"] == 1
assert rows[0]["title"] == "第一块正文"
assert rows[0]["reviewer_id"] == 9003
assert rows[1]["reviewer_id"] == 9003
@pytest.mark.parametrize("knowledge_type", ["faq", "product", "policy"])
async def test_every_allowlisted_type_maps_to_its_collection(knowledge_type: str) -> None:
service, session, _, _ = _service(ids=[11, 12])
await service.ingest(filename="a.txt", content=b"x", knowledge_type=knowledge_type)
assert _by_table(session, "fin_knowledge_meta")[0]["milvus_collection"] == (
TYPE_TO_COLLECTION[knowledge_type]
)
async def test_ingest_rejects_collection_outside_the_type_map() -> None:
"""全局约束:集合名不得由调用方随意指定,只能取白名单。"""
service, session, _, _ = _service()
with pytest.raises(ValidationAgentError):
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq",
collection="fin_policy_collection")
assert session.statements == []
async def test_ingest_rejects_unknown_knowledge_type() -> None:
service, _, _, _ = _service()
with pytest.raises(ValidationAgentError):
await service.ingest(filename="faq.md", content=b"x", knowledge_type="secret")
async def test_ingest_enqueues_sync_event_with_string_knowledge_id() -> None:
"""生产口径:payload 与 aggregate_id 都是 `str(knowledge_id)`(裁定 4)。"""
service, session, _, _ = _service(ids=[11, 12])
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq")
events = _by_table(session, "domain_event_outbox")
assert [event["aggregate_id"] for event in events] == ["11", "12"]
assert [json.loads(event["payload"])["knowledge_id"] for event in events] == ["11", "12"]
assert {event["event_type"] for event in events} == {"knowledge.vector_sync_requested"}
assert {event["aggregate_type"] for event in events} == {"knowledge_meta"}
async def test_event_type_constant_is_shared_with_worker() -> None:
from app.worker.knowledge_vector_worker import VECTOR_SYNC_EVENT
service, session, _, _ = _service(ids=[11, 12])
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq")
assert _by_table(session, "domain_event_outbox")[0]["event_type"] == VECTOR_SYNC_EVENT
async def test_ingest_saves_document_before_splitting() -> None:
service, _, storage, parser = _service(ids=[11, 12])
await service.ingest(filename=r"C:\kb\faq.md", content=b"body", knowledge_type="faq")
assert storage.saved[0]["content"] == b"body"
assert storage.saved[0]["key"].startswith("kb/faq/")
assert storage.saved[0]["key"].endswith("faq.md")
assert storage.saved[0]["key"].count("-") == 1
assert parser.calls == [{"filename": r"C:\kb\faq.md", "content": b"body"}]
async def test_ingest_uses_caller_created_by_when_given() -> None:
service, session, _, _ = _service(ids=[11, 12], created_by=9003)
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq", created_by=9004)
assert {row["reviewer_id"] for row in _by_table(session, "fin_knowledge_meta")} == {9004}
async def test_ingest_fails_when_no_created_by_anywhere() -> None:
service, session, _, _ = _service(ids=[11, 12], created_by=None)
with pytest.raises(ValidationAgentError):
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq")
assert session.statements == []
async def test_ingest_does_not_commit_or_embed() -> None:
"""事务与服务边界:ingest 只挂 SQL,commit 归调用方;向量化归 Outbox Worker。"""
service, session, _, _ = _service(ids=[11, 12])
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq")
assert session.commits == 0
async def test_empty_document_produces_no_rows() -> None:
service, session, _, _ = _service(ids=[], chunks=[])
assert await service.ingest(filename="empty.md", content=b"", knowledge_type="faq") == []
assert session.statements == []
class FailingInsertSession(FakeSession):
"""模拟"INSERT 没生效":`rowcount == 0`,且 `LAST_INSERT_ID()` 是 0。
这与 `mysql` 真机行为一致 —— 插入 0 行时 `LAST_INSERT_ID()` 保持**上一次**的值
(新连接上是 0),因此"拿 0"的真相是"插入没生效",报错文案必须这么说。
"""
async def execute(self, statement: Any, params: Any = None) -> Any:
self.statements.append((str(statement), dict(params or {})))
class Result:
rowcount = 0
return Result()
class ZeroIdSession(FakeSession):
"""模拟 `LAST_INSERT_ID()` 返回 0(连接上还没有过自增插入)但 INSERT 确实影响了 1 行。"""
async def execute(self, statement: Any, params: Any = None) -> Any:
self.statements.append((str(statement), dict(params or {})))
class Result:
rowcount = 1
return Result()
async def scalar(self, statement: Any, params: Any = None) -> Any:
self.statements.append((str(statement), dict(params or {})))
return 0
async def test_insert_not_effective_is_reported_as_such() -> None:
"""`rowcount == 0` 时必须在读 `LAST_INSERT_ID()` 之前就失败,且文案指向"插入没生效"。"""
service, _, _, _ = _service(ids=[11, 12])
service._session = FailingInsertSession([11, 12]) # type: ignore[assignment]
with pytest.raises(ValidationAgentError, match="插入未生效"):
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq")
async def test_zero_last_insert_id_is_reported_as_zero() -> None:
"""`LAST_INSERT_ID()` 为 0 时的文案必须点出 0,不能笼统说"取不到"(否则误导排障)。"""
service, _, _, _ = _service(ids=[11, 12])
service._session = ZeroIdSession([11, 12]) # type: ignore[assignment]
with pytest.raises(ValidationAgentError, match="LAST_INSERT_ID"):
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq")
async def test_last_insert_id_is_read_once_per_chunk_in_the_same_statement_order() -> None:
"""硬耦合护栏:每个 chunk 的 `INSERT` 后**紧随**一次 `LAST_INSERT_ID()`,中间不得有别的语句。
这是 `LAST_INSERT_ID()` 是**连接级隐式状态**的直接后果:两者之间只要插入任何其它 SQL,
拿到的就是别人的 id。本用例把"成对且相邻"固定下来,防止以后有人在中间加一条查询。
"""
service, session, _, _ = _service(ids=[11, 12])
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq")
shape = [
"INSERT INTO fin_knowledge_meta" if sql.startswith("INSERT INTO fin_knowledge_meta")
else "SELECT LAST_INSERT_ID()" if "LAST_INSERT_ID" in sql
else "INSERT INTO domain_event_outbox" if sql.startswith("INSERT INTO domain_event_outbox")
else "SELECT previous active ids" if "fin_knowledge_meta.id" in sql
else sql
for sql, _ in session.statements
]
assert shape == [
# 导入侧幂等的第一步:查同 source_file + 集合的 active 旧版(本次为空,故没有 UPDATE)。
"SELECT previous active ids",
"INSERT INTO fin_knowledge_meta",
"SELECT LAST_INSERT_ID()",
"INSERT INTO domain_event_outbox",
"INSERT INTO fin_knowledge_meta",
"SELECT LAST_INSERT_ID()",
"INSERT INTO domain_event_outbox",
]
# --- 导入侧幂等:同一份文档重传 = 覆盖上一版 ---------------------------------------
def _vector_events(session: FakeSession, event_type: str) -> list[dict[str, Any]]:
return [
params
for _, params in session.statements
if params.get("event_type") == event_type
]
async def test_reingesting_the_same_file_expires_the_previous_version() -> None:
"""重传同名文档必须**先下线上一版并逐块投出向量删除事件**。
不这么做的话,旧版向量会留在 Milvus 里继续参与排序(检索侧不看 status),
表现就是"同一份内容有 N 个副本互相抢答"、客服的"领先次优 ≥0.07"门槛被永远卡死。
"""
service, session, _, _ = _service(ids=[11, 12], previous_ids=[3, 4, 5])
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq")
updates = [
(sql, params) for sql, params in session.statements if sql.startswith("UPDATE")
]
assert len(updates) == 1
_, params = updates[0]
assert params["status"] == "expired"
# `id.in_([...])` 编译后是一个列表参数(参数名由 SQLAlchemy 生成,不写死名字)。
id_lists = [value for value in params.values() if isinstance(value, list)]
assert id_lists == [[3, 4, 5]]
deletes = _vector_events(session, "knowledge.vector_delete_requested")
assert [event["aggregate_id"] for event in deletes] == ["3", "4", "5"]
assert [json.loads(event["payload"]) for event in deletes] == [
{"knowledge_id": "3"}, {"knowledge_id": "4"}, {"knowledge_id": "5"}
]
assert {event["aggregate_type"] for event in deletes} == {"knowledge_meta"}
# 新版自己的同步事件不受影响。
assert [event["aggregate_id"]
for event in _vector_events(session, "knowledge.vector_sync_requested")] == ["11", "12"]
async def test_supersede_happens_before_the_new_rows_are_written() -> None:
"""顺序是硬约束:先下线旧版(含投删除事件)、再写新版。
反过来的话,"新行插入失败"会留下一份"旧版已下线、新版没写成"的文档(检索侧彻底查不到);
代价不对称,所以顺序要在测试里固定下来。
"""
service, session, _, _ = _service(ids=[11, 12], previous_ids=[3])
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq")
def kind(index: int) -> str:
sql = session.statements[index][0]
if sql.startswith("UPDATE"):
return "update"
if "fin_knowledge_meta.id" in sql:
return "select_previous"
if sql.startswith("INSERT INTO fin_knowledge_meta"):
return "insert_row"
if sql.startswith("INSERT INTO domain_event_outbox"):
return "insert_event"
return "other"
kinds = [kind(index) for index in range(len(session.statements))]
assert kinds[0] == "select_previous"
assert kinds.index("update") < kinds.index("insert_row")
# 旧版那条删除事件也在新版第一行之前投出。
delete_index = next(
index for index, (_, params) in enumerate(session.statements)
if params.get("event_type") == "knowledge.vector_delete_requested"
)
assert delete_index < kinds.index("insert_row")
async def test_first_import_enqueues_no_vector_delete() -> None:
"""首次导入(没有旧版)不得投出任何删除事件——否则会把别人的向量删掉。"""
service, session, _, _ = _service(ids=[11, 12], previous_ids=[])
await service.ingest(filename="faq.md", content=b"x", knowledge_type="faq")
assert _vector_events(session, "knowledge.vector_delete_requested") == []
assert not [sql for sql, _ in session.statements if sql.startswith("UPDATE")]
async def test_empty_document_must_not_expire_the_previous_version() -> None:
"""空文档(解析出 0 块)不得触发下线:否则"旧版下架了、新版一行没写",文档凭空消失。"""
service, session, _, _ = _service(ids=[], chunks=[], previous_ids=[3, 4])
assert await service.ingest(filename="empty.md", content=b"", knowledge_type="faq") == []
assert session.statements == []