Files
XingHuo/tests/test_chat_tools.py
T
GaoYiYuan_0626 301c78edc3 feat: T-04 Core RO Tool 节点——chat Tool 接入+归属校验+agent_tool_call 落库
- app/tool/core_tools.py: 三只读 Tool(query_customer_profile/holdings/recent_trades, Core RO 仅 SELECT) + TOOL_REGISTRY 白名单(requires_customer); JSON 安全化(Decimal 两位/时间 isoformat)
- app/service/tool_service.py: match_intent 关键词意图(仅 customer/advisor; risk/analyst 空转归 C1/C2) + assert_tool_access 归属断言(customer 本人/advisor assigned/risk_officer 全量/其余拒, 口径对齐 deps.assert_customer_access) + run_tool 编排(白名单→校验→执行→agent_tool_call 落库, 落库失败降级 warning)
- agent_service: 图 START→tool→llm→guard; Tool 结果注入 LLM 上下文(SystemMessage); 降级回复带 Tool 摘要; chat() 可选 session 上下文(缺省空转, T-07 兼容)
- session_repository: insert_tool_call(message_id 一期 NULL, session_id+trace_id 可还原)
- 归属拒绝口径: Tool 层不抛 403 改 blocked 留痕(AUTH_403_* 同码), 对话内呈现——API 层 deny 铁律不变
- tests: test_chat_tools 20 例(三态+落库字段+降级+图注入) + test_chat 4 例(全链路/JWT 外 debug 通道), _ddl 补 agent_tool_call/core_holding, 297 绿
- 真库冒烟: uvicorn+真 MySQL/Redis chat 触发持仓查询 success 落痕 47ms, 现场已清理
2026-09-07 08:43:36 +08:00

