Files
group_xinghuo_jinrong/app/repository/session_repository.py
T
zhanghongyu_0626 9d4d4aaa6d feat(chat): Enhance session management with new API endpoints and filtering options
- Added `status` query parameter to `list_sessions_api` for filtering sessions by their status (active/closed).
- Introduced `close_all_sessions_api` endpoint to allow users to close all active sessions for the current actor.
- Updated `SessionRepository` to support status filtering in session listing and implemented logic for closing active sessions.
- Improved Redis connection settings for better performance and reliability.

This update enhances the chat functionality by providing more control over session management, improving user experience and system efficiency.
2026-09-10 22:30:07 +08:00

369 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.
"""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 list_sessions(
self,
*,
actor_id: str,
agent_type: str,
limit: int = 20,
offset: int = 0,
status: str | None = None,
) -> tuple[list[dict[str, Any]], int]:
"""前端会话列表(方案 B):仅本人 + 本 Agent 线,created_at 倒序分页。
返回 (items, total);total 供前端分页器。id 倒序兜底同秒并发建的
会话排序稳定(created_at 精度秒级时并列)。datetime 统一转 str——
sqlite 返 str、MySQL 返 datetime,响应体跨库同构。
status 可选:active / closed;缺省返回全部状态(兼容旧客户端)。
"""
where = "WHERE actor_id = :actor AND agent_type = :atype"
params: dict[str, Any] = {
"actor": actor_id,
"atype": agent_type,
"lim": limit,
"off": offset,
}
if status is not None:
where += " AND status = :status"
params["status"] = status
with self._engine.connect() as conn:
total = int(
conn.execute(
text(f"SELECT COUNT(*) FROM agent_session {where}"),
params,
).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
"""
),
params,
).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 close_all_active(self, *, actor_id: str, agent_type: str) -> int:
"""关闭该 actor 在某 Agent 线上的全部 active 会话(前端「清空历史」)。"""
with self._engine.begin() as conn:
result = conn.execute(
text(
"UPDATE agent_session SET status = 'closed', closed_at = CURRENT_TIMESTAMP"
" WHERE actor_id = :actor AND agent_type = :atype AND status = 'active'"
),
{"actor": actor_id, "atype": agent_type},
)
return int(result.rowcount)
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