From 9a2ddb30d7c320d465976cd16619eed61aafa2a5 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 28 Nov 2023 14:06:51 +0000 Subject: [PATCH] Add test --- permchain/__init__.py | 2 -- permchain/pregel/__init__.py | 2 +- permchain/pregel/write.py | 13 +++++++++---- tests/test_pregel.py | 1 + 4 files changed, 11 insertions(+), 7 deletions(-) diff --git a/permchain/__init__.py b/permchain/__init__.py index 3e82fd557..f72ff7219 100644 --- a/permchain/__init__.py +++ b/permchain/__init__.py @@ -1,12 +1,10 @@ from permchain.checkpoint.base import BaseCheckpointAdapter, CheckpointAt from permchain.pregel import Channel, Pregel, ReservedChannels -from permchain.pregel.read import ChannelRead __all__ = [ "Channel", "Pregel", "ReservedChannels", - "ChannelRead", "BaseCheckpointAdapter", "CheckpointAt", ] diff --git a/permchain/pregel/__init__.py b/permchain/pregel/__init__.py index 8545adfec..5145f842e 100644 --- a/permchain/pregel/__init__.py +++ b/permchain/pregel/__init__.py @@ -117,7 +117,7 @@ class Channel: """Writes to channels the result of the lambda, or None to skip writing.""" return ChannelWrite( channels=( - [(c, RunnablePassthrough()) for c in channels] + [(c, None) for c in channels] + [(k, _coerce_write_value(v)) for k, v in kwargs.items()] ) ) diff --git a/permchain/pregel/write.py b/permchain/pregel/write.py index 27480aa33..8a0c1b253 100644 --- a/permchain/pregel/write.py +++ b/permchain/pregel/write.py @@ -15,7 +15,7 @@ TYPE_SEND = Callable[[Sequence[tuple[str, Any]]], None] class ChannelWrite(RunnablePassthrough): - channels: Sequence[tuple[str, Runnable]] + channels: Sequence[tuple[str, Runnable | None]] """ Mapping of write channels to Runnables that return the value to be written, or None to skip writing. @@ -27,7 +27,7 @@ class ChannelWrite(RunnablePassthrough): def __init__( self, *, - channels: Sequence[tuple[str, Runnable]], + channels: Sequence[tuple[str, Runnable | None]], ): super().__init__(func=self._write, afunc=self._awrite, channels=channels) @@ -44,12 +44,17 @@ class ChannelWrite(RunnablePassthrough): ] def _write(self, input: Any, config: RunnableConfig) -> None: - values = [(chan, r.invoke(input, config)) for chan, r in self.channels] + values = [ + (chan, r.invoke(input, config) if r else input) for chan, r in self.channels + ] self.do_write(config, **dict(values)) async def _awrite(self, input: Any, config: RunnableConfig) -> None: - values = [(chan, await r.ainvoke(input, config)) for chan, r in self.channels] + values = [ + (chan, await r.ainvoke(input, config) if r else input) + for chan, r in self.channels + ] self.do_write(config, **dict(values)) diff --git a/tests/test_pregel.py b/tests/test_pregel.py index a2777eec0..1a2f127db 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -37,6 +37,7 @@ def test_invoke_single_process_in_out(mocker: MockerFixture) -> None: assert app.input_schema.schema() == {"title": "PregelInput", "type": "integer"} assert app.output_schema.schema() == {"title": "PregelOutput", "type": "integer"} assert app.invoke(2) == 3 + assert repr(app), "does not raise recursion error" def test_invoke_single_process_in_out_implicit_channels(mocker: MockerFixture) -> None: