From 68f7a3acc4db7b120b395b1f3ac2af630fb59f03 Mon Sep 17 00:00:00 2001 From: Sydney Runkle Date: Wed, 29 Apr 2026 14:11:17 -0400 Subject: [PATCH] test: add add_messages migration test; move all imports to module level MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Add test_add_messages_to_delta_migration_preserves_message_history (sync + async) covering the primary real-world BinaryOperatorAggregate → DeltaChannel migration path with real Message objects and IDs - Hoist all in-function imports to module level in test_channels.py and fix _delta_channel_with_type helper accordingly - Add section headers in test_channels.py for better navigation --- libs/langgraph/tests/test_channels.py | 146 ++++-------------- .../tests/test_delta_channel_migration.py | 102 ++++++++++++ 2 files changed, 134 insertions(+), 114 deletions(-) diff --git a/libs/langgraph/tests/test_channels.py b/libs/langgraph/tests/test_channels.py index 2f1cf6d15..f967cbdbd 100644 --- a/libs/langgraph/tests/test_channels.py +++ b/libs/langgraph/tests/test_channels.py @@ -1,9 +1,13 @@ import operator from collections.abc import Sequence +from typing import Annotated import pytest -from langchain_core.messages import AIMessage, HumanMessage +from langchain_core.messages import AIMessage, HumanMessage, RemoveMessage from langgraph.checkpoint.base import DELTA_SENTINEL +from langgraph.checkpoint.memory import InMemorySaver +from langgraph.checkpoint.serde.types import _DeltaSnapshot +from typing_extensions import NotRequired, TypedDict from langgraph._internal._typing import MISSING from langgraph.channels.binop import BinaryOperatorAggregate @@ -12,11 +16,19 @@ from langgraph.channels.last_value import LastValue from langgraph.channels.topic import Topic from langgraph.channels.untracked_value import UntrackedValue from langgraph.errors import EmptyChannelError, InvalidUpdateError +from langgraph.graph import START, StateGraph from langgraph.graph.message import _messages_delta_reducer +from langgraph.graph.state import _get_channel +from langgraph.types import Overwrite pytestmark = pytest.mark.anyio +# --------------------------------------------------------------------------- +# Core channel primitives +# --------------------------------------------------------------------------- + + def test_last_value() -> None: channel = LastValue(int).from_checkpoint(MISSING) assert channel.ValueType is int @@ -99,49 +111,41 @@ def test_untracked_value() -> None: assert channel.ValueType is dict assert channel.UpdateType is dict - # UntrackedValue should start empty with pytest.raises(EmptyChannelError): channel.get() - # Should be able to update with a value test_data = {"session": "test", "temp": "dir"} channel.update([test_data]) assert channel.get() == test_data - # Update with new value new_data = {"session": "updated", "temp": "newdir"} channel.update([new_data]) assert channel.get() == new_data - # On checkpoint, UntrackedValue should return MISSING checkpoint = channel.checkpoint() assert checkpoint is MISSING - # Creating from checkpoint with MISSING should start empty new_channel = UntrackedValue(dict).from_checkpoint(checkpoint) with pytest.raises(EmptyChannelError): new_channel.get() +# --------------------------------------------------------------------------- +# DeltaChannel — message reducer +# --------------------------------------------------------------------------- + + def test_delta_channel_basic_two_steps() -> None: - from langchain_core.messages import AIMessage, HumanMessage - from langgraph.checkpoint.base import DELTA_SENTINEL - - from langgraph.graph.message import _messages_delta_reducer - ch = DeltaChannel(_messages_delta_reducer, list).from_checkpoint(MISSING) - # Step 1: one message added ch.update([HumanMessage(content="hi", id="h1")]) d1 = ch.checkpoint() assert d1 is DELTA_SENTINEL - # Step 2: another message ch.update([AIMessage(content="hello", id="a1")]) d2 = ch.checkpoint() assert d2 is DELTA_SENTINEL - # Full accumulated value is preserved in memory assert len(ch.get()) == 2 assert ch.get()[0].content == "hi" assert ch.get()[1].content == "hello" @@ -149,10 +153,6 @@ def test_delta_channel_basic_two_steps() -> None: def test_delta_channel_from_checkpoint_writes_list() -> None: """replay_writes on a fresh channel replays through the operator.""" - from langchain_core.messages import AIMessage, HumanMessage - - from langgraph.graph.message import _messages_delta_reducer - spec = DeltaChannel(_messages_delta_reducer, list) ch = spec.from_checkpoint(DELTA_SENTINEL) ch.replay_writes( @@ -170,11 +170,6 @@ def test_delta_channel_from_checkpoint_writes_list() -> None: def test_delta_channel_from_checkpoint_backwards_compat() -> None: - from langchain_core.messages import HumanMessage - - from langgraph.graph.message import _messages_delta_reducer - - # Old BinaryOperatorAggregate checkpoint: plain list treated as backward compat spec = DeltaChannel(_messages_delta_reducer, list) old_value = [HumanMessage(content="old", id="h1")] ch = spec.from_checkpoint(old_value) @@ -182,33 +177,21 @@ def test_delta_channel_from_checkpoint_backwards_compat() -> None: def test_delta_channel_overwrite() -> None: - from langchain_core.messages import HumanMessage - from langgraph.checkpoint.base import DELTA_SENTINEL - - from langgraph.graph.message import _messages_delta_reducer - from langgraph.types import Overwrite - ch = DeltaChannel(_messages_delta_reducer, list).from_checkpoint(MISSING) ch.update([HumanMessage(content="old", id="h1")]) ch.update([Overwrite([HumanMessage(content="new", id="h2")])]) d = ch.checkpoint() assert d is DELTA_SENTINEL - # After overwrite, value is reset to only the new message assert len(ch.get()) == 1 assert ch.get()[0].content == "new" def test_delta_channel_remove_message_and_replay() -> None: """RemoveMessage must round-trip correctly when writes are replayed.""" - from langchain_core.messages import AIMessage, HumanMessage, RemoveMessage - - from langgraph.graph.message import _messages_delta_reducer - spec = DeltaChannel(_messages_delta_reducer, list) ch = spec.from_checkpoint(MISSING) - # Step 1: add two messages ch.update([HumanMessage(content="hi", id="h1")]) ch.update([AIMessage(content="hello", id="a1")]) assert ch.get() == [ @@ -216,11 +199,9 @@ def test_delta_channel_remove_message_and_replay() -> None: AIMessage(content="hello", id="a1"), ] - # Step 2: remove the AI message ch.update([RemoveMessage(id="a1")]) assert ch.get() == [HumanMessage(content="hi", id="h1")] - # Replay the writes list from scratch — must reproduce the post-remove state ch2 = spec.from_checkpoint(DELTA_SENTINEL) ch2.replay_writes( [ @@ -234,21 +215,13 @@ def test_delta_channel_remove_message_and_replay() -> None: def test_delta_channel_update_by_id_and_replay() -> None: """Updating a message by ID must round-trip correctly through writes replay.""" - from langchain_core.messages import HumanMessage - - from langgraph.graph.message import _messages_delta_reducer - spec = DeltaChannel(_messages_delta_reducer, list) ch = spec.from_checkpoint(MISSING) - # Step 1: add a message ch.update([HumanMessage(content="original", id="h1")]) - - # Step 2: update the same message by ID ch.update([HumanMessage(content="updated", id="h1")]) assert ch.get() == [HumanMessage(content="updated", id="h1")] - # Replay writes — must produce the updated message, not the original ch2 = spec.from_checkpoint(DELTA_SENTINEL) ch2.replay_writes( [ @@ -262,19 +235,18 @@ def test_delta_channel_update_by_id_and_replay() -> None: def test_delta_channel_checkpoint_returns_sentinel() -> None: """checkpoint() always returns DELTA_SENTINEL regardless of state.""" - from langgraph.checkpoint.base import DELTA_SENTINEL - - from langgraph.graph.message import _messages_delta_reducer - ch = DeltaChannel(_messages_delta_reducer, list).from_checkpoint(MISSING) assert ch.checkpoint() is DELTA_SENTINEL - from langchain_core.messages import HumanMessage - ch.update([HumanMessage(content="hi", id="h1")]) assert ch.checkpoint() is DELTA_SENTINEL +# --------------------------------------------------------------------------- +# DeltaChannel — snapshot frequency +# --------------------------------------------------------------------------- + + def test_delta_channel_snapshot_step_based() -> None: """Snapshots fire on every Nth step regardless of whether the channel was written. @@ -282,17 +254,7 @@ def test_delta_channel_snapshot_step_based() -> None: blob — even if the channel had no write that step (eager snapshot). This bounds the ancestor walk to at most N steps on any read. """ - from typing import Annotated - from langchain_core.messages import AIMessage, HumanMessage - from langgraph.checkpoint.memory import InMemorySaver - from langgraph.checkpoint.serde.types import _DeltaSnapshot - from typing_extensions import TypedDict - - from langgraph.graph import START, StateGraph - from langgraph.graph.message import _messages_delta_reducer - - # snapshot_frequency=5: snapshot every 5 pregel steps class State(TypedDict): messages: Annotated[ list, DeltaChannel(_messages_delta_reducer, snapshot_frequency=5) @@ -300,12 +262,10 @@ def test_delta_channel_snapshot_step_based() -> None: other: str def node_a(state: State) -> dict: - # writes to messages i = len(state["messages"]) // 2 return {"messages": [AIMessage(content=f"a{i}", id=f"a{i}")]} def node_b(state: State) -> dict: - # writes ONLY to other, not messages — snapshot must still fire at step N return {"other": "y"} g = StateGraph(State) @@ -323,7 +283,6 @@ def test_delta_channel_snapshot_step_based() -> None: config, ) - # Confirm at least one snapshot blob exists for messages msg_blob_values = [ saver.serde.loads_typed((type_tag, blob)) for k, (type_tag, blob) in saver.blobs.items() @@ -332,7 +291,6 @@ def test_delta_channel_snapshot_step_based() -> None: snapshots = [v for v in msg_blob_values if isinstance(v, _DeltaSnapshot)] assert snapshots, "expected at least one _DeltaSnapshot blob for messages" - # Final state must be correct regardless of snapshot cadence state = graph.get_state(config) assert len(state.values["messages"]) == 12 # 6 human + 6 AI @@ -341,15 +299,6 @@ def test_delta_channel_snapshot_fires_even_when_not_written() -> None: """Eager snapshot: _DeltaSnapshot stored at snapshot step even when the channel had no write that step (node_b doesn't touch messages). """ - from typing import Annotated - - from langchain_core.messages import AIMessage, HumanMessage - from langgraph.checkpoint.memory import InMemorySaver - from langgraph.checkpoint.serde.types import _DeltaSnapshot - from typing_extensions import TypedDict - - from langgraph.graph import START, StateGraph - from langgraph.graph.message import _messages_delta_reducer class State(TypedDict): messages: Annotated[ @@ -362,7 +311,6 @@ def test_delta_channel_snapshot_fires_even_when_not_written() -> None: return {"messages": [AIMessage(content=f"a{i}", id=f"a{i}")]} def ticker(state: State) -> dict: - # never writes messages return {"tick": state["tick"] + 1} g = StateGraph(State) @@ -380,33 +328,27 @@ def test_delta_channel_snapshot_fires_even_when_not_written() -> None: config, ) - # Count distinct message channel blob versions msg_blobs = { k: saver.serde.loads_typed((t, b)) for k, (t, b) in saver.blobs.items() if k[2] == "messages" and t == "msgpack" and b } snapshots = {k: v for k, v in msg_blobs.items() if isinstance(v, _DeltaSnapshot)} - # There must be snapshots (ticker steps are snapshot steps too) assert snapshots, ( "eager snapshots must fire even on steps where messages wasn't written" ) - # All get_state calls must return the correct accumulated value state = graph.get_state(config) assert len(state.values["messages"]) == 10 # 5 human + 5 AI +# --------------------------------------------------------------------------- +# DeltaChannel — end-to-end (InMemorySaver) +# --------------------------------------------------------------------------- + + def test_delta_channel_inmemory_saver_assembles_writes() -> None: """InMemorySaver assembles writes from checkpoint_writes inside get_tuple.""" - from typing import Annotated - - from langchain_core.messages import AIMessage, HumanMessage - from langgraph.checkpoint.memory import InMemorySaver - from typing_extensions import TypedDict - - from langgraph.graph import START, StateGraph - from langgraph.graph.message import _messages_delta_reducer class State(TypedDict): messages: Annotated[list, DeltaChannel(_messages_delta_reducer, list)] @@ -427,9 +369,6 @@ def test_delta_channel_inmemory_saver_assembles_writes() -> None: graph.invoke({"messages": [HumanMessage(content="hi", id="h1")]}, config) graph.invoke({"messages": [HumanMessage(content="bye", id="h2")]}, config) - # get_tuple returns raw storage shape — channel_values stores DELTA_SENTINEL - # for delta channels; the reconstructed writes flow separately via - # saver._get_channel_writes_history. saved = saver.get_tuple(config) assert saved is not None assert "messages" in saved.checkpoint["channel_values"] @@ -440,18 +379,13 @@ def test_delta_channel_inmemory_saver_assembles_writes() -> None: # --------------------------------------------------------------------------- -# Dict-reducer tests +# DeltaChannel — dict reducer # --------------------------------------------------------------------------- -def _delta_channel_with_type(operator, typ): +def _delta_channel_with_type(op, typ): """Build a DeltaChannel with an explicit type via the Annotated injection path.""" - from typing import Annotated - - from langgraph.channels.delta import DeltaChannel - from langgraph.graph.state import _get_channel - - return _get_channel("_test", Annotated[typ, DeltaChannel(operator)]) + return _get_channel("_test", Annotated[typ, DeltaChannel(op)]) def test_delta_channel_dict_reducer_fresh_channel() -> None: @@ -542,7 +476,6 @@ def test_delta_channel_dict_reducer_with_deletions() -> None: def test_delta_channel_dict_reducer_overwrite_in_update() -> None: """Overwrite(dict) in update() must preserve dict shape, not coerce to list.""" - from langgraph.types import Overwrite def merge_dicts(state: dict, writes: list) -> dict: result = dict(state) @@ -558,7 +491,6 @@ def test_delta_channel_dict_reducer_overwrite_in_update() -> None: def test_delta_channel_dict_reducer_overwrite_in_writes_replay() -> None: """Overwrite(dict) embedded in replayed writes must reconstruct as dict.""" - from langgraph.types import Overwrite def merge_dicts(state: dict, writes: list) -> dict: result = dict(state) @@ -580,12 +512,6 @@ def test_delta_channel_dict_reducer_overwrite_in_writes_replay() -> None: def test_delta_channel_dict_reducer_with_notrequired_annotation() -> None: """DeltaChannel infers dict type through `Annotated[NotRequired[dict[...]], ch]`.""" - from typing import Annotated - - from typing_extensions import NotRequired - - from langgraph.channels.delta import DeltaChannel - from langgraph.graph.state import _get_channel def merge_dicts(state: dict, writes: list) -> dict: result = dict(state) @@ -603,13 +529,6 @@ def test_delta_channel_dict_reducer_with_notrequired_annotation() -> None: def test_delta_channel_dict_reducer_end_to_end_filesystem() -> None: """End-to-end: graph with dict-reducer (filesystem-style) channel wrapped in DeltaChannel.""" - from typing import Annotated - - from langgraph.checkpoint.memory import InMemorySaver - from typing_extensions import TypedDict - - from langgraph.channels.delta import DeltaChannel - from langgraph.graph import START, StateGraph def merge_files(state: dict, writes: list) -> dict: result = dict(state) @@ -684,7 +603,7 @@ def test_delta_channel_dict_reducer_backwards_compat() -> None: # --------------------------------------------------------------------------- -# seed / pre-delta migration +# DeltaChannel — seed / pre-delta migration # --------------------------------------------------------------------------- @@ -731,5 +650,4 @@ def test_delta_channel_from_checkpoint_seed_none_is_distinct_from_sentinel() -> spec = DeltaChannel(replace, list) ch = spec.from_checkpoint(None) ch.replay_writes([("t0", "x", "after")]) - # Reducer replaces; seed=None → first write produces "after". assert ch.get() == "after" diff --git a/libs/langgraph/tests/test_delta_channel_migration.py b/libs/langgraph/tests/test_delta_channel_migration.py index 9793fd831..46fdc0658 100644 --- a/libs/langgraph/tests/test_delta_channel_migration.py +++ b/libs/langgraph/tests/test_delta_channel_migration.py @@ -46,12 +46,14 @@ import operator from typing import Annotated, Any import pytest +from langchain_core.messages import AIMessage, HumanMessage from langgraph.checkpoint.memory import InMemorySaver from typing_extensions import TypedDict from langgraph.channels.binop import BinaryOperatorAggregate from langgraph.channels.delta import DeltaChannel from langgraph.graph import END, START, StateGraph +from langgraph.graph.message import _messages_delta_reducer, add_messages pytestmark = pytest.mark.anyio @@ -509,3 +511,103 @@ def test_fork_from_update_state_checkpoint() -> None: f"fork lost update_state base: base={base_items}, forked={forked_items}" ) assert forked_items[-1] == "fork0", f"fork delta not appended: {forked_items}" + + +# --------------------------------------------------------------------------- +# 9. Migration from `add_messages` → `DeltaChannel(_messages_delta_reducer)` +# +# `add_messages` is the primary real-world use case: it creates a +# BinaryOperatorAggregate with dedup-by-ID and RemoveMessage semantics. +# After swapping the annotation to DeltaChannel, pre-migration blobs +# (plain lists of Message objects) must be used directly as the seed. +# --------------------------------------------------------------------------- + + +def _add_messages_graph(checkpointer: Any) -> Any: + class MessagesState(TypedDict): + messages: Annotated[list, add_messages] + + return ( + StateGraph(MessagesState) + .add_node("noop", _noop) + .add_edge(START, "noop") + .add_edge("noop", END) + .compile(checkpointer=checkpointer) + ) + + +def _delta_messages_graph(checkpointer: Any) -> Any: + class DeltaMessagesState(TypedDict): + messages: Annotated[list, DeltaChannel(_messages_delta_reducer)] + + return ( + StateGraph(DeltaMessagesState) + .add_node("noop", _noop) + .add_edge(START, "noop") + .add_edge("noop", END) + .compile(checkpointer=checkpointer) + ) + + +def test_add_messages_to_delta_migration_preserves_message_history() -> None: + """Migration from `add_messages` to `DeltaChannel(_messages_delta_reducer)` + preserves message ordering and IDs at both the tip and settled ancestor + boundaries. + + The pre-migration blob is a plain list of Message objects; DeltaChannel + must use it directly as the seed without walking ancestors past it. + """ + checkpointer = InMemorySaver() + config = {"configurable": {"thread_id": "add-messages-migration"}} + + pre_graph = _add_messages_graph(checkpointer) + pre_graph.invoke({"messages": [HumanMessage(content="hello", id="h1")]}, config) + pre_graph.invoke({"messages": [AIMessage(content="hi", id="a1")]}, config) + pre_graph.invoke({"messages": [HumanMessage(content="thanks", id="h2")]}, config) + + pre_tip = pre_graph.get_state(config) + assert [m.id for m in pre_tip.values["messages"]] == ["h1", "a1", "h2"] + + delta_graph = _delta_messages_graph(checkpointer) + + # Tip: latest checkpoint has a full list blob — must use it directly. + snap = delta_graph.get_state(config) + assert [m.id for m in snap.values["messages"]] == ["h1", "a1", "h2"], ( + f"tip hydration mismatch: got {[m.id for m in snap.values['messages']]}" + ) + + # Settled ancestor boundaries must also match. + pre_settled = [ + [m.id for m in s.values.get("messages", [])] + for s in pre_graph.get_state_history(config) + if s.next == ("__start__",) + ] + delta_settled = [ + [m.id for m in s.values.get("messages", [])] + for s in delta_graph.get_state_history(config) + if s.next == ("__start__",) + ] + assert delta_settled == pre_settled, ( + f"settled boundary mismatch after migration: " + f"pre={pre_settled}, delta={delta_settled}" + ) + + +async def test_add_messages_to_delta_migration_preserves_message_history_async() -> None: + """Async variant of the add_messages migration test.""" + checkpointer = InMemorySaver() + config = {"configurable": {"thread_id": "add-messages-migration-async"}} + + pre_graph = _add_messages_graph(checkpointer) + await pre_graph.ainvoke( + {"messages": [HumanMessage(content="hello", id="h1")]}, config + ) + await pre_graph.ainvoke( + {"messages": [AIMessage(content="hi", id="a1")]}, config + ) + + delta_graph = _delta_messages_graph(checkpointer) + snap = await delta_graph.aget_state(config) + assert [m.id for m in snap.values["messages"]] == ["h1", "a1"], ( + f"async tip hydration mismatch: got {[m.id for m in snap.values['messages']]}" + )