Use batched async kv inside loop

This commit is contained in:
Nuno Campos
2024-08-21 09:30:22 -07:00
parent 4e1db854f6
commit c3794f1fd3
2 changed files with 8 additions and 5 deletions
+7 -3
View File
@@ -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)
+1 -2
View File
@@ -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"],
)