diff --git a/libs/langgraph/langgraph/pregel/_checkpoint.py b/libs/langgraph/langgraph/pregel/_checkpoint.py index c336f75a6..028a53e1e 100644 --- a/libs/langgraph/langgraph/pregel/_checkpoint.py +++ b/libs/langgraph/langgraph/pregel/_checkpoint.py @@ -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], diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index f4ca48024..ccfd9eadb 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -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" + }, }, {}, ) diff --git a/libs/langgraph/tests/test_delta_channel_supersteps_bound.py b/libs/langgraph/tests/test_delta_channel_supersteps_bound.py index a2c3a2776..cf2fc8513 100644 --- a/libs/langgraph/tests/test_delta_channel_supersteps_bound.py +++ b/libs/langgraph/tests/test_delta_channel_supersteps_bound.py @@ -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