231 lines
8.3 KiB
Python
231 lines
8.3 KiB
Python
"""`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"]
|