From cd732c1725bd3ed802ec8ffb867724b55564c598 Mon Sep 17 00:00:00 2001 From: zhujiawei Date: Thu, 6 Aug 2026 17:47:25 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20=E9=87=8D=E6=9E=84worker=E7=BB=93?= =?UTF-8?q?=E6=9E=84?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- AGENTS.md | 15 +- README.md | 14 +- ...00\345\217\221\346\227\245\345\277\227.md" | 8 + app/agent/core.py | 611 +----------------- app/agent/worker.py | 333 ---------- app/agent/worker/__init__.py | 9 + app/agent/worker/__main__.py | 7 + app/agent/worker/bootstrap.py | 84 +++ app/agent/worker/event_adapter.py | 285 ++++++++ app/agent/worker/event_writer.py | 186 ++++++ app/agent/worker/execution.py | 102 +++ app/agent/worker/graph_runner.py | 121 ++++ app/agent/worker/result_presenter.py | 78 +++ app/agent/worker/runtime.py | 136 ++++ tests/unit/admin/test_admin_write_service.py | 6 +- .../test_final_result_graph_selection.py | 2 +- tests/unit/agent/test_ordered_event_writer.py | 2 +- tests/unit/agent/test_stream_events.py | 2 +- tests/unit/agent/test_worker_runtime.py | 87 +++ 19 files changed, 1149 insertions(+), 939 deletions(-) delete mode 100644 app/agent/worker.py create mode 100644 app/agent/worker/__init__.py create mode 100644 app/agent/worker/__main__.py create mode 100644 app/agent/worker/bootstrap.py create mode 100644 app/agent/worker/event_adapter.py create mode 100644 app/agent/worker/event_writer.py create mode 100644 app/agent/worker/execution.py create mode 100644 app/agent/worker/graph_runner.py create mode 100644 app/agent/worker/result_presenter.py create mode 100644 app/agent/worker/runtime.py create mode 100644 tests/unit/agent/test_worker_runtime.py diff --git a/AGENTS.md b/AGENTS.md index 37b9729..880f15e 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -60,6 +60,11 @@ │ ├── main/ # 通用页面相关路由 │ ├── auth/ # 登录、注册等认证相关路由 │ ├── admin/ # 管理员 API、审计服务与受保护 Vue 入口 +│ ├── agent/ # 分析任务 API、共享队列服务与独立 worker +│ │ ├── routes.py # Web 任务创建与 SSE 订阅 +│ │ ├── job_service.py # Web、monitor、worker 共享持久化服务 +│ │ ├── core.py # 无运行时状态的兼容导入门面 +│ │ └── worker/ # worker package 与 python -m 入口 │ ├── chat/ # 聊天与会话相关路由和服务 │ ├── files/ # 文件上传与管理相关路由 │ └── static/ # 前端静态资源 @@ -115,8 +120,8 @@ - `create_app()` 会先执行 `app/db.py` 中的 `check_database_readiness()`,确认数据库和关键表已就绪,然后再注册蓝图。 - 当前实际注册的蓝图有 7 个:`auth`、`chat`、`files`、`agent`、`main`、`admin`、`admin_page`。 - Web 进程只负责登录态校验、短请求、analysis job 入队和 SSE 推送;Agent/RAG/MCP 长任务不在 Web 进程内执行,而是由独立 worker 进程处理。 -- 后台 worker 入口是 `python -m app.agent.worker`;worker 启动流程是:数据库就绪检查 -> 初始化 LLM -> 检查 RAG 可用性 -> 按 `JOB_WORKERS` 启动多个 slot。 -- 每个 worker slot 会独占一组 MCP server process、一个通过 `MultiServerMCPClient.session("causal")` 打开的持久 `ClientSession`、一组由 `load_mcp_tools(session)` 生成的 LangChain tools,以及一个编译好的 Agent graph;真实执行单元是 slot,不是 Flask 请求线程。旧 `open_mcp_session()` / 手写 `list_tools()` 包装仅保留作历史兼容入口。 +- 后台 worker 入口仍是 `python -m app.agent.worker`,实际由 `app/agent/worker/__main__.py` 调用 `bootstrap.main()`;启动流程是:数据库就绪检查 -> PostgreSQL checkpoint 检查 -> 创建显式进程 runtime(LLM、RAG 可用性)-> 按 `JOB_WORKERS` 启动多个 slot。 +- 每个 worker slot 会独占一组 MCP server process、一个通过 `MultiServerMCPClient.session("causal")` 打开的持久 `ClientSession`、一组由 `load_mcp_tools(session)` 生成的 LangChain tools,以及一个编译好的 Agent graph;`runtime.py` 通过 `ProcessRuntime` 和 `SlotRuntime` 显式返回这些依赖,执行函数不读取 `app.agent.core` 的全局 LLM 或 graph。真实执行单元是 slot,不是 Flask 请求线程。 - 父图当前只暴露 `mcp`、`rag` 两个工具阶段节点:`mcp` 子图正常路径执行 `mcp_planner -> mcp_tool_node -> mcp_result_parser`,planner、ToolNode 和 parser 的失败路径会在子图内生成标准 `success=False` 结果并结束子图;`rag` 子图内部执行 `rag_question_planner -> rag_tool_node -> rag_result_parser`。worker 使用 LangGraph v2 `updates/messages/custom/tasks` 多流:根图 `tasks` 形成用户时间线,子图工具事件折叠到 `mcp`/`rag` 阶段,只有 `normal_chat` 和 `inquiry_answer` 的文字进入 `text_delta`,原始 Prompt、ToolMessage、完整工具结果和内部 attempt 不进入普通用户 SSE 协议。 - Pydantic 结构化输出统一通过 `Agent/llm_structured_output.py` 的同步/异步入口执行,固定使用普通 `function_calling`;调用器仅对结构化请求发送 `thinking.type=disabled`,避免 DeepSeek Thinking 与固定 `tool_choice` 冲突。MCP 继续使用原生 Tool Calls;只有 MCP planner 使用关闭 Thinking 的 LLM 副本和 `tool_choice="required"`,确保模型必须自行选择一个已加载工具。 - `agent` 与 `fold` 的条件路由只读取 `route_decision`、`fold_decision` 显式 State 字段;展示消息仅用于用户可见内容和审计,不参与控制流。 @@ -339,7 +344,7 @@ MYSQL_USER/MYSQL_PASSWORD:现在主要是兼容兜底,主从开发里不依 - 新增数据库表但忘了更新 `check_database_readiness` - 修改接口返回结构但没有检查前端 `script.js` - 改了上传或聊天附件结构却没同步恢复逻辑 -- 修改 MCP 或 RAG 初始化路径但没检查 `CausalAgent.py` 和 `app/agent/core.py` +- 修改 MCP、RAG 或 worker 初始化路径但没检查 `app/agent/worker/runtime.py`、`bootstrap.py` 和 Docker Compose 入口 ## 6. 修改后的验证要求 @@ -385,7 +390,9 @@ python CausalAgent.py 至少核对: -- `app/agent/core.py` +- `app/agent/worker/runtime.py` +- `app/agent/worker/bootstrap.py` +- `app/agent/worker/graph_runner.py` - `Agent/causal_agent/` - `Agent/tool_node/` - `Agent/knowledge_base/` diff --git a/README.md b/README.md index 3c0082e..e5fbf25 100644 --- a/README.md +++ b/README.md @@ -56,12 +56,12 @@ CausalAgent - [快速开始 | Quick Start](#快速开始--quick-start) - [Docker部署](#docker部署) - [数据库生产化配置](#数据库生产化配置) + - [管理员后台](#管理员后台) - [后端单元测试](#后端单元测试) - [windows部署](#windows部署) - [贡献](#贡献) - [Star 趋势](#star-趋势) - [项目结构](#项目结构) -- [更新日志](./README/开发日志.md) @@ -357,6 +357,18 @@ docker compose -f docker-compose.test.yml run --rm unit-test sh │ ├── main/ # 通用页面相关路由 │ ├── auth/ # 登录、注册等认证相关路由 │ ├── admin/ # 管理 API、审计服务与受保护 Vue 入口 +│ ├── agent/ # 分析任务 API、队列服务与独立 worker +│ │ ├── routes.py # Web 进程创建任务与订阅 SSE +│ │ ├── job_service.py # Web、monitor、worker 共享的任务持久化服务 +│ │ ├── core.py # 不持有运行时状态的兼容导入门面 +│ │ └── worker/ # python -m app.agent.worker 包入口 +│ │ ├── bootstrap.py # 启动检查与 slot 编排 +│ │ ├── runtime.py # 显式进程/slot runtime +│ │ ├── execution.py # 单 job 执行与 heartbeat +│ │ ├── event_writer.py # 顺序事件持久化 +│ │ ├── graph_runner.py # LangGraph 流式执行 +│ │ ├── event_adapter.py # 内部流到公开事件协议 +│ │ └── result_presenter.py # 最终结果展示结构 │ ├── chat/ # 聊天 & 会话相关路由与服务 │ ├── files/ # 文件上传/管理相关路由 │ └── static/ # 前端静态资源 diff --git "a/README/\345\274\200\345\217\221\346\227\245\345\277\227.md" "b/README/\345\274\200\345\217\221\346\227\245\345\277\227.md" index f8de667..6aa2069 100644 --- "a/README/\345\274\200\345\217\221\346\227\245\345\277\227.md" +++ "b/README/\345\274\200\345\217\221\346\227\245\345\277\227.md" @@ -548,3 +548,11 @@ - 【修复:Agent 创建请求幂等】 - `/api/agent/jobs` 要求客户端提供 `Idempotency-Key`,服务端将请求指纹和幂等键与 job 创建放在同一个 MySQL 事务中。 - 网络重试使用同一幂等键时返回原 job;同一幂等键对应不同会话或消息时返回冲突,避免终态 job 释放 active 锁后重复保存聊天记录。 + +--- +2026.8.6 +- 【Agent Worker包结构重构】 + - 将单文件 `app/agent/worker.py` 拆分为可通过 `python -m app.agent.worker` 启动的 package,按启动编排、运行时、单任务执行、事件写入、图执行、事件适配和结果展示划分职责。 + - 新增显式 `ProcessRuntime` 与 `SlotRuntime`:LLM 在进程级创建,MCP session、tools 与 graph 在 slot 级创建;任务执行不再读取 `app.agent.core.llm` 等模块全局变量。 + - `job_service.py` 和 `routes.py` 保持在 worker package 外,继续供 Web、monitor、管理员看板与 worker 共用;管理员任务、checkpoint 和 SSE 数据契约不变。 + - 测试改为直接导入新职责模块,并将管理员 checkpoint 约束检查指向实际写入 `job_id` metadata 的 `graph_runner.py`。 diff --git a/app/agent/core.py b/app/agent/core.py index 67c90ab..0aa2796 100644 --- a/app/agent/core.py +++ b/app/agent/core.py @@ -1,601 +1,20 @@ -""" -app.agent.core - agent核心模块 +"""Agent worker 旧导入路径的轻量兼容服务。 -- 初始化llm -- 初始化mcp -- 启动期检查rag -- 初始化agent +仓库内部代码应直接依赖 ``app.agent.worker`` 下的职责模块。这里不保存 +LLM、MCP 或 graph 全局状态,也不会在导入时启动运行时资源。 """ -import asyncio, threading, logging, sys, os, time -import hashlib -from contextlib import AsyncExitStack -from dataclasses import dataclass -from mcp import ClientSession, StdioServerParameters -from mcp.client.stdio import stdio_client -from config.settings import settings -from langchain_openai import ChatOpenAI -from langchain_core.messages import HumanMessage, AIMessage -from langchain_core.tools import BaseTool -from typing import Any, Type, List -from pydantic import BaseModel, create_model -from langgraph.types import Command -from Agent.causal_agent.state import CausalAgentState -from app.chat.response_storage import render_summary_for_display - -## die manager -BASE_DIR = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) -mcp_dir = os.path.join(BASE_DIR, "Agent") -mcp_server_path = os.path.join(mcp_dir, "CausalAgentMCP", "mcp_server.py") -knowledge_base_dir = os.path.join(BASE_DIR, "Agent","knowledge_base") - - -@dataclass -class McpClientResources: - """worker slot 级 MCP 资源,生命周期由传入的 AsyncExitStack 管理。""" - client: Any - session: Any - tools: list - -# 将 MCP 和事件循环,llm和rag链的相关的状态集中管理 -mcp_session: ClientSession | None = None -mcp_tools: list = [] -mcp_process_stack = AsyncExitStack() -background_loop: asyncio.AbstractEventLoop | None = None -llm = None -agent_graph = None - -NODE_DESCRIPTIONS = { - "agent": "分析用户意图", - "fold": "加载文件并验证数据", - "preprocess": "预处理数据", - "mcp": "执行因果分析", - "rag": "检索知识库", - "postprocess": "校验并修正因果图", - "report": "生成分析报告", - "normal_chat": "生成回答", - "inquiry_answer": "回答报告追问", -} -TEXT_STREAM_NODES = {"normal_chat", "inquiry_answer"} -TOOL_STAGE_NODES = { - "mcp_planner": "mcp", - "mcp_tool_node": "mcp", - "mcp_result_parser": "mcp", - "rag_question_planner": "rag", - "rag_tool_node": "rag", - "rag_result_parser": "rag", -} -DECISION_PROGRESS = { - "fold": "已识别为因果分析请求", - "normal_chat": "已识别为普通问答", - "inquiry_answer": "已识别为报告追问", - "postprocess": "已识别为已有结果的后续处理", - "preprocess": "文件与数据验证完成", - "agent": "需要补充分析输入", -} -DECISION_FIELDS = { - "agent": "route_decision", - "fold": "fold_decision", -} - - -def _opaque_id(*parts: Any) -> str: - """根据内部标识生成不暴露 attempt/task 内容的稳定 ID。""" - raw = "\x1f".join(str(part) for part in parts) - return hashlib.sha256(raw.encode("utf-8")).hexdigest()[:24] - - -def sanitize_public_error(error: Any) -> str: - """把内部异常归类为有限的用户可见错误,避免泄露连接和路径信息。""" - name = error if isinstance(error, str) else type(error).__name__ - normalized = str(name).lower() - if "timeout" in normalized: - return "调用超时" - if any(token in normalized for token in ("connection", "connect", "network")): - return "服务连接失败" - if any(token in normalized for token in ("rate", "limit")): - return "服务当前繁忙" - if any(token in normalized for token in ("permission", "auth")): - return "服务授权失败" - return "节点执行失败" - - -class LangGraphEventAdapter: - """把 LangGraph v2 多流事件转换为稳定、可持久化的公开事件协议。""" - - def __init__(self, job_id: str | None, job_attempt: int): - """初始化单次 job attempt 的阶段、重试和文字流状态。""" - self.job_id = job_id or "untracked" - self.job_attempt = job_attempt - self.steps: dict[str, dict[str, Any]] = {} - self.active_by_node: dict[str, str] = {} - self.failed_attempts: dict[str, str] = {} - self.streams: dict[str, dict[str, Any]] = {} - - def _base(self, event_type: str, step: dict[str, Any]) -> dict[str, Any]: - """构造所有阶段事件共享的持久化字段。""" - return { - "type": event_type, - "step_id": step["step_id"], - "node_name": step["node_name"], - "title": NODE_DESCRIPTIONS[step["node_name"]], - "attempt": self.job_attempt, - } - - def _active_step(self, node_name: str) -> dict[str, Any] | None: - """获取指定父图节点当前尚未结束的阶段实例。""" - task_id = self.active_by_node.get(node_name) - return self.steps.get(task_id) if task_id else None - - def _parent_tool_step(self, node_name: str) -> dict[str, Any] | None: - """把子图内部节点映射到当前 MCP/RAG 父阶段。""" - parent_name = TOOL_STAGE_NODES.get(node_name) - return self._active_step(parent_name) if parent_name else None - - def _task_event(self, namespace: Any, data: Any) -> list[dict[str, Any]]: - """将根图 tasks 开始/结束转换为阶段生命周期事件。""" - if namespace or not isinstance(data, dict): - return [] - task_id = str(data.get("id") or "") - node_name = data.get("name") - if not task_id or node_name not in NODE_DESCRIPTIONS: - return [] - if "input" in data: - step = { - "step_id": _opaque_id(self.job_id, self.job_attempt, task_id), - "node_name": node_name, - "started_at": time.monotonic(), - } - self.steps[task_id] = step - self.active_by_node[node_name] = task_id - return [self._base("node_start", step)] - step = self.steps.get(task_id) - if not step: - return [] - event = self._base("node_end", step) - event["duration"] = round(max(0.0, time.monotonic() - step["started_at"]), 2) - event["status"] = "failed" if data.get("error") else "completed" - if data.get("error"): - event["message"] = sanitize_public_error(data.get("error")) - if self.active_by_node.get(node_name) == task_id: - self.active_by_node.pop(node_name, None) - return [event] - - def _custom_event(self, data: Any) -> list[dict[str, Any]]: - """只有失败后再次开始时,才把内部 attempt 转换为真正的重试。""" - if not isinstance(data, dict): - return [] - event_type = data.get("type") - task_id = str(data.get("task_id") or "") - node_name = data.get("node_name") - if event_type == "node_attempt_failed" and task_id: - self.failed_attempts[task_id] = sanitize_public_error(data.get("error_kind")) - return [] - if event_type != "node_attempt_start" or not task_id: - return [] - failure = self.failed_attempts.pop(task_id, None) - if not failure: - return [] - step = ( - self.steps.get(task_id) - or self._active_step(node_name) - or self._parent_tool_step(node_name) - ) - if not step: - return [] - discarded = self.streams.pop(step["step_id"], None) - event = self._base("node_retry", step) - event["message"] = failure - if discarded: - event["discard_stream_id"] = discarded["stream_id"] - return [event] - - @staticmethod - def _tool_call(message: Any) -> tuple[str | None, list[str]]: - """只读取工具名和参数字段名,不返回参数值。""" - calls = getattr(message, "tool_calls", None) or [] - if not calls: - return None, [] - call = calls[0] - if isinstance(call, dict): - name = call.get("name") - args = call.get("args") - else: - name = getattr(call, "name", None) - args = getattr(call, "args", None) - keys = sorted(str(key) for key in args)[:12] if isinstance(args, dict) else [] - return name, keys - - def _update_event(self, namespace: Any, data: Any) -> list[dict[str, Any]]: - """从显式 State 和规范化结果生成 decision、progress 与工具摘要。""" - if not isinstance(data, dict): - return [] - events: list[dict[str, Any]] = [] - for node_name, output in data.items(): - if not isinstance(output, dict): - continue - step = self._active_step(node_name) - if step: - decision_field = DECISION_FIELDS.get(node_name) - decision = output.get(decision_field) if decision_field else None - if decision: - progress = self._base("progress", step) - progress["summary"] = DECISION_PROGRESS.get(decision, "路由判断完成") - events.append(progress) - event = self._base("decision", step) - event["summary"] = f"已决定进入:{NODE_DESCRIPTIONS.get(decision, decision)}" - events.append(event) - tool_step = self._parent_tool_step(node_name) - if not tool_step: - continue - messages = output.get("messages") or [] - latest = messages[-1] if messages else None - tool_name, argument_keys = self._tool_call(latest) - if tool_name: - event = self._base("tool_call_start", tool_step) - event.update({"tool_name": tool_name, "argument_keys": argument_keys}) - events.append(event) - result_key = "causal_analysis_result" if TOOL_STAGE_NODES[node_name] == "mcp" else "knowledge_base_result" - result = output.get(result_key) - if isinstance(result, dict): - metadata = result.get("_tool_call") or {} - event = self._base("tool_call_result", tool_step) - event.update({ - "tool_name": metadata.get("name") or tool_name or "工具", - "summary": "调用完成" if result.get("success") is not False else "调用失败", - }) - events.append(event) - return events - - def _message_event(self, data: Any) -> list[dict[str, Any]]: - """仅转换普通问答和报告追问的非空字符串 token。""" - if not isinstance(data, (tuple, list)) or len(data) != 2: - return [] - chunk, metadata = data - if not isinstance(metadata, dict): - return [] - node_name = metadata.get("langgraph_node") - content = getattr(chunk, "content", None) - if node_name not in TEXT_STREAM_NODES or not isinstance(content, str) or not content: - return [] - step = self._active_step(node_name) - if not step: - return [] - stream = self.streams.get(step["step_id"]) - if not stream: - stream = { - "stream_id": _opaque_id(step["step_id"], time.monotonic_ns()), - "sequence": 0, - } - self.streams[step["step_id"]] = stream - stream["sequence"] += 1 - event = self._base("text_chunk", step) - event.update({ - "stream_id": stream["stream_id"], - "sequence": stream["sequence"], - "delta": content, - }) - return [event] - - def convert(self, chunk: Any) -> list[dict[str, Any]]: - """转换一个 v2 StreamPart;无法识别的内部流会被忽略。""" - if not isinstance(chunk, dict): - return [] - stream_type = chunk.get("type") - namespace = chunk.get("ns") or () - data = chunk.get("data") - if stream_type == "tasks": - return self._task_event(namespace, data) - if stream_type == "custom": - return self._custom_event(data) - if stream_type == "updates": - return self._update_event(namespace, data) - if stream_type == "messages": - return self._message_event(data) - return [] - -def initialize_llm(): - """在应用启动时初始化全局LLM实例。""" - global llm - # 使用新的配置对象 - if not all([settings.MODEL, settings.BASE_URL, settings.API_KEY]): - logging.error("LLM 配置不完整,无法初始化。") - return False - - logging.info(f"正在初始化 LLM 模型: {settings.MODEL}") - llm = ChatOpenAI( - model=settings.MODEL, - base_url=settings.BASE_URL, - api_key=settings.API_KEY, - streaming=False, - ) - logging.info("LLM 实例初始化成功。") - - return True - - - -async def open_mcp_client_resources(process_stack: AsyncExitStack) -> McpClientResources: - """ - 使用 LangChain MCP adapter 打开 slot 级持久 session 并加载 tools。 - - 这里显式使用 client.session("causal"),避免 MultiServerMCPClient 默认 - stateless get_tools() 路径在每次工具调用时重新创建 ClientSession。 - """ - try: - from langchain_mcp_adapters.client import MultiServerMCPClient - from langchain_mcp_adapters.tools import load_mcp_tools - except ImportError as exc: - raise RuntimeError( - "缺少 langchain-mcp-adapters,无法初始化 LangChain MCP adapter。" - ) from exc - - client = MultiServerMCPClient( - { - "causal": { - "transport": "stdio", - "command": sys.executable, - "args": [mcp_server_path], - } - } - ) - logging.info("MCP 初始化。") - session = await process_stack.enter_async_context(client.session("causal")) - tools = await load_mcp_tools(session) - return McpClientResources(client=client, session=session, tools=tools) - - -def initialize_rag_system(): - """ - 启动期仅做知识库可用性检查。 - 并在首次查询时延迟初始化 LLM、Embedding 和 Chroma。 - """ - logging.info("正在检查 RAG 知识库目录...") - persist_directory = os.path.join(knowledge_base_dir, "db") - if not os.path.exists(persist_directory): - logging.warning( - "知识库持久化目录不存在。请先运行 Agent/knowledge_base/build_knowledge.py 构建知识库。", - persist_directory, - ) - return False - - logging.info("RAG 启动检查通过;向量库将在首次实际查询时由 query_rag.py 延迟初始化。") - return True - - -def _snapshot_interrupts(snapshot) -> list[Any]: - """汇总 LangGraph StateSnapshot 中尚待恢复的 task interrupts。""" - return [ - item - for task in (getattr(snapshot, "tasks", None) or ()) - for item in (getattr(task, "interrupts", None) or ()) - ] - - -async def ai_call_stream( - text, - user_id, - username, - session_id, - *, - job_id=None, - job_attempt=1, - graph=None, -): - """ - 流式版本的 ai_call,使用 astream() 捕获节点执行更新。 - 这是一个生成器函数,会yield SSE格式的事件数据。 - """ - logging.info(f"[流式] 处理用户 {username} 的消息,会话ID: {session_id}") - target_graph = graph or agent_graph - if target_graph is None: - raise RuntimeError("Agent Graph 尚未初始化") - - # 配置:使用 session_id 作为 thread_id - config = { - "configurable": { - "thread_id": session_id, - "user_id": user_id - }, - "metadata": { - "job_id": job_id, - } if job_id else {}, - } - - # 检查当前状态,判断是否是恢复中断的会话 - try: - state = await target_graph.aget_state(config) - is_interrupted = bool(_snapshot_interrupts(state)) - - if is_interrupted: - logging.info(f"[流式] 检测到会话 {session_id} 处于中断状态,使用Command(resume=...)恢复") - input_data = Command(resume=text) - else: - logging.info(f"[流式] 正常对话或第一次对话") - input_data = { - "messages": [HumanMessage(content=text)], - "user_id": user_id, - "username": username, - "session_id": session_id - } - except Exception as e: - logging.warning(f"[流式] 无法获取状态,假设为新对话: {e}") - input_data = { - "messages": [HumanMessage(content=text)], - "user_id": user_id, - "username": username, - "session_id": session_id - } - - adapter = LangGraphEventAdapter(job_id, job_attempt) - streamed_interrupts = [] - - try: - async for chunk in target_graph.astream( - input_data, - config, - stream_mode=["updates", "messages", "custom", "tasks"], - subgraphs=True, - version="v2", - ): - if chunk.get("type") == "updates" and isinstance(chunk.get("data"), dict): - interrupt_data = chunk["data"].get("__interrupt__") - if interrupt_data: - streamed_interrupts.extend( - interrupt_data if isinstance(interrupt_data, (list, tuple)) else [interrupt_data] - ) - for event_data in adapter.convert(chunk): - yield event_data - - # 获取最终状态以检查interrupt - state = await target_graph.aget_state(config) - final_state_data = state.values - pending_interrupts = _snapshot_interrupts(state) - interrupts = pending_interrupts or streamed_interrupts - - # 检查是否有interrupt - if interrupts: - # 提取问题文本 - interrupt_obj = interrupts[0] - question = interrupt_obj.value if hasattr(interrupt_obj, 'value') else str(interrupt_obj) - - event_data = { - "type": "interrupt", - "message": question - } - event_data["attempt"] = job_attempt - yield event_data - logging.info(f"[SSE] 图已暂停,等待用户输入") - return - else: - # 发送最终结果 - result = process_final_result(final_state_data) - event_data = { - "type": "final_result", - "data": result - } - event_data["attempt"] = job_attempt - yield event_data - logging.info(f"[SSE] 发送最终结果") - - except Exception as e: - logging.error(f"[流式] 执行 LangGraph Agent 时发生错误: {e}", exc_info=True) - yield { - "type": "error", - "message": sanitize_public_error(e), - "attempt": job_attempt, - } - -def process_final_result(final_state_data): - """ - 处理图正常完成后的最终结果 - - 优先级策略: - 1. 优先返回最新的 AI 消息(messages[-1]) - 2. 如果最新消息是决策消息,则检查是否有 final_report - 3. 如果有因果图数据,返回结构化响应 - """ - - # 检查最后一条消息 - messages = final_state_data.get("messages", []) - if messages: - last_message = messages[-1] - - # 如果最后一条是 AI 消息 - if isinstance(last_message, AIMessage): - message_name = getattr(last_message, 'name', None) - - # 检查是否是有实际内容的回复节点 - # (不是决策消息,而是真正的回复) - if message_name in ['normal_chat', 'inquiry_answer']: - logging.info(f"返回 {message_name} 节点的回复") - return { - "type": "text", - "summary": last_message.content - } - - # 如果是 report 节点生成的决策消息 - # 检查是否同时有 final_report - if message_name == 'report' and final_state_data.get("final_report"): - logging.info("返回完整的因果分析报告") - result = { - "summary": final_state_data["final_report"], - "layout": "report" - } - # 检查是否有因果图数据(结构化返回) - if final_state_data.get("causal_analysis_result"): - analysis_data = final_state_data["causal_analysis_result"] - if analysis_data.get("success"): - original_graph = analysis_data.get("data") - postprocess_result = final_state_data.get("postprocess_result") or {} - revised_graph = postprocess_result.get("revised_graph") - has_valid_revised_graph = ( - isinstance(revised_graph, dict) - and isinstance(revised_graph.get("nodes"), list) - and isinstance(revised_graph.get("edges"), list) - and not postprocess_result.get("error") - ) - result["type"] = "causal_graph" - result["data"] = revised_graph if has_valid_revised_graph else original_graph - result["graph_source"] = ( - "postprocessed" if has_valid_revised_graph else "original" - ) - result["revision_summary"] = postprocess_result.get( - "revision_summary", - "", - ) - logging.info("返回因果图数据: source=%s", result["graph_source"]) - - if "type" not in result: - result["type"] = "text" - - # 检查是否有可视化映射,并替换占位符 - if final_state_data.get("visualization_mapping"): - visualization_mapping = final_state_data["visualization_mapping"] - if visualization_mapping: # 确保不是空字典 - # 保存映射数据(用于数据库存储) - result["raw_summary"] = result["summary"] - result["visualization_mapping"] = visualization_mapping - logging.info(f"包含 {len(visualization_mapping)} 个可视化图表") - - # 替换 summary 中的占位符为真实图表 - result["summary"] = render_summary_for_display( - result["summary"], - visualization_mapping - ) - logging.info("已替换报告中的占位符为真实图表") - - return result - - # 返回 final_report(如果有) - final_report = final_state_data.get("final_report") - if final_report: - logging.info("未找到最新消息,降级返回 final_report") - return {"type": "text", "summary": final_report, "layout": "report"} - - logging.warning("未找到任何可返回的内容,返回默认消息") - return {"type": "text", "summary": "抱歉,我在处理时遇到了问题。"} -async def open_mcp_session(process_stack: AsyncExitStack): - """ - 旧函数:创建一组独立 MCP 进程和 ClientSession。 +from app.agent.worker.event_adapter import ( + LangGraphEventAdapter, + sanitize_public_error, +) +from app.agent.worker.graph_runner import ai_call_stream +from app.agent.worker.result_presenter import process_final_result - Web 旧模式只使用一个全局会话;worker 池会为每个 slot 调用本函数, - 确保 slot = MCP session/process = graph instance。 - """ - server_params = StdioServerParameters(command=sys.executable, args=[mcp_server_path]) - ## 这里enter_async_context(...) 会立刻执行它的 __aenter__(),真正打开资源 - ## 同时,它的 __aexit__() 被登记到 process_stack这个栈,也就是只有__aexit__()被登记到栈里 - read_stream, write_stream = await process_stack.enter_async_context(stdio_client(server_params)) - session = await process_stack.enter_async_context(ClientSession(read_stream, write_stream)) - await session.initialize() - tools_response = await session.list_tools() - tools = [{ - "type": "function", - "function": { - "name": tool.name, - "description": tool.description, - "parameters": tool.inputSchema, - } - } for tool in tools_response.tools] - return session, tools +__all__ = [ + "LangGraphEventAdapter", + "ai_call_stream", + "process_final_result", + "sanitize_public_error", +] diff --git a/app/agent/worker.py b/app/agent/worker.py deleted file mode 100644 index e983107..0000000 --- a/app/agent/worker.py +++ /dev/null @@ -1,333 +0,0 @@ -""" -后台 analysis job worker。 - -启动方式: - python -m app.agent.worker -""" - -from __future__ import annotations - -import asyncio -from contextlib import AsyncExitStack -import logging -import socket -import sys -from typing import Any - -from Agent.causal_agent.postgres_checkpointer import ( - build_checkpointer, - open_checkpoint_pool, - verify_checkpoint_schema, -) -from app.agent import core as agent_core -from app.agent import job_service -from app.db import check_database_readiness -from config.settings import settings - - -TEXT_FLUSH_INTERVAL_SECONDS = 0.150 -TEXT_FLUSH_CHARACTER_LIMIT = 384 - - -async def _heartbeat_until_stopped( - job_id: str, - worker_id: str, - attempt_count: int, - stop: asyncio.Event, -) -> None: - """在 job 执行期间定期刷新 heartbeat_at,直到 stop 被设置。""" - # 持续检查stop是否为true - while not stop.is_set(): - await asyncio.sleep(settings.JOB_HEARTBEAT_INTERVAL_SECONDS) - if stop.is_set(): - break - await asyncio.to_thread( - job_service.update_heartbeat, - job_id, - worker_id, - attempt_count, - ) - logging.info("[worker] heartbeat job=%s worker=%s", job_id, worker_id) - - -async def _complete_terminal_event( - job: dict[str, Any], - worker_id: str, - payload: dict[str, Any], -) -> bool: - """在一个事务内完成终态事件、聊天保存和 job 成功更新。""" - event_type = payload.get("type") - if event_type == "final_result": - response_data = payload.get("data", {}) - elif event_type == "interrupt": - response_data = { - "type": "human_input_required", - "summary": payload.get("message", ""), - } - else: - return False - - result = payload.get("data") if event_type == "final_result" else payload - return await asyncio.to_thread( - job_service.complete_job_with_chat, - job, - worker_id, - event_type, - payload, - response_data, - result, - ) - - -class OrderedEventWriter: - """按源事件顺序持久化,并按时间或字符数批量合并文字增量。""" - - def __init__(self, job: dict[str, Any], worker_id: str): - """为一个 job 创建独占队列和单一消费协程。""" - self.job = job - self.worker_id = worker_id - self.queue: asyncio.Queue = asyncio.Queue() - self.buffer: dict[str, Any] | None = None - self.persisted_sequences: dict[str, int] = {} - self.terminal_seen = False - self.error: BaseException | None = None - self.task = asyncio.create_task(self._consume()) - - async def submit(self, payload: dict[str, Any]) -> None: - """把一个结构化事件交给单写协程。""" - if self.error: - raise self.error - accepted = asyncio.get_running_loop().create_future() - await self.queue.put((payload, accepted)) - await accepted - - async def close(self) -> None: - """刷新剩余文字并等待所有数据库写入完成。""" - if self.error: - raise self.error - accepted = asyncio.get_running_loop().create_future() - await self.queue.put((None, accepted)) - await accepted - await self.task - if self.error: - raise self.error - - async def _persist(self, payload: dict[str, Any]) -> None: - """按事件类型调用 fenced 普通写入或现有事务终态入口。""" - event_type = payload.get("type", "message") - attempt_count = int(self.job["attempt_count"]) - if event_type in {"final_result", "interrupt"}: - await _complete_terminal_event(self.job, self.worker_id, payload) - self.terminal_seen = True - return - if event_type == "error": - await asyncio.to_thread( - job_service.fail_job, - self.job["job_id"], - self.worker_id, - attempt_count, - payload.get("message", "任务执行失败"), - ) - self.terminal_seen = True - return - await asyncio.to_thread( - job_service.write_event, - self.job["job_id"], - self.worker_id, - attempt_count, - event_type, - payload, - ) - - async def _flush_text(self) -> None: - """把当前文字缓冲合并为一个有序 text_delta。""" - if not self.buffer: - return - stream_id = self.buffer["stream_id"] - sequence = self.persisted_sequences.get(stream_id, 0) + 1 - self.persisted_sequences[stream_id] = sequence - payload = { - key: value - for key, value in self.buffer.items() - if key not in {"chunks", "character_count", "started_at"} - } - payload.update({ - "type": "text_delta", - "sequence": sequence, - "delta": "".join(self.buffer["chunks"]), - }) - self.buffer = None - await self._persist(payload) - - async def _accept_text(self, payload: dict[str, Any]) -> None: - """接收内部 token,并在切流或字符阈值时刷新。""" - stream_id = payload["stream_id"] - if self.buffer and self.buffer["stream_id"] != stream_id: - await self._flush_text() - if not self.buffer: - self.buffer = { - key: value - for key, value in payload.items() - if key not in {"type", "sequence", "delta"} - } - self.buffer.update({ - "stream_id": stream_id, - "chunks": [], - "character_count": 0, - "started_at": asyncio.get_running_loop().time(), - }) - self.buffer["chunks"].append(payload["delta"]) - self.buffer["character_count"] += len(payload["delta"]) - if self.buffer["character_count"] >= TEXT_FLUSH_CHARACTER_LIMIT: - await self._flush_text() - - async def _consume(self) -> None: - """串行消费队列,并按批次起始时间触发 150ms 刷新。""" - try: - while True: - try: - if self.buffer: - elapsed = asyncio.get_running_loop().time() - self.buffer["started_at"] - payload, accepted = await asyncio.wait_for( - self.queue.get(), - timeout=max(0.001, TEXT_FLUSH_INTERVAL_SECONDS - elapsed), - ) - else: - payload, accepted = await self.queue.get() - except asyncio.TimeoutError: - await self._flush_text() - continue - - if payload is None: - await self._flush_text() - accepted.set_result(None) - return - if payload.get("type") == "text_chunk": - await self._accept_text(payload) - else: - await self._flush_text() - await self._persist(payload) - accepted.set_result(None) - except BaseException as exc: - self.error = exc - if "accepted" in locals() and not accepted.done(): - accepted.set_exception(exc) - - -async def _run_job(job: dict[str, Any], graph, worker_id: str) -> None: - """执行单个 job,将 Agent 流式事件落库,并处理终态保存。""" - job_id = job["job_id"] - stop_heartbeat = asyncio.Event() - # 心跳检测协程是否正常进行 - heartbeat_task = asyncio.create_task( - _heartbeat_until_stopped(job_id, worker_id, int(job["attempt_count"]), stop_heartbeat) - ) - writer = OrderedEventWriter(job, worker_id) - - try: - ## 执行 AI流式传输 - logging.info("[worker] start job=%s worker=%s session=%s", job_id, worker_id, job["session_id"]) - async for payload in agent_core.ai_call_stream( - job["message"], - job["user_id"], - f"user-{job['user_id']}", - job["session_id"], - job_id=job_id, - job_attempt=int(job["attempt_count"]), - graph=graph, - ): - await writer.submit(payload) - await writer.close() - - if not writer.terminal_seen: - await asyncio.to_thread( - job_service.complete_job, - job_id, - worker_id, - int(job["attempt_count"]), - {}, - ) - logging.info("[worker] finish job=%s worker=%s", job_id, worker_id) - except Exception as exc: - logging.error("[worker] job failed job=%s worker=%s error=%s", job_id, worker_id, exc, exc_info=True) - safe_message = agent_core.sanitize_public_error(exc) - await asyncio.to_thread( - job_service.fail_job, - job_id, - worker_id, - int(job["attempt_count"]), - safe_message, - ) - ## 如果结束,停止心跳 - finally: - stop_heartbeat.set() - await heartbeat_task - - -async def _run_slot(slot_index: int, checkpoint_pool) -> None: - """启动一个 worker slot,并让它独占一组 MCP session/process 和 graph。""" - from Agent.causal_agent.graph import create_graph_from_tools - - worker_id = f"{socket.gethostname()}:{slot_index}" - stack = AsyncExitStack() - try: - # 独占session和graph - mcp_resources = await agent_core.open_mcp_client_resources(stack) - checkpointer = build_checkpointer(checkpoint_pool) - graph = create_graph_from_tools( - agent_core.llm, - mcp_resources.tools, - checkpointer, - ) - logging.info( - "[worker] slot ready worker=%s tools=%s", - worker_id, - [tool.name for tool in mcp_resources.tools], - ) - - # 死循环领取job - while True: - # 线程池 - job = await asyncio.to_thread(job_service.claim_next_job, worker_id) - if not job: - # 休止JOB_POLL_INTERVAL_SECONDS秒 - await asyncio.sleep(settings.JOB_POLL_INTERVAL_SECONDS) - continue - await _run_job(job, graph, worker_id) - finally: - # 安全释放stack - await stack.aclose() - - -async def _main_async() -> None: - """初始化配置、数据库、LLM/RAG,然后启动固定数量的 worker slots。""" - check_database_readiness() - async with open_checkpoint_pool() as checkpoint_pool: - await verify_checkpoint_schema(checkpoint_pool) - if not agent_core.initialize_llm(): - raise RuntimeError("LLM 初始化失败") - if not agent_core.initialize_rag_system(): - logging.warning("RAG 系统初始化失败,worker 将以无知识库模式运行。") - - slot_count = max(1, settings.JOB_WORKERS) - logging.info("[worker] starting slot_count=%s", slot_count) - # 获取_run_slot所有返回,并解包,单并不结束 - await asyncio.gather( - *[_run_slot(i + 1, checkpoint_pool) for i in range(slot_count)] - ) - - -def main() -> None: - """命令行入口:python -m app.agent.worker。""" - if sys.platform == "win32": - asyncio.set_event_loop_policy(asyncio.WindowsProactorEventLoopPolicy()) - logging.basicConfig( - level=logging.INFO, - format="%(asctime)s - %(levelname)s - %(message)s", - force=True, - ) - asyncio.run(_main_async()) - - -if __name__ == "__main__": - main() diff --git a/app/agent/worker/__init__.py b/app/agent/worker/__init__.py new file mode 100644 index 0000000..f4a6108 --- /dev/null +++ b/app/agent/worker/__init__.py @@ -0,0 +1,9 @@ +"""analysis job worker package。""" + +from app.agent.worker.event_writer import ( + OrderedEventWriter, + TEXT_FLUSH_CHARACTER_LIMIT, +) + + +__all__ = ["OrderedEventWriter", "TEXT_FLUSH_CHARACTER_LIMIT"] diff --git a/app/agent/worker/__main__.py b/app/agent/worker/__main__.py new file mode 100644 index 0000000..f5c3bc6 --- /dev/null +++ b/app/agent/worker/__main__.py @@ -0,0 +1,7 @@ +"""支持 ``python -m app.agent.worker`` 的模块入口。""" + +from app.agent.worker.bootstrap import main + + +if __name__ == "__main__": + main() diff --git a/app/agent/worker/bootstrap.py b/app/agent/worker/bootstrap.py new file mode 100644 index 0000000..3754bba --- /dev/null +++ b/app/agent/worker/bootstrap.py @@ -0,0 +1,84 @@ +"""编排 worker 启动检查、runtime 创建和 slot 生命周期。""" + +from __future__ import annotations + +import asyncio +from contextlib import AsyncExitStack +import logging +import socket +import sys +from typing import Any + +from Agent.causal_agent.postgres_checkpointer import ( + open_checkpoint_pool, + verify_checkpoint_schema, +) +from app.agent import job_service +from app.agent.worker.execution import run_job +from app.agent.worker.runtime import ( + ProcessRuntime, + create_process_runtime, + create_slot_runtime, +) +from app.db import check_database_readiness +from config.settings import settings + + +async def run_slot( + slot_index: int, + checkpoint_pool: Any, + process_runtime: ProcessRuntime, +) -> None: + """启动一个 slot,并让其独占 MCP session、tools 和 graph。""" + worker_id = f"{socket.gethostname()}:{slot_index}" + stack = AsyncExitStack() + try: + slot_runtime = await create_slot_runtime( + process_runtime, + stack, + checkpoint_pool, + ) + logging.info( + "[worker] slot ready worker=%s tools=%s", + worker_id, + [tool.name for tool in slot_runtime.mcp_tools], + ) + while True: + job = await asyncio.to_thread(job_service.claim_next_job, worker_id) + if not job: + await asyncio.sleep(settings.JOB_POLL_INTERVAL_SECONDS) + continue + await run_job(job, slot_runtime, worker_id) + finally: + await stack.aclose() + + +async def main_async() -> None: + """执行数据库、checkpoint、LLM/RAG 检查并启动固定数量的 slots。""" + check_database_readiness() + async with open_checkpoint_pool() as checkpoint_pool: + await verify_checkpoint_schema(checkpoint_pool) + process_runtime = create_process_runtime() + if not process_runtime.rag_available: + logging.warning("RAG 系统不可用,worker 将以无知识库模式运行。") + + slot_count = max(1, settings.JOB_WORKERS) + logging.info("[worker] starting slot_count=%s", slot_count) + await asyncio.gather( + *[ + run_slot(index + 1, checkpoint_pool, process_runtime) + for index in range(slot_count) + ] + ) + + +def main() -> None: + """命令行入口:``python -m app.agent.worker``。""" + if sys.platform == "win32": + asyncio.set_event_loop_policy(asyncio.WindowsProactorEventLoopPolicy()) + logging.basicConfig( + level=logging.INFO, + format="%(asctime)s - %(levelname)s - %(message)s", + force=True, + ) + asyncio.run(main_async()) diff --git a/app/agent/worker/event_adapter.py b/app/agent/worker/event_adapter.py new file mode 100644 index 0000000..de2849c --- /dev/null +++ b/app/agent/worker/event_adapter.py @@ -0,0 +1,285 @@ +"""把 LangGraph v2 内部流转换为稳定的公开任务事件。""" + +from __future__ import annotations + +import hashlib +import time +from typing import Any + + +NODE_DESCRIPTIONS = { + "agent": "分析用户意图", + "fold": "加载文件并验证数据", + "preprocess": "预处理数据", + "mcp": "执行因果分析", + "rag": "检索知识库", + "postprocess": "校验并修正因果图", + "report": "生成分析报告", + "normal_chat": "生成回答", + "inquiry_answer": "回答报告追问", +} +TEXT_STREAM_NODES = {"normal_chat", "inquiry_answer"} +TOOL_STAGE_NODES = { + "mcp_planner": "mcp", + "mcp_tool_node": "mcp", + "mcp_result_parser": "mcp", + "rag_question_planner": "rag", + "rag_tool_node": "rag", + "rag_result_parser": "rag", +} +DECISION_PROGRESS = { + "fold": "已识别为因果分析请求", + "normal_chat": "已识别为普通问答", + "inquiry_answer": "已识别为报告追问", + "postprocess": "已识别为已有结果的后续处理", + "preprocess": "文件与数据验证完成", + "agent": "需要补充分析输入", +} +DECISION_FIELDS = { + "agent": "route_decision", + "fold": "fold_decision", +} + + +def _opaque_id(*parts: Any) -> str: + """根据内部标识生成不暴露 attempt/task 内容的稳定 ID。""" + raw = "\x1f".join(str(part) for part in parts) + return hashlib.sha256(raw.encode("utf-8")).hexdigest()[:24] + + +def sanitize_public_error(error: Any) -> str: + """把内部异常归类为有限的用户可见错误,避免泄露连接和路径信息。""" + name = error if isinstance(error, str) else type(error).__name__ + normalized = str(name).lower() + if "timeout" in normalized: + return "调用超时" + if any(token in normalized for token in ("connection", "connect", "network")): + return "服务连接失败" + if any(token in normalized for token in ("rate", "limit")): + return "服务当前繁忙" + if any(token in normalized for token in ("permission", "auth")): + return "服务授权失败" + return "节点执行失败" + + +class LangGraphEventAdapter: + """把 LangGraph v2 多流事件转换为稳定、可持久化的公开事件协议。""" + + def __init__(self, job_id: str | None, job_attempt: int): + """初始化单次 job attempt 的阶段、重试和文字流状态。""" + self.job_id = job_id or "untracked" + self.job_attempt = job_attempt + self.steps: dict[str, dict[str, Any]] = {} + self.active_by_node: dict[str, str] = {} + self.failed_attempts: dict[str, str] = {} + self.streams: dict[str, dict[str, Any]] = {} + + def _base(self, event_type: str, step: dict[str, Any]) -> dict[str, Any]: + """构造所有阶段事件共享的持久化字段。""" + return { + "type": event_type, + "step_id": step["step_id"], + "node_name": step["node_name"], + "title": NODE_DESCRIPTIONS[step["node_name"]], + "attempt": self.job_attempt, + } + + def _active_step(self, node_name: str) -> dict[str, Any] | None: + """获取指定父图节点当前尚未结束的阶段实例。""" + task_id = self.active_by_node.get(node_name) + return self.steps.get(task_id) if task_id else None + + def _parent_tool_step(self, node_name: str) -> dict[str, Any] | None: + """把子图内部节点映射到当前 MCP/RAG 父阶段。""" + parent_name = TOOL_STAGE_NODES.get(node_name) + return self._active_step(parent_name) if parent_name else None + + def _task_event(self, namespace: Any, data: Any) -> list[dict[str, Any]]: + """将根图 tasks 开始/结束转换为阶段生命周期事件。""" + if namespace or not isinstance(data, dict): + return [] + task_id = str(data.get("id") or "") + node_name = data.get("name") + if not task_id or node_name not in NODE_DESCRIPTIONS: + return [] + if "input" in data: + step = { + "step_id": _opaque_id(self.job_id, self.job_attempt, task_id), + "node_name": node_name, + "started_at": time.monotonic(), + } + self.steps[task_id] = step + self.active_by_node[node_name] = task_id + return [self._base("node_start", step)] + step = self.steps.get(task_id) + if not step: + return [] + event = self._base("node_end", step) + event["duration"] = round( + max(0.0, time.monotonic() - step["started_at"]), + 2, + ) + event["status"] = "failed" if data.get("error") else "completed" + if data.get("error"): + event["message"] = sanitize_public_error(data.get("error")) + if self.active_by_node.get(node_name) == task_id: + self.active_by_node.pop(node_name, None) + return [event] + + def _custom_event(self, data: Any) -> list[dict[str, Any]]: + """只有失败后再次开始时,才把内部 attempt 转换为真正的重试。""" + if not isinstance(data, dict): + return [] + event_type = data.get("type") + task_id = str(data.get("task_id") or "") + node_name = data.get("node_name") + if event_type == "node_attempt_failed" and task_id: + self.failed_attempts[task_id] = sanitize_public_error( + data.get("error_kind") + ) + return [] + if event_type != "node_attempt_start" or not task_id: + return [] + failure = self.failed_attempts.pop(task_id, None) + if not failure: + return [] + step = ( + self.steps.get(task_id) + or self._active_step(node_name) + or self._parent_tool_step(node_name) + ) + if not step: + return [] + discarded = self.streams.pop(step["step_id"], None) + event = self._base("node_retry", step) + event["message"] = failure + if discarded: + event["discard_stream_id"] = discarded["stream_id"] + return [event] + + @staticmethod + def _tool_call(message: Any) -> tuple[str | None, list[str]]: + """只读取工具名和参数字段名,不返回参数值。""" + calls = getattr(message, "tool_calls", None) or [] + if not calls: + return None, [] + call = calls[0] + if isinstance(call, dict): + name = call.get("name") + args = call.get("args") + else: + name = getattr(call, "name", None) + args = getattr(call, "args", None) + keys = sorted(str(key) for key in args)[:12] if isinstance(args, dict) else [] + return name, keys + + def _update_event(self, namespace: Any, data: Any) -> list[dict[str, Any]]: + """从显式 State 和规范化结果生成 decision、progress 与工具摘要。""" + if not isinstance(data, dict): + return [] + events: list[dict[str, Any]] = [] + for node_name, output in data.items(): + if not isinstance(output, dict): + continue + step = self._active_step(node_name) + if step: + decision_field = DECISION_FIELDS.get(node_name) + decision = output.get(decision_field) if decision_field else None + if decision: + progress = self._base("progress", step) + progress["summary"] = DECISION_PROGRESS.get( + decision, + "路由判断完成", + ) + events.append(progress) + event = self._base("decision", step) + event["summary"] = ( + f"已决定进入:{NODE_DESCRIPTIONS.get(decision, decision)}" + ) + events.append(event) + tool_step = self._parent_tool_step(node_name) + if not tool_step: + continue + messages = output.get("messages") or [] + latest = messages[-1] if messages else None + tool_name, argument_keys = self._tool_call(latest) + if tool_name: + event = self._base("tool_call_start", tool_step) + event.update( + {"tool_name": tool_name, "argument_keys": argument_keys} + ) + events.append(event) + result_key = ( + "causal_analysis_result" + if TOOL_STAGE_NODES[node_name] == "mcp" + else "knowledge_base_result" + ) + result = output.get(result_key) + if isinstance(result, dict): + metadata = result.get("_tool_call") or {} + event = self._base("tool_call_result", tool_step) + event.update( + { + "tool_name": metadata.get("name") or tool_name or "工具", + "summary": ( + "调用完成" + if result.get("success") is not False + else "调用失败" + ), + } + ) + events.append(event) + return events + + def _message_event(self, data: Any) -> list[dict[str, Any]]: + """仅转换普通问答和报告追问的非空字符串 token。""" + if not isinstance(data, (tuple, list)) or len(data) != 2: + return [] + chunk, metadata = data + if not isinstance(metadata, dict): + return [] + node_name = metadata.get("langgraph_node") + content = getattr(chunk, "content", None) + if ( + node_name not in TEXT_STREAM_NODES + or not isinstance(content, str) + or not content + ): + return [] + step = self._active_step(node_name) + if not step: + return [] + stream = self.streams.get(step["step_id"]) + if not stream: + stream = { + "stream_id": _opaque_id(step["step_id"], time.monotonic_ns()), + "sequence": 0, + } + self.streams[step["step_id"]] = stream + stream["sequence"] += 1 + event = self._base("text_chunk", step) + event.update( + { + "stream_id": stream["stream_id"], + "sequence": stream["sequence"], + "delta": content, + } + ) + return [event] + + def convert(self, chunk: Any) -> list[dict[str, Any]]: + """转换一个 v2 StreamPart;无法识别的内部流会被忽略。""" + if not isinstance(chunk, dict): + return [] + stream_type = chunk.get("type") + namespace = chunk.get("ns") or () + data = chunk.get("data") + if stream_type == "tasks": + return self._task_event(namespace, data) + if stream_type == "custom": + return self._custom_event(data) + if stream_type == "updates": + return self._update_event(namespace, data) + if stream_type == "messages": + return self._message_event(data) + return [] diff --git a/app/agent/worker/event_writer.py b/app/agent/worker/event_writer.py new file mode 100644 index 0000000..9d458d9 --- /dev/null +++ b/app/agent/worker/event_writer.py @@ -0,0 +1,186 @@ +"""按 LangGraph 源事件顺序持久化 worker 事件。""" + +from __future__ import annotations + +import asyncio +from typing import Any + +from app.agent import job_service + + +TEXT_FLUSH_INTERVAL_SECONDS = 0.150 +TEXT_FLUSH_CHARACTER_LIMIT = 384 + + +async def _complete_terminal_event( + job: dict[str, Any], + worker_id: str, + payload: dict[str, Any], +) -> bool: + """在一个事务内完成终态事件、聊天保存和 job 成功更新。""" + event_type = payload.get("type") + if event_type == "final_result": + response_data = payload.get("data", {}) + elif event_type == "interrupt": + response_data = { + "type": "human_input_required", + "summary": payload.get("message", ""), + } + else: + return False + + result = payload.get("data") if event_type == "final_result" else payload + return await asyncio.to_thread( + job_service.complete_job_with_chat, + job, + worker_id, + event_type, + payload, + response_data, + result, + ) + + +class OrderedEventWriter: + """按源事件顺序持久化,并按时间或字符数批量合并文字增量。""" + + def __init__(self, job: dict[str, Any], worker_id: str): + """为一个 job 创建独占队列和单一消费协程。""" + self.job = job + self.worker_id = worker_id + self.queue: asyncio.Queue = asyncio.Queue() + self.buffer: dict[str, Any] | None = None + self.persisted_sequences: dict[str, int] = {} + self.terminal_seen = False + self.error: BaseException | None = None + self.task = asyncio.create_task(self._consume()) + + async def submit(self, payload: dict[str, Any]) -> None: + """把一个结构化事件交给单写协程。""" + if self.error: + raise self.error + accepted = asyncio.get_running_loop().create_future() + await self.queue.put((payload, accepted)) + await accepted + + async def close(self) -> None: + """刷新剩余文字并等待所有数据库写入完成。""" + if self.error: + raise self.error + accepted = asyncio.get_running_loop().create_future() + await self.queue.put((None, accepted)) + await accepted + await self.task + if self.error: + raise self.error + + async def _persist(self, payload: dict[str, Any]) -> None: + """按事件类型调用 fenced 普通写入或现有事务终态入口。""" + event_type = payload.get("type", "message") + attempt_count = int(self.job["attempt_count"]) + if event_type in {"final_result", "interrupt"}: + await _complete_terminal_event(self.job, self.worker_id, payload) + self.terminal_seen = True + return + if event_type == "error": + await asyncio.to_thread( + job_service.fail_job, + self.job["job_id"], + self.worker_id, + attempt_count, + payload.get("message", "任务执行失败"), + ) + self.terminal_seen = True + return + await asyncio.to_thread( + job_service.write_event, + self.job["job_id"], + self.worker_id, + attempt_count, + event_type, + payload, + ) + + async def _flush_text(self) -> None: + """把当前文字缓冲合并为一个有序 text_delta。""" + if not self.buffer: + return + stream_id = self.buffer["stream_id"] + sequence = self.persisted_sequences.get(stream_id, 0) + 1 + self.persisted_sequences[stream_id] = sequence + payload = { + key: value + for key, value in self.buffer.items() + if key not in {"chunks", "character_count", "started_at"} + } + payload.update( + { + "type": "text_delta", + "sequence": sequence, + "delta": "".join(self.buffer["chunks"]), + } + ) + self.buffer = None + await self._persist(payload) + + async def _accept_text(self, payload: dict[str, Any]) -> None: + """接收内部 token,并在切流或字符阈值时刷新。""" + stream_id = payload["stream_id"] + if self.buffer and self.buffer["stream_id"] != stream_id: + await self._flush_text() + if not self.buffer: + self.buffer = { + key: value + for key, value in payload.items() + if key not in {"type", "sequence", "delta"} + } + self.buffer.update( + { + "stream_id": stream_id, + "chunks": [], + "character_count": 0, + "started_at": asyncio.get_running_loop().time(), + } + ) + self.buffer["chunks"].append(payload["delta"]) + self.buffer["character_count"] += len(payload["delta"]) + if self.buffer["character_count"] >= TEXT_FLUSH_CHARACTER_LIMIT: + await self._flush_text() + + async def _consume(self) -> None: + """串行消费队列,并按批次起始时间触发定时刷新。""" + try: + while True: + try: + if self.buffer: + elapsed = ( + asyncio.get_running_loop().time() + - self.buffer["started_at"] + ) + payload, accepted = await asyncio.wait_for( + self.queue.get(), + timeout=max( + 0.001, + TEXT_FLUSH_INTERVAL_SECONDS - elapsed, + ), + ) + else: + payload, accepted = await self.queue.get() + except asyncio.TimeoutError: + await self._flush_text() + continue + + if payload is None: + await self._flush_text() + accepted.set_result(None) + return + if payload.get("type") == "text_chunk": + await self._accept_text(payload) + else: + await self._flush_text() + await self._persist(payload) + accepted.set_result(None) + except BaseException as exc: + self.error = exc + if "accepted" in locals() and not accepted.done(): + accepted.set_exception(exc) diff --git a/app/agent/worker/execution.py b/app/agent/worker/execution.py new file mode 100644 index 0000000..423274c --- /dev/null +++ b/app/agent/worker/execution.py @@ -0,0 +1,102 @@ +"""执行单个 analysis job 并维护其 heartbeat。""" + +from __future__ import annotations + +import asyncio +import logging +from typing import Any + +from app.agent import job_service +from app.agent.worker.event_adapter import sanitize_public_error +from app.agent.worker.event_writer import OrderedEventWriter +from app.agent.worker.graph_runner import ai_call_stream +from app.agent.worker.runtime import SlotRuntime +from config.settings import settings + + +async def heartbeat_until_stopped( + job_id: str, + worker_id: str, + attempt_count: int, + stop: asyncio.Event, +) -> None: + """在 job 执行期间定期刷新 heartbeat,直到收到停止信号。""" + while not stop.is_set(): + await asyncio.sleep(settings.JOB_HEARTBEAT_INTERVAL_SECONDS) + if stop.is_set(): + break + await asyncio.to_thread( + job_service.update_heartbeat, + job_id, + worker_id, + attempt_count, + ) + logging.info("[worker] heartbeat job=%s worker=%s", job_id, worker_id) + + +async def run_job( + job: dict[str, Any], + slot_runtime: SlotRuntime, + worker_id: str, +) -> None: + """使用显式 slot runtime 执行 job、写事件并处理终态。""" + job_id = job["job_id"] + attempt_count = int(job["attempt_count"]) + stop_heartbeat = asyncio.Event() + heartbeat_task = asyncio.create_task( + heartbeat_until_stopped( + job_id, + worker_id, + attempt_count, + stop_heartbeat, + ) + ) + writer = OrderedEventWriter(job, worker_id) + + try: + logging.info( + "[worker] start job=%s worker=%s session=%s tools=%s", + job_id, + worker_id, + job["session_id"], + len(slot_runtime.mcp_tools), + ) + async for payload in ai_call_stream( + job["message"], + job["user_id"], + f"user-{job['user_id']}", + job["session_id"], + job_id=job_id, + job_attempt=attempt_count, + graph=slot_runtime.graph, + ): + await writer.submit(payload) + await writer.close() + + if not writer.terminal_seen: + await asyncio.to_thread( + job_service.complete_job, + job_id, + worker_id, + attempt_count, + {}, + ) + logging.info("[worker] finish job=%s worker=%s", job_id, worker_id) + except Exception as exc: + logging.error( + "[worker] job failed job=%s worker=%s error=%s", + job_id, + worker_id, + exc, + exc_info=True, + ) + await asyncio.to_thread( + job_service.fail_job, + job_id, + worker_id, + attempt_count, + sanitize_public_error(exc), + ) + finally: + stop_heartbeat.set() + await heartbeat_task diff --git a/app/agent/worker/graph_runner.py b/app/agent/worker/graph_runner.py new file mode 100644 index 0000000..43742f6 --- /dev/null +++ b/app/agent/worker/graph_runner.py @@ -0,0 +1,121 @@ +"""执行单个 LangGraph graph 并生成稳定的内部流事件。""" + +from __future__ import annotations + +import logging +from typing import Any, AsyncIterator + +from langchain_core.messages import HumanMessage +from langgraph.types import Command + +from app.agent.worker.event_adapter import ( + LangGraphEventAdapter, + sanitize_public_error, +) +from app.agent.worker.result_presenter import process_final_result + + +def _snapshot_interrupts(snapshot: Any) -> list[Any]: + """汇总 LangGraph StateSnapshot 中尚待恢复的 task interrupts。""" + return [ + item + for task in (getattr(snapshot, "tasks", None) or ()) + for item in (getattr(task, "interrupts", None) or ()) + ] + + +async def ai_call_stream( + text: str, + user_id: int, + username: str, + session_id: str, + *, + job_id: str | None, + job_attempt: int, + graph: Any, +) -> AsyncIterator[dict[str, Any]]: + """在显式传入的 graph 上执行一次调用并产出公开事件。""" + if graph is None: + raise RuntimeError("Agent Graph 尚未初始化") + + logging.info("[流式] 处理用户 %s 的消息,会话ID: %s", username, session_id) + config = { + "configurable": { + "thread_id": session_id, + "user_id": user_id, + }, + "metadata": {"job_id": job_id} if job_id else {}, + } + + try: + state = await graph.aget_state(config) + if _snapshot_interrupts(state): + logging.info("[流式] 检测到会话 %s 处于中断状态", session_id) + input_data: Any = Command(resume=text) + else: + input_data = { + "messages": [HumanMessage(content=text)], + "user_id": user_id, + "username": username, + "session_id": session_id, + } + except Exception as exc: + logging.warning("[流式] 无法获取状态,按新对话处理: %s", exc) + input_data = { + "messages": [HumanMessage(content=text)], + "user_id": user_id, + "username": username, + "session_id": session_id, + } + + adapter = LangGraphEventAdapter(job_id, job_attempt) + streamed_interrupts: list[Any] = [] + try: + async for chunk in graph.astream( + input_data, + config, + stream_mode=["updates", "messages", "custom", "tasks"], + subgraphs=True, + version="v2", + ): + if chunk.get("type") == "updates" and isinstance(chunk.get("data"), dict): + interrupt_data = chunk["data"].get("__interrupt__") + if interrupt_data: + streamed_interrupts.extend( + interrupt_data + if isinstance(interrupt_data, (list, tuple)) + else [interrupt_data] + ) + for event_data in adapter.convert(chunk): + yield event_data + + state = await graph.aget_state(config) + interrupts = _snapshot_interrupts(state) or streamed_interrupts + if interrupts: + interrupt_obj = interrupts[0] + question = ( + interrupt_obj.value + if hasattr(interrupt_obj, "value") + else str(interrupt_obj) + ) + yield { + "type": "interrupt", + "message": question, + "attempt": job_attempt, + } + logging.info("[SSE] 图已暂停,等待用户输入") + return + + yield { + "type": "final_result", + "data": process_final_result(state.values), + "attempt": job_attempt, + } + logging.info("[SSE] 发送最终结果") + except Exception as exc: + logging.error("[流式] 执行 LangGraph Agent 时发生错误: %s", exc, exc_info=True) + yield { + "type": "error", + "message": sanitize_public_error(exc), + "attempt": job_attempt, + } diff --git a/app/agent/worker/result_presenter.py b/app/agent/worker/result_presenter.py new file mode 100644 index 0000000..a94f337 --- /dev/null +++ b/app/agent/worker/result_presenter.py @@ -0,0 +1,78 @@ +"""把 Agent 最终状态转换为聊天与 SSE 共用的公开结果。""" + +from __future__ import annotations + +import logging +from typing import Any + +from langchain_core.messages import AIMessage + +from app.chat.response_storage import render_summary_for_display + + +def process_final_result(final_state_data: dict[str, Any]) -> dict[str, Any]: + """按消息、报告和因果图优先级生成稳定的最终响应。""" + messages = final_state_data.get("messages", []) + if messages: + last_message = messages[-1] + if isinstance(last_message, AIMessage): + message_name = getattr(last_message, "name", None) + if message_name in {"normal_chat", "inquiry_answer"}: + logging.info("返回 %s 节点的回复", message_name) + return {"type": "text", "summary": last_message.content} + + if message_name == "report" and final_state_data.get("final_report"): + logging.info("返回完整的因果分析报告") + result: dict[str, Any] = { + "summary": final_state_data["final_report"], + "layout": "report", + } + analysis_data = final_state_data.get("causal_analysis_result") + if isinstance(analysis_data, dict) and analysis_data.get("success"): + original_graph = analysis_data.get("data") + postprocess_result = final_state_data.get("postprocess_result") or {} + revised_graph = postprocess_result.get("revised_graph") + has_valid_revised_graph = ( + isinstance(revised_graph, dict) + and isinstance(revised_graph.get("nodes"), list) + and isinstance(revised_graph.get("edges"), list) + and not postprocess_result.get("error") + ) + result["type"] = "causal_graph" + result["data"] = ( + revised_graph if has_valid_revised_graph else original_graph + ) + result["graph_source"] = ( + "postprocessed" if has_valid_revised_graph else "original" + ) + result["revision_summary"] = postprocess_result.get( + "revision_summary", + "", + ) + logging.info( + "返回因果图数据: source=%s", + result["graph_source"], + ) + + result.setdefault("type", "text") + visualization_mapping = final_state_data.get("visualization_mapping") + if visualization_mapping: + result["raw_summary"] = result["summary"] + result["visualization_mapping"] = visualization_mapping + result["summary"] = render_summary_for_display( + result["summary"], + visualization_mapping, + ) + logging.info( + "已替换报告中的 %s 个可视化占位符", + len(visualization_mapping), + ) + return result + + final_report = final_state_data.get("final_report") + if final_report: + logging.info("未找到最新消息,降级返回 final_report") + return {"type": "text", "summary": final_report, "layout": "report"} + + logging.warning("未找到任何可返回的内容,返回默认消息") + return {"type": "text", "summary": "抱歉,我在处理时遇到了问题。"} diff --git a/app/agent/worker/runtime.py b/app/agent/worker/runtime.py new file mode 100644 index 0000000..b021f74 --- /dev/null +++ b/app/agent/worker/runtime.py @@ -0,0 +1,136 @@ +"""构建 worker 进程级与 slot 级运行时依赖。""" + +from __future__ import annotations + +from contextlib import AsyncExitStack +from dataclasses import dataclass +import logging +from pathlib import Path +import sys +from typing import Any + +from langchain_openai import ChatOpenAI + +from Agent.causal_agent.postgres_checkpointer import build_checkpointer +from config.settings import settings + + +PROJECT_ROOT = Path(__file__).resolve().parents[3] +MCP_SERVER_PATH = PROJECT_ROOT / "Agent" / "CausalAgentMCP" / "mcp_server.py" +KNOWLEDGE_BASE_DIRECTORY = PROJECT_ROOT / "Agent" / "knowledge_base" + + +@dataclass(frozen=True) +class ProcessRuntime: + """保存 worker 进程内共享、且不依赖 slot 生命周期的对象。""" + + llm: ChatOpenAI + rag_available: bool + + +@dataclass(frozen=True) +class McpClientResources: + """保存由 slot 的 ``AsyncExitStack`` 管理的 MCP 资源。""" + + client: Any + session: Any + tools: list[Any] + + +@dataclass(frozen=True) +class SlotRuntime: + """显式保存一个 slot 的 LLM、MCP 生命周期资源、tools 与 graph。""" + + llm: ChatOpenAI + mcp_resources: McpClientResources + mcp_tools: list[Any] + graph: Any + + +def create_llm() -> ChatOpenAI: + """根据当前配置创建 LLM;配置缺失时立即失败。""" + if not all([settings.MODEL, settings.BASE_URL, settings.API_KEY]): + raise RuntimeError("LLM 配置不完整,无法初始化") + + logging.info("正在初始化 LLM 模型: %s", settings.MODEL) + llm = ChatOpenAI( + model=settings.MODEL, + base_url=settings.BASE_URL, + api_key=settings.API_KEY, + streaming=False, + ) + logging.info("LLM 实例初始化成功。") + return llm + + +def check_rag_availability() -> bool: + """只检查知识库目录,实际向量库仍在首次查询时延迟加载。""" + logging.info("正在检查 RAG 知识库目录...") + persist_directory = KNOWLEDGE_BASE_DIRECTORY / "db" + if not persist_directory.exists(): + logging.warning( + "知识库持久化目录不存在。请先运行 " + "Agent/knowledge_base/build_knowledge.py 构建知识库。" + ) + return False + + logging.info("RAG 启动检查通过;向量库将在首次实际查询时延迟初始化。") + return True + + +def create_process_runtime() -> ProcessRuntime: + """创建一次进程级运行时,消除对 ``app.agent.core`` 全局变量的依赖。""" + return ProcessRuntime( + llm=create_llm(), + rag_available=check_rag_availability(), + ) + + +async def open_mcp_client_resources( + process_stack: AsyncExitStack, +) -> McpClientResources: + """为一个 slot 打开持久 MCP session 并加载 LangChain tools。""" + try: + from langchain_mcp_adapters.client import MultiServerMCPClient + from langchain_mcp_adapters.tools import load_mcp_tools + except ImportError as exc: + raise RuntimeError( + "缺少 langchain-mcp-adapters,无法初始化 LangChain MCP adapter。" + ) from exc + + client = MultiServerMCPClient( + { + "causal": { + "transport": "stdio", + "command": sys.executable, + "args": [str(MCP_SERVER_PATH)], + } + } + ) + logging.info("MCP 初始化。") + session = await process_stack.enter_async_context(client.session("causal")) + tools = await load_mcp_tools(session) + return McpClientResources(client=client, session=session, tools=tools) + + +async def create_slot_runtime( + process_runtime: ProcessRuntime, + process_stack: AsyncExitStack, + checkpoint_pool: Any, +) -> SlotRuntime: + """创建一个 slot 独占的 MCP 资源和 graph,并显式返回其依赖。""" + from Agent.causal_agent.graph import create_graph_from_tools + + mcp_resources = await open_mcp_client_resources(process_stack) + checkpointer = build_checkpointer(checkpoint_pool) + graph = create_graph_from_tools( + process_runtime.llm, + mcp_resources.tools, + checkpointer, + ) + return SlotRuntime( + llm=process_runtime.llm, + mcp_resources=mcp_resources, + mcp_tools=mcp_resources.tools, + graph=graph, + ) diff --git a/tests/unit/admin/test_admin_write_service.py b/tests/unit/admin/test_admin_write_service.py index 7b5e32d..6ceaaf4 100644 --- a/tests/unit/admin/test_admin_write_service.py +++ b/tests/unit/admin/test_admin_write_service.py @@ -129,12 +129,14 @@ def test_checkpoint_runtime_and_admin_reads_use_postgres(self): inspection = Path("Database/checkpoint_inspection.py").read_text( encoding="utf-8" ) - core = Path("app/agent/core.py").read_text(encoding="utf-8") + graph_runner = Path("app/agent/worker/graph_runner.py").read_text( + encoding="utf-8" + ) self.assertIn("CREATE TABLE checkpoint_cleanup_outbox", migration) self.assertIn("DROP TABLE IF EXISTS checkpoints", migration) self.assertIn("metadata ->> 'job_id' = %s", inspection) self.assertIn("ORDER BY checkpoint_id DESC", inspection) - self.assertIn('"job_id": job_id', core) + self.assertIn('"job_id": job_id', graph_runner) def test_lifecycle_repair_is_dry_run_and_requires_database_confirmation(self): """孤立修复 CLI 默认 dry-run,apply 必须精确确认数据库。""" diff --git a/tests/unit/agent/test_final_result_graph_selection.py b/tests/unit/agent/test_final_result_graph_selection.py index 50ee978..6d86707 100644 --- a/tests/unit/agent/test_final_result_graph_selection.py +++ b/tests/unit/agent/test_final_result_graph_selection.py @@ -19,7 +19,7 @@ from langchain_core.messages import AIMessage -from app.agent.core import process_final_result +from app.agent.worker.result_presenter import process_final_result from app.chat.response_storage import prepare_ai_response_for_storage diff --git a/tests/unit/agent/test_ordered_event_writer.py b/tests/unit/agent/test_ordered_event_writer.py index cc3d9f8..b06da68 100644 --- a/tests/unit/agent/test_ordered_event_writer.py +++ b/tests/unit/agent/test_ordered_event_writer.py @@ -17,7 +17,7 @@ for key, value in TEST_ENV.items(): os.environ.setdefault(key, value) -from app.agent.worker import ( # noqa: E402 +from app.agent.worker.event_writer import ( # noqa: E402 OrderedEventWriter, TEXT_FLUSH_CHARACTER_LIMIT, ) diff --git a/tests/unit/agent/test_stream_events.py b/tests/unit/agent/test_stream_events.py index 1ede407..0463ebf 100644 --- a/tests/unit/agent/test_stream_events.py +++ b/tests/unit/agent/test_stream_events.py @@ -19,7 +19,7 @@ for key, value in TEST_ENV.items(): os.environ.setdefault(key, value) -from app.agent.core import LangGraphEventAdapter # noqa: E402 +from app.agent.worker.event_adapter import LangGraphEventAdapter # noqa: E402 from app.agent.routes import _public_event_payload # noqa: E402 from Agent.causal_agent.graph_utils import bind_node, bind_runnable_node # noqa: E402 from Agent.causal_agent.state import CausalAgentState # noqa: E402 diff --git a/tests/unit/agent/test_worker_runtime.py b/tests/unit/agent/test_worker_runtime.py new file mode 100644 index 0000000..a80eb1d --- /dev/null +++ b/tests/unit/agent/test_worker_runtime.py @@ -0,0 +1,87 @@ +"""验证 worker runtime 的依赖显式性和 slot 隔离边界。""" + +import asyncio +from contextlib import AsyncExitStack +import runpy +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock, patch + +from app.agent.worker.runtime import ( + McpClientResources, + ProcessRuntime, + create_process_runtime, + create_slot_runtime, +) + + +def test_worker_package_entrypoint_calls_bootstrap_main(): + """包入口应继续把 ``python -m app.agent.worker`` 转交给 bootstrap。""" + with patch("app.agent.worker.bootstrap.main") as main: + runpy.run_module("app.agent.worker", run_name="__main__") + + main.assert_called_once_with() + + +def test_core_facade_does_not_hold_runtime_globals(): + """兼容门面可以转发纯接口,但不能重新成为可变运行时容器。""" + from app.agent import core + + assert not hasattr(core, "llm") + assert not hasattr(core, "agent_graph") + assert not hasattr(core, "mcp_session") + + +def test_process_runtime_returns_explicit_llm_and_rag_state(): + """进程初始化应返回值对象,而不是写入 core 模块全局变量。""" + llm = Mock(name="llm") + with ( + patch("app.agent.worker.runtime.create_llm", return_value=llm), + patch("app.agent.worker.runtime.check_rag_availability", return_value=False), + ): + runtime = create_process_runtime() + + assert runtime.llm is llm + assert runtime.rag_available is False + + +def test_slot_runtime_builds_graph_from_explicit_dependencies(): + """slot graph 应由传入 LLM、该 slot tools 与 checkpoint 共同构建。""" + llm = Mock(name="llm") + tools = [SimpleNamespace(name="causal_test")] + graph = Mock(name="graph") + checkpointer = Mock(name="checkpointer") + resources = McpClientResources( + client=Mock(name="client"), + session=Mock(name="session"), + tools=tools, + ) + process_runtime = ProcessRuntime(llm=llm, rag_available=True) + + with ( + patch( + "app.agent.worker.runtime.open_mcp_client_resources", + new=AsyncMock(return_value=resources), + ), + patch( + "app.agent.worker.runtime.build_checkpointer", + return_value=checkpointer, + ) as build_checkpointer, + patch( + "Agent.causal_agent.graph.create_graph_from_tools", + return_value=graph, + ) as create_graph, + ): + slot_runtime = asyncio.run( + create_slot_runtime( + process_runtime, + AsyncExitStack(), + Mock(name="checkpoint_pool"), + ) + ) + + build_checkpointer.assert_called_once() + create_graph.assert_called_once_with(llm, tools, checkpointer) + assert slot_runtime.llm is llm + assert slot_runtime.mcp_resources is resources + assert slot_runtime.mcp_tools is tools + assert slot_runtime.graph is graph