diff --git a/libs/prebuilt/langgraph/prebuilt/tool_node.py b/libs/prebuilt/langgraph/prebuilt/tool_node.py index 12f751d6f..db45f8f46 100644 --- a/libs/prebuilt/langgraph/prebuilt/tool_node.py +++ b/libs/prebuilt/langgraph/prebuilt/tool_node.py @@ -52,6 +52,7 @@ from typing import ( from langchain_core.messages import ( AIMessage, AnyMessage, + RemoveMessage, ToolCall, ToolMessage, convert_to_messages, @@ -72,6 +73,7 @@ from typing_extensions import Annotated, get_args, get_origin from langgraph._internal._runnable import RunnableCallable from langgraph.errors import GraphBubbleUp +from langgraph.graph.message import REMOVE_ALL_MESSAGES from langgraph.prebuilt._internal import ToolCallWithContext from langgraph.store.base import BaseStore from langgraph.types import Command, Send @@ -754,6 +756,11 @@ class ToolNode(RunnableCallable): # convert to message objects if updates are in a dict format messages_update = convert_to_messages(messages_update) + + # no validation needed if all messages are being removed + if messages_update == [RemoveMessage(id=REMOVE_ALL_MESSAGES)]: + return updated_command + has_matching_tool_message = False for message in messages_update: if not isinstance(message, ToolMessage): diff --git a/libs/prebuilt/tests/test_tool_node.py b/libs/prebuilt/tests/test_tool_node.py index ec7c91fdd..2b6dfbebe 100644 --- a/libs/prebuilt/tests/test_tool_node.py +++ b/libs/prebuilt/tests/test_tool_node.py @@ -7,6 +7,7 @@ from typing import ( import pytest from langchain_core.messages import ( AIMessage, + RemoveMessage, ToolMessage, ) from langchain_core.tools import BaseTool, ToolException @@ -15,6 +16,7 @@ from pydantic import BaseModel, ValidationError from pydantic.v1 import ValidationError as ValidationErrorV1 from langgraph.errors import GraphBubbleUp, GraphInterrupt +from langgraph.graph.message import REMOVE_ALL_MESSAGES from langgraph.prebuilt import ToolNode from langgraph.prebuilt.tool_node import TOOL_CALL_ERROR_TEMPLATE from langgraph.types import Command, Send @@ -1129,3 +1131,28 @@ def test_tool_node_parent_command_with_send(): graph=Command.PARENT, ) ] + + +async def test_tool_node_command_remove_all_messages(): + from langchain_core.tools.base import InjectedToolCallId + + @dec_tool + def remove_all_messages_tool(tool_call_id: Annotated[str, InjectedToolCallId]): + """A tool that removes all messages.""" + return Command(update={"messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES)]}) + + tool_node = ToolNode([remove_all_messages_tool]) + tool_call = { + "name": "remove_all_messages_tool", + "args": {}, + "id": "tool_call_123", + } + result = await tool_node.ainvoke( + {"messages": [AIMessage(content="", tool_calls=[tool_call])]} + ) + + assert isinstance(result, list) + assert len(result) == 1 + command = result[0] + assert isinstance(command, Command) + assert command.update == {"messages": [RemoveMessage(id=REMOVE_ALL_MESSAGES)]}