diff --git a/libs/langchain_v1/langchain/agents/factory.py b/libs/langchain_v1/langchain/agents/factory.py index 0810bdc4b4..529344673c 100644 --- a/libs/langchain_v1/langchain/agents/factory.py +++ b/libs/langchain_v1/langchain/agents/factory.py @@ -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 diff --git a/libs/langchain_v1/tests/unit_tests/agents/test_invalid_tool_calls.py b/libs/langchain_v1/tests/unit_tests/agents/test_invalid_tool_calls.py new file mode 100644 index 0000000000..99304779be --- /dev/null +++ b/libs/langchain_v1/tests/unit_tests/agents/test_invalid_tool_calls.py @@ -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