1. 新增会话通信工具集,支持列出会话、查看会话历史、发送会话消息和查询智能体进度 2. 实现子智能体嵌套深度限制与父子会话关系管理 3. 新增数据库字段与索引支持会话层级存储 4. 添加会话工具单元测试与LITE模式适配
319 lines
12 KiB
Python
319 lines
12 KiB
Python
"""会话通信工具集 — Agent 间会话通信能力"""
|
||
|
||
import os
|
||
|
||
from langgraph.prebuilt.tool_node import ToolRuntime
|
||
from pydantic import BaseModel, Field
|
||
|
||
from yuxi.agents.toolkits.registry import tool
|
||
from yuxi.utils import logger
|
||
|
||
_LITE_MODE = os.environ.get("LITE_MODE", "").lower() in ("true", "1")
|
||
|
||
# ── 常量 ──────────────────────────────────────────────────
|
||
|
||
_HISTORY_CONTENT_MAX_CHARS = 4000
|
||
_HISTORY_BYTES_MAX = 80 * 1024
|
||
|
||
|
||
# ── 辅助函数 ──────────────────────────────────────────────
|
||
|
||
|
||
def _get_runtime_uid(runtime: ToolRuntime) -> str:
|
||
"""从 ToolRuntime.context 获取当前用户 uid。"""
|
||
uid = getattr(runtime.context, "uid", None)
|
||
if not uid:
|
||
raise ValueError("当前运行时缺少 uid")
|
||
return str(uid)
|
||
|
||
|
||
def _get_runtime_thread_id(runtime: ToolRuntime) -> str:
|
||
"""从 ToolRuntime.context 获取当前线程 ID。
|
||
|
||
注意:子智能体的 file_thread_id 继承自父会话,不能用于标识当前会话。
|
||
会话通信场景应使用 thread_id(当前智能体自身的线程)。
|
||
"""
|
||
thread_id = getattr(runtime.context, "thread_id", None)
|
||
if not thread_id:
|
||
raise ValueError("当前运行时缺少 thread_id")
|
||
return str(thread_id)
|
||
|
||
|
||
# ── 权限校验 ──────────────────────────────────────────────
|
||
|
||
|
||
async def _check_session_access(uid: str, thread_id: str) -> None:
|
||
"""
|
||
检查用户是否有权访问指定会话(tree 级可见性)。
|
||
1. 会话的 uid 与当前用户一致 → 允许
|
||
2. 沿 parent_thread_id 链上溯,任一祖先属于当前用户 → 允许
|
||
3. 否则 → 拒绝
|
||
"""
|
||
from yuxi.repositories.conversation_repository import ConversationRepository
|
||
from yuxi.storage.postgres.manager import pg_manager
|
||
|
||
async with pg_manager.get_async_session_context() as db:
|
||
repo = ConversationRepository(db)
|
||
visited: set[str] = set()
|
||
current_tid = thread_id
|
||
|
||
while current_tid and current_tid not in visited:
|
||
visited.add(current_tid)
|
||
conv = await repo.get_conversation_by_thread_id(current_tid)
|
||
if not conv:
|
||
raise ValueError(f"会话不存在: {thread_id}")
|
||
if conv.uid == uid:
|
||
return
|
||
parent_tid = conv.parent_thread_id
|
||
if not parent_tid:
|
||
break
|
||
current_tid = parent_tid
|
||
|
||
raise PermissionError(f"无权访问会话: {thread_id}")
|
||
|
||
|
||
# ── LITE 模式下不注册会话通信工具 ─────────────────────────
|
||
|
||
if _LITE_MODE:
|
||
logger.info("LITE_MODE enabled, session communication tools not registered")
|
||
else:
|
||
|
||
# ── list_sessions ─────────────────────────────────────────
|
||
|
||
class ListSessionsInput(BaseModel):
|
||
status: str | None = Field(default=None, description="按状态过滤: active / archived")
|
||
agent_id: str | None = Field(default=None, description="按智能体 ID 过滤")
|
||
limit: int = Field(default=20, ge=1, le=50, description="返回数量上限")
|
||
|
||
@tool(
|
||
category="buildin",
|
||
tags=["会话"],
|
||
display_name="列出会话",
|
||
args_schema=ListSessionsInput,
|
||
)
|
||
async def list_sessions(
|
||
status: str | None = None,
|
||
agent_id: str | None = None,
|
||
limit: int = 20,
|
||
runtime: ToolRuntime = None,
|
||
) -> str:
|
||
"""
|
||
列出当前用户可见的对话会话,包括自己创建的和子智能体的会话。
|
||
返回每个会话的 thread_id、agent_id、状态、标题、最近更新时间。
|
||
"""
|
||
from yuxi.repositories.conversation_repository import ConversationRepository
|
||
from yuxi.storage.postgres.manager import pg_manager
|
||
|
||
uid = _get_runtime_uid(runtime)
|
||
|
||
async with pg_manager.get_async_session_context() as db:
|
||
repo = ConversationRepository(db)
|
||
conversations = await repo.list_conversations(
|
||
uid=uid,
|
||
agent_id=agent_id,
|
||
status=status or "active",
|
||
limit=limit,
|
||
)
|
||
|
||
if not conversations:
|
||
return "当前无可见的会话"
|
||
|
||
lines = []
|
||
for conv in conversations:
|
||
parent_marker = " [子智能体]" if conv.parent_thread_id else ""
|
||
lines.append(
|
||
f"- thread_id: {conv.thread_id} | agent: {conv.agent_id} | "
|
||
f"状态: {conv.status} | 标题: {conv.title or '无'}{parent_marker} | "
|
||
f"更新: {conv.updated_at}"
|
||
)
|
||
return "\n".join(lines)
|
||
|
||
# ── get_session_history ───────────────────────────────────
|
||
|
||
class GetSessionHistoryInput(BaseModel):
|
||
thread_id: str = Field(description="目标会话线程 ID")
|
||
limit: int = Field(default=10, ge=1, le=50, description="返回消息数量")
|
||
include_tools: bool = Field(default=False, description="是否包含工具调用消息")
|
||
|
||
@tool(
|
||
category="buildin",
|
||
tags=["会话"],
|
||
display_name="查看会话历史",
|
||
args_schema=GetSessionHistoryInput,
|
||
)
|
||
async def get_session_history(
|
||
thread_id: str,
|
||
limit: int = 10,
|
||
include_tools: bool = False,
|
||
runtime: ToolRuntime = None,
|
||
) -> str:
|
||
"""
|
||
获取指定会话的最近消息。
|
||
仅可查看自己或子智能体的会话历史。
|
||
超长内容会被截断,工具调用消息默认不包含。
|
||
"""
|
||
from yuxi.repositories.conversation_repository import ConversationRepository
|
||
from yuxi.storage.postgres.manager import pg_manager
|
||
|
||
uid = _get_runtime_uid(runtime)
|
||
await _check_session_access(uid, thread_id)
|
||
|
||
async with pg_manager.get_async_session_context() as db:
|
||
repo = ConversationRepository(db)
|
||
# get_messages_by_thread_id 按 created_at.asc() 排序 + limit 返回最旧消息,
|
||
# 但我们需要最近消息,所以先降序查再反转。
|
||
conversation = await repo.get_conversation_by_thread_id(thread_id)
|
||
if not conversation:
|
||
return f"会话 {thread_id} 暂无消息"
|
||
from sqlalchemy import select
|
||
|
||
from yuxi.storage.postgres.models_business import Message
|
||
|
||
result = await db.execute(
|
||
select(Message)
|
||
.where(Message.conversation_id == conversation.id)
|
||
.order_by(Message.created_at.desc())
|
||
.limit(limit)
|
||
)
|
||
messages = list(reversed(result.scalars().unique().all()))
|
||
|
||
if not messages:
|
||
return f"会话 {thread_id} 暂无消息"
|
||
|
||
total_bytes = 0
|
||
lines = []
|
||
truncated = False
|
||
|
||
for msg in messages:
|
||
# 过滤工具消息
|
||
if not include_tools and msg.message_type in ("tool_call", "tool_result"):
|
||
continue
|
||
|
||
content = msg.content or ""
|
||
# 截断超长内容
|
||
if len(content) > _HISTORY_CONTENT_MAX_CHARS:
|
||
content = content[:_HISTORY_CONTENT_MAX_CHARS] + "...[已截断]"
|
||
truncated = True
|
||
|
||
line = f"[{msg.role}] {content}"
|
||
line_bytes = len(line.encode("utf-8"))
|
||
if total_bytes + line_bytes > _HISTORY_BYTES_MAX:
|
||
truncated = True
|
||
break
|
||
total_bytes += line_bytes
|
||
lines.append(line)
|
||
|
||
suffix = "\n(部分内容已截断)" if truncated else ""
|
||
return "\n".join(lines) + suffix
|
||
|
||
# ── send_to_session ───────────────────────────────────────
|
||
|
||
class SendToSessionInput(BaseModel):
|
||
thread_id: str = Field(description="目标会话线程 ID")
|
||
message: str = Field(description="要发送的消息内容")
|
||
|
||
@tool(
|
||
category="buildin",
|
||
tags=["会话"],
|
||
display_name="发送会话消息",
|
||
args_schema=SendToSessionInput,
|
||
)
|
||
async def send_to_session(
|
||
thread_id: str,
|
||
message: str,
|
||
runtime: ToolRuntime = None,
|
||
) -> str:
|
||
"""
|
||
向指定会话发送消息。
|
||
消息会以 user 角色注入目标会话的历史记录。
|
||
如果目标会话正在运行,消息将作为下一轮输入;否则仅写入历史。
|
||
不能向自己所在的会话发送消息(避免循环)。
|
||
"""
|
||
from yuxi.repositories.conversation_repository import ConversationRepository
|
||
from yuxi.storage.postgres.manager import pg_manager
|
||
|
||
uid = _get_runtime_uid(runtime)
|
||
current_thread_id = _get_runtime_thread_id(runtime)
|
||
|
||
# 防止向自身会话发送消息
|
||
if thread_id == current_thread_id:
|
||
return "不能向当前所在的会话发送消息"
|
||
|
||
await _check_session_access(uid, thread_id)
|
||
|
||
# 标注消息来源
|
||
annotated_message = f"[来自会话 {current_thread_id}] {message}"
|
||
|
||
async with pg_manager.get_async_session_context() as db:
|
||
repo = ConversationRepository(db)
|
||
await repo.add_message_by_thread_id(
|
||
thread_id=thread_id,
|
||
role="user",
|
||
content=annotated_message,
|
||
message_type="text",
|
||
extra_metadata={
|
||
"source_thread_id": current_thread_id,
|
||
"source_type": "inter_session",
|
||
},
|
||
)
|
||
|
||
return f"消息已发送到会话 {thread_id}"
|
||
|
||
# ── get_agent_progress ────────────────────────────────────
|
||
|
||
class GetAgentProgressInput(BaseModel):
|
||
thread_id: str = Field(description="会话线程 ID")
|
||
|
||
@tool(
|
||
category="buildin",
|
||
tags=["会话"],
|
||
display_name="查看智能体进度",
|
||
args_schema=GetAgentProgressInput,
|
||
)
|
||
async def get_agent_progress(
|
||
thread_id: str,
|
||
runtime: ToolRuntime = None,
|
||
) -> str:
|
||
"""
|
||
查看指定智能体的运行进度。
|
||
返回运行状态、开始时间、运行类型等信息。
|
||
"""
|
||
from sqlalchemy import select
|
||
|
||
from yuxi.storage.postgres.manager import pg_manager
|
||
from yuxi.storage.postgres.models_business import AgentRun
|
||
|
||
uid = _get_runtime_uid(runtime)
|
||
await _check_session_access(uid, thread_id)
|
||
|
||
async with pg_manager.get_async_session_context() as db:
|
||
result = await db.execute(
|
||
select(AgentRun)
|
||
.where(
|
||
AgentRun.thread_id == thread_id,
|
||
AgentRun.uid == uid,
|
||
)
|
||
.order_by(AgentRun.created_at.desc())
|
||
.limit(1)
|
||
)
|
||
run = result.scalar_one_or_none()
|
||
|
||
if not run:
|
||
return f"会话 {thread_id} 暂无运行记录"
|
||
|
||
lines = [
|
||
f"运行 ID: {run.id}",
|
||
f"状态: {run.status}",
|
||
f"类型: {run.run_type}",
|
||
f"智能体: {run.agent_id}",
|
||
f"创建时间: {run.created_at}",
|
||
]
|
||
if run.started_at:
|
||
lines.append(f"开始时间: {run.started_at}")
|
||
if run.finished_at:
|
||
lines.append(f"完成时间: {run.finished_at}")
|
||
if run.error_message:
|
||
lines.append(f"错误: {run.error_message}")
|
||
|
||
return "\n".join(lines)
|