Compare commits

...
Author SHA1 Message Date
Elior Nataf Lackritz 56fd3458cc refactor(langgraph): check for put_writes' task_path in one place
Both loop setups and bulk_update_state inspected the saver's signature
themselves; they now share put_writes_accepts_task_path.
2026-10-06 18:29:03 -04:00
Elior Nataf Lackritz a7a465877d fix(langgraph): store a task path for each bulk_update_state update
Updates applied together in one superstep are applied in the order
given, but were stored without a task path, so a saver replaying by
task path fell back to their task ids, which are hashes. Each update now
gets the path (__interrupt__, i) and its writes store it.
2026-10-06 16:03:21 -04:00
Elior Nataf Lackritz b1e167401e Merge branch 'main' into fix/bulk-update-state-shared-task-id 2026-10-06 16:00:47 -04:00
Elior Nataf Lackritz 43a227b584 test(langgraph): check that explicit task ids key bulk_update_state writes
The test ran on a fresh thread, where bulk_update_state stores no writes per
task, so it passed whatever ids it was given and whether or not each task's
writes were stored. It now runs after a first invoke and checks that the base
holds both updates' writes under the given ids.
2026-10-05 11:46:29 -04:00
Elior Nataf Lackritz a8e732c879 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.
2026-09-30 12:47:55 -04:00
4 changed files with 202 additions and 22 deletions
@@ -3,6 +3,7 @@ from __future__ import annotations
import uuid
from collections.abc import Callable, Iterable, Mapping
from datetime import datetime, timezone
from inspect import signature
from typing import Any, Literal, cast
from langchain_core.runnables import RunnableConfig
@@ -49,6 +50,15 @@ def empty_checkpoint() -> Checkpoint:
)
def put_writes_accepts_task_path(put_writes: Callable[..., Any]) -> bool:
"""Whether a saver's `put_writes` or `aput_writes` takes `task_path`.
Savers written before the parameter existed don't, so it is passed only
when this is true.
"""
return signature(put_writes).parameters.get("task_path") is not None
def exit_delta_task_id(step: int, task_id: str) -> str:
"""Synthetic task id for exit-mode DeltaChannel writes.
+3 -5
View File
@@ -12,7 +12,6 @@ from contextlib import (
ExitStack,
)
from datetime import datetime, timezone
from inspect import signature
from types import TracebackType
from typing import (
Any,
@@ -107,6 +106,7 @@ from langgraph.pregel._checkpoint import (
delta_channels_with_pending_writes,
empty_checkpoint,
exit_delta_task_id,
put_writes_accepts_task_path,
)
from langgraph.pregel._executor import (
AsyncBackgroundExecutor,
@@ -1584,8 +1584,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
self.checkpointer_get_next_version = checkpointer.get_next_version
self.checkpointer_put_writes = checkpointer.put_writes
self.checkpointer_put_writes_accepts_task_path = (
signature(checkpointer.put_writes).parameters.get("task_path")
is not None
put_writes_accepts_task_path(checkpointer.put_writes)
)
else:
self.checkpointer_get_next_version = increment
@@ -1840,8 +1839,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
self.checkpointer_get_next_version = checkpointer.get_next_version
self.checkpointer_put_writes = checkpointer.aput_writes
self.checkpointer_put_writes_accepts_task_path = (
signature(checkpointer.aput_writes).parameters.get("task_path")
is not None
put_writes_accepts_task_path(checkpointer.aput_writes)
)
else:
self.checkpointer_get_next_version = increment
+37 -8
View File
@@ -126,6 +126,7 @@ from langgraph.pregel._algo import (
apply_writes,
local_read,
prepare_next_tasks,
task_path_str,
)
from langgraph.pregel._call import identifier
from langgraph.pregel._checkpoint import (
@@ -139,6 +140,7 @@ from langgraph.pregel._checkpoint import (
delta_channels_with_pending_writes,
empty_checkpoint,
get_updated_channels_from_tasks,
put_writes_accepts_task_path,
versions_seen_without_bumps,
)
from langgraph.pregel._draw import draw_graph
@@ -1983,13 +1985,13 @@ 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:
raise InvalidUpdateError(f"Node {as_node} has no writers")
writes: deque[tuple[str, Any]] = deque()
task = PregelTaskWrites((), as_node, writes, [INTERRUPT])
task = PregelTaskWrites((INTERRUPT, i), as_node, writes, [INTERRUPT])
# get the task ids that were prepared for this node
# if a task id was provided in the StateUpdate, we use it
# otherwise, we use the next available task id
@@ -1997,7 +1999,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)
@@ -2050,7 +2052,10 @@ class Pregel(
channel_writes = [w for w in task.writes if w[0] != PUSH]
if channel_writes:
checkpointer.put_writes(
checkpoint_config, channel_writes, task_id
checkpoint_config,
channel_writes,
task_id,
**_task_path_kwarg(checkpointer.put_writes, task),
)
apply_writes(
checkpoint,
@@ -2471,13 +2476,13 @@ 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:
raise InvalidUpdateError(f"Node {as_node} has no writers")
writes: deque[tuple[str, Any]] = deque()
task = PregelTaskWrites((), as_node, writes, [INTERRUPT])
task = PregelTaskWrites((INTERRUPT, i), as_node, writes, [INTERRUPT])
# get the task ids that were prepared for this node
# if a task id was provided in the StateUpdate, we use it
# otherwise, we use the next available task id
@@ -2485,7 +2490,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)
@@ -2538,7 +2543,10 @@ class Pregel(
channel_writes = [w for w in task.writes if w[0] != PUSH]
if channel_writes:
await checkpointer.aput_writes(
checkpoint_config, channel_writes, task_id
checkpoint_config,
channel_writes,
task_id,
**_task_path_kwarg(checkpointer.aput_writes, task),
)
apply_writes(
checkpoint,
@@ -4237,6 +4245,27 @@ class Pregel(
await self.cache.aclear(namespaces)
def _task_path_kwarg(put_writes: Callable[..., Any], task: PregelTaskWrites) -> dict:
"""Pass the task's path to savers whose `put_writes` takes one.
Savers that replay a checkpoint's writes in task path order then give back
updates applied together in the order they were given.
"""
if not put_writes_accepts_task_path(put_writes):
return {}
return {"task_path": task_path_str(task.path)}
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)
# ---------------------------------------------------------------------------
@@ -304,15 +306,14 @@ 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()
graph = _build_graph(saver)
config = {"configurable": {"thread_id": "bulk-multi-task"}}
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
base = saver.get_tuple(config)
assert base is not None
graph.bulk_update_state(
config,
@@ -332,13 +333,155 @@ def test_bulk_update_state_multi_task_per_superstep_delta_channel() -> None:
],
)
stored = saver.get_tuple(base.config)
assert stored is not None
assert {task_id for task_id, _, _ in stored.pending_writes or []} == {
"task-1",
"task-2",
}, "explicit task ids must key the stored writes"
state = graph.get_state(config)
contents = [m.content for m in state.values["messages"]]
ids = [m.id for m in state.values["messages"]]
assert sorted(contents) == ["first", "second"], (
assert sorted(contents) == ["first", "hi", "second"], (
f"both updates' writes must persist; got {contents}"
)
assert sorted(ids) == ["m1", "m2"]
assert sorted(ids) == ["hi", "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}"
)
class _TaskPathOrderSaver(InMemorySaver):
"""Replays each checkpoint's writes by `(task_path, task_id, idx)`."""
def get_tuple(self, config: Any) -> Any:
tup = super().get_tuple(config)
if tup is None or not tup.pending_writes:
return tup
conf = tup.config["configurable"]
stored = self.writes[
(conf["thread_id"], conf["checkpoint_ns"], conf["checkpoint_id"])
]
rows = sorted(
zip(stored.items(), tup.pending_writes),
key=lambda row: (row[0][1][3], *row[0][0]),
)
return tup._replace(pending_writes=[write for _, write in rows])
get_delta_channel_history = BaseCheckpointSaver.get_delta_channel_history
aget_delta_channel_history = BaseCheckpointSaver.aget_delta_channel_history
GIVEN = ["u1", "u2", "u3", "u4", "u5", "u6"]
def _updates_in_given_order() -> list[list[StateUpdate]]:
return [
[_update(c, "assistant" if i % 2 else "model") for i, c in enumerate(GIVEN)]
]
def test_bulk_update_state_replays_updates_in_the_order_given() -> None:
graph = _build_graph(_TaskPathOrderSaver(), two_nodes=True)
config = {"configurable": {"thread_id": "bulk-order"}}
graph.invoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
graph.bulk_update_state(config, _updates_in_given_order())
assert _contents(graph.get_state(config)) == ["hi", *GIVEN]
async def test_abulk_update_state_replays_updates_in_the_order_given() -> None:
graph = _build_graph(_TaskPathOrderSaver(), two_nodes=True)
config = {"configurable": {"thread_id": "bulk-order"}}
await graph.ainvoke({"messages": [HumanMessage(content="hi", id="hi")]}, config)
await graph.abulk_update_state(config, _updates_in_given_order())
assert _contents(await graph.aget_state(config)) == ["hi", *GIVEN]
# ---------------------------------------------------------------------------