mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
fix(langchain): repair invalid tool calls in create_agent (#40530)
Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
1 parent
fbac0b2d0c
commit
bfa03c8cf8
2 files changed
+327
-9
No files matched your search
@@ -19,10 +19,17 @@ from typing import (
|
||||
)
|
||||
|
||||
from langchain_core.language_models.chat_models import BaseChatModel
|
||||
from langchain_core.messages import AIMessage, AnyMessage, SystemMessage, ToolMessage
|
||||
from langchain_core.messages import (
|
||||
AIMessage,
|
||||
AnyMessage,
|
||||
RemoveMessage,
|
||||
SystemMessage,
|
||||
ToolMessage,
|
||||
)
|
||||
from langchain_core.tools import BaseTool
|
||||
from langgraph._internal._runnable import RunnableCallable
|
||||
from langgraph.constants import END, START
|
||||
from langgraph.graph.message import REMOVE_ALL_MESSAGES
|
||||
from langgraph.graph.state import StateGraph
|
||||
from langgraph.prebuilt import ToolCallTransformer
|
||||
from langgraph.prebuilt.tool_node import ToolNode
|
||||
@@ -83,6 +90,7 @@ class _ComposedExtendedModelResponse(Generic[ResponseT]):
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Awaitable, Callable, Iterable, Sequence
|
||||
|
||||
from langchain_core.messages import InvalidToolCall
|
||||
from langchain_core.runnables import Runnable, RunnableConfig
|
||||
from langgraph.cache.base import BaseCache
|
||||
from langgraph.graph.state import CompiledStateGraph
|
||||
@@ -215,6 +223,7 @@ def _build_commands(
|
||||
middleware_commands: list[Command[Any]] | None = None,
|
||||
*,
|
||||
has_structured_output: bool = False,
|
||||
repaired_messages: list[AnyMessage] | None = None,
|
||||
) -> list[Command[Any]]:
|
||||
"""Build a list of Commands from a model response and middleware commands.
|
||||
|
||||
@@ -230,11 +239,19 @@ def _build_commands(
|
||||
`response_format`. When `True` and no structured response was
|
||||
produced, `structured_response` is explicitly cleared to avoid a
|
||||
stale value from a previous checkpointed turn.
|
||||
repaired_messages: Complete message history after repairing invalid tool calls.
|
||||
|
||||
Returns:
|
||||
List of `Command` objects ready to be returned from a model node.
|
||||
"""
|
||||
state: dict[str, Any] = {"messages": model_response.result}
|
||||
messages = model_response.result
|
||||
if repaired_messages is not None:
|
||||
messages = [
|
||||
RemoveMessage(id=REMOVE_ALL_MESSAGES),
|
||||
*repaired_messages,
|
||||
*model_response.result,
|
||||
]
|
||||
state: dict[str, Any] = {"messages": messages}
|
||||
|
||||
if model_response.structured_response is not None:
|
||||
state["structured_response"] = model_response.structured_response
|
||||
@@ -655,6 +672,40 @@ def _handle_structured_output_error(
|
||||
return True, handle_errors(exception)
|
||||
|
||||
|
||||
def _invalid_tool_call_message(tool_call: InvalidToolCall) -> ToolMessage | None:
|
||||
tool_call_id = tool_call.get("id")
|
||||
if tool_call_id is None:
|
||||
return None
|
||||
name = tool_call.get("name") or "unknown"
|
||||
return ToolMessage(
|
||||
content=(
|
||||
f"Tool call {name} with id {tool_call_id} could not be executed - "
|
||||
"arguments were malformed or truncated."
|
||||
),
|
||||
name=name,
|
||||
tool_call_id=tool_call_id,
|
||||
status="error",
|
||||
)
|
||||
|
||||
|
||||
def _patch_invalid_tool_calls(messages: Sequence[AnyMessage]) -> list[AnyMessage]:
|
||||
answered_ids = {
|
||||
message.tool_call_id for message in messages if isinstance(message, ToolMessage)
|
||||
}
|
||||
patched_messages: list[AnyMessage] = []
|
||||
for message in messages:
|
||||
patched_messages.append(message)
|
||||
if not isinstance(message, AIMessage):
|
||||
continue
|
||||
for tool_call in message.invalid_tool_calls:
|
||||
if tool_call.get("id") in answered_ids:
|
||||
continue
|
||||
if tool_message := _invalid_tool_call_message(tool_call):
|
||||
patched_messages.append(tool_message)
|
||||
answered_ids.add(tool_message.tool_call_id)
|
||||
return patched_messages
|
||||
|
||||
|
||||
def _chain_tool_call_wrappers(
|
||||
wrappers: Sequence[ToolCallWrapper],
|
||||
) -> ToolCallWrapper | None:
|
||||
@@ -1467,12 +1518,13 @@ def create_agent(
|
||||
|
||||
def model_node(state: AgentState[Any], runtime: Runtime[ContextT]) -> list[Command[Any]]:
|
||||
"""Sync model request handler with sequential middleware processing."""
|
||||
messages = _patch_invalid_tool_calls(state["messages"])
|
||||
request = ModelRequest(
|
||||
model=model,
|
||||
tools=default_tools,
|
||||
system_message=system_message,
|
||||
response_format=initial_response_format,
|
||||
messages=state["messages"],
|
||||
messages=messages,
|
||||
tool_choice=None,
|
||||
state=state,
|
||||
runtime=runtime,
|
||||
@@ -1481,11 +1533,18 @@ def create_agent(
|
||||
has_structured_output = initial_response_format is not None
|
||||
if wrap_model_call_handler is None:
|
||||
model_response = _execute_model_sync(request)
|
||||
return _build_commands(model_response, has_structured_output=has_structured_output)
|
||||
return _build_commands(
|
||||
model_response,
|
||||
has_structured_output=has_structured_output,
|
||||
repaired_messages=messages if messages != state["messages"] else None,
|
||||
)
|
||||
|
||||
result = wrap_model_call_handler(request, _execute_model_sync)
|
||||
return _build_commands(
|
||||
result.model_response, result.commands, has_structured_output=has_structured_output
|
||||
result.model_response,
|
||||
result.commands,
|
||||
has_structured_output=has_structured_output,
|
||||
repaired_messages=messages if messages != state["messages"] else None,
|
||||
)
|
||||
|
||||
async def _execute_model_async(request: ModelRequest[ContextT]) -> ModelResponse:
|
||||
@@ -1518,12 +1577,13 @@ def create_agent(
|
||||
|
||||
async def amodel_node(state: AgentState[Any], runtime: Runtime[ContextT]) -> list[Command[Any]]:
|
||||
"""Async model request handler with sequential middleware processing."""
|
||||
messages = _patch_invalid_tool_calls(state["messages"])
|
||||
request = ModelRequest(
|
||||
model=model,
|
||||
tools=default_tools,
|
||||
system_message=system_message,
|
||||
response_format=initial_response_format,
|
||||
messages=state["messages"],
|
||||
messages=messages,
|
||||
tool_choice=None,
|
||||
state=state,
|
||||
runtime=runtime,
|
||||
@@ -1532,11 +1592,18 @@ def create_agent(
|
||||
has_structured_output = initial_response_format is not None
|
||||
if awrap_model_call_handler is None:
|
||||
model_response = await _execute_model_async(request)
|
||||
return _build_commands(model_response, has_structured_output=has_structured_output)
|
||||
return _build_commands(
|
||||
model_response,
|
||||
has_structured_output=has_structured_output,
|
||||
repaired_messages=messages if messages != state["messages"] else None,
|
||||
)
|
||||
|
||||
result = await awrap_model_call_handler(request, _execute_model_async)
|
||||
return _build_commands(
|
||||
result.model_response, result.commands, has_structured_output=has_structured_output
|
||||
result.model_response,
|
||||
result.commands,
|
||||
has_structured_output=has_structured_output,
|
||||
repaired_messages=messages if messages != state["messages"] else None,
|
||||
)
|
||||
|
||||
# Use sync or async based on model capabilities
|
||||
@@ -2032,7 +2099,7 @@ def _make_tools_to_model_edge(
|
||||
return end_destination
|
||||
|
||||
# 3. Exit condition: A structured output tool was executed
|
||||
if any(t.name in structured_output_tools for t in tool_messages):
|
||||
if any(t.name in structured_output_tools and t.status != "error" for t in tool_messages):
|
||||
return end_destination
|
||||
|
||||
# 4. Default: Continue the loop
|
||||
|
||||
@@ -0,0 +1,251 @@
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
from langchain_core.callbacks import CallbackManagerForLLMRun
|
||||
from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, ToolMessage
|
||||
from langchain_core.outputs import ChatGeneration, ChatResult
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from langchain.agents import create_agent
|
||||
from langchain.agents.structured_output import ToolStrategy
|
||||
from langchain.tools import tool
|
||||
from tests.unit_tests.agents.model import FakeToolCallingModel
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
|
||||
|
||||
class InvalidToolCallingModel(FakeToolCallingModel):
|
||||
invalid_tool_call_id: str | None = "call_1"
|
||||
received_messages: list[BaseMessage] = Field(default_factory=list)
|
||||
|
||||
def _generate(
|
||||
self,
|
||||
messages: list[BaseMessage],
|
||||
stop: list[str] | None = None,
|
||||
run_manager: CallbackManagerForLLMRun | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ChatResult:
|
||||
_ = (stop, run_manager, kwargs)
|
||||
self.received_messages = messages
|
||||
message = AIMessage(
|
||||
content="",
|
||||
invalid_tool_calls=[
|
||||
{
|
||||
"name": "get_weather",
|
||||
"args": '{"city":',
|
||||
"id": self.invalid_tool_call_id,
|
||||
"error": "Invalid JSON",
|
||||
}
|
||||
],
|
||||
)
|
||||
self.index += 1
|
||||
return ChatResult(generations=[ChatGeneration(message=message)])
|
||||
|
||||
|
||||
@tool
|
||||
def get_weather(city: str = "Paris") -> str:
|
||||
"""Get the weather for a city."""
|
||||
return city
|
||||
|
||||
|
||||
class WeatherResponse(BaseModel):
|
||||
city: str
|
||||
|
||||
|
||||
class MixedToolCallingModel(FakeToolCallingModel):
|
||||
received_messages: list[BaseMessage] = Field(default_factory=list)
|
||||
|
||||
def _generate(
|
||||
self,
|
||||
messages: list[BaseMessage],
|
||||
stop: list[str] | None = None,
|
||||
run_manager: CallbackManagerForLLMRun | None = None,
|
||||
**kwargs: Any,
|
||||
) -> ChatResult:
|
||||
_ = (stop, run_manager, kwargs)
|
||||
self.received_messages = messages
|
||||
if self.index == 0:
|
||||
message = AIMessage(
|
||||
content="",
|
||||
tool_calls=[{"name": "get_weather", "args": {}, "id": "weather"}],
|
||||
invalid_tool_calls=[
|
||||
{
|
||||
"name": "WeatherResponse",
|
||||
"args": '{"city":',
|
||||
"id": "structured",
|
||||
"error": "Invalid JSON",
|
||||
}
|
||||
],
|
||||
)
|
||||
else:
|
||||
message = AIMessage(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"name": "WeatherResponse",
|
||||
"args": {"city": "Paris"},
|
||||
"id": "response",
|
||||
}
|
||||
],
|
||||
)
|
||||
self.index += 1
|
||||
return ChatResult(generations=[ChatGeneration(message=message)])
|
||||
|
||||
|
||||
def test_create_agent_does_not_patch_model_output() -> None:
|
||||
model = InvalidToolCallingModel()
|
||||
agent = create_agent(model, [get_weather])
|
||||
|
||||
result = agent.invoke({"messages": [HumanMessage("Weather?")]})
|
||||
|
||||
assert model.index == 1
|
||||
assert len(result["messages"]) == 2
|
||||
assert isinstance(result["messages"][-1], AIMessage)
|
||||
|
||||
|
||||
def test_create_agent_answers_invalid_tool_calls_on_next_turn() -> None:
|
||||
model = InvalidToolCallingModel()
|
||||
agent = create_agent(model, [get_weather], checkpointer=InMemorySaver())
|
||||
config: RunnableConfig = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
agent.invoke({"messages": [HumanMessage("Weather?")]}, config)
|
||||
model.invalid_tool_call_id = None
|
||||
result = agent.invoke({"messages": [HumanMessage("Try again")]}, config)
|
||||
|
||||
tool_message = model.received_messages[2]
|
||||
assert isinstance(tool_message, ToolMessage)
|
||||
assert tool_message.tool_call_id == "call_1"
|
||||
assert tool_message.name == "get_weather"
|
||||
assert tool_message.status == "error"
|
||||
assert "malformed or truncated" in tool_message.text
|
||||
assert [type(message) for message in result["messages"]] == [
|
||||
HumanMessage,
|
||||
AIMessage,
|
||||
ToolMessage,
|
||||
HumanMessage,
|
||||
AIMessage,
|
||||
]
|
||||
|
||||
|
||||
async def test_create_agent_answers_invalid_tool_calls_async() -> None:
|
||||
model = InvalidToolCallingModel()
|
||||
agent = create_agent(model, [get_weather], checkpointer=InMemorySaver())
|
||||
config: RunnableConfig = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
await agent.ainvoke({"messages": [HumanMessage("Weather?")]}, config)
|
||||
model.invalid_tool_call_id = None
|
||||
result = await agent.ainvoke({"messages": [HumanMessage("Try again")]}, config)
|
||||
|
||||
tool_message = model.received_messages[2]
|
||||
assert isinstance(tool_message, ToolMessage)
|
||||
assert tool_message.tool_call_id == "call_1"
|
||||
assert tool_message.status == "error"
|
||||
assert isinstance(result["messages"][2], ToolMessage)
|
||||
|
||||
|
||||
def test_create_agent_answers_historical_invalid_tool_calls() -> None:
|
||||
model = InvalidToolCallingModel(invalid_tool_call_id=None)
|
||||
agent = create_agent(model, [get_weather])
|
||||
invalid_message = AIMessage(
|
||||
content="",
|
||||
invalid_tool_calls=[
|
||||
{
|
||||
"name": "get_weather",
|
||||
"args": '{"city":',
|
||||
"id": "historical_call",
|
||||
"error": "Invalid JSON",
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
result = agent.invoke({"messages": [HumanMessage("Weather?"), invalid_message]})
|
||||
|
||||
historical_result = model.received_messages[-1]
|
||||
assert isinstance(historical_result, ToolMessage)
|
||||
assert historical_result.tool_call_id == "historical_call"
|
||||
assert (
|
||||
sum(
|
||||
isinstance(message, ToolMessage) and message.tool_call_id == "historical_call"
|
||||
for message in result["messages"]
|
||||
)
|
||||
== 1
|
||||
)
|
||||
|
||||
|
||||
def test_create_agent_does_not_duplicate_historical_tool_messages() -> None:
|
||||
model = InvalidToolCallingModel(invalid_tool_call_id=None)
|
||||
agent = create_agent(model, [get_weather])
|
||||
invalid_message = AIMessage(
|
||||
content="",
|
||||
invalid_tool_calls=[
|
||||
{
|
||||
"name": "get_weather",
|
||||
"args": '{"city":',
|
||||
"id": "answered_call",
|
||||
"error": "Invalid JSON",
|
||||
}
|
||||
],
|
||||
)
|
||||
existing_result = ToolMessage(
|
||||
content="Already answered",
|
||||
tool_call_id="answered_call",
|
||||
status="error",
|
||||
)
|
||||
|
||||
result = agent.invoke(
|
||||
{"messages": [HumanMessage("Weather?"), invalid_message, existing_result]}
|
||||
)
|
||||
|
||||
assert model.received_messages[-1] is existing_result
|
||||
assert (
|
||||
sum(
|
||||
isinstance(message, ToolMessage) and message.tool_call_id == "answered_call"
|
||||
for message in result["messages"]
|
||||
)
|
||||
== 1
|
||||
)
|
||||
|
||||
|
||||
def test_invalid_structured_tool_call_does_not_end_agent() -> None:
|
||||
model = MixedToolCallingModel()
|
||||
agent = create_agent(
|
||||
model,
|
||||
[get_weather],
|
||||
response_format=ToolStrategy(WeatherResponse),
|
||||
)
|
||||
|
||||
result = agent.invoke({"messages": [HumanMessage("Weather?")]})
|
||||
|
||||
assert model.index == 2
|
||||
assert result["structured_response"] == WeatherResponse(city="Paris")
|
||||
answered = {
|
||||
message.tool_call_id: message.status
|
||||
for message in model.received_messages
|
||||
if isinstance(message, ToolMessage)
|
||||
}
|
||||
assert answered == {"structured": "error", "weather": "success"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("tool_call_id", "patched"), [(None, False), ("", True)])
|
||||
def test_create_agent_patches_only_invalid_tool_calls_with_ids(
|
||||
tool_call_id: str | None, *, patched: bool
|
||||
) -> None:
|
||||
model = InvalidToolCallingModel(invalid_tool_call_id=None)
|
||||
agent = create_agent(model, [get_weather])
|
||||
invalid_message = AIMessage(
|
||||
content="",
|
||||
invalid_tool_calls=[
|
||||
{
|
||||
"name": "get_weather",
|
||||
"args": '{"city":',
|
||||
"id": tool_call_id,
|
||||
"error": "Invalid JSON",
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
agent.invoke({"messages": [HumanMessage("Weather?"), invalid_message]})
|
||||
|
||||
assert isinstance(model.received_messages[-1], ToolMessage) is patched
|
||||
Reference in new issue
Block a user