mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-22 23:52:23 +02:00
Allow checkpointer to override logic that creates the next version for each channel
This commit is contained in:
@@ -1,12 +1,10 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from contextlib import asynccontextmanager, contextmanager
|
||||
from datetime import datetime, timezone
|
||||
from typing import (
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
Generator,
|
||||
Generic,
|
||||
Mapping,
|
||||
Optional,
|
||||
Sequence,
|
||||
TypeVar,
|
||||
@@ -14,8 +12,6 @@ from typing import (
|
||||
|
||||
from typing_extensions import Self
|
||||
|
||||
from langgraph.checkpoint.base import Checkpoint
|
||||
from langgraph.checkpoint.id import uuid6
|
||||
from langgraph.errors import EmptyChannelError, InvalidUpdateError
|
||||
|
||||
Value = TypeVar("Value")
|
||||
@@ -83,68 +79,8 @@ class BaseChannel(Generic[Value, Update, C], ABC):
|
||||
pass
|
||||
|
||||
|
||||
@contextmanager
|
||||
def ChannelsManager(
|
||||
channels: Mapping[str, BaseChannel],
|
||||
checkpoint: Checkpoint,
|
||||
) -> Generator[Mapping[str, BaseChannel], None, None]:
|
||||
"""Manage channels for the lifetime of a Pregel invocation (multiple steps)."""
|
||||
# TODO use https://docs.python.org/3/library/contextlib.html#contextlib.ExitStack
|
||||
empty = {
|
||||
k: v.from_checkpoint(checkpoint["channel_values"].get(k))
|
||||
for k, v in channels.items()
|
||||
}
|
||||
try:
|
||||
yield {k: v.__enter__() for k, v in empty.items()}
|
||||
finally:
|
||||
for v in empty.values():
|
||||
v.__exit__(None, None, None)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def AsyncChannelsManager(
|
||||
channels: Mapping[str, BaseChannel],
|
||||
checkpoint: Checkpoint,
|
||||
) -> AsyncGenerator[Mapping[str, BaseChannel], None]:
|
||||
"""Manage channels for the lifetime of a Pregel invocation (multiple steps)."""
|
||||
empty = {
|
||||
k: v.afrom_checkpoint(checkpoint["channel_values"].get(k))
|
||||
for k, v in channels.items()
|
||||
}
|
||||
try:
|
||||
yield {k: await v.__aenter__() for k, v in empty.items()}
|
||||
finally:
|
||||
for v in empty.values():
|
||||
await v.__aexit__(None, None, None)
|
||||
|
||||
|
||||
def create_checkpoint(
|
||||
checkpoint: Checkpoint, channels: Mapping[str, BaseChannel], step: int
|
||||
) -> Checkpoint:
|
||||
"""Create a checkpoint for the given channels."""
|
||||
ts = datetime.now(timezone.utc).isoformat()
|
||||
values: dict[str, Any] = {}
|
||||
for k, v in channels.items():
|
||||
try:
|
||||
values[k] = v.checkpoint()
|
||||
except EmptyChannelError:
|
||||
pass
|
||||
return Checkpoint(
|
||||
v=1,
|
||||
ts=ts,
|
||||
id=str(uuid6(clock_seq=step)),
|
||||
channel_values=values,
|
||||
channel_versions=checkpoint["channel_versions"],
|
||||
versions_seen=checkpoint["versions_seen"],
|
||||
pending_sends=checkpoint.get("pending_sends", []),
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"BaseChannel",
|
||||
"ChannelsManager",
|
||||
"AsyncChannelsManager",
|
||||
"create_checkpoint",
|
||||
"EmptyChannelError",
|
||||
"InvalidUpdateError",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
from contextlib import asynccontextmanager, contextmanager
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, AsyncGenerator, Generator, Mapping
|
||||
|
||||
from langgraph.channels.base import BaseChannel
|
||||
from langgraph.checkpoint.base import Checkpoint
|
||||
from langgraph.checkpoint.id import uuid6
|
||||
from langgraph.errors import EmptyChannelError
|
||||
|
||||
|
||||
@contextmanager
|
||||
def ChannelsManager(
|
||||
channels: Mapping[str, BaseChannel],
|
||||
checkpoint: Checkpoint,
|
||||
) -> Generator[Mapping[str, BaseChannel], None, None]:
|
||||
"""Manage channels for the lifetime of a Pregel invocation (multiple steps)."""
|
||||
# TODO use https://docs.python.org/3/library/contextlib.html#contextlib.ExitStack
|
||||
empty = {
|
||||
k: v.from_checkpoint(checkpoint["channel_values"].get(k))
|
||||
for k, v in channels.items()
|
||||
}
|
||||
try:
|
||||
yield {k: v.__enter__() for k, v in empty.items()}
|
||||
finally:
|
||||
for v in empty.values():
|
||||
v.__exit__(None, None, None)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def AsyncChannelsManager(
|
||||
channels: Mapping[str, BaseChannel],
|
||||
checkpoint: Checkpoint,
|
||||
) -> AsyncGenerator[Mapping[str, BaseChannel], None]:
|
||||
"""Manage channels for the lifetime of a Pregel invocation (multiple steps)."""
|
||||
empty = {
|
||||
k: v.afrom_checkpoint(checkpoint["channel_values"].get(k))
|
||||
for k, v in channels.items()
|
||||
}
|
||||
try:
|
||||
yield {k: await v.__aenter__() for k, v in empty.items()}
|
||||
finally:
|
||||
for v in empty.values():
|
||||
await v.__aexit__(None, None, None)
|
||||
|
||||
|
||||
def create_checkpoint(
|
||||
checkpoint: Checkpoint, channels: Mapping[str, BaseChannel], step: int
|
||||
) -> Checkpoint:
|
||||
"""Create a checkpoint for the given channels."""
|
||||
ts = datetime.now(timezone.utc).isoformat()
|
||||
values: dict[str, Any] = {}
|
||||
for k, v in channels.items():
|
||||
try:
|
||||
values[k] = v.checkpoint()
|
||||
except EmptyChannelError:
|
||||
pass
|
||||
return Checkpoint(
|
||||
v=1,
|
||||
ts=ts,
|
||||
id=str(uuid6(clock_seq=step)),
|
||||
channel_values=values,
|
||||
channel_versions=checkpoint["channel_versions"],
|
||||
versions_seen=checkpoint["versions_seen"],
|
||||
pending_sends=checkpoint.get("pending_sends", []),
|
||||
)
|
||||
@@ -15,6 +15,7 @@ from typing import (
|
||||
|
||||
from langchain_core.runnables import ConfigurableFieldSpec, RunnableConfig
|
||||
|
||||
from langgraph.channels.base import BaseChannel
|
||||
from langgraph.checkpoint.id import uuid6
|
||||
from langgraph.constants import Send
|
||||
from langgraph.serde.base import SerializerProtocol
|
||||
@@ -187,6 +188,7 @@ class BaseCheckpointSaver(ABC):
|
||||
self,
|
||||
config: Optional[RunnableConfig],
|
||||
*,
|
||||
filter: Optional[Dict[str, Any]] = None,
|
||||
before: Optional[RunnableConfig] = None,
|
||||
limit: Optional[int] = None,
|
||||
) -> AsyncIterator[CheckpointTuple]:
|
||||
@@ -200,3 +202,6 @@ class BaseCheckpointSaver(ABC):
|
||||
metadata: CheckpointMetadata,
|
||||
) -> RunnableConfig:
|
||||
raise NotImplementedError
|
||||
|
||||
def get_next_version(self, current: int, channel: BaseChannel) -> int:
|
||||
return current + 1
|
||||
|
||||
@@ -17,7 +17,7 @@ from typing import (
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from typing_extensions import Self
|
||||
|
||||
from langgraph.channels.base import AsyncChannelsManager, ChannelsManager
|
||||
from langgraph.channels.manager import AsyncChannelsManager, ChannelsManager
|
||||
from langgraph.managed.base import ConfiguredManagedValue, ManagedValue, V
|
||||
from langgraph.pregel import Pregel
|
||||
from langgraph.pregel.io import read_channels
|
||||
|
||||
@@ -52,10 +52,12 @@ from langchain_core.tracers._streaming import _StreamingCallbackHandler
|
||||
from typing_extensions import Self
|
||||
|
||||
from langgraph.channels.base import (
|
||||
AsyncChannelsManager,
|
||||
BaseChannel,
|
||||
ChannelsManager,
|
||||
EmptyChannelError,
|
||||
)
|
||||
from langgraph.channels.manager import (
|
||||
AsyncChannelsManager,
|
||||
ChannelsManager,
|
||||
create_checkpoint,
|
||||
)
|
||||
from langgraph.checkpoint.base import (
|
||||
@@ -564,7 +566,9 @@ class Pregel(
|
||||
),
|
||||
)
|
||||
# apply to checkpoint and save
|
||||
_apply_writes(checkpoint, channels, task.writes)
|
||||
_apply_writes(
|
||||
checkpoint, channels, task.writes, self.checkpointer.get_next_version
|
||||
)
|
||||
step = saved.metadata.get("step", -2) + 1 if saved else -1
|
||||
|
||||
# merge configurable fields with previous checkpoint config
|
||||
@@ -650,7 +654,9 @@ class Pregel(
|
||||
),
|
||||
)
|
||||
# apply to checkpoint and save
|
||||
_apply_writes(checkpoint, channels, task.writes)
|
||||
_apply_writes(
|
||||
checkpoint, channels, task.writes, self.checkpointer.get_next_version
|
||||
)
|
||||
step = saved.metadata.get("step", -2) + 1 if saved else -1
|
||||
|
||||
# merge configurable fields with previous checkpoint config
|
||||
@@ -849,7 +855,14 @@ class Pregel(
|
||||
for_execution=True,
|
||||
)
|
||||
# apply input writes
|
||||
_apply_writes(checkpoint, channels, input_writes)
|
||||
_apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
input_writes,
|
||||
self.checkpointer.get_next_version
|
||||
if self.checkpointer
|
||||
else _increment,
|
||||
)
|
||||
# save input checkpoint
|
||||
yield from put_checkpoint(
|
||||
{
|
||||
@@ -975,7 +988,14 @@ class Pregel(
|
||||
)
|
||||
|
||||
# apply writes to channels
|
||||
_apply_writes(checkpoint, channels, pending_writes)
|
||||
_apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
pending_writes,
|
||||
self.checkpointer.get_next_version
|
||||
if self.checkpointer
|
||||
else _increment,
|
||||
)
|
||||
|
||||
# yield values output
|
||||
if "values" in stream_modes:
|
||||
@@ -1178,7 +1198,14 @@ class Pregel(
|
||||
for_execution=True,
|
||||
)
|
||||
# apply input writes
|
||||
_apply_writes(checkpoint, channels, input_writes)
|
||||
_apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
input_writes,
|
||||
self.checkpointer.get_next_version
|
||||
if self.checkpointer
|
||||
else _increment,
|
||||
)
|
||||
# save input checkpoint
|
||||
for chunk in put_checkpoint(
|
||||
{"source": "input", "step": start, "writes": input}
|
||||
@@ -1304,7 +1331,14 @@ class Pregel(
|
||||
)
|
||||
|
||||
# apply writes to channels
|
||||
_apply_writes(checkpoint, channels, pending_writes)
|
||||
_apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
pending_writes,
|
||||
self.checkpointer.get_next_version
|
||||
if self.checkpointer
|
||||
else _increment,
|
||||
)
|
||||
|
||||
# yield current values
|
||||
if "values" in stream_modes:
|
||||
@@ -1534,7 +1568,7 @@ def _local_read(
|
||||
if fresh:
|
||||
checkpoint = create_checkpoint(checkpoint, channels, -1)
|
||||
with ChannelsManager(channels, checkpoint) as channels:
|
||||
_apply_writes(copy_checkpoint(checkpoint), channels, writes)
|
||||
_apply_writes(copy_checkpoint(checkpoint), channels, writes, _increment)
|
||||
return read_channels(channels, select)
|
||||
else:
|
||||
return read_channels(channels, select)
|
||||
@@ -1559,10 +1593,15 @@ def _local_write(
|
||||
commit(writes)
|
||||
|
||||
|
||||
def _increment(current: int, channel: BaseChannel) -> int:
|
||||
return current + 1
|
||||
|
||||
|
||||
def _apply_writes(
|
||||
checkpoint: Checkpoint,
|
||||
channels: Mapping[str, BaseChannel],
|
||||
pending_writes: Sequence[tuple[str, Any]],
|
||||
get_next_version: Callable[[int, BaseChannel], int],
|
||||
) -> None:
|
||||
if checkpoint["pending_sends"]:
|
||||
checkpoint["pending_sends"].clear()
|
||||
@@ -1591,7 +1630,9 @@ def _apply_writes(
|
||||
raise InvalidUpdateError(
|
||||
f"Invalid update for channel {chan} with values {vals}"
|
||||
) from e
|
||||
checkpoint["channel_versions"][chan] = max_version + 1
|
||||
checkpoint["channel_versions"][chan] = get_next_version(
|
||||
max_version, channels[chan]
|
||||
)
|
||||
updated_channels.add(chan)
|
||||
# Channels that weren't updated in this step are notified of a new step
|
||||
for chan in channels:
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import pytest
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
|
||||
from langgraph.channels.base import create_checkpoint
|
||||
from langgraph.channels.manager import create_checkpoint
|
||||
from langgraph.checkpoint.aiosqlite import AsyncSqliteSaver
|
||||
from langgraph.checkpoint.base import Checkpoint, CheckpointMetadata, empty_checkpoint
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import pytest
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
|
||||
from langgraph.channels.base import create_checkpoint
|
||||
from langgraph.channels.manager import create_checkpoint
|
||||
from langgraph.checkpoint.base import Checkpoint, CheckpointMetadata, empty_checkpoint
|
||||
from langgraph.checkpoint.memory import MemorySaver
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import pytest
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
|
||||
from langgraph.channels.base import create_checkpoint
|
||||
from langgraph.channels.manager import create_checkpoint
|
||||
from langgraph.checkpoint.base import Checkpoint, CheckpointMetadata, empty_checkpoint
|
||||
from langgraph.checkpoint.sqlite import (
|
||||
_AIO_ERROR_MSG,
|
||||
|
||||
Reference in New Issue
Block a user