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.
This commit is contained in:
Elior Nataf Lackritz
2026-09-30 12:47:55 -04:00
parent eb69f67b65
commit a8e732c879
2 changed files with 105 additions and 11 deletions
+14 -4
View File
@@ -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)
@@ -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
# ---------------------------------------------------------------------------