Update context manager interface for channels, add channel that exposes a context manager

This commit is contained in:
Nuno Campos
2023-10-18 17:05:38 +01:00
parent c933f66a04
commit 0048704af9
3 changed files with 264 additions and 90 deletions
+130 -81
View File
@@ -1,15 +1,19 @@
import json
from abc import ABC, abstractmethod
from types import TracebackType
from contextlib import asynccontextmanager, contextmanager
from typing import (
AsyncContextManager,
AsyncGenerator,
Callable,
FrozenSet,
Generator,
Generic,
Optional,
Sequence,
Type,
TypeVar,
)
from typing import ContextManager as ContextManagerType
from typing_extensions import Self
@@ -28,38 +32,26 @@ class InvalidUpdateError(Exception):
class Channel(Generic[Value, Update], ABC):
@property
@abstractmethod
def ValueType(self) -> type[Value]:
def ValueType(self) -> Type[Value]:
"""The type of the value stored in the channel."""
@property
@abstractmethod
def UpdateType(self) -> type[Update]:
def UpdateType(self) -> Type[Update]:
"""The type of the update received by the channel."""
@contextmanager
@abstractmethod
def __enter__(self, checkpoint: Optional[str] = None) -> Self:
def _empty(self, checkpoint: Optional[str] = None) -> Generator[Self, None, None]:
"""Return a new identical channel, optionally initialized from a checkpoint."""
@abstractmethod
def __exit__(
self,
__exc_type: type[BaseException] | None,
__exc_value: BaseException | None,
__traceback: TracebackType | None,
) -> bool | None:
"""Clean up the channel, or deallocate resources as needed."""
...
async def __aenter__(self, checkpoint: Optional[str] = None) -> Self:
return self.__enter__(checkpoint)
async def __aexit__(
self,
__exc_type: type[BaseException] | None,
__exc_value: BaseException | None,
__traceback: TracebackType | None,
) -> None:
self.__exit__(__exc_type, __exc_value, __traceback)
@asynccontextmanager
async def _aempty(
self, checkpoint: Optional[str] = None
) -> AsyncGenerator[Self, None]:
"""Return a new identical channel, optionally initialized from a checkpoint."""
with self._empty(checkpoint) as value:
yield value
@abstractmethod
def _update(self, values: Sequence[Update]) -> None:
@@ -70,7 +62,7 @@ class Channel(Generic[Value, Update], ABC):
...
@abstractmethod
def _checkpoint(self) -> str:
def _checkpoint(self) -> str | None:
...
@@ -89,31 +81,27 @@ class BinaryOperatorAggregate(Generic[Value], Channel[Value, Value]):
self.operator = operator
@property
def ValueType(self) -> type[Value]:
def ValueType(self) -> Type[Value]:
"""The type of the value stored in the channel."""
return self.typ
@property
def UpdateType(self) -> type[Value]:
def UpdateType(self) -> Type[Value]:
"""The type of the update received by the channel."""
return self.typ
def __enter__(self, checkpoint: Optional[str] = None) -> Self:
@contextmanager
def _empty(self, checkpoint: Optional[str] = None) -> Generator[Self, None, None]:
empty = self.__class__(self.typ, self.operator)
if checkpoint is not None:
empty.value = json.loads(checkpoint)
return empty
def __exit__(
self,
__exc_type: type[BaseException] | None,
__exc_value: BaseException | None,
__traceback: TracebackType | None,
) -> None:
try:
del self.value
except AttributeError:
pass
yield empty
finally:
try:
del empty.value
except AttributeError:
pass
def _update(self, values: Sequence[Value]) -> None:
if not hasattr(self, "value"):
@@ -140,31 +128,27 @@ class LastValue(Generic[Value], Channel[Value, Value]):
self.typ = typ
@property
def ValueType(self) -> type[Value]:
def ValueType(self) -> Type[Value]:
"""The type of the value stored in the channel."""
return self.typ
@property
def UpdateType(self) -> type[Value]:
def UpdateType(self) -> Type[Value]:
"""The type of the update received by the channel."""
return self.typ
def __enter__(self, checkpoint: Optional[str] = None) -> Self:
@contextmanager
def _empty(self, checkpoint: Optional[str] = None) -> Generator[Self, None, None]:
empty = self.__class__(self.typ)
if checkpoint is not None:
empty.value = json.loads(checkpoint)
return empty
def __exit__(
self,
__exc_type: type[BaseException] | None,
__exc_value: BaseException | None,
__traceback: TracebackType | None,
) -> None:
try:
del self.value
except AttributeError:
pass
yield empty
finally:
try:
del empty.value
except AttributeError:
pass
def _update(self, values: Sequence[Value]) -> None:
if len(values) != 1:
@@ -189,31 +173,27 @@ class Inbox(Generic[Value], Channel[Sequence[Value], Value]):
self.typ = typ
@property
def ValueType(self) -> type[Sequence[Value]]:
def ValueType(self) -> Type[Sequence[Value]]:
"""The type of the value stored in the channel."""
return Sequence[self.typ] # type: ignore[name-defined]
@property
def UpdateType(self) -> type[Value]:
def UpdateType(self) -> Type[Value]:
"""The type of the update received by the channel."""
return self.typ
def __enter__(self, checkpoint: Optional[str] = None) -> Self:
@contextmanager
def _empty(self, checkpoint: Optional[str] = None) -> Generator[Self, None, None]:
empty = self.__class__(self.typ)
if checkpoint is not None:
empty.queue = tuple(json.loads(checkpoint))
return empty
def __exit__(
self,
__exc_type: type[BaseException] | None,
__exc_value: BaseException | None,
__traceback: TracebackType | None,
) -> None:
try:
del self.queue
except AttributeError:
pass
yield empty
finally:
try:
del empty.queue
except AttributeError:
pass
def _update(self, values: Sequence[Value]) -> None:
self.queue = tuple(values)
@@ -235,31 +215,27 @@ class Set(Generic[Value], Channel[FrozenSet[Value], Value]):
self.typ = typ
@property
def ValueType(self) -> type[FrozenSet[Value]]:
def ValueType(self) -> Type[FrozenSet[Value]]:
"""The type of the value stored in the channel."""
return FrozenSet[self.typ] # type: ignore[name-defined]
@property
def UpdateType(self) -> type[Value]:
def UpdateType(self) -> Type[Value]:
"""The type of the update received by the channel."""
return self.typ
def __enter__(self, checkpoint: Optional[str] = None) -> Self:
@contextmanager
def _empty(self, checkpoint: Optional[str] = None) -> Generator[Self, None, None]:
empty = self.__class__(self.typ)
if checkpoint is not None:
empty.set = set(json.loads(checkpoint))
return empty
def __exit__(
self,
__exc_type: type[BaseException] | None,
__exc_value: BaseException | None,
__traceback: TracebackType | None,
) -> None:
try:
del self.set
except AttributeError:
pass
yield empty
finally:
try:
del empty.set
except AttributeError:
pass
def _update(self, values: Sequence[Value]) -> None:
if not hasattr(self, "set"):
@@ -274,3 +250,76 @@ class Set(Generic[Value], Channel[FrozenSet[Value], Value]):
def _checkpoint(self) -> str:
return json.dumps(list(self.set))
AsyncValue = TypeVar("AsyncValue")
class ContextManager(Generic[Value], Channel[Value, None]):
value: Value
def __init__(
self,
typ: Type[Value],
ctx: Optional[Callable[[], ContextManagerType[Value]]] = None,
actx: Optional[Callable[[], AsyncContextManager[Value]]] = None,
) -> None:
if ctx is None and actx is None:
raise ValueError("Must provide either sync or async context manager.")
self.typ = typ
self.ctx = ctx
self.actx = actx
@property
def ValueType(self) -> Type[Value]:
"""The type of the value stored in the channel."""
return self.typ
@property
def UpdateType(self) -> Type[None]:
"""The type of the update received by the channel."""
raise InvalidUpdateError()
@contextmanager
def _empty(self, checkpoint: Optional[str] = None) -> Generator[Self, None, None]:
if self.ctx is None:
raise ValueError("Cannot enter sync context manager.")
empty = self.__class__(self.typ, ctx=self.ctx, actx=self.actx)
# ContextManager doesn't have a checkpoint
ctx = self.ctx()
empty.value = ctx.__enter__()
try:
yield empty
finally:
ctx.__exit__(None, None, None)
@asynccontextmanager
async def _aempty(
self, checkpoint: Optional[str] = None
) -> AsyncGenerator[Self, None]:
if self.actx is not None:
empty = self.__class__(self.typ, ctx=self.ctx, actx=self.actx)
# ContextManager doesn't have a checkpoint
actx = self.actx()
empty.value = await actx.__aenter__()
try:
yield empty
finally:
await actx.__aexit__(None, None, None)
else:
with self._empty() as empty:
yield empty
def _update(self, values: Sequence[None]) -> None:
raise InvalidUpdateError()
def _get(self) -> Value:
try:
return self.value
except AttributeError:
raise EmptyChannelError()
def _checkpoint(self) -> None:
return None
+4 -4
View File
@@ -245,9 +245,9 @@ class PregelSink(RunnableLambda):
def ChannelsManager(
channels: Mapping[str, Channel]
) -> Generator[Mapping[str, Channel], None, None]:
empty = {k: v.__enter__() for k, v in channels.items()}
empty = {k: v._empty() for k, v in channels.items()}
try:
yield empty
yield {k: v.__enter__() for k, v in empty.items()}
finally:
for v in empty.values():
v.__exit__(None, None, None)
@@ -257,9 +257,9 @@ def ChannelsManager(
async def AsyncChannelsManager(
channels: Mapping[str, Channel]
) -> AsyncGenerator[Mapping[str, Channel], None]:
empty = {k: await v.__aenter__() for k, v in channels.items()}
empty = {k: v._aempty() for k, v in channels.items()}
try:
yield empty
yield {k: await v.__aenter__() for k, v in empty.items()}
finally:
for v in empty.values():
await v.__aexit__(None, None, None)
+130 -5
View File
@@ -1,13 +1,31 @@
import operator
from typing import FrozenSet, Sequence
from contextlib import asynccontextmanager, contextmanager
from typing import AsyncGenerator, FrozenSet, Generator, Sequence
import pytest
from pytest_mock import MockerFixture
import permchain.channels as channels
def test_last_value() -> None:
with channels.LastValue(int) as channel:
with channels.LastValue(int)._empty() as channel:
assert channel.ValueType is int
assert channel.UpdateType is int
with pytest.raises(channels.EmptyChannelError):
channel._get()
with pytest.raises(channels.InvalidUpdateError):
channel._update([5, 6])
channel._update([3])
assert channel._get() == 3
channel._update([4])
assert channel._get() == 4
async def test_last_value_async() -> None:
async with channels.LastValue(int)._aempty() as channel:
assert channel.ValueType is int
assert channel.UpdateType is int
@@ -23,7 +41,21 @@ def test_last_value() -> None:
def test_inbox() -> None:
with channels.Inbox(str) as channel:
with channels.Inbox(str)._empty() as channel:
assert channel.ValueType is Sequence[str]
assert channel.UpdateType is str
with pytest.raises(channels.EmptyChannelError):
channel._get()
channel._update(["a", "b"])
assert channel._get() == ("a", "b")
channel._update(["c"])
assert channel._get() == ("c",)
async def test_inbox_async() -> None:
async with channels.Inbox(str)._aempty() as channel:
assert channel.ValueType is Sequence[str]
assert channel.UpdateType is str
@@ -37,7 +69,21 @@ def test_inbox() -> None:
def test_set() -> None:
with channels.Set(str) as channel:
with channels.Set(str)._empty() as channel:
assert channel.ValueType is FrozenSet[str]
assert channel.UpdateType is str
with pytest.raises(channels.EmptyChannelError):
channel._get()
channel._update(["a", "b"])
assert channel._get() == frozenset(("a", "b"))
channel._update(["b", "c"])
assert channel._get() == frozenset(("a", "b", "c"))
async def test_set_async() -> None:
async with channels.Set(str)._aempty() as channel:
assert channel.ValueType is FrozenSet[str]
assert channel.UpdateType is str
@@ -51,7 +97,7 @@ def test_set() -> None:
def test_binop() -> None:
with channels.BinaryOperatorAggregate(int, operator.add) as channel:
with channels.BinaryOperatorAggregate(int, operator.add)._empty() as channel:
assert channel.ValueType is int
assert channel.UpdateType is int
@@ -62,3 +108,82 @@ def test_binop() -> None:
assert channel._get() == 6
channel._update([4])
assert channel._get() == 10
async def test_binop_async() -> None:
async with channels.BinaryOperatorAggregate(int, operator.add)._aempty() as channel:
assert channel.ValueType is int
assert channel.UpdateType is int
with pytest.raises(channels.EmptyChannelError):
channel._get()
channel._update([1, 2, 3])
assert channel._get() == 6
channel._update([4])
assert channel._get() == 10
def test_ctx_manager(mocker: MockerFixture) -> None:
setup = mocker.Mock()
cleanup = mocker.Mock()
@contextmanager
def an_int() -> Generator[int, None, None]:
setup()
try:
yield 5
finally:
cleanup()
with channels.ContextManager(int, an_int)._empty() as channel:
assert setup.call_count == 1
assert cleanup.call_count == 0
assert channel.ValueType is int
with pytest.raises(channels.InvalidUpdateError):
assert channel.UpdateType is None
assert channel._get() == 5
with pytest.raises(channels.InvalidUpdateError):
channel._update([5])
assert setup.call_count == 1
assert cleanup.call_count == 1
async def test_ctx_manager_async(mocker: MockerFixture) -> None:
setup = mocker.Mock()
cleanup = mocker.Mock()
@contextmanager
def an_int_sync() -> Generator[int, None, None]:
try:
yield 5
finally:
pass
@asynccontextmanager
async def an_int() -> AsyncGenerator[int, None]:
setup()
try:
yield 5
finally:
cleanup()
async with channels.ContextManager(int, an_int_sync, an_int)._aempty() as channel:
assert setup.call_count == 1
assert cleanup.call_count == 0
assert channel.ValueType is int
with pytest.raises(channels.InvalidUpdateError):
assert channel.UpdateType is None
assert channel._get() == 5
with pytest.raises(channels.InvalidUpdateError):
channel._update([5])
assert setup.call_count == 1
assert cleanup.call_count == 1