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:
ccurmeandopen-swe[bot] authored and GitHub committed 2026-09-22 15:22:59 -04:00
1 parent fbac0b2d0c
commit bfa03c8cf8
2 files changed
+327 -9

No files matched your search

+76 -9
View File
@@ -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