442 lines
16 KiB
Python
442 lines
16 KiB
Python
"""知识库管理服务契约测试(Task 11)。
|
||||
|
|
|
|||
|
|
被测的是**接口层用例**:鉴权 → 开事务 → 复用入库链 → 提交;删除则是"标记 + 投递向量删除事件"。
|
|||
|
|
|
|||
|
|
全部用替身,**不连真库、不连 Milvus、不连 Redis**(本环境 Milvus/Redis 不可用,
|
|||
|
|
上传与删除的向量落地由 Outbox Worker 负责,本层只保证事件被正确投出)。
|
|||
|
|
|
|||
|
|
固定口径(见 `app/service/knowledge_management_service.py` 的模块 docstring):
|
|||
|
|
- `created_by` 只来自认证上下文,方法不提供该参数;
|
|||
|
|
- 删除是**标记** `status='expired'` + 投 `knowledge.vector_delete_requested`,不物理删除;
|
|||
|
|
- 删不存在的 id 抛 404 语义异常,不静默成功。
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
import base64
|
|||
|
|
from datetime import UTC, datetime
|
|||
|
|
from pathlib import Path
|
|||
|
|
from typing import Any
|
|||
|
|
|
|||
|
|
import pytest
|
|||
|
|
|
|||
|
|
from app.api.controllers.knowledge_management import knowledge_management_service
|
|||
|
|
from app.core.contracts import RequestContext
|
|||
|
|
from app.core.errors import ForbiddenAgentError, GenericResourceNotFoundError, ValidationAgentError
|
|||
|
|
from app.model.knowledge import KnowledgeMeta
|
|||
|
|
from app.model.platform import DomainEventOutbox
|
|||
|
|
from app.service.knowledge_management_service import (
|
|||
|
|
EXPIRED_STATUS,
|
|||
|
|
REQUIRED_PERMISSION,
|
|||
|
|
KnowledgeManagementService,
|
|||
|
|
decode_content,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
PERMISSION = REQUIRED_PERMISSION
|
|||
|
|
|
|||
|
|
|
|||
|
|
class _AsyncContext:
|
|||
|
|
def __init__(self, value: Any = None, error: BaseException | None = None) -> None:
|
|||
|
|
self._value = value
|
|||
|
|
self._error = error
|
|||
|
|
|
|||
|
|
async def __aenter__(self) -> Any:
|
|||
|
|
return self._value
|
|||
|
|
|
|||
|
|
async def __aexit__(self, *exc: object) -> bool:
|
|||
|
|
if self._error is not None:
|
|||
|
|
raise self._error
|
|||
|
|
return False
|
|||
|
|
|
|||
|
|
|
|||
|
|
class FakeIngest:
|
|||
|
|
"""`KnowledgeIngestService` 替身:记录调用参数,返回预设的 `knowledge_id` 列表。"""
|
|||
|
|
|
|||
|
|
def __init__(self, ids: list[int] | None = None, error: BaseException | None = None) -> None:
|
|||
|
|
self._ids = ids if ids is not None else [11, 12]
|
|||
|
|
self._error = error
|
|||
|
|
self.calls: list[dict[str, Any]] = []
|
|||
|
|
|
|||
|
|
async def ingest(self, **kwargs: Any) -> list[int]:
|
|||
|
|
self.calls.append(kwargs)
|
|||
|
|
if self._error is not None:
|
|||
|
|
raise self._error
|
|||
|
|
return list(self._ids)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def meta_row(
|
|||
|
|
knowledge_id: int = 7,
|
|||
|
|
*,
|
|||
|
|
status: str = "active",
|
|||
|
|
knowledge_type: str = "faq",
|
|||
|
|
storage_key: str | None = "kb/faq/abc-faq.md",
|
|||
|
|
content: str = "正文",
|
|||
|
|
) -> KnowledgeMeta:
|
|||
|
|
now = datetime.now(UTC).replace(tzinfo=None)
|
|||
|
|
return KnowledgeMeta(
|
|||
|
|
id=knowledge_id,
|
|||
|
|
knowledge_type=knowledge_type,
|
|||
|
|
title="标题",
|
|||
|
|
source_file="faq.md",
|
|||
|
|
minio_path=storage_key,
|
|||
|
|
milvus_collection="fin_faq_collection",
|
|||
|
|
version="v1",
|
|||
|
|
content_text=content,
|
|||
|
|
tags={"chunk_index": 0},
|
|||
|
|
reviewer_id=9003,
|
|||
|
|
review_status="published",
|
|||
|
|
status=status,
|
|||
|
|
created_at=now,
|
|||
|
|
updated_at=now,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
|
|||
|
|
class FakeSession:
|
|||
|
|
"""替身 session:记录语句、`add()` 与**编译后的绑定参数**,并把 UPDATE 真的应用到内存行上。
|
|||
|
|
|
|||
|
|
两点是刻意的:
|
|||
|
|
|
|||
|
|
1. `session.execute(update(...).values(...))` 这类 SQLAlchemy 构造式语句的参数**不在**
|
|||
|
|
`execute(params=...)` 里,而是编译进语句本身(`compile().params`)。只记录 `params`
|
|||
|
|
会得到空 dict,让"UPDATE 到底改了什么"变成盲区;
|
|||
|
|
2. 把 UPDATE 落地到内存行:删除的验收点是"提交后这行确实是 `expired`",
|
|||
|
|
只断言"发过一条 UPDATE"接不住"WHERE 条件写错、什么也没改"这类错误。
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
def __init__(self, rows: list[KnowledgeMeta] | None = None) -> None:
|
|||
|
|
self.rows = rows if rows is not None else []
|
|||
|
|
self.statements: list[tuple[str, dict[str, Any]]] = []
|
|||
|
|
self.added: list[Any] = []
|
|||
|
|
self.begin_calls = 0
|
|||
|
|
self.commits = 0
|
|||
|
|
|
|||
|
|
def begin(self) -> _AsyncContext:
|
|||
|
|
self.begin_calls += 1
|
|||
|
|
return _AsyncContext()
|
|||
|
|
|
|||
|
|
async def __aenter__(self) -> "FakeSession":
|
|||
|
|
return self
|
|||
|
|
|
|||
|
|
async def __aexit__(self, *exc: object) -> bool:
|
|||
|
|
return False
|
|||
|
|
|
|||
|
|
async def get(self, model: Any, ident: Any) -> Any:
|
|||
|
|
return next((row for row in self.rows if row.id == ident), None)
|
|||
|
|
|
|||
|
|
async def scalars(self, statement: Any) -> Any:
|
|||
|
|
params = _bound_params(statement)
|
|||
|
|
self.statements.append((str(statement), params))
|
|||
|
|
|
|||
|
|
class Rows:
|
|||
|
|
"""最小 `ScalarResult` 替身:只提供服务实际用到的 `.all()`。"""
|
|||
|
|
|
|||
|
|
def __init__(self, rows: list[KnowledgeMeta]) -> None:
|
|||
|
|
self._rows = rows
|
|||
|
|
|
|||
|
|
def all(self) -> list[KnowledgeMeta]:
|
|||
|
|
return list(self._rows)
|
|||
|
|
|
|||
|
|
return Rows(self.rows_for(statement, params))
|
|||
|
|
|
|||
|
|
def rows_for(self, statement: Any, params: dict[str, Any]) -> list[KnowledgeMeta]:
|
|||
|
|
"""按编译后的 WHERE 语义做**最小**过滤:目前只认"排除某个 status"这一个条件。
|
|||
|
|
|
|||
|
|
为什么值得实现这一层:列表的验收点是"只返回未过期行"。如果替身把整表原样返回,
|
|||
|
|
`count` 就只能证明"计数=行数",证明不了任何过滤;而真机 Milvus/Redis 不可用、
|
|||
|
|
这里也不连 MySQL,所以必须让替身在内存里执行这一条语义。
|
|||
|
|
"""
|
|||
|
|
sql = str(statement)
|
|||
|
|
rows = list(self.rows)
|
|||
|
|
if "fin_knowledge_meta.status !=" in sql:
|
|||
|
|
excluded = next(
|
|||
|
|
(value for key, value in params.items() if key.startswith("status")), None
|
|||
|
|
)
|
|||
|
|
rows = [row for row in rows if row.status != excluded]
|
|||
|
|
type_key = next(
|
|||
|
|
(key for key in params if key.startswith("knowledge_type")), None
|
|||
|
|
)
|
|||
|
|
if type_key is not None:
|
|||
|
|
rows = [row for row in rows if row.knowledge_type == params[type_key]]
|
|||
|
|
return rows
|
|||
|
|
|
|||
|
|
async def execute(self, statement: Any, params: Any = None) -> Any:
|
|||
|
|
values = dict(params or {}) or _bound_params(statement)
|
|||
|
|
self.statements.append((str(statement), values))
|
|||
|
|
if str(statement).upper().startswith("UPDATE"):
|
|||
|
|
for row in self.rows:
|
|||
|
|
for column, value in values.items():
|
|||
|
|
if hasattr(row, column):
|
|||
|
|
setattr(row, column, value)
|
|||
|
|
|
|||
|
|
class Result:
|
|||
|
|
rowcount = len(values) and 1
|
|||
|
|
|
|||
|
|
return Result()
|
|||
|
|
|
|||
|
|
def add(self, value: Any) -> None:
|
|||
|
|
self.added.append(value)
|
|||
|
|
|
|||
|
|
async def commit(self) -> None:
|
|||
|
|
self.commits += 1
|
|||
|
|
|
|||
|
|
|
|||
|
|
def _bound_params(statement: Any) -> dict[str, Any]:
|
|||
|
|
"""取 SQLAlchemy 构造式语句编译后的绑定参数(纯 SQL 字符串没有可编译对象)。"""
|
|||
|
|
compile_ = getattr(statement, "compile", None)
|
|||
|
|
if compile_ is None:
|
|||
|
|
return {}
|
|||
|
|
return dict(compile_().params)
|
|||
|
|
|
|||
|
|
|
|||
|
|
class FakeStorage:
|
|||
|
|
def __init__(self, error: BaseException | None = None) -> None:
|
|||
|
|
self.archived: list[str] = []
|
|||
|
|
self._error = error
|
|||
|
|
|
|||
|
|
async def archive(self, *, key: str) -> None:
|
|||
|
|
if self._error is not None:
|
|||
|
|
raise self._error
|
|||
|
|
self.archived.append(key)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def context(
|
|||
|
|
*, permissions: tuple[str, ...] = (PERMISSION,), user_id: str = "9003"
|
|||
|
|
) -> RequestContext:
|
|||
|
|
return RequestContext(user_id=user_id, trace_id="trace-task11", permissions=permissions)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def build_service(
|
|||
|
|
*,
|
|||
|
|
rows: list[KnowledgeMeta] | None = None,
|
|||
|
|
ingest: FakeIngest | None = None,
|
|||
|
|
storage: FakeStorage | None = None,
|
|||
|
|
) -> tuple[KnowledgeManagementService, FakeSession, FakeIngest, FakeStorage]:
|
|||
|
|
session = FakeSession(rows)
|
|||
|
|
fake_ingest = ingest or FakeIngest()
|
|||
|
|
fake_storage = storage or FakeStorage()
|
|||
|
|
service = KnowledgeManagementService(
|
|||
|
|
ingest_factory=lambda _session: fake_ingest,
|
|||
|
|
session_factory=lambda: session,
|
|||
|
|
storage=fake_storage,
|
|||
|
|
archive_on_delete=False,
|
|||
|
|
)
|
|||
|
|
return service, session, fake_ingest, fake_storage
|
|||
|
|
|
|||
|
|
|
|||
|
|
def delete_events(session: FakeSession) -> list[Any]:
|
|||
|
|
return [
|
|||
|
|
event for event in session.added
|
|||
|
|
if isinstance(event, DomainEventOutbox)
|
|||
|
|
and event.event_type == "knowledge.vector_delete_requested"
|
|||
|
|
]
|
|||
|
|
|
|||
|
|
|
|||
|
|
# --- ① 上传 -----------------------------------------------------------------
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_upload_returns_knowledge_ids_and_takes_created_by_from_context() -> None:
|
|||
|
|
service, session, ingest, _ = build_service(ingest=FakeIngest([11, 12]))
|
|||
|
|
|
|||
|
|
result = await service.upload(
|
|||
|
|
context(user_id="9003"),
|
|||
|
|
filename="faq.md",
|
|||
|
|
content=b"# \xe7\x94\xb3\xe8\xb4\xad",
|
|||
|
|
knowledge_type="faq",
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
assert result["knowledge_ids"] == [11, 12]
|
|||
|
|
assert result["chunk_count"] == 2
|
|||
|
|
assert result["created_by"] == "9003"
|
|||
|
|
assert ingest.calls[0]["created_by"] == 9003 # 取自 context.user_id,不是硬编码
|
|||
|
|
assert ingest.calls[0]["filename"] == "faq.md"
|
|||
|
|
assert ingest.calls[0]["content"] == b"# \xe7\x94\xb3\xe8\xb4\xad"
|
|||
|
|
assert session.begin_calls == 1 # 事务归本服务:提交由 session.begin() 的退出负责
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_upload_rejects_unknown_knowledge_type_without_touching_the_session() -> None:
|
|||
|
|
service, session, ingest, _ = build_service(
|
|||
|
|
ingest=FakeIngest(error=ValidationAgentError("knowledge_type 非法:secret"))
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
with pytest.raises(ValidationAgentError):
|
|||
|
|
await service.upload(
|
|||
|
|
context(), filename="a.txt", content=b"x", knowledge_type="secret"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
assert session.begin_calls == 1
|
|||
|
|
assert session.added == []
|
|||
|
|
assert len(ingest.calls) == 1 # 校验发生在入库链里(白名单只有一份真相)
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_upload_requires_the_management_permission(monkeypatch: pytest.MonkeyPatch) -> None:
|
|||
|
|
"""无权限时必须在开事务之前就被拦下:替身 session 一旦被使用就报错。"""
|
|||
|
|
service, _, _, _ = build_service()
|
|||
|
|
|
|||
|
|
def explode() -> None:
|
|||
|
|
raise AssertionError("未授权请求不得打开数据库会话")
|
|||
|
|
|
|||
|
|
service._session_factory = explode # type: ignore[assignment]
|
|||
|
|
|
|||
|
|
with pytest.raises(ForbiddenAgentError):
|
|||
|
|
await service.upload(context(permissions=()), filename="a.txt", content=b"x",
|
|||
|
|
knowledge_type="faq")
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_upload_rejects_filename_with_path() -> None:
|
|||
|
|
service, _, _, _ = build_service()
|
|||
|
|
|
|||
|
|
with pytest.raises(ValidationAgentError):
|
|||
|
|
await service.upload(context(), filename=r"C:\kb\faq.md", content=b"x",
|
|||
|
|
knowledge_type="faq")
|
|||
|
|
|
|||
|
|
|
|||
|
|
def test_decode_content_rejects_non_base64() -> None:
|
|||
|
|
"""宽容解码会把 `!!!` 解成空字节 → "上传成功但内容为空",必须失败关闭。"""
|
|||
|
|
with pytest.raises(ValidationAgentError):
|
|||
|
|
decode_content("!!!not base64!!!")
|
|||
|
|
|
|||
|
|
assert decode_content(base64.b64encode("正文".encode()).decode()) == "正文".encode()
|
|||
|
|
|
|||
|
|
|
|||
|
|
# --- ② 列表 -----------------------------------------------------------------
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_list_only_returns_rows_that_are_not_expired() -> None:
|
|||
|
|
rows = [meta_row(7, status="active"), meta_row(8, status=EXPIRED_STATUS)]
|
|||
|
|
service, session, _, _ = build_service(rows=rows)
|
|||
|
|
|
|||
|
|
result = await service.list_documents(context(), limit=20, offset=0)
|
|||
|
|
|
|||
|
|
assert [item["knowledge_id"] for item in result["items"]] == [7]
|
|||
|
|
assert result["count"] == 1
|
|||
|
|
# 过滤必须发生在 **SQL 侧**:断言编译后语句里的 WHERE 条件与绑定值。
|
|||
|
|
sql, params = session.statements[0]
|
|||
|
|
assert "fin_knowledge_meta.status !=" in sql
|
|||
|
|
assert EXPIRED_STATUS in params.values()
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_list_never_returns_the_full_content_text() -> None:
|
|||
|
|
service, _, _, _ = build_service(rows=[meta_row(7, content="正" * 500)])
|
|||
|
|
|
|||
|
|
item = (await service.list_documents(context(), limit=20, offset=0))["items"][0]
|
|||
|
|
|
|||
|
|
assert "content_text" not in item
|
|||
|
|
assert item["content_preview"] == "正" * 200
|
|||
|
|
assert item["content_length"] == 500
|
|||
|
|
assert item["created_by"] == 9003
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_list_rejects_unknown_knowledge_type_filter() -> None:
|
|||
|
|
service, session, _, _ = build_service(rows=[meta_row(7)])
|
|||
|
|
|
|||
|
|
with pytest.raises(ValidationAgentError):
|
|||
|
|
await service.list_documents(context(), limit=20, offset=0, knowledge_type="secret")
|
|||
|
|
assert session.statements == []
|
|||
|
|
|
|||
|
|
|
|||
|
|
# --- ③ 删除 -----------------------------------------------------------------
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_delete_marks_expired_and_enqueues_vector_delete_event() -> None:
|
|||
|
|
row = meta_row(7, status="active")
|
|||
|
|
service, session, _, _ = build_service(rows=[row])
|
|||
|
|
|
|||
|
|
result = await service.delete_document(context(), 7)
|
|||
|
|
|
|||
|
|
assert row.status == EXPIRED_STATUS # UPDATE 真的落到该行上
|
|||
|
|
events = delete_events(session)
|
|||
|
|
assert len(events) == 1
|
|||
|
|
event = events[0]
|
|||
|
|
assert event.event_type == "knowledge.vector_delete_requested"
|
|||
|
|
assert event.aggregate_type == "knowledge_meta"
|
|||
|
|
assert event.aggregate_id == "7"
|
|||
|
|
assert event.trace_id == "trace-task11"
|
|||
|
|
assert event.status == "pending"
|
|||
|
|
assert event.retry_count == 0
|
|||
|
|
# payload 是 JSON 列口径:`knowledge_id` 必须是**字符串**(与投递侧同口径,
|
|||
|
|
# 消费侧 `_knowledge_id` 只接受纯数字字符串)。
|
|||
|
|
assert event.payload == {"knowledge_id": "7"}
|
|||
|
|
assert result == {"knowledge_id": 7, "status": EXPIRED_STATUS,
|
|||
|
|
"vector_delete_event": "knowledge.vector_delete_requested"}
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_delete_unknown_id_raises_not_found_and_emits_no_event() -> None:
|
|||
|
|
service, session, _, _ = build_service(rows=[meta_row(7)])
|
|||
|
|
|
|||
|
|
with pytest.raises(GenericResourceNotFoundError):
|
|||
|
|
await service.delete_document(context(), 999)
|
|||
|
|
|
|||
|
|
assert delete_events(session) == []
|
|||
|
|
assert session.commits == 0
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_deleting_an_already_expired_row_is_also_not_found() -> None:
|
|||
|
|
"""重复删除不得再投一次删除事件(否则向量删除会被反复排队)。"""
|
|||
|
|
service, session, _, _ = build_service(rows=[meta_row(7, status=EXPIRED_STATUS)])
|
|||
|
|
|
|||
|
|
with pytest.raises(GenericResourceNotFoundError):
|
|||
|
|
await service.delete_document(context(), 7)
|
|||
|
|
|
|||
|
|
assert delete_events(session) == []
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_archive_is_best_effort_and_does_not_roll_back_the_delete() -> None:
|
|||
|
|
"""归档失败不得影响已提交的删除(否则会出现"库还是 active、文件已归档"的不一致)。"""
|
|||
|
|
row = meta_row(7)
|
|||
|
|
storage = FakeStorage(error=OSError("磁盘满"))
|
|||
|
|
service, session, _, _ = build_service(rows=[row], storage=storage)
|
|||
|
|
service._archive_on_delete = True
|
|||
|
|
|
|||
|
|
await service.delete_document(context(), 7)
|
|||
|
|
|
|||
|
|
assert row.status == EXPIRED_STATUS
|
|||
|
|
assert len(delete_events(session)) == 1
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_archive_key_comes_from_the_knowledge_row() -> None:
|
|||
|
|
row = meta_row(7, storage_key="kb/faq/deadbeef-faq.md")
|
|||
|
|
storage = FakeStorage()
|
|||
|
|
service, _, _, _ = build_service(rows=[row], storage=storage)
|
|||
|
|
service._archive_on_delete = True
|
|||
|
|
|
|||
|
|
await service.delete_document(context(), 7)
|
|||
|
|
|
|||
|
|
assert storage.archived == ["kb/faq/deadbeef-faq.md"]
|
|||
|
|
|
|||
|
|
|
|||
|
|
# --- 装配与路由 -------------------------------------------------------------
|
|||
|
|
|
|||
|
|
|
|||
|
|
async def test_service_factory_tolerates_a_cold_composition_root() -> None:
|
|||
|
|
"""接口工厂应当可构造(Milvus/Redis 不可用不影响装配:这层不连它们)。"""
|
|||
|
|
service = knowledge_management_service()
|
|||
|
|
|
|||
|
|
assert isinstance(service, KnowledgeManagementService)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def test_router_registers_the_three_required_paths() -> None:
|
|||
|
|
"""老师 Phase 1 第 7 条点名的那三个端点必须在 `/api/v1` 下注册。"""
|
|||
|
|
from app.api.controllers.knowledge_management import router
|
|||
|
|
|
|||
|
|
paths = {(route.path, method) for route in router.routes for method in route.methods}
|
|||
|
|
assert ("/api/v1/knowledge/upload", "POST") in paths
|
|||
|
|
assert ("/api/v1/knowledge/list", "GET") in paths
|
|||
|
|
assert ("/api/v1/knowledge/{knowledge_id}", "DELETE") in paths
|
|||
|
|
|
|||
|
|
|
|||
|
|
def test_every_endpoint_depends_on_the_authentication_gate() -> None:
|
|||
|
|
"""硬约束:三个端点都必须依赖 `build_request_context`,禁止匿名入口。"""
|
|||
|
|
from app.api.controllers.knowledge_management import router
|
|||
|
|
|
|||
|
|
for route in router.routes:
|
|||
|
|
names = {
|
|||
|
|
getattr(dependency.call, "__name__", "")
|
|||
|
|
for dependency in route.dependant.dependencies
|
|||
|
|
}
|
|||
|
|
assert "build_request_context" in names, route.path
|
|||
|
|
|
|||
|
|
|
|||
|
|
def test_management_permission_is_registered_in_the_rbac_seed() -> None:
|
|||
|
|
"""权限码必须与 `tools/seed_test_rbac.py` 一致,否则"实现了却一直 403"。"""
|
|||
|
|
seed = Path("tools/seed_test_rbac.py").read_text(encoding="utf-8")
|
|||
|
|
|
|||
|
|
assert f'"{PERMISSION}"' in seed
|