From a8e732c87930134420c53b2b3d9af2f41b34a857 Mon Sep 17 00:00:00 2001 From: Elior Nataf Lackritz Date: Wed, 30 Sep 2026 12:44:18 -0400 Subject: [PATCH] fix(langgraph): give each bulk_update_state update its own task id An update whose node has no pending task to reuse was stored under uuid5(checkpoint_id, INTERRUPT), so every such update in one superstep shared a task id. Savers keep one write per (task_id, idx), so all but the first update's writes were dropped. Plain channels were unaffected, since their value is stored in the new checkpoint, but a DeltaChannel replays those writes and lost every update after the first. The ith update now gets uuid5(checkpoint_id, f"{INTERRUPT}:{i}"). The first keeps the old id, so a single update stores exactly what it did before. --- libs/langgraph/langgraph/pregel/main.py | 18 +++- .../tests/test_delta_channel_update_state.py | 98 +++++++++++++++++-- 2 files changed, 105 insertions(+), 11 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index f4ca48024..11168d13b 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -1950,7 +1950,7 @@ class Pregel( run_tasks: list[PregelTaskWrites] = [] run_task_ids: list[str] = [] - for as_node, values, provided_task_id in valid_updates: + for i, (as_node, values, provided_task_id) in enumerate(valid_updates): # create task to run all writers of the chosen node writers = self.nodes[as_node].flat_writers if not writers: @@ -1964,7 +1964,7 @@ class Pregel( task_id = provided_task_id or ( prepared_task_ids.popleft() if prepared_task_ids - else str(uuid5(UUID(checkpoint["id"]), INTERRUPT)) + else _update_task_id(checkpoint["id"], i) ) run_tasks.append(task) run_task_ids.append(task_id) @@ -2410,7 +2410,7 @@ class Pregel( run_tasks: list[PregelTaskWrites] = [] run_task_ids: list[str] = [] - for as_node, values, provided_task_id in valid_updates: + for i, (as_node, values, provided_task_id) in enumerate(valid_updates): # create task to run all writers of the chosen node writers = self.nodes[as_node].flat_writers if not writers: @@ -2424,7 +2424,7 @@ class Pregel( task_id = provided_task_id or ( prepared_task_ids.popleft() if prepared_task_ids - else str(uuid5(UUID(checkpoint["id"]), INTERRUPT)) + else _update_task_id(checkpoint["id"], i) ) run_tasks.append(task) run_task_ids.append(task_id) @@ -4172,6 +4172,16 @@ class Pregel( await self.cache.aclear(namespaces) +def _update_task_id(checkpoint_id: str, i: int) -> str: + """Task id for the `i`th update of a superstep that has no task to reuse. + + Savers keep one write per `(task_id, idx)`, so updates sharing an id lose + all but the first one's writes, which a `DeltaChannel` replays from. The + first update keeps the id a lone update has always had. + """ + return str(uuid5(UUID(checkpoint_id), INTERRUPT if i == 0 else f"{INTERRUPT}:{i}")) + + def _trigger_to_nodes(nodes: dict[str, PregelNode]) -> Mapping[str, Sequence[str]]: """Index from a trigger to nodes that depend on it.""" trigger_to_nodes: defaultdict[str, list[str]] = defaultdict(list) diff --git a/libs/langgraph/tests/test_delta_channel_update_state.py b/libs/langgraph/tests/test_delta_channel_update_state.py index da6db40f6..66e163961 100644 --- a/libs/langgraph/tests/test_delta_channel_update_state.py +++ b/libs/langgraph/tests/test_delta_channel_update_state.py @@ -20,6 +20,7 @@ from typing import Annotated, Any import pytest from langchain_core.messages import HumanMessage +from langgraph.checkpoint.base import BaseCheckpointSaver from langgraph.checkpoint.memory import InMemorySaver from langgraph.checkpoint.serde.types import _DeltaSnapshot from typing_extensions import TypedDict @@ -27,16 +28,17 @@ from typing_extensions import TypedDict from langgraph.channels.delta import DeltaChannel from langgraph.graph import START, StateGraph from langgraph.graph.message import _messages_delta_reducer -from langgraph.types import StateUpdate +from langgraph.types import StateSnapshot, StateUpdate pytestmark = pytest.mark.anyio def _build_graph( - checkpointer: InMemorySaver, + checkpointer: BaseCheckpointSaver, *, two_nodes: bool = False, snapshot_frequency: int = 1000, + interrupt_before: list[str] | None = None, ) -> Any: """Compile a minimal DeltaChannel-backed `messages` graph. @@ -63,7 +65,7 @@ def _build_graph( builder.set_finish_point("assistant") else: builder.set_finish_point("model") - return builder.compile(checkpointer=checkpointer) + return builder.compile(checkpointer=checkpointer, interrupt_before=interrupt_before) # --------------------------------------------------------------------------- @@ -273,10 +275,6 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None: that each call `put_writes`. Guards the regression where moving `put_writes` outside the per-task loop would persist only the last task's writes. - - Explicit `task_id`s are required to disambiguate writes belonging to - different `StateUpdate`s targeting the same node — otherwise both share - the deterministic interrupt-derived id and collide in the saver. """ saver = InMemorySaver() @@ -310,6 +308,92 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None: assert sorted(ids) == ["m1", "m2"] +def _update(content: str, as_node: str) -> StateUpdate: + return StateUpdate( + values={"messages": [HumanMessage(content=content, id=content)]}, + as_node=as_node, + ) + + +def _contents(state: StateSnapshot) -> list[str]: + return [m.content for m in state.values["messages"]] + + +def test_bulk_update_state_keeps_every_update_without_task_ids( + sync_checkpointer: BaseCheckpointSaver, +) -> None: + graph = _build_graph(sync_checkpointer, two_nodes=True) + config = {"configurable": {"thread_id": "bulk-no-task-ids"}} + graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config) + + graph.bulk_update_state( + config, + [ + [ + _update("first", "model"), + _update("second", "model"), + _update("third", "assistant"), + ] + ], + ) + + contents = _contents(graph.get_state(config)) + assert sorted(contents) == ["first", "hi", "second", "third"], ( + f"every update's writes must persist; got {contents}" + ) + + +async def test_abulk_update_state_keeps_every_update_without_task_ids( + async_checkpointer: BaseCheckpointSaver, +) -> None: + graph = _build_graph(async_checkpointer, two_nodes=True) + config = {"configurable": {"thread_id": "bulk-no-task-ids"}} + await graph.ainvoke({"messages": [HumanMessage(content="hi", id="hi")]}, config) + + await graph.abulk_update_state( + config, + [ + [ + _update("first", "model"), + _update("second", "model"), + _update("third", "assistant"), + ] + ], + ) + + contents = _contents(await graph.aget_state(config)) + assert sorted(contents) == ["first", "hi", "second", "third"], ( + f"every update's writes must persist; got {contents}" + ) + + +def test_bulk_update_state_keeps_every_update_next_to_a_pending_task( + sync_checkpointer: BaseCheckpointSaver, +) -> None: + graph = _build_graph( + sync_checkpointer, two_nodes=True, interrupt_before=["assistant"] + ) + config = {"configurable": {"thread_id": "bulk-pending-task"}} + graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config) + assert graph.get_state(config).next == ("assistant",) + + graph.bulk_update_state( + config, + [ + [ + _update("first", "assistant"), + _update("second", "model"), + _update("third", "model"), + ] + ], + ) + + contents = _contents(graph.get_state(config)) + assert sorted(contents) == ["first", "hi", "second", "third"], ( + f"every update's writes must persist; got {contents}" + ) + + # --------------------------------------------------------------------------- # Public-API observation of fresh-thread checkpoint shape # ---------------------------------------------------------------------------