Files
group_xinghuo_jinrong/app/repository/session_repository.py
T

340 lines
13 KiB
Python
Raw Normal View History

"""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, "rw")
# ---------- 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 list_sessions(
self, *, actor_id: str, agent_type: str, limit: int = 20, offset: int = 0
) -> tuple[list[dict[str, Any]], int]:
"""前端会话列表(方案 B):仅本人 + 本 Agent 线,created_at 倒序分页。
返回 (items, total);total 供前端分页器。id 倒序兜底同秒并发建的
会话排序稳定(created_at 精度秒级时并列)。datetime 统一转 str——
sqlite 返 str、MySQL 返 datetime,响应体跨库同构。
"""
where = "WHERE actor_id = :actor AND agent_type = :atype"
with self._engine.connect() as conn:
total = int(
conn.execute(
text(f"SELECT COUNT(*) FROM agent_session {where}"),
{"actor": actor_id, "atype": agent_type},
).scalar_one()
)
rows = conn.execute(
text(
f"""
SELECT session_id, agent_type, actor_id, actor_role, customer_id,
advisor_id, title, status, created_at, closed_at
FROM agent_session {where}
ORDER BY created_at DESC, id DESC
LIMIT :lim OFFSET :off
"""
),
{"actor": actor_id, "atype": agent_type, "lim": limit, "off": offset},
).mappings().all()
items = [
{
**{k: r[k] for k in (
"session_id", "agent_type", "actor_id", "actor_role",
"customer_id", "advisor_id", "title", "status",
)},
"created_at": str(r["created_at"]) if r["created_at"] is not None else None,
"closed_at": str(r["closed_at"]) if r["closed_at"] is not None else None,
}
for r in rows
]
return items, total
def close_session(self, session_id: str) -> bool:
"""关闭会话(方案 B):active → closed + closed_at 落时间。
条件更新(WHERE status='active')保证幂等语义由路由层判定——重复
关闭返回 False 由路由转 409,不会出现并发双击把 closed_at 刷新的
静默写。审计表只 INSERT 铁律不适用于本表(agent_session 是对话域
业务表,status 流转是既有设计,见 _ddl/01-mysql-共用底座)。
"""
with self._engine.begin() as conn:
result = conn.execute(
text(
"UPDATE agent_session SET status = 'closed', closed_at = CURRENT_TIMESTAMP"
" WHERE session_id = :sid AND status = 'active'"
),
{"sid": session_id},
)
return result.rowcount > 0
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 insert_turn(
self,
*,
session_id: str,
trace_id: str,
user_content: str,
assistant_content: str,
has_disclaimer: bool = False,
) -> int:
"""一轮对话(user + assistant)**同事务**落库,返回起始 seq_no。
方案 C 评审 P0/P1:两条消息必须原子。分两次写时中途故障会留下
「user 已落、assistant 未落」的半截历史——下次请求会把它当上下文
读进 LLM,属于难排查的数据污染。同事务还顺带解决 seq 取号竞态:
max(seq_no)+1 在事务内计算,并发同会话不重号(并发写锁归后续)。
"""
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, NULL)
"""
)
with self._engine.begin() as conn:
seq = (
int(
conn.execute(
text(
"SELECT COALESCE(MAX(seq_no), 0) FROM agent_message WHERE session_id = :sid"
),
{"sid": session_id},
).scalar_one()
)
+ 1
)
conn.execute(
sql,
{
"sid": session_id,
"trace_id": trace_id,
"seq_no": seq,
"role": "user",
"content": user_content,
"has_disclaimer": 0,
},
)
conn.execute(
sql,
{
"sid": session_id,
"trace_id": trace_id,
"seq_no": seq + 1,
"role": "assistant",
"content": assistant_content,
"has_disclaimer": 1 if has_disclaimer else 0,
},
)
return seq
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
]
def list_messages_page(
self, session_id: str, *, limit: int = 50, offset: int = 0
) -> tuple[list[dict[str, Any]], int]:
"""前端历史消息分页(方案 B):seq_no 升序 + offset,返回 (items, total)。
与 list_messages(LLM 窗口「最近 N 条」)语义不同:前端要全量可翻页,
从第一条开始升序拉取。created_at 一并返回供前端展示时间。
"""
with self._engine.connect() as conn:
total = int(
conn.execute(
text("SELECT COUNT(*) FROM agent_message WHERE session_id = :sid"),
{"sid": session_id},
).scalar_one()
)
rows = conn.execute(
text(
"""
SELECT seq_no, role, content, has_disclaimer, created_at
FROM agent_message WHERE session_id = :sid
ORDER BY seq_no ASC
LIMIT :lim OFFSET :off
"""
),
{"sid": session_id, "lim": limit, "off": offset},
).mappings().all()
items = [
{
"seq_no": r["seq_no"],
"role": r["role"],
"content": r["content"],
"has_disclaimer": bool(r["has_disclaimer"]),
"created_at": str(r["created_at"]) if r["created_at"] is not None else None,
}
for r in rows
]
return items, total