From 48c4283554e5c2b2c241060c2e4fadf3240a6c30 Mon Sep 17 00:00:00 2001 From: miluELK <2636626273@qq.com> Date: Mon, 27 Oct 2025 15:32:12 +0800 Subject: [PATCH 1/4] =?UTF-8?q?feat:=20=E6=B7=BB=E5=8A=A0=E4=BA=86ReactAge?= =?UTF-8?q?nt,MultiAgent,ToolAgent=E3=80=82=E4=B8=BAState=E6=B7=BB?= =?UTF-8?q?=E5=8A=A0=E4=BA=86=E5=9F=BA=E7=B1=BBBaseState=E3=80=82common/to?= =?UTF-8?q?ols=E4=B8=AD=E6=B7=BB=E5=8A=A0=E4=BA=86=E4=BA=BA=E5=B7=A5?= =?UTF-8?q?=E5=AE=A1=E6=89=B9=E5=B7=A5=E5=85=B7=EF=BC=88=E7=9B=AE=E5=89=8D?= =?UTF-8?q?=E6=97=A0=E6=B3=95=E6=AD=A3=E5=B8=B8=E8=B0=83=E7=94=A8=EF=BC=8C?= =?UTF-8?q?=E5=89=8D=E7=AB=AF=E9=9C=80=E8=A6=81=E5=AE=8C=E5=96=84=E4=BA=BA?= =?UTF-8?q?=E5=B7=A5=E5=AE=A1=E6=89=B9=E6=B5=81=E7=A8=8B=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 5 +- src/agents/__init__.py | 3 +- src/agents/common/base.py | 8 +++ src/agents/common/state.py | 20 +++++++ src/agents/common/toolagent.py | 87 +++++++++++++++++++++++++++++++ src/agents/common/tools.py | 44 ++++++++++++++++ src/agents/multiAgent/__init__.py | 3 ++ src/agents/multiAgent/context.py | 25 +++++++++ src/agents/multiAgent/graph.py | 61 ++++++++++++++++++++++ src/agents/multiAgent/state.py | 20 +++++++ src/agents/multiAgent/tools.py | 63 ++++++++++++++++++++++ src/agents/react/context.py | 25 +++++++++ src/agents/react/graph.py | 81 ++++++++++++++++------------ src/agents/react/state.py | 22 ++++++++ src/agents/react/tools.py | 48 +++++++++++++++++ 15 files changed, 479 insertions(+), 36 deletions(-) create mode 100644 src/agents/common/state.py create mode 100644 src/agents/common/toolagent.py create mode 100644 src/agents/multiAgent/__init__.py create mode 100644 src/agents/multiAgent/context.py create mode 100644 src/agents/multiAgent/graph.py create mode 100644 src/agents/multiAgent/state.py create mode 100644 src/agents/multiAgent/tools.py create mode 100644 src/agents/react/context.py create mode 100644 src/agents/react/state.py create mode 100644 src/agents/react/tools.py diff --git a/.gitignore b/.gitignore index ba8806c8..89f4c324 100644 --- a/.gitignore +++ b/.gitignore @@ -36,11 +36,10 @@ cache .trae .pytest_cache -### (企业私有代码 - 仅忽略敏感配置,不忽略代码文件) -# 移除了 *.private* 和 *_private 规则,允许 Git 本地管理 -# 通过 .git/info/exclude 或本地分支管理私有代码 *.secret* *.nogit* +*_private +*.private # *.local* 保留用于本地配置文件 *.local.py *.local.js diff --git a/src/agents/__init__.py b/src/agents/__init__.py index 5b91bbd4..2b2746fb 100644 --- a/src/agents/__init__.py +++ b/src/agents/__init__.py @@ -3,11 +3,12 @@ import importlib import inspect from pathlib import Path +from server.utils.singleton import SingletonMeta from src.agents.common.base import BaseAgent from src.utils import logger -class AgentManager: +class AgentManager(metaclass=SingletonMeta): def __init__(self): self._classes = {} self._instances = {} # 存储已创建的 agent 实例 diff --git a/src/agents/common/base.py b/src/agents/common/base.py index b3d4fd0d..702368d0 100644 --- a/src/agents/common/base.py +++ b/src/agents/common/base.py @@ -67,6 +67,14 @@ class BaseAgent: ): yield msg, metadata + async def invoke_messages(self, messages: list[str], input_context=None, **kwargs): + graph = await self.get_graph() + context = self.context_schema.from_file(module_name=self.module_name, input_context=input_context) + logger.debug(f"invoke_messages: {context}") + input_config = {"configurable": input_context, "recursion_limit": 100} + msg = await graph.ainvoke({"messages": messages}, context=context, config=input_config) + return msg + async def check_checkpointer(self): app = await self.get_graph() if not hasattr(app, "checkpointer") or app.checkpointer is None: diff --git a/src/agents/common/state.py b/src/agents/common/state.py new file mode 100644 index 00000000..981d013f --- /dev/null +++ b/src/agents/common/state.py @@ -0,0 +1,20 @@ +"""Define the state structures for the agent.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass, field +from typing import Annotated + +from langchain.messages import AnyMessage +from langgraph.graph import add_messages + + +@dataclass +class BaseState: + """Defines the input state for the agent, representing a narrower interface to the outside world. + + This class is used to define the initial state and structure of incoming data. + """ + + messages: Annotated[Sequence[AnyMessage], add_messages] = field(default_factory=list) diff --git a/src/agents/common/toolagent.py b/src/agents/common/toolagent.py new file mode 100644 index 00000000..38e6dfe4 --- /dev/null +++ b/src/agents/common/toolagent.py @@ -0,0 +1,87 @@ +from abc import abstractmethod +from typing import Any, cast + +from langchain.messages import AIMessage, ToolMessage +from langgraph.prebuilt import ToolNode +from langgraph.runtime import Runtime + +from src.agents.common.base import BaseAgent +from src.agents.common.mcp import get_mcp_tools +from src.agents.common.models import load_chat_model +from src.utils import logger + +from .state import BaseState +from .context import BaseContext + +class ToolAgent(BaseAgent): + name = "ToolAgent" + description = "具有工具调用能力的Agent" + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.graph = None + self.checkpointer = None + self.context_schema = BaseContext + self.agent_tools = None + + + # TODO:[修改建议] _get_invoke_tools,llm_call,dynamic_tools_node这类针对工具调用的功能大多数Agent都能用得到 + # 可以通过一个ToolAgent类继承BaseAgent,通过重写抽象方法获取tools,通过继承BaseState和BaseContext获取配置 + # 必要时可通过重写以下方法实现其他逻辑 + @abstractmethod + def get_tools(self): + logger.error(f"get_tools() is not implemented in {self.__class__.__name__}") + return [] + + + async def _get_invoke_tools(self, selected_tools: list[str], selected_mcps: list[str]): + """根据配置获取工具。 + 默认不使用任何工具。 + 如果配置为列表,则使用列表中的工具。 + """ + enabled_tools = [] + self.agent_tools = self.agent_tools or self.get_tools() + if selected_tools and isinstance(selected_tools, list) and len(selected_tools) > 0: + # 使用配置中指定的工具 + enabled_tools = [tool for tool in self.agent_tools if tool.name in selected_tools] + + if selected_mcps and isinstance(selected_mcps, list) and len(selected_mcps) > 0: + for mcp in selected_mcps: + enabled_tools.extend(await get_mcp_tools(mcp)) + + return enabled_tools + + async def llm_call(self, state: BaseState, runtime: Runtime[BaseContext] = None) -> dict[str, Any]: + """调用 llm 模型 - 异步版本以支持异步工具""" + model = load_chat_model(runtime.context.model) + + # 这里要根据配置动态获取工具 + available_tools = await self._get_invoke_tools(runtime.context.tools, runtime.context.mcps) + logger.info(f"LLM binded ({len(available_tools)}) available_tools: {[tool.name for tool in available_tools]}") + + if available_tools: + model = model.bind_tools(available_tools) + + # 使用异步调用 + response = cast( + AIMessage, + await model.ainvoke([{"role": "system", "content": runtime.context.system_prompt}, *state.messages]), + ) + return {"messages": [response]} + + async def dynamic_tools_node(self, state: BaseState, runtime: Runtime[BaseContext]) -> dict[str, list[ToolMessage]]: + """Execute tools dynamically based on configuration. + + This function gets the available tools based on the current configuration + and executes the requested tool calls from the last message. + """ + # Get available tools based on configuration + available_tools = await self._get_invoke_tools(runtime.context.tools, runtime.context.mcps) + + # Create a ToolNode with the available tools + tool_node = ToolNode(available_tools) + + # Execute the tool node + result = await tool_node.ainvoke(state) + + return cast(dict[str, list[ToolMessage]], result) \ No newline at end of file diff --git a/src/agents/common/tools.py b/src/agents/common/tools.py index 7d9fa2aa..4a681687 100644 --- a/src/agents/common/tools.py +++ b/src/agents/common/tools.py @@ -5,11 +5,54 @@ from typing import Annotated, Any from langchain.tools import tool from langchain_core.tools import StructuredTool from langchain_tavily import TavilySearch +from langgraph.types import interrupt from pydantic import BaseModel, Field from src import config, graph_base, knowledge_base from src.utils import logger +# TODO[修改建议]:前端需要通过interrupt进行交互,点击是或否来批准执行 +# 返回中断点: +# is_approved : bool = True 或者 False +# resume_command = Command(resume=is_approved) +# stream = graph.stream(resume_command, config=config, stream_mode="messages") +# graph.invoke(resume_command, config=config) +@tool(name_or_callable="人工审批工具", description="请求人工审批工具,用于在执行重要操作前获得人类确认。") +def get_approved_user_goal( + operation_description: str, +)->dict: + """ + 请求人工审批,在执行重要操作前获得人类确认。 + + Args: + operation_description: 需要审批的操作描述,例如 "调用知识库工具" + Returns: + dict: 包含审批结果的字典,格式为 {"approved": bool, "message": str} + """ + # 构建详细的中断信息 + interrupt_info = { + "question": f"是否批准以下操作?", + "operation": operation_description, + } + + # 触发人工审批 + is_approved = interrupt(interrupt_info) + + # 返回审批结果 + if is_approved: + result = { + "approved": True, + "message": f"✅ 操作已批准:{operation_description}", + } + print(f"✅ 人工审批通过: {operation_description}") + else: + result = { + "approved": False, + "message": f"❌ 操作被拒绝:{operation_description}", + } + print(f"❌ 人工审批被拒绝: {operation_description}") + + return result @tool(name_or_callable="查询知识图谱", description="使用这个工具可以查询知识图谱中包含的三元组信息。") def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge graph."]) -> Any: @@ -31,6 +74,7 @@ def get_static_tools() -> list: """注册静态工具""" static_tools = [ query_knowledge_graph, + get_approved_user_goal ] # 检查是否启用网页搜索 diff --git a/src/agents/multiAgent/__init__.py b/src/agents/multiAgent/__init__.py new file mode 100644 index 00000000..740c89a0 --- /dev/null +++ b/src/agents/multiAgent/__init__.py @@ -0,0 +1,3 @@ +from .graph import SampleMultiAgent + +__all__ = ["SampleMultiAgent"] diff --git a/src/agents/multiAgent/context.py b/src/agents/multiAgent/context.py new file mode 100644 index 00000000..7bb1adad --- /dev/null +++ b/src/agents/multiAgent/context.py @@ -0,0 +1,25 @@ +from dataclasses import dataclass, field +from typing import Annotated + +from src.agents.common.context import BaseContext +from src.agents.common.mcp import MCP_SERVERS +from src.agents.common.tools import gen_tool_info + +from .tools import get_tools + + +@dataclass(kw_only=True) +class Context(BaseContext): + tools: Annotated[list[dict], {"__template_metadata__": {"kind": "tools"}}] = field( + default_factory=list, + metadata={ + "name": "工具", + "options": gen_tool_info(get_tools()), # 这里的选择是所有的工具 + "description": "工具列表", + }, + ) + + mcps: list[str] = field( + default_factory=list, + metadata={"name": "MCP服务器", "options": list(MCP_SERVERS.keys()), "description": "MCP服务器列表"}, + ) diff --git a/src/agents/multiAgent/graph.py b/src/agents/multiAgent/graph.py new file mode 100644 index 00000000..0d6d7679 --- /dev/null +++ b/src/agents/multiAgent/graph.py @@ -0,0 +1,61 @@ +from langgraph.graph import END, START, StateGraph +from langgraph.prebuilt import tools_condition + +from src.agents.common.toolagent import ToolAgent + +from .context import Context +from .state import State +from .tools import get_tools + + +class SampleMultiAgent(ToolAgent): + name = "MultiAgent智能体" + description = "Supervisor智能体,具有调用其他子智能体的能力(在工具中添加)" + + # TODO[已完成]: 通过将其他agent封装为工具的方式添加了多智能体调度 + ''' + 你是一个多智能体核心,通过多智能体调用的方式帮助用户完成一系列任务: + + 1.当你需要知识库问答功能时,请调用对话聊天智能体实现 + 2.当你需要加密计算的时候,请调用加密计算智能体实现 + ''' + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.graph = None + self.checkpointer = None + self.context_schema = Context + self.agent_tools = None + + def get_tools(self): + return get_tools() + + async def get_graph(self, **kwargs): + """构建图""" + if self.graph: + return self.graph + + builder = StateGraph(State, context_schema=self.context_schema) + builder.add_node("chatbot", self.llm_call) + builder.add_node("tools", self.dynamic_tools_node) + builder.add_edge(START, "chatbot") + builder.add_conditional_edges( + "chatbot", + tools_condition, + ) + builder.add_edge("tools", "chatbot") + builder.add_edge("chatbot", END) + + self.checkpointer = await self._get_checkpointer() + graph = builder.compile(checkpointer=self.checkpointer, name=self.name) + self.graph = graph + return graph + + +def main(): + pass + + +if __name__ == "__main__": + main() + # asyncio.run(main()) diff --git a/src/agents/multiAgent/state.py b/src/agents/multiAgent/state.py new file mode 100644 index 00000000..f46bfb9c --- /dev/null +++ b/src/agents/multiAgent/state.py @@ -0,0 +1,20 @@ +"""Define the state structures for the agent.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass, field +from typing import Annotated + +from langchain.messages import AnyMessage +from langgraph.graph import add_messages + + +@dataclass +class State: + """Defines the input state for the agent, representing a narrower interface to the outside world. + + This class is used to define the initial state and structure of incoming data. + """ + + messages: Annotated[Sequence[AnyMessage], add_messages] = field(default_factory=list) diff --git a/src/agents/multiAgent/tools.py b/src/agents/multiAgent/tools.py new file mode 100644 index 00000000..246d8ddd --- /dev/null +++ b/src/agents/multiAgent/tools.py @@ -0,0 +1,63 @@ +import os +from typing import Any + +from langchain.tools import tool + +from src.agents import agent_manager +from src.agents.common.toolkits.mysql import get_mysql_tools +from src.agents.common.tools import get_buildin_tools +from src.utils import logger + +# TODO[修改建议]:能不能通过前端直接指定子智能体? +# 调用子智能体后的日志是输出到tool_calls的 +@tool(name_or_callable="对话聊天智能体", description="调用指定智能体进行对话聊天的功能") +async def call_chatbot(query: str) -> str: + """ + 调用指定chatbot智能体进行对话聊天的功能 + + Args: + query: 根据需要构造的提问 + Returns: + str: 最终的回答结果 + """ + try: + input = [{"role": "user", "content": query}] + chatbot = agent_manager.get_agent("ChatbotAgent") + message = await chatbot.invoke_messages(input) + # 直接获取最后一个消息的内容 + final_answer = message.get('messages', [])[-1].content + logger.info(f"ChatbotAgent: {final_answer}") + return final_answer + except Exception as e: + logger.error(f"CallAgent error: {e}") + raise + +@tool(name_or_callable="加密计算智能体", description="调用指定智能体进行加密计算的功能") +async def call_react_agent(query: str) -> str: + """ + 调用指定智能体进行加密计算的功能 + + Args: + query: 根据需要构造的提问 + Returns: + str: 最终的回答结果 + """ + try: + input = [{"role": "user", "content": query}] + chatbot = agent_manager.get_agent("ReActAgent") + message = await chatbot.invoke_messages(input) + # 直接获取最后一个消息的内容 + final_answer = message.get('messages', [])[-1].content + logger.info(f"ReActAgent: {final_answer}") + return final_answer + except Exception as e: + logger.error(f"CallAgent error: {e}") + raise + + +def get_tools() -> list[Any]: + """获取所有可运行的工具(给大模型使用)""" + tools = get_buildin_tools() + tools.append(call_chatbot) + tools.append(call_react_agent) + return tools diff --git a/src/agents/react/context.py b/src/agents/react/context.py new file mode 100644 index 00000000..7bb1adad --- /dev/null +++ b/src/agents/react/context.py @@ -0,0 +1,25 @@ +from dataclasses import dataclass, field +from typing import Annotated + +from src.agents.common.context import BaseContext +from src.agents.common.mcp import MCP_SERVERS +from src.agents.common.tools import gen_tool_info + +from .tools import get_tools + + +@dataclass(kw_only=True) +class Context(BaseContext): + tools: Annotated[list[dict], {"__template_metadata__": {"kind": "tools"}}] = field( + default_factory=list, + metadata={ + "name": "工具", + "options": gen_tool_info(get_tools()), # 这里的选择是所有的工具 + "description": "工具列表", + }, + ) + + mcps: list[str] = field( + default_factory=list, + metadata={"name": "MCP服务器", "options": list(MCP_SERVERS.keys()), "description": "MCP服务器列表"}, + ) diff --git a/src/agents/react/graph.py b/src/agents/react/graph.py index eb79c87a..b4abb844 100644 --- a/src/agents/react/graph.py +++ b/src/agents/react/graph.py @@ -1,51 +1,68 @@ -from pathlib import Path +from langgraph.constants import START, END +from langgraph.graph import StateGraph -from langchain.agents import create_agent -from langchain.agents.middleware import ModelRequest, ModelResponse, dynamic_prompt, wrap_model_call +from src.agents.common.toolagent import ToolAgent -from src.agents.common.base import BaseAgent -from src.agents.common.models import load_chat_model -from src.agents.common.tools import get_buildin_tools -from src.utils import logger +from .state import State +from .context import Context +from .tools import get_tools + +def tools_branch_continue(state: State): + messages = state.messages + last_message = messages[-1] + # 如果无工具调用则结束 + if not last_message.tool_calls: + return "end" + else: + return "continue" -@dynamic_prompt -def context_aware_prompt(request: ModelRequest) -> str: - runtime = request.runtime - return runtime.context.system_prompt - - -@wrap_model_call -async def context_based_model(request: ModelRequest, handler) -> ModelResponse: - # 从 runtime context 读取配置 - model_spec = request.runtime.context.model - model = load_chat_model(model_spec) - - request = request.override(model=model) - return await handler(request) - - -class ReActAgent(BaseAgent): +class ReActAgent(ToolAgent): name = "智能体 Demo" description = "A react agent that can answer questions and help with tasks." + # TODO:[已完成] React智能体 + ''' + 提示词示例: + 你是一个智能体助手 + + 你的所有需要计算的工作不能自己计算,只能通过调用加密计算机工具进行,你只能通过调用该工具的方式进行计算才能得到正确答案。计算顺序是正常的。 + + 例如 2+3+2,你必须调用两次加密计算机工具进行运算,最终结果为9。 + ''' + def __init__(self, **kwargs): super().__init__(**kwargs) + self.graph = None + self.checkpointer = None + self.context_schema = Context + self.agent_tools = None def get_tools(self): - return get_buildin_tools() + return get_tools() async def get_graph(self, **kwargs): + # 创建 ReActAgent + """构建图""" if self.graph: return self.graph - # 创建 ReActAgent - graph = create_agent( - model=load_chat_model("siliconflow/Qwen/Qwen3-235B-A22B-Instruct-2507"), # 实际会被覆盖 - tools=self.get_tools(), - middleware=[context_aware_prompt, context_based_model], - checkpointer=await self._get_checkpointer(), + builder = StateGraph(State, context_schema=self.context_schema) + builder.add_node("agent", self.llm_call) + builder.add_node("tools", self.dynamic_tools_node) + builder.set_entry_point("agent") + # 添加条件边:agent 决定是否调用工具继续还是结束对话 + builder.add_conditional_edges( + "agent", + tools_branch_continue, + { + "continue": "tools", # 调用工具 + "end": END, # 结束对话 + }, ) - + builder.add_edge("tools", "agent") + self.checkpointer = await self._get_checkpointer() + graph = builder.compile(checkpointer=self.checkpointer, name=self.name) self.graph = graph return graph + diff --git a/src/agents/react/state.py b/src/agents/react/state.py new file mode 100644 index 00000000..927bfa95 --- /dev/null +++ b/src/agents/react/state.py @@ -0,0 +1,22 @@ +"""Define the state structures for the agent.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass, field +from typing import Annotated + +from langchain.messages import AnyMessage +from langgraph.graph import add_messages + +from src.agents.common.state import BaseState + + +@dataclass +class State(BaseState): + """Defines the input state for the agent, representing a narrower interface to the outside world. + + This class is used to define the initial state and structure of incoming data. + """ + + messages: Annotated[Sequence[AnyMessage], add_messages] = field(default_factory=list) diff --git a/src/agents/react/tools.py b/src/agents/react/tools.py new file mode 100644 index 00000000..f3ec28cc --- /dev/null +++ b/src/agents/react/tools.py @@ -0,0 +1,48 @@ +import os +from typing import Any + +import requests +from langchain.tools import tool + +from src.agents.common.toolkits.mysql import get_mysql_tools +from src.agents.common.tools import get_buildin_tools +from src.storage.minio import upload_image_to_minio +from src.utils import logger + +@tool(name_or_callable="加密计算器", description="可以对给定的2个数字选择进行加减乘除四种加密计算") +def calculator(a: float, b: float, operation: str) -> float: + """ + 可以对给定的2个数字选择进行加减乘除四种加密计算 + + Args: + a: 第一个数字 + b: 第二个数字 + operation: 计算操作符号,可以是add,subtract,multiply,divide + + Returns: + float: 最终的计算结果 + """ + try: + if operation == "add": + return a + b + 1 + elif operation == "subtract": + return a - b - 1 + elif operation == "multiply": + return a * b * 2 + elif operation == "divide": + if b == 0: + raise ZeroDivisionError("除数不能为零") + return a / b - 1 + else: + raise ValueError(f"不支持的运算类型: {operation},仅支持 add, subtract, multiply, divide") + except Exception as e: + logger.error(f"Calculator error: {e}") + raise + + +def get_tools() -> list[Any]: + """获取所有可运行的工具(给大模型使用)""" + tools = get_buildin_tools() + tools.append(calculator) + tools.extend(get_mysql_tools()) + return tools From ad8e26c6548573b6dc923d16686f1f21373aeb90 Mon Sep 17 00:00:00 2001 From: miluELK <2636626273@qq.com> Date: Tue, 28 Oct 2025 08:22:52 +0800 Subject: [PATCH 2/4] =?UTF-8?q?feat:=20=E4=B8=BAagent=E5=B0=81=E8=A3=85?= =?UTF-8?q?=E4=B8=BAtool=E8=B0=83=E7=94=A8=E6=B7=BB=E5=8A=A0=E4=BA=86runti?= =?UTF-8?q?meconfig=EF=BC=8C=E5=8F=AF=E4=BB=A5=E4=BC=A0=E9=80=92=E8=AE=B0?= =?UTF-8?q?=E5=BF=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/agents/multiAgent/state.py | 4 +++- src/agents/multiAgent/tools.py | 21 +++++++++++++++++---- 2 files changed, 20 insertions(+), 5 deletions(-) diff --git a/src/agents/multiAgent/state.py b/src/agents/multiAgent/state.py index f46bfb9c..927bfa95 100644 --- a/src/agents/multiAgent/state.py +++ b/src/agents/multiAgent/state.py @@ -9,9 +9,11 @@ from typing import Annotated from langchain.messages import AnyMessage from langgraph.graph import add_messages +from src.agents.common.state import BaseState + @dataclass -class State: +class State(BaseState): """Defines the input state for the agent, representing a narrower interface to the outside world. This class is used to define the initial state and structure of incoming data. diff --git a/src/agents/multiAgent/tools.py b/src/agents/multiAgent/tools.py index 246d8ddd..19c4e417 100644 --- a/src/agents/multiAgent/tools.py +++ b/src/agents/multiAgent/tools.py @@ -2,6 +2,7 @@ import os from typing import Any from langchain.tools import tool +from langchain_core.runnables import RunnableConfig from src.agents import agent_manager from src.agents.common.toolkits.mysql import get_mysql_tools @@ -11,19 +12,25 @@ from src.utils import logger # TODO[修改建议]:能不能通过前端直接指定子智能体? # 调用子智能体后的日志是输出到tool_calls的 @tool(name_or_callable="对话聊天智能体", description="调用指定智能体进行对话聊天的功能") -async def call_chatbot(query: str) -> str: +async def call_chatbot(query: str, config: RunnableConfig) -> str: """ 调用指定chatbot智能体进行对话聊天的功能 Args: query: 根据需要构造的提问 + config: LangGraph运行时配置(自动注入) Returns: str: 最终的回答结果 """ try: input = [{"role": "user", "content": query}] chatbot = agent_manager.get_agent("ChatbotAgent") - message = await chatbot.invoke_messages(input) + configurable = config.get("configurable",{}) + input_context = { + "thread_id":configurable.get("thread_id"), + "user_id": configurable.get("user_id"), + } + message = await chatbot.invoke_messages(input,input_context=input_context) # 直接获取最后一个消息的内容 final_answer = message.get('messages', [])[-1].content logger.info(f"ChatbotAgent: {final_answer}") @@ -33,19 +40,25 @@ async def call_chatbot(query: str) -> str: raise @tool(name_or_callable="加密计算智能体", description="调用指定智能体进行加密计算的功能") -async def call_react_agent(query: str) -> str: +async def call_react_agent(query: str, config: RunnableConfig) -> str: """ 调用指定智能体进行加密计算的功能 Args: query: 根据需要构造的提问 + config: LangGraph运行时配置(自动注入) Returns: str: 最终的回答结果 """ try: input = [{"role": "user", "content": query}] chatbot = agent_manager.get_agent("ReActAgent") - message = await chatbot.invoke_messages(input) + configurable = config.get("configurable",{}) + input_context = { + "thread_id":configurable.get("thread_id"), + "user_id": configurable.get("user_id"), + } + message = await chatbot.invoke_messages(input,input_context=input_context) # 直接获取最后一个消息的内容 final_answer = message.get('messages', [])[-1].content logger.info(f"ReActAgent: {final_answer}") From 127c73e180dbd6a20b87eb3bf4b0ab901b7fc8c8 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Fri, 31 Oct 2025 14:16:07 +0800 Subject: [PATCH 3/4] =?UTF-8?q?refactor:=20=E6=99=BA=E8=83=BD=E4=BD=93?= =?UTF-8?q?=E9=87=8D=E5=91=BD=E5=90=8D=E5=B9=B6=E6=B7=BB=E5=8A=A0=E6=A0=BC?= =?UTF-8?q?=E5=BC=8F=E6=A3=80=E6=9F=A5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/agents/common/toolagent.py | 2 +- src/agents/common/tools.py | 2 +- src/agents/mini_agent/__init__.py | 3 ++ src/agents/mini_agent/graph.py | 49 +++++++++++++++++++ .../{multiAgent => multi_agent}/__init__.py | 0 .../{multiAgent => multi_agent}/context.py | 0 .../{multiAgent => multi_agent}/graph.py | 0 .../{multiAgent => multi_agent}/state.py | 0 .../{multiAgent => multi_agent}/tools.py | 2 - src/agents/react/graph.py | 7 ++- src/agents/react/tools.py | 3 -- src/agents/reporter/graph.py | 3 +- src/knowledge/base.py | 1 - 13 files changed, 58 insertions(+), 14 deletions(-) create mode 100644 src/agents/mini_agent/__init__.py create mode 100644 src/agents/mini_agent/graph.py rename src/agents/{multiAgent => multi_agent}/__init__.py (100%) rename src/agents/{multiAgent => multi_agent}/context.py (100%) rename src/agents/{multiAgent => multi_agent}/graph.py (100%) rename src/agents/{multiAgent => multi_agent}/state.py (100%) rename src/agents/{multiAgent => multi_agent}/tools.py (97%) diff --git a/src/agents/common/toolagent.py b/src/agents/common/toolagent.py index 38e6dfe4..c43c988c 100644 --- a/src/agents/common/toolagent.py +++ b/src/agents/common/toolagent.py @@ -84,4 +84,4 @@ class ToolAgent(BaseAgent): # Execute the tool node result = await tool_node.ainvoke(state) - return cast(dict[str, list[ToolMessage]], result) \ No newline at end of file + return cast(dict[str, list[ToolMessage]], result) diff --git a/src/agents/common/tools.py b/src/agents/common/tools.py index 4a681687..81ede3e7 100644 --- a/src/agents/common/tools.py +++ b/src/agents/common/tools.py @@ -31,7 +31,7 @@ def get_approved_user_goal( """ # 构建详细的中断信息 interrupt_info = { - "question": f"是否批准以下操作?", + "question": "是否批准以下操作?", "operation": operation_description, } diff --git a/src/agents/mini_agent/__init__.py b/src/agents/mini_agent/__init__.py new file mode 100644 index 00000000..b3f7b398 --- /dev/null +++ b/src/agents/mini_agent/__init__.py @@ -0,0 +1,3 @@ +from .graph import MiniAgent + +__all__ = ["MiniAgent"] diff --git a/src/agents/mini_agent/graph.py b/src/agents/mini_agent/graph.py new file mode 100644 index 00000000..68f2f02b --- /dev/null +++ b/src/agents/mini_agent/graph.py @@ -0,0 +1,49 @@ + +from langchain.agents import create_agent +from langchain.agents.middleware import ModelRequest, ModelResponse, dynamic_prompt, wrap_model_call + +from src.agents.common.base import BaseAgent +from src.agents.common.models import load_chat_model +from src.agents.common.tools import get_buildin_tools + + +@dynamic_prompt +def context_aware_prompt(request: ModelRequest) -> str: + runtime = request.runtime + return runtime.context.system_prompt + + +@wrap_model_call +async def context_based_model(request: ModelRequest, handler) -> ModelResponse: + # 从 runtime context 读取配置 + model_spec = request.runtime.context.model + model = load_chat_model(model_spec) + + request = request.override(model=model) + return await handler(request) + + +class MiniAgent(BaseAgent): + name = "智能体 Demo" + description = "一个基于内置工具的智能体示例" + + def __init__(self, **kwargs): + super().__init__(**kwargs) + + def get_tools(self): + return get_buildin_tools() + + async def get_graph(self, **kwargs): + if self.graph: + return self.graph + + # 创建 MiniAgent + graph = create_agent( + model=load_chat_model("siliconflow/Qwen/Qwen3-235B-A22B-Instruct-2507"), # 实际会被覆盖 + tools=self.get_tools(), + middleware=[context_aware_prompt, context_based_model], + checkpointer=await self._get_checkpointer(), + ) + + self.graph = graph + return graph diff --git a/src/agents/multiAgent/__init__.py b/src/agents/multi_agent/__init__.py similarity index 100% rename from src/agents/multiAgent/__init__.py rename to src/agents/multi_agent/__init__.py diff --git a/src/agents/multiAgent/context.py b/src/agents/multi_agent/context.py similarity index 100% rename from src/agents/multiAgent/context.py rename to src/agents/multi_agent/context.py diff --git a/src/agents/multiAgent/graph.py b/src/agents/multi_agent/graph.py similarity index 100% rename from src/agents/multiAgent/graph.py rename to src/agents/multi_agent/graph.py diff --git a/src/agents/multiAgent/state.py b/src/agents/multi_agent/state.py similarity index 100% rename from src/agents/multiAgent/state.py rename to src/agents/multi_agent/state.py diff --git a/src/agents/multiAgent/tools.py b/src/agents/multi_agent/tools.py similarity index 97% rename from src/agents/multiAgent/tools.py rename to src/agents/multi_agent/tools.py index 19c4e417..f5dcef34 100644 --- a/src/agents/multiAgent/tools.py +++ b/src/agents/multi_agent/tools.py @@ -1,11 +1,9 @@ -import os from typing import Any from langchain.tools import tool from langchain_core.runnables import RunnableConfig from src.agents import agent_manager -from src.agents.common.toolkits.mysql import get_mysql_tools from src.agents.common.tools import get_buildin_tools from src.utils import logger diff --git a/src/agents/react/graph.py b/src/agents/react/graph.py index b4abb844..f311cdbd 100644 --- a/src/agents/react/graph.py +++ b/src/agents/react/graph.py @@ -1,4 +1,4 @@ -from langgraph.constants import START, END +from langgraph.constants import END from langgraph.graph import StateGraph from src.agents.common.toolagent import ToolAgent @@ -18,10 +18,9 @@ def tools_branch_continue(state: State): class ReActAgent(ToolAgent): - name = "智能体 Demo" - description = "A react agent that can answer questions and help with tasks." + name = "ReActAgent" + description = "符合 ReAct 范式的智能体,可以通过调用工具来完成复杂任务。" - # TODO:[已完成] React智能体 ''' 提示词示例: 你是一个智能体助手 diff --git a/src/agents/react/tools.py b/src/agents/react/tools.py index f3ec28cc..5dd1c766 100644 --- a/src/agents/react/tools.py +++ b/src/agents/react/tools.py @@ -1,12 +1,9 @@ -import os from typing import Any -import requests from langchain.tools import tool from src.agents.common.toolkits.mysql import get_mysql_tools from src.agents.common.tools import get_buildin_tools -from src.storage.minio import upload_image_to_minio from src.utils import logger @tool(name_or_callable="加密计算器", description="可以对给定的2个数字选择进行加减乘除四种加密计算") diff --git a/src/agents/reporter/graph.py b/src/agents/reporter/graph.py index 8402722a..351e17ba 100644 --- a/src/agents/reporter/graph.py +++ b/src/agents/reporter/graph.py @@ -1,5 +1,4 @@ import textwrap -from pathlib import Path from langchain.agents import create_agent from langchain.agents.middleware import ModelRequest, ModelResponse, dynamic_prompt, wrap_model_call @@ -64,4 +63,4 @@ class SqlReporterAgent(BaseAgent): self.graph = graph logger.info("SqlReporterAgent 构建成功") - return graph \ No newline at end of file + return graph diff --git a/src/knowledge/base.py b/src/knowledge/base.py index 2eb643a4..329cb964 100644 --- a/src/knowledge/base.py +++ b/src/knowledge/base.py @@ -136,7 +136,6 @@ class KnowledgeBase(ABC): """ from src.utils import hashstr - from src.utils import hashstr # 从 kwargs 中获取 is_private 配置 is_private = kwargs.get('is_private', False) From 181341db48a85cd5b3db357245ca4e51aeebf538 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Sat, 1 Nov 2025 21:34:16 +0800 Subject: [PATCH 4/4] =?UTF-8?q?feat:=20=E4=B8=BA=E6=99=BA=E8=83=BD?= =?UTF-8?q?=E4=BD=93=E6=93=8D=E4=BD=9C=E5=AE=9E=E7=8E=B0=E4=BA=BA=E5=B7=A5?= =?UTF-8?q?=E5=AE=A1=E6=89=B9=E6=9C=BA=E5=88=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 HumanApprovalModal 组件,用于处理用户对关键操作的审批。 - 引入 useApproval 可组合项,用于管理审批状态和逻辑。 - 更新 AgentChatComponent 以显示审批模态框并处理审批操作。 - 增强消息处理功能,支持工具调用合并并改进对 AI 消息块的处理。 - 重构各种组件和 API,以整合新的审批流程,确保代理交互期间的流畅用户体验。 --- CLAUDE.md | 1 + scripts/batch_upload.py | 23 +- server/routers/chat_router.py | 418 ++++++++++++++----- src/agents/common/mcp.py | 1 + src/agents/common/toolagent.py | 5 +- src/agents/common/tools.py | 9 +- src/agents/mini_agent/graph.py | 1 - src/agents/multi_agent/graph.py | 4 +- src/agents/multi_agent/tools.py | 18 +- src/agents/react/graph.py | 8 +- src/agents/react/tools.py | 1 + src/agents/reporter/graph.py | 3 +- src/knowledge/base.py | 3 +- web/src/apis/agent_api.js | 28 +- web/src/components/AgentChatComponent.vue | 209 +++++++++- web/src/components/AgentMessageComponent.vue | 15 + web/src/components/HumanApprovalModal.vue | 238 +++++++++++ web/src/composables/useApproval.js | 128 ++++++ web/src/utils/messageProcessor.js | 49 ++- 19 files changed, 989 insertions(+), 173 deletions(-) create mode 100644 CLAUDE.md create mode 100644 web/src/components/HumanApprovalModal.vue create mode 100644 web/src/composables/useApproval.js diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 00000000..460e1518 --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1 @@ +See AGENTS.md \ No newline at end of file diff --git a/scripts/batch_upload.py b/scripts/batch_upload.py index a369e42a..54d1eabc 100644 --- a/scripts/batch_upload.py +++ b/scripts/batch_upload.py @@ -314,7 +314,10 @@ def upload( directory: pathlib.Path = typer.Option( ..., help="The directory containing files to upload.", exists=True, file_okay=False ), - pattern: list[str] = typer.Option(["*.md"], help="The glob patterns for files to upload (e.g., '*.pdf', '**/*.txt'). Can be specified multiple times."), + pattern: list[str] = typer.Option( + ["*.md"], + help="The glob patterns for files to upload (e.g., '*.pdf', '**/*.txt'). Can be specified multiple times.", + ), base_url: str = typer.Option("http://127.0.0.1:5050/api", help="The base URL of the API server."), username: str = typer.Option(..., help="Admin username for login."), password: str = typer.Option(..., help="Admin password for login."), @@ -354,7 +357,9 @@ def upload( if not all_files: patterns_str = "', '".join(pattern) - console.print(f"[bold yellow]No files found in '{directory}' matching patterns: '{patterns_str}'. Aborting.[/bold yellow]") + console.print( + f"[bold yellow]No files found in '{directory}' matching patterns: '{patterns_str}'. Aborting.[/bold yellow]" + ) raise typer.Exit() # 过滤掉macos的隐藏文件 @@ -398,11 +403,13 @@ def upload( # Split all files into batches for batch_num in range(0, len(files_to_upload), batch_size): - batch_files = files_to_upload[batch_num:batch_num + batch_size] + batch_files = files_to_upload[batch_num : batch_num + batch_size] batch_start = batch_num + 1 batch_end = min(batch_num + batch_size, len(files_to_upload)) - console.print(f"\n[bold yellow]=== Batch {batch_start}-{batch_end} of {len(files_to_upload)} ===[/bold yellow]") + console.print( + f"\n[bold yellow]=== Batch {batch_start}-{batch_end} of {len(files_to_upload)} ===[/bold yellow]" + ) # Step 1: Upload this batch of files sequentially console.print(f"[blue]Step 1: Uploading {len(batch_files)} files...[/blue]") @@ -420,7 +427,9 @@ def upload( console=console, transient=True, ) as progress: - upload_task_id = progress.add_task(f"Uploading batch {batch_start}-{batch_end}...", total=len(batch_files), postfix="") + upload_task_id = progress.add_task( + f"Uploading batch {batch_start}-{batch_end}...", total=len(batch_files), postfix="" + ) for file_path, file_hash in batch_files: server_file_path = await upload_single_file( @@ -458,7 +467,9 @@ def upload( # Step 3: Wait for this batch to complete if wait_for_completion and task_id: - console.print(f"[cyan]Step 3: Waiting for batch {batch_start}-{batch_end} to complete...[/cyan]") + console.print( + f"[cyan]Step 3: Waiting for batch {batch_start}-{batch_end} to complete...[/cyan]" + ) await wait_for_tasks_completion(client, base_url, [task_id], poll_interval) console.print(f"[green]Batch {batch_start}-{batch_end} completed![/green]") else: diff --git a/server/routers/chat_router.py b/server/routers/chat_router.py index f0d93b33..9f6beece 100644 --- a/server/routers/chat_router.py +++ b/server/routers/chat_router.py @@ -8,6 +8,7 @@ from pathlib import Path from fastapi import APIRouter, Body, Depends, HTTPException from fastapi.responses import StreamingResponse from langchain.messages import AIMessageChunk, HumanMessage +from langgraph.types import Command from pydantic import BaseModel from sqlalchemy.orm import Session @@ -81,6 +82,200 @@ async def set_default_agent(request_data: dict = Body(...), current_user=Depends # ============================================================================= +async def _get_langgraph_messages(agent_instance, config_dict): + """获取LangGraph中的消息""" + graph = await agent_instance.get_graph() + state = await graph.aget_state(config_dict) + + if not state or not state.values: + logger.warning("No state found in LangGraph") + return None + + return state.values.get("messages", []) + + +def _get_existing_message_ids(conv_mgr, thread_id): + """获取已保存的消息ID集合""" + existing_messages = conv_mgr.get_messages_by_thread_id(thread_id) + return {msg.extra_metadata["id"] for msg in existing_messages if msg.extra_metadata and "id" in msg.extra_metadata} + + +async def _save_ai_message(conv_mgr, thread_id, msg_dict): + """保存AI消息和相关的工具调用""" + content = msg_dict.get("content", "") + tool_calls_data = msg_dict.get("tool_calls", []) + + # 保存AI消息 + ai_msg = conv_mgr.add_message_by_thread_id( + thread_id=thread_id, + role="assistant", + content=content, + message_type="text", + extra_metadata=msg_dict, + ) + + # 保存工具调用 + if tool_calls_data: + logger.debug(f"Saving {len(tool_calls_data)} tool calls from AI message") + for tc in tool_calls_data: + conv_mgr.add_tool_call( + message_id=ai_msg.id, + tool_name=tc.get("name", "unknown"), + tool_input=tc.get("args", {}), + status="pending", + langgraph_tool_call_id=tc.get("id"), + ) + + logger.debug(f"Saved AI message {ai_msg.id} with {len(tool_calls_data)} tool calls") + + +def _save_tool_message(conv_mgr, msg_dict): + """保存工具执行结果""" + tool_call_id = msg_dict.get("tool_call_id") + content = msg_dict.get("content", "") + name = msg_dict.get("name", "") + + if not tool_call_id: + return + + # 确保tool_output是字符串类型 + if isinstance(content, list): + tool_output = json.dumps(content) if content else "" + else: + tool_output = str(content) + + # 更新工具调用结果 + updated_tc = conv_mgr.update_tool_call_output( + langgraph_tool_call_id=tool_call_id, + tool_output=tool_output, + status="success", + ) + + if updated_tc: + logger.debug(f"Updated tool_call {tool_call_id} ({name}) with output") + else: + logger.warning(f"Tool call {tool_call_id} not found for update") + + +async def save_messages_from_langgraph_state( + agent_instance, + thread_id, + conv_mgr, + config_dict, +): + """ + 从 LangGraph state 中读取完整消息并保存到数据库 + 这样可以获得完整的 tool_calls 参数 + """ + try: + messages = await _get_langgraph_messages(agent_instance, config_dict) + if messages is None: + return + + logger.debug(f"Retrieved {len(messages)} messages from LangGraph state") + existing_ids = _get_existing_message_ids(conv_mgr, thread_id) + + for msg in messages: + msg_dict = msg.model_dump() if hasattr(msg, "model_dump") else {} + msg_type = msg_dict.get("type", "unknown") + + if msg_type == "human" or msg.id in existing_ids: + continue + + if msg_type == "ai": + await _save_ai_message(conv_mgr, thread_id, msg_dict) + elif msg_type == "tool": + _save_tool_message(conv_mgr, msg_dict) + else: + logger.warning(f"Unknown message type: {msg_type}, skipping") + continue + + logger.debug(f"Processed message type={msg_type}") + + logger.info("Saved messages from LangGraph state") + + except Exception as e: + logger.error(f"Error saving messages from LangGraph state: {e}") + logger.error(traceback.format_exc()) + + +async def check_and_handle_interrupts(agent, langgraph_config, make_chunk, meta, thread_id): + """检查并处理 LangGraph 中断状态,发送人工审批请求到前端""" + try: + # 获取 agent 的 graph 对象 + graph = await agent.get_graph() + + # 获取当前状态,检查是否有中断 + state = await graph.aget_state(langgraph_config) + + if not state or not state.values: + logger.debug("No state found when checking for interrupts") + return + + # 检查是否有中断信息 + # LangGraph 中断信息通常在 state.tasks 或 __interrupt__ 字段中 + interrupt_info = None + + # 方法1: 检查 state.tasks 中的中断 + if hasattr(state, "tasks") and state.tasks: + for task in state.tasks: + if hasattr(task, "interrupts") and task.interrupts: + interrupt_info = task.interrupts[0] # 取第一个中断 + break + + # 方法2: 检查 state.values 中的 __interrupt__ 字段 + if not interrupt_info and state.values: + interrupt_data = state.values.get("__interrupt__") + if interrupt_data and isinstance(interrupt_data, list) and len(interrupt_data) > 0: + interrupt_info = interrupt_data[0] + + # 方法3: 检查 state.next 字段,如果指向中断节点 + if not interrupt_info and hasattr(state, "next") and state.next: + # 如果 next 指向某个需要审批的节点,可能需要额外处理 + logger.debug(f"State next nodes: {state.next}") + + if interrupt_info: + logger.info(f"Human approval interrupt detected: {interrupt_info}") + + # 提取中断信息 + question = "是否批准以下操作?" + operation = "需要人工审批的操作" + + if isinstance(interrupt_info, dict): + question = interrupt_info.get("question", question) + operation = interrupt_info.get("operation", operation) + elif isinstance(interrupt_info, (list, tuple)) and len(interrupt_info) > 0: + # 有些情况下中断信息可能是元组形式 + first_interrupt = interrupt_info[0] + if isinstance(first_interrupt, dict): + question = first_interrupt.get("question", question) + operation = first_interrupt.get("operation", operation) + else: + operation = str(first_interrupt) + else: + operation = str(interrupt_info) + + # 发送人工审批请求到前端 + logger.info(f"Sending human approval request - question: {question}, operation: {operation}") + + yield make_chunk( + status="human_approval_required", + thread_id=thread_id, + interrupt_info={"question": question, "operation": operation}, + ) + + else: + logger.debug("No human approval interrupt detected") + + except Exception as e: + logger.error(f"Error checking for interrupts: {e}") + logger.error(traceback.format_exc()) + # 不抛出异常,避免影响主流程 + + +# ============================================================================= + + @chat.post("/call") async def call(query: str = Body(...), meta: dict = Body(None), current_user: User = Depends(get_required_user)): """调用模型进行简单问答(需要登录)""" @@ -147,114 +342,6 @@ async def chat_agent( + b"\n" ) - async def save_messages_from_langgraph_state( - agent_instance, - thread_id, - conv_mgr, - config_dict, - ): - """ - 从 LangGraph state 中读取完整消息并保存到数据库 - 这样可以获得完整的 tool_calls 参数 - """ - try: - graph = await agent_instance.get_graph() - state = await graph.aget_state(config_dict) - - if not state or not state.values: - logger.warning("No state found in LangGraph") - return - - messages = state.values.get("messages", []) - logger.debug(f"Retrieved {len(messages)} messages from LangGraph state") - - # 获取已保存的消息数量,避免重复保存 - existing_messages = conv_mgr.get_messages_by_thread_id(thread_id) - existing_ids = { - msg.extra_metadata["id"] - for msg in existing_messages - if msg.extra_metadata and "id" in msg.extra_metadata - } - - for msg in messages: - msg_dict = msg.model_dump() if hasattr(msg, "model_dump") else {} - msg_type = msg_dict.get("type", "unknown") - - if msg_type == "human" or msg.id in existing_ids: - continue - - elif msg_type == "ai": - # AI 消息 - content = msg_dict.get("content", "") - tool_calls_data = msg_dict.get("tool_calls", []) - - # 格式清洗 - if finish_reason := msg_dict.get("response_metadata", {}).get("finish_reason"): - if "tool_call" in finish_reason and len(finish_reason) > len("tool_call"): - model_name = msg_dict.get("response_metadata", {}).get("model_name", "") - repeat_count = len(finish_reason) // len("tool_call") - msg_dict["response_metadata"]["finish_reason"] = "tool_call" - msg_dict["response_metadata"]["model_name"] = model_name[: len(model_name) // repeat_count] - - # 保存 AI 消息 - ai_msg = conv_mgr.add_message_by_thread_id( - thread_id=thread_id, - role="assistant", - content=content, - message_type="text", - extra_metadata=msg_dict, # 保存原始 model_dump - ) - - # 保存 tool_calls(如果有)- 使用 LangGraph 的 tool_call_id - if tool_calls_data: - logger.debug(f"Saving {len(tool_calls_data)} tool calls from AI message") - for tc in tool_calls_data: - conv_mgr.add_tool_call( - message_id=ai_msg.id, - tool_name=tc.get("name", "unknown"), - tool_input=tc.get("args", {}), # 完整的参数 - status="pending", # 工具还未执行 - langgraph_tool_call_id=tc.get("id"), # 保存 LangGraph tool_call_id - ) - - logger.debug(f"Saved AI message {ai_msg.id} with {len(tool_calls_data)} tool calls") - - elif msg_type == "tool": - # 工具执行结果消息 - 使用 tool_call_id 精确匹配 - tool_call_id = msg_dict.get("tool_call_id") - content = msg_dict.get("content", "") - name = msg_dict.get("name", "") - - if tool_call_id: - # 确保tool_output是字符串类型,避免SQLite不支持列表类型 - if isinstance(content, list): - tool_output = json.dumps(content) if content else "" - else: - tool_output = str(content) - - # 通过 LangGraph tool_call_id 精确匹配并更新 - updated_tc = conv_mgr.update_tool_call_output( - langgraph_tool_call_id=tool_call_id, - tool_output=tool_output, - status="success", - ) - if updated_tc: - logger.debug(f"Updated tool_call {tool_call_id} ({name}) with output") - else: - logger.warning(f"Tool call {tool_call_id} not found for update") - - else: - logger.warning(f"Unknown message type: {msg_type}, skipping") - continue - - logger.debug(f"Processed message type={msg_type}") - - logger.info("Saved messages from LangGraph state") - - except Exception as e: - logger.error(f"Error saving messages from LangGraph state: {e}") - logger.error(traceback.format_exc()) - # TODO:[功能建议]针对需要人工审批后再执行的工具, # 可以使用langgraph的interrupt方法中断对话,等待用户输入后再使用command跳转回去 async def stream_messages(): @@ -323,11 +410,17 @@ async def chat_agent( yield make_chunk(message="检测到敏感内容,已中断输出", status="error") return + # After streaming finished, check for interrupts and save messages + langgraph_config = {"configurable": input_context} + + # Check for human approval interrupts + async for chunk in check_and_handle_interrupts(agent, langgraph_config, make_chunk, meta, thread_id): + yield chunk + meta["time_cost"] = asyncio.get_event_loop().time() - start_time yield make_chunk(status="finished", meta=meta) - # After streaming finished, save all messages from LangGraph state - langgraph_config = {"configurable": input_context} + # Save all messages from LangGraph state await save_messages_from_langgraph_state( agent_instance=agent, thread_id=thread_id, @@ -336,8 +429,17 @@ async def chat_agent( ) except (asyncio.CancelledError, ConnectionError) as e: - # 客户端主动中断连接,尝试保存已生成的部分内容 + # 客户端主动中断连接,检查中断并保存已生成的部分内容 logger.warning(f"Client disconnected, cancelling stream: {e}") + + # 即使在断开连接时也检查中断,确保状态一致性 + langgraph_config = {"configurable": input_context} + try: + async for chunk in check_and_handle_interrupts(agent, langgraph_config, make_chunk, meta, thread_id): + yield chunk + except Exception as interrupt_error: + logger.error(f"Error checking interrupts during disconnect: {interrupt_error}") + if full_msg: # 创建新的 db session,因为原 session 可能已关闭 new_db = db_manager.get_session() @@ -360,6 +462,15 @@ async def chat_agent( except Exception as e: logger.error(f"Error streaming messages: {e}, {traceback.format_exc()}") + + # 即使在异常情况下也检查中断,确保状态一致性 + langgraph_config = {"configurable": input_context} + try: + async for chunk in check_and_handle_interrupts(agent, langgraph_config, make_chunk, meta, thread_id): + yield chunk + except Exception as interrupt_error: + logger.error(f"Error checking interrupts during exception: {interrupt_error}") + if full_msg: # 创建新的 db session,因为原 session 可能已关闭 new_db = db_manager.get_session() @@ -420,6 +531,91 @@ async def get_tools(agent_id: str, current_user: User = Depends(get_required_use return {"tools": {tool["id"]: tool for tool in tools_info}} +@chat.post("/agent/{agent_id}/resume") +async def resume_agent_chat( + agent_id: str, + thread_id: str = Body(...), + approved: bool = Body(...), + current_user: User = Depends(get_required_user), + db: Session = Depends(get_db), +): + """恢复被人工审批中断的对话(需要登录)""" + start_time = asyncio.get_event_loop().time() + logger.info(f"Resuming agent_id: {agent_id}, thread_id: {thread_id}, approved: {approved}") + + meta = { + "agent_id": agent_id, + "thread_id": thread_id, + "user_id": current_user.id, + "approved": approved, + } + if "request_id" not in meta or not meta.get("request_id"): + meta["request_id"] = str(uuid.uuid4()) + + async def stream_resume(): + # 定义resume专用的make_chunk函数,与主聊天端点保持一致 + def make_resume_chunk(content=None, **kwargs): + return ( + json.dumps( + {"request_id": meta.get("request_id"), "response": content, **kwargs}, ensure_ascii=False + ).encode("utf-8") + + b"\n" + ) + + try: + agent = agent_manager.get_agent(agent_id) + except Exception as e: + logger.error(f"Error getting agent {agent_id}: {e}, {traceback.format_exc()}") + yield ( + f'{{"request_id": "{meta.get("request_id")}", "message": ' + f'"Error getting agent {agent_id}: {e}", "status": "error"}}\n' + ) + return + + # 发送init状态块,与主聊天端点保持一致 + init_msg = {"type": "system", "content": f"Resume with approved: {approved}"} + yield make_resume_chunk(status="init", meta=meta, msg=init_msg) + + # 使用 Command(resume=approved) 恢复执行 + resume_command = Command(resume=approved) + graph = await agent.get_graph() + + # 加载 context(包含 tools, model 等配置) + input_context = {"user_id": str(current_user.id), "thread_id": thread_id} + context = agent.context_schema.from_file(module_name=agent.module_name, input_context=input_context) + logger.debug(f"Resume with context: {context}") + + # 创建流式数据源 + stream_source = graph.astream( + resume_command, context=context, config={"configurable": input_context}, stream_mode="messages" + ) + + async for msg, metadata in stream_source: + # 确保msg有正确的ID结构 + msg_dict = msg.model_dump() + if "id" not in msg_dict: + msg_dict["id"] = str(uuid.uuid4()) + + yield make_resume_chunk( + content=getattr(msg, "content", ""), msg=msg_dict, metadata=metadata, status="loading" + ) + + meta["time_cost"] = asyncio.get_event_loop().time() - start_time + yield make_resume_chunk(status="finished", meta=meta) + + # 保存消息到数据库 + langgraph_config = {"configurable": input_context} + conv_manager = ConversationManager(db) + await save_messages_from_langgraph_state( + agent_instance=agent, + thread_id=thread_id, + conv_mgr=conv_manager, + config_dict=langgraph_config, + ) + + return StreamingResponse(stream_resume(), media_type="application/json") + + @chat.post("/agent/{agent_id}/config") async def save_agent_config(agent_id: str, config: dict = Body(...), current_user: User = Depends(get_required_user)): """保存智能体配置到YAML文件(需要登录)""" diff --git a/src/agents/common/mcp.py b/src/agents/common/mcp.py index 1fa36171..e470b64e 100644 --- a/src/agents/common/mcp.py +++ b/src/agents/common/mcp.py @@ -88,6 +88,7 @@ async def get_mcp_tools(server_name: str, additional_servers: dict[str, dict] = logger.error(f"Failed to load tools from MCP server '{server_name}': {e}") return [] + async def get_all_mcp_tools() -> list[Callable[..., Any]]: """Get all tools from all configured MCP servers.""" all_tools = [] diff --git a/src/agents/common/toolagent.py b/src/agents/common/toolagent.py index c43c988c..de3ec218 100644 --- a/src/agents/common/toolagent.py +++ b/src/agents/common/toolagent.py @@ -10,8 +10,9 @@ from src.agents.common.mcp import get_mcp_tools from src.agents.common.models import load_chat_model from src.utils import logger -from .state import BaseState from .context import BaseContext +from .state import BaseState + class ToolAgent(BaseAgent): name = "ToolAgent" @@ -24,7 +25,6 @@ class ToolAgent(BaseAgent): self.context_schema = BaseContext self.agent_tools = None - # TODO:[修改建议] _get_invoke_tools,llm_call,dynamic_tools_node这类针对工具调用的功能大多数Agent都能用得到 # 可以通过一个ToolAgent类继承BaseAgent,通过重写抽象方法获取tools,通过继承BaseState和BaseContext获取配置 # 必要时可通过重写以下方法实现其他逻辑 @@ -33,7 +33,6 @@ class ToolAgent(BaseAgent): logger.error(f"get_tools() is not implemented in {self.__class__.__name__}") return [] - async def _get_invoke_tools(self, selected_tools: list[str], selected_mcps: list[str]): """根据配置获取工具。 默认不使用任何工具。 diff --git a/src/agents/common/tools.py b/src/agents/common/tools.py index 81ede3e7..c569e96b 100644 --- a/src/agents/common/tools.py +++ b/src/agents/common/tools.py @@ -11,6 +11,7 @@ from pydantic import BaseModel, Field from src import config, graph_base, knowledge_base from src.utils import logger + # TODO[修改建议]:前端需要通过interrupt进行交互,点击是或否来批准执行 # 返回中断点: # is_approved : bool = True 或者 False @@ -20,7 +21,7 @@ from src.utils import logger @tool(name_or_callable="人工审批工具", description="请求人工审批工具,用于在执行重要操作前获得人类确认。") def get_approved_user_goal( operation_description: str, -)->dict: +) -> dict: """ 请求人工审批,在执行重要操作前获得人类确认。 @@ -54,6 +55,7 @@ def get_approved_user_goal( return result + @tool(name_or_callable="查询知识图谱", description="使用这个工具可以查询知识图谱中包含的三元组信息。") def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge graph."]) -> Any: """Use this to query knowledge graph, which include some food domain knowledge.""" @@ -72,10 +74,7 @@ def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge def get_static_tools() -> list: """注册静态工具""" - static_tools = [ - query_knowledge_graph, - get_approved_user_goal - ] + static_tools = [query_knowledge_graph, get_approved_user_goal] # 检查是否启用网页搜索 if config.enable_web_search: diff --git a/src/agents/mini_agent/graph.py b/src/agents/mini_agent/graph.py index 68f2f02b..e11f7751 100644 --- a/src/agents/mini_agent/graph.py +++ b/src/agents/mini_agent/graph.py @@ -1,4 +1,3 @@ - from langchain.agents import create_agent from langchain.agents.middleware import ModelRequest, ModelResponse, dynamic_prompt, wrap_model_call diff --git a/src/agents/multi_agent/graph.py b/src/agents/multi_agent/graph.py index 0d6d7679..8db3db0c 100644 --- a/src/agents/multi_agent/graph.py +++ b/src/agents/multi_agent/graph.py @@ -13,12 +13,12 @@ class SampleMultiAgent(ToolAgent): description = "Supervisor智能体,具有调用其他子智能体的能力(在工具中添加)" # TODO[已完成]: 通过将其他agent封装为工具的方式添加了多智能体调度 - ''' + """ 你是一个多智能体核心,通过多智能体调用的方式帮助用户完成一系列任务: 1.当你需要知识库问答功能时,请调用对话聊天智能体实现 2.当你需要加密计算的时候,请调用加密计算智能体实现 - ''' + """ def __init__(self, **kwargs): super().__init__(**kwargs) diff --git a/src/agents/multi_agent/tools.py b/src/agents/multi_agent/tools.py index f5dcef34..fc3fe951 100644 --- a/src/agents/multi_agent/tools.py +++ b/src/agents/multi_agent/tools.py @@ -7,6 +7,7 @@ from src.agents import agent_manager from src.agents.common.tools import get_buildin_tools from src.utils import logger + # TODO[修改建议]:能不能通过前端直接指定子智能体? # 调用子智能体后的日志是输出到tool_calls的 @tool(name_or_callable="对话聊天智能体", description="调用指定智能体进行对话聊天的功能") @@ -23,20 +24,21 @@ async def call_chatbot(query: str, config: RunnableConfig) -> str: try: input = [{"role": "user", "content": query}] chatbot = agent_manager.get_agent("ChatbotAgent") - configurable = config.get("configurable",{}) + configurable = config.get("configurable", {}) input_context = { - "thread_id":configurable.get("thread_id"), + "thread_id": configurable.get("thread_id"), "user_id": configurable.get("user_id"), } - message = await chatbot.invoke_messages(input,input_context=input_context) + message = await chatbot.invoke_messages(input, input_context=input_context) # 直接获取最后一个消息的内容 - final_answer = message.get('messages', [])[-1].content + final_answer = message.get("messages", [])[-1].content logger.info(f"ChatbotAgent: {final_answer}") return final_answer except Exception as e: logger.error(f"CallAgent error: {e}") raise + @tool(name_or_callable="加密计算智能体", description="调用指定智能体进行加密计算的功能") async def call_react_agent(query: str, config: RunnableConfig) -> str: """ @@ -51,14 +53,14 @@ async def call_react_agent(query: str, config: RunnableConfig) -> str: try: input = [{"role": "user", "content": query}] chatbot = agent_manager.get_agent("ReActAgent") - configurable = config.get("configurable",{}) + configurable = config.get("configurable", {}) input_context = { - "thread_id":configurable.get("thread_id"), + "thread_id": configurable.get("thread_id"), "user_id": configurable.get("user_id"), } - message = await chatbot.invoke_messages(input,input_context=input_context) + message = await chatbot.invoke_messages(input, input_context=input_context) # 直接获取最后一个消息的内容 - final_answer = message.get('messages', [])[-1].content + final_answer = message.get("messages", [])[-1].content logger.info(f"ReActAgent: {final_answer}") return final_answer except Exception as e: diff --git a/src/agents/react/graph.py b/src/agents/react/graph.py index f311cdbd..a9d200c6 100644 --- a/src/agents/react/graph.py +++ b/src/agents/react/graph.py @@ -3,10 +3,11 @@ from langgraph.graph import StateGraph from src.agents.common.toolagent import ToolAgent -from .state import State from .context import Context +from .state import State from .tools import get_tools + def tools_branch_continue(state: State): messages = state.messages last_message = messages[-1] @@ -21,14 +22,14 @@ class ReActAgent(ToolAgent): name = "ReActAgent" description = "符合 ReAct 范式的智能体,可以通过调用工具来完成复杂任务。" - ''' + """ 提示词示例: 你是一个智能体助手 你的所有需要计算的工作不能自己计算,只能通过调用加密计算机工具进行,你只能通过调用该工具的方式进行计算才能得到正确答案。计算顺序是正常的。 例如 2+3+2,你必须调用两次加密计算机工具进行运算,最终结果为9。 - ''' + """ def __init__(self, **kwargs): super().__init__(**kwargs) @@ -64,4 +65,3 @@ class ReActAgent(ToolAgent): graph = builder.compile(checkpointer=self.checkpointer, name=self.name) self.graph = graph return graph - diff --git a/src/agents/react/tools.py b/src/agents/react/tools.py index 5dd1c766..42529bdf 100644 --- a/src/agents/react/tools.py +++ b/src/agents/react/tools.py @@ -6,6 +6,7 @@ from src.agents.common.toolkits.mysql import get_mysql_tools from src.agents.common.tools import get_buildin_tools from src.utils import logger + @tool(name_or_callable="加密计算器", description="可以对给定的2个数字选择进行加减乘除四种加密计算") def calculator(a: float, b: float, operation: str) -> float: """ diff --git a/src/agents/reporter/graph.py b/src/agents/reporter/graph.py index 351e17ba..51563f58 100644 --- a/src/agents/reporter/graph.py +++ b/src/agents/reporter/graph.py @@ -4,8 +4,8 @@ from langchain.agents import create_agent from langchain.agents.middleware import ModelRequest, ModelResponse, dynamic_prompt, wrap_model_call from src.agents.common.base import BaseAgent -from src.agents.common.models import load_chat_model from src.agents.common.mcp import get_mcp_tools +from src.agents.common.models import load_chat_model from src.agents.common.toolkits.mysql import get_mysql_tools from src.utils import logger @@ -16,6 +16,7 @@ _mcp_servers = { }, } + @dynamic_prompt def context_aware_prompt(request: ModelRequest) -> str: user_prompt = request.runtime.context.system_prompt diff --git a/src/knowledge/base.py b/src/knowledge/base.py index 329cb964..f58884da 100644 --- a/src/knowledge/base.py +++ b/src/knowledge/base.py @@ -136,9 +136,8 @@ class KnowledgeBase(ABC): """ from src.utils import hashstr - # 从 kwargs 中获取 is_private 配置 - is_private = kwargs.get('is_private', False) + is_private = kwargs.get("is_private", False) prefix = "kb_private_" if is_private else "kb_" db_id = f"{prefix}{hashstr(database_name, with_salt=True)}" diff --git a/web/src/apis/agent_api.js b/web/src/apis/agent_api.js index 5b9c33d5..32643400 100644 --- a/web/src/apis/agent_api.js +++ b/web/src/apis/agent_api.js @@ -137,7 +137,33 @@ export const agentApi = { * 获取所有可用工具的信息 * @returns {Promise} - 工具信息列表 */ - getTools: (agentId) => apiGet(`/api/chat/tools?agent_id=${agentId}`) + getTools: (agentId) => apiGet(`/api/chat/tools?agent_id=${agentId}`), + + /** + * 恢复被人工审批中断的对话(流式响应) + * @param {string} agentId - 智能体ID + * @param {Object} data - 恢复数据 { thread_id, approved } + * @param {Object} options - 可选参数(signal, headers等) + * @returns {Promise} - 恢复响应流 + */ + resumeAgentChat: (agentId, data, options = {}) => { + const { signal, headers: extraHeaders, ...restOptions } = options || {}; + const baseHeaders = { + 'Content-Type': 'application/json', + ...useUserStore().getAuthHeaders() + }; + + return fetch(`/api/chat/agent/${agentId}/resume`, { + method: 'POST', + body: JSON.stringify(data), + signal, + headers: { + ...baseHeaders, + ...(extraHeaders || {}) + }, + ...restOptions + }) + } } diff --git a/web/src/components/AgentChatComponent.vue b/web/src/components/AgentChatComponent.vue index df714f76..55b781fc 100644 --- a/web/src/components/AgentChatComponent.vue +++ b/web/src/components/AgentChatComponent.vue @@ -29,7 +29,7 @@
-
+
新对话
@@ -97,14 +97,14 @@
- +
@@ -117,6 +117,15 @@
+ + +
{ })); }); +// Keep per-thread streaming scratch data in a consistent shape. +const createOnGoingConvState = () => ({ + msgChunks: {}, + currentRequestKey: null, + currentAssistantKey: null, + toolCallBuffers: {} +}); + const chatState = reactive({ currentThreadId: null, isLoadingThreads: false, @@ -227,6 +246,18 @@ const currentThread = computed(() => { const currentThreadMessages = computed(() => threadMessages.value[currentChatId.value] || []); +// 计算是否显示Refs组件的条件 +const shouldShowRefs = computed(() => { + return (conv) => { + return getLastMessage(conv) && + conv.status !== 'streaming' && + !approvalState.showModal && + !(approvalState.threadId && + chatState.currentThreadId === approvalState.threadId && + isProcessing.value); + }; +}); + // 当前线程状态的computed属性 const currentThreadState = computed(() => { return getThreadState(currentChatId.value); @@ -316,7 +347,7 @@ const getThreadState = (threadId) => { chatState.threadStates[threadId] = { isStreaming: false, streamAbortController: null, - onGoingConv: { msgChunks: {} } + onGoingConv: createOnGoingConvState() }; } return chatState.threadStates[threadId]; @@ -349,11 +380,11 @@ const resetOnGoingConv = (threadId = null, preserveMessages = false) => { // 延迟清空消息,给历史记录加载足够时间 setTimeout(() => { if (threadState.onGoingConv) { - threadState.onGoingConv = { msgChunks: {} }; + threadState.onGoingConv = createOnGoingConvState(); } }, 100); } else { - threadState.onGoingConv = { msgChunks: {} }; + threadState.onGoingConv = createOnGoingConvState(); } } } else { @@ -369,11 +400,11 @@ const resetOnGoingConv = (threadId = null, preserveMessages = false) => { if (preserveMessages) { setTimeout(() => { if (threadState.onGoingConv) { - threadState.onGoingConv = { msgChunks: {} }; + threadState.onGoingConv = createOnGoingConvState(); } }, 100); } else { - threadState.onGoingConv = { msgChunks: {} }; + threadState.onGoingConv = createOnGoingConvState(); } } } else { @@ -388,6 +419,7 @@ const resetOnGoingConv = (threadId = null, preserveMessages = false) => { const _processStreamChunk = (chunk, threadId) => { const { status, msg, request_id, message } = chunk; const threadState = getThreadState(threadId); + // console.log('Processing stream chunk:', chunk, 'for thread:', threadId); if (!threadState) return false; @@ -399,10 +431,10 @@ const _processStreamChunk = (chunk, threadId) => { if (msg.id) { if (!threadState.onGoingConv.msgChunks[msg.id]) { threadState.onGoingConv.msgChunks[msg.id] = []; - } + } threadState.onGoingConv.msgChunks[msg.id].push(msg); } - return false; + return false; case 'error': handleChatError({ message }, 'stream'); // Stop the loading indicator @@ -420,6 +452,9 @@ const _processStreamChunk = (chunk, threadId) => { fetchThreadMessages({ agentId: currentAgentId.value, threadId: threadId }); resetOnGoingConv(threadId); return true; + case 'human_approval_required': + // 使用审批 composable 处理审批请求 + return processApprovalInStream(chunk, threadId, currentAgentId.value); case 'finished': // 先标记流式结束,但保持消息显示直到历史记录加载完成 if (threadState) { @@ -542,6 +577,13 @@ const fetchThreadMessages = async ({ agentId, threadId }) => { } }; +// ==================== 审批功能管理 ==================== +const { approvalState, handleApproval, processApprovalInStream } = useApproval({ + getThreadState, + resetOnGoingConv, + fetchThreadMessages +}); + // 发送消息并处理流式响应 const sendMessage = async ({ agentId, threadId, text, signal = undefined }) => { if (!agentId || !threadId || !text) { @@ -590,7 +632,7 @@ const switchToFirstChatIfEmpty = async () => { }; const createNewChat = async () => { - if (!AgentValidator.validateAgentId(currentAgentId.value, '创建对话') || isProcessing.value) return; + if (!AgentValidator.validateAgentId(currentAgentId.value, '创建对话') || chatState.creatingNewChat) return; // 如果第一个对话为空,直接切换到第一个对话而不是创建新对话 if (await switchToFirstChatIfEmpty()) return; @@ -603,6 +645,17 @@ const createNewChat = async () => { try { const newThread = await createThread(currentAgentId.value, '新的对话'); if (newThread) { + // 中断之前线程的流式输出(如果存在) + const previousThreadId = chatState.currentThreadId; + if (previousThreadId) { + const previousThreadState = getThreadState(previousThreadId); + if (previousThreadState?.isStreaming && previousThreadState.streamAbortController) { + previousThreadState.streamAbortController.abort(); + previousThreadState.isStreaming = false; + previousThreadState.streamAbortController = null; + } + } + chatState.currentThreadId = newThread.id; } } catch (error) { @@ -615,8 +668,17 @@ const createNewChat = async () => { const selectChat = async (chatId) => { if (!AgentValidator.validateAgentIdWithError(currentAgentId.value, '选择对话', handleValidationError)) return; - // 切换线程时,不再中断上一个线程的流式输出 - // resetOnGoingConv(chatState.currentThreadId); + // 中断之前线程的流式输出(如果存在) + const previousThreadId = chatState.currentThreadId; + if (previousThreadId && previousThreadId !== chatId) { + const previousThreadState = getThreadState(previousThreadId); + if (previousThreadState?.isStreaming && previousThreadState.streamAbortController) { + previousThreadState.streamAbortController.abort(); + previousThreadState.isStreaming = false; + previousThreadState.streamAbortController = null; + } + } + chatState.currentThreadId = chatId; chatState.isLoadingMessages = true; try { @@ -763,6 +825,106 @@ const handleSendOrStop = async () => { await handleSendMessage(); }; +// ==================== 人工审批处理 ==================== +const handleApprovalWithStream = async (approved) => { + console.log('🔄 [STREAM] Starting resume stream processing'); + + const threadId = approvalState.threadId; + if (!threadId) { + message.error('无效的审批请求'); + approvalState.showModal = false; + return; + } + + const threadState = getThreadState(threadId); + if (!threadState) { + message.error('无法找到对应的对话线程'); + approvalState.showModal = false; + return; + } + + try { + // 使用审批 composable 处理审批 + const response = await handleApproval(approved, currentAgentId.value); + + if (!response) return; // 如果 handleApproval 抛出错误,这里不会执行 + + console.log('🔄 [STREAM] Processing resume streaming response'); + + // 处理流式响应 + const reader = response.body.getReader(); + const decoder = new TextDecoder(); + let buffer = ''; + let stopReading = false; + + while (!stopReading) { + const { done, value } = await reader.read(); + if (done) break; + + buffer += decoder.decode(value, { stream: true }); + const lines = buffer.split('\n'); + buffer = lines.pop() || ''; + + for (const line of lines) { + const trimmedLine = line.trim(); + if (trimmedLine) { + try { + const chunk = JSON.parse(trimmedLine); + console.log('🔄 [STREAM] Processing chunk:', chunk); + + // 处理chunk并更新对话 - _processStreamChunk 已经处理了所有必要的逻辑 + if (_processStreamChunk(chunk, threadId)) { + stopReading = true; + break; + } + + } catch (e) { + console.warn('Failed to parse stream chunk JSON:', e, 'Line:', trimmedLine); + } + } + } + } + + if (!stopReading && buffer.trim()) { + try { + const chunk = JSON.parse(buffer.trim()); + console.log('🔄 [STREAM] Processing final chunk:', chunk); + + // 处理最终chunk - _processStreamChunk 已经处理了所有必要的逻辑 + if (_processStreamChunk(chunk, threadId)) { + stopReading = true; + } + + } catch (e) { + console.warn('Failed to parse final stream chunk JSON:', e); + } + } + + console.log('🔄 [STREAM] Resume stream processing completed'); + + } catch (error) { + console.error('❌ [STREAM] Resume stream failed:', error); + if (error.name !== 'AbortError') { + console.error('Resume approval error:', error); + // handleChatError 已在 useApproval 中调用 + } + } finally { + console.log('🔄 [STREAM] Cleaning up streaming state'); + if (threadState) { + threadState.isStreaming = false; + threadState.streamAbortController = null; + } + } +}; + +const handleApprove = () => { + handleApprovalWithStream(true); +}; + +const handleReject = () => { + handleApprovalWithStream(false); +}; + // ==================== UI HANDLERS ==================== const handleKeyDown = (e) => { if (e.key === 'Enter' && !e.shiftKey) { @@ -795,7 +957,6 @@ defineExpose({ getExportPayload: buildExportPayload }); -const retryMessage = (msg) => { /* TODO */ }; const toggleSidebar = () => { uiState.isSidebarOpen = !uiState.isSidebarOpen; localStorage.setItem('chat_sidebar_open', uiState.isSidebarOpen); @@ -812,7 +973,24 @@ const getLastMessage = (conv) => { }; const showMsgRefs = (msg) => { - if (msg.isLast) return ['copy']; + // 如果正在审批中,不显示 refs + if (approvalState.showModal) { + return false; + } + + // 如果当前线程ID与审批线程ID匹配,但审批框已关闭(说明刚刚处理完审批) + // 且当前有新的流式处理正在进行,则不显示之前被中断的消息的 refs + if (approvalState.threadId && + chatState.currentThreadId === approvalState.threadId && + !approvalState.showModal && + isProcessing) { + return false; + } + + // 只有真正完成的消息才显示 refs + if (msg.isLast && msg.status === 'finished') { + return ['copy']; + } return false; }; @@ -875,6 +1053,7 @@ watch(currentAgentId, async (newAgentId, oldAgentId) => { } }, { immediate: true }); + watch(conversations, () => { if (isProcessing.value) { scrollController.scrollToBottom(); diff --git a/web/src/components/AgentMessageComponent.vue b/web/src/components/AgentMessageComponent.vue index cf0913a9..2e27f4d6 100644 --- a/web/src/components/AgentMessageComponent.vue +++ b/web/src/components/AgentMessageComponent.vue @@ -3,6 +3,8 @@

