From 7bf325d8a0e432590b6088145e76038755a269e1 Mon Sep 17 00:00:00 2001 From: Quanzheng Long Date: Wed, 6 May 2026 15:28:10 -0700 Subject: [PATCH] fix lint: narrow checkpointer types + suppress UP013 in tests - _put_exit_delta_writes: narrow self.checkpointer / put_after_previous / put_writes to non-None at the top so mypy accepts submit() calls. - test_exit_delta_persistence.py: suppress UP013 on functional TypedDict() uses (class form can't reference local variables in Annotated). Co-authored-by: Cursor --- libs/langgraph/langgraph/pregel/_loop.py | 37 ++++++++----------- .../tests/test_exit_delta_persistence.py | 20 ++++------ 2 files changed, 23 insertions(+), 34 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/_loop.py b/libs/langgraph/langgraph/pregel/_loop.py index 01b812e82..b3e4b8bdb 100644 --- a/libs/langgraph/langgraph/pregel/_loop.py +++ b/libs/langgraph/langgraph/pregel/_loop.py @@ -902,9 +902,7 @@ class PregelLoop: if self._exit_delta_writes is not None: for c, v in input_writes: if isinstance(self.specs.get(c), DeltaChannel): - self._exit_delta_writes.append( - (self.step, NULL_TASK_ID, c, v) - ) + self._exit_delta_writes.append((self.step, NULL_TASK_ID, c, v)) # Persist delta-channel input writes so sub-freq inputs are # recoverable via ancestor walk (mirrors the Command input path). if self.durability != "exit": @@ -967,7 +965,7 @@ class PregelLoop: def _put_checkpoint(self, metadata: CheckpointMetadata) -> None: # `is` (object identity) — not `==`. Three of four call sites pass a - # fresh dict ({"source":"input"|"loop"|"fork"}); only + # fresh dict ({"source":"input"|"loop"|"fork"}); only # `_suppress_interrupt`(will rename to _on_loop_exit soon) # at exit reuses the existing `self.checkpoint_metadata` instance. So # `metadata is self.checkpoint_metadata` is True only on the exit call, @@ -985,7 +983,7 @@ class PregelLoop: # `_put_checkpoint` is called once per superstep with a fresh # metadata dict (source="input"|"loop"|"fork") — those are the # intermediate calls that bump the count by +1 for each delta - # channel touched that step. In exit mode, + # channel touched that step. In exit mode, # `_suppress_interrupt`(will rename to _on_loop_exit soon) # additionally calls `_put_checkpoint(self.checkpoint_metadata)` AT # EXIT to commit the final checkpoint — this runs *after* the last @@ -996,8 +994,7 @@ class PregelLoop: # used to mask this latent bug by resetting every count to 0.) if not exiting: prev_counts = dict( - self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) - or {} + self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {} ) new_counts = dict(prev_counts) if self.updated_channels: @@ -1009,8 +1006,7 @@ class PregelLoop: self.checkpoint_metadata = metadata else: new_counts = dict( - self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) - or {} + self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {} ) # do checkpoint? do_checkpoint = self._checkpointer_put_after_previous is not None and ( @@ -1101,12 +1097,15 @@ class PregelLoop: Stub is created lazily — only when no persisted parent exists AND at least one delta channel has writes that won't be snapshotted. """ - if not self._exit_delta_writes or self.checkpointer is None: + if ( + not self._exit_delta_writes + or self.checkpointer is None + or self._checkpointer_put_after_previous is None + or self.checkpointer_put_writes is None + ): return - counts = ( - self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {} - ) + counts = self.checkpoint_metadata.get("delta_updates_since_snapshot", {}) or {} will_snapshot = decide_delta_snapshots(self.channels, counts) pending = [ @@ -1158,9 +1157,7 @@ class PregelLoop: CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get( CONFIG_KEY_CHECKPOINT_NS, "" ), - CONFIG_KEY_CHECKPOINT_ID: anchor_config[CONF][ - CONFIG_KEY_CHECKPOINT_ID - ], + CONFIG_KEY_CHECKPOINT_ID: anchor_config[CONF][CONFIG_KEY_CHECKPOINT_ID], }, ) for (step, tid), entries in grouped.items(): @@ -1555,9 +1552,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager): ) self._delta_write_futs = [] self._exit_delta_writes = ( - [] - if self.durability == "exit" and self.checkpointer is not None - else None + [] if self.durability == "exit" and self.checkpointer is not None else None ) self.submit = self.stack.enter_context(BackgroundExecutor(self.config)) self.channels, self.managed = channels_from_checkpoint( @@ -1815,9 +1810,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager): ) self._delta_write_futs = [] self._exit_delta_writes = ( - [] - if self.durability == "exit" and self.checkpointer is not None - else None + [] if self.durability == "exit" and self.checkpointer is not None else None ) self.submit = await self.stack.enter_async_context( AsyncBackgroundExecutor(self.config) diff --git a/libs/langgraph/tests/test_exit_delta_persistence.py b/libs/langgraph/tests/test_exit_delta_persistence.py index 3667e28a0..2ab09b521 100644 --- a/libs/langgraph/tests/test_exit_delta_persistence.py +++ b/libs/langgraph/tests/test_exit_delta_persistence.py @@ -27,7 +27,9 @@ def _build_graph( freq: int = 1000, ) -> Any: channel = DeltaChannel(_messages_delta_reducer, snapshot_frequency=freq) - State = TypedDict("State", {"messages": Annotated[list, channel]}) # type: ignore[call-overload] + # Functional TypedDict form: class form can't reference `channel` (a + # local variable) inside Annotated due to forward-ref evaluation rules. + State = TypedDict("State", {"messages": Annotated[list, channel]}) # type: ignore[call-overload] # noqa: UP013 def respond(state: dict) -> dict: i = len(state["messages"]) @@ -47,7 +49,7 @@ def _build_graph( async def test_exit_first_run_no_delta_writes() -> None: """Graph with delta channel invoked with input that doesn't touch it. Only one checkpoint row, no stub.""" - State = TypedDict( + State = TypedDict( # noqa: UP013 "State", { "messages": Annotated[list, DeltaChannel(_messages_delta_reducer)], @@ -93,9 +95,7 @@ async def test_exit_first_run_all_snapshot() -> None: head = saver.get_tuple(config) assert head is not None - assert isinstance( - head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot - ) + assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot) state = graph.get_state(config) assert [m.content for m in state.values["messages"]] == ["hi", "reply-1"] @@ -208,9 +208,7 @@ async def test_exit_snapshot_fires_at_frequency() -> None: assert head is not None count2 = head.metadata.get("delta_updates_since_snapshot", {}).get("messages", 0) assert count2 == 0, f"Expected reset to 0 after snapshot, got {count2}" - assert isinstance( - head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot - ) + assert isinstance(head.checkpoint["channel_values"].get("messages"), _DeltaSnapshot) async def test_exit_mixed_snapshot_and_non_snapshot() -> None: @@ -219,7 +217,7 @@ async def test_exit_mixed_snapshot_and_non_snapshot() -> None: fast_ch = DeltaChannel(_messages_delta_reducer, snapshot_frequency=1) slow_ch = DeltaChannel(_messages_delta_reducer, snapshot_frequency=1000) - State = TypedDict( + State = TypedDict( # noqa: UP013 "State", {"fast": Annotated[list, fast_ch], "slow": Annotated[list, slow_ch]}, ) # type: ignore[call-overload] @@ -300,9 +298,7 @@ async def test_exit_metadata_round_trip() -> None: ) head = saver.get_tuple(config) assert head is not None - count = head.metadata.get("delta_updates_since_snapshot", {}).get( - "messages", 0 - ) + count = head.metadata.get("delta_updates_since_snapshot", {}).get("messages", 0) cumulative = i * 2 if cumulative >= freq: assert count == 0 or count == cumulative % freq or count < freq, (