408 lines
14 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.
"""T-04 对话 Tool:意图匹配 / 归属校验 / agent_tool_call 落库 / 图节点注入。
单测层:run_tool 直调(sqlite 注入,success/blocked/error 三态与落库字段、
落库降级);图集成:FakeLLM 捕获注入的 [工具查询结果] 上下文、降级回复
带摘要、无会话上下文空转。全链路(TestClient 真路由栈)见 test_chat.py。
"""
from __future__ import annotations
import datetime as _dt
import pytest
from langchain_core.messages import AIMessage
from sqlalchemy import text
from _ddl import create_sqlite_engine
from app.config.settings import settings
from app.repository.core_ro import CoreReadOnlyRepository
from app.repository.session_repository import SessionRepository
from app.service import agent_service, tool_service
class FakeLLM:
def __init__(self, reply: str = "模拟回复"):
self.reply = reply
self.calls: list[list] = []
def invoke(self, messages):
self.calls.append(list(messages))
return AIMessage(content=self.reply)
ACTOR_CUSTOMER = {"actor_id": "CUST-9527", "roles": ["customer"], "token_type": "customer"}
ACTOR_ADVISOR = {"actor_id": "STAFF-10086", "roles": ["advisor"], "token_type": "staff"}
ACTOR_RISK = {"actor_id": "STAFF-30001", "roles": ["risk_officer"], "token_type": "staff"}
ACTOR_ANALYST = {"actor_id": "STAFF-40001", "roles": ["analyst"], "token_type": "staff"}
@pytest.fixture()
def tool_env(monkeypatch):
engine = create_sqlite_engine()
session_repo = SessionRepository(engine=engine)
core_ro = CoreReadOnlyRepository(engine=engine)
with engine.begin() as conn:
conn.execute(
text(
"INSERT INTO core_customer (customer_id, display_name, age, occupation, is_active)"
" VALUES ('CUST-9527', '张三', 45, '工程师', 1)"
)
)
conn.execute(
text(
"INSERT INTO core_customer_risk (customer_id, risk_code, evaluated_at)"
" VALUES ('CUST-9527', 'A', '2026-08-01 10:00:00')"
)
)
conn.execute(
text(
"INSERT INTO core_product (product_id, product_name, min_risk_code, product_type)"
" VALUES ('P-001', '稳健一号', 'A', 'fund')"
)
)
conn.execute(
text(
"INSERT INTO core_holding (customer_id, product_id, market_value, quantity)"
" VALUES ('CUST-9527', 'P-001', 50000.00, 100)"
)
)
conn.execute(
text(
"INSERT INTO core_trade (trade_id, customer_id, product_id, trade_type, amount,"
" trade_status, traded_at) VALUES ('TRD-001', 'CUST-9527', 'P-001', 'subscribe',"
" 10000.00, 'confirmed', :ts)"
),
{"ts": _dt.datetime.now()},
)
conn.execute(
text(
"INSERT INTO core_customer_advisor (advisor_id, customer_id, rel_status)"
" VALUES ('STAFF-10086', 'CUST-9527', 'active')"
)
)
monkeypatch.setattr(tool_service, "_session_repo", lambda: session_repo)
monkeypatch.setattr(tool_service, "_core_ro", lambda: core_ro)
yield {"engine": engine, "session_repo": session_repo}
engine.dispose()
def _tool_rows(engine):
with engine.connect() as conn:
return [
dict(r)
for r in conn.execute(
text("SELECT * FROM agent_tool_call ORDER BY id")
).mappings().all()
]
# ---------- 意图匹配 ----------
def test_match_intent_hits():
assert tool_service.match_intent("customer", "查一下我的持仓") == "query_holdings"
assert tool_service.match_intent("customer", "我的风险测评结果是什么") == "query_customer_profile"
assert tool_service.match_intent("advisor", "看看客户交易记录") == "query_recent_trades"
def test_match_intent_miss_or_agent():
assert tool_service.match_intent("customer", "你好呀") is None
assert tool_service.match_intent("risk", "查一下我的持仓") is None # risk 分支归 C1
assert tool_service.match_intent("analyst", "交易流水") is None
# ---------- run_tool:success / blocked / error 与落库 ----------
def test_run_tool_holdings_success_and_audit(tool_env):
record = tool_service.run_tool(
tool_name="query_holdings",
agent_type="customer",
actor=ACTOR_CUSTOMER,
customer_id="CUST-9527",
session_id="sess-t1",
trace_id="trace-t1",
)
assert record["status"] == "success" and record["error_code"] is None
assert record["data"]["total_count"] == 1
assert record["data"]["sum_market_value"] == 50000.0
assert record["data"]["items"][0]["product_name"] == "稳健一号" # JOIN 生效
rows = _tool_rows(tool_env["engine"])
assert len(rows) == 1
row = rows[0]
assert (row["session_id"], row["trace_id"], row["tool_name"], row["status"]) == (
"sess-t1",
"trace-t1",
"query_holdings",
"success",
)
assert row["error_code"] is None and row["latency_ms"] is not None
assert '"customer_id": "CUST-9527"' in row["tool_input"]
assert '"total_count": 1' in row["tool_output"]
def test_run_tool_customer_profile(tool_env):
record = tool_service.run_tool(
tool_name="query_customer_profile",
agent_type="customer",
actor=ACTOR_CUSTOMER,
customer_id="CUST-9527",
session_id="sess-t2",
)
assert record["status"] == "success"
assert record["data"]["found"] is True
assert record["data"]["risk_code"] == "A"
def test_run_tool_recent_trades_params(tool_env):
record = tool_service.run_tool(
tool_name="query_recent_trades",
agent_type="customer",
actor=ACTOR_CUSTOMER,
customer_id="CUST-9527",
tool_input={"days": 7},
session_id="sess-t3",
)
assert record["status"] == "success"
assert record["data"]["days"] == 7 and record["data"]["total_count"] == 1
def test_run_tool_blocked_not_owner(tool_env):
"""customer 查他人 → blocked(AUTH_403_NOT_OWNER,与 deps 同码)+ 留痕。"""
record = tool_service.run_tool(
tool_name="query_holdings",
agent_type="customer",
actor=ACTOR_CUSTOMER,
customer_id="CUST-1001",
session_id="sess-b1",
)
assert record["status"] == "blocked" and record["error_code"] == "AUTH_403_NOT_OWNER"
assert record["data"] is None
row = _tool_rows(tool_env["engine"])[0]
assert (row["status"], row["error_code"]) == ("blocked", "AUTH_403_NOT_OWNER")
def test_run_tool_blocked_advisor_not_assigned(tool_env):
record = tool_service.run_tool(
tool_name="query_holdings",
agent_type="advisor",
actor=ACTOR_ADVISOR,
customer_id="CUST-1010",
session_id="sess-b2",
)
assert record["status"] == "blocked" and record["error_code"] == "AUTH_403_NOT_ASSIGNED"
def test_run_tool_blocked_scope(tool_env):
"""analyst 等未授权角色 fail-closed(AUTH_403_SCOPE)。"""
record = tool_service.run_tool(
tool_name="query_holdings",
agent_type="customer",
actor=ACTOR_ANALYST,
customer_id="CUST-9527",
session_id="sess-b3",
)
assert record["status"] == "blocked" and record["error_code"] == "AUTH_403_SCOPE"
def test_run_tool_risk_officer_full_access(tool_env):
"""risk_officer 全量(对齐 assert_customer_access 口径)。"""
record = tool_service.run_tool(
tool_name="query_holdings",
agent_type="risk",
actor=ACTOR_RISK,
customer_id="CUST-9527",
session_id="sess-r1",
)
assert record["status"] == "success"
def test_run_tool_blocked_no_customer(tool_env):
record = tool_service.run_tool(
tool_name="query_holdings",
agent_type="customer",
actor=ACTOR_CUSTOMER,
customer_id="",
session_id="sess-n1",
)
assert record["status"] == "blocked" and record["error_code"] == "TOOL_BLOCKED_NO_CUSTOMER"
def test_run_tool_unknown_tool(tool_env):
record = tool_service.run_tool(
tool_name="drop_database",
agent_type="customer",
actor=ACTOR_CUSTOMER,
customer_id="CUST-9527",
session_id="sess-u1",
)
assert record["status"] == "blocked" and record["error_code"] == "TOOL_UNKNOWN"
def test_run_tool_error_swallowed_and_audited(tool_env, monkeypatch):
"""Tool 执行异常 → error 落痕且不向对话链路抛。"""
from app.tool import core_tools
original = core_tools.TOOL_REGISTRY["query_holdings"]["func"]
def boom(customer_id, core_ro):
raise RuntimeError("db exploded")
core_tools.TOOL_REGISTRY["query_holdings"]["func"] = boom
try:
record = tool_service.run_tool(
tool_name="query_holdings",
agent_type="customer",
actor=ACTOR_CUSTOMER,
customer_id="CUST-9527",
session_id="sess-e1",
)
finally:
core_tools.TOOL_REGISTRY["query_holdings"]["func"] = original
assert record["status"] == "error" and record["error_code"] == "TOOL_ERROR"
row = _tool_rows(tool_env["engine"])[0]
assert (row["status"], row["error_code"]) == ("error", "TOOL_ERROR")
def test_run_tool_audit_insert_degrades(tool_env, monkeypatch):
"""留痕失败降级 warning(不阻塞对话,口径同 T-02 审计降级)。"""
def broken_repo():
raise RuntimeError("repo down")
monkeypatch.setattr(tool_service, "_session_repo", broken_repo)
record = tool_service.run_tool(
tool_name="query_holdings",
agent_type="customer",
actor=ACTOR_CUSTOMER,
customer_id="CUST-9527",
session_id="sess-d1",
)
assert record["status"] == "success" # 对话链路不受留痕故障影响
# ---------- 图集成:tool 节点注入与降级 ----------
@pytest.fixture()
def fake_llm(monkeypatch):
llm = FakeLLM()
monkeypatch.setattr(agent_service, "_llm", llm)
monkeypatch.setattr(settings, "deepseek_api_key", "test-key")
yield llm
agent_service.reset_cache()
def test_graph_tool_result_injected_into_llm(fake_llm, tool_env):
out = agent_service.chat(
"customer",
[],
"查一下我的持仓",
session_id="sess-g1",
trace_id="trace-g1",
actor=ACTOR_CUSTOMER,
customer_id="CUST-9527",
)
msgs = fake_llm.calls[0]
assert msgs[1].__class__.__name__ == "SystemMessage"
assert "[工具查询结果]" in msgs[1].content
assert "合计市值 50000" in msgs[1].content
assert '"total_count": 1' in msgs[1].content # 数据 JSON 一并注入
assert len(_tool_rows(tool_env["engine"])) == 1
assert out["tool_results"][0]["status"] == "success"
def test_graph_no_intent_skips_tool(fake_llm, tool_env):
agent_service.chat(
"customer",
[],
"你好",
session_id="sess-g2",
actor=ACTOR_CUSTOMER,
customer_id="CUST-9527",
)
tool_msgs = [
m
for m in fake_llm.calls[0]
if m.__class__.__name__ == "SystemMessage" and "[工具查询结果]" in m.content
]
assert tool_msgs == []
assert _tool_rows(tool_env["engine"]) == []
def test_graph_blocked_result_visible_to_llm(fake_llm, tool_env):
"""归属拒绝以 blocked 结果注入(对话内呈现,非 403)。"""
agent_service.chat(
"customer",
[],
"查一下我的持仓",
session_id="sess-g3",
actor=ACTOR_CUSTOMER,
customer_id="CUST-1001",
)
content = fake_llm.calls[0][1].content
assert "拒绝" in content and "AUTH_403_NOT_OWNER" in content
def test_graph_without_session_context_no_tool(fake_llm, tool_env):
"""无会话上下文(旧调用方式)Tool 空转——T-07 兼容。"""
out = agent_service.chat("customer", [], "查一下我的持仓")
assert out["tool_results"] == []
assert "[工具查询结果]" not in fake_llm.calls[0][1].content
assert _tool_rows(tool_env["engine"]) == []
def test_graph_degraded_reply_includes_summary(tool_env, monkeypatch):
"""无 key 降级:回复携带 Tool 查询摘要(查询不白跑)。"""
monkeypatch.setattr(settings, "deepseek_api_key", "")
agent_service.reset_cache()
out = agent_service.chat(
"customer",
[],
"查一下我的持仓",
session_id="sess-g4",
actor=ACTOR_CUSTOMER,
customer_id="CUST-9527",
)
assert "LLM 未配置" in out["reply"]
assert "合计市值 50000" in out["reply"]
assert out["has_disclaimer"] is True # 降级回复同样经 guard
# ---------- summarize / context_text ----------
def test_summarize_variants(tool_env):
ok = tool_service.run_tool(
tool_name="query_customer_profile",
agent_type="customer",
actor=ACTOR_CUSTOMER,
customer_id="CUST-9527",
session_id="sess-s1",
)
text = tool_service.summarize(ok)
assert "张三" in text and "A" in text
blocked = tool_service.run_tool(
tool_name="query_holdings",
agent_type="customer",
actor=ACTOR_CUSTOMER,
customer_id="CUST-1001",
session_id="sess-s2",
)
assert "无权访问" in tool_service.summarize(blocked)
not_found = tool_service.run_tool(
tool_name="query_customer_profile",
agent_type="risk",
actor=ACTOR_RISK,
customer_id="CUST-NOPE",
session_id="sess-s3",
)
assert "未找到" in tool_service.summarize(not_found)
def test_context_text_empty():
assert tool_service.context_text([]) == ""