From a8b14a4e04853d27ecc8098dd36fe49570eb0801 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Thu, 13 Jun 2024 17:49:18 -0700 Subject: [PATCH] Allow checkpointer to override logic that creates the next version for each channel --- langgraph/channels/base.py | 64 ----------------------------- langgraph/channels/manager.py | 65 ++++++++++++++++++++++++++++++ langgraph/checkpoint/base.py | 5 +++ langgraph/managed/few_shot.py | 2 +- langgraph/pregel/__init__.py | 61 +++++++++++++++++++++++----- tests/checkpoint/test_aiosqlite.py | 2 +- tests/checkpoint/test_memory.py | 2 +- tests/checkpoint/test_sqlite.py | 2 +- 8 files changed, 125 insertions(+), 78 deletions(-) create mode 100644 langgraph/channels/manager.py diff --git a/langgraph/channels/base.py b/langgraph/channels/base.py index 01d595464..b45d1548c 100644 --- a/langgraph/channels/base.py +++ b/langgraph/channels/base.py @@ -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", ] diff --git a/langgraph/channels/manager.py b/langgraph/channels/manager.py new file mode 100644 index 000000000..1c230ac00 --- /dev/null +++ b/langgraph/channels/manager.py @@ -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", []), + ) diff --git a/langgraph/checkpoint/base.py b/langgraph/checkpoint/base.py index a639bc336..3ef13b23f 100644 --- a/langgraph/checkpoint/base.py +++ b/langgraph/checkpoint/base.py @@ -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 diff --git a/langgraph/managed/few_shot.py b/langgraph/managed/few_shot.py index 9d46739eb..a8ba4da82 100644 --- a/langgraph/managed/few_shot.py +++ b/langgraph/managed/few_shot.py @@ -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 diff --git a/langgraph/pregel/__init__.py b/langgraph/pregel/__init__.py index 2eae9398e..f8fd0facc 100644 --- a/langgraph/pregel/__init__.py +++ b/langgraph/pregel/__init__.py @@ -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: diff --git a/tests/checkpoint/test_aiosqlite.py b/tests/checkpoint/test_aiosqlite.py index 99e24b722..abcf6cb71 100644 --- a/tests/checkpoint/test_aiosqlite.py +++ b/tests/checkpoint/test_aiosqlite.py @@ -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 diff --git a/tests/checkpoint/test_memory.py b/tests/checkpoint/test_memory.py index e41e53528..bac957180 100644 --- a/tests/checkpoint/test_memory.py +++ b/tests/checkpoint/test_memory.py @@ -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 diff --git a/tests/checkpoint/test_sqlite.py b/tests/checkpoint/test_sqlite.py index 6dd7dcbae..cd2aa8e1e 100644 --- a/tests/checkpoint/test_sqlite.py +++ b/tests/checkpoint/test_sqlite.py @@ -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,