From ee7cab664ad1934283121e95bdd29f501d47e6e7 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Mon, 22 Jul 2024 17:52:02 -0700 Subject: [PATCH] Fix checkpoint put ordering issue --- libs/langgraph/langgraph/pregel/loop.py | 62 ++++++++++++++++++----- libs/langgraph/tests/memory_assert.py | 6 +++ libs/langgraph/tests/test_pregel.py | 2 +- libs/langgraph/tests/test_pregel_async.py | 2 +- 4 files changed, 57 insertions(+), 15 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index bfb8020b0..f9c00c62a 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -1,4 +1,5 @@ import asyncio +import concurrent.futures from collections import deque from contextlib import AsyncExitStack, ExitStack from types import TracebackType @@ -68,7 +69,7 @@ if TYPE_CHECKING: V = TypeVar("V") INPUT_DONE = object() INPUT_RESUMING = object() -EMPTY_LIST = [] +EMPTY_SEQ = () class PregelLoop: @@ -79,8 +80,16 @@ class PregelLoop: checkpointer_put_writes: Optional[ Callable[[RunnableConfig, Sequence[tuple[str, Any]], str], Any] ] - checkpointer_put: Optional[ - Callable[[RunnableConfig, Checkpoint, CheckpointMetadata], Any] + _checkpointer_put_after_previous: Optional[ + Callable[ + [ + Optional[concurrent.futures.Future], + RunnableConfig, + Sequence[tuple[str, Any]], + str, + ], + Any, + ] ] graph: "Pregel" @@ -143,9 +152,9 @@ class PregelLoop: def tick( self, *, - output_keys: Union[str, Sequence[str]] = EMPTY_LIST, - interrupt_after: Sequence[str] = EMPTY_LIST, - interrupt_before: Sequence[str] = EMPTY_LIST, + output_keys: Union[str, Sequence[str]] = EMPTY_SEQ, + interrupt_after: Sequence[str] = EMPTY_SEQ, + interrupt_before: Sequence[str] = EMPTY_SEQ, manager: Union[None, AsyncParentRunManager, ParentRunManager] = None, ) -> bool: """Execute a single iteration of the Pregel loop. @@ -294,7 +303,7 @@ class PregelLoop: # assign step metadata["step"] = self.step # bail if no checkpointer - if self.checkpointer_put is not None: + if self._checkpointer_put_after_previous is not None: # create new checkpoint self.checkpoint_metadata = metadata self.checkpoint = create_checkpoint( @@ -308,8 +317,11 @@ class PregelLoop: id=self.config["configurable"]["thread_ts"] if self.is_nested else None, ) # save it, without blocking - self.submit( - self.checkpointer_put, + # if there's a previous checkpoint save in progress, wait for it + # ensuring checkpointers receive checkpoints in order + self._put_checkpoint_fut = self.submit( + self._checkpointer_put_after_previous, + getattr(self, "_put_checkpoint_fut", None), self.checkpoint_config, copy_checkpoint(self.checkpoint), self.checkpoint_metadata, @@ -360,11 +372,23 @@ class SyncPregelLoop(PregelLoop, ContextManager): if checkpointer: self.checkpointer_get_next_version = checkpointer.get_next_version self.checkpointer_put_writes = checkpointer.put_writes - self.checkpointer_put = checkpointer.put else: self.checkpointer_get_next_version = increment + self._checkpointer_put_after_previous = None self.checkpointer_put_writes = None - self.checkpointer_put = None + + def _checkpointer_put_after_previous( + self, + prev: Optional[concurrent.futures.Future], + config: RunnableConfig, + checkpoint: Checkpoint, + metadata: CheckpointMetadata, + ) -> RunnableConfig: + try: + if prev is not None: + prev.result() + finally: + self.checkpointer.put(config, checkpoint, metadata) # context manager @@ -424,11 +448,23 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): if checkpointer: self.checkpointer_get_next_version = checkpointer.get_next_version self.checkpointer_put_writes = checkpointer.aput_writes - self.checkpointer_put = checkpointer.aput else: self.checkpointer_get_next_version = increment + self._checkpointer_put_after_previous = None self.checkpointer_put_writes = None - self.checkpointer_put = None + + async def _checkpointer_put_after_previous( + self, + prev: Optional[asyncio.Task], + config: RunnableConfig, + checkpoint: Checkpoint, + metadata: CheckpointMetadata, + ) -> RunnableConfig: + try: + if prev is not None: + await prev + finally: + await self.checkpointer.aput(config, checkpoint, metadata) # context manager diff --git a/libs/langgraph/tests/memory_assert.py b/libs/langgraph/tests/memory_assert.py index 6709a404a..c0ede2f20 100644 --- a/libs/langgraph/tests/memory_assert.py +++ b/libs/langgraph/tests/memory_assert.py @@ -30,9 +30,11 @@ class MemorySaverAssertImmutable(MemorySaver): self, *, serde: Optional[SerializerProtocol] = None, + put_sleep: Optional[float] = None, ) -> None: super().__init__(serde=serde) self.storage_for_copies = defaultdict(dict) + self.put_sleep = put_sleep def put( self, @@ -40,6 +42,10 @@ class MemorySaverAssertImmutable(MemorySaver): checkpoint: Checkpoint, metadata: Optional[CheckpointMetadata] = None, ) -> None: + if self.put_sleep: + import time + + time.sleep(self.put_sleep) # assert checkpoint hasn't been modified since last written thread_id = config["configurable"]["thread_id"] if saved := super().get(config): diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index b0a822209..d90a83ab8 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -7572,7 +7572,7 @@ def test_nested_graph(snapshot: SnapshotAssertion) -> None: @pytest.mark.parametrize( "checkpointer_fct", [ - lambda: MemorySaverAssertImmutable(), + lambda: MemorySaverAssertImmutable(put_sleep=0.2), lambda: SqliteSaver.from_conn_string(":memory:"), ], ids=[ diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index d5cdfc472..424d63af8 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -6083,7 +6083,7 @@ async def test_nested_graph(snapshot: SnapshotAssertion) -> None: @pytest.mark.parametrize( "checkpointer_fct", [ - lambda: MemorySaverAssertImmutable(), + lambda: MemorySaverAssertImmutable(put_sleep=0.2), lambda: AsyncSqliteSaver.from_conn_string(":memory:"), ], ids=[