mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-08 17:08:35 +08:00
fix agent tool call history ordering
This commit is contained in:
@@ -1,7 +1,7 @@
|
|||||||
from typing import Any
|
from typing import Any, Optional
|
||||||
|
|
||||||
from langchain.agents.middleware import AgentMiddleware, AgentState
|
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.runtime import Runtime
|
||||||
from langgraph.types import Overwrite
|
from langgraph.types import Overwrite
|
||||||
|
|
||||||
@@ -9,35 +9,65 @@ from langgraph.types import Overwrite
|
|||||||
class PatchToolCallsMiddleware(AgentMiddleware):
|
class PatchToolCallsMiddleware(AgentMiddleware):
|
||||||
"""修复消息历史中悬空工具调用的中间件。"""
|
"""修复消息历史中悬空工具调用的中间件。"""
|
||||||
|
|
||||||
def before_agent(self, state: AgentState, runtime: Runtime[Any]) -> dict[str, Any] | None: # noqa: ARG002
|
@staticmethod
|
||||||
"""在代理运行之前,处理任何 AIMessage 中悬空的工具调用。"""
|
def _build_cancelled_tool_message(tool_call: dict[str, Any]) -> ToolMessage:
|
||||||
messages = state["messages"]
|
"""构造取消状态的工具响应消息。"""
|
||||||
if not messages or len(messages) == 0:
|
tool_name = tool_call.get("name") or "unknown_tool"
|
||||||
return None
|
tool_call_id = tool_call.get("id") or ""
|
||||||
|
|
||||||
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 = (
|
tool_msg = (
|
||||||
f"Tool call {tool_call['name']} with id {tool_call['id']} was "
|
f"Tool call {tool_name} with id {tool_call_id} was "
|
||||||
"cancelled - another message came in before it could be completed."
|
"cancelled - another message came in before it could be completed."
|
||||||
)
|
)
|
||||||
patched_messages.append(
|
return ToolMessage(
|
||||||
ToolMessage(
|
|
||||||
content=tool_msg,
|
content=tool_msg,
|
||||||
name=tool_call["name"],
|
name=tool_name,
|
||||||
tool_call_id=tool_call["id"],
|
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
|
||||||
|
|
||||||
|
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)}
|
return {"messages": Overwrite(patched_messages)}
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user