diff --git a/app/agent/middleware/patch_tool_calls.py b/app/agent/middleware/patch_tool_calls.py index 0bf6a48f1..4e4870d3d 100644 --- a/app/agent/middleware/patch_tool_calls.py +++ b/app/agent/middleware/patch_tool_calls.py @@ -1,7 +1,7 @@ -from typing import Any +from typing import Any, Optional from langchain.agents.middleware import AgentMiddleware, AgentState -from langchain_core.messages import AIMessage, ToolMessage +from langchain_core.messages import AIMessage, BaseMessage, ToolMessage from langgraph.runtime import Runtime from langgraph.types import Overwrite @@ -9,35 +9,65 @@ from langgraph.types import Overwrite class PatchToolCallsMiddleware(AgentMiddleware): """修复消息历史中悬空工具调用的中间件。""" - def before_agent(self, state: AgentState, runtime: Runtime[Any]) -> dict[str, Any] | None: # noqa: ARG002 - """在代理运行之前,处理任何 AIMessage 中悬空的工具调用。""" - messages = state["messages"] + @staticmethod + def _build_cancelled_tool_message(tool_call: dict[str, Any]) -> ToolMessage: + """构造取消状态的工具响应消息。""" + tool_name = tool_call.get("name") or "unknown_tool" + tool_call_id = tool_call.get("id") or "" + tool_msg = ( + f"Tool call {tool_name} with id {tool_call_id} was " + "cancelled - another message came in before it could be completed." + ) + return ToolMessage( + content=tool_msg, + name=tool_name, + tool_call_id=tool_call_id, + ) + + @classmethod + def _normalize_messages(cls, messages: list[BaseMessage]) -> list[BaseMessage]: + """规范化工具调用消息顺序,满足 OpenAI tool_calls 协议要求。""" if not messages or len(messages) == 0: + return messages + + tool_messages = { + msg.tool_call_id: msg + for msg in messages + if isinstance(msg, ToolMessage) and msg.tool_call_id + } + patched_messages = [] + for msg in messages: + if isinstance(msg, ToolMessage): + continue + + patched_messages.append(msg) + if not isinstance(msg, AIMessage) or not msg.tool_calls: + continue + + for tool_call in msg.tool_calls: + tool_call_id = tool_call.get("id") + corresponding_tool_msg = tool_messages.get(tool_call_id) + if corresponding_tool_msg: + patched_messages.append(corresponding_tool_msg) + else: + patched_messages.append(cls._build_cancelled_tool_message(tool_call)) + + return patched_messages + + def before_agent(self, state: AgentState, runtime: Runtime[Any]) -> Optional[dict[str, Any]]: # noqa: ARG002 + """在代理运行之前,处理任何 AIMessage 中悬空或乱序的工具调用。""" + messages = state["messages"] + patched_messages = self._normalize_messages(messages) + if patched_messages == messages: return None - patched_messages = [] - # 遍历消息并添加任何悬空的工具调用 - for i, msg in enumerate(messages): - patched_messages.append(msg) - if isinstance(msg, AIMessage) and msg.tool_calls: - for tool_call in msg.tool_calls: - corresponding_tool_msg = next( - (msg for msg in messages[i:] if msg.type == "tool" and msg.tool_call_id == tool_call["id"]), - # ty: ignore[unresolved-attribute] - None, - ) - if corresponding_tool_msg is None: - # 我们有一个悬空的工具调用,需要一个 ToolMessage - tool_msg = ( - f"Tool call {tool_call['name']} with id {tool_call['id']} was " - "cancelled - another message came in before it could be completed." - ) - patched_messages.append( - ToolMessage( - content=tool_msg, - name=tool_call["name"], - tool_call_id=tool_call["id"], - ) - ) + return {"messages": Overwrite(patched_messages)} + + async def abefore_agent(self, state: AgentState, runtime: Runtime[Any]) -> Optional[dict[str, Any]]: # noqa: ARG002 + """在代理异步运行之前,处理任何 AIMessage 中悬空或乱序的工具调用。""" + messages = state["messages"] + patched_messages = self._normalize_messages(messages) + if patched_messages == messages: + return None return {"messages": Overwrite(patched_messages)} diff --git a/tests/test_agent_patch_tool_calls.py b/tests/test_agent_patch_tool_calls.py new file mode 100644 index 000000000..570f3115c --- /dev/null +++ b/tests/test_agent_patch_tool_calls.py @@ -0,0 +1,89 @@ +import asyncio +import unittest + +from langchain_core.messages import AIMessage, HumanMessage, ToolMessage + +from app.agent.middleware.patch_tool_calls import PatchToolCallsMiddleware + + +def _build_tool_call(tool_call_id: str = "call_1", name: str = "search") -> dict: + """构造测试用工具调用。""" + return { + "id": tool_call_id, + "type": "tool_call", + "name": name, + "args": {}, + } + + +class TestPatchToolCallsMiddleware(unittest.TestCase): + """测试工具调用历史修复中间件。""" + + def test_adds_missing_tool_messages_immediately_after_ai_message(self): + """缺失工具响应时应立即补齐 ToolMessage。""" + middleware = PatchToolCallsMiddleware() + messages = [ + HumanMessage(content="查天气"), + AIMessage(content="", tool_calls=[_build_tool_call()]), + HumanMessage(content="不用查了"), + ] + + result = middleware.before_agent({"messages": messages}, runtime=None) + + patched_messages = result["messages"].value + self.assertIs(patched_messages[1], messages[1]) + self.assertIsInstance(patched_messages[2], ToolMessage) + self.assertEqual(patched_messages[2].tool_call_id, "call_1") + self.assertIs(patched_messages[3], messages[2]) + + def test_moves_late_tool_messages_next_to_matching_ai_message(self): + """乱序工具响应应移动到对应 assistant 消息之后。""" + middleware = PatchToolCallsMiddleware() + tool_message = ToolMessage(content="晴天", tool_call_id="call_1") + messages = [ + HumanMessage(content="查天气"), + AIMessage(content="", tool_calls=[_build_tool_call()]), + HumanMessage(content="再问一句"), + tool_message, + ] + + result = middleware.before_agent({"messages": messages}, runtime=None) + + patched_messages = result["messages"].value + self.assertIs(patched_messages[1], messages[1]) + self.assertIs(patched_messages[2], tool_message) + self.assertIs(patched_messages[3], messages[2]) + self.assertNotIn(tool_message, patched_messages[4:]) + + def test_drops_orphan_tool_messages(self): + """孤立工具响应不应继续进入模型请求历史。""" + middleware = PatchToolCallsMiddleware() + orphan_tool_message = ToolMessage(content="晴天", tool_call_id="call_orphan") + messages = [ + HumanMessage(content="查天气"), + orphan_tool_message, + HumanMessage(content="继续"), + ] + + result = middleware.before_agent({"messages": messages}, runtime=None) + + patched_messages = result["messages"].value + self.assertEqual([msg.type for msg in patched_messages], ["human", "human"]) + self.assertNotIn(orphan_tool_message, patched_messages) + + def test_async_hook_normalizes_messages(self): + """异步 Agent 执行入口也应修复工具调用历史。""" + middleware = PatchToolCallsMiddleware() + messages = [ + HumanMessage(content="查天气"), + AIMessage(content="", tool_calls=[_build_tool_call()]), + ] + + result = asyncio.run(middleware.abefore_agent({"messages": messages}, runtime=None)) + + patched_messages = result["messages"].value + self.assertEqual([msg.type for msg in patched_messages], ["human", "ai", "tool"]) + + +if __name__ == "__main__": + unittest.main()