Files
XingHuo/app/tool/core_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

115 lines
4.0 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.
"""Core 模拟库只读 Tool(T-04 · FLOW §2「CoreReadOnlyRepository:持仓/流水/L0」)。
Tool 定义层(纯查询,无业务流程):customer_id 由调用方(tool_service)
注入,**不接受 LLM 生成**——归属防线之一(A-01 语义,T-03 之前的止损)。
注册表 TOOL_REGISTRY 是对话 Tool 的唯一白名单(tool_service 校验)。
requires_customer:Tool 是否必须绑定会话客户(True → 无归属主体即
blocked,不触发查询)。风控分支 Tool(C1 chat_tools)注册时对台账类
统计用 requires_customer=False,复用同一 runner。
返回值约定:JSON 安全 dict(Decimal→float 两位、datetime/date→isoformat),
落库 tool_output 与 LLM 上下文共用同一结构,不做第二套序列化。
"""
from __future__ import annotations
import datetime as _dt
from decimal import Decimal
from typing import Any, Callable
from app.repository.core_ro import CoreReadOnlyRepository
DETAIL_LIMIT = 20 # 明细条数上限(上下文与落库同限,防超长)
DEFAULT_TRADE_DAYS = 30
def _jsonable(value: Any) -> Any:
"""MySQL 行 → JSON 安全结构(递归;Decimal 两位小数、时间 isoformat)。"""
if isinstance(value, Decimal):
return float(round(value, 2))
if isinstance(value, (_dt.datetime, _dt.date)):
return value.isoformat()
if isinstance(value, dict):
return {k: _jsonable(v) for k, v in value.items()}
if isinstance(value, (list, tuple)):
return [_jsonable(v) for v in value]
return value
def query_customer_profile(
customer_id: str, core_ro: CoreReadOnlyRepository | None = None
) -> dict[str, Any]:
"""客户档案与风险测评(L0):core_customer + core_customer_risk。"""
repo = core_ro or CoreReadOnlyRepository()
row = repo.get_customer_l0(customer_id)
if row is None:
return {"found": False, "customer_id": customer_id}
return {"found": True, **_jsonable(row)}
def query_holdings(
customer_id: str, core_ro: CoreReadOnlyRepository | None = None
) -> dict[str, Any]:
"""持仓明细(按市值降序)+ 合计(sum_market_value)。"""
repo = core_ro or CoreReadOnlyRepository()
rows = repo.list_holdings(customer_id)
total = sum((r.get("market_value") or 0) for r in rows)
return {
"total_count": len(rows),
"sum_market_value": _jsonable(total),
"items": [_jsonable(r) for r in rows[:DETAIL_LIMIT]],
}
def query_recent_trades(
customer_id: str,
days: int = DEFAULT_TRADE_DAYS,
core_ro: CoreReadOnlyRepository | None = None,
) -> dict[str, Any]:
"""近 N 天 confirmed 申赎流水(时间升序截断至 DETAIL_LIMIT)。"""
repo = core_ro or CoreReadOnlyRepository()
end = _dt.datetime.now()
start = end - _dt.timedelta(days=days)
rows = repo.list_trades_range(customer_id, start, end)
total = sum((r.get("amount") or 0) for r in rows)
return {
"days": days,
"total_count": len(rows),
"sum_amount": _jsonable(total),
"items": [_jsonable(r) for r in rows[:DETAIL_LIMIT]],
}
class ToolSpec(dict):
"""注册表条目:func + 描述 + 是否必须绑定会话客户。"""
TOOL_REGISTRY: dict[str, ToolSpec] = {
"query_customer_profile": ToolSpec(
func=query_customer_profile,
description="查询客户档案与风险测评等级(L0)",
requires_customer=True,
),
"query_holdings": ToolSpec(
func=query_holdings,
description="查询客户持仓明细与合计市值",
requires_customer=True,
),
"query_recent_trades": ToolSpec(
func=query_recent_trades,
description="查询客户近期申赎交易流水",
requires_customer=True,
),
}
def get_tool(name: str) -> ToolSpec | None:
"""白名单查找(未知 Tool 一律 None,由 runner 拒绝)。"""
return TOOL_REGISTRY.get(name)
def tool_func(name: str) -> Callable[..., dict[str, Any]] | None:
spec = TOOL_REGISTRY.get(name)
return spec["func"] if spec else None