From b266b73ab90e92b1ea748ff5804e3802d594fec4 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 14 Nov 2023 11:27:31 +0000 Subject: [PATCH] Make channel managers aware of checkpoints --- permchain/channels/base.py | 18 ++++++++++++++---- permchain/pregel/__init__.py | 4 ++-- 2 files changed, 16 insertions(+), 6 deletions(-) diff --git a/permchain/channels/base.py b/permchain/channels/base.py index ee53993f5..1983147f5 100644 --- a/permchain/channels/base.py +++ b/permchain/channels/base.py @@ -80,11 +80,13 @@ class BaseChannel(Generic[Value, Update, Checkpoint], ABC): @contextmanager def ChannelsManager( - channels: Mapping[str, BaseChannel] + channels: Mapping[str, BaseChannel], + checkpoint: Optional[Mapping[str, Any]], ) -> 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.empty() for k, v in channels.items()} + checkpoint = checkpoint or {} + empty = {k: v.empty(checkpoint.get(k)) for k, v in channels.items()} try: yield {k: v.__enter__() for k, v in empty.items()} finally: @@ -94,12 +96,20 @@ def ChannelsManager( @asynccontextmanager async def AsyncChannelsManager( - channels: Mapping[str, BaseChannel] + channels: Mapping[str, BaseChannel], + checkpoint: Optional[Mapping[str, Any]], ) -> AsyncGenerator[Mapping[str, BaseChannel], None]: """Manage channels for the lifetime of a Pregel invocation (multiple steps).""" - empty = {k: v.aempty() for k, v in channels.items()} + checkpoint = checkpoint or {} + empty = {k: v.aempty(checkpoint.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(channels: Mapping[str, BaseChannel]) -> Mapping[str, Any]: + """Create a checkpoint for the given channels.""" + checkpoint = {k: v.checkpoint() for k, v in channels.items()} + return {k: v for k, v in checkpoint.items() if v is not None} diff --git a/permchain/pregel/__init__.py b/permchain/pregel/__init__.py index dcf48c73c..374d390c1 100644 --- a/permchain/pregel/__init__.py +++ b/permchain/pregel/__init__.py @@ -164,7 +164,7 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): ) -> Iterator[dict[str, Any] | Any]: processes = {**self.chains} # TODO this is where we'd restore from checkpoint - with ChannelsManager(self.channels) as channels, get_executor_for_config( + with ChannelsManager(self.channels, None) as channels, get_executor_for_config( config ) as executor: next_tasks = _apply_writes_and_prepare_next_tasks( @@ -243,7 +243,7 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): ) -> AsyncIterator[dict[str, Any] | Any]: processes = {**self.chains} # TODO this is where we'd restore from checkpoint - async with AsyncChannelsManager(self.channels) as channels: + async with AsyncChannelsManager(self.channels, None) as channels: next_tasks = _apply_writes_and_prepare_next_tasks( processes, channels,