Preparation for 0.5 release: langgraph-checkpoint (#5124)

Prepare langgraph-checkpoint for 0.5

- Given we have no upper bound on langgraph-checkpoint dep need to undo all changes in langgraph-checkpoint that might break previous versions of langgraph
This commit is contained in:
Nuno Campos
2025-06-16 21:57:11 +00:00
committed by GitHub
parent 4fec8e9dec
commit 1134017d07
22 changed files with 109 additions and 219 deletions
@@ -168,7 +168,7 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
checkpoint["channel_versions"][TASKS] = (
max(checkpoint["channel_versions"].values())
if checkpoint["channel_versions"]
else self.get_next_version(None)
else self.get_next_version(None, None)
)
def _load_blobs(
@@ -246,7 +246,7 @@ class BasePostgresSaver(BaseCheckpointSaver[str]):
for idx, (channel, value) in enumerate(writes)
]
def get_next_version(self, current: str | None) -> str:
def get_next_version(self, current: str | None, channel: None) -> str:
if current is None:
current_v = 0
elif isinstance(current, int):
@@ -1,53 +0,0 @@
from __future__ import annotations
from collections.abc import Mapping
from datetime import datetime, timezone
from typing import Any, Protocol
from langgraph.checkpoint.base import Checkpoint, EmptyChannelError
from langgraph.checkpoint.base.id import uuid6
class ChannelProtocol(Protocol):
def checkpoint(self) -> Any | None: ...
def empty_checkpoint() -> Checkpoint:
return Checkpoint(
v=1,
id=str(uuid6(clock_seq=-2)),
ts=datetime.now(timezone.utc).isoformat(),
channel_values={},
channel_versions={},
versions_seen={},
)
def create_checkpoint(
checkpoint: Checkpoint,
channels: Mapping[str, ChannelProtocol] | None,
step: int,
*,
id: str | None = None,
) -> Checkpoint:
"""Create a checkpoint for the given channels."""
ts = datetime.now(timezone.utc).isoformat()
if channels is None:
values = checkpoint["channel_values"]
else:
values = {}
for k, v in channels.items():
if k not in checkpoint["channel_versions"]:
continue
try:
values[k] = v.checkpoint()
except EmptyChannelError:
pass
return Checkpoint(
v=1,
ts=ts,
id=id or str(uuid6(clock_seq=step)),
channel_values=values,
channel_versions=checkpoint["channel_versions"],
versions_seen=checkpoint["versions_seen"],
)
+2 -1
View File
@@ -14,13 +14,14 @@ from langgraph.checkpoint.base import (
EXCLUDED_METADATA_KEYS,
Checkpoint,
CheckpointMetadata,
create_checkpoint,
empty_checkpoint,
)
from langgraph.checkpoint.postgres.aio import (
AsyncPostgresSaver,
AsyncShallowPostgresSaver,
)
from langgraph.checkpoint.serde.types import TASKS
from tests.checkpoint_utils import create_checkpoint, empty_checkpoint
from tests.conftest import DEFAULT_POSTGRES_URI
+2 -1
View File
@@ -15,10 +15,11 @@ from langgraph.checkpoint.base import (
EXCLUDED_METADATA_KEYS,
Checkpoint,
CheckpointMetadata,
create_checkpoint,
empty_checkpoint,
)
from langgraph.checkpoint.postgres import PostgresSaver, ShallowPostgresSaver
from langgraph.checkpoint.serde.types import TASKS
from tests.checkpoint_utils import create_checkpoint, empty_checkpoint
from tests.conftest import DEFAULT_POSTGRES_URI
@@ -536,7 +536,7 @@ class SqliteSaver(BaseCheckpointSaver[str]):
"""
raise NotImplementedError(_AIO_ERROR_MSG)
def get_next_version(self, current: str | None) -> str:
def get_next_version(self, current: str | None, channel: None) -> str:
"""Generate the next version ID for a channel.
This method creates a new version identifier for a channel based on its current version.
@@ -591,7 +591,7 @@ class AsyncSqliteSaver(BaseCheckpointSaver[str]):
)
await self.conn.commit()
def get_next_version(self, current: str | None) -> str:
def get_next_version(self, current: str | None, channel: None) -> str:
"""Generate the next version ID for a channel.
This method creates a new version identifier for a channel based on its current version.
@@ -1,53 +0,0 @@
from __future__ import annotations
from collections.abc import Mapping
from datetime import datetime, timezone
from typing import Any, Protocol
from langgraph.checkpoint.base import Checkpoint, EmptyChannelError
from langgraph.checkpoint.base.id import uuid6
class ChannelProtocol(Protocol):
def checkpoint(self) -> Any | None: ...
def empty_checkpoint() -> Checkpoint:
return Checkpoint(
v=1,
id=str(uuid6(clock_seq=-2)),
ts=datetime.now(timezone.utc).isoformat(),
channel_values={},
channel_versions={},
versions_seen={},
)
def create_checkpoint(
checkpoint: Checkpoint,
channels: Mapping[str, ChannelProtocol] | None,
step: int,
*,
id: str | None = None,
) -> Checkpoint:
"""Create a checkpoint for the given channels."""
ts = datetime.now(timezone.utc).isoformat()
if channels is None:
values = checkpoint["channel_values"]
else:
values = {}
for k, v in channels.items():
if k not in checkpoint["channel_versions"]:
continue
try:
values[k] = v.checkpoint()
except EmptyChannelError:
pass
return Checkpoint(
v=1,
ts=ts,
id=id or str(uuid6(clock_seq=step)),
channel_values=values,
channel_versions=checkpoint["channel_versions"],
versions_seen=checkpoint["versions_seen"],
)
@@ -6,9 +6,10 @@ from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import (
Checkpoint,
CheckpointMetadata,
create_checkpoint,
empty_checkpoint,
)
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
from tests.checkpoint_utils import create_checkpoint, empty_checkpoint
class TestAsyncSqliteSaver:
+2 -1
View File
@@ -6,10 +6,11 @@ from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import (
Checkpoint,
CheckpointMetadata,
create_checkpoint,
empty_checkpoint,
)
from langgraph.checkpoint.sqlite import SqliteSaver
from langgraph.checkpoint.sqlite.utils import _metadata_predicate, search_where
from tests.checkpoint_utils import create_checkpoint, empty_checkpoint
class TestSqliteSaver:
@@ -1,10 +1,8 @@
from __future__ import annotations
from collections.abc import AsyncIterator, Iterator, Sequence
from inspect import signature
from collections.abc import AsyncIterator, Iterator, Mapping, Sequence
from typing import ( # noqa: UP035
Any,
ClassVar,
Generic,
Literal,
NamedTuple,
@@ -15,6 +13,7 @@ from typing import ( # noqa: UP035
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base.id import uuid6
from langgraph.checkpoint.serde.base import SerializerProtocol, maybe_add_typed_methods
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
from langgraph.checkpoint.serde.types import (
@@ -22,6 +21,7 @@ from langgraph.checkpoint.serde.types import (
INTERRUPT,
RESUME,
SCHEDULED,
ChannelProtocol,
)
V = TypeVar("V", int, float, str)
@@ -91,6 +91,7 @@ def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint:
channel_values=checkpoint["channel_values"].copy(),
channel_versions=checkpoint["channel_versions"].copy(),
versions_seen={k: v.copy() for k, v in checkpoint["versions_seen"].items()},
pending_sends=checkpoint.get("pending_sends", []).copy(),
)
@@ -118,19 +119,8 @@ class BaseCheckpointSaver(Generic[V]):
versions to avoid blocking the main thread.
"""
_get_next_version_legacy: ClassVar[bool] = False
"""Flag indicating if get_next_version method is legacy (takes two parameters)."""
serde: SerializerProtocol = JsonPlusSerializer()
def __init_subclass__(cls) -> None:
cls._get_next_version_legacy = (
len(signature(cls.get_next_version).parameters) > 2 # self + current
if hasattr(cls, "get_next_version")
else False
)
return super().__init_subclass__()
def __init__(
self,
*,
@@ -138,6 +128,15 @@ class BaseCheckpointSaver(Generic[V]):
) -> None:
self.serde = maybe_add_typed_methods(serde or self.serde)
@property
def config_specs(self) -> list:
"""Define the configuration options for the checkpoint saver.
Returns:
list: List of configuration field specs.
"""
return []
def get(self, config: RunnableConfig) -> Checkpoint | None:
"""Fetch a checkpoint using the given configuration.
@@ -347,7 +346,7 @@ class BaseCheckpointSaver(Generic[V]):
"""
raise NotImplementedError
def get_next_version(self, current: V | None) -> V:
def get_next_version(self, current: V | None, channel: None) -> V:
"""Generate the next version ID for a channel.
Default is to use integer versions, incrementing by 1. If you override, you can use str/int/float versions,
@@ -355,6 +354,7 @@ class BaseCheckpointSaver(Generic[V]):
Args:
current: The current version identifier (int, float, or str).
channel: Deprecated argument, kept for backwards compatibility.
Returns:
V: The next version identifier, which must be increasing.
@@ -417,3 +417,54 @@ EXCLUDED_METADATA_KEYS = {
"checkpoint_ns",
"checkpoint_map",
}
# --- below are deprecated utilities used by past versions of LangGraph ---
LATEST_VERSION = 2
def empty_checkpoint() -> Checkpoint:
from datetime import datetime, timezone
return Checkpoint(
v=LATEST_VERSION,
id=str(uuid6(clock_seq=-2)),
ts=datetime.now(timezone.utc).isoformat(),
channel_values={},
channel_versions={},
versions_seen={},
pending_sends=[],
)
def create_checkpoint(
checkpoint: Checkpoint,
channels: Mapping[str, ChannelProtocol] | None,
step: int,
*,
id: str | None = None,
) -> Checkpoint:
"""Create a checkpoint for the given channels."""
from datetime import datetime, timezone
ts = datetime.now(timezone.utc).isoformat()
if channels is None:
values = checkpoint["channel_values"]
else:
values = {}
for k, v in channels.items():
if k not in checkpoint["channel_versions"]:
continue
try:
values[k] = v.checkpoint()
except EmptyChannelError:
pass
return Checkpoint(
v=LATEST_VERSION,
ts=ts,
id=id or str(uuid6(clock_seq=step)),
channel_values=values,
channel_versions=checkpoint["channel_versions"],
versions_seen=checkpoint["versions_seen"],
pending_sends=checkpoint.get("pending_sends", []),
)
@@ -512,7 +512,7 @@ class InMemorySaver(
"""
return self.delete_thread(thread_id)
def get_next_version(self, current: str | None) -> str:
def get_next_version(self, current: str | None, channel: None) -> str:
if current is None:
current_v = 0
elif isinstance(current, int):
-53
View File
@@ -1,53 +0,0 @@
from __future__ import annotations
from collections.abc import Mapping
from datetime import datetime, timezone
from typing import Any, Protocol
from langgraph.checkpoint.base import Checkpoint, EmptyChannelError
from langgraph.checkpoint.base.id import uuid6
class ChannelProtocol(Protocol):
def checkpoint(self) -> Any | None: ...
def empty_checkpoint() -> Checkpoint:
return Checkpoint(
v=1,
id=str(uuid6(clock_seq=-2)),
ts=datetime.now(timezone.utc).isoformat(),
channel_values={},
channel_versions={},
versions_seen={},
)
def create_checkpoint(
checkpoint: Checkpoint,
channels: Mapping[str, ChannelProtocol] | None,
step: int,
*,
id: str | None = None,
) -> Checkpoint:
"""Create a checkpoint for the given channels."""
ts = datetime.now(timezone.utc).isoformat()
if channels is None:
values = checkpoint["channel_values"]
else:
values = {}
for k, v in channels.items():
if k not in checkpoint["channel_versions"]:
continue
try:
values[k] = v.checkpoint()
except EmptyChannelError:
pass
return Checkpoint(
v=1,
ts=ts,
id=id or str(uuid6(clock_seq=step)),
channel_values=values,
channel_versions=checkpoint["channel_versions"],
versions_seen=checkpoint["versions_seen"],
)
+1 -3
View File
@@ -6,12 +6,10 @@ from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import (
Checkpoint,
CheckpointMetadata,
)
from langgraph.checkpoint.memory import InMemorySaver
from tests.checkpoint_utils import (
create_checkpoint,
empty_checkpoint,
)
from langgraph.checkpoint.memory import InMemorySaver
class TestMemorySaver:
+1 -1
View File
@@ -32,7 +32,6 @@ from langgraph.checkpoint.base import (
BaseCheckpointSaver,
Checkpoint,
CheckpointTuple,
copy_checkpoint,
)
from langgraph.config import get_config
from langgraph.constants import (
@@ -79,6 +78,7 @@ from langgraph.pregel.algo import (
from langgraph.pregel.call import identifier
from langgraph.pregel.checkpoint import (
channels_from_checkpoint,
copy_checkpoint,
create_checkpoint,
empty_checkpoint,
)
+4 -3
View File
@@ -83,7 +83,7 @@ from langgraph.types import (
)
from langgraph.utils.config import merge_configs, patch_config
GetNextVersion = Callable[[Optional[V]], V]
GetNextVersion = Callable[[Optional[V], None], V]
SUPPORTS_EXC_NOTES = sys.version_info >= (3, 11)
@@ -214,7 +214,7 @@ def local_read(
return values
def increment(current: int | None) -> int:
def increment(current: int | None, channel: None) -> int:
"""Default channel versioning function, increments the current int version."""
return current + 1 if current is not None else 1
@@ -265,7 +265,8 @@ def apply_writes(
next_version = get_next_version(
max(checkpoint["channel_versions"].values())
if checkpoint["channel_versions"]
else None
else None,
None,
)
# Consume all channels that were read
@@ -71,3 +71,14 @@ def channels_from_checkpoint(
},
managed_specs,
)
def copy_checkpoint(checkpoint: Checkpoint) -> Checkpoint:
return Checkpoint(
v=checkpoint["v"],
ts=checkpoint["ts"],
id=checkpoint["id"],
channel_values=checkpoint["channel_values"].copy(),
channel_versions=checkpoint["channel_versions"].copy(),
versions_seen={k: v.copy() for k, v in checkpoint["versions_seen"].items()},
)
+3 -16
View File
@@ -29,7 +29,6 @@ from typing_extensions import ParamSpec, Self
from langgraph.cache.base import BaseCache
from langgraph.channels.base import BaseChannel
from langgraph.channels.last_value import LastValue
from langgraph.checkpoint.base import (
EXCLUDED_METADATA_KEYS,
WRITES_IDX_MAP,
@@ -39,7 +38,6 @@ from langgraph.checkpoint.base import (
CheckpointMetadata,
CheckpointTuple,
PendingWrite,
copy_checkpoint,
)
from langgraph.constants import (
CONF,
@@ -86,6 +84,7 @@ from langgraph.pregel.algo import (
)
from langgraph.pregel.checkpoint import (
channels_from_checkpoint,
copy_checkpoint,
create_checkpoint,
empty_checkpoint,
)
@@ -963,13 +962,7 @@ class SyncPregelLoop(PregelLoop, AbstractContextManager):
)
self.stack = ExitStack()
if checkpointer:
if checkpointer._get_next_version_legacy:
empty_channel: LastValue[Any] = LastValue(Any)
self.checkpointer_get_next_version = (
lambda c: checkpointer.get_next_version(c, empty_channel) # type: ignore[call-arg]
)
else:
self.checkpointer_get_next_version = checkpointer.get_next_version
self.checkpointer_get_next_version = checkpointer.get_next_version
self.checkpointer_put_writes = checkpointer.put_writes
self.checkpointer_put_writes_accepts_task_path = (
signature(checkpointer.put_writes).parameters.get("task_path")
@@ -1142,13 +1135,7 @@ class AsyncPregelLoop(PregelLoop, AbstractAsyncContextManager):
)
self.stack = AsyncExitStack()
if checkpointer:
if checkpointer._get_next_version_legacy:
empty_channel: LastValue[Any] = LastValue(Any)
self.checkpointer_get_next_version = (
lambda c: checkpointer.get_next_version(c, empty_channel) # type: ignore[call-arg]
)
else:
self.checkpointer_get_next_version = checkpointer.get_next_version
self.checkpointer_get_next_version = checkpointer.get_next_version
self.checkpointer_put_writes = checkpointer.aput_writes
self.checkpointer_put_writes_accepts_task_path = (
signature(checkpointer.aput_writes).parameters.get("task_path")
@@ -7,12 +7,9 @@ from typing import Annotated, Literal, Optional, Union
import pytest
from typing_extensions import TypedDict
from langgraph.checkpoint.base import (
BaseCheckpointSaver,
CheckpointTuple,
copy_checkpoint,
)
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointTuple
from langgraph.graph.state import StateGraph
from langgraph.pregel.checkpoint import copy_checkpoint
from langgraph.types import Command, Interrupt, PregelTask, StateSnapshot, interrupt
from langgraph.utils.config import patch_configurable
from tests.any_int import AnyInt
+1 -1
View File
@@ -159,7 +159,7 @@ def test_checkpoint_errors() -> None:
raise ValueError("Faulty put_writes")
class FaultyVersionCheckpointer(InMemorySaver):
def get_next_version(self, current: Optional[int]) -> int:
def get_next_version(self, current: Optional[int], channel: None) -> int:
raise ValueError("Faulty get_next_version")
def logic(inp: str) -> str:
+1 -1
View File
@@ -103,7 +103,7 @@ async def test_checkpoint_errors() -> None:
raise ValueError("Faulty put_writes")
class FaultyVersionCheckpointer(InMemorySaver):
def get_next_version(self, current: Optional[int]) -> int:
def get_next_version(self, current: Optional[int], channel: None) -> int:
raise ValueError("Faulty get_next_version")
def logic(inp: str) -> str:
@@ -591,7 +591,7 @@ def create_react_agent(
workflow = StateGraph(state_schema, config_schema=config_schema)
workflow.add_node(
"agent",
RunnableCallable(call_model, acall_model), # type: ignore[call-overload]
RunnableCallable(call_model, acall_model),
input_schema=input_schema,
)
if pre_model_hook is not None:
@@ -610,7 +610,7 @@ def create_react_agent(
if response_format is not None:
workflow.add_node(
"generate_structured_response",
RunnableCallable( # type: ignore[call-overload]
RunnableCallable(
generate_structured_response,
agenerate_structured_response,
),
@@ -660,10 +660,10 @@ def create_react_agent(
# Define the two nodes we will cycle between
workflow.add_node(
"agent",
RunnableCallable(call_model, acall_model), # type: ignore[call-overload]
RunnableCallable(call_model, acall_model),
input_schema=input_schema,
)
workflow.add_node("tools", tool_node) # type: ignore[call-overload]
workflow.add_node("tools", tool_node)
# Optionally add a pre-model hook node that will be called
# every time before the "agent" (LLM-calling node)
@@ -693,7 +693,7 @@ def create_react_agent(
if response_format is not None:
workflow.add_node(
"generate_structured_response",
RunnableCallable( # type: ignore[call-overload]
RunnableCallable(
generate_structured_response,
agenerate_structured_response,
),
+1 -1
View File
@@ -13,9 +13,9 @@ from langgraph.checkpoint.base import (
CheckpointMetadata,
CheckpointTuple,
SerializerProtocol,
copy_checkpoint,
)
from langgraph.checkpoint.memory import InMemorySaver, PersistentDict
from langgraph.pregel.checkpoint import copy_checkpoint
class NoopSerializer(SerializerProtocol):