diff --git a/libs/langgraph/langgraph/pregel/_checkpoint_writer.py b/libs/langgraph/langgraph/pregel/_checkpoint_writer.py index 33bf507d1..b1109a705 100644 --- a/libs/langgraph/langgraph/pregel/_checkpoint_writer.py +++ b/libs/langgraph/langgraph/pregel/_checkpoint_writer.py @@ -22,7 +22,21 @@ CHECKPOINT_BACKLOG_ENV_VAR = "LANGGRAPH_CHECKPOINT_BACKLOG" DEFAULT_CHECKPOINT_BACKLOG = 10 -@dataclass(frozen=True) +def _resolve_checkpoint_backlog() -> int: + if raw := os.getenv(CHECKPOINT_BACKLOG_ENV_VAR): + try: + backlog = int(raw) + except ValueError: + return DEFAULT_CHECKPOINT_BACKLOG + if backlog > 0: + return backlog + return DEFAULT_CHECKPOINT_BACKLOG + + +CHECKPOINT_BACKLOG = _resolve_checkpoint_backlog() + + +@dataclass(frozen=True, slots=True) class CheckpointRequest: config: RunnableConfig checkpoint: Checkpoint @@ -34,18 +48,9 @@ def _raise(error: BaseException) -> None: raise error -def resolve_checkpoint_backlog() -> int: - if raw := os.getenv(CHECKPOINT_BACKLOG_ENV_VAR): - try: - backlog = int(raw) - except ValueError: - return DEFAULT_CHECKPOINT_BACKLOG - if backlog > 0: - return backlog - return DEFAULT_CHECKPOINT_BACKLOG - - class SyncCheckpointWriter(AbstractContextManager): + __slots__ = ("put", "queue", "error", "closed", "thread") + def __init__( self, put: Callable[ @@ -55,9 +60,7 @@ class SyncCheckpointWriter(AbstractContextManager): max_pending: int | None = None, ) -> None: self.put = put - max_pending = ( - resolve_checkpoint_backlog() if max_pending is None else max_pending - ) + max_pending = CHECKPOINT_BACKLOG if max_pending is None else max_pending self.queue: queue.Queue[CheckpointRequest | None] = queue.Queue(max_pending) self.error: BaseException | None = None self.closed = False @@ -130,6 +133,8 @@ class SyncCheckpointWriter(AbstractContextManager): class AsyncCheckpointWriter(AbstractAsyncContextManager): + __slots__ = ("put", "queue", "error", "closed", "task") + def __init__( self, put: Callable[ @@ -139,9 +144,7 @@ class AsyncCheckpointWriter(AbstractAsyncContextManager): max_pending: int | None = None, ) -> None: self.put = put - max_pending = ( - resolve_checkpoint_backlog() if max_pending is None else max_pending - ) + max_pending = CHECKPOINT_BACKLOG if max_pending is None else max_pending self.queue: asyncio.Queue[CheckpointRequest | None] = asyncio.Queue(max_pending) self.error: BaseException | None = None self.closed = False diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 200e5b710..d056f6b0a 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -57,9 +57,7 @@ from langgraph.pregel import ( from langgraph.pregel._checkpoint_writer import ( CHECKPOINT_BACKLOG_ENV_VAR, DEFAULT_CHECKPOINT_BACKLOG, - AsyncCheckpointWriter, - SyncCheckpointWriter, - resolve_checkpoint_backlog, + _resolve_checkpoint_backlog, ) from langgraph.pregel._loop import SyncPregelLoop from langgraph.pregel._runner import PregelRunner @@ -3828,20 +3826,17 @@ def test_sync_durability_applies_checkpoint_backpressure() -> None: def test_checkpoint_backlog_uses_env_override(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv(CHECKPOINT_BACKLOG_ENV_VAR, "7") - - assert resolve_checkpoint_backlog() == 7 - assert SyncCheckpointWriter(lambda *_args: None).queue.maxsize == 7 - assert AsyncCheckpointWriter(lambda *_args: None).queue.maxsize == 7 + assert _resolve_checkpoint_backlog() == 7 def test_checkpoint_backlog_invalid_env_uses_default( monkeypatch: pytest.MonkeyPatch, ) -> None: monkeypatch.setenv(CHECKPOINT_BACKLOG_ENV_VAR, "not-an-int") - assert resolve_checkpoint_backlog() == DEFAULT_CHECKPOINT_BACKLOG + assert _resolve_checkpoint_backlog() == DEFAULT_CHECKPOINT_BACKLOG monkeypatch.setenv(CHECKPOINT_BACKLOG_ENV_VAR, "0") - assert resolve_checkpoint_backlog() == DEFAULT_CHECKPOINT_BACKLOG + assert _resolve_checkpoint_backlog() == DEFAULT_CHECKPOINT_BACKLOG def test_checkpoint_metadata(sync_checkpointer: BaseCheckpointSaver) -> None: