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