Files
XingHuo/app/repository/session_repository.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

183 lines
6.4 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.
"""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
]