diff --git a/libs/langgraph/langgraph/pregel/_checkpoint.py b/libs/langgraph/langgraph/pregel/_checkpoint.py index c336f75a6..c77a8b283 100644 --- a/libs/langgraph/langgraph/pregel/_checkpoint.py +++ b/libs/langgraph/langgraph/pregel/_checkpoint.py @@ -242,6 +242,10 @@ def channels_from_checkpoint( (`_DeltaSnapshot` blob or pre-migration plain value) and accumulate the writes between it and the target. All delta channels needing replay are batched into a single saver call. + + A delta channel with no version at the checkpoint was never written, so + it is empty without a walk. A walk for it would find no snapshot to stop + at and read every ancestor, every time the thread is loaded. """ channel_specs: dict[str, BaseChannel] = {} managed_specs: dict[str, ManagedValueSpec] = {} @@ -254,7 +258,8 @@ def channels_from_checkpoint( delta_channels: list[str] = [ k for k, spec in channel_specs.items() - if _needs_replay(spec, checkpoint["channel_values"].get(k, MISSING)) + if k in checkpoint["channel_versions"] + and _needs_replay(spec, checkpoint["channel_values"].get(k, MISSING)) ] histories: Mapping[str, Any] = {} if delta_channels and saver is not None and config is not None: @@ -296,7 +301,8 @@ async def achannels_from_checkpoint( delta_channels: list[str] = [ k for k, spec in channel_specs.items() - if _needs_replay(spec, checkpoint["channel_values"].get(k, MISSING)) + if k in checkpoint["channel_versions"] + and _needs_replay(spec, checkpoint["channel_values"].get(k, MISSING)) ] histories: Mapping[str, Any] = {} if delta_channels and saver is not None and config is not None: diff --git a/libs/langgraph/tests/test_delta_channel_supersteps_bound.py b/libs/langgraph/tests/test_delta_channel_supersteps_bound.py index a2c3a2776..6d96e6169 100644 --- a/libs/langgraph/tests/test_delta_channel_supersteps_bound.py +++ b/libs/langgraph/tests/test_delta_channel_supersteps_bound.py @@ -193,3 +193,45 @@ async def test_counter_reset_after_supersteps_snapshot() -> None: state = graph.get_state(config) assert state.values["b"] == ["seed-b"] + + +class _HistoryRequestSaver(InMemorySaver): + def __init__(self) -> None: + super().__init__() + self.requested: list[list[str]] = [] + + def get_delta_channel_history(self, *, config: Any, channels: Any) -> Any: + self.requested.append(sorted(channels)) + return super().get_delta_channel_history(config=config, channels=channels) + + +def test_never_written_channel_is_not_walked() -> None: + saver = _HistoryRequestSaver() + graph = _build_two_channel_graph(saver, n_loops=3) + config = {"configurable": {"thread_id": "never-written"}} + graph.invoke({"a": ["seed-a"]}, config) + saver.requested.clear() + + graph.invoke({"a": ["more-a"]}, config) + state = graph.get_state(config) + + assert state.values["b"] == [] + assert saver.requested and all(r == ["a"] for r in saver.requested), ( + f"only the written channel needs a walk; asked for {saver.requested}" + ) + + +async def test_anever_written_channel_is_not_walked() -> None: + saver = _HistoryRequestSaver() + graph = _build_two_channel_graph(saver, n_loops=3) + config = {"configurable": {"thread_id": "never-written"}} + await graph.ainvoke({"a": ["seed-a"]}, config) + saver.requested.clear() + + await graph.ainvoke({"a": ["more-a"]}, config) + state = await graph.aget_state(config) + + assert state.values["b"] == [] + assert saver.requested and all(r == ["a"] for r in saver.requested), ( + f"only the written channel needs a walk; asked for {saver.requested}" + )