Simplify api of context mngr channel

This commit is contained in:
Nuno Campos
2023-10-19 12:13:18 +01:00
parent 73a5ff0044
commit b07f8d87e5
4 changed files with 50 additions and 26 deletions
+1 -4
View File
@@ -6,7 +6,6 @@ from langchain.schema.runnable import RunnableLambda, RunnablePassthrough
from langchain.utils.html import extract_sub_links
from permchain import Pregel, channels
from permchain.pregel import PregelRead
# Load url with sync httpx client
@@ -107,9 +106,7 @@ def recursive_web_loader(
"next_urls": channels.UniqueInbox(str),
"documents": channels.Stream(Document),
"visited": channels.Set(str),
"client": channels.ContextManager(
httpx.Client | httpx.AsyncClient, httpx.Client, httpx.AsyncClient
),
"client": channels.ContextManager(httpx.Client, httpx.AsyncClient),
},
# this will accept a string as input
input="base_url",
+24 -15
View File
@@ -11,7 +11,6 @@ from typing import (
Generic,
Optional,
Sequence,
Tuple,
Type,
TypeVar,
Union,
@@ -203,7 +202,11 @@ class Inbox(Generic[Value], Channel[Sequence[Value], Value | Sequence[Value]]):
self.queue = tuple(
cast(Value, v)
for value in values
for v in ((value,) if isinstance(value, self.typ) else value)
for v in (
(value,)
if isinstance(value, self.typ)
else cast(Sequence[Value], value)
)
)
def get(self) -> Sequence[Value]:
@@ -250,7 +253,11 @@ class UniqueInbox(Generic[Value], Channel[Sequence[Value], Value | Sequence[Valu
set(
cast(Value, v)
for value in values
for v in ((value,) if isinstance(value, self.typ) else value)
for v in (
(value,)
if isinstance(value, self.typ)
else cast(Sequence[Value], value)
)
)
)
@@ -304,7 +311,7 @@ class Set(Generic[Value], Channel[FrozenSet[Value], Value]):
return json.dumps(list(self.set))
class Stream(Generic[Value], Channel[Tuple[Value], Value]):
class Stream(Generic[Value], Channel[Sequence[Value], Value]):
"""Stores all unique values received."""
def __init__(self, typ: Type[Value]) -> None:
@@ -312,9 +319,9 @@ class Stream(Generic[Value], Channel[Tuple[Value], Value]):
self.set = list[Value]()
@property
def ValueType(self) -> Type[Tuple[Value]]:
def ValueType(self) -> Any:
"""The type of the value stored in the channel."""
return Tuple[self.typ] # type: ignore[name-defined]
return Sequence[self.typ] # type: ignore[name-defined]
@property
def UpdateType(self) -> Type[Value]:
@@ -334,7 +341,7 @@ class Stream(Generic[Value], Channel[Tuple[Value], Value]):
def update(self, values: Sequence[Value]) -> None:
self.set.extend(values)
def get(self) -> Tuple[Value]:
def get(self) -> Sequence[Value]:
try:
return tuple(self.set)
except AttributeError:
@@ -344,17 +351,14 @@ class Stream(Generic[Value], Channel[Tuple[Value], Value]):
return json.dumps(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,
typ: Optional[Type[Value]] = None,
) -> None:
if ctx is None and actx is None:
raise ValueError("Must provide either sync or async context manager.")
@@ -364,9 +368,14 @@ class ContextManager(Generic[Value], Channel[Value, None]):
self.actx = actx
@property
def ValueType(self) -> Type[Value]:
def ValueType(self) -> Any:
"""The type of the value stored in the channel."""
return self.typ
return (
self.typ
or (self.ctx if hasattr(self.ctx, "__enter__") else None)
or (self.actx if hasattr(self.actx, "__aenter__") else None)
or None
)
@property
def UpdateType(self) -> Type[None]:
@@ -378,7 +387,7 @@ class ContextManager(Generic[Value], Channel[Value, None]):
if self.ctx is None:
raise ValueError("Cannot enter sync context manager.")
empty = self.__class__(self.typ, ctx=self.ctx, actx=self.actx)
empty = self.__class__(ctx=self.ctx, actx=self.actx, typ=self.typ)
# ContextManager doesn't have a checkpoint
ctx = self.ctx()
empty.value = ctx.__enter__()
@@ -392,7 +401,7 @@ class ContextManager(Generic[Value], Channel[Value, None]):
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)
empty = self.__class__(ctx=self.ctx, actx=self.actx, typ=self.typ)
# ContextManager doesn't have a checkpoint
actx = self.actx()
empty.value = await actx.__aenter__()
+10 -5
View File
@@ -129,6 +129,7 @@ class PregelInvoke(RunnableBinding):
if isinstance(self.bound, RunnablePassthrough):
return PregelInvoke(channels=self.channels, bound=coerce_to_runnable(other))
else:
# delegate to __or__ in self.bound
return PregelInvoke(channels=self.channels, bound=self.bound | other)
def __ror__(
@@ -164,7 +165,7 @@ class PregelBatch(RunnableEach):
channel=self.channel, key=self.key, bound=self.bound | joiner
)
def __or__(
def __or__( # type: ignore[override]
self,
other: Runnable[Any, Other]
| Callable[[Any], Other]
@@ -175,6 +176,7 @@ class PregelBatch(RunnableEach):
channel=self.channel, key=self.key, bound=coerce_to_runnable(other)
)
else:
# delegate to __or__ in self.bound
return PregelBatch(
channel=self.channel, key=self.key, bound=self.bound | other
)
@@ -203,7 +205,7 @@ class PregelSink(RunnableLambda):
channels: Sequence[tuple[str, Runnable]],
max_steps: Optional[int] = None,
):
super().__init__(func=self._write, afunc=self._awrite)
super().__init__(func=self._write, afunc=self._awrite) # type: ignore[arg-type]
self.channels = channels
self.max_steps = max_steps
@@ -356,9 +358,12 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]):
"""Runs process.invoke() each time channels are updated,
with a dict of the channel values as input."""
return PregelInvoke(
channels={None: channels}
if isinstance(channels, str)
else {chan: chan for chan in channels}
channels=cast(
Mapping[None, str] | Mapping[str, str],
{None: channels}
if isinstance(channels, str)
else {chan: chan for chan in channels},
)
)
@classmethod
+15 -2
View File
@@ -2,6 +2,7 @@ import operator
from contextlib import asynccontextmanager, contextmanager
from typing import AsyncGenerator, FrozenSet, Generator, Sequence, Union
import httpx
import pytest
from pytest_mock import MockerFixture
@@ -133,7 +134,7 @@ def test_ctx_manager(mocker: MockerFixture) -> None:
finally:
cleanup()
with channels.ContextManager(int, an_int).empty() as channel:
with channels.ContextManager(an_int, None, int).empty() as channel:
assert setup.call_count == 1
assert cleanup.call_count == 0
@@ -150,6 +151,18 @@ def test_ctx_manager(mocker: MockerFixture) -> None:
assert cleanup.call_count == 1
def test_ctx_manager_ctx(mocker: MockerFixture) -> None:
with channels.ContextManager(httpx.Client).empty() as channel:
assert channel.ValueType is httpx.Client
with pytest.raises(channels.InvalidUpdateError):
assert channel.UpdateType is None
assert isinstance(channel.get(), httpx.Client)
with pytest.raises(channels.InvalidUpdateError):
channel.update([5]) # type: ignore
async def test_ctx_manager_async(mocker: MockerFixture) -> None:
setup = mocker.Mock()
cleanup = mocker.Mock()
@@ -169,7 +182,7 @@ async def test_ctx_manager_async(mocker: MockerFixture) -> None:
finally:
cleanup()
async with channels.ContextManager(int, an_int_sync, an_int).aempty() as channel:
async with channels.ContextManager(an_int_sync, an_int, int).aempty() as channel:
assert setup.call_count == 1
assert cleanup.call_count == 0