diff --git a/langgraph/graph/message.py b/langgraph/graph/message.py index d41a2e747..ad537b5c0 100644 --- a/langgraph/graph/message.py +++ b/langgraph/graph/message.py @@ -1,3 +1,4 @@ +import uuid from typing import Annotated, Union from langchain_core.messages import AnyMessage @@ -8,11 +9,27 @@ Messages = Union[list[AnyMessage], AnyMessage] def add_messages(left: Messages, right: Messages) -> Messages: + # coerce to list if not isinstance(left, list): left = [left] if not isinstance(right, list): right = [right] - return left + right + # assign missing ids + for m in left: + if m.id is None: + m.id = str(uuid.uuid4()) + for m in right: + if m.id is None: + m.id = str(uuid.uuid4()) + # merge + left_idx_by_id = {m.id: i for i, m in enumerate(left)} + merged = left.copy() + for m in right: + if (existing_idx := left_idx_by_id.get(m.id)) is not None: + merged[existing_idx] = m + else: + merged.append(m) + return merged class MessageGraph(StateGraph): diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 000000000..9b6b04c01 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,12 @@ +from uuid import UUID + +import pytest +from pytest_mock import MockerFixture + + +@pytest.fixture() +def deterministic_uuids(mocker: MockerFixture) -> MockerFixture: + side_effect = ( + UUID(f"00000000-0000-4000-8000-{i:012}", version=4) for i in range(10000) + ) + return mocker.patch("uuid.uuid4", side_effect=side_effect) diff --git a/tests/test_pregel.py b/tests/test_pregel.py index 9558dbe60..37a57d90b 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -2321,7 +2321,9 @@ def test_prebuilt_chat(snapshot: SnapshotAssertion) -> None: ] -def test_message_graph(snapshot: SnapshotAssertion) -> None: +def test_message_graph( + snapshot: SnapshotAssertion, deterministic_uuids: MockerFixture +) -> None: from langchain.chat_models.fake import FakeMessagesListChatModel from langchain_community.tools import tool from langchain_core.agents import AgentAction @@ -2348,6 +2350,7 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None: "arguments": json.dumps("query"), } }, + id="ai1", ), AIMessage( content="", @@ -2357,8 +2360,9 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None: "arguments": json.dumps("another"), } }, + id="ai2", ), - AIMessage(content="answer"), + AIMessage(content="answer", id="ai3"), ] ) @@ -2438,22 +2442,35 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None: assert app.get_graph().draw_ascii() == snapshot assert app.invoke(HumanMessage(content="what is weather in sf")) == [ - HumanMessage(content="what is weather in sf"), + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000002", # adds missing ids + ), AIMessage( content="", additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"query"'} }, + id="ai1", # respects ids passed in + ), + FunctionMessage( + content="result for query", + name="search_api", + id="00000000-0000-4000-8000-000000000014", ), - FunctionMessage(content="result for query", name="search_api"), AIMessage( content="", additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"another"'} }, + id="ai2", ), - FunctionMessage(content="result for another", name="search_api"), - AIMessage(content="answer"), + FunctionMessage( + content="result for another", + name="search_api", + id="00000000-0000-4000-8000-000000000026", + ), + AIMessage(content="answer", id="ai3"), ] assert [*app.stream([HumanMessage(content="what is weather in sf")])] == [ @@ -2463,29 +2480,51 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None: additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"query"'} }, + id="ai1", + ) + }, + { + "action": FunctionMessage( + content="result for query", + name="search_api", + id="00000000-0000-4000-8000-000000000047", ) }, - {"action": FunctionMessage(content="result for query", name="search_api")}, { "agent": AIMessage( content="", additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"another"'} }, + id="ai2", ) }, - {"action": FunctionMessage(content="result for another", name="search_api")}, - {"agent": AIMessage(content="answer")}, + { + "action": FunctionMessage( + content="result for another", + name="search_api", + id="00000000-0000-4000-8000-000000000059", + ) + }, + {"agent": AIMessage(content="answer", id="ai3")}, { "__end__": [ - HumanMessage(content="what is weather in sf"), + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000035", + ), AIMessage( content="", additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"query"'} }, + id="ai1", + ), + FunctionMessage( + content="result for query", + name="search_api", + id="00000000-0000-4000-8000-000000000047", ), - FunctionMessage(content="result for query", name="search_api"), AIMessage( content="", additional_kwargs={ @@ -2494,9 +2533,14 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None: "arguments": '"another"', } }, + id="ai2", ), - FunctionMessage(content="result for another", name="search_api"), - AIMessage(content="answer"), + FunctionMessage( + content="result for another", + name="search_api", + id="00000000-0000-4000-8000-000000000059", + ), + AIMessage(content="answer", id="ai3"), ] }, ] @@ -2518,25 +2562,167 @@ def test_message_graph(snapshot: SnapshotAssertion) -> None: additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"query"'} }, + id="ai1", ) } ] assert app_w_interrupt.get_state(config) == StateSnapshot( values=[ - HumanMessage(content="what is weather in sf"), + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000068", + ), AIMessage( content="", additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"query"'} }, + id="ai1", ), ], next=("agent:edges",), config=app_w_interrupt.checkpointer.get_tuple(config).config, ) - # TODO use update_state once we have message ids + # modify ai message + last_message = app_w_interrupt.get_state(config).values[-1] + last_message.additional_kwargs["function_call"]["arguments"] = '"a different query"' + app_w_interrupt.update_state(config, last_message) + + # message was replaced instead of appended + assert app_w_interrupt.get_state(config) == StateSnapshot( + values=[ + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000068", + ), + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"a different query"', + } + }, + id="ai1", + ), + ], + next=("agent:edges",), + config=app_w_interrupt.checkpointer.get_tuple(config).config, + ) + + assert [c for c in app_w_interrupt.stream(None, config)] == [ + { + "action": FunctionMessage( + content="result for a different query", + name="search_api", + id="00000000-0000-4000-8000-000000000081", + ) + }, + { + "agent": AIMessage( + content="", + additional_kwargs={ + "function_call": {"name": "search_api", "arguments": '"another"'} + }, + id="ai2", + ) + }, + ] + + assert app_w_interrupt.get_state(config) == StateSnapshot( + values=[ + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000068", + ), + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"a different query"', + } + }, + id="ai1", + ), + FunctionMessage( + content="result for a different query", + name="search_api", + id="00000000-0000-4000-8000-000000000081", + ), + AIMessage( + content="", + additional_kwargs={ + "function_call": {"name": "search_api", "arguments": '"another"'} + }, + id="ai2", + ), + ], + next=("agent:edges",), + config=app_w_interrupt.checkpointer.get_tuple(config).config, + ) + + app_w_interrupt.update_state( + config, + AIMessage(content="answer", id="ai2"), + ) + + # replaces message even if object identity is different, as long as id is the same + assert app_w_interrupt.get_state(config) == StateSnapshot( + values=[ + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000068", + ), + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"a different query"', + } + }, + id="ai1", + ), + FunctionMessage( + content="result for a different query", + name="search_api", + id="00000000-0000-4000-8000-000000000081", + ), + AIMessage(content="answer", id="ai2"), + ], + next=("agent:edges",), + config=app_w_interrupt.checkpointer.get_tuple(config).config, + ) + + assert [c for c in app_w_interrupt.stream(None, config)] == [ + { + "__end__": [ + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000068", + ), + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"a different query"', + } + }, + id="ai1", + ), + FunctionMessage( + content="result for a different query", + name="search_api", + id="00000000-0000-4000-8000-000000000081", + ), + AIMessage(content="answer", id="ai2"), + ] + } + ] def test_in_one_fan_out_out_one_graph_state() -> None: diff --git a/tests/test_pregel_async.py b/tests/test_pregel_async.py index 0fbb65994..3ebb79910 100644 --- a/tests/test_pregel_async.py +++ b/tests/test_pregel_async.py @@ -2356,7 +2356,7 @@ async def test_prebuilt_chat() -> None: ] -async def test_message_graph() -> None: +async def test_message_graph(deterministic_uuids: MockerFixture) -> None: from langchain.chat_models.fake import FakeMessagesListChatModel from langchain_community.tools import tool from langchain_core.agents import AgentAction @@ -2383,6 +2383,7 @@ async def test_message_graph() -> None: "arguments": json.dumps("query"), } }, + id="ai1", ), AIMessage( content="", @@ -2392,8 +2393,9 @@ async def test_message_graph() -> None: "arguments": json.dumps("another"), } }, + id="ai2", ), - AIMessage(content="answer"), + AIMessage(content="answer", id="ai3"), ] ) @@ -2468,22 +2470,34 @@ async def test_message_graph() -> None: app = workflow.compile() assert await app.ainvoke(HumanMessage(content="what is weather in sf")) == [ - HumanMessage(content="what is weather in sf"), + HumanMessage( + content="what is weather in sf", id="00000000-0000-4000-8000-000000000002" + ), AIMessage( content="", additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"query"'} }, + id="ai1", + ), + FunctionMessage( + content="result for query", + name="search_api", + id="00000000-0000-4000-8000-000000000014", ), - FunctionMessage(content="result for query", name="search_api"), AIMessage( content="", additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"another"'} }, + id="ai2", ), - FunctionMessage(content="result for another", name="search_api"), - AIMessage(content="answer"), + FunctionMessage( + content="result for another", + name="search_api", + id="00000000-0000-4000-8000-000000000026", + ), + AIMessage(content="answer", id="ai3"), ] assert [ @@ -2495,29 +2509,51 @@ async def test_message_graph() -> None: additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"query"'} }, + id="ai1", + ) + }, + { + "action": FunctionMessage( + content="result for query", + name="search_api", + id="00000000-0000-4000-8000-000000000047", ) }, - {"action": FunctionMessage(content="result for query", name="search_api")}, { "agent": AIMessage( content="", additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"another"'} }, + id="ai2", ) }, - {"action": FunctionMessage(content="result for another", name="search_api")}, - {"agent": AIMessage(content="answer")}, + { + "action": FunctionMessage( + content="result for another", + name="search_api", + id="00000000-0000-4000-8000-000000000059", + ) + }, + {"agent": AIMessage(content="answer", id="ai3")}, { "__end__": [ - HumanMessage(content="what is weather in sf"), + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000035", + ), AIMessage( content="", additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"query"'} }, + id="ai1", + ), + FunctionMessage( + content="result for query", + name="search_api", + id="00000000-0000-4000-8000-000000000047", ), - FunctionMessage(content="result for query", name="search_api"), AIMessage( content="", additional_kwargs={ @@ -2526,9 +2562,14 @@ async def test_message_graph() -> None: "arguments": '"another"', } }, + id="ai2", ), - FunctionMessage(content="result for another", name="search_api"), - AIMessage(content="answer"), + FunctionMessage( + content="result for another", + name="search_api", + id="00000000-0000-4000-8000-000000000059", + ), + AIMessage(content="answer", id="ai3"), ] }, ] @@ -2550,25 +2591,167 @@ async def test_message_graph() -> None: additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"query"'} }, + id="ai1", ) } ] assert await app_w_interrupt.aget_state(config) == StateSnapshot( values=[ - HumanMessage(content="what is weather in sf"), + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000068", + ), AIMessage( content="", additional_kwargs={ "function_call": {"name": "search_api", "arguments": '"query"'} }, + id="ai1", ), ], next=("agent:edges",), config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config, ) - # TODO use update_state once we have message ids + # modify ai message + last_message = (await app_w_interrupt.aget_state(config)).values[-1] + last_message.additional_kwargs["function_call"]["arguments"] = '"a different query"' + await app_w_interrupt.aupdate_state(config, last_message) + + # message was replaced instead of appended + assert await app_w_interrupt.aget_state(config) == StateSnapshot( + values=[ + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000068", + ), + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"a different query"', + } + }, + id="ai1", + ), + ], + next=("agent:edges",), + config=app_w_interrupt.checkpointer.get_tuple(config).config, + ) + + assert [c async for c in app_w_interrupt.astream(None, config)] == [ + { + "action": FunctionMessage( + content="result for a different query", + name="search_api", + id="00000000-0000-4000-8000-000000000081", + ) + }, + { + "agent": AIMessage( + content="", + additional_kwargs={ + "function_call": {"name": "search_api", "arguments": '"another"'} + }, + id="ai2", + ) + }, + ] + + assert await app_w_interrupt.aget_state(config) == StateSnapshot( + values=[ + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000068", + ), + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"a different query"', + } + }, + id="ai1", + ), + FunctionMessage( + content="result for a different query", + name="search_api", + id="00000000-0000-4000-8000-000000000081", + ), + AIMessage( + content="", + additional_kwargs={ + "function_call": {"name": "search_api", "arguments": '"another"'} + }, + id="ai2", + ), + ], + next=("agent:edges",), + config=app_w_interrupt.checkpointer.get_tuple(config).config, + ) + + await app_w_interrupt.aupdate_state( + config, + AIMessage(content="answer", id="ai2"), + ) + + # replaces message even if object identity is different, as long as id is the same + assert await app_w_interrupt.aget_state(config) == StateSnapshot( + values=[ + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000068", + ), + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"a different query"', + } + }, + id="ai1", + ), + FunctionMessage( + content="result for a different query", + name="search_api", + id="00000000-0000-4000-8000-000000000081", + ), + AIMessage(content="answer", id="ai2"), + ], + next=("agent:edges",), + config=app_w_interrupt.checkpointer.get_tuple(config).config, + ) + + assert [c async for c in app_w_interrupt.astream(None, config)] == [ + { + "__end__": [ + HumanMessage( + content="what is weather in sf", + id="00000000-0000-4000-8000-000000000068", + ), + AIMessage( + content="", + additional_kwargs={ + "function_call": { + "name": "search_api", + "arguments": '"a different query"', + } + }, + id="ai1", + ), + FunctionMessage( + content="result for a different query", + name="search_api", + id="00000000-0000-4000-8000-000000000081", + ), + AIMessage(content="answer", id="ai2"), + ] + } + ] async def test_in_one_fan_out_out_one_graph_state() -> None: