diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index f0d9e46eb..7081402c7 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -1101,6 +1101,7 @@ class CompiledStateGraph( def _migrate_checkpoint(self, checkpoint: Checkpoint) -> None: """Migrate a checkpoint to new channel layout.""" + super()._migrate_checkpoint(checkpoint) values = checkpoint["channel_values"] versions = checkpoint["channel_versions"] diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 3ee6420fa..aef5c7ce3 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -908,7 +908,12 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou def _migrate_checkpoint(self, checkpoint: Checkpoint) -> None: """Migrate a saved checkpoint to new channel layout.""" - pass + if checkpoint["v"] < 4 and checkpoint.get("pending_sends"): + pending_sends: list[Send] = checkpoint.pop("pending_sends") + checkpoint["channel_values"][TASKS] = pending_sends + checkpoint["channel_versions"][TASKS] = max( + checkpoint["channel_versions"].values() + ) def _prepare_state_snapshot( self, diff --git a/libs/langgraph/tests/conftest.py b/libs/langgraph/tests/conftest.py index c6226cb72..269c3e66b 100644 --- a/libs/langgraph/tests/conftest.py +++ b/libs/langgraph/tests/conftest.py @@ -12,6 +12,7 @@ from langgraph.checkpoint.base import BaseCheckpointSaver from langgraph.store.base import BaseStore from tests.conftest_checkpointer import ( _checkpointer_memory, + _checkpointer_memory_migrate_sends, _checkpointer_postgres, _checkpointer_postgres_aio, _checkpointer_postgres_aio_pipe, @@ -125,6 +126,7 @@ async def async_store(request: pytest.FixtureRequest) -> AsyncIterator[BaseStore if NO_DOCKER else [ "memory", + "memory_migrate_sends", "sqlite", "sqlite_aes", "postgres", @@ -139,6 +141,9 @@ def sync_checkpointer( if checkpointer_name == "memory": with _checkpointer_memory() as checkpointer: yield checkpointer + elif checkpointer_name == "memory_migrate_sends": + with _checkpointer_memory_migrate_sends() as checkpointer: + yield checkpointer elif checkpointer_name == "sqlite": with _checkpointer_sqlite() as checkpointer: yield checkpointer diff --git a/libs/langgraph/tests/conftest_checkpointer.py b/libs/langgraph/tests/conftest_checkpointer.py index ba15a8251..deb8802f4 100644 --- a/libs/langgraph/tests/conftest_checkpointer.py +++ b/libs/langgraph/tests/conftest_checkpointer.py @@ -14,7 +14,10 @@ from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver pytest.register_assert_rewrite("tests.memory_assert") -from tests.memory_assert import MemorySaverAssertImmutable # noqa: E402 +from tests.memory_assert import ( # noqa: E402 + MemorySaverAssertImmutable, + MemorySaverNeedsPendingSendsMigration, +) DEFAULT_POSTGRES_URI = "postgres://postgres:postgres@localhost:5442/" @@ -24,6 +27,11 @@ def _checkpointer_memory(): yield MemorySaverAssertImmutable() +@contextmanager +def _checkpointer_memory_migrate_sends(): + yield MemorySaverNeedsPendingSendsMigration() + + @contextmanager def _checkpointer_sqlite(): with SqliteSaver.from_conn_string(":memory:") as checkpointer: @@ -187,6 +195,7 @@ async def _checkpointer_postgres_aio_pool(): __all__ = [ "_checkpointer_memory", + "_checkpointer_memory_migrate_sends", "_checkpointer_sqlite", "_checkpointer_sqlite_aes", "_checkpointer_postgres", diff --git a/libs/langgraph/tests/memory_assert.py b/libs/langgraph/tests/memory_assert.py index 3a1ef4536..43eb1aee6 100644 --- a/libs/langgraph/tests/memory_assert.py +++ b/libs/langgraph/tests/memory_assert.py @@ -7,6 +7,7 @@ from typing import Any, Optional from langchain_core.runnables import RunnableConfig from langgraph.checkpoint.base import ( + BaseCheckpointSaver, ChannelVersions, Checkpoint, CheckpointMetadata, @@ -14,6 +15,7 @@ from langgraph.checkpoint.base import ( SerializerProtocol, ) from langgraph.checkpoint.memory import InMemorySaver, PersistentDict +from langgraph.constants import TASKS class NoopSerializer(SerializerProtocol): @@ -24,6 +26,28 @@ class NoopSerializer(SerializerProtocol): return "type", obj +class MemorySaverNeedsPendingSendsMigration(BaseCheckpointSaver): + def __init__(self) -> None: + self.saver = InMemorySaver() + + def __getattribute__(self, name): + if name in ("saver", "__class__", "get_tuple"): + return object.__getattribute__(self, name) + return getattr(self.saver, name) + + def get_tuple(self, config): + if tup := self.saver.get_tuple(config): + if tup.checkpoint["v"] == 4 and tup.checkpoint["channel_values"].get(TASKS): + tup.checkpoint["v"] = 3 + tup.checkpoint["pending_sends"] = tup.checkpoint["channel_values"].pop( + TASKS + ) + tup.checkpoint["channel_versions"].pop(TASKS) + for seen in tup.checkpoint["versions_seen"].values(): + seen.pop(TASKS, None) + return tup + + class MemorySaverAssertImmutable(InMemorySaver): storage_for_copies: defaultdict[str, dict[str, dict[str, Checkpoint]]]