ForcePilot/backend/package/yuxi/agents/toolkits/buildin/session_tools.py
Kris 32b213593e feat: 新增会话通信工具与子会话嵌套管理能力
1. 新增会话通信工具集,支持列出会话、查看会话历史、发送会话消息和查询智能体进度
2. 实现子智能体嵌套深度限制与父子会话关系管理
3. 新增数据库字段与索引支持会话层级存储
4. 添加会话工具单元测试与LITE模式适配
2026-06-13 20:19:07 +08:00

319 lines
12 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.

"""会话通信工具集 — 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)