Compare commits

..
Author SHA1 Message Date
Elior Nataf Lackritz d5238f83cd fix(langgraph): don't walk history for a DeltaChannel that was never written
A DeltaChannel with no version at the checkpoint was never written, so it
is empty. channels_from_checkpoint still asked the saver for its history,
and with no snapshot to stop at the walk read every ancestor of the thread
on every load. On SQLite that fails outright once a thread has about 32k
checkpoints: stage 2 binds one variable per ancestor. Only channels with a
version are walked now.
2026-09-30 20:41:06 -04:00
2 changed files with 50 additions and 2 deletions
@@ -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:
@@ -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}"
)