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

442 lines
16 KiB
Python
Raw Normal View History

"""知识库管理服务契约测试(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