From c3794f1fd3211bbc55bfec225d6ae3aa64f4e0d5 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 20 Aug 2024 15:03:09 -0700 Subject: [PATCH] Use batched async kv inside loop --- libs/langgraph/langgraph/pregel/loop.py | 10 +++++++--- libs/langgraph/tests/test_pregel_async.py | 3 +-- 2 files changed, 8 insertions(+), 5 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index aaa43be67..903a3a644 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -50,6 +50,8 @@ from langgraph.constants import ( Interrupt, ) from langgraph.errors import EmptyInputError, GraphInterrupt +from langgraph.kv.base import BaseKV +from langgraph.kv.batch import AsyncBatchedKV from langgraph.managed.base import ( AsyncManagedValuesManager, ManagedValueMapping, @@ -103,7 +105,7 @@ class PregelLoop: ] ] graph: "Pregel" - + kv: Optional[BaseKV] submit: Submit channels: Mapping[str, BaseChannel] managed: ManagedValueMapping @@ -425,6 +427,7 @@ class SyncPregelLoop(PregelLoop, ContextManager): graph: "Pregel", ) -> None: super().__init__(input, config=config, checkpointer=checkpointer, graph=graph) + self.kv = graph.kv self.stack = ExitStack() if checkpointer: self.checkpointer_get_next_version = checkpointer.get_next_version @@ -476,7 +479,7 @@ class SyncPregelLoop(PregelLoop, ContextManager): self.managed = self.stack.enter_context( ManagedValuesManager( self.graph.managed_values_dict, - patch_config(self.config, configurable={CONFIG_KEY_KV: self.graph.kv}), + patch_config(self.config, configurable={CONFIG_KEY_KV: self.kv}), ) ) self.stack.push(self._suppress_interrupt) @@ -508,6 +511,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): graph: "Pregel", ) -> None: super().__init__(input, config=config, checkpointer=checkpointer, graph=graph) + self.kv = AsyncBatchedKV(graph.kv) if graph.kv else None self.stack = AsyncExitStack() if checkpointer: self.checkpointer_get_next_version = checkpointer.get_next_version @@ -563,7 +567,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): self.managed = await self.stack.enter_async_context( AsyncManagedValuesManager( self.graph.managed_values_dict, - patch_config(self.config, configurable={CONFIG_KEY_KV: self.graph.kv}), + patch_config(self.config, configurable={CONFIG_KEY_KV: self.kv}), ) ) self.stack.push(self._suppress_interrupt) diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 02cd637ac..b383a4d3f 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -52,7 +52,6 @@ from langgraph.errors import InvalidUpdateError, NodeInterrupt from langgraph.graph import END, Graph, StateGraph from langgraph.graph.graph import START from langgraph.graph.message import MessageGraph, add_messages -from langgraph.kv.batch import AsyncBatchedKV from langgraph.kv.memory import MemoryKV from langgraph.managed.shared_value import SharedValue from langgraph.prebuilt.chat_agent_executor import ( @@ -4824,7 +4823,7 @@ async def test_start_branch_then() -> None: async with AsyncSqliteSaver.from_conn_string(":memory:") as saver: tool_two = tool_two_graph.compile( - kv=AsyncBatchedKV(MemoryKV()), + kv=MemoryKV(), checkpointer=saver, interrupt_before=["tool_two_fast", "tool_two_slow"], )