Compare commits

...
Author SHA1 Message Date
Elior Nataf Lackritz 3129e521ea fix(langgraph): keep DeltaChannel counters on every update_state path
update_state(None, as_node=END), as_node="__input__" and "__copy__" saved
their checkpoint without counters_since_delta_snapshot, so the next
checkpoint restarted every delta channel's snapshot cadence from zero.
END and __input__ now advance the counters by one superstep, like the
other update_state paths; __copy__ keeps the copied checkpoint's.
2026-09-30 21:00:26 -04:00
3 changed files with 127 additions and 6 deletions
@@ -114,6 +114,26 @@ def create_metadata_for_update_state_api(
return new_counters
def advance_delta_counters(
channels: Mapping[str, BaseChannel],
updated_channels: set[str],
*,
prev_metadata: Mapping[str, Any] | None,
) -> dict[str, Any]:
"""The `counters_since_delta_snapshot` entry for an update_state
checkpoint saved one superstep after `prev_metadata`'s, for the paths
that skip `create_checkpoint_plan_for_update_state_api`.
Without it, the next checkpoint restarts every delta channel's snapshot
cadence from zero.
"""
counters = create_metadata_for_update_state_api(
channels, updated_channels, prev_metadata=prev_metadata
)
non_zero = {k: v for k, v in counters.items() if v != (0, 0)}
return {"counters_since_delta_snapshot": non_zero} if non_zero else {}
def create_checkpoint_plan_for_update_state_api(
channels: Mapping[str, BaseChannel],
updated_channels: set[str],
+53 -6
View File
@@ -129,6 +129,7 @@ from langgraph.pregel._algo import (
from langgraph.pregel._call import identifier
from langgraph.pregel._checkpoint import (
achannels_from_checkpoint,
advance_delta_counters,
channels_from_checkpoint,
copy_checkpoint,
create_checkpoint,
@@ -1681,6 +1682,7 @@ class Pregel(
"Cannot apply multiple updates when clearing state"
)
updated_channels: set[str] = set()
if saved is not None:
# tasks for this checkpoint
next_tasks = prepare_next_tasks(
@@ -1703,7 +1705,7 @@ class Pregel(
for w in saved.pending_writes or []
if w[0] == NULL_TASK_ID
]:
apply_writes(
updated_channels |= apply_writes(
checkpoint,
channels,
[PregelTaskWrites((), INPUT, null_writes, [])],
@@ -1718,7 +1720,7 @@ class Pregel(
continue
next_tasks[tid].writes.append((k, v))
# clear all current tasks
apply_writes(
updated_channels |= apply_writes(
checkpoint,
channels,
next_tasks.values(),
@@ -1733,6 +1735,13 @@ class Pregel(
"source": "update",
"step": step + 1,
"parents": saved.metadata.get("parents", {}) if saved else {},
**(
advance_delta_counters(
channels, updated_channels, prev_metadata=saved.metadata
)
if saved
else {}
),
},
get_new_channel_versions(
checkpoint_previous_versions,
@@ -1751,7 +1760,7 @@ class Pregel(
)
if input_writes := deque(map_input(self.input_channels, values)):
apply_writes(
updated_channels = apply_writes(
checkpoint,
channels,
[PregelTaskWrites((), INPUT, input_writes, [])],
@@ -1774,6 +1783,15 @@ class Pregel(
"parents": saved.metadata.get("parents", {})
if saved
else {},
**(
advance_delta_counters(
channels,
updated_channels,
prev_metadata=saved.metadata,
)
if saved
else {}
),
},
get_new_channel_versions(
checkpoint_previous_versions,
@@ -1819,6 +1837,12 @@ class Pregel(
"source": "fork",
"step": step + 1,
"parents": saved.metadata.get("parents", {}),
# The copy has the same values and the same parent.
**{
k: v
for k, v in saved.metadata.items()
if k == "counters_since_delta_snapshot"
},
},
{},
)
@@ -2145,6 +2169,7 @@ class Pregel(
raise InvalidUpdateError(
"Cannot apply multiple updates when clearing state"
)
updated_channels: set[str] = set()
if saved is not None:
# tasks for this checkpoint
next_tasks = prepare_next_tasks(
@@ -2167,7 +2192,7 @@ class Pregel(
for w in saved.pending_writes or []
if w[0] == NULL_TASK_ID
]:
apply_writes(
updated_channels |= apply_writes(
checkpoint,
channels,
[PregelTaskWrites((), INPUT, null_writes, [])],
@@ -2182,7 +2207,7 @@ class Pregel(
continue
next_tasks[tid].writes.append((k, v))
# clear all current tasks
apply_writes(
updated_channels |= apply_writes(
checkpoint,
channels,
next_tasks.values(),
@@ -2197,6 +2222,13 @@ class Pregel(
"source": "update",
"step": step + 1,
"parents": saved.metadata.get("parents", {}) if saved else {},
**(
advance_delta_counters(
channels, updated_channels, prev_metadata=saved.metadata
)
if saved
else {}
),
},
get_new_channel_versions(
checkpoint_previous_versions, checkpoint["channel_versions"]
@@ -2214,7 +2246,7 @@ class Pregel(
)
if input_writes := deque(map_input(self.input_channels, values)):
apply_writes(
updated_channels = apply_writes(
checkpoint,
channels,
[PregelTaskWrites((), INPUT, input_writes, [])],
@@ -2237,6 +2269,15 @@ class Pregel(
"parents": saved.metadata.get("parents", {})
if saved
else {},
**(
advance_delta_counters(
channels,
updated_channels,
prev_metadata=saved.metadata,
)
if saved
else {}
),
},
get_new_channel_versions(
checkpoint_previous_versions,
@@ -2282,6 +2323,12 @@ class Pregel(
"source": "fork",
"step": step + 1,
"parents": saved.metadata.get("parents", {}),
# The copy has the same values and the same parent.
**{
k: v
for k, v in saved.metadata.items()
if k == "counters_since_delta_snapshot"
},
},
{},
)
@@ -147,6 +147,60 @@ async def test_predicate_fires_on_supersteps_overflow() -> None:
assert "x" not in result2
def _delta_counters(saver: InMemorySaver, config: Any) -> dict[str, list[int]]:
tup = saver.get_tuple(config)
assert tup is not None
counters = tup.metadata.get("counters_since_delta_snapshot") or {}
return {ch: list(c) for ch, c in counters.items()}
_UPDATE_PATHS_WITHOUT_THE_SNAPSHOT_PLAN = pytest.mark.parametrize(
("values", "as_node", "supersteps"),
[
(None, END, 1),
(None, "__copy__", 0),
({"a": []}, "__input__", 1),
],
ids=["clear as END", "copy", "update as input"],
)
@_UPDATE_PATHS_WITHOUT_THE_SNAPSHOT_PLAN
def test_update_state_path_keeps_delta_counters(
values: Any, as_node: str, supersteps: int
) -> None:
saver = InMemorySaver()
graph = _build_two_channel_graph(saver)
config = {"configurable": {"thread_id": "counters"}}
graph.invoke({"a": ["seed-a"], "b": ["seed-b"]}, config)
before = _delta_counters(saver, config)
assert set(before) == {"a", "b"}, f"both channels need live counters: {before}"
updated = graph.update_state(config, values, as_node=as_node)
assert _delta_counters(saver, updated) == {
ch: [u, s + supersteps] for ch, (u, s) in before.items()
}
@_UPDATE_PATHS_WITHOUT_THE_SNAPSHOT_PLAN
async def test_aupdate_state_path_keeps_delta_counters(
values: Any, as_node: str, supersteps: int
) -> None:
saver = InMemorySaver()
graph = _build_two_channel_graph(saver)
config = {"configurable": {"thread_id": "counters"}}
await graph.ainvoke({"a": ["seed-a"], "b": ["seed-b"]}, config)
before = _delta_counters(saver, config)
assert set(before) == {"a", "b"}, f"both channels need live counters: {before}"
updated = await graph.aupdate_state(config, values, as_node=as_node)
assert _delta_counters(saver, updated) == {
ch: [u, s + supersteps] for ch, (u, s) in before.items()
}
async def test_counter_reset_after_supersteps_snapshot() -> None:
"""After the supersteps bound triggers a snapshot, the counters for
that channel reset. Verify by using a bound higher than one run's