{{ message.content }}

+

{{ message.content }}

+
@@ -232,6 +234,19 @@ const toggleToolCall = (toolCallId) => { white-space: pre-line; } + .message-text-system { + max-width: 100%; + margin-bottom: 0; + white-space: pre-line; + color: var(--gray-600); + font-style: italic; + font-size: 14px; + padding: 8px 12px; + background-color: var(--gray-50); + border-left: 3px solid var(--gray-300); + border-radius: 4px; + } + .err-msg { color: #d15252; border: 1px solid #f19999; diff --git a/web/src/components/HumanApprovalModal.vue b/web/src/components/HumanApprovalModal.vue new file mode 100644 index 00000000..6d8e0e57 --- /dev/null +++ b/web/src/components/HumanApprovalModal.vue @@ -0,0 +1,238 @@ + + + + + diff --git a/web/src/composables/useApproval.js b/web/src/composables/useApproval.js new file mode 100644 index 00000000..1eab8730 --- /dev/null +++ b/web/src/composables/useApproval.js @@ -0,0 +1,128 @@ +import { reactive } from 'vue'; +import { message } from 'ant-design-vue'; +import { handleChatError } from '@/utils/errorHandler'; +import { agentApi } from '@/apis'; + +export function useApproval({ getThreadState, resetOnGoingConv, fetchThreadMessages }) { + // 审批状态 + const approvalState = reactive({ + showModal: false, + question: '', + operation: '', + threadId: null, + interruptInfo: null + }); + + // 处理审批逻辑 + const handleApproval = async (approved, currentAgentId) => { + const threadId = approvalState.threadId; + if (!threadId) { + message.error('无效的审批请求'); + approvalState.showModal = false; + return; + } + + const threadState = getThreadState(threadId); + if (!threadState) { + message.error('无法找到对应的对话线程'); + approvalState.showModal = false; + return; + } + + // 关闭弹窗 + approvalState.showModal = false; + + // 清理旧的流式控制器(如果存在) + if (threadState.streamAbortController) { + threadState.streamAbortController.abort(); + threadState.streamAbortController = null; + } + + // 标记为处理中 + threadState.isStreaming = true; + resetOnGoingConv(threadId); + threadState.streamAbortController = new AbortController(); + + console.log('🔄 [APPROVAL] Starting resume process:', { approved, threadId, currentAgentId }); + + try { + // 调用恢复接口 + const response = await agentApi.resumeAgentChat( + currentAgentId, + { + thread_id: threadId, + approved: approved + }, + { + signal: threadState.streamAbortController?.signal + } + ); + + console.log('🔄 [APPROVAL] Resume API response received'); + + if (!response.ok) { + const errorText = await response.text(); + console.error('Resume API error:', response.status, errorText); + throw new Error(`HTTP error! status: ${response.status}, details: ${errorText}`); + } + + console.log('🔄 [APPROVAL] Resume API successful, returning response for stream processing'); + return response; // 返回响应供调用方处理流式数据 + + } catch (error) { + console.error('❌ [APPROVAL] Resume failed:', error); + if (error.name !== 'AbortError') { + handleChatError(error, 'resume'); + message.error(`恢复对话失败: ${error.message || '未知错误'}`); + } + // 重置状态 - 只在错误时重置 + threadState.isStreaming = false; + threadState.streamAbortController = null; + throw error; // 重新抛出错误让调用方处理 + } + // 移除 finally 块 - 让组件管理流式状态的生命周期 + }; + + // 在流式处理中处理审批请求 + const processApprovalInStream = (chunk, threadId, currentAgentId) => { + if (chunk.status !== 'human_approval_required') { + return false; + } + + const { interrupt_info } = chunk; + const threadState = getThreadState(threadId); + + if (!threadState) return false; + + // 停止显示"处理中"状态,让用户可以看到并操作审批弹窗 + threadState.isStreaming = false; + + // 显示审批弹窗 + approvalState.showModal = true; + approvalState.question = interrupt_info?.question || '是否批准以下操作?'; + approvalState.operation = interrupt_info?.operation || '未知操作'; + approvalState.threadId = chunk.thread_id || threadId; + approvalState.interruptInfo = interrupt_info; + + // 刷新消息历史显示已执行的部分 + fetchThreadMessages({ agentId: currentAgentId, threadId: threadId }); + + return true; // 表示已处理审批请求,应停止流式处理 + }; + + // 重置审批状态 + const resetApprovalState = () => { + approvalState.showModal = false; + approvalState.question = ''; + approvalState.operation = ''; + approvalState.threadId = null; + approvalState.interruptInfo = null; + }; + + return { + approvalState, + handleApproval, + processApprovalInStream, + resetApprovalState + }; +} \ No newline at end of file diff --git a/web/src/utils/messageProcessor.js b/web/src/utils/messageProcessor.js index a892c42e..fdb89278 100644 --- a/web/src/utils/messageProcessor.js +++ b/web/src/utils/messageProcessor.js @@ -134,9 +134,6 @@ export class MessageProcessor { // 处理AIMessageChunk类型 if (result.type === 'AIMessageChunk') { result.type = 'ai'; - if (result.additional_kwargs?.tool_calls) { - result.tool_calls = result.additional_kwargs.tool_calls; - } } return result; @@ -149,23 +146,47 @@ export class MessageProcessor { * @param {Object} chunk - 当前块 */ static _mergeToolCalls(result, chunk) { - if (chunk.additional_kwargs?.tool_calls) { - if (!result.additional_kwargs) result.additional_kwargs = {}; - if (!result.additional_kwargs.tool_calls) result.additional_kwargs.tool_calls = []; + if (chunk.tool_call_chunks && chunk.tool_call_chunks.length > 0) { + // 确保 result 有 tool_calls 数组 + if (!result.tool_calls) result.tool_calls = []; - for (const toolCall of chunk.additional_kwargs.tool_calls) { - const existingToolCall = result.additional_kwargs.tool_calls.find( - t => (t.id === toolCall.id || t.index === toolCall.index) + for (const toolCallChunk of chunk.tool_call_chunks) { + // 使用 index 来标识工具调用(因为可能有多个工具调用) + const existingToolCallIndex = result.tool_calls.findIndex( + t => t.index === toolCallChunk.index ); - if (existingToolCall) { - // 合并相同ID的tool call - if (existingToolCall.function && toolCall.function) { - existingToolCall.function.arguments += toolCall.function.arguments; + if (existingToolCallIndex !== -1) { + // 合并相同index的tool call + const existingToolCall = result.tool_calls[existingToolCallIndex]; + + // 更新名称和ID(如果存在) + if (toolCallChunk.name && !existingToolCall.function?.name) { + if (!existingToolCall.function) existingToolCall.function = {}; + existingToolCall.function.name = toolCallChunk.name; + } + + if (toolCallChunk.id && !existingToolCall.id) { + existingToolCall.id = toolCallChunk.id; + } + + // 合并参数 + if (toolCallChunk.args) { + if (!existingToolCall.function) existingToolCall.function = {}; + if (!existingToolCall.function.arguments) existingToolCall.function.arguments = ''; + existingToolCall.function.arguments += toolCallChunk.args; } } else { // 添加新的tool call - result.additional_kwargs.tool_calls.push(JSON.parse(JSON.stringify(toolCall))); + const newToolCall = { + index: toolCallChunk.index, + id: toolCallChunk.id, + function: { + name: toolCallChunk.name || null, + arguments: toolCallChunk.args || '' + } + }; + result.tool_calls.push(newToolCall); } } }