Files
group_xinghuo_jinrong/tests/test_kb_tools.py
T

230 lines
9.5 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.
"""T21-5 知识库 Tool 单测:注册表合并 / skip_access_check / 意图扩展 / 带参注入。
覆盖:
1. kb_tools.search_knowledge Tool 函数(rag_service mock,返回结构);
2. get_registered_tool 三层合并(core → risk → kb);
3. run_tool 对 skip_access_check 的处理(advisor 无绑定客户 → success 而非
AUTH_403 blocked;未知角色也放行——公开知识无归属语义,意图层已限范围);
4. match_intent:customer/advisor 命中 kb 词、Core 词优先、risk 不命中;
5. tool_node 按 spec 白名单注入 query(agent_service 集成,FakeLLM 模式同既有测试)。
"""
from __future__ import annotations
import pytest
from app.service import rag_service, tool_service
from app.service.agent_service import tool_node
from app.tool import kb_tools
@pytest.fixture()
def fake_rag(monkeypatch):
"""mock rag_service.search_knowledge(不连 Ollama/Milvus)。"""
calls: dict = {}
def fake_search(query, *, product_id=None, doc_type=None, top_k=3):
calls.update({"query": query, "top_k": top_k})
return {
"query": query,
"results": [
{
"id": "PROD-110022_1",
"score": 0.82,
"product_id": "PROD-110022",
"product_name": "稳健债基 A",
"doc_type": "rule",
"risk_level": "R1",
"source_doc_id": "KB-PROD-110022",
"source_version": "2026.09",
"effective_date": "2026-09-01",
"chunk_text": "申购 T 日 15:00 前提交按当日净值确认,T+1 确认份额。",
"chunk_no": 1,
}
],
"source_refs": [
{
"source_doc_id": "KB-PROD-110022",
"source_version": "2026.09",
"product_id": "PROD-110022",
"product_name": "稳健债基 A",
}
],
}
monkeypatch.setattr(rag_service, "search_knowledge", fake_search)
return calls
CUSTOMER_ACTOR = {"actor_id": "CUST-1", "roles": ["customer"], "token_type": "dev"}
ADVISOR_ACTOR = {"actor_id": "ADV-1", "roles": ["advisor"], "token_type": "dev"}
class TestToolFunction:
def test_returns_hits_and_refs(self, fake_rag):
out = kb_tools.search_knowledge(query="申购规则")
assert out["hit_count"] == 1
assert out["results"][0]["source_doc_id"] == "KB-PROD-110022"
assert fake_rag["query"] == "申购规则"
def test_registry_shape(self):
spec = kb_tools.KB_TOOL_REGISTRY["search_knowledge"]
assert spec["requires_customer"] is False
assert spec["skip_access_check"] is True
assert "query" in spec["param_whitelist"]
class TestRegistryMerge:
def test_three_layer_merge(self):
# core / risk / kb 各取一个代表
assert tool_service.get_registered_tool("query_holdings") is not None
assert tool_service.get_registered_tool("alert_query") is not None
assert tool_service.get_registered_tool("search_knowledge") is not None
assert tool_service.get_registered_tool("no_such_tool") is None
class TestRunToolSkipAccess:
def test_advisor_without_customer_succeeds(self, fake_rag, monkeypatch):
# 核心:advisor 无绑定客户(customer_id="")查知识库必须 success
# (无 skip_access_check 时会被 assert_tool_access 拒 AUTH_403_NOT_ASSIGNED)
monkeypatch.setattr(tool_service, "_core_ro", lambda: object())
monkeypatch.setattr(tool_service, "_risk_repo", lambda: object())
monkeypatch.setattr(
tool_service, "_session_repo", lambda: type("R", (), {"insert_tool_call": lambda *a, **k: None})()
)
record = tool_service.run_tool(
tool_name="search_knowledge",
agent_type="advisor",
actor=ADVISOR_ACTOR,
customer_id="",
tool_input={"query": "基金申购费率"},
session_id="S1",
)
assert record["status"] == "success"
assert record["data"]["hit_count"] == 1
def test_unknown_role_also_passes(self, fake_rag, monkeypatch):
# 公开知识:非四角色(如 platform 误调)也放行——意图层已限定开放范围
monkeypatch.setattr(tool_service, "_core_ro", lambda: object())
monkeypatch.setattr(tool_service, "_risk_repo", lambda: object())
monkeypatch.setattr(
tool_service, "_session_repo", lambda: type("R", (), {"insert_tool_call": lambda *a, **k: None})()
)
record = tool_service.run_tool(
tool_name="search_knowledge",
agent_type="customer",
actor={"actor_id": "X", "roles": ["unknown_role"], "token_type": "dev"},
customer_id="",
tool_input={"query": "费率"},
session_id="S1",
)
assert record["status"] == "success"
def test_bad_param_still_blocked(self, fake_rag, monkeypatch):
# 白名单外参数仍拒(skip_access_check 只豁免归属,不豁免入参校验)
monkeypatch.setattr(tool_service, "_core_ro", lambda: object())
monkeypatch.setattr(tool_service, "_risk_repo", lambda: object())
monkeypatch.setattr(
tool_service, "_session_repo", lambda: type("R", (), {"insert_tool_call": lambda *a, **k: None})()
)
record = tool_service.run_tool(
tool_name="search_knowledge",
agent_type="customer",
actor=CUSTOMER_ACTOR,
customer_id="CUST-1",
tool_input={"customer_id": "HACK"},
session_id="S1",
)
assert record["status"] == "blocked"
assert record["error_code"] == "TOOL_BAD_PARAM"
class TestIntent:
def test_customer_kb_keywords(self):
assert tool_service.match_intent("customer", "稳健债基的申购费率是多少") == "search_knowledge"
assert tool_service.match_intent("customer", "赎回几天到账") == "search_knowledge"
assert tool_service.match_intent("advisor", "这只基金的定投起点") == "search_knowledge"
def test_core_keywords_take_priority(self):
# Core 词更特异(组序在前):持仓类问句不被 kb 词抢走
assert tool_service.match_intent("customer", "看一下我的持仓") == "query_holdings"
assert tool_service.match_intent("customer", "我的风险测评结果") == "query_customer_profile"
def test_risk_branch_no_kb(self):
# 拍板:风控不开放知识检索
assert tool_service.match_intent("risk", "基金申购费率") is None
assert tool_service.match_intent("analyst", "基金申购费率") is None
class TestToolNodeInjection:
"""tool_node 按 spec 白名单注入 query(state→tool_input)。"""
def _state(self, agent_type="customer", msg="基金申购费率"):
return {
"agent_type": agent_type,
"history": [],
"user_message": msg,
"messages": [],
"reply": "",
"has_disclaimer": False,
"session_id": "S1",
"trace_id": "T1",
"actor": CUSTOMER_ACTOR,
"customer_id": "CUST-1",
"tool_results": [],
}
def test_query_injected_from_user_message(self, fake_rag, monkeypatch):
captured: dict = {}
def fake_run_tool(**kwargs):
captured.update(kwargs)
return {"tool_name": kwargs["tool_name"], "status": "success", "error_code": None,
"data": {"hit_count": 0, "results": [], "source_refs": []}, "latency_ms": 1}
monkeypatch.setattr(tool_service, "run_tool", fake_run_tool)
tool_node(self._state())
assert captured["tool_input"] == {"query": "基金申购费率"}
def test_no_input_for_core_tools(self, fake_rag, monkeypatch):
# Core Tool(白名单无 query)不注入 tool_input,维持 T-04 口径
captured: dict = {}
def fake_run_tool(**kwargs):
captured.update(kwargs)
return {"tool_name": kwargs["tool_name"], "status": "success", "error_code": None,
"data": {}, "latency_ms": 1}
monkeypatch.setattr(tool_service, "run_tool", fake_run_tool)
tool_node(self._state(msg="看一下我的持仓"))
assert captured["tool_input"] is None
class TestSummarize:
def test_kb_summary_with_refs(self):
record = {
"tool_name": "search_knowledge",
"status": "success",
"error_code": None,
"data": {
"hit_count": 1,
"results": [{
"product_name": "稳健债基 A", "doc_type": "rule", "score": 0.82,
"chunk_text": "申购 T 日 15:00 前提交按当日净值确认,T+1 确认份额。本金部分。",
}],
"source_refs": [{"source_doc_id": "KB-PROD-110022", "source_version": "2026.09"}],
},
}
text = tool_service.summarize(record)
assert "命中 1 条" in text
assert "稳健债基 A" in text
assert "KB-PROD-110022@2026.09" in text
def test_kb_summary_empty(self):
record = {
"tool_name": "search_knowledge", "status": "success", "error_code": None,
"data": {"hit_count": 0, "results": [], "source_refs": []},
}
text = tool_service.summarize(record)
assert "未命中" in text and "不要编造" in text