From c6ae8d25b9b0c22bfe12c468905a8746b9e4229d Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 6 Aug 2025 19:09:33 +0100 Subject: [PATCH] perf: Save updated_channels to checkpoint (#5828) - This makes prepare_next_tasks constant on number of nodes in all cases, whereas before we were falling back to node iteration when resuming from an existing checkpoint --- .../langgraph/checkpoint/base/__init__.py | 6 +++++ .../langgraph/langgraph/pregel/_checkpoint.py | 3 +++ libs/langgraph/langgraph/pregel/_loop.py | 27 ++++++++++++++----- .../tests/test_checkpoint_migration.py | 6 +++++ libs/langgraph/tests/test_pregel.py | 3 +++ libs/langgraph/tests/test_pregel_async.py | 3 +++ 6 files changed, 42 insertions(+), 6 deletions(-) diff --git a/libs/checkpoint/langgraph/checkpoint/base/__init__.py b/libs/checkpoint/langgraph/checkpoint/base/__init__.py index a5704e13f..a2a768d0c 100644 --- a/libs/checkpoint/langgraph/checkpoint/base/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/base/__init__.py @@ -81,6 +81,9 @@ class Checkpoint(TypedDict): This keeps track of the versions of the channels that each node has seen. Used to determine which nodes to execute next. """ + updated_channels: list[str] | None + """The channels that were updated in this checkpoint. + """ def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint: @@ -92,6 +95,7 @@ def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint: channel_versions=checkpoint["channel_versions"].copy(), versions_seen={k: v.copy() for k, v in checkpoint["versions_seen"].items()}, pending_sends=checkpoint.get("pending_sends", []).copy(), + updated_channels=checkpoint.get("updated_channels", None), ) @@ -437,6 +441,7 @@ def empty_checkpoint() -> Checkpoint: channel_versions={}, versions_seen={}, pending_sends=[], + updated_channels=None, ) @@ -470,4 +475,5 @@ def create_checkpoint( channel_versions=checkpoint["channel_versions"], versions_seen=checkpoint["versions_seen"], pending_sends=checkpoint.get("pending_sends", []), + updated_channels=None, ) diff --git a/libs/langgraph/langgraph/pregel/_checkpoint.py b/libs/langgraph/langgraph/pregel/_checkpoint.py index 50eb254b8..9afd67af9 100644 --- a/libs/langgraph/langgraph/pregel/_checkpoint.py +++ b/libs/langgraph/langgraph/pregel/_checkpoint.py @@ -29,6 +29,7 @@ def create_checkpoint( step: int, *, id: str | None = None, + updated_channels: set[str] | None = None, ) -> Checkpoint: """Create a checkpoint for the given channels.""" ts = datetime.now(timezone.utc).isoformat() @@ -49,6 +50,7 @@ def create_checkpoint( channel_values=values, channel_versions=checkpoint["channel_versions"], versions_seen=checkpoint["versions_seen"], + updated_channels=None if updated_channels is None else sorted(updated_channels), ) @@ -81,4 +83,5 @@ def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint: channel_values=checkpoint["channel_values"].copy(), channel_versions=checkpoint["channel_versions"].copy(), versions_seen={k: v.copy() for k, v in checkpoint["versions_seen"].items()}, + updated_channels=checkpoint.get("updated_channels", None), ) diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index 9720a39cb..c9300f3d4 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -568,7 +568,9 @@ class PregelLoop: if task := tasks.get(tid): task.writes.append((k, v)) - def _first(self, *, input_keys: str | Sequence[str]) -> set[str] | None: + def _first( + self, *, input_keys: str | Sequence[str], updated_channels: set[str] | None + ) -> set[str] | None: # resuming from previous checkpoint requires # - finding a previous checkpoint # - receiving None input (outer graph) or RESUMING flag (subgraph) @@ -585,8 +587,6 @@ class PregelLoop: ), ) ) - # this can be set only when there are input_writes - updated_channels: set[str] | None = None # map command to writes if isinstance(self.input, Command): @@ -614,13 +614,15 @@ class PregelLoop: if null_writes := [ w[1:] for w in self.checkpoint_pending_writes if w[0] == NULL_TASK_ID ]: - apply_writes( + null_updated_channels = apply_writes( self.checkpoint, self.channels, [PregelTaskWrites((), INPUT, null_writes, [])], self.checkpointer_get_next_version, self.trigger_to_nodes, ) + if updated_channels is not None: + updated_channels.update(null_updated_channels) # proceed past previous checkpoint if is_resuming: self.checkpoint["versions_seen"].setdefault(INTERRUPT, {}) @@ -648,6 +650,7 @@ class PregelLoop: store=None, checkpointer=None, manager=None, + updated_channels=updated_channels, ) # apply input writes updated_channels = apply_writes( @@ -661,6 +664,7 @@ class PregelLoop: self.trigger_to_nodes, ) # save input checkpoint + self.updated_channels = updated_channels self._put_checkpoint({"source": "input"}) elif CONFIG_KEY_RESUMING not in configurable: raise EmptyInputError(f"Received no input for {input_keys}") @@ -693,6 +697,7 @@ class PregelLoop: self.channels if do_checkpoint else None, self.step, id=self.checkpoint["id"] if exiting else None, + updated_channels=self.updated_channels, ) # bail if no checkpointer if do_checkpoint and self._checkpointer_put_after_previous is not None: @@ -1036,7 +1041,12 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): self.step = self.checkpoint_metadata["step"] + 1 self.stop = self.step + self.config["recursion_limit"] + 1 self.checkpoint_previous_versions = self.checkpoint["channel_versions"].copy() - self.updated_channels = self._first(input_keys=self.input_keys) + self.updated_channels = self._first( + input_keys=self.input_keys, + updated_channels=set(self.checkpoint.get("updated_channels")) # type: ignore[arg-type] + if self.checkpoint.get("updated_channels") + else None, + ) return self @@ -1212,7 +1222,12 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): self.step = self.checkpoint_metadata["step"] + 1 self.stop = self.step + self.config["recursion_limit"] + 1 self.checkpoint_previous_versions = self.checkpoint["channel_versions"].copy() - self.updated_channels = self._first(input_keys=self.input_keys) + self.updated_channels = self._first( + input_keys=self.input_keys, + updated_channels=set(self.checkpoint.get("updated_channels")) # type: ignore[arg-type] + if self.checkpoint.get("updated_channels") + else None, + ) return self diff --git a/libs/langgraph/tests/test_checkpoint_migration.py b/libs/langgraph/tests/test_checkpoint_migration.py index f8b4adbdc..efc75735f 100644 --- a/libs/langgraph/tests/test_checkpoint_migration.py +++ b/libs/langgraph/tests/test_checkpoint_migration.py @@ -330,6 +330,7 @@ SAVED_CHECKPOINTS = { "docs": ["doc1", "doc2", "doc3", "doc4"], "answer": "doc1,doc2,doc3,doc4", }, + "updated_channels": None, }, metadata={ "source": "loop", @@ -390,6 +391,7 @@ SAVED_CHECKPOINTS = { "docs": ["doc1", "doc2", "doc3", "doc4"], "branch:to:qa": None, }, + "updated_channels": None, }, metadata={ "source": "loop", @@ -465,6 +467,7 @@ SAVED_CHECKPOINTS = { "branch:to:retriever_one": None, "docs": ["doc3", "doc4"], }, + "updated_channels": None, }, metadata={ "source": "loop", @@ -516,6 +519,7 @@ SAVED_CHECKPOINTS = { "branch:to:analyzer_one": None, "branch:to:retriever_two": None, }, + "updated_channels": None, }, metadata={ "source": "loop", @@ -570,6 +574,7 @@ SAVED_CHECKPOINTS = { "query": "what is weather in sf", "branch:to:rewrite_query": None, }, + "updated_channels": None, }, metadata={ "source": "loop", @@ -618,6 +623,7 @@ SAVED_CHECKPOINTS = { }, "versions_seen": {"__input__": {}}, "channel_values": {"__start__": {"query": "what is weather in sf"}}, + "updated_channels": None, }, metadata={ "source": "input", diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index e69e4e2f2..a1ec594fd 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -968,6 +968,7 @@ def test_pending_writes_resume( "branch:to:two": AnyVersion(), }, "channel_values": {"value": 6}, + "updated_channels": ["value"], }, metadata={ "parents": {}, @@ -1015,6 +1016,7 @@ def test_pending_writes_resume( "branch:to:one": None, "branch:to:two": None, }, + "updated_channels": ["branch:to:one", "branch:to:two", "value"], }, metadata={ "parents": {}, @@ -1066,6 +1068,7 @@ def test_pending_writes_resume( "__start__": AnyVersion(), }, "channel_values": {"__start__": {"value": 1}}, + "updated_channels": ["__start__"], }, metadata={ "parents": {}, diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 9b7fa4628..c834ebbdd 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -1908,6 +1908,7 @@ async def test_pending_writes_resume( "branch:to:two": AnyVersion(), }, "channel_values": {"value": 6}, + "updated_channels": ["value"], }, metadata={ "parents": {}, @@ -1955,6 +1956,7 @@ async def test_pending_writes_resume( "branch:to:one": None, "branch:to:two": None, }, + "updated_channels": ["branch:to:one", "branch:to:two", "value"], }, metadata={ "parents": {}, @@ -2002,6 +2004,7 @@ async def test_pending_writes_resume( "__start__": AnyVersion(), }, "channel_values": {"__start__": {"value": 1}}, + "updated_channels": ["__start__"], }, metadata={ "parents": {},