diff --git a/permchain/pregel/__init__.py b/permchain/pregel/__init__.py index d7444b5d8..29deb0606 100644 --- a/permchain/pregel/__init__.py +++ b/permchain/pregel/__init__.py @@ -215,9 +215,13 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): input: Iterator[dict[str, Any] | Any], run_manager: CallbackManagerForChainRun, config: RunnableConfig, + *, + output: str | Sequence[str] | None = None, ) -> Iterator[tuple[dict[str, Any] | Any, CheckpointView]]: if config["recursion_limit"] < 1: raise ValueError("recursion_limit must be at least 1") + # assign defaults + output = output if output is not None else self.output # copy nodes to ignore mutations during execution processes = {**self.nodes} # get checkpoint from saver, or create an empty one @@ -256,24 +260,31 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): # collect all writes to channels, without applying them yet pending_writes = deque[tuple[str, Any]]() + # prepare tasks with config + tasks_w_config = [ + ( + proc, + input, + patch_config( + config, + run_name=name, + callbacks=run_manager.get_child(f"graph:step:{step}"), + configurable={ + # deque.extend is thread-safe + CONFIG_KEY_SEND: pending_writes.extend, + CONFIG_KEY_READ: read, + }, + ), + ) + for proc, input, name in next_tasks + ] + # execute tasks, and wait for one to fail or all to finish. # each task is independent from all other concurrent tasks done, inflight = concurrent.futures.wait( [ - executor.submit( - proc.invoke, - input, - patch_config( - config, - callbacks=run_manager.get_child(f"pregel:step:{step}"), - configurable={ - # deque.extend is thread-safe - CONFIG_KEY_SEND: pending_writes.extend, - CONFIG_KEY_READ: read, - }, - ), - ) - for proc, input, _ in next_tasks + executor.submit(proc.invoke, input, config) + for proc, input, config in tasks_w_config ], return_when=concurrent.futures.FIRST_EXCEPTION, timeout=self.step_timeout, @@ -293,7 +304,7 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): values=_updateable_channel_values(channels), step=step + 1, ) - yield map_output(self.output, pending_writes, channels), view + yield map_output(output, pending_writes, channels), view # if view was updated, apply writes to channels _apply_writes_from_view(checkpoint, channels, view) @@ -312,9 +323,13 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): input: AsyncIterator[dict[str, Any] | Any], run_manager: AsyncCallbackManagerForChainRun, config: RunnableConfig, + *, + output: str | Sequence[str] | None = None, ) -> AsyncIterator[tuple[dict[str, Any] | Any, CheckpointView]]: if config["recursion_limit"] < 1: raise ValueError("recursion_limit must be at least 1") + # assign defaults + output = output if output is not None else self.output # copy nodes to ignore mutations during execution processes = {**self.nodes} # get checkpoint from saver, or create an empty one @@ -351,27 +366,31 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): # collect all writes to channels, without applying them yet pending_writes = deque[tuple[str, Any]]() + # prepare tasks with config + tasks_w_config = [ + ( + proc, + input, + patch_config( + config, + run_name=name, + callbacks=run_manager.get_child(f"graph:step:{step}"), + configurable={ + # deque.extend is thread-safe + CONFIG_KEY_SEND: pending_writes.extend, + CONFIG_KEY_READ: read, + }, + ), + ) + for proc, input, name in next_tasks + ] + # execute tasks, and wait for one to fail or all to finish. # each task is independent from all other concurrent tasks done, inflight = await asyncio.wait( [ - asyncio.create_task( - proc.ainvoke( - input, - patch_config( - config, - callbacks=run_manager.get_child( - f"pregel:step:{step}" - ), - configurable={ - # deque.extend is thread-safe - CONFIG_KEY_SEND: pending_writes.extend, - CONFIG_KEY_READ: read, - }, - ), - ) - ) - for proc, input, _ in next_tasks + asyncio.create_task(proc.ainvoke(input, config)) + for proc, input, config in tasks_w_config ], return_when=asyncio.FIRST_EXCEPTION, timeout=self.step_timeout, @@ -391,7 +410,7 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): values=_updateable_channel_values(channels), step=step + 1, ) - yield map_output(self.output, pending_writes, channels), view + yield map_output(output, pending_writes, channels), view # if view was updated, apply writes to channels _apply_writes_from_view(checkpoint, channels, view) @@ -409,10 +428,12 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): self, input: dict[str, Any] | Any, config: RunnableConfig | None = None, + *, + output: str | Sequence[str] | None = None, **kwargs: Any, ) -> dict[str, Any] | Any: latest: dict[str, Any] | Any = None - for chunk in self.stream(input, config, **kwargs): + for chunk in self.stream(input, config, output=output, **kwargs): latest = chunk return latest @@ -420,30 +441,36 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): self, input: dict[str, Any] | Any, config: RunnableConfig | None = None, + *, + output: str | Sequence[str] | None = None, **kwargs: Any, ) -> Iterator[dict[str, Any] | Any]: - return self.transform(iter([input]), config, **kwargs) + return self.transform(iter([input]), config, output=output, **kwargs) def transform( self, input: Iterator[dict[str, Any] | Any], config: RunnableConfig | None = None, + *, + output: str | Sequence[str] | None = None, **kwargs: Any | None, ) -> Iterator[dict[str, Any] | Any]: - for output, _ in self._transform_stream_with_config( - input, self._transform, config, **kwargs + for out, _ in self._transform_stream_with_config( + input, self._transform, config, output=output, **kwargs ): - if output is not None: - yield output + if out is not None: + yield cast(dict[str, Any] | Any, out) def step( self, input: dict[str, Any] | Any, config: RunnableConfig | None = None, + *, + output: str | Sequence[str] | None = None, **kwargs: Any, ) -> Iterator[tuple[dict[str, Any] | Any, CheckpointView]]: for tup in self._transform_stream_with_config( - iter([input]), self._transform, config, **kwargs + iter([input]), self._transform, config, output=output, **kwargs ): yield cast(tuple[dict[str, Any] | Any, CheckpointView], tup) @@ -451,10 +478,12 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): self, input: dict[str, Any] | Any, config: RunnableConfig | None = None, + *, + output: str | Sequence[str] | None = None, **kwargs: Any, ) -> dict[str, Any] | Any: latest: dict[str, Any] | Any = None - async for chunk in self.astream(input, config, **kwargs): + async for chunk in self.astream(input, config, output=output, **kwargs): latest = chunk return latest @@ -462,37 +491,45 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]): self, input: dict[str, Any] | Any, config: RunnableConfig | None = None, + *, + output: str | Sequence[str] | None = None, **kwargs: Any, ) -> AsyncIterator[dict[str, Any] | Any]: async def input_stream() -> AsyncIterator[dict[str, Any] | Any]: yield input - async for chunk in self.atransform(input_stream(), config, **kwargs): + async for chunk in self.atransform( + input_stream(), config, output=output, **kwargs + ): yield chunk async def atransform( self, input: AsyncIterator[dict[str, Any] | Any], config: RunnableConfig | None = None, + *, + output: str | Sequence[str] | None = None, **kwargs: Any | None, ) -> AsyncIterator[dict[str, Any] | Any]: - async for output, _ in self._atransform_stream_with_config( - input, self._atransform, config, **kwargs + async for out, _ in self._atransform_stream_with_config( + input, self._atransform, config, output=output, **kwargs ): - if output is not None: - yield output + if out is not None: + yield out async def astep( self, input: dict[str, Any] | Any, config: RunnableConfig | None = None, + *, + output: str | Sequence[str] | None = None, **kwargs: Any, ) -> AsyncIterator[tuple[dict[str, Any] | Any, CheckpointView]]: async def input_stream() -> AsyncIterator[dict[str, Any] | Any]: yield input async for tup in self._atransform_stream_with_config( - input_stream(), self._atransform, config, **kwargs + input_stream(), self._atransform, config, output=output, **kwargs ): yield cast(tuple[dict[str, Any] | Any, CheckpointView], tup) diff --git a/tests/test_pregel.py b/tests/test_pregel.py index 1a1187367..8e5979746 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -42,6 +42,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 app.invoke(2, output=["output"]) == {"output": 3} assert repr(app), "does not raise recursion error" assert gapp.invoke(2) == 3 @@ -243,6 +244,10 @@ def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None: ) assert [*app.stream({"input": 2, "inbox": 12})] == [13, 4] # [12 + 1, 2 + 1 + 1] + assert [*app.stream({"input": 2, "inbox": 12}, output=["output"])] == [ + {"output": 13}, + {"output": 4}, + ] def test_batch_two_processes_in_out() -> None: @@ -256,6 +261,13 @@ def test_batch_two_processes_in_out() -> None: app = Pregel(nodes={"one": one, "two": two}) assert app.batch([3, 2, 1, 3, 5]) == [5, 4, 3, 5, 7] + assert app.batch([3, 2, 1, 3, 5], output=["output"]) == [ + {"output": 5}, + {"output": 4}, + {"output": 3}, + {"output": 5}, + {"output": 7}, + ] graph = Graph() graph.add_node("add_one", add_one_with_delay) diff --git a/tests/test_pregel_async.py b/tests/test_pregel_async.py index 879df9dfe..1e97c0f13 100644 --- a/tests/test_pregel_async.py +++ b/tests/test_pregel_async.py @@ -36,6 +36,7 @@ async 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 await app.ainvoke(2) == 3 + assert await app.ainvoke(2, output=["output"]) == {"output": 3} async def test_invoke_single_process_in_out_implicit_channels( @@ -195,6 +196,9 @@ async def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None: # [12 + 1, 2 + 1 + 1] assert [c async for c in pubsub.astream({"input": 2, "inbox": 12})] == [13, 4] + assert [ + c async for c in pubsub.astream({"input": 2, "inbox": 12}, output=["output"]) + ] == [{"output": 13}, {"output": 4}] async def test_batch_two_processes_in_out() -> None: @@ -211,6 +215,13 @@ async def test_batch_two_processes_in_out() -> None: ) assert await app.abatch([3, 2, 1, 3, 5]) == [5, 4, 3, 5, 7] + assert await app.abatch([3, 2, 1, 3, 5], output=["output"]) == [ + {"output": 5}, + {"output": 4}, + {"output": 3}, + {"output": 5}, + {"output": 7}, + ] async def test_invoke_many_processes_in_out(mocker: MockerFixture) -> None: