- 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, 现场已清理
183 lines
6.4 KiB
Python
183 lines
6.4 KiB
Python
"""jinrong_agent 会话与消息读写(T-06 · agent_session / agent_message,共用底座)。
|
||
|
||
会话权威副本在 MySQL(FLOW §6 / redis-keys §2.1:Redis 只保滑动窗口,丢失
|
||
可重建);seq_no 由调用方经 next_seq_no 取号(同会话并发写防重锁归
|
||
sess:{agent}:{id}:lock,最小闭环单演示进程暂未启用,留痕)。
|
||
T-04 增 agent_tool_call 落库(对话域第三张表;Tool 调用审计,与
|
||
audit_log/input_guard_log 合规审计表区分——本表由对话链路读写)。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from typing import Any
|
||
|
||
from sqlalchemy import text
|
||
from sqlalchemy.engine import Engine
|
||
|
||
from app.config.settings import settings
|
||
from app.utils.db import get_engine
|
||
|
||
|
||
class SessionRepository:
|
||
"""agent_session / agent_message 读写;audit 类表不在本层(T-02 中间件)。"""
|
||
|
||
def __init__(self, engine: Engine | None = None) -> None:
|
||
self._engine = engine or get_engine(settings.mysql_database)
|
||
|
||
# ---------- agent_session ----------
|
||
|
||
def get_session(self, session_id: str) -> dict | None:
|
||
with self._engine.connect() as conn:
|
||
row = conn.execute(
|
||
text("SELECT * FROM agent_session WHERE session_id = :sid"), {"sid": session_id}
|
||
).mappings().first()
|
||
return dict(row) if row else None
|
||
|
||
def create_session(
|
||
self,
|
||
*,
|
||
session_id: str,
|
||
trace_id: str,
|
||
agent_type: str,
|
||
actor_id: str,
|
||
actor_role: str,
|
||
customer_id: str | None,
|
||
advisor_id: str | None = None,
|
||
title: str | None = None,
|
||
) -> dict:
|
||
"""手册 §9:agent_type/actor_id/actor_role/customer_id/advisor_id/trace_id 绑定。"""
|
||
sql = text(
|
||
"""
|
||
INSERT INTO agent_session
|
||
(session_id, trace_id, agent_type, actor_id, actor_role, customer_id, advisor_id, title)
|
||
VALUES (:sid, :trace_id, :atype, :actor_id, :actor_role, :customer_id, :advisor_id, :title)
|
||
"""
|
||
)
|
||
with self._engine.begin() as conn:
|
||
conn.execute(
|
||
sql,
|
||
{
|
||
"sid": session_id,
|
||
"trace_id": trace_id,
|
||
"atype": agent_type,
|
||
"actor_id": actor_id,
|
||
"actor_role": actor_role,
|
||
"customer_id": customer_id,
|
||
"advisor_id": advisor_id,
|
||
"title": title,
|
||
},
|
||
)
|
||
return self.get_session(session_id) # type: ignore[return-value]
|
||
|
||
# ---------- agent_message ----------
|
||
|
||
def next_seq_no(self, session_id: str) -> int:
|
||
with self._engine.connect() as conn:
|
||
current = conn.execute(
|
||
text("SELECT MAX(seq_no) FROM agent_message WHERE session_id = :sid"),
|
||
{"sid": session_id},
|
||
).scalar_one()
|
||
return int(current or 0) + 1
|
||
|
||
def insert_message(
|
||
self,
|
||
*,
|
||
session_id: str,
|
||
trace_id: str,
|
||
seq_no: int,
|
||
role: str,
|
||
content: str,
|
||
has_disclaimer: bool = False,
|
||
token_est: int | None = None,
|
||
) -> None:
|
||
sql = text(
|
||
"""
|
||
INSERT INTO agent_message
|
||
(session_id, trace_id, seq_no, role, content, has_disclaimer, token_est)
|
||
VALUES (:sid, :trace_id, :seq_no, :role, :content, :has_disclaimer, :token_est)
|
||
"""
|
||
)
|
||
with self._engine.begin() as conn:
|
||
conn.execute(
|
||
sql,
|
||
{
|
||
"sid": session_id,
|
||
"trace_id": trace_id,
|
||
"seq_no": seq_no,
|
||
"role": role,
|
||
"content": content,
|
||
"has_disclaimer": 1 if has_disclaimer else 0,
|
||
"token_est": token_est,
|
||
},
|
||
)
|
||
|
||
# ---------- agent_tool_call ----------
|
||
|
||
def insert_tool_call(
|
||
self,
|
||
*,
|
||
session_id: str,
|
||
trace_id: str,
|
||
tool_name: str,
|
||
tool_input: str,
|
||
tool_output: str | None,
|
||
status: str,
|
||
error_code: str | None = None,
|
||
latency_ms: int | None = None,
|
||
) -> None:
|
||
"""Tool 调用留痕(表设计 §01 共用底座;只 INSERT,不更新)。
|
||
|
||
message_id 一期 NULL:Tool 节点先于 LLM 执行,user 消息尚未落库,
|
||
经 session_id + trace_id 已可审计还原;消息级关联归后续增强。
|
||
tool_input/tool_output 由调用方序列化为 JSON 字符串(列类型 JSON)。
|
||
"""
|
||
sql = text(
|
||
"""
|
||
INSERT INTO agent_tool_call
|
||
(session_id, trace_id, message_id, tool_name, tool_input,
|
||
tool_output, status, error_code, latency_ms)
|
||
VALUES (:sid, :trace_id, NULL, :tool_name, :tool_input,
|
||
:tool_output, :status, :error_code, :latency_ms)
|
||
"""
|
||
)
|
||
with self._engine.begin() as conn:
|
||
conn.execute(
|
||
sql,
|
||
{
|
||
"sid": session_id,
|
||
"trace_id": trace_id,
|
||
"tool_name": tool_name,
|
||
"tool_input": tool_input,
|
||
"tool_output": tool_output,
|
||
"status": status,
|
||
"error_code": error_code,
|
||
"latency_ms": latency_ms,
|
||
},
|
||
)
|
||
|
||
def list_messages(self, session_id: str, limit: int = 20) -> list[dict[str, Any]]:
|
||
"""最近 limit 条(seq_no 升序返回,供 LLM 窗口直用)。
|
||
|
||
派生表必须带 alias(MySQL 1248;sqlite 宽松——测试与生产同 SQL 防双语义)。
|
||
"""
|
||
sql = text(
|
||
"""
|
||
SELECT * FROM (
|
||
SELECT * FROM agent_message WHERE session_id = :sid
|
||
ORDER BY seq_no DESC LIMIT :lim
|
||
) AS recent
|
||
ORDER BY recent.seq_no ASC
|
||
"""
|
||
)
|
||
with self._engine.connect() as conn:
|
||
rows = conn.execute(sql, {"sid": session_id, "lim": limit}).mappings().all()
|
||
return [
|
||
{
|
||
"seq_no": r["seq_no"],
|
||
"role": r["role"],
|
||
"content": r["content"],
|
||
"has_disclaimer": bool(r["has_disclaimer"]),
|
||
}
|
||
for r in rows
|
||
]
|