mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-17 21:25:46 +02:00
Follow-up to #8540, which turned on `PLC0415` (import-outside-top-level) for checkpoint-postgres and checkpoint-sqlite. This does the remaining six packages: checkpoint, checkpoint-conformance, langgraph, prebuilt, cli, sdk-py. Scoped to tests, per @sydney-runkle's call on #8540: library code is exempted with `per-file-ignores`, since it still has deferred imports nobody has reviewed and mixing that in would make this hard to read. ## What changed Function-level imports across 56 test files moved to module level. Nine could not move and carry an explicit `# noqa: PLC0415` with a reason: | File | Why it stays local | |---|---| | `libs/langgraph/tests/test_deprecation.py` (4) | the import has to run inside `pytest.warns` for the warning to be observed | | `libs/langgraph/tests/test_serde_allowlist.py` | try/except guard, skips when langchain_core is absent | | `libs/langgraph/tests/test_delta_channel_benchmark.py` | optional psycopg probe | | `libs/checkpoint/tests/test_conformance_delta.py` (3) | protected by a module-level `pytest.importorskip`; hoisting past the guard turns a skip into a collection error | That last one is the trap: an import moved above `pytest.importorskip` silently defeats the guard. I hit it locally and it turned the skip into a `ModuleNotFoundError` at collection. Every file with an `importorskip` or `except ImportError` was checked by hand for this. ## Verification `make lint` and `make test` in each of the six: | Package | Tests | |---|---| | checkpoint | 156 passed, 17 skipped | | checkpoint-conformance | 1 passed | | langgraph | 1968 passed, 4 skipped | | prebuilt | 284 passed | | cli | 336 passed | | sdk-py | 493 passed | Also confirmed the rule actually fires: a throwaway test file with a function-level import is flagged in all six packages, and the source exemption holds.
58 lines
1.8 KiB
Python
58 lines
1.8 KiB
Python
import os
|
|
import tempfile
|
|
import time
|
|
from collections import defaultdict
|
|
from functools import partial
|
|
|
|
from langgraph.checkpoint.base import (
|
|
ChannelVersions,
|
|
Checkpoint,
|
|
CheckpointMetadata,
|
|
SerializerProtocol,
|
|
)
|
|
from langgraph.checkpoint.memory import InMemorySaver, PersistentDict
|
|
from langgraph.pregel._checkpoint import copy_checkpoint
|
|
|
|
|
|
class MemorySaverAssertImmutable(InMemorySaver):
|
|
storage_for_copies: defaultdict[str, dict[str, dict[str, Checkpoint]]]
|
|
|
|
def __init__(
|
|
self,
|
|
*,
|
|
serde: SerializerProtocol | None = None,
|
|
put_sleep: float | None = None,
|
|
) -> None:
|
|
_, filename = tempfile.mkstemp()
|
|
super().__init__(
|
|
serde=serde, factory=partial(PersistentDict, filename=filename)
|
|
)
|
|
self.storage_for_copies = defaultdict(lambda: defaultdict(dict))
|
|
self.put_sleep = put_sleep
|
|
self.stack.callback(os.remove, filename)
|
|
|
|
def put(
|
|
self,
|
|
config: dict,
|
|
checkpoint: Checkpoint,
|
|
metadata: CheckpointMetadata,
|
|
new_versions: ChannelVersions,
|
|
) -> None:
|
|
if self.put_sleep:
|
|
time.sleep(self.put_sleep)
|
|
# assert checkpoint hasn't been modified since last written
|
|
thread_id = config["configurable"]["thread_id"]
|
|
checkpoint_ns = config["configurable"]["checkpoint_ns"]
|
|
if saved := super().get(config):
|
|
assert (
|
|
self.serde.loads_typed(
|
|
self.storage_for_copies[thread_id][checkpoint_ns][saved["id"]]
|
|
)
|
|
== saved
|
|
)
|
|
self.storage_for_copies[thread_id][checkpoint_ns][checkpoint["id"]] = (
|
|
self.serde.dumps_typed(copy_checkpoint(checkpoint))
|
|
)
|
|
# call super to write checkpoint
|
|
return super().put(config, checkpoint, metadata, new_versions)
|