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

231 lines
8.3 KiB
Python
Raw Normal View History

"""`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"]