mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 05:25:05 +02:00
Update context manager interface for channels, add channel that exposes a context manager
This commit is contained in:
+130
-81
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user