From 1cc02825ea2cc47fc53ccc83761d6f94db2cf329 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 14 Aug 2024 15:53:40 -0700 Subject: [PATCH 01/21] Add ScopedValue - state shared between threads --- libs/langgraph/langgraph/channels/base.py | 5 + libs/langgraph/langgraph/constants.py | 2 + libs/langgraph/langgraph/graph/state.py | 30 ++++- libs/langgraph/langgraph/kv/__init__.py | 0 libs/langgraph/langgraph/kv/base.py | 113 ++++++++++++++++++ libs/langgraph/langgraph/kv/memory.py | 33 +++++ libs/langgraph/langgraph/managed/base.py | 43 ++++++- .../langgraph/managed/scoped_value.py | 99 +++++++++++++++ libs/langgraph/langgraph/pregel/__init__.py | 8 +- libs/langgraph/langgraph/pregel/algo.py | 11 +- libs/langgraph/langgraph/pregel/loop.py | 34 +++++- libs/langgraph/tests/test_kv.py | 32 +++++ libs/langgraph/tests/test_pregel.py | 41 ++++++- 13 files changed, 424 insertions(+), 27 deletions(-) create mode 100644 libs/langgraph/langgraph/kv/__init__.py create mode 100644 libs/langgraph/langgraph/kv/base.py create mode 100644 libs/langgraph/langgraph/kv/memory.py create mode 100644 libs/langgraph/langgraph/managed/scoped_value.py create mode 100644 libs/langgraph/tests/test_kv.py diff --git a/libs/langgraph/langgraph/channels/base.py b/libs/langgraph/langgraph/channels/base.py index fe47f0d8f..9c7794c46 100644 --- a/libs/langgraph/langgraph/channels/base.py +++ b/libs/langgraph/langgraph/channels/base.py @@ -33,6 +33,11 @@ class BaseChannel(Generic[Value, Update, C], ABC): # serialize/deserialize methods + def tap(self) -> Optional[C]: + """Return the current checkpoint of the channel, without consuming it. + By default, it just calls checkpoint().""" + return self.checkpoint() + @abstractmethod def checkpoint(self) -> Optional[C]: """Return a serializable representation of the channel's current state. diff --git a/libs/langgraph/langgraph/constants.py b/libs/langgraph/langgraph/constants.py index f85b33ba3..cf87be337 100644 --- a/libs/langgraph/langgraph/constants.py +++ b/libs/langgraph/langgraph/constants.py @@ -5,6 +5,7 @@ INPUT = "__input__" CONFIG_KEY_SEND = "__pregel_send" CONFIG_KEY_READ = "__pregel_read" CONFIG_KEY_CHECKPOINTER = "__pregel_checkpointer" +CONFIG_KEY_KV = "__pregel_kv" CONFIG_KEY_RESUMING = "__pregel_resuming" CONFIG_KEY_TASK_ID = "__pregel_task_id" INTERRUPT = "__interrupt__" @@ -17,6 +18,7 @@ RESERVED = { CONFIG_KEY_SEND, CONFIG_KEY_READ, CONFIG_KEY_CHECKPOINTER, + CONFIG_KEY_KV, CONFIG_KEY_RESUMING, CONFIG_KEY_TASK_ID, INPUT, diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 5810aaf29..cda357607 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -40,7 +40,14 @@ from langgraph.graph.graph import ( Graph, Send, ) -from langgraph.managed.base import ManagedValue, is_managed_value +from langgraph.kv.base import BaseKV +from langgraph.managed.base import ( + ChannelKeyPlaceholder, + ConfiguredManagedValue, + ManagedValue, + is_managed_value, + is_writable_managed_value, +) from langgraph.pregel.read import ChannelRead, PregelNode from langgraph.pregel.types import All, RetryPolicy from langgraph.pregel.write import SKIP_WRITE, ChannelWrite, ChannelWriteEntry @@ -373,6 +380,8 @@ class StateGraph(Graph): def compile( self, + *, + kv: Optional[BaseKV] = None, checkpointer: Optional[BaseCheckpointSaver] = None, interrupt_before: Optional[Union[All, Sequence[str]]] = None, interrupt_after: Optional[Union[All, Sequence[str]]] = None, @@ -442,6 +451,7 @@ class StateGraph(Graph): interrupt_after_nodes=interrupt_after, auto_validate=False, debug=debug, + kv=kv, ) compiled.attach_node(START, None) @@ -511,7 +521,11 @@ class CompiledStateGraph(CompiledGraph): if not isinstance(v, Context) and not is_managed_value(v) ] else: - output_keys = list(self.builder.channels) + output_keys = list(self.builder.channels) + [ + k + for k, v in self.builder.managed.items() + if is_writable_managed_value(v) + ] def _get_state_key( input: Union[None, dict, Any], config: RunnableConfig, *, key: str @@ -684,7 +698,7 @@ def _get_channels( return {"__root__": _get_channel(schema, allow_managed=False)}, {} all_keys = { - name: _get_channel(typ) + name: _get_channel(name, typ) for name, typ in get_type_hints(schema, include_extras=True).items() if name != "__slots__" } @@ -695,9 +709,9 @@ def _get_channels( def _get_channel( - annotation: Any, *, allow_managed: bool = True + name: str, annotation: Any, *, allow_managed: bool = True ) -> Union[BaseChannel, Type[ManagedValue]]: - if manager := _is_field_managed_value(annotation): + if manager := _is_field_managed_value(name, annotation): if allow_managed: return manager else: @@ -736,12 +750,16 @@ def _is_field_binop(typ: Type[Any]) -> Optional[BinaryOperatorAggregate]: return None -def _is_field_managed_value(typ: Type[Any]) -> Optional[Type[ManagedValue]]: +def _is_field_managed_value(name: str, typ: Type[Any]) -> Optional[Type[ManagedValue]]: if hasattr(typ, "__metadata__"): meta = typ.__metadata__ if len(meta) >= 1: decoration = get_origin(meta[-1]) or meta[-1] if is_managed_value(decoration): + if isinstance(decoration, ConfiguredManagedValue): + for k, v in decoration.kwargs.items(): + if v is ChannelKeyPlaceholder: + decoration.kwargs[k] = name return decoration return None diff --git a/libs/langgraph/langgraph/kv/__init__.py b/libs/langgraph/langgraph/kv/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/libs/langgraph/langgraph/kv/base.py b/libs/langgraph/langgraph/kv/base.py new file mode 100644 index 000000000..570783ec1 --- /dev/null +++ b/libs/langgraph/langgraph/kv/base.py @@ -0,0 +1,113 @@ +import asyncio +from typing import Any, List, NamedTuple, Optional, Union + +V = dict[str, Any] + + +class BaseKV: + def get(self, pairs: List[tuple[str, str]]) -> dict[tuple[str, str], Optional[V]]: + # list[(namespace, key)] -> dict[(namespace, key), value | none] + raise NotImplementedError + + def list(self, prefixes: List[str]) -> dict[str, dict[str, V]]: + # list[namespace] -> dict[namespace, list[value]] + raise NotImplementedError + + def put(self, writes: List[tuple[str, str, Optional[V]]]) -> None: + # list[(namespace, key, value | none)] -> None + raise NotImplementedError + + async def aget( + self, pairs: List[tuple[str, str]] + ) -> dict[tuple[str, str], Optional[V]]: + # list[(namespace, key)] -> dict[(namespace, key), value | none] + raise NotImplementedError + + async def alist(self, prefixes: List[str]) -> dict[str, dict[str, V]]: + # list[namespace] -> dict[namespace, list[value]] + raise NotImplementedError + + async def aput(self, writes: List[tuple[str, str, Optional[V]]]) -> None: + # list[(namespace, key, value | none)] -> None + raise NotImplementedError + + +class GetOp(NamedTuple): + pairs: List[tuple[str, str]] + + +class ListOp(NamedTuple): + prefixes: List[str] + + +class PutOp(NamedTuple): + writes: List[tuple[str, str, Optional[V]]] + + +class KeyValueStore(BaseKV): + def __init__(self, kv: BaseKV) -> None: + self.kv = kv + self.aqueue: dict[asyncio.Future, Union[GetOp, ListOp, PutOp]] = {} + self.task = asyncio.create_task(_run(self.aqueue, self.kv)) + + def __del__(self) -> None: + self.task.cancel() + + async def aget( + self, pairs: List[tuple[str, str]] + ) -> dict[tuple[str, str], Optional[V]]: + fut = asyncio.get_running_loop().create_future() + self.aqueue[fut] = GetOp(pairs) + return await fut + + async def alist(self, prefixes: List[str]) -> dict[str, dict[str, V]]: + fut = asyncio.get_running_loop().create_future() + self.aqueue[fut] = ListOp(prefixes) + return await fut + + async def aput(self, writes: List[tuple[str, str, Optional[V]]]) -> None: + fut = asyncio.get_running_loop().create_future() + self.aqueue[fut] = PutOp(writes) + return await fut + + +async def _run( + aqueue: dict[asyncio.Future, Union[GetOp, ListOp, PutOp]], kv: BaseKV +) -> None: + while True: + await asyncio.sleep(0) + if not aqueue: + continue + # this could use a lock, if we want thread safety + taken = aqueue.copy() + aqueue.clear() + # action each operation + gets = {f: o for f, o in taken.items() if isinstance(o, GetOp)} + if gets: + try: + results = await kv.aget([p for op in gets.values() for p in op.pairs]) + for fut, op in gets.items(): + fut.set_result({k: results.get(k) for k in op.pairs}) + except Exception as e: + for fut in gets: + fut.set_exception(e) + lists = {f: o for f, o in taken.items() if isinstance(o, ListOp)} + if lists: + try: + results = await kv.alist( + [p for op in lists.values() for p in op.prefixes] + ) + for fut, op in lists.items(): + fut.set_result({k: results.get(k) for k in op.prefixes}) + except Exception as e: + for fut in lists: + fut.set_exception(e) + puts = {f: o for f, o in taken.items() if isinstance(o, PutOp)} + if puts: + try: + await kv.aput([w for op in puts.values() for w in op.writes]) + for fut in puts: + fut.set_result(None) + except Exception as e: + for fut in puts: + fut.set_exception(e) diff --git a/libs/langgraph/langgraph/kv/memory.py b/libs/langgraph/langgraph/kv/memory.py new file mode 100644 index 000000000..387ea1e7a --- /dev/null +++ b/libs/langgraph/langgraph/kv/memory.py @@ -0,0 +1,33 @@ +from collections import defaultdict +from typing import List + +from langgraph.kv.base import BaseKV, V + + +class MemoryKV(BaseKV): + def __init__(self) -> None: + self.data: dict[str, dict[str, V]] = defaultdict(dict) + + def get(self, pairs: List[tuple[str, str]]) -> dict[tuple[str, str], V | None]: + return {pair: self.data[pair[0]].get(pair[1]) for pair in pairs} + + async def aget( + self, pairs: List[tuple[str, str]] + ) -> dict[tuple[str, str], V | None]: + return self.get(pairs) + + def list(self, prefixes: List[str]) -> dict[str, dict[str, V]]: + return {prefix: self.data[prefix] for prefix in prefixes} + + async def alist(self, prefixes: List[str]) -> dict[str, dict[str, V]]: + return self.list(prefixes) + + def put(self, writes: List[tuple[str, str, V | None]]) -> None: + for namespace, key, value in writes: + if value is None: + self.data[namespace].pop(key, None) + else: + self.data[namespace][key] = value + + async def aput(self, writes: List[tuple[str, str, V | None]]) -> None: + self.put(writes) diff --git a/libs/langgraph/langgraph/managed/base.py b/libs/langgraph/langgraph/managed/base.py index 0455ed58b..5820383f8 100644 --- a/libs/langgraph/langgraph/managed/base.py +++ b/libs/langgraph/langgraph/managed/base.py @@ -5,9 +5,12 @@ from inspect import isclass from typing import ( Any, AsyncGenerator, + AsyncIterator, Generator, Generic, + Iterator, NamedTuple, + Sequence, Type, TypeVar, Union, @@ -17,6 +20,7 @@ from langchain_core.runnables import RunnableConfig from typing_extensions import Self, TypeGuard V = TypeVar("V") +U = TypeVar("U") class ManagedValue(ABC, Generic[V]): @@ -25,9 +29,7 @@ class ManagedValue(ABC, Generic[V]): @classmethod @contextmanager - def enter( - cls, config: RunnableConfig, **kwargs: Any - ) -> Generator[Self, None, None]: + def enter(cls, config: RunnableConfig, **kwargs: Any) -> Iterator[Self]: try: value = cls(config, **kwargs) yield value @@ -41,9 +43,7 @@ class ManagedValue(ABC, Generic[V]): @classmethod @asynccontextmanager - async def aenter( - cls, config: RunnableConfig, **kwargs: Any - ) -> AsyncGenerator[Self, None]: + async def aenter(cls, config: RunnableConfig, **kwargs: Any) -> AsyncIterator[Self]: try: value = cls(config, **kwargs) yield value @@ -60,6 +60,16 @@ class ManagedValue(ABC, Generic[V]): ... +class WritableManagedValue(Generic[V, U], ManagedValue[V], ABC): + @abstractmethod + def update(self, writes: Sequence[U]) -> None: + ... + + @abstractmethod + async def aupdate(self, writes: Sequence[U]) -> None: + ... + + class ConfiguredManagedValue(NamedTuple): cls: Type[ManagedValue] kwargs: dict[str, Any] @@ -76,6 +86,24 @@ def is_managed_value(value: Any) -> TypeGuard[ManagedValueSpec]: ) +def is_readonly_managed_value(value: Any) -> TypeGuard[Type[ManagedValue]]: + return ( + isclass(value) + and issubclass(value, ManagedValue) + and not issubclass(value, WritableManagedValue) + ) or ( + isinstance(value, ConfiguredManagedValue) + and not issubclass(value.cls, WritableManagedValue) + ) + + +def is_writable_managed_value(value: Any) -> TypeGuard[Type[WritableManagedValue]]: + return (isclass(value) and issubclass(value, WritableManagedValue)) or ( + isinstance(value, ConfiguredManagedValue) + and issubclass(value.cls, WritableManagedValue) + ) + + @contextmanager def ManagedValuesManager( values: dict[str, ManagedValueSpec], @@ -119,3 +147,6 @@ async def AsyncManagedValuesManager( yield {tasks[task]: task.result() for task in done} else: yield {} + + +ChannelKeyPlaceholder = object() diff --git a/libs/langgraph/langgraph/managed/scoped_value.py b/libs/langgraph/langgraph/managed/scoped_value.py new file mode 100644 index 000000000..adb16e1b4 --- /dev/null +++ b/libs/langgraph/langgraph/managed/scoped_value.py @@ -0,0 +1,99 @@ +from contextlib import asynccontextmanager, contextmanager +from typing import ( + Any, + AsyncIterator, + Iterator, + Optional, + Self, + Sequence, +) + +from langchain_core.runnables import RunnableConfig + +from langgraph.constants import CONFIG_KEY_KV +from langgraph.errors import InvalidUpdateError +from langgraph.kv.base import BaseKV +from langgraph.managed.base import ( + ChannelKeyPlaceholder, + ConfiguredManagedValue, + WritableManagedValue, +) +from langgraph.pregel.types import PregelTaskDescription + +V = dict[str, Any] + + +Value = dict[str, V] +Update = dict[str, Optional[V]] + + +class ScopedValue(WritableManagedValue[Value, Update]): + @staticmethod + def configure(scope: str) -> ConfiguredManagedValue: + return ConfiguredManagedValue( + ScopedValue, {"scope": scope, "key": ChannelKeyPlaceholder} + ) + + @classmethod + @contextmanager + def enter(cls, config: RunnableConfig, **kwargs: Any) -> Iterator[Self]: + with super().enter(config, **kwargs) as value: + if value.kv is not None: + saved = value.kv.list([value.ns]) + value.value = saved[value.ns] + yield value + + @classmethod + @asynccontextmanager + async def aenter(cls, config: RunnableConfig, **kwargs: Any) -> AsyncIterator[Self]: + async with super().aenter(config, **kwargs) as value: + if value.kv is not None: + saved = await value.kv.alist([value.ns]) + value.value = saved[value.ns] + yield value + + def __init__(self, config: RunnableConfig, *, scope: str, key: str) -> None: + self.scope = scope + self.config = config + self.value: Value = {} + self.kv: BaseKV = config["configurable"].get(CONFIG_KEY_KV) + if self.kv is None: + self.ns: Optional[str] = None + elif scope_value := config["configurable"].get(self.scope): + self.ns = f"scoped:{scope}:{key}:{scope_value}" + else: + raise ValueError( + f"Scope {scope} for shared state key not in config.configurable" + ) + + def __call__(self, step: int, task: PregelTaskDescription) -> Value: + return self.value.copy() + + def _process_update( + self, values: Sequence[Update] + ) -> list[tuple[str, str, Optional[dict[str, Any]]]]: + writes = [] + for vv in values: + for k, v in vv.items(): + if v is None: + if k in self.value: + self.value[k] = None + writes.append((self.ns, k, None)) + elif not isinstance(v, dict): + raise InvalidUpdateError("Received a non-dict value") + else: + self.value[k] = v + writes.append((self.ns, k, v)) + return writes + + def update(self, values: Sequence[Update]) -> None: + if self.kv is None: + self._process_update(values) + else: + return self.kv.put(self._process_update(values)) + + async def aupdate(self, writes: Sequence[Update]) -> None: + if self.kv is None: + self._process_update(writes) + else: + return await self.kv.aput(self._process_update(writes)) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index d2f0790cc..833453197 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -73,6 +73,7 @@ from langgraph.constants import ( Interrupt, ) from langgraph.errors import GraphInterrupt, GraphRecursionError, InvalidUpdateError +from langgraph.kv.base import BaseKV from langgraph.managed.base import ( AsyncManagedValuesManager, ManagedValuesManager, @@ -223,6 +224,9 @@ class Pregel( checkpointer: Optional[BaseCheckpointSaver] = None """Checkpointer used to save and load graph state. Defaults to None.""" + kv: Optional[BaseKV] = None + """Key-value store to use. Defaults to None.""" + retry_policy: Optional[RetryPolicy] = None """Retry policy to use when running tasks. Set to None to disable.""" @@ -644,7 +648,7 @@ class Pregel( ), ) # apply to checkpoint and save - apply_writes( + assert not apply_writes( checkpoint, channels, [task], self.checkpointer.get_next_version ) checkpoint = create_checkpoint(checkpoint, channels, step + 1) @@ -788,7 +792,7 @@ class Pregel( ), ) # apply to checkpoint and save - apply_writes( + assert not apply_writes( checkpoint, channels, [task], self.checkpointer.get_next_version ) checkpoint = create_checkpoint(checkpoint, channels, step + 1) diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 5922fd14e..994af018e 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -145,7 +145,7 @@ def apply_writes( channels: Mapping[str, BaseChannel], tasks: Sequence[WritesProtocol], get_next_version: Optional[Callable[[int, BaseChannel], int]], -) -> None: +) -> dict[str, list[Any]]: # update seen versions for task in tasks: checkpoint["versions_seen"].setdefault(task.name, {}).update( @@ -161,6 +161,7 @@ def apply_writes( max_version = max(checkpoint["channel_versions"].values()) else: max_version = None + # Consume all channels that were read for chan in { chan for task in tasks for chan in task.triggers if chan not in RESERVED @@ -177,12 +178,15 @@ def apply_writes( # Group writes by channel pending_writes_by_channel: dict[str, list[Any]] = defaultdict(list) + pending_writes_by_managed: dict[str, list[Any]] = defaultdict(list) for task in tasks: for chan, val in task.writes: if chan == TASKS: checkpoint["pending_sends"].append(val) - else: + elif chan in channels: pending_writes_by_channel[chan].append(val) + else: + pending_writes_by_managed[chan].append(val) # Find the highest version of all channels if checkpoint["channel_versions"]: @@ -214,6 +218,9 @@ def apply_writes( max_version, channels[chan] ) + # Return managed values writes to be applied externally + return pending_writes_by_managed + @overload def prepare_next_tasks( diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index f48f86fcd..aaa43be67 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -18,10 +18,11 @@ from typing import ( Type, TypeVar, Union, + cast, ) from langchain_core.callbacks import AsyncParentRunManager, ParentRunManager -from langchain_core.runnables import RunnableConfig +from langchain_core.runnables import RunnableConfig, patch_config from typing_extensions import Self from langgraph.channels.base import BaseChannel @@ -40,6 +41,7 @@ from langgraph.checkpoint.base import ( empty_checkpoint, ) from langgraph.constants import ( + CONFIG_KEY_KV, CONFIG_KEY_READ, CONFIG_KEY_RESUMING, ERROR, @@ -52,6 +54,7 @@ from langgraph.managed.base import ( AsyncManagedValuesManager, ManagedValueMapping, ManagedValuesManager, + WritableManagedValue, ) from langgraph.pregel.algo import ( PregelTaskWrites, @@ -182,12 +185,15 @@ class PregelLoop: elif all(task.writes for task in self.tasks): writes = [w for t in self.tasks for w in t.writes] # all tasks have finished - apply_writes( + mv_writes = apply_writes( self.checkpoint, self.channels, self.tasks, self.checkpointer_get_next_version, ) + # apply writes to managed values + for key, values in mv_writes.items(): + self._update_mv(key, values) # produce values output self.stream.extend( ("values", v) @@ -324,12 +330,13 @@ class PregelLoop: manager=None, ) # apply input writes - apply_writes( + mv_writes = apply_writes( self.checkpoint, self.channels, discard_tasks + [PregelTaskWrites(INPUT, input_writes, [])], self.checkpointer_get_next_version, ) + assert not mv_writes # save input checkpoint self._put_checkpoint({"source": "input", "writes": self.input}) else: @@ -395,6 +402,9 @@ class PregelLoop: # increment step self.step += 1 + def _update_mv(self, key: str, values: Sequence[Any]) -> None: + raise NotImplementedError + def _suppress_interrupt( self, exc_type: Optional[Type[BaseException]], @@ -438,6 +448,9 @@ class SyncPregelLoop(PregelLoop, ContextManager): finally: self.checkpointer.put(config, checkpoint, metadata, new_versions) + def _update_mv(self, key: str, values: Sequence[Any]) -> None: + return self.submit(cast(WritableManagedValue, self.managed[key]).update, values) + # context manager def __enter__(self) -> Self: @@ -461,7 +474,10 @@ class SyncPregelLoop(PregelLoop, ContextManager): ChannelsManager(self.graph.channels, self.checkpoint, self.config) ) self.managed = self.stack.enter_context( - ManagedValuesManager(self.graph.managed_values_dict, self.config) + ManagedValuesManager( + self.graph.managed_values_dict, + patch_config(self.config, configurable={CONFIG_KEY_KV: self.graph.kv}), + ) ) self.stack.push(self._suppress_interrupt) self.status = "pending" @@ -515,6 +531,11 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): finally: await self.checkpointer.aput(config, checkpoint, metadata, new_versions) + def _update_mv(self, key: str, values: Sequence[Any]) -> None: + return self.submit( + cast(WritableManagedValue, self.managed[key]).aupdate, values + ) + # context manager async def __aenter__(self) -> Self: @@ -540,7 +561,10 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): AsyncChannelsManager(self.graph.channels, self.checkpoint, self.config) ) self.managed = await self.stack.enter_async_context( - AsyncManagedValuesManager(self.graph.managed_values_dict, self.config) + AsyncManagedValuesManager( + self.graph.managed_values_dict, + patch_config(self.config, configurable={CONFIG_KEY_KV: self.graph.kv}), + ) ) self.stack.push(self._suppress_interrupt) self.status = "pending" diff --git a/libs/langgraph/tests/test_kv.py b/libs/langgraph/tests/test_kv.py new file mode 100644 index 000000000..2e24ee07d --- /dev/null +++ b/libs/langgraph/tests/test_kv.py @@ -0,0 +1,32 @@ +import asyncio +from typing import Any, List + +from pytest_mock import MockerFixture + +from langgraph.kv.base import BaseKV, KeyValueStore + + +async def test_kv_queue(mocker: MockerFixture) -> None: + aget = mocker.stub() + + class MockKV(BaseKV): + async def aget( + self, pairs: List[tuple[str, str]] + ) -> dict[tuple[str, str], dict[str, Any] | None]: + aget(pairs) + return {pair: {0: pair[0], 1: pair[1]} for pair in pairs} + + store = KeyValueStore(MockKV()) + + # concurrent calls are batched + results = await asyncio.gather( + store.aget([("a", "b")]), + store.aget([("c", "d")]), + ) + assert results == [ + {("a", "b"): {0: "a", 1: "b"}}, + {("c", "d"): {0: "c", 1: "d"}}, + ] + assert [c.args for c in aget.call_args_list] == [ + ([("a", "b"), ("c", "d")],), + ] diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 138d4dba5..3cc1c4d23 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -58,6 +58,8 @@ from langgraph.graph import END, Graph from langgraph.graph.graph import START from langgraph.graph.message import MessageGraph, add_messages from langgraph.graph.state import StateGraph +from langgraph.kv.memory import MemoryKV +from langgraph.managed.scoped_value import ScopedValue from langgraph.prebuilt.chat_agent_executor import ( create_tool_calling_executor, ) @@ -165,6 +167,7 @@ def test_graph_validation() -> None: class State(TypedDict): hello: str + shared_things: Annotated[dict[str, dict[str, Any]], ScopedValue("assistant_id")] def node_a(state: State) -> State: # typo @@ -6202,10 +6205,34 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: class State(TypedDict): my_key: Annotated[str, operator.add] market: str + shared: Annotated[ + dict[str, dict[str, Any]], ScopedValue.configure("assistant_id") + ] + + def assert_shared_value(data: State, config: RunnableConfig) -> State: + assert "shared" in data + if thread_id := config["configurable"].get("thread_id"): + if thread_id == "1": + # this is the first thread, so should not see a value + assert data["shared"] == {} + return {"shared": {"1": {"hello": "world"}}} + elif thread_id == "2": + # this should get value saved by thread 1 + assert data["shared"] == {"1": {"hello": "world"}} + elif thread_id == "3": + # this is a different assistant, so should not see previous value + assert data["shared"] == {} + return {} + + def tool_two_slow(data: State, config: RunnableConfig) -> State: + return {"my_key": " slow", **assert_shared_value(data, config)} + + def tool_two_fast(data: State, config: RunnableConfig) -> State: + return {"my_key": " fast", **assert_shared_value(data, config)} tool_two_graph = StateGraph(State) - tool_two_graph.add_node("tool_two_slow", lambda s: {"my_key": " slow"}) - tool_two_graph.add_node("tool_two_fast", lambda s: {"my_key": " fast"}) + tool_two_graph.add_node("tool_two_slow", tool_two_slow) + tool_two_graph.add_node("tool_two_fast", tool_two_fast) tool_two_graph.set_conditional_entry_point( lambda s: "tool_two_slow" if s["market"] == "DE" else "tool_two_fast", then=END ) @@ -6223,14 +6250,16 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: with SqliteSaver.from_conn_string(":memory:") as saver: tool_two = tool_two_graph.compile( - checkpointer=saver, interrupt_before=["tool_two_fast", "tool_two_slow"] + kv=MemoryKV(), + checkpointer=saver, + interrupt_before=["tool_two_fast", "tool_two_slow"], ) # missing thread_id with pytest.raises(ValueError, match="thread_id"): tool_two.invoke({"my_key": "value", "market": "DE"}) - thread1 = {"configurable": {"thread_id": "1"}} + thread1 = {"configurable": {"thread_id": "1", "assistant_id": "a"}} # stop when about to enter node assert tool_two.invoke({"my_key": "value ⛰️", "market": "DE"}, thread1) == { "my_key": "value ⛰️", @@ -6282,7 +6311,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config, ) - thread2 = {"configurable": {"thread_id": "2"}} + thread2 = {"configurable": {"thread_id": "2", "assistant_id": "a"}} # stop when about to enter node assert tool_two.invoke({"my_key": "value", "market": "US"}, thread2) == { "my_key": "value", @@ -6322,7 +6351,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config, ) - thread3 = {"configurable": {"thread_id": "3"}} + thread3 = {"configurable": {"thread_id": "3", "assistant_id": "b"}} # stop when about to enter node assert tool_two.invoke({"my_key": "value", "market": "US"}, thread3) == { "my_key": "value", From 1dc09dc45fd53ed9c7a795f344f14fd23c9c13d2 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 16 Aug 2024 09:58:49 -0700 Subject: [PATCH 02/21] Remove warning on write to managed channel --- libs/langgraph/langgraph/pregel/algo.py | 15 ++++++++++++--- 1 file changed, 12 insertions(+), 3 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 994af018e..a52d15f4d 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -122,6 +122,7 @@ def local_write( processes: Mapping[str, PregelNode], channels: Mapping[str, BaseChannel], writes: Sequence[tuple[str, Any]], + managed: ManagedValueMapping, ) -> None: for chan, value in writes: if chan == TASKS: @@ -131,7 +132,7 @@ def local_write( ) if value.node not in processes: raise InvalidUpdateError(f"Invalid node name {value.node} in packet") - elif chan not in channels: + elif chan not in channels and chan not in managed: logger.warning(f"Skipping write for channel '{chan}' which has no readers") commit(writes) @@ -321,7 +322,11 @@ def prepare_next_tasks( CONFIG_KEY_TASK_ID: task_id, # deque.extend is thread-safe CONFIG_KEY_SEND: partial( - local_write, writes.extend, processes, channels + local_write, + writes.extend, + processes, + channels, + managed, ), CONFIG_KEY_READ: partial( local_read, @@ -412,7 +417,11 @@ def prepare_next_tasks( CONFIG_KEY_TASK_ID: task_id, # deque.extend is thread-safe CONFIG_KEY_SEND: partial( - local_write, writes.extend, processes, channels + local_write, + writes.extend, + processes, + channels, + managed, ), CONFIG_KEY_READ: partial( local_read, From 2bc0e2df4299259fc7e3187c3c288a0c5145dfc1 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 16 Aug 2024 15:53:05 -0700 Subject: [PATCH 03/21] Fix --- libs/langgraph/langgraph/channels/base.py | 5 ----- libs/langgraph/langgraph/managed/scoped_value.py | 3 +-- 2 files changed, 1 insertion(+), 7 deletions(-) diff --git a/libs/langgraph/langgraph/channels/base.py b/libs/langgraph/langgraph/channels/base.py index 9c7794c46..fe47f0d8f 100644 --- a/libs/langgraph/langgraph/channels/base.py +++ b/libs/langgraph/langgraph/channels/base.py @@ -33,11 +33,6 @@ class BaseChannel(Generic[Value, Update, C], ABC): # serialize/deserialize methods - def tap(self) -> Optional[C]: - """Return the current checkpoint of the channel, without consuming it. - By default, it just calls checkpoint().""" - return self.checkpoint() - @abstractmethod def checkpoint(self) -> Optional[C]: """Return a serializable representation of the channel's current state. diff --git a/libs/langgraph/langgraph/managed/scoped_value.py b/libs/langgraph/langgraph/managed/scoped_value.py index adb16e1b4..997bc9469 100644 --- a/libs/langgraph/langgraph/managed/scoped_value.py +++ b/libs/langgraph/langgraph/managed/scoped_value.py @@ -18,7 +18,6 @@ from langgraph.managed.base import ( ConfiguredManagedValue, WritableManagedValue, ) -from langgraph.pregel.types import PregelTaskDescription V = dict[str, Any] @@ -66,7 +65,7 @@ class ScopedValue(WritableManagedValue[Value, Update]): f"Scope {scope} for shared state key not in config.configurable" ) - def __call__(self, step: int, task: PregelTaskDescription) -> Value: + def __call__(self, step: int) -> Value: return self.value.copy() def _process_update( From 7ca37afc74095c86ad0e86f36ce948d0ed7e33b7 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 16 Aug 2024 16:03:51 -0700 Subject: [PATCH 04/21] Fix --- libs/langgraph/langgraph/graph/state.py | 2 +- libs/langgraph/langgraph/managed/scoped_value.py | 6 +++--- libs/langgraph/langgraph/pregel/algo.py | 2 +- libs/langgraph/tests/test_pregel.py | 10 +++++----- libs/langgraph/tests/test_pregel_async.py | 2 +- 5 files changed, 11 insertions(+), 11 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index cda357607..1c9242ea1 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -695,7 +695,7 @@ def _get_channels( schema: Type[dict], ) -> tuple[dict[str, BaseChannel], dict[str, Type[ManagedValue]]]: if not hasattr(schema, "__annotations__"): - return {"__root__": _get_channel(schema, allow_managed=False)}, {} + return {"__root__": _get_channel("__root__", schema, allow_managed=False)}, {} all_keys = { name: _get_channel(name, typ) diff --git a/libs/langgraph/langgraph/managed/scoped_value.py b/libs/langgraph/langgraph/managed/scoped_value.py index 997bc9469..38a666163 100644 --- a/libs/langgraph/langgraph/managed/scoped_value.py +++ b/libs/langgraph/langgraph/managed/scoped_value.py @@ -26,11 +26,11 @@ Value = dict[str, V] Update = dict[str, Optional[V]] -class ScopedValue(WritableManagedValue[Value, Update]): +class SharedValue(WritableManagedValue[Value, Update]): @staticmethod - def configure(scope: str) -> ConfiguredManagedValue: + def on(scope: str) -> ConfiguredManagedValue: return ConfiguredManagedValue( - ScopedValue, {"scope": scope, "key": ChannelKeyPlaceholder} + SharedValue, {"scope": scope, "key": ChannelKeyPlaceholder} ) @classmethod diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index a52d15f4d..37f168bf7 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -121,8 +121,8 @@ def local_write( commit: Callable[[Sequence[tuple[str, Any]]], None], processes: Mapping[str, PregelNode], channels: Mapping[str, BaseChannel], - writes: Sequence[tuple[str, Any]], managed: ManagedValueMapping, + writes: Sequence[tuple[str, Any]], ) -> None: for chan, value in writes: if chan == TASKS: diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 3cc1c4d23..4bcce7ccd 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -59,7 +59,7 @@ from langgraph.graph.graph import START from langgraph.graph.message import MessageGraph, add_messages from langgraph.graph.state import StateGraph from langgraph.kv.memory import MemoryKV -from langgraph.managed.scoped_value import ScopedValue +from langgraph.managed.scoped_value import SharedValue from langgraph.prebuilt.chat_agent_executor import ( create_tool_calling_executor, ) @@ -167,7 +167,9 @@ def test_graph_validation() -> None: class State(TypedDict): hello: str - shared_things: Annotated[dict[str, dict[str, Any]], ScopedValue("assistant_id")] + shared_things: Annotated[ + dict[str, dict[str, Any]], SharedValue.on("assistant_id") + ] def node_a(state: State) -> State: # typo @@ -6205,9 +6207,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: class State(TypedDict): my_key: Annotated[str, operator.add] market: str - shared: Annotated[ - dict[str, dict[str, Any]], ScopedValue.configure("assistant_id") - ] + shared: Annotated[dict[str, dict[str, Any]], SharedValue.on("assistant_id")] def assert_shared_value(data: State, config: RunnableConfig) -> State: assert "shared" in data diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 7c5742aa0..3048ddbb4 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -563,7 +563,7 @@ async def test_invoke_single_process_in_out(mocker: MockerFixture) -> None: assert app.input_schema.schema() == {"title": "LangGraphInput", "type": "integer"} assert app.output_schema.schema() == {"title": "LangGraphOutput", "type": "integer"} - assert await app.ainvoke(2) == 3 + assert await app.ainvoke(2, debug=True) == 3 assert await app.ainvoke(2, output_keys=["output"]) == {"output": 3} assert await gapp.ainvoke(2) == 3 From 77d7deb033820834dc431b735cd8b7c09b07aa98 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 16 Aug 2024 16:23:57 -0700 Subject: [PATCH 05/21] Rename --- .../langgraph/managed/{scoped_value.py => shared_value.py} | 2 +- libs/langgraph/tests/test_pregel.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) rename libs/langgraph/langgraph/managed/{scoped_value.py => shared_value.py} (98%) diff --git a/libs/langgraph/langgraph/managed/scoped_value.py b/libs/langgraph/langgraph/managed/shared_value.py similarity index 98% rename from libs/langgraph/langgraph/managed/scoped_value.py rename to libs/langgraph/langgraph/managed/shared_value.py index 38a666163..6853306f3 100644 --- a/libs/langgraph/langgraph/managed/scoped_value.py +++ b/libs/langgraph/langgraph/managed/shared_value.py @@ -4,11 +4,11 @@ from typing import ( AsyncIterator, Iterator, Optional, - Self, Sequence, ) from langchain_core.runnables import RunnableConfig +from typing_extensions import Self from langgraph.constants import CONFIG_KEY_KV from langgraph.errors import InvalidUpdateError diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 4bcce7ccd..fef936c80 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -59,7 +59,7 @@ from langgraph.graph.graph import START from langgraph.graph.message import MessageGraph, add_messages from langgraph.graph.state import StateGraph from langgraph.kv.memory import MemoryKV -from langgraph.managed.scoped_value import SharedValue +from langgraph.managed.shared_value import SharedValue from langgraph.prebuilt.chat_agent_executor import ( create_tool_calling_executor, ) From 9b90a24d9460ca70ee15a26541c7af706d82a886 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 16 Aug 2024 16:34:02 -0700 Subject: [PATCH 06/21] Lint --- libs/langgraph/langgraph/managed/base.py | 1 + .../langgraph/managed/shared_value.py | 26 +++++++++++++++++-- 2 files changed, 25 insertions(+), 2 deletions(-) diff --git a/libs/langgraph/langgraph/managed/base.py b/libs/langgraph/langgraph/managed/base.py index 5820383f8..8dbf5c30c 100644 --- a/libs/langgraph/langgraph/managed/base.py +++ b/libs/langgraph/langgraph/managed/base.py @@ -150,3 +150,4 @@ async def AsyncManagedValuesManager( ChannelKeyPlaceholder = object() +ChannelTypePlaceholder = object() diff --git a/libs/langgraph/langgraph/managed/shared_value.py b/libs/langgraph/langgraph/managed/shared_value.py index 6853306f3..f1b8d9516 100644 --- a/libs/langgraph/langgraph/managed/shared_value.py +++ b/libs/langgraph/langgraph/managed/shared_value.py @@ -1,3 +1,4 @@ +import collections.abc from contextlib import asynccontextmanager, contextmanager from typing import ( Any, @@ -5,10 +6,11 @@ from typing import ( Iterator, Optional, Sequence, + Type, ) from langchain_core.runnables import RunnableConfig -from typing_extensions import Self +from typing_extensions import NotRequired, Required, Self from langgraph.constants import CONFIG_KEY_KV from langgraph.errors import InvalidUpdateError @@ -26,6 +28,17 @@ Value = dict[str, V] Update = dict[str, Optional[V]] +# Adapted from typing_extensions +def _strip_extras(t): + """Strips Annotated, Required and NotRequired from a given type.""" + if hasattr(t, "__origin__"): + return _strip_extras(t.__origin__) + if hasattr(t, "__origin__") and t.__origin__ in (Required, NotRequired): + return _strip_extras(t.__args__[0]) + + return t + + class SharedValue(WritableManagedValue[Value, Update]): @staticmethod def on(scope: str) -> ConfiguredManagedValue: @@ -51,7 +64,16 @@ class SharedValue(WritableManagedValue[Value, Update]): value.value = saved[value.ns] yield value - def __init__(self, config: RunnableConfig, *, scope: str, key: str) -> None: + def __init__( + self, config: RunnableConfig, *, typ: Type[Any], scope: str, key: str + ) -> None: + if typ := _strip_extras(typ): + if typ not in ( + dict, + collections.abc.Mapping, + collections.abc.MutableMapping, + ): + raise ValueError("SharedValue must be a dict") self.scope = scope self.config = config self.value: Value = {} From 656f89e16ab7074c972d1920d774bf0852346996 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 16 Aug 2024 16:38:46 -0700 Subject: [PATCH 07/21] Lint --- libs/langgraph/langgraph/graph/state.py | 3 +++ libs/langgraph/langgraph/managed/shared_value.py | 9 ++++++++- 2 files changed, 11 insertions(+), 1 deletion(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 1c9242ea1..261fbf1d5 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -43,6 +43,7 @@ from langgraph.graph.graph import ( from langgraph.kv.base import BaseKV from langgraph.managed.base import ( ChannelKeyPlaceholder, + ChannelTypePlaceholder, ConfiguredManagedValue, ManagedValue, is_managed_value, @@ -760,6 +761,8 @@ def _is_field_managed_value(name: str, typ: Type[Any]) -> Optional[Type[ManagedV for k, v in decoration.kwargs.items(): if v is ChannelKeyPlaceholder: decoration.kwargs[k] = name + if v is ChannelTypePlaceholder: + decoration.kwargs[k] = typ.__origin__ return decoration return None diff --git a/libs/langgraph/langgraph/managed/shared_value.py b/libs/langgraph/langgraph/managed/shared_value.py index f1b8d9516..7c55fc858 100644 --- a/libs/langgraph/langgraph/managed/shared_value.py +++ b/libs/langgraph/langgraph/managed/shared_value.py @@ -17,6 +17,7 @@ from langgraph.errors import InvalidUpdateError from langgraph.kv.base import BaseKV from langgraph.managed.base import ( ChannelKeyPlaceholder, + ChannelTypePlaceholder, ConfiguredManagedValue, WritableManagedValue, ) @@ -43,7 +44,12 @@ class SharedValue(WritableManagedValue[Value, Update]): @staticmethod def on(scope: str) -> ConfiguredManagedValue: return ConfiguredManagedValue( - SharedValue, {"scope": scope, "key": ChannelKeyPlaceholder} + SharedValue, + { + "scope": scope, + "key": ChannelKeyPlaceholder, + "typ": ChannelTypePlaceholder, + }, ) @classmethod @@ -68,6 +74,7 @@ class SharedValue(WritableManagedValue[Value, Update]): self, config: RunnableConfig, *, typ: Type[Any], scope: str, key: str ) -> None: if typ := _strip_extras(typ): + print(typ) if typ not in ( dict, collections.abc.Mapping, From 8f8f3849fcfd55eefcdaf316ac0d0554e244d9bf Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 20 Aug 2024 13:30:37 -0700 Subject: [PATCH 08/21] Lint --- libs/langgraph/langgraph/kv/memory.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/libs/langgraph/langgraph/kv/memory.py b/libs/langgraph/langgraph/kv/memory.py index 387ea1e7a..265dfe77a 100644 --- a/libs/langgraph/langgraph/kv/memory.py +++ b/libs/langgraph/langgraph/kv/memory.py @@ -1,5 +1,5 @@ from collections import defaultdict -from typing import List +from typing import List, Optional from langgraph.kv.base import BaseKV, V @@ -8,12 +8,12 @@ class MemoryKV(BaseKV): def __init__(self) -> None: self.data: dict[str, dict[str, V]] = defaultdict(dict) - def get(self, pairs: List[tuple[str, str]]) -> dict[tuple[str, str], V | None]: + def get(self, pairs: List[tuple[str, str]]) -> dict[tuple[str, str], Optional[V]]: return {pair: self.data[pair[0]].get(pair[1]) for pair in pairs} async def aget( self, pairs: List[tuple[str, str]] - ) -> dict[tuple[str, str], V | None]: + ) -> dict[tuple[str, str], Optional[V]]: return self.get(pairs) def list(self, prefixes: List[str]) -> dict[str, dict[str, V]]: @@ -22,12 +22,12 @@ class MemoryKV(BaseKV): async def alist(self, prefixes: List[str]) -> dict[str, dict[str, V]]: return self.list(prefixes) - def put(self, writes: List[tuple[str, str, V | None]]) -> None: + def put(self, writes: List[tuple[str, str, Optional[V]]]) -> None: for namespace, key, value in writes: if value is None: self.data[namespace].pop(key, None) else: self.data[namespace][key] = value - async def aput(self, writes: List[tuple[str, str, V | None]]) -> None: - self.put(writes) + async def aput(self, writes: List[tuple[str, str, Optional[V]]]) -> None: + return self.put(writes) From 630d9c79edca85ef0417afd78727487c973e25e2 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 20 Aug 2024 13:55:28 -0700 Subject: [PATCH 09/21] Split out async batch to sep file --- libs/langgraph/langgraph/graph/state.py | 2 +- libs/langgraph/langgraph/kv/base.py | 84 +----------------------- libs/langgraph/langgraph/kv/batch.py | 85 +++++++++++++++++++++++++ libs/langgraph/tests/test_kv.py | 31 ++++++--- 4 files changed, 110 insertions(+), 92 deletions(-) create mode 100644 libs/langgraph/langgraph/kv/batch.py diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 261fbf1d5..c116e6d47 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -381,9 +381,9 @@ class StateGraph(Graph): def compile( self, + checkpointer: Optional[BaseCheckpointSaver] = None, *, kv: Optional[BaseKV] = None, - checkpointer: Optional[BaseCheckpointSaver] = None, interrupt_before: Optional[Union[All, Sequence[str]]] = None, interrupt_after: Optional[Union[All, Sequence[str]]] = None, debug: bool = False, diff --git a/libs/langgraph/langgraph/kv/base.py b/libs/langgraph/langgraph/kv/base.py index 570783ec1..37e07ee84 100644 --- a/libs/langgraph/langgraph/kv/base.py +++ b/libs/langgraph/langgraph/kv/base.py @@ -1,5 +1,4 @@ -import asyncio -from typing import Any, List, NamedTuple, Optional, Union +from typing import Any, List, Optional V = dict[str, Any] @@ -30,84 +29,3 @@ class BaseKV: async def aput(self, writes: List[tuple[str, str, Optional[V]]]) -> None: # list[(namespace, key, value | none)] -> None raise NotImplementedError - - -class GetOp(NamedTuple): - pairs: List[tuple[str, str]] - - -class ListOp(NamedTuple): - prefixes: List[str] - - -class PutOp(NamedTuple): - writes: List[tuple[str, str, Optional[V]]] - - -class KeyValueStore(BaseKV): - def __init__(self, kv: BaseKV) -> None: - self.kv = kv - self.aqueue: dict[asyncio.Future, Union[GetOp, ListOp, PutOp]] = {} - self.task = asyncio.create_task(_run(self.aqueue, self.kv)) - - def __del__(self) -> None: - self.task.cancel() - - async def aget( - self, pairs: List[tuple[str, str]] - ) -> dict[tuple[str, str], Optional[V]]: - fut = asyncio.get_running_loop().create_future() - self.aqueue[fut] = GetOp(pairs) - return await fut - - async def alist(self, prefixes: List[str]) -> dict[str, dict[str, V]]: - fut = asyncio.get_running_loop().create_future() - self.aqueue[fut] = ListOp(prefixes) - return await fut - - async def aput(self, writes: List[tuple[str, str, Optional[V]]]) -> None: - fut = asyncio.get_running_loop().create_future() - self.aqueue[fut] = PutOp(writes) - return await fut - - -async def _run( - aqueue: dict[asyncio.Future, Union[GetOp, ListOp, PutOp]], kv: BaseKV -) -> None: - while True: - await asyncio.sleep(0) - if not aqueue: - continue - # this could use a lock, if we want thread safety - taken = aqueue.copy() - aqueue.clear() - # action each operation - gets = {f: o for f, o in taken.items() if isinstance(o, GetOp)} - if gets: - try: - results = await kv.aget([p for op in gets.values() for p in op.pairs]) - for fut, op in gets.items(): - fut.set_result({k: results.get(k) for k in op.pairs}) - except Exception as e: - for fut in gets: - fut.set_exception(e) - lists = {f: o for f, o in taken.items() if isinstance(o, ListOp)} - if lists: - try: - results = await kv.alist( - [p for op in lists.values() for p in op.prefixes] - ) - for fut, op in lists.items(): - fut.set_result({k: results.get(k) for k in op.prefixes}) - except Exception as e: - for fut in lists: - fut.set_exception(e) - puts = {f: o for f, o in taken.items() if isinstance(o, PutOp)} - if puts: - try: - await kv.aput([w for op in puts.values() for w in op.writes]) - for fut in puts: - fut.set_result(None) - except Exception as e: - for fut in puts: - fut.set_exception(e) diff --git a/libs/langgraph/langgraph/kv/batch.py b/libs/langgraph/langgraph/kv/batch.py new file mode 100644 index 000000000..fa6cfc59b --- /dev/null +++ b/libs/langgraph/langgraph/kv/batch.py @@ -0,0 +1,85 @@ +import asyncio +from typing import NamedTuple, Optional, Union + +from langgraph.kv.base import BaseKV, V + + +class GetOp(NamedTuple): + pairs: list[tuple[str, str]] + + +class ListOp(NamedTuple): + prefixes: list[str] + + +class PutOp(NamedTuple): + writes: list[tuple[str, str, Optional[V]]] + + +class AsyncBatchedKV(BaseKV): + def __init__(self, kv: BaseKV) -> None: + self.kv = kv + self.aqueue: dict[asyncio.Future, Union[GetOp, ListOp, PutOp]] = {} + self.task = asyncio.create_task(_run(self.aqueue, self.kv)) + + def __del__(self) -> None: + self.task.cancel() + + async def aget( + self, pairs: list[tuple[str, str]] + ) -> dict[tuple[str, str], Optional[V]]: + fut = asyncio.get_running_loop().create_future() + self.aqueue[fut] = GetOp(pairs) + return await fut + + async def alist(self, prefixes: list[str]) -> dict[str, dict[str, V]]: + fut = asyncio.get_running_loop().create_future() + self.aqueue[fut] = ListOp(prefixes) + return await fut + + async def aput(self, writes: list[tuple[str, str, Optional[V]]]) -> None: + fut = asyncio.get_running_loop().create_future() + self.aqueue[fut] = PutOp(writes) + return await fut + + +async def _run( + aqueue: dict[asyncio.Future, Union[GetOp, ListOp, PutOp]], kv: BaseKV +) -> None: + while True: + await asyncio.sleep(0) + if not aqueue: + continue + # this could use a lock, if we want thread safety + taken = aqueue.copy() + aqueue.clear() + # action each operation + gets = {f: o for f, o in taken.items() if isinstance(o, GetOp)} + if gets: + try: + results = await kv.aget([p for op in gets.values() for p in op.pairs]) + for fut, op in gets.items(): + fut.set_result({k: results.get(k) for k in op.pairs}) + except Exception as e: + for fut in gets: + fut.set_exception(e) + lists = {f: o for f, o in taken.items() if isinstance(o, ListOp)} + if lists: + try: + results = await kv.alist( + [p for op in lists.values() for p in op.prefixes] + ) + for fut, op in lists.items(): + fut.set_result({k: results.get(k) for k in op.prefixes}) + except Exception as e: + for fut in lists: + fut.set_exception(e) + puts = {f: o for f, o in taken.items() if isinstance(o, PutOp)} + if puts: + try: + await kv.aput([w for op in puts.values() for w in op.writes]) + for fut in puts: + fut.set_result(None) + except Exception as e: + for fut in puts: + fut.set_exception(e) diff --git a/libs/langgraph/tests/test_kv.py b/libs/langgraph/tests/test_kv.py index 2e24ee07d..5af46ef9c 100644 --- a/libs/langgraph/tests/test_kv.py +++ b/libs/langgraph/tests/test_kv.py @@ -1,22 +1,28 @@ import asyncio -from typing import Any, List +from typing import Any from pytest_mock import MockerFixture -from langgraph.kv.base import BaseKV, KeyValueStore +from langgraph.kv.base import BaseKV +from langgraph.kv.batch import AsyncBatchedKV -async def test_kv_queue(mocker: MockerFixture) -> None: +async def test_kv_async_batch(mocker: MockerFixture) -> None: aget = mocker.stub() + alist = mocker.stub() class MockKV(BaseKV): async def aget( - self, pairs: List[tuple[str, str]] + self, pairs: list[tuple[str, str]] ) -> dict[tuple[str, str], dict[str, Any] | None]: aget(pairs) - return {pair: {0: pair[0], 1: pair[1]} for pair in pairs} + return {pair: 1 for pair in pairs} - store = KeyValueStore(MockKV()) + async def alist(self, prefixes: list[str]) -> dict[str, dict[str, Any]]: + alist(prefixes) + return {prefix: {prefix: 1} for prefix in prefixes} + + store = AsyncBatchedKV(MockKV()) # concurrent calls are batched results = await asyncio.gather( @@ -24,9 +30,18 @@ async def test_kv_queue(mocker: MockerFixture) -> None: store.aget([("c", "d")]), ) assert results == [ - {("a", "b"): {0: "a", 1: "b"}}, - {("c", "d"): {0: "c", 1: "d"}}, + {("a", "b"): 1}, + {("c", "d"): 1}, ] assert [c.args for c in aget.call_args_list] == [ ([("a", "b"), ("c", "d")],), ] + + results = await asyncio.gather( + store.alist(["a", "b"]), + store.alist(["c", "d"]), + ) + assert results == [{"a": {"a": 1}, "b": {"b": 1}}, {"c": {"c": 1}, "d": {"d": 1}}] + assert [c.args for c in alist.call_args_list] == [ + (["a", "b", "c", "d"],), + ] From 4e1db854f6773758a06b0c288cc79f1cbfbd47c8 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 20 Aug 2024 14:55:23 -0700 Subject: [PATCH 10/21] Add async test --- .../langgraph/managed/shared_value.py | 1 - libs/langgraph/tests/test_pregel_async.py | 40 ++++++++++++++++--- 2 files changed, 34 insertions(+), 7 deletions(-) diff --git a/libs/langgraph/langgraph/managed/shared_value.py b/libs/langgraph/langgraph/managed/shared_value.py index 7c55fc858..c12482c5d 100644 --- a/libs/langgraph/langgraph/managed/shared_value.py +++ b/libs/langgraph/langgraph/managed/shared_value.py @@ -74,7 +74,6 @@ class SharedValue(WritableManagedValue[Value, Update]): self, config: RunnableConfig, *, typ: Type[Any], scope: str, key: str ) -> None: if typ := _strip_extras(typ): - print(typ) if typ not in ( dict, collections.abc.Mapping, diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 3048ddbb4..02cd637ac 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -52,6 +52,9 @@ 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 ( create_tool_calling_executor, ) @@ -4778,10 +4781,33 @@ async def test_start_branch_then() -> None: class State(TypedDict): my_key: Annotated[str, operator.add] market: str + shared: Annotated[dict[str, dict[str, Any]], SharedValue.on("assistant_id")] + other: Annotated[dict[str, dict[str, Any]], SharedValue.on("assistant_id")] + + def assert_shared_value(data: State, config: RunnableConfig) -> State: + assert "shared" in data + if thread_id := config["configurable"].get("thread_id"): + if thread_id == "1": + # this is the first thread, so should not see a value + assert data["shared"] == {} + return {"shared": {"1": {"hello": "world"}}, "other": {"2": {1: 2}}} + elif thread_id == "2": + # this should get value saved by thread 1 + assert data["shared"] == {"1": {"hello": "world"}} + elif thread_id == "3": + # this is a different assistant, so should not see previous value + assert data["shared"] == {} + return {} + + def tool_two_slow(data: State, config: RunnableConfig) -> State: + return {"my_key": " slow", **assert_shared_value(data, config)} + + def tool_two_fast(data: State, config: RunnableConfig) -> State: + return {"my_key": " fast", **assert_shared_value(data, config)} tool_two_graph = StateGraph(State) - tool_two_graph.add_node("tool_two_slow", lambda s, config: {"my_key": " slow"}) - tool_two_graph.add_node("tool_two_fast", lambda s: {"my_key": " fast"}) + tool_two_graph.add_node("tool_two_slow", tool_two_slow) + tool_two_graph.add_node("tool_two_fast", tool_two_fast) tool_two_graph.set_conditional_entry_point( lambda s: "tool_two_slow" if s["market"] == "DE" else "tool_two_fast", then=END ) @@ -4798,14 +4824,16 @@ async def test_start_branch_then() -> None: async with AsyncSqliteSaver.from_conn_string(":memory:") as saver: tool_two = tool_two_graph.compile( - checkpointer=saver, interrupt_before=["tool_two_fast", "tool_two_slow"] + kv=AsyncBatchedKV(MemoryKV()), + checkpointer=saver, + interrupt_before=["tool_two_fast", "tool_two_slow"], ) # missing thread_id with pytest.raises(ValueError, match="thread_id"): await tool_two.ainvoke({"my_key": "value", "market": "DE"}) - thread1 = {"configurable": {"thread_id": "1"}} + thread1 = {"configurable": {"thread_id": "1", "assistant_id": "a"}} # stop when about to enter node assert await tool_two.ainvoke({"my_key": "value", "market": "DE"}, thread1) == { "my_key": "value", @@ -4865,7 +4893,7 @@ async def test_start_branch_then() -> None: ][-1].config, ) - thread2 = {"configurable": {"thread_id": "2"}} + thread2 = {"configurable": {"thread_id": "2", "assistant_id": "a"}} # stop when about to enter node assert await tool_two.ainvoke({"my_key": "value", "market": "US"}, thread2) == { "my_key": "value", @@ -4913,7 +4941,7 @@ async def test_start_branch_then() -> None: ][-1].config, ) - thread3 = {"configurable": {"thread_id": "3"}} + thread3 = {"configurable": {"thread_id": "3", "assistant_id": "b"}} # stop when about to enter node assert await tool_two.ainvoke({"my_key": "value", "market": "US"}, thread3) == { "my_key": "value", From c3794f1fd3211bbc55bfec225d6ae3aa64f4e0d5 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 20 Aug 2024 15:03:09 -0700 Subject: [PATCH 11/21] 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"], ) From b228fc1a9bf75fcdf99de3694d12eb930a23293c Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 20 Aug 2024 16:36:35 -0700 Subject: [PATCH 12/21] Add serde --- libs/langgraph/langgraph/kv/base.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/libs/langgraph/langgraph/kv/base.py b/libs/langgraph/langgraph/kv/base.py index 37e07ee84..bbd2e0899 100644 --- a/libs/langgraph/langgraph/kv/base.py +++ b/libs/langgraph/langgraph/kv/base.py @@ -1,9 +1,15 @@ from typing import Any, List, Optional +from langgraph.checkpoint.serde.base import SerializerProtocol +from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer + V = dict[str, Any] class BaseKV: + def __init__(self, *, serde: SerializerProtocol = JsonPlusSerializer()) -> None: + self.serde = serde + def get(self, pairs: List[tuple[str, str]]) -> dict[tuple[str, str], Optional[V]]: # list[(namespace, key)] -> dict[(namespace, key), value | none] raise NotImplementedError From bd7b9cca21bb74868de6ff9b9407b9793215ba62 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 21 Aug 2024 09:04:43 -0700 Subject: [PATCH 13/21] WIP --- libs/langgraph/langgraph/graph/state.py | 4 +-- libs/langgraph/langgraph/kv/base.py | 18 +---------- libs/langgraph/langgraph/kv/batch.py | 30 ++++--------------- libs/langgraph/langgraph/kv/memory.py | 12 ++------ .../langgraph/managed/shared_value.py | 4 +-- libs/langgraph/langgraph/pregel/__init__.py | 4 +-- libs/langgraph/langgraph/pregel/loop.py | 4 +-- libs/langgraph/tests/test_kv.py | 4 +-- 8 files changed, 18 insertions(+), 62 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index c116e6d47..26770bd20 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -40,7 +40,7 @@ from langgraph.graph.graph import ( Graph, Send, ) -from langgraph.kv.base import BaseKV +from langgraph.kv.base import BaseMemory from langgraph.managed.base import ( ChannelKeyPlaceholder, ChannelTypePlaceholder, @@ -383,7 +383,7 @@ class StateGraph(Graph): self, checkpointer: Optional[BaseCheckpointSaver] = None, *, - kv: Optional[BaseKV] = None, + kv: Optional[BaseMemory] = None, interrupt_before: Optional[Union[All, Sequence[str]]] = None, interrupt_after: Optional[Union[All, Sequence[str]]] = None, debug: bool = False, diff --git a/libs/langgraph/langgraph/kv/base.py b/libs/langgraph/langgraph/kv/base.py index bbd2e0899..0c5791eb6 100644 --- a/libs/langgraph/langgraph/kv/base.py +++ b/libs/langgraph/langgraph/kv/base.py @@ -1,19 +1,9 @@ from typing import Any, List, Optional -from langgraph.checkpoint.serde.base import SerializerProtocol -from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer - V = dict[str, Any] -class BaseKV: - def __init__(self, *, serde: SerializerProtocol = JsonPlusSerializer()) -> None: - self.serde = serde - - def get(self, pairs: List[tuple[str, str]]) -> dict[tuple[str, str], Optional[V]]: - # list[(namespace, key)] -> dict[(namespace, key), value | none] - raise NotImplementedError - +class BaseMemory: def list(self, prefixes: List[str]) -> dict[str, dict[str, V]]: # list[namespace] -> dict[namespace, list[value]] raise NotImplementedError @@ -22,12 +12,6 @@ class BaseKV: # list[(namespace, key, value | none)] -> None raise NotImplementedError - async def aget( - self, pairs: List[tuple[str, str]] - ) -> dict[tuple[str, str], Optional[V]]: - # list[(namespace, key)] -> dict[(namespace, key), value | none] - raise NotImplementedError - async def alist(self, prefixes: List[str]) -> dict[str, dict[str, V]]: # list[namespace] -> dict[namespace, list[value]] raise NotImplementedError diff --git a/libs/langgraph/langgraph/kv/batch.py b/libs/langgraph/langgraph/kv/batch.py index fa6cfc59b..971d50df4 100644 --- a/libs/langgraph/langgraph/kv/batch.py +++ b/libs/langgraph/langgraph/kv/batch.py @@ -1,11 +1,7 @@ import asyncio from typing import NamedTuple, Optional, Union -from langgraph.kv.base import BaseKV, V - - -class GetOp(NamedTuple): - pairs: list[tuple[str, str]] +from langgraph.kv.base import BaseMemory, V class ListOp(NamedTuple): @@ -16,22 +12,15 @@ class PutOp(NamedTuple): writes: list[tuple[str, str, Optional[V]]] -class AsyncBatchedKV(BaseKV): - def __init__(self, kv: BaseKV) -> None: +class AsyncBatchedKV(BaseMemory): + def __init__(self, kv: BaseMemory) -> None: self.kv = kv - self.aqueue: dict[asyncio.Future, Union[GetOp, ListOp, PutOp]] = {} + self.aqueue: dict[asyncio.Future, Union[ListOp, PutOp]] = {} self.task = asyncio.create_task(_run(self.aqueue, self.kv)) def __del__(self) -> None: self.task.cancel() - async def aget( - self, pairs: list[tuple[str, str]] - ) -> dict[tuple[str, str], Optional[V]]: - fut = asyncio.get_running_loop().create_future() - self.aqueue[fut] = GetOp(pairs) - return await fut - async def alist(self, prefixes: list[str]) -> dict[str, dict[str, V]]: fut = asyncio.get_running_loop().create_future() self.aqueue[fut] = ListOp(prefixes) @@ -44,7 +33,7 @@ class AsyncBatchedKV(BaseKV): async def _run( - aqueue: dict[asyncio.Future, Union[GetOp, ListOp, PutOp]], kv: BaseKV + aqueue: dict[asyncio.Future, Union[ListOp, PutOp]], kv: BaseMemory ) -> None: while True: await asyncio.sleep(0) @@ -54,15 +43,6 @@ async def _run( taken = aqueue.copy() aqueue.clear() # action each operation - gets = {f: o for f, o in taken.items() if isinstance(o, GetOp)} - if gets: - try: - results = await kv.aget([p for op in gets.values() for p in op.pairs]) - for fut, op in gets.items(): - fut.set_result({k: results.get(k) for k in op.pairs}) - except Exception as e: - for fut in gets: - fut.set_exception(e) lists = {f: o for f, o in taken.items() if isinstance(o, ListOp)} if lists: try: diff --git a/libs/langgraph/langgraph/kv/memory.py b/libs/langgraph/langgraph/kv/memory.py index 265dfe77a..4b0e92a2d 100644 --- a/libs/langgraph/langgraph/kv/memory.py +++ b/libs/langgraph/langgraph/kv/memory.py @@ -1,21 +1,13 @@ from collections import defaultdict from typing import List, Optional -from langgraph.kv.base import BaseKV, V +from langgraph.kv.base import BaseMemory, V -class MemoryKV(BaseKV): +class MemoryKV(BaseMemory): def __init__(self) -> None: self.data: dict[str, dict[str, V]] = defaultdict(dict) - def get(self, pairs: List[tuple[str, str]]) -> dict[tuple[str, str], Optional[V]]: - return {pair: self.data[pair[0]].get(pair[1]) for pair in pairs} - - async def aget( - self, pairs: List[tuple[str, str]] - ) -> dict[tuple[str, str], Optional[V]]: - return self.get(pairs) - def list(self, prefixes: List[str]) -> dict[str, dict[str, V]]: return {prefix: self.data[prefix] for prefix in prefixes} diff --git a/libs/langgraph/langgraph/managed/shared_value.py b/libs/langgraph/langgraph/managed/shared_value.py index c12482c5d..a1659bb8b 100644 --- a/libs/langgraph/langgraph/managed/shared_value.py +++ b/libs/langgraph/langgraph/managed/shared_value.py @@ -14,7 +14,7 @@ from typing_extensions import NotRequired, Required, Self from langgraph.constants import CONFIG_KEY_KV from langgraph.errors import InvalidUpdateError -from langgraph.kv.base import BaseKV +from langgraph.kv.base import BaseMemory from langgraph.managed.base import ( ChannelKeyPlaceholder, ChannelTypePlaceholder, @@ -83,7 +83,7 @@ class SharedValue(WritableManagedValue[Value, Update]): self.scope = scope self.config = config self.value: Value = {} - self.kv: BaseKV = config["configurable"].get(CONFIG_KEY_KV) + self.kv: BaseMemory = config["configurable"].get(CONFIG_KEY_KV) if self.kv is None: self.ns: Optional[str] = None elif scope_value := config["configurable"].get(self.scope): diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 833453197..a5049d36f 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -73,7 +73,7 @@ from langgraph.constants import ( Interrupt, ) from langgraph.errors import GraphInterrupt, GraphRecursionError, InvalidUpdateError -from langgraph.kv.base import BaseKV +from langgraph.kv.base import BaseMemory from langgraph.managed.base import ( AsyncManagedValuesManager, ManagedValuesManager, @@ -224,7 +224,7 @@ class Pregel( checkpointer: Optional[BaseCheckpointSaver] = None """Checkpointer used to save and load graph state. Defaults to None.""" - kv: Optional[BaseKV] = None + kv: Optional[BaseMemory] = None """Key-value store to use. Defaults to None.""" retry_policy: Optional[RetryPolicy] = None diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 903a3a644..41234e9c7 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -50,7 +50,7 @@ from langgraph.constants import ( Interrupt, ) from langgraph.errors import EmptyInputError, GraphInterrupt -from langgraph.kv.base import BaseKV +from langgraph.kv.base import BaseMemory from langgraph.kv.batch import AsyncBatchedKV from langgraph.managed.base import ( AsyncManagedValuesManager, @@ -105,7 +105,7 @@ class PregelLoop: ] ] graph: "Pregel" - kv: Optional[BaseKV] + kv: Optional[BaseMemory] submit: Submit channels: Mapping[str, BaseChannel] managed: ManagedValueMapping diff --git a/libs/langgraph/tests/test_kv.py b/libs/langgraph/tests/test_kv.py index 5af46ef9c..68294f42c 100644 --- a/libs/langgraph/tests/test_kv.py +++ b/libs/langgraph/tests/test_kv.py @@ -3,7 +3,7 @@ from typing import Any from pytest_mock import MockerFixture -from langgraph.kv.base import BaseKV +from langgraph.kv.base import BaseMemory from langgraph.kv.batch import AsyncBatchedKV @@ -11,7 +11,7 @@ async def test_kv_async_batch(mocker: MockerFixture) -> None: aget = mocker.stub() alist = mocker.stub() - class MockKV(BaseKV): + class MockKV(BaseMemory): async def aget( self, pairs: list[tuple[str, str]] ) -> dict[tuple[str, str], dict[str, Any] | None]: From c9e6ee6da70c3880df5ae54fd8f54e5c546a5666 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 21 Aug 2024 09:34:12 -0700 Subject: [PATCH 14/21] Rename --- libs/langgraph/langgraph/graph/state.py | 6 +++--- libs/langgraph/langgraph/managed/shared_value.py | 4 ++-- libs/langgraph/langgraph/pregel/__init__.py | 4 ++-- libs/langgraph/langgraph/pregel/loop.py | 14 +++++++------- libs/langgraph/langgraph/{kv => store}/__init__.py | 0 libs/langgraph/langgraph/{kv => store}/base.py | 2 +- libs/langgraph/langgraph/{kv => store}/batch.py | 8 ++++---- libs/langgraph/langgraph/{kv => store}/memory.py | 4 ++-- libs/langgraph/tests/test_kv.py | 8 ++++---- libs/langgraph/tests/test_pregel.py | 4 ++-- libs/langgraph/tests/test_pregel_async.py | 4 ++-- 11 files changed, 29 insertions(+), 29 deletions(-) rename libs/langgraph/langgraph/{kv => store}/__init__.py (100%) rename libs/langgraph/langgraph/{kv => store}/base.py (97%) rename libs/langgraph/langgraph/{kv => store}/batch.py (93%) rename libs/langgraph/langgraph/{kv => store}/memory.py (91%) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 26770bd20..8c85b7760 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -40,7 +40,6 @@ from langgraph.graph.graph import ( Graph, Send, ) -from langgraph.kv.base import BaseMemory from langgraph.managed.base import ( ChannelKeyPlaceholder, ChannelTypePlaceholder, @@ -52,6 +51,7 @@ from langgraph.managed.base import ( from langgraph.pregel.read import ChannelRead, PregelNode from langgraph.pregel.types import All, RetryPolicy from langgraph.pregel.write import SKIP_WRITE, ChannelWrite, ChannelWriteEntry +from langgraph.store.base import BaseStore from langgraph.utils import RunnableCallable, coerce_to_runnable logger = logging.getLogger(__name__) @@ -383,7 +383,7 @@ class StateGraph(Graph): self, checkpointer: Optional[BaseCheckpointSaver] = None, *, - kv: Optional[BaseMemory] = None, + store: Optional[BaseStore] = None, interrupt_before: Optional[Union[All, Sequence[str]]] = None, interrupt_after: Optional[Union[All, Sequence[str]]] = None, debug: bool = False, @@ -452,7 +452,7 @@ class StateGraph(Graph): interrupt_after_nodes=interrupt_after, auto_validate=False, debug=debug, - kv=kv, + store=store, ) compiled.attach_node(START, None) diff --git a/libs/langgraph/langgraph/managed/shared_value.py b/libs/langgraph/langgraph/managed/shared_value.py index a1659bb8b..74a7c1ef4 100644 --- a/libs/langgraph/langgraph/managed/shared_value.py +++ b/libs/langgraph/langgraph/managed/shared_value.py @@ -14,13 +14,13 @@ from typing_extensions import NotRequired, Required, Self from langgraph.constants import CONFIG_KEY_KV from langgraph.errors import InvalidUpdateError -from langgraph.kv.base import BaseMemory from langgraph.managed.base import ( ChannelKeyPlaceholder, ChannelTypePlaceholder, ConfiguredManagedValue, WritableManagedValue, ) +from langgraph.store.base import BaseStore V = dict[str, Any] @@ -83,7 +83,7 @@ class SharedValue(WritableManagedValue[Value, Update]): self.scope = scope self.config = config self.value: Value = {} - self.kv: BaseMemory = config["configurable"].get(CONFIG_KEY_KV) + self.kv: BaseStore = config["configurable"].get(CONFIG_KEY_KV) if self.kv is None: self.ns: Optional[str] = None elif scope_value := config["configurable"].get(self.scope): diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index a5049d36f..5e058fd16 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -73,7 +73,6 @@ from langgraph.constants import ( Interrupt, ) from langgraph.errors import GraphInterrupt, GraphRecursionError, InvalidUpdateError -from langgraph.kv.base import BaseMemory from langgraph.managed.base import ( AsyncManagedValuesManager, ManagedValuesManager, @@ -109,6 +108,7 @@ from langgraph.pregel.types import ( from langgraph.pregel.utils import get_new_channel_versions from langgraph.pregel.validate import validate_graph, validate_keys from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry +from langgraph.store.base import BaseStore WriteValue = Union[ Runnable[Input, Output], @@ -224,7 +224,7 @@ class Pregel( checkpointer: Optional[BaseCheckpointSaver] = None """Checkpointer used to save and load graph state. Defaults to None.""" - kv: Optional[BaseMemory] = None + store: Optional[BaseStore] = None """Key-value store to use. Defaults to None.""" retry_policy: Optional[RetryPolicy] = None diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 41234e9c7..8079c0dc2 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -50,8 +50,6 @@ from langgraph.constants import ( Interrupt, ) from langgraph.errors import EmptyInputError, GraphInterrupt -from langgraph.kv.base import BaseMemory -from langgraph.kv.batch import AsyncBatchedKV from langgraph.managed.base import ( AsyncManagedValuesManager, ManagedValueMapping, @@ -74,6 +72,8 @@ from langgraph.pregel.executor import ( from langgraph.pregel.io import map_input, map_output_updates, map_output_values, single from langgraph.pregel.types import PregelExecutableTask from langgraph.pregel.utils import get_new_channel_versions +from langgraph.store.base import BaseStore +from langgraph.store.batch import AsyncBatchedStore if TYPE_CHECKING: from langgraph.pregel import Pregel @@ -105,7 +105,7 @@ class PregelLoop: ] ] graph: "Pregel" - kv: Optional[BaseMemory] + store: Optional[BaseStore] submit: Submit channels: Mapping[str, BaseChannel] managed: ManagedValueMapping @@ -427,7 +427,7 @@ class SyncPregelLoop(PregelLoop, ContextManager): graph: "Pregel", ) -> None: super().__init__(input, config=config, checkpointer=checkpointer, graph=graph) - self.kv = graph.kv + self.store = graph.store self.stack = ExitStack() if checkpointer: self.checkpointer_get_next_version = checkpointer.get_next_version @@ -479,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.kv}), + patch_config(self.config, configurable={CONFIG_KEY_KV: self.store}), ) ) self.stack.push(self._suppress_interrupt) @@ -511,7 +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.store = AsyncBatchedStore(graph.store) if graph.store else None self.stack = AsyncExitStack() if checkpointer: self.checkpointer_get_next_version = checkpointer.get_next_version @@ -567,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.kv}), + patch_config(self.config, configurable={CONFIG_KEY_KV: self.store}), ) ) self.stack.push(self._suppress_interrupt) diff --git a/libs/langgraph/langgraph/kv/__init__.py b/libs/langgraph/langgraph/store/__init__.py similarity index 100% rename from libs/langgraph/langgraph/kv/__init__.py rename to libs/langgraph/langgraph/store/__init__.py diff --git a/libs/langgraph/langgraph/kv/base.py b/libs/langgraph/langgraph/store/base.py similarity index 97% rename from libs/langgraph/langgraph/kv/base.py rename to libs/langgraph/langgraph/store/base.py index 0c5791eb6..7f0030f56 100644 --- a/libs/langgraph/langgraph/kv/base.py +++ b/libs/langgraph/langgraph/store/base.py @@ -3,7 +3,7 @@ from typing import Any, List, Optional V = dict[str, Any] -class BaseMemory: +class BaseStore: def list(self, prefixes: List[str]) -> dict[str, dict[str, V]]: # list[namespace] -> dict[namespace, list[value]] raise NotImplementedError diff --git a/libs/langgraph/langgraph/kv/batch.py b/libs/langgraph/langgraph/store/batch.py similarity index 93% rename from libs/langgraph/langgraph/kv/batch.py rename to libs/langgraph/langgraph/store/batch.py index 971d50df4..cfd05e548 100644 --- a/libs/langgraph/langgraph/kv/batch.py +++ b/libs/langgraph/langgraph/store/batch.py @@ -1,7 +1,7 @@ import asyncio from typing import NamedTuple, Optional, Union -from langgraph.kv.base import BaseMemory, V +from langgraph.store.base import BaseStore, V class ListOp(NamedTuple): @@ -12,8 +12,8 @@ class PutOp(NamedTuple): writes: list[tuple[str, str, Optional[V]]] -class AsyncBatchedKV(BaseMemory): - def __init__(self, kv: BaseMemory) -> None: +class AsyncBatchedStore(BaseStore): + def __init__(self, kv: BaseStore) -> None: self.kv = kv self.aqueue: dict[asyncio.Future, Union[ListOp, PutOp]] = {} self.task = asyncio.create_task(_run(self.aqueue, self.kv)) @@ -33,7 +33,7 @@ class AsyncBatchedKV(BaseMemory): async def _run( - aqueue: dict[asyncio.Future, Union[ListOp, PutOp]], kv: BaseMemory + aqueue: dict[asyncio.Future, Union[ListOp, PutOp]], kv: BaseStore ) -> None: while True: await asyncio.sleep(0) diff --git a/libs/langgraph/langgraph/kv/memory.py b/libs/langgraph/langgraph/store/memory.py similarity index 91% rename from libs/langgraph/langgraph/kv/memory.py rename to libs/langgraph/langgraph/store/memory.py index 4b0e92a2d..48fa2884f 100644 --- a/libs/langgraph/langgraph/kv/memory.py +++ b/libs/langgraph/langgraph/store/memory.py @@ -1,10 +1,10 @@ from collections import defaultdict from typing import List, Optional -from langgraph.kv.base import BaseMemory, V +from langgraph.store.base import BaseStore, V -class MemoryKV(BaseMemory): +class MemoryStore(BaseStore): def __init__(self) -> None: self.data: dict[str, dict[str, V]] = defaultdict(dict) diff --git a/libs/langgraph/tests/test_kv.py b/libs/langgraph/tests/test_kv.py index 68294f42c..0c2862b48 100644 --- a/libs/langgraph/tests/test_kv.py +++ b/libs/langgraph/tests/test_kv.py @@ -3,15 +3,15 @@ from typing import Any from pytest_mock import MockerFixture -from langgraph.kv.base import BaseMemory -from langgraph.kv.batch import AsyncBatchedKV +from langgraph.store.base import BaseStore +from langgraph.store.batch import AsyncBatchedStore async def test_kv_async_batch(mocker: MockerFixture) -> None: aget = mocker.stub() alist = mocker.stub() - class MockKV(BaseMemory): + class MockKV(BaseStore): async def aget( self, pairs: list[tuple[str, str]] ) -> dict[tuple[str, str], dict[str, Any] | None]: @@ -22,7 +22,7 @@ async def test_kv_async_batch(mocker: MockerFixture) -> None: alist(prefixes) return {prefix: {prefix: 1} for prefix in prefixes} - store = AsyncBatchedKV(MockKV()) + store = AsyncBatchedStore(MockKV()) # concurrent calls are batched results = await asyncio.gather( diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index fef936c80..8cf6aafa8 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -58,7 +58,6 @@ from langgraph.graph import END, Graph from langgraph.graph.graph import START from langgraph.graph.message import MessageGraph, add_messages from langgraph.graph.state import StateGraph -from langgraph.kv.memory import MemoryKV from langgraph.managed.shared_value import SharedValue from langgraph.prebuilt.chat_agent_executor import ( create_tool_calling_executor, @@ -67,6 +66,7 @@ from langgraph.prebuilt.tool_node import ToolNode from langgraph.pregel import Channel, GraphRecursionError, Pregel, StateSnapshot from langgraph.pregel.retry import RetryPolicy from langgraph.pregel.types import PregelTask +from langgraph.store.memory import MemoryStore from tests.any_str import AnyStr, ExceptionLike from tests.memory_assert import ( MemorySaverAssertCheckpointMetadata, @@ -6250,7 +6250,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None: with SqliteSaver.from_conn_string(":memory:") as saver: tool_two = tool_two_graph.compile( - kv=MemoryKV(), + store=MemoryStore(), checkpointer=saver, interrupt_before=["tool_two_fast", "tool_two_slow"], ) diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index b383a4d3f..d7a532489 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.memory import MemoryKV from langgraph.managed.shared_value import SharedValue from langgraph.prebuilt.chat_agent_executor import ( create_tool_calling_executor, @@ -62,6 +61,7 @@ from langgraph.prebuilt.tool_node import ToolNode from langgraph.pregel import Channel, GraphRecursionError, Pregel, StateSnapshot from langgraph.pregel.retry import RetryPolicy from langgraph.pregel.types import PregelTask +from langgraph.store.memory import MemoryStore from tests.any_str import AnyStr, ExceptionLike from tests.memory_assert import ( MemorySaverAssertCheckpointMetadata, @@ -4823,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=MemoryKV(), + store=MemoryStore(), checkpointer=saver, interrupt_before=["tool_two_fast", "tool_two_slow"], ) From aa1a6be160c624aa2f3e2eb233cb10852aa05169 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 21 Aug 2024 09:36:21 -0700 Subject: [PATCH 15/21] Rename more --- libs/langgraph/langgraph/constants.py | 4 ++-- .../langgraph/managed/shared_value.py | 22 +++++++++---------- libs/langgraph/langgraph/pregel/loop.py | 6 ++--- libs/langgraph/langgraph/store/batch.py | 12 +++++----- .../tests/{test_kv.py => test_store.py} | 18 +++------------ 5 files changed, 25 insertions(+), 37 deletions(-) rename libs/langgraph/tests/{test_kv.py => test_store.py} (68%) diff --git a/libs/langgraph/langgraph/constants.py b/libs/langgraph/langgraph/constants.py index cf87be337..51b6e1437 100644 --- a/libs/langgraph/langgraph/constants.py +++ b/libs/langgraph/langgraph/constants.py @@ -5,7 +5,7 @@ INPUT = "__input__" CONFIG_KEY_SEND = "__pregel_send" CONFIG_KEY_READ = "__pregel_read" CONFIG_KEY_CHECKPOINTER = "__pregel_checkpointer" -CONFIG_KEY_KV = "__pregel_kv" +CONFIG_KEY_STORE = "__pregel_store" CONFIG_KEY_RESUMING = "__pregel_resuming" CONFIG_KEY_TASK_ID = "__pregel_task_id" INTERRUPT = "__interrupt__" @@ -18,7 +18,7 @@ RESERVED = { CONFIG_KEY_SEND, CONFIG_KEY_READ, CONFIG_KEY_CHECKPOINTER, - CONFIG_KEY_KV, + CONFIG_KEY_STORE, CONFIG_KEY_RESUMING, CONFIG_KEY_TASK_ID, INPUT, diff --git a/libs/langgraph/langgraph/managed/shared_value.py b/libs/langgraph/langgraph/managed/shared_value.py index 74a7c1ef4..7f647d006 100644 --- a/libs/langgraph/langgraph/managed/shared_value.py +++ b/libs/langgraph/langgraph/managed/shared_value.py @@ -12,7 +12,7 @@ from typing import ( from langchain_core.runnables import RunnableConfig from typing_extensions import NotRequired, Required, Self -from langgraph.constants import CONFIG_KEY_KV +from langgraph.constants import CONFIG_KEY_STORE from langgraph.errors import InvalidUpdateError from langgraph.managed.base import ( ChannelKeyPlaceholder, @@ -56,8 +56,8 @@ class SharedValue(WritableManagedValue[Value, Update]): @contextmanager def enter(cls, config: RunnableConfig, **kwargs: Any) -> Iterator[Self]: with super().enter(config, **kwargs) as value: - if value.kv is not None: - saved = value.kv.list([value.ns]) + if value.store is not None: + saved = value.store.list([value.ns]) value.value = saved[value.ns] yield value @@ -65,8 +65,8 @@ class SharedValue(WritableManagedValue[Value, Update]): @asynccontextmanager async def aenter(cls, config: RunnableConfig, **kwargs: Any) -> AsyncIterator[Self]: async with super().aenter(config, **kwargs) as value: - if value.kv is not None: - saved = await value.kv.alist([value.ns]) + if value.store is not None: + saved = await value.store.alist([value.ns]) value.value = saved[value.ns] yield value @@ -83,8 +83,8 @@ class SharedValue(WritableManagedValue[Value, Update]): self.scope = scope self.config = config self.value: Value = {} - self.kv: BaseStore = config["configurable"].get(CONFIG_KEY_KV) - if self.kv is None: + self.store: BaseStore = config["configurable"].get(CONFIG_KEY_STORE) + if self.store is None: self.ns: Optional[str] = None elif scope_value := config["configurable"].get(self.scope): self.ns = f"scoped:{scope}:{key}:{scope_value}" @@ -114,13 +114,13 @@ class SharedValue(WritableManagedValue[Value, Update]): return writes def update(self, values: Sequence[Update]) -> None: - if self.kv is None: + if self.store is None: self._process_update(values) else: - return self.kv.put(self._process_update(values)) + return self.store.put(self._process_update(values)) async def aupdate(self, writes: Sequence[Update]) -> None: - if self.kv is None: + if self.store is None: self._process_update(writes) else: - return await self.kv.aput(self._process_update(writes)) + return await self.store.aput(self._process_update(writes)) diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 8079c0dc2..a1eadb697 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -41,9 +41,9 @@ from langgraph.checkpoint.base import ( empty_checkpoint, ) from langgraph.constants import ( - CONFIG_KEY_KV, CONFIG_KEY_READ, CONFIG_KEY_RESUMING, + CONFIG_KEY_STORE, ERROR, INPUT, INTERRUPT, @@ -479,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.store}), + patch_config(self.config, configurable={CONFIG_KEY_STORE: self.store}), ) ) self.stack.push(self._suppress_interrupt) @@ -567,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.store}), + patch_config(self.config, configurable={CONFIG_KEY_STORE: self.store}), ) ) self.stack.push(self._suppress_interrupt) diff --git a/libs/langgraph/langgraph/store/batch.py b/libs/langgraph/langgraph/store/batch.py index cfd05e548..54eb20d47 100644 --- a/libs/langgraph/langgraph/store/batch.py +++ b/libs/langgraph/langgraph/store/batch.py @@ -13,10 +13,10 @@ class PutOp(NamedTuple): class AsyncBatchedStore(BaseStore): - def __init__(self, kv: BaseStore) -> None: - self.kv = kv + def __init__(self, store: BaseStore) -> None: + self.store = store self.aqueue: dict[asyncio.Future, Union[ListOp, PutOp]] = {} - self.task = asyncio.create_task(_run(self.aqueue, self.kv)) + self.task = asyncio.create_task(_run(self.aqueue, self.store)) def __del__(self) -> None: self.task.cancel() @@ -33,7 +33,7 @@ class AsyncBatchedStore(BaseStore): async def _run( - aqueue: dict[asyncio.Future, Union[ListOp, PutOp]], kv: BaseStore + aqueue: dict[asyncio.Future, Union[ListOp, PutOp]], store: BaseStore ) -> None: while True: await asyncio.sleep(0) @@ -46,7 +46,7 @@ async def _run( lists = {f: o for f, o in taken.items() if isinstance(o, ListOp)} if lists: try: - results = await kv.alist( + results = await store.alist( [p for op in lists.values() for p in op.prefixes] ) for fut, op in lists.items(): @@ -57,7 +57,7 @@ async def _run( puts = {f: o for f, o in taken.items() if isinstance(o, PutOp)} if puts: try: - await kv.aput([w for op in puts.values() for w in op.writes]) + await store.aput([w for op in puts.values() for w in op.writes]) for fut in puts: fut.set_result(None) except Exception as e: diff --git a/libs/langgraph/tests/test_kv.py b/libs/langgraph/tests/test_store.py similarity index 68% rename from libs/langgraph/tests/test_kv.py rename to libs/langgraph/tests/test_store.py index 0c2862b48..a9fd0d1d4 100644 --- a/libs/langgraph/tests/test_kv.py +++ b/libs/langgraph/tests/test_store.py @@ -7,11 +7,11 @@ from langgraph.store.base import BaseStore from langgraph.store.batch import AsyncBatchedStore -async def test_kv_async_batch(mocker: MockerFixture) -> None: +async def test_async_batch_store(mocker: MockerFixture) -> None: aget = mocker.stub() alist = mocker.stub() - class MockKV(BaseStore): + class MockStore(BaseStore): async def aget( self, pairs: list[tuple[str, str]] ) -> dict[tuple[str, str], dict[str, Any] | None]: @@ -22,21 +22,9 @@ async def test_kv_async_batch(mocker: MockerFixture) -> None: alist(prefixes) return {prefix: {prefix: 1} for prefix in prefixes} - store = AsyncBatchedStore(MockKV()) + store = AsyncBatchedStore(MockStore()) # concurrent calls are batched - results = await asyncio.gather( - store.aget([("a", "b")]), - store.aget([("c", "d")]), - ) - assert results == [ - {("a", "b"): 1}, - {("c", "d"): 1}, - ] - assert [c.args for c in aget.call_args_list] == [ - ([("a", "b"), ("c", "d")],), - ] - results = await asyncio.gather( store.alist(["a", "b"]), store.alist(["c", "d"]), From f037a2e9cb6ae5dc75f22c2f0e1138521c8dca84 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 21 Aug 2024 09:39:58 -0700 Subject: [PATCH 16/21] Update docstring --- libs/langgraph/langgraph/pregel/__init__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 5e058fd16..436ed5ed0 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -225,7 +225,7 @@ class Pregel( """Checkpointer used to save and load graph state. Defaults to None.""" store: Optional[BaseStore] = None - """Key-value store to use. Defaults to None.""" + """Memory store to use for SharedValues. Defaults to None.""" retry_policy: Optional[RetryPolicy] = None """Retry policy to use when running tasks. Set to None to disable.""" From 1af0367b3459beb19008dfd0463a81d6f5c37082 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 21 Aug 2024 09:41:53 -0700 Subject: [PATCH 17/21] Add error message --- libs/langgraph/langgraph/pregel/__init__.py | 4 ++-- libs/langgraph/langgraph/pregel/loop.py | 5 ++--- 2 files changed, 4 insertions(+), 5 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 436ed5ed0..429226214 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -650,7 +650,7 @@ class Pregel( # apply to checkpoint and save assert not apply_writes( checkpoint, channels, [task], self.checkpointer.get_next_version - ) + ), "Can't write to SharedValues from update_state" checkpoint = create_checkpoint(checkpoint, channels, step + 1) # check interrupt before if tasks := should_interrupt( @@ -794,7 +794,7 @@ class Pregel( # apply to checkpoint and save assert not apply_writes( checkpoint, channels, [task], self.checkpointer.get_next_version - ) + ), "Can't write to SharedValues from update_state" checkpoint = create_checkpoint(checkpoint, channels, step + 1) # check interrupt before if tasks := should_interrupt( diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index a1eadb697..672d42049 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -332,13 +332,12 @@ class PregelLoop: manager=None, ) # apply input writes - mv_writes = apply_writes( + assert not apply_writes( self.checkpoint, self.channels, discard_tasks + [PregelTaskWrites(INPUT, input_writes, [])], self.checkpointer_get_next_version, - ) - assert not mv_writes + ), "Can't write to SharedValues in graph input" # save input checkpoint self._put_checkpoint({"source": "input", "writes": self.input}) else: From 2106b5e4a6bed75c284146458a07217f5e5bbf06 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 21 Aug 2024 10:48:04 -0700 Subject: [PATCH 18/21] Lint --- libs/langgraph/tests/test_pregel.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 8cf6aafa8..57df91165 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -167,9 +167,6 @@ def test_graph_validation() -> None: class State(TypedDict): hello: str - shared_things: Annotated[ - dict[str, dict[str, Any]], SharedValue.on("assistant_id") - ] def node_a(state: State) -> State: # typo From 7fd4a9ed300fac144e0c07b66e898d02d84bb069 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 21 Aug 2024 10:48:52 -0700 Subject: [PATCH 19/21] Lint --- libs/langgraph/tests/test_pregel_async.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index d7a532489..70ce65daa 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -565,7 +565,7 @@ async def test_invoke_single_process_in_out(mocker: MockerFixture) -> None: assert app.input_schema.schema() == {"title": "LangGraphInput", "type": "integer"} assert app.output_schema.schema() == {"title": "LangGraphOutput", "type": "integer"} - assert await app.ainvoke(2, debug=True) == 3 + assert await app.ainvoke(2) == 3 assert await app.ainvoke(2, output_keys=["output"]) == {"output": 3} assert await gapp.ainvoke(2) == 3 From 4250ff92b8dbbaea3fb8fff31be10a25a9adaa00 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 21 Aug 2024 11:25:37 -0700 Subject: [PATCH 20/21] Fix --- libs/langgraph/langgraph/managed/shared_value.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/libs/langgraph/langgraph/managed/shared_value.py b/libs/langgraph/langgraph/managed/shared_value.py index 7f647d006..7bb6e23b7 100644 --- a/libs/langgraph/langgraph/managed/shared_value.py +++ b/libs/langgraph/langgraph/managed/shared_value.py @@ -58,7 +58,7 @@ class SharedValue(WritableManagedValue[Value, Update]): with super().enter(config, **kwargs) as value: if value.store is not None: saved = value.store.list([value.ns]) - value.value = saved[value.ns] + value.value = saved[value.ns] or {} yield value @classmethod @@ -67,7 +67,7 @@ class SharedValue(WritableManagedValue[Value, Update]): async with super().aenter(config, **kwargs) as value: if value.store is not None: saved = await value.store.alist([value.ns]) - value.value = saved[value.ns] + value.value = saved[value.ns] or {} yield value def __init__( From 5fb2c2c6f804fc55961610483ed8ed95338d767d Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 21 Aug 2024 11:27:04 -0700 Subject: [PATCH 21/21] Lint --- libs/langgraph/tests/test_store.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/libs/langgraph/tests/test_store.py b/libs/langgraph/tests/test_store.py index a9fd0d1d4..cd53ee407 100644 --- a/libs/langgraph/tests/test_store.py +++ b/libs/langgraph/tests/test_store.py @@ -1,5 +1,5 @@ import asyncio -from typing import Any +from typing import Any, Optional from pytest_mock import MockerFixture @@ -14,7 +14,7 @@ async def test_async_batch_store(mocker: MockerFixture) -> None: class MockStore(BaseStore): async def aget( self, pairs: list[tuple[str, str]] - ) -> dict[tuple[str, str], dict[str, Any] | None]: + ) -> dict[tuple[str, str], Optional[dict[str, Any]]]: aget(pairs) return {pair: 1 for pair in pairs}