Allow checkpointer to override logic that creates the next version for each channel

This commit is contained in:
Nuno Campos
2024-06-13 17:49:18 -07:00
parent 3f1e47d519
commit a8b14a4e04
8 changed files with 125 additions and 78 deletions
-64
View File
@@ -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",
]
+65
View File
@@ -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", []),
)
+5
View File
@@ -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
+1 -1
View File
@@ -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
+51 -10
View File
@@ -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 -1
View File
@@ -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 -1
View File
@@ -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 -1
View File
@@ -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,