diff --git a/CLAUDE.md b/CLAUDE.md index a865c741..5c11c218 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1 +1 @@ -Read AGENTS.md \ No newline at end of file +See AGENTS.md 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 38e6dfe4..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]): """根据配置获取工具。 默认不使用任何工具。 @@ -84,4 +83,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..524b1bd4 100644 --- a/src/agents/common/tools.py +++ b/src/agents/common/tools.py @@ -54,6 +54,50 @@ def get_approved_user_goal( return result +# 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": "是否批准以下操作?", + "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: """Use this to query knowledge graph, which include some food domain knowledge.""" @@ -72,10 +116,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/__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..e11f7751 --- /dev/null +++ b/src/agents/mini_agent/graph.py @@ -0,0 +1,48 @@ +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/multi_agent/__init__.py b/src/agents/multi_agent/__init__.py new file mode 100644 index 00000000..740c89a0 --- /dev/null +++ b/src/agents/multi_agent/__init__.py @@ -0,0 +1,3 @@ +from .graph import SampleMultiAgent + +__all__ = ["SampleMultiAgent"] diff --git a/src/agents/multi_agent/context.py b/src/agents/multi_agent/context.py new file mode 100644 index 00000000..7bb1adad --- /dev/null +++ b/src/agents/multi_agent/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/multi_agent/graph.py b/src/agents/multi_agent/graph.py new file mode 100644 index 00000000..8db3db0c --- /dev/null +++ b/src/agents/multi_agent/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/multi_agent/state.py b/src/agents/multi_agent/state.py new file mode 100644 index 00000000..927bfa95 --- /dev/null +++ b/src/agents/multi_agent/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/multi_agent/tools.py b/src/agents/multi_agent/tools.py new file mode 100644 index 00000000..fc3fe951 --- /dev/null +++ b/src/agents/multi_agent/tools.py @@ -0,0 +1,76 @@ +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.tools import get_buildin_tools +from src.utils import logger + + +# TODO[修改建议]:能不能通过前端直接指定子智能体? +# 调用子智能体后的日志是输出到tool_calls的 +@tool(name_or_callable="对话聊天智能体", description="调用指定智能体进行对话聊天的功能") +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") + 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}") + 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: + """ + 调用指定智能体进行加密计算的功能 + + Args: + query: 根据需要构造的提问 + config: LangGraph运行时配置(自动注入) + Returns: + str: 最终的回答结果 + """ + try: + input = [{"role": "user", "content": query}] + chatbot = agent_manager.get_agent("ReActAgent") + 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}") + 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/graph.py b/src/agents/react/graph.py index b4abb844..50fad7a9 100644 --- a/src/agents/react/graph.py +++ b/src/agents/react/graph.py @@ -1,12 +1,13 @@ -from langgraph.constants import START, END +from langgraph.constants import END 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] @@ -18,8 +19,17 @@ 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 范式的智能体,可以通过调用工具来完成复杂任务。" + + """ + 提示词示例: + 你是一个智能体助手 + + 你的所有需要计算的工作不能自己计算,只能通过调用加密计算机工具进行,你只能通过调用该工具的方式进行计算才能得到正确答案。计算顺序是正常的。 + + 例如 2+3+2,你必须调用两次加密计算机工具进行运算,最终结果为9。 + """ # TODO:[已完成] React智能体 ''' diff --git a/src/agents/react/tools.py b/src/agents/react/tools.py index f3ec28cc..42529bdf 100644 --- a/src/agents/react/tools.py +++ b/src/agents/react/tools.py @@ -1,14 +1,12 @@ -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: """ diff --git a/src/agents/reporter/graph.py b/src/agents/reporter/graph.py index 8402722a..51563f58 100644 --- a/src/agents/reporter/graph.py +++ b/src/agents/reporter/graph.py @@ -1,12 +1,11 @@ import textwrap -from pathlib import Path 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 @@ -17,6 +16,7 @@ _mcp_servers = { }, } + @dynamic_prompt def context_aware_prompt(request: ModelRequest) -> str: user_prompt = request.runtime.context.system_prompt @@ -64,4 +64,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..f58884da 100644 --- a/src/knowledge/base.py +++ b/src/knowledge/base.py @@ -136,10 +136,8 @@ class KnowledgeBase(ABC): """ from src.utils import hashstr - 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 a8209206..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,7 +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('msg:', msg); + // console.log('Processing stream chunk:', chunk, 'for thread:', threadId); if (!threadState) return false; @@ -400,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 @@ -421,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) { @@ -543,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) { @@ -591,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; @@ -604,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) { @@ -616,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 { @@ -764,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) { @@ -796,7 +957,6 @@ defineExpose({ getExportPayload: buildExportPayload }); -const retryMessage = (msg) => { /* TODO */ }; const toggleSidebar = () => { uiState.isSidebarOpen = !uiState.isSidebarOpen; localStorage.setItem('chat_sidebar_open', uiState.isSidebarOpen); @@ -813,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; }; @@ -876,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 45ec6009..796931a7 100644 --- a/web/src/components/AgentMessageComponent.vue +++ b/web/src/components/AgentMessageComponent.vue @@ -3,6 +3,8 @@

{{ message.content }}

+

{{ message.content }}

+
@@ -272,6 +274,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 0fc75fa3..d7fe3799 100644 --- a/web/src/utils/messageProcessor.js +++ b/web/src/utils/messageProcessor.js @@ -138,10 +138,6 @@ export class MessageProcessor { // 处理AIMessageChunk类型 if (result.type === 'AIMessageChunk') { result.type = 'ai'; - // 将tool_calls从additional_kwargs移到顶层,并确保格式正确 - if (result.additional_kwargs?.tool_calls) { - result.tool_calls = result.additional_kwargs.tool_calls; - } } return result; @@ -154,78 +150,47 @@ export class MessageProcessor { * @param {Object} chunk - 当前块 */ static _mergeToolCalls(result, chunk) { - // 1. 处理 additional_kwargs.tool_calls (旧格式,保持兼容性) - if (chunk.additional_kwargs?.tool_calls) { - if (!result.additional_kwargs) result.additional_kwargs = {}; - if (!result.additional_kwargs.tool_calls) result.additional_kwargs.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) - ); - - if (existingToolCall) { - // 合并相同ID的tool call - if (existingToolCall.function && toolCall.function) { - existingToolCall.function.arguments += toolCall.function.arguments; - } - } else { - // 添加新的tool call - result.additional_kwargs.tool_calls.push(JSON.parse(JSON.stringify(toolCall))); - } - } - } - - // 2. 处理顶层的 tool_calls (新格式) - if (chunk.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.tool_calls) { - // 过滤掉无效的工具调用(没有id或name的空对象) - if (!toolCall.id && !toolCall.name) { - continue; - } + for (const toolCallChunk of chunk.tool_call_chunks) { + // 使用 index 来标识工具调用(因为可能有多个工具调用) + const existingToolCallIndex = result.tool_calls.findIndex( + t => t.index === toolCallChunk.index + ); - const existingToolCall = result.tool_calls.find(t => t.id === toolCall.id); + if (existingToolCallIndex !== -1) { + // 合并相同index的tool call + const existingToolCall = result.tool_calls[existingToolCallIndex]; - if (existingToolCall) { - // 合并相同ID的tool call的args - if (toolCall.args && existingToolCall.args !== undefined) { - existingToolCall.args += toolCall.args; + // 更新名称和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.tool_calls.push(JSON.parse(JSON.stringify(toolCall))); - } - } - } - - // 3. 处理 tool_call_chunks (分片的工具调用参数) - if (chunk.tool_call_chunks) { - if (!result.tool_call_chunks) result.tool_call_chunks = []; - - for (const toolCallChunk of chunk.tool_call_chunks) { - // 过滤掉无效的chunk(没有id、name和args的空对象) - if (!toolCallChunk.id && !toolCallChunk.name && !toolCallChunk.args) { - continue; - } - - const existingChunk = result.tool_call_chunks.find( - t => (t.id === toolCallChunk.id || (t.index === toolCallChunk.index && t.index !== null)) - ); - - if (existingChunk) { - // 合并参数字符串 - if (toolCallChunk.args && existingChunk.args !== undefined) { - existingChunk.args += toolCallChunk.args; - } - // 更新工具名称 - if (toolCallChunk.name && !existingChunk.name) { - existingChunk.name = toolCallChunk.name; - } - } else { - // 添加新的chunk - result.tool_call_chunks.push(JSON.parse(JSON.stringify(toolCallChunk))); + const newToolCall = { + index: toolCallChunk.index, + id: toolCallChunk.id, + function: { + name: toolCallChunk.name || null, + arguments: toolCallChunk.args || '' + } + }; + result.tool_calls.push(newToolCall); } } }