From 05c89338a8fc7723e1aa230c44e88efdfe122cce Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Thu, 9 Nov 2023 20:43:20 +0000 Subject: [PATCH] Expose imperative channel write api --- permchain/pregel/write.py | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/permchain/pregel/write.py b/permchain/pregel/write.py index 34a3e16bb..1cf9a5ccb 100644 --- a/permchain/pregel/write.py +++ b/permchain/pregel/write.py @@ -44,15 +44,16 @@ class ChannelWrite(RunnablePassthrough): ] def _write(self, input: Any, config: RunnableConfig) -> None: - write: TYPE_SEND = config["configurable"][CONFIG_KEY_SEND] - values = [(chan, r.invoke(input, config)) for chan, r in self.channels] - write([(chan, val) for chan, val in values if val is not None]) + self.do_write(config, **dict(values)) async def _awrite(self, input: Any, config: RunnableConfig) -> None: - write: TYPE_SEND = config["configurable"][CONFIG_KEY_SEND] - values = [(chan, await r.ainvoke(input, config)) for chan, r in self.channels] - write([(chan, val) for chan, val in values if val is not None]) + self.do_write(config, **dict(values)) + + @staticmethod + def do_write(config: RunnableConfig, **values: Any) -> None: + write: TYPE_SEND = config["configurable"][CONFIG_KEY_SEND] + write([(chan, val) for chan, val in values.items() if val is not None])