Make channel managers aware of checkpoints

This commit is contained in:
Nuno Campos
2023-11-14 11:27:31 +00:00
parent 56a38ee859
commit b266b73ab9
2 changed files with 16 additions and 6 deletions
+14 -4
View File
@@ -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}
+2 -2
View File
@@ -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,