"""`query_knowledge` 工具单测(Task 7):声明自洽、只读、出参形状与序列化、端点筛选。 全部用 fake retrieval service / fake session —— **不连真 Milvus、不连 MySQL**。 """ from __future__ import annotations import json from typing import Any import pytest from app.core.contracts import RequestContext from app.core.knowledge_contracts import KnowledgeHit, KnowledgeQuery, KnowledgeSearchResult from app.service import knowledge_tool #: 与 `app/service/knowledge_tool.py` 的契约逐字对照的期望值(不从被测模块取,避免自证)。 EXPECTED_NAME = "query_knowledge" EXPECTED_PERMISSION = "knowledge:query" EXPECTED_ROLES = ("customer", "operator", "advisor", "risk_operator", "admin") # --- fakes ------------------------------------------------------------------- class FakeEndpoint: def __init__(self, code: str, capabilities: list[str]) -> None: self.endpoint_code = code self.capabilities = capabilities self.timeout_ms = 5000 class FakeScalars: def __init__(self, rows: list[Any]) -> None: self._rows = rows def __iter__(self) -> Any: return iter(self._rows) class FakeSession: def __init__(self, rows: list[Any]) -> None: self._rows = rows async def __aenter__(self) -> FakeSession: return self async def __aexit__(self, *args: object) -> None: return None async def scalars(self, statement: Any) -> FakeScalars: del statement return FakeScalars(self._rows) class FakeService: """鸭子类型的检索服务:记录 search 实参,返回预置结果。""" def __init__(self, result: KnowledgeSearchResult) -> None: self.result = result self.embedder = object() self.calls: list[dict[str, Any]] = [] async def search( self, query: KnowledgeQuery, *, embedding_endpoints: Any = None, embedder: Any = None, ) -> KnowledgeSearchResult: self.calls.append( {"query": query, "embedding_endpoints": embedding_endpoints, "embedder": embedder} ) return self.result def hit(**overrides: Any) -> KnowledgeHit: payload: dict[str, Any] = { "knowledge_id": "101", "collection": "fin_faq_collection", "title": "场内基金申购费率", "snippet": "正文:场内基金申购费率按成交金额收取。", "score": 0.83, "tags": ("费率", "申购"), "version": "v1", "intent": "faq", } payload.update(overrides) return KnowledgeHit(**payload) def context() -> RequestContext: return RequestContext(user_id="9001", trace_id="task7-trace", roles=("customer",)) def wire( monkeypatch: pytest.MonkeyPatch, *, result: KnowledgeSearchResult | None = None, endpoints: list[Any] | None = None, service: FakeService | None = None, ) -> tuple[FakeService, list[Any]]: """把 `_build_service` / `_embedding_endpoints` 换成 fake,返回服务与端点记录。""" active = service or FakeService( result if result is not None else KnowledgeSearchResult(hits=(hit(),)) ) seen_endpoints: list[Any] = endpoints if endpoints is not None else [FakeEndpoint("e1", [])] def build(config: Any, client: Any = None) -> FakeService: del config, client return active async def embedding_endpoints() -> list[Any]: return seen_endpoints monkeypatch.setattr(knowledge_tool, "_build_service", build) monkeypatch.setattr(knowledge_tool, "_embedding_endpoints", embedding_endpoints) return active, seen_endpoints # --- ① 声明自洽 --------------------------------------------------------------- def test_contract_constants_match_dispatch_contract() -> None: assert knowledge_tool.TOOL_NAME == EXPECTED_NAME assert knowledge_tool.REQUIRED_PERMISSION == EXPECTED_PERMISSION assert knowledge_tool.ALLOWED_ROLES == EXPECTED_ROLES assert knowledge_tool.TIMEOUT_SECONDS > 0 def test_bootstrap_registers_tool_with_contract_permission_and_roles() -> None: """工具必须真的进生产注册表:注册只声明上限,键名/权限/角色错一个字都会失败关闭。""" from app.service.agent.bootstrap import get_agent_factory registry = get_agent_factory()._tool_executor.registry definition = registry.get(EXPECTED_NAME) assert definition.input_model is KnowledgeQuery assert definition.required_permission == EXPECTED_PERMISSION assert tuple(definition.allowed_roles) == EXPECTED_ROLES # 注册表本身拒绝非只读工具,这里再显式断言一次(验收要求)。 assert definition.read_only is True assert definition.timeout_seconds == knowledge_tool.TIMEOUT_SECONDS # --- ② handler 出参 ---------------------------------------------------------- async def test_handler_returns_json_serializable_payload_with_one_hit( monkeypatch: pytest.MonkeyPatch, ) -> None: service, endpoints = wire(monkeypatch) query = KnowledgeQuery(query="申购费率是多少", intents=("faq",)) output = await knowledge_tool.query_knowledge_tool(query, context()) # 出参必须能直接 JSON 化(Agent 结果与审计都要序列化它)。 dumped = json.loads(json.dumps(output, ensure_ascii=False)) assert dumped["degraded"] is False assert dumped["degradation_reason"] is None assert dumped["hits"][0]["knowledge_id"] == "101" assert dumped["hits"][0]["collection"] == "fin_faq_collection" assert dumped["hits"][0]["title"] == "场内基金申购费率" assert dumped["hits"][0]["score"] == pytest.approx(0.83) assert dumped["hits"][0]["tags"] == ["费率", "申购"] # tuple → JSON 数组 assert dumped["hits"][0]["intent"] == "faq" # 入参与端点原样透传给检索服务,工具自身不改写集合路由。 assert service.calls[0]["query"] is query assert service.calls[0]["embedding_endpoints"] == endpoints assert service.calls[0]["embedder"] is service.embedder async def test_handler_shape_is_stable_on_empty_result( monkeypatch: pytest.MonkeyPatch, ) -> None: result = KnowledgeSearchResult(hits=(), searched_collections=("fin_faq_collection",)) wire(monkeypatch, result=result) output = await knowledge_tool.query_knowledge_tool( KnowledgeQuery(query="不存在的知识", intents=("faq",)), context() ) assert output == { "hits": [], "degraded": False, "degradation_reason": None, "searched_collections": ["fin_faq_collection"], } async def test_handler_surfaces_degradation_flag(monkeypatch: pytest.MonkeyPatch) -> None: """Milvus 不可用时检索降级为 MySQL LIKE:调用方必须能看见,不能静默当完整召回。""" result = KnowledgeSearchResult( hits=(hit(intent=None, score=None),), degraded=True, degradation_reason="知识向量检索失败:fin_faq_collection", searched_collections=("fin_faq_collection",), ) wire(monkeypatch, result=result) output = await knowledge_tool.query_knowledge_tool( KnowledgeQuery(query="申购费率是多少", intents=("faq",)), context() ) assert output["degraded"] is True assert output["degradation_reason"] assert output["hits"][0]["score"] is None assert output["hits"][0]["intent"] is None # --- ③ 端点筛选 --------------------------------------------------------------- async def test_embedding_endpoints_filters_out_non_embedding_capability( monkeypatch: pytest.MonkeyPatch, ) -> None: """不筛能力会把向量化请求打到聊天端点上(真库同时存在 chat / embedding 端点)。""" rows = [ FakeEndpoint("chat-primary", ["chat"]), FakeEndpoint("embedding-primary", ["embedding"]), FakeEndpoint("no-capability", []), ] # 生产契约:`_session_factory()` 返回的是**会话工厂本身**(`SessionFactory`), # 调用处再调一次拿到会话。所以替身也必须返回"可调用的工厂",不是会话对象 —— # 否则这个替身会掩盖"少调一层"的真实缺陷(该缺陷曾让 query_knowledge 100% 失败)。 monkeypatch.setattr(knowledge_tool, "_session_factory", lambda: (lambda: FakeSession(rows))) endpoints = await knowledge_tool._embedding_endpoints() assert [endpoint.endpoint_code for endpoint in endpoints] == ["embedding-primary"]