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

231 lines
8.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""`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"]