From b1a42430555c4ff31f3430ede7978f297f328cd4 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Sun, 29 Sep 2024 17:06:42 -0700 Subject: [PATCH] Implement new Store interface (#1834) --- .../langgraph/checkpoint/serde/jsonplus.py | 13 + libs/checkpoint/langgraph/store/base.py | 410 +++++++++++++++++ libs/checkpoint/langgraph/store/batch.py | 88 ++++ libs/checkpoint/langgraph/store/memory.py | 119 +++++ .../langgraph/store/py.typed} | 0 libs/checkpoint/tests/test_jsonplus.py | 8 + libs/checkpoint/tests/test_store.py | 261 +++++++++++ .../langgraph/managed/shared_value.py | 26 +- libs/langgraph/langgraph/pregel/__init__.py | 13 +- libs/langgraph/langgraph/pregel/algo.py | 14 + libs/langgraph/langgraph/pregel/loop.py | 4 +- libs/langgraph/langgraph/store/base.py | 21 - libs/langgraph/langgraph/store/batch.py | 65 --- libs/langgraph/langgraph/store/memory.py | 25 -- libs/langgraph/langgraph/utils/runnable.py | 28 +- libs/langgraph/tests/test_algo.py | 11 +- libs/langgraph/tests/test_pregel.py | 63 ++- libs/langgraph/tests/test_pregel_async.py | 85 +++- libs/langgraph/tests/test_runnable.py | 43 ++ libs/langgraph/tests/test_store.py | 38 -- .../langgraph/scheduler/kafka/executor.py | 1 + libs/scheduler-kafka/tests/test_subgraph.py | 6 + .../tests/test_subgraph_sync.py | 6 + libs/sdk-py/langgraph_sdk/client.py | 414 +++++++++++++++++- libs/sdk-py/langgraph_sdk/schema.py | 24 + 25 files changed, 1594 insertions(+), 192 deletions(-) create mode 100644 libs/checkpoint/langgraph/store/base.py create mode 100644 libs/checkpoint/langgraph/store/batch.py create mode 100644 libs/checkpoint/langgraph/store/memory.py rename libs/{langgraph/langgraph/store/__init__.py => checkpoint/langgraph/store/py.typed} (100%) create mode 100644 libs/checkpoint/tests/test_store.py delete mode 100644 libs/langgraph/langgraph/store/base.py delete mode 100644 libs/langgraph/langgraph/store/batch.py delete mode 100644 libs/langgraph/langgraph/store/memory.py create mode 100644 libs/langgraph/tests/test_runnable.py delete mode 100644 libs/langgraph/tests/test_store.py diff --git a/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py b/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py index d9796b0c3..71b411f3d 100644 --- a/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py +++ b/libs/checkpoint/langgraph/checkpoint/serde/jsonplus.py @@ -26,6 +26,7 @@ from zoneinfo import ZoneInfo from langgraph.checkpoint.serde.base import SerializerProtocol from langgraph.checkpoint.serde.types import SendProtocol +from langgraph.store.base import Item LC_REVIVER = Reviver() @@ -414,6 +415,18 @@ def _msgpack_default(obj: Any) -> Union[str, msgpack.ExtType]: ), ), ) + elif isinstance(obj, Item): + return msgpack.ExtType( + EXT_CONSTRUCTOR_KW_ARGS, + _msgpack_enc( + ( + obj.__class__.__module__, + obj.__class__.__name__, + {k: getattr(obj, k) for k in obj.__slots__}, + ), + ), + ) + elif isinstance(obj, BaseException): return repr(obj) else: diff --git a/libs/checkpoint/langgraph/store/base.py b/libs/checkpoint/langgraph/store/base.py new file mode 100644 index 000000000..d1ad914a6 --- /dev/null +++ b/libs/checkpoint/langgraph/store/base.py @@ -0,0 +1,410 @@ +"""Base classes and types for persistent key-value stores. + +Stores enable persistence and memory that can be shared across threads, +scoped to user IDs, assistant IDs, or other arbitrary namespaces. +""" + +from abc import ABC, abstractmethod +from datetime import datetime +from typing import Any, Iterable, Literal, NamedTuple, Optional, Union, cast + + +class Item: + """Represents a stored item with metadata. + + Args: + value (dict[str, Any]): The stored data as a dictionary. Keys are filterable. + (str): Unique identifier within the namespace. + namespace (tuple[str, ...]): Hierarchical path defining the collection in which this document resides. + Represented as a tuple of strings, allowing for nested categorization. + For example: ("documents", 'user123') + created_at (datetime): Timestamp of item creation. + updated_at (datetime): Timestamp of last update. + """ + + __slots__ = ("value", "key", "namespace", "created_at", "updated_at") + + def __init__( + self, + *, + value: dict[str, Any], + key: str, + namespace: tuple[str, ...], + created_at: datetime, + updated_at: datetime, + ): + self.value = value + self.key = key + # The casting from json-like types is for if this object is + # deserialized. + self.namespace = tuple(namespace) + self.created_at = ( + datetime.fromisoformat(cast(str, created_at)) + if isinstance(created_at, str) + else created_at + ) + self.updated_at = ( + datetime.fromisoformat(cast(str, created_at)) + if isinstance(updated_at, str) + else updated_at + ) + + def __eq__(self, other: object) -> bool: + if not isinstance(other, Item): + return False + return ( + self.value == other.value + and self.key == other.key + and self.namespace == other.namespace + and self.created_at == other.created_at + and self.updated_at == other.updated_at + ) + + def __hash__(self) -> int: + return hash((self.namespace, self.key)) + + def dict(self) -> dict: + return { + "value": self.value, + "key": self.key, + "namespace": list(self.namespace), + "created_at": self.created_at.isoformat(), + "updated_at": self.updated_at.isoformat(), + } + + +class GetOp(NamedTuple): + """Operation to retrieve an item by namespace and key.""" + + namespace: tuple[str, ...] + """Hierarchical path for the item.""" + key: str + """Unique identifier within the namespace.""" + + +class SearchOp(NamedTuple): + """Operation to search for items within a namespace prefix.""" + + namespace_prefix: tuple[str, ...] + """Hierarchical path prefix to search within.""" + filter: Optional[dict[str, Any]] = None + """Key-value pairs to filter results.""" + limit: int = 10 + """Maximum number of items to return.""" + offset: int = 0 + """Number of items to skip before returning results.""" + + +class PutOp(NamedTuple): + """Operation to store, update, or delete an item.""" + + namespace: tuple[str, ...] + """Hierarchical path for the item. + + Represented as a tuple of strings, allowing for nested categorization. + For example: ("documents", "user123") + """ + + key: str + """Unique identifier for the document. + + Should be distinct within its namespace. + """ + + value: Optional[dict[str, Any]] + """Data to be stored, or None to delete the item. + + Schema: + - Should be a dictionary where: + - Keys are strings representing field names + - Values can be of any serializable type + - If None, it indicates that the item should be deleted + """ + + +NameSpacePath = tuple[Union[str, Literal["*"]], ...] + +NamespaceMatchType = Literal["prefix", "suffix"] + + +class MatchCondition(NamedTuple): + """Represents a single match condition.""" + + match_type: NamespaceMatchType + path: NameSpacePath + + +class ListNamespacesOp(NamedTuple): + """Operation to list namespaces with optional match conditions.""" + + match_conditions: Optional[tuple[MatchCondition, ...]] = None + """A tuple of match conditions to apply to namespaces.""" + + max_depth: Optional[int] = None + """Return namespaces up to this depth in the hierarchy.""" + + limit: int = 100 + """Maximum number of namespaces to return.""" + + offset: int = 0 + """Number of namespaces to skip before returning results.""" + + +Op = Union[GetOp, SearchOp, PutOp, ListNamespacesOp] +Result = Union[Item, list[Item], list[tuple[str, ...]], None] + + +class InvalidNamespaceError(ValueError): + """Provided namespace is invalid.""" + + +def _validate_namespace(namespace: tuple[str, ...]) -> None: + for label in namespace: + if "." in label: + raise InvalidNamespaceError( + f"Invalid namespace label '{label}'. Namespace labels cannot contain periods ('.')." + ) + + +class BaseStore(ABC): + """Abstract base class for key-value stores.""" + + __slots__ = ("__weakref__",) + + @abstractmethod + def batch(self, ops: Iterable[Op]) -> list[Result]: + """Execute multiple operations synchronously in a single batch. + + Args: + ops: An iterable of operations to execute. + + Returns: + A list of results, where each result corresponds to an operation in the input. + The order of results matches the order of input operations. + """ + + @abstractmethod + async def abatch(self, ops: Iterable[Op]) -> list[Result]: + """Execute multiple operations asynchronously in a single batch. + + Args: + ops: An iterable of operations to execute. + + Returns: + A list of results, where each result corresponds to an operation in the input. + The order of results matches the order of input operations. + """ + + def get(self, namespace: tuple[str, ...], key: str) -> Optional[Item]: + """Retrieve a single item. + + Args: + namespace: Hierarchical path for the item. + key: Unique identifier within the namespace. + + Returns: + The retrieved item or None if not found. + """ + return self.batch([GetOp(namespace, key)])[0] + + def search( + self, + namespace_prefix: tuple[str, ...], + /, + *, + filter: Optional[dict[str, Any]] = None, + limit: int = 10, + offset: int = 0, + ) -> list[Item]: + """Search for items within a namespace prefix. + + Args: + namespace_prefix: Hierarchical path prefix to search within. + filter: Key-value pairs to filter results. + limit: Maximum number of items to return. + offset: Number of items to skip before returning results. + + Returns: + List of items matching the search criteria. + """ + return self.batch([SearchOp(namespace_prefix, filter, limit, offset)])[0] + + def put(self, namespace: tuple[str, ...], key: str, value: dict[str, Any]) -> None: + """Store or update an item. + + Args: + namespace: Hierarchical path for the item. + key: Unique identifier within the namespace. + value: Dictionary containing the item's data. + """ + _validate_namespace(namespace) + self.batch([PutOp(namespace, key, value)]) + + def delete(self, namespace: tuple[str, ...], key: str) -> None: + """Delete an item. + + Args: + namespace: Hierarchical path for the item. + key: Unique identifier within the namespace. + """ + self.batch([PutOp(namespace, key, None)]) + + def list_namespaces( + self, + *, + prefix: Optional[NameSpacePath] = None, + suffix: Optional[NameSpacePath] = None, + max_depth: Optional[int] = None, + limit: int = 100, + offset: int = 0, + ) -> list[tuple[str, ...]]: + """List and filter namespaces in the store. + + Used to explore the organization of data, + find specific collections, or navigate the namespace hierarchy. + + Args: + prefix (Optional[Tuple[str, ...]]): Filter namespaces that start with this path. + suffix (Optional[Tuple[str, ...]]): Filter namespaces that end with this path. + max_depth (Optional[int]): Return namespaces up to this depth in the hierarchy. + Namespaces deeper than this level will be truncated to this depth. + limit (int): Maximum number of namespaces to return (default 100). + offset (int): Number of namespaces to skip for pagination (default 0). + + Returns: + List[Tuple[str, ...]]: A list of namespace tuples that match the criteria. + Each tuple represents a full namespace path up to `max_depth`. + + Examples: + + Setting max_depth=3. Given the namespaces: + # ("a", "b", "c") + # ("a", "b", "d", "e") + # ("a", "b", "d", "i") + # ("a", "b", "f") + # ("a", "c", "f") + store.list_namespaces(prefix=("a", "b"), max_depth=3) + # [("a", "b", "c"), ("a", "b", "d"), ("a", "b", "f")] + """ + match_conditions = [] + if prefix: + match_conditions.append(MatchCondition(match_type="prefix", path=prefix)) + if suffix: + match_conditions.append(MatchCondition(match_type="suffix", path=suffix)) + + op = ListNamespacesOp( + match_conditions=tuple(match_conditions), + max_depth=max_depth, + limit=limit, + offset=offset, + ) + return self.batch([op])[0] + + async def aget(self, namespace: tuple[str, ...], key: str) -> Optional[Item]: + """Asynchronously retrieve a single item. + + Args: + namespace: Hierarchical path for the item. + key: Unique identifier within the namespace. + + Returns: + The retrieved item or None if not found. + """ + return (await self.abatch([GetOp(namespace, key)]))[0] + + async def asearch( + self, + namespace_prefix: tuple[str, ...], + /, + *, + filter: Optional[dict[str, Any]] = None, + limit: int = 10, + offset: int = 0, + ) -> list[Item]: + """Asynchronously search for items within a namespace prefix. + + Args: + namespace_prefix: Hierarchical path prefix to search within. + filter: Key-value pairs to filter results. + limit: Maximum number of items to return. + offset: Number of items to skip before returning results. + + Returns: + List of items matching the search criteria. + """ + return (await self.abatch([SearchOp(namespace_prefix, filter, limit, offset)]))[ + 0 + ] + + async def aput( + self, namespace: tuple[str, ...], key: str, value: dict[str, Any] + ) -> None: + """Asynchronously store or update an item. + + Args: + namespace: Hierarchical path for the item. + key: Unique identifier within the namespace. + value: Dictionary containing the item's data. + """ + _validate_namespace(namespace) + await self.abatch([PutOp(namespace, key, value)]) + + async def adelete(self, namespace: tuple[str, ...], key: str) -> None: + """Asynchronously delete an item. + + Args: + namespace: Hierarchical path for the item. + key: Unique identifier within the namespace. + """ + await self.abatch([PutOp(namespace, key, None)]) + + async def alist_namespaces( + self, + *, + prefix: Optional[NameSpacePath] = None, + suffix: Optional[NameSpacePath] = None, + max_depth: Optional[int] = None, + limit: int = 100, + offset: int = 0, + ) -> list[tuple[str, ...]]: + """List and filter namespaces in the store asynchronously. + + Used to explore the organization of data, + find specific collections, or navigate the namespace hierarchy. + + Args: + prefix (Optional[Tuple[str, ...]]): Filter namespaces that start with this path. + suffix (Optional[Tuple[str, ...]]): Filter namespaces that end with this path. + max_depth (Optional[int]): Return namespaces up to this depth in the hierarchy. + Namespaces deeper than this level will be truncated to this depth. + limit (int): Maximum number of namespaces to return (default 100). + offset (int): Number of namespaces to skip for pagination (default 0). + + Returns: + List[Tuple[str, ...]]: A list of namespace tuples that match the criteria. + Each tuple represents a full namespace path up to `max_depth`. + + Examples: + + Setting max_depth=3. Given the namespaces: + # ("a", "b", "c") + # ("a", "b", "d", "e") + # ("a", "b", "d", "i") + # ("a", "b", "f") + # ("a", "c", "f") + await store.alist_namespaces(prefix=("a", "b"), max_depth=3) + # [("a", "b", "c"), ("a", "b", "d"), ("a", "b", "f")] + """ + match_conditions = [] + if prefix: + match_conditions.append(MatchCondition(match_type="prefix", path=prefix)) + if suffix: + match_conditions.append(MatchCondition(match_type="suffix", path=suffix)) + + op = ListNamespacesOp( + match_conditions=tuple(match_conditions), + max_depth=max_depth, + limit=limit, + offset=offset, + ) + return (await self.abatch([op]))[0] diff --git a/libs/checkpoint/langgraph/store/batch.py b/libs/checkpoint/langgraph/store/batch.py new file mode 100644 index 000000000..8283a7a66 --- /dev/null +++ b/libs/checkpoint/langgraph/store/batch.py @@ -0,0 +1,88 @@ +import asyncio +import weakref +from typing import Any, Optional + +from langgraph.store.base import BaseStore, GetOp, Item, Op, PutOp, SearchOp + + +class AsyncBatchedBaseStore(BaseStore): + """Efficiently batch operations in a background task.""" + + __slots__ = ("_loop", "_aqueue", "_task") + + def __init__(self) -> None: + self._loop = asyncio.get_running_loop() + self._aqueue: dict[asyncio.Future, Op] = {} + self._task = self._loop.create_task(_run(self._aqueue, weakref.ref(self))) + + def __del__(self) -> None: + self._task.cancel() + + async def aget( + self, + namespace: tuple[str, ...], + key: str, + ) -> Optional[Item]: + fut = self._loop.create_future() + self._aqueue[fut] = GetOp(namespace, key) + return await fut + + async def asearch( + self, + namespace_prefix: tuple[str, ...], + /, + *, + filter: Optional[dict[str, Any]] = None, + limit: int = 10, + offset: int = 0, + ) -> list[Item]: + fut = self._loop.create_future() + self._aqueue[fut] = SearchOp(namespace_prefix, filter, limit, offset) + return await fut + + async def aput( + self, + namespace: tuple[str, ...], + key: str, + value: dict[str, Any], + ) -> None: + fut = self._loop.create_future() + self._aqueue[fut] = PutOp(namespace, key, value) + return await fut + + async def adelete( + self, + namespace: tuple[str, ...], + key: str, + ) -> None: + fut = self._loop.create_future() + self._aqueue[fut] = PutOp(namespace, key, None) + return await fut + + +async def _run( + aqueue: dict[asyncio.Future, Op], store: weakref.ReferenceType[BaseStore] +) -> None: + while True: + await asyncio.sleep(0) + if not aqueue: + continue + if s := store(): + # get the operations to run + taken = aqueue.copy() + # action each operation + try: + results = await s.abatch(taken.values()) + # set the results of each operation + for fut, result in zip(taken, results): + fut.set_result(result) + except Exception as e: + for fut in taken: + fut.set_exception(e) + # remove the operations from the queue + for fut in taken: + del aqueue[fut] + else: + break + # remove strong ref to store + del s diff --git a/libs/checkpoint/langgraph/store/memory.py b/libs/checkpoint/langgraph/store/memory.py new file mode 100644 index 000000000..69a315096 --- /dev/null +++ b/libs/checkpoint/langgraph/store/memory.py @@ -0,0 +1,119 @@ +from collections import defaultdict +from datetime import datetime, timezone +from typing import Iterable + +from langgraph.store.base import ( + BaseStore, + GetOp, + Item, + ListNamespacesOp, + MatchCondition, + Op, + PutOp, + Result, + SearchOp, +) + + +class InMemoryStore(BaseStore): + """A KV store backed by an in-memory python dictionary. + + Useful for testing/experimentation and lightweight PoC's. + For actual persistence, use a Store backed by a proper database. + """ + + __slots__ = ("_data",) + + def __init__(self) -> None: + self._data: dict[tuple[str, ...], dict[str, Item]] = defaultdict(dict) + + def batch(self, ops: Iterable[Op]) -> list[Result]: + results: list[Result] = [] + for op in ops: + if isinstance(op, GetOp): + item = self._data[op.namespace].get(op.key) + results.append(item) + elif isinstance(op, SearchOp): + candidates = [ + item + for namespace, items in self._data.items() + if ( + namespace[: len(op.namespace_prefix)] == op.namespace_prefix + if len(namespace) >= len(op.namespace_prefix) + else False + ) + for item in items.values() + ] + if op.filter: + candidates = [ + item + for item in candidates + if item.value.items() >= op.filter.items() + ] + results.append(candidates[op.offset : op.offset + op.limit]) + elif isinstance(op, PutOp): + if op.value is None: + self._data[op.namespace].pop(op.key, None) + elif op.key in self._data[op.namespace]: + self._data[op.namespace][op.key].value = op.value + self._data[op.namespace][op.key].updated_at = datetime.now( + timezone.utc + ) + else: + self._data[op.namespace][op.key] = Item( + value=op.value, + key=op.key, + namespace=op.namespace, + created_at=datetime.now(timezone.utc), + updated_at=datetime.now(timezone.utc), + ) + results.append(None) + elif isinstance(op, ListNamespacesOp): + results.append(self._handle_list_namespaces(op)) + return results + + async def abatch(self, ops: Iterable[Op]) -> list[Result]: + return self.batch(ops) + + def _handle_list_namespaces(self, op: ListNamespacesOp) -> list[tuple[str, ...]]: + all_namespaces = list( + self._data.keys() + ) # Avoid collection size changing while iterating + namespaces = all_namespaces + if op.match_conditions: + namespaces = [ + ns + for ns in namespaces + if all(_does_match(condition, ns) for condition in op.match_conditions) + ] + + if op.max_depth is not None: + namespaces = sorted({ns[: op.max_depth] for ns in namespaces}) + else: + namespaces = sorted(namespaces) + return namespaces[op.offset : op.offset + op.limit] + + +def _does_match(match_condition: MatchCondition, key: tuple[str, ...]) -> bool: + match_type = match_condition.match_type + path = match_condition.path + + if len(key) < len(path): + return False + + if match_type == "prefix": + for k_elem, p_elem in zip(key, path): + if p_elem == "*": + continue # Wildcard matches any element + if k_elem != p_elem: + return False + return True + elif match_type == "suffix": + for k_elem, p_elem in zip(reversed(key), reversed(path)): + if p_elem == "*": + continue # Wildcard matches any element + if k_elem != p_elem: + return False + return True + else: + raise ValueError(f"Unsupported match type: {match_type}") diff --git a/libs/langgraph/langgraph/store/__init__.py b/libs/checkpoint/langgraph/store/py.typed similarity index 100% rename from libs/langgraph/langgraph/store/__init__.py rename to libs/checkpoint/langgraph/store/py.typed diff --git a/libs/checkpoint/tests/test_jsonplus.py b/libs/checkpoint/tests/test_jsonplus.py index 1621a0c16..869e7732e 100644 --- a/libs/checkpoint/tests/test_jsonplus.py +++ b/libs/checkpoint/tests/test_jsonplus.py @@ -16,6 +16,7 @@ from pydantic.v1 import SecretStr as SecretStrV1 from zoneinfo import ZoneInfo from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer +from langgraph.store.base import Item class InnerPydantic(BaseModel): @@ -124,6 +125,13 @@ def test_serde_jsonplus() -> None: "a_float": 1.1, "a_bytes": b"my bytes", "a_bytearray": bytearray([42]), + "my_item": Item( + value={}, + key="my-key", + namespace=("a", "name", " "), + created_at=datetime(2024, 9, 24, 17, 29, 10, 128397), + updated_at=datetime(2024, 9, 24, 17, 29, 10, 128397), + ), } serde = JsonPlusSerializer() diff --git a/libs/checkpoint/tests/test_store.py b/libs/checkpoint/tests/test_store.py new file mode 100644 index 000000000..317d9a720 --- /dev/null +++ b/libs/checkpoint/tests/test_store.py @@ -0,0 +1,261 @@ +import asyncio +from datetime import datetime +from typing import Iterable + +from pytest_mock import MockerFixture + +from langgraph.store.base import GetOp, Item, Op, Result +from langgraph.store.batch import AsyncBatchedBaseStore +from langgraph.store.memory import InMemoryStore + + +async def test_async_batch_store(mocker: MockerFixture) -> None: + abatch = mocker.stub() + + class MockStore(AsyncBatchedBaseStore): + def batch(self, ops: Iterable[Op]) -> list[Result]: + raise NotImplementedError + + async def abatch(self, ops: Iterable[Op]) -> list[Result]: + assert all(isinstance(op, GetOp) for op in ops) + abatch(ops) + return [ + Item( + value={}, + key=getattr(op, "key", ""), + namespace=getattr(op, "namespace", ()), + created_at=datetime(2024, 9, 24, 17, 29, 10, 128397), + updated_at=datetime(2024, 9, 24, 17, 29, 10, 128397), + ) + for op in ops + ] + + store = MockStore() + + # concurrent calls are batched + results = await asyncio.gather( + store.aget(namespace=("a",), key="b"), + store.aget(namespace=("c",), key="d"), + ) + assert results == [ + Item( + value={}, + key="b", + namespace=("a",), + created_at=datetime(2024, 9, 24, 17, 29, 10, 128397), + updated_at=datetime(2024, 9, 24, 17, 29, 10, 128397), + ), + Item( + value={}, + key="d", + namespace=("c",), + created_at=datetime(2024, 9, 24, 17, 29, 10, 128397), + updated_at=datetime(2024, 9, 24, 17, 29, 10, 128397), + ), + ] + assert abatch.call_count == 1 + assert [tuple(c.args[0]) for c in abatch.call_args_list] == [ + ( + GetOp(("a",), "b"), + GetOp(("c",), "d"), + ), + ] + + +def test_list_namespaces_basic() -> None: + store = InMemoryStore() + + namespaces = [ + ("a", "b", "c"), + ("a", "b", "d", "e"), + ("a", "b", "d", "i"), + ("a", "b", "f"), + ("a", "c", "f"), + ("b", "a", "f"), + ("users", "123"), + ("users", "456", "settings"), + ("admin", "users", "789"), + ] + + for i, ns in enumerate(namespaces): + store.put(namespace=ns, key=f"id_{i}", value={"data": f"value_{i:02d}"}) + + result = store.list_namespaces(prefix=("a", "b")) + expected = [ + ("a", "b", "c"), + ("a", "b", "d", "e"), + ("a", "b", "d", "i"), + ("a", "b", "f"), + ] + assert sorted(result) == sorted(expected) + + result = store.list_namespaces(suffix=("f",)) + expected = [ + ("a", "b", "f"), + ("a", "c", "f"), + ("b", "a", "f"), + ] + assert sorted(result) == sorted(expected) + + result = store.list_namespaces(prefix=("a",), suffix=("f",)) + expected = [ + ("a", "b", "f"), + ("a", "c", "f"), + ] + assert sorted(result) == sorted(expected) + + # Test max_depth + result = store.list_namespaces(prefix=("a", "b"), max_depth=3) + expected = [ + ("a", "b", "c"), + ("a", "b", "d"), + ("a", "b", "f"), + ] + assert sorted(result) == sorted(expected) + + # Test limit and offset + result = store.list_namespaces(prefix=("a", "b"), limit=2) + expected = [ + ("a", "b", "c"), + ("a", "b", "d", "e"), + ] + assert result == expected + + result = store.list_namespaces(prefix=("a", "b"), offset=2) + expected = [ + ("a", "b", "d", "i"), + ("a", "b", "f"), + ] + assert result == expected + + result = store.list_namespaces(prefix=("a", "*", "f")) + expected = [ + ("a", "b", "f"), + ("a", "c", "f"), + ] + assert sorted(result) == sorted(expected) + + result = store.list_namespaces(suffix=("*", "f")) + expected = [ + ("a", "b", "f"), + ("a", "c", "f"), + ("b", "a", "f"), + ] + assert sorted(result) == sorted(expected) + + result = store.list_namespaces(prefix=("nonexistent",)) + assert result == [] + + result = store.list_namespaces(prefix=("users", "123")) + expected = [("users", "123")] + assert result == expected + + +def test_list_namespaces_with_wildcards() -> None: + store = InMemoryStore() + + namespaces = [ + ("users", "123"), + ("users", "456"), + ("users", "789", "settings"), + ("admin", "users", "789"), + ("guests", "123"), + ("guests", "456", "preferences"), + ] + + for i, ns in enumerate(namespaces): + store.put(namespace=ns, key=f"id_{i}", value={"data": f"value_{i:02d}"}) + + result = store.list_namespaces(prefix=("users", "*")) + expected = [ + ("users", "123"), + ("users", "456"), + ("users", "789", "settings"), + ] + assert sorted(result) == sorted(expected) + + result = store.list_namespaces(suffix=("*", "preferences")) + expected = [ + ("guests", "456", "preferences"), + ] + assert result == expected + + result = store.list_namespaces(prefix=("*", "users"), suffix=("*", "settings")) + + assert result == [] + + store.put( + namespace=("admin", "users", "settings", "789"), + key="foo", + value={"data": "some_val"}, + ) + expected = [ + ("admin", "users", "settings", "789"), + ] + + +def test_list_namespaces_pagination() -> None: + store = InMemoryStore() + + for i in range(20): + ns = ("namespace", f"sub_{i:02d}") + store.put(namespace=ns, key=f"id_{i:02d}", value={"data": f"value_{i:02d}"}) + + result = store.list_namespaces(prefix=("namespace",), limit=5, offset=0) + expected = [("namespace", f"sub_{i:02d}") for i in range(5)] + assert result == expected + + result = store.list_namespaces(prefix=("namespace",), limit=5, offset=5) + expected = [("namespace", f"sub_{i:02d}") for i in range(5, 10)] + assert result == expected + + result = store.list_namespaces(prefix=("namespace",), limit=5, offset=15) + expected = [("namespace", f"sub_{i:02d}") for i in range(15, 20)] + assert result == expected + + +def test_list_namespaces_max_depth() -> None: + store = InMemoryStore() + + namespaces = [ + ("a", "b", "c", "d"), + ("a", "b", "c", "e"), + ("a", "b", "f"), + ("a", "g"), + ("h", "i", "j", "k"), + ] + + for i, ns in enumerate(namespaces): + store.put(namespace=ns, key=f"id_{i}", value={"data": f"value_{i:02d}"}) + + result = store.list_namespaces(max_depth=2) + expected = [ + ("a", "b"), + ("a", "g"), + ("h", "i"), + ] + assert sorted(result) == sorted(expected) + + +def test_list_namespaces_no_conditions() -> None: + store = InMemoryStore() + + namespaces = [ + ("a", "b"), + ("c", "d"), + ("e", "f", "g"), + ] + + for i, ns in enumerate(namespaces): + store.put(namespace=ns, key=f"id_{i}", value={"data": f"value_{i:02d}"}) + + result = store.list_namespaces() + expected = namespaces + assert sorted(result) == sorted(expected) + + +def test_list_namespaces_empty_store() -> None: + store = InMemoryStore() + + result = store.list_namespaces() + assert result == [] diff --git a/libs/langgraph/langgraph/managed/shared_value.py b/libs/langgraph/langgraph/managed/shared_value.py index 0f94dd1a1..7c8f45b14 100644 --- a/libs/langgraph/langgraph/managed/shared_value.py +++ b/libs/langgraph/langgraph/managed/shared_value.py @@ -21,7 +21,7 @@ from langgraph.managed.base import ( ConfiguredManagedValue, WritableManagedValue, ) -from langgraph.store.base import BaseStore +from langgraph.store.base import BaseStore, PutOp V = dict[str, Any] @@ -58,8 +58,8 @@ class SharedValue(WritableManagedValue[Value, Update]): def enter(cls, config: RunnableConfig, **kwargs: Any) -> Iterator[Self]: with super().enter(config, **kwargs) as value: if value.store is not None: - saved = value.store.list([value.ns]) - value.value = saved[value.ns] or {} + saved = value.store.search(value.ns) + value.value = {it.key: it.value for it in saved} yield value @classmethod @@ -67,8 +67,8 @@ class SharedValue(WritableManagedValue[Value, Update]): async def aenter(cls, config: RunnableConfig, **kwargs: Any) -> AsyncIterator[Self]: 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] or {} + saved = await value.store.asearch(value.ns) + value.value = {it.key: it.value for it in saved} yield value def __init__( @@ -87,7 +87,7 @@ class SharedValue(WritableManagedValue[Value, Update]): if self.store is None: pass elif scope_value := config[CONF].get(self.scope): - self.ns = f"scoped:{scope}:{key}:{scope_value}" + self.ns = ("scoped", scope, key, scope_value) else: raise ValueError( f"Scope {scope} for shared state key not in config.configurable" @@ -96,31 +96,29 @@ class SharedValue(WritableManagedValue[Value, Update]): def __call__(self, step: int) -> Value: return self.value.copy() - def _process_update( - self, values: Sequence[Update] - ) -> list[tuple[str, str, Optional[dict[str, Any]]]]: - writes: list[tuple[str, str, Optional[dict[str, Any]]]] = [] + def _process_update(self, values: Sequence[Update]) -> list[PutOp]: + writes: list[PutOp] = [] for vv in values: for k, v in vv.items(): if v is None: if k in self.value: del self.value[k] - writes.append((self.ns, k, None)) + writes.append(PutOp(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)) + writes.append(PutOp(self.ns, k, v)) return writes def update(self, values: Sequence[Update]) -> None: if self.store is None: self._process_update(values) else: - return self.store.put(self._process_update(values)) + return self.store.batch(self._process_update(values)) async def aupdate(self, writes: Sequence[Update]) -> None: if self.store is None: self._process_update(writes) else: - return await self.store.aput(self._process_update(writes)) + return await self.store.abatch(self._process_update(writes)) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index cd3c8150d..3d7414494 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -61,6 +61,7 @@ from langgraph.constants import ( CONFIG_KEY_READ, CONFIG_KEY_RESUMING, CONFIG_KEY_SEND, + CONFIG_KEY_STORE, CONFIG_KEY_STREAM, CONFIG_KEY_STREAM_WRITER, CONFIG_KEY_TASK_ID, @@ -1070,6 +1071,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): Union[All, Sequence[str]], Union[All, Sequence[str]], Optional[BaseCheckpointSaver], + Optional[BaseStore], ]: if config["recursion_limit"] < 1: raise ValueError("recursion_limit must be at least 1") @@ -1096,6 +1098,10 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): raise ValueError( f"Checkpointer requires one or more of the following 'configurable' keys: {[s.id for s in checkpointer.config_specs]}" ) + if CONFIG_KEY_STORE in config.get(CONF, {}): + store: Optional[BaseStore] = config[CONF][CONFIG_KEY_STORE] + else: + store = self.store return ( debug, set(stream_mode), @@ -1103,6 +1109,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): interrupt_before, interrupt_after, checkpointer, + store, ) def stream( @@ -1219,6 +1226,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): interrupt_before_, interrupt_after_, checkpointer, + store, ) = self._defaults( config, stream_mode=stream_mode, @@ -1241,7 +1249,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): input, stream=StreamProtocol(stream.put, stream_modes), config=config, - store=self.store, + store=store, checkpointer=checkpointer, nodes=self.nodes, specs=self.channels, @@ -1434,6 +1442,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): interrupt_before_, interrupt_after_, checkpointer, + store, ) = self._defaults( config, stream_mode=stream_mode, @@ -1456,7 +1465,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): input, stream=StreamProtocol(stream.put_nowait, stream_modes), config=config, - store=self.store, + store=store, checkpointer=checkpointer, nodes=self.nodes, specs=self.channels, diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 99de1f130..daa942a39 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -35,6 +35,7 @@ from langgraph.constants import ( CONFIG_KEY_CHECKPOINTER, CONFIG_KEY_READ, CONFIG_KEY_SEND, + CONFIG_KEY_STORE, CONFIG_KEY_TASK_ID, EMPTY_SEQ, INTERRUPT, @@ -54,6 +55,7 @@ from langgraph.pregel.io import read_channel, read_channels from langgraph.pregel.log import logger from langgraph.pregel.manager import ChannelsManager from langgraph.pregel.read import PregelNode +from langgraph.store.base import BaseStore from langgraph.types import All, PregelExecutableTask, PregelTask from langgraph.utils.config import merge_configs, patch_config @@ -274,6 +276,7 @@ def prepare_next_tasks( step: int, *, for_execution: Literal[False], + store: Literal[None] = None, checkpointer: Literal[None] = None, manager: Literal[None] = None, ) -> dict[str, PregelTask]: ... @@ -289,6 +292,7 @@ def prepare_next_tasks( step: int, *, for_execution: Literal[True], + store: Optional[BaseStore], checkpointer: Optional[BaseCheckpointSaver], manager: Union[None, ParentRunManager, AsyncParentRunManager], ) -> dict[str, PregelExecutableTask]: ... @@ -303,6 +307,7 @@ def prepare_next_tasks( step: int, *, for_execution: bool, + store: Optional[BaseStore] = None, checkpointer: Optional[BaseCheckpointSaver] = None, manager: Union[None, ParentRunManager, AsyncParentRunManager] = None, ) -> Union[dict[str, PregelTask], dict[str, PregelExecutableTask]]: @@ -322,6 +327,7 @@ def prepare_next_tasks( config=config, step=step, for_execution=for_execution, + store=store, checkpointer=checkpointer, manager=manager, ): @@ -339,6 +345,7 @@ def prepare_next_tasks( config=config, step=step, for_execution=for_execution, + store=store, checkpointer=checkpointer, manager=manager, ): @@ -357,6 +364,7 @@ def prepare_single_task( config: RunnableConfig, step: int, for_execution: bool, + store: Optional[BaseStore] = None, checkpointer: Optional[BaseCheckpointSaver] = None, manager: Union[None, ParentRunManager, AsyncParentRunManager] = None, ) -> Union[None, PregelTask, PregelExecutableTask]: @@ -438,6 +446,9 @@ def prepare_single_task( PregelTaskWrites(packet.node, writes, triggers), config, ), + CONFIG_KEY_STORE: ( + store or configurable.get(CONFIG_KEY_STORE) + ), CONFIG_KEY_CHECKPOINTER: ( checkpointer or configurable.get(CONFIG_KEY_CHECKPOINTER) @@ -545,6 +556,9 @@ def prepare_single_task( PregelTaskWrites(name, writes, triggers), config, ), + CONFIG_KEY_STORE: ( + store or configurable.get(CONFIG_KEY_STORE) + ), CONFIG_KEY_CHECKPOINTER: ( checkpointer or configurable.get(CONFIG_KEY_CHECKPOINTER) diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 071f798d7..6353680c3 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -101,7 +101,6 @@ from langgraph.pregel.manager import AsyncChannelsManager, ChannelsManager from langgraph.pregel.read import PregelNode from langgraph.pregel.utils import get_new_channel_versions from langgraph.store.base import BaseStore -from langgraph.store.batch import AsyncBatchedStore from langgraph.types import All, PregelExecutableTask, StreamMode from langgraph.utils.config import patch_configurable @@ -367,6 +366,7 @@ class PregelLoop: self.step, for_execution=True, manager=manager, + store=self.store, checkpointer=self.checkpointer, ) # we don't need to save the writes for the last task that completes @@ -495,6 +495,7 @@ class PregelLoop: self.config, self.step, for_execution=True, + store=None, checkpointer=None, manager=None, ) @@ -783,7 +784,6 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): check_subgraphs=check_subgraphs, debug=debug, ) - self.store = AsyncBatchedStore(self.store) if self.store else None self.stack = AsyncExitStack() if checkpointer: self.checkpointer_get_next_version = checkpointer.get_next_version diff --git a/libs/langgraph/langgraph/store/base.py b/libs/langgraph/langgraph/store/base.py deleted file mode 100644 index 046483f2e..000000000 --- a/libs/langgraph/langgraph/store/base.py +++ /dev/null @@ -1,21 +0,0 @@ -from typing import Any, List, Optional - -V = Any - - -class BaseStore: - 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 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 diff --git a/libs/langgraph/langgraph/store/batch.py b/libs/langgraph/langgraph/store/batch.py deleted file mode 100644 index 54eb20d47..000000000 --- a/libs/langgraph/langgraph/store/batch.py +++ /dev/null @@ -1,65 +0,0 @@ -import asyncio -from typing import NamedTuple, Optional, Union - -from langgraph.store.base import BaseStore, V - - -class ListOp(NamedTuple): - prefixes: list[str] - - -class PutOp(NamedTuple): - writes: list[tuple[str, str, Optional[V]]] - - -class AsyncBatchedStore(BaseStore): - 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.store)) - - def __del__(self) -> None: - self.task.cancel() - - 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[ListOp, PutOp]], store: BaseStore -) -> 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 - lists = {f: o for f, o in taken.items() if isinstance(o, ListOp)} - if lists: - try: - results = await store.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 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: - for fut in puts: - fut.set_exception(e) diff --git a/libs/langgraph/langgraph/store/memory.py b/libs/langgraph/langgraph/store/memory.py deleted file mode 100644 index 48fa2884f..000000000 --- a/libs/langgraph/langgraph/store/memory.py +++ /dev/null @@ -1,25 +0,0 @@ -from collections import defaultdict -from typing import List, Optional - -from langgraph.store.base import BaseStore, V - - -class MemoryStore(BaseStore): - def __init__(self) -> None: - self.data: dict[str, dict[str, V]] = defaultdict(dict) - - 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, 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, Optional[V]]]) -> None: - return self.put(writes) diff --git a/libs/langgraph/langgraph/utils/runnable.py b/libs/langgraph/langgraph/utils/runnable.py index 0a90217a0..2545eb49e 100644 --- a/libs/langgraph/langgraph/utils/runnable.py +++ b/libs/langgraph/langgraph/utils/runnable.py @@ -34,7 +34,8 @@ from langchain_core.runnables.utils import Input from langchain_core.tracers._streaming import _StreamingCallbackHandler from typing_extensions import TypeGuard -from langgraph.constants import CONF, CONFIG_KEY_STREAM_WRITER +from langgraph.constants import CONF, CONFIG_KEY_STORE, CONFIG_KEY_STREAM_WRITER +from langgraph.store.base import BaseStore from langgraph.types import StreamWriter from langgraph.utils.config import ( ensure_config, @@ -62,10 +63,16 @@ ASYNCIO_ACCEPTS_CONTEXT = sys.version_info >= (3, 11) KWARGS_CONFIG_KEYS: tuple[tuple[str, tuple[Any, ...], str, Any], ...] = ( ( sys.intern("writer"), - (StreamWriter, inspect.Parameter.empty), + (StreamWriter, "StreamWriter", inspect.Parameter.empty), CONFIG_KEY_STREAM_WRITER, lambda _: None, ), + ( + sys.intern("store"), + (BaseStore, "BaseStore", inspect.Parameter.empty), + CONFIG_KEY_STORE, + inspect.Parameter.empty, + ), ) """List of kwargs that can be passed to functions, and their corresponding config keys, default values and type annotations.""" @@ -110,6 +117,7 @@ class RunnableCallable(Runnable): if func is None and afunc is None: raise ValueError("At least one of func or afunc must be provided.") params = inspect.signature(cast(Callable, func or afunc)).parameters + self.func_accepts_config = "config" in params self.func_accepts: dict[str, bool] = {} for kw, typ, _, _ in KWARGS_CONFIG_KEYS: @@ -140,9 +148,15 @@ class RunnableCallable(Runnable): kwargs = {**self.kwargs, **kwargs} if self.func_accepts_config: kwargs["config"] = config + _conf = config[CONF] for kw, _, ck, defv in KWARGS_CONFIG_KEYS: if self.func_accepts[kw]: - kwargs[kw] = config[CONF].get(ck, defv) + if defv is inspect.Parameter.empty and ck not in _conf: + raise ValueError( + f"Missing required config key '{ck}' for '{self.name}'." + ) + else: + kwargs[kw] = _conf.get(ck, defv) context = copy_context() if self.trace: callback_manager = get_callback_manager_for_config(config, self.tags) @@ -179,9 +193,15 @@ class RunnableCallable(Runnable): kwargs = {**self.kwargs, **kwargs} if self.func_accepts_config: kwargs["config"] = config + _conf = config[CONF] for kw, _, ck, defv in KWARGS_CONFIG_KEYS: if self.func_accepts[kw]: - kwargs[kw] = config[CONF].get(ck, defv) + if defv is inspect.Parameter.empty and ck not in _conf: + raise ValueError( + f"Missing required config key '{ck}' for '{self.name}'." + ) + else: + kwargs[kw] = _conf.get(ck, defv) context = copy_context() if self.trace: callback_manager = get_async_callback_manager_for_config(config, self.tags) diff --git a/libs/langgraph/tests/test_algo.py b/libs/langgraph/tests/test_algo.py index 203280488..4e259f29e 100644 --- a/libs/langgraph/tests/test_algo.py +++ b/libs/langgraph/tests/test_algo.py @@ -17,7 +17,16 @@ def test_prepare_next_tasks() -> None: ) assert ( prepare_next_tasks( - checkpoint, processes, channels, managed, config, 0, for_execution=True + checkpoint, + processes, + channels, + managed, + config, + 0, + for_execution=True, + checkpointer=None, + store=None, + manager=None, ) == {} ) diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 3a7c3c5bf..636f450ee 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -2,6 +2,7 @@ import json import operator import re import time +import uuid import warnings from collections import Counter from concurrent.futures import ThreadPoolExecutor @@ -70,7 +71,8 @@ from langgraph.pregel import ( StateSnapshot, ) from langgraph.pregel.retry import RetryPolicy -from langgraph.store.memory import MemoryStore +from langgraph.store.base import BaseStore +from langgraph.store.memory import InMemoryStore from langgraph.types import Interrupt, PregelTask, Send, StreamWriter from tests.any_str import AnyDict, AnyStr, AnyVersion, FloatBetween, UnsortedSequence from tests.conftest import ALL_CHECKPOINTERS_SYNC, SHOULD_CHECK_SNAPSHOTS @@ -6806,7 +6808,7 @@ def test_start_branch_then( } tool_two = tool_two_graph.compile( - store=MemoryStore(), + store=InMemoryStore(), checkpointer=checkpointer, interrupt_before=["tool_two_fast", "tool_two_slow"], ) @@ -11362,3 +11364,60 @@ def test_subgraph_retries(): app = parent.compile(checkpointer=checkpointer) with pytest.raises(RandomError): app.invoke({"count": 0}, {"configurable": {"thread_id": "foo"}}) + + +@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) +def test_store_injected(request: pytest.FixtureRequest, checkpointer_name: str) -> None: + checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") + + class State(TypedDict): + count: Annotated[int, operator.add] + + doc_id = str(uuid.uuid4()) + doc = {"some-key": "this-is-a-val"} + + def node(input: State, config: RunnableConfig, store: BaseStore): + assert isinstance(store, BaseStore) + assert isinstance(store, InMemoryStore) + store.put( + ("foo", "bar"), + doc_id, + { + **doc, + "from_thread": config["configurable"]["thread_id"], + "some_val": input["count"], + }, + ) + return {"count": 1} + + builder = StateGraph(State) + builder.add_node("node", node) + builder.add_edge("__start__", "node") + the_store = InMemoryStore() + graph = builder.compile(store=the_store, checkpointer=checkpointer) + + thread_1 = str(uuid.uuid4()) + result = graph.invoke({"count": 0}, {"configurable": {"thread_id": thread_1}}) + assert result == {"count": 1} + returned_doc = the_store.get(("foo", "bar"), doc_id).value + assert returned_doc == {**doc, "from_thread": thread_1, "some_val": 0} + assert len(the_store.search(("foo", "bar"))) == 1 + + # Check update on existing thread + result = graph.invoke({"count": 0}, {"configurable": {"thread_id": thread_1}}) + assert result == {"count": 2} + returned_doc = the_store.get(("foo", "bar"), doc_id).value + assert returned_doc == {**doc, "from_thread": thread_1, "some_val": 1} + assert len(the_store.search(("foo", "bar"))) == 1 + + thread_2 = str(uuid.uuid4()) + + result = graph.invoke({"count": 0}, {"configurable": {"thread_id": thread_2}}) + assert result == {"count": 1} + returned_doc = the_store.get(("foo", "bar"), doc_id).value + assert returned_doc == { + **doc, + "from_thread": thread_2, + "some_val": 0, + } # Overwrites the whole doc + assert len(the_store.search(("foo", "bar"))) == 1 # still overwriting the same one diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index ccc38877a..6a7252391 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -2,6 +2,7 @@ import asyncio import operator import re import sys +import uuid from collections import Counter from contextlib import asynccontextmanager, contextmanager from time import perf_counter @@ -24,9 +25,7 @@ from uuid import UUID import httpx import pytest -from langchain_core.messages import ( - ToolCall, -) +from langchain_core.messages import ToolCall from langchain_core.runnables import ( RunnableConfig, RunnableLambda, @@ -57,18 +56,12 @@ from langgraph.graph import END, Graph, StateGraph from langgraph.graph.graph import START from langgraph.graph.message import MessageGraph, add_messages from langgraph.managed.shared_value import SharedValue -from langgraph.prebuilt.chat_agent_executor import ( - create_tool_calling_executor, -) +from langgraph.prebuilt.chat_agent_executor import create_tool_calling_executor from langgraph.prebuilt.tool_node import ToolNode -from langgraph.pregel import ( - Channel, - GraphRecursionError, - Pregel, - StateSnapshot, -) +from langgraph.pregel import Channel, GraphRecursionError, Pregel, StateSnapshot from langgraph.pregel.retry import RetryPolicy -from langgraph.store.memory import MemoryStore +from langgraph.store.base import BaseStore +from langgraph.store.memory import InMemoryStore from langgraph.types import Interrupt, PregelTask, Send, StreamWriter from tests.any_str import AnyDict, AnyStr, AnyVersion, FloatBetween, UnsortedSequence from tests.conftest import ( @@ -5423,7 +5416,7 @@ async def test_start_branch_then(checkpointer_name: str) -> None: async with awith_checkpointer(checkpointer_name) as checkpointer: tool_two = tool_two_graph.compile( - store=MemoryStore(), + store=InMemoryStore(), checkpointer=checkpointer, interrupt_before=["tool_two_fast", "tool_two_slow"], ) @@ -9645,3 +9638,67 @@ async def test_checkpointer_null_pending_writes() -> None: assert (await graph.ainvoke([], {"configurable": {"thread_id": "foo"}})) == [ "1" ] * 4 + + +@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) +async def test_store_injected_async(checkpointer_name: str) -> None: + class State(TypedDict): + count: Annotated[int, operator.add] + + doc_id = str(uuid.uuid4()) + doc = {"some-key": "this-is-a-val"} + + async def node(input: State, config: RunnableConfig, store: BaseStore): + assert isinstance(store, BaseStore) + assert isinstance(store, InMemoryStore) + await store.aput( + ("foo", "bar"), + doc_id, + { + **doc, + "from_thread": config["configurable"]["thread_id"], + "some_val": input["count"], + }, + ) + return {"count": 1} + + builder = StateGraph(State) + builder.add_node("node", node) + builder.add_edge("__start__", "node") + the_store = InMemoryStore() + async with awith_checkpointer(checkpointer_name) as checkpointer: + graph = builder.compile(store=the_store, checkpointer=checkpointer) + + thread_1 = str(uuid.uuid4()) + result = await graph.ainvoke( + {"count": 0}, {"configurable": {"thread_id": thread_1}} + ) + assert result == {"count": 1} + returned_doc = (await the_store.aget(("foo", "bar"), doc_id)).value + assert returned_doc == {**doc, "from_thread": thread_1, "some_val": 0} + assert len((await the_store.asearch(("foo", "bar")))) == 1 + + # Check update on existing thread + result = await graph.ainvoke( + {"count": 0}, {"configurable": {"thread_id": thread_1}} + ) + assert result == {"count": 2} + returned_doc = (await the_store.aget(("foo", "bar"), doc_id)).value + assert returned_doc == {**doc, "from_thread": thread_1, "some_val": 1} + assert len((await the_store.asearch(("foo", "bar")))) == 1 + + thread_2 = str(uuid.uuid4()) + + result = await graph.ainvoke( + {"count": 0}, {"configurable": {"thread_id": thread_2}} + ) + assert result == {"count": 1} + returned_doc = (await the_store.aget(("foo", "bar"), doc_id)).value + assert returned_doc == { + **doc, + "from_thread": thread_2, + "some_val": 0, + } # Overwrites the whole doc + assert ( + len((await the_store.asearch(("foo", "bar")))) == 1 + ) # still overwriting the same one diff --git a/libs/langgraph/tests/test_runnable.py b/libs/langgraph/tests/test_runnable.py new file mode 100644 index 000000000..64d858fd4 --- /dev/null +++ b/libs/langgraph/tests/test_runnable.py @@ -0,0 +1,43 @@ +from __future__ import annotations + +from typing import Any + +from langgraph.store.base import BaseStore +from langgraph.types import StreamWriter +from langgraph.utils.runnable import RunnableCallable + + +def test_runnable_callable_func_accepts(): + def sync_func(x: Any) -> str: + return f"{x}" + + async def async_func(x: Any) -> str: + return f"{x}" + + def func_with_store(x: Any, store: BaseStore) -> str: + return f"{x}" + + def func_with_writer(x: Any, writer: StreamWriter) -> str: + return f"{x}" + + async def afunc_with_store(x: Any, store: BaseStore) -> str: + return f"{x}" + + async def afunc_with_writer(x: Any, writer: StreamWriter) -> str: + return f"{x}" + + runnables = { + "sync": RunnableCallable(sync_func), + "async": RunnableCallable(func=None, afunc=async_func), + "with_store": RunnableCallable(func_with_store), + "with_writer": RunnableCallable(func_with_writer), + "awith_store": RunnableCallable(afunc_with_store), + "awith_writer": RunnableCallable(afunc_with_writer), + } + + expected_store = {"with_store": True, "awith_store": True} + expected_writer = {"with_writer": True, "awith_writer": True} + + for name, runnable in runnables.items(): + assert runnable.func_accepts["writer"] == expected_writer.get(name, False) + assert runnable.func_accepts["store"] == expected_store.get(name, False) diff --git a/libs/langgraph/tests/test_store.py b/libs/langgraph/tests/test_store.py deleted file mode 100644 index 71494adaf..000000000 --- a/libs/langgraph/tests/test_store.py +++ /dev/null @@ -1,38 +0,0 @@ -import asyncio -from typing import Any, Optional - -import pytest -from pytest_mock import MockerFixture - -from langgraph.store.base import BaseStore -from langgraph.store.batch import AsyncBatchedStore - -pytestmark = pytest.mark.anyio - - -async def test_async_batch_store(mocker: MockerFixture) -> None: - aget = mocker.stub() - alist = mocker.stub() - - class MockStore(BaseStore): - async def aget( - self, pairs: list[tuple[str, str]] - ) -> dict[tuple[str, str], Optional[dict[str, Any]]]: - aget(pairs) - return {pair: 1 for pair in pairs} - - async def alist(self, prefixes: list[str]) -> dict[str, dict[str, Any]]: - alist(prefixes) - return {prefix: {prefix: 1} for prefix in prefixes} - - store = AsyncBatchedStore(MockStore()) - - # concurrent calls are batched - 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"],), - ] diff --git a/libs/scheduler-kafka/langgraph/scheduler/kafka/executor.py b/libs/scheduler-kafka/langgraph/scheduler/kafka/executor.py index c803239e8..d8e150f5b 100644 --- a/libs/scheduler-kafka/langgraph/scheduler/kafka/executor.py +++ b/libs/scheduler-kafka/langgraph/scheduler/kafka/executor.py @@ -194,6 +194,7 @@ class AsyncKafkaExecutor(AbstractAsyncContextManager): step=saved.metadata["step"] + 1, for_execution=True, checkpointer=self.graph.checkpointer, + store=self.graph.store, ): # execute task, saving writes runner = PregelRunner( diff --git a/libs/scheduler-kafka/tests/test_subgraph.py b/libs/scheduler-kafka/tests/test_subgraph.py index 54172ef9b..fd2530843 100644 --- a/libs/scheduler-kafka/tests/test_subgraph.py +++ b/libs/scheduler-kafka/tests/test_subgraph.py @@ -194,6 +194,7 @@ async def test_subgraph_w_interrupt( "__pregel_ensure_latest": True, "__pregel_dedupe_tasks": True, "__pregel_resuming": False, + '__pregel_store': None, "__pregel_task_id": history[0].tasks[0].id, "checkpoint_id": None, "checkpoint_map": { @@ -257,6 +258,7 @@ async def test_subgraph_w_interrupt( "__pregel_ensure_latest": True, "__pregel_dedupe_tasks": True, "__pregel_resuming": False, + '__pregel_store': None, "__pregel_task_id": history[0].tasks[0].id, "checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_map": { @@ -350,6 +352,7 @@ async def test_subgraph_w_interrupt( "__pregel_ensure_latest": True, "__pregel_dedupe_tasks": True, "__pregel_resuming": False, + '__pregel_store': None, "__pregel_task_id": history[0].tasks[0].id, "checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_map": { @@ -453,6 +456,7 @@ async def test_subgraph_w_interrupt( "__pregel_ensure_latest": True, "__pregel_dedupe_tasks": True, "__pregel_resuming": True, + '__pregel_store': None, "__pregel_task_id": history[1].tasks[0].id, "checkpoint_id": None, "checkpoint_map": { @@ -511,6 +515,7 @@ async def test_subgraph_w_interrupt( "__pregel_ensure_latest": True, "__pregel_dedupe_tasks": True, "__pregel_resuming": True, + '__pregel_store': None, "__pregel_task_id": history[1].tasks[0].id, "checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_map": { @@ -625,6 +630,7 @@ async def test_subgraph_w_interrupt( "__pregel_ensure_latest": True, "__pregel_dedupe_tasks": True, "__pregel_resuming": True, + '__pregel_store': None, "__pregel_task_id": history[1].tasks[0].id, "checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_map": { diff --git a/libs/scheduler-kafka/tests/test_subgraph_sync.py b/libs/scheduler-kafka/tests/test_subgraph_sync.py index 16576b343..32d0ceea0 100644 --- a/libs/scheduler-kafka/tests/test_subgraph_sync.py +++ b/libs/scheduler-kafka/tests/test_subgraph_sync.py @@ -193,6 +193,7 @@ def test_subgraph_w_interrupt( "__pregel_ensure_latest": True, "__pregel_dedupe_tasks": True, "__pregel_resuming": False, + '__pregel_store': None, "__pregel_task_id": history[0].tasks[0].id, "checkpoint_id": None, "checkpoint_map": { @@ -254,6 +255,7 @@ def test_subgraph_w_interrupt( "__pregel_read": None, "__pregel_send": None, "__pregel_ensure_latest": True, + '__pregel_store': None, "__pregel_dedupe_tasks": True, "__pregel_resuming": False, "__pregel_task_id": history[0].tasks[0].id, @@ -348,6 +350,7 @@ def test_subgraph_w_interrupt( "__pregel_send": None, "__pregel_ensure_latest": True, "__pregel_dedupe_tasks": True, + '__pregel_store': None, "__pregel_resuming": False, "__pregel_task_id": history[0].tasks[0].id, "checkpoint_id": c.config["configurable"]["checkpoint_id"], @@ -450,6 +453,7 @@ def test_subgraph_w_interrupt( "__pregel_send": None, "__pregel_ensure_latest": True, "__pregel_dedupe_tasks": True, + '__pregel_store': None, "__pregel_resuming": True, "__pregel_task_id": history[1].tasks[0].id, "checkpoint_id": None, @@ -508,6 +512,7 @@ def test_subgraph_w_interrupt( "__pregel_send": None, "__pregel_ensure_latest": True, "__pregel_dedupe_tasks": True, + '__pregel_store': None, "__pregel_resuming": True, "__pregel_task_id": history[1].tasks[0].id, "checkpoint_id": c.config["configurable"]["checkpoint_id"], @@ -623,6 +628,7 @@ def test_subgraph_w_interrupt( "__pregel_ensure_latest": True, "__pregel_dedupe_tasks": True, "__pregel_resuming": True, + '__pregel_store': None, "__pregel_task_id": history[1].tasks[0].id, "checkpoint_id": c.config["configurable"]["checkpoint_id"], "checkpoint_map": { diff --git a/libs/sdk-py/langgraph_sdk/client.py b/libs/sdk-py/langgraph_sdk/client.py index a466c4dc8..a75305479 100644 --- a/libs/sdk-py/langgraph_sdk/client.py +++ b/libs/sdk-py/langgraph_sdk/client.py @@ -11,6 +11,7 @@ from typing import ( Iterator, List, Optional, + Sequence, Union, overload, ) @@ -29,12 +30,15 @@ from langgraph_sdk.schema import ( Cron, DisconnectMode, GraphSchema, + Item, Json, + ListNamespaceResponse, MultitaskStrategy, OnCompletionBehavior, OnConflictBehavior, Run, RunCreate, + SearchItemsResponse, StreamMode, StreamPart, Subgraphs, @@ -143,6 +147,7 @@ class LangGraphClient: self.threads = ThreadsClient(self.http) self.runs = RunsClient(self.http) self.crons = CronClient(self.http) + self.store = StoreClient(self.http) class HttpClient: @@ -211,9 +216,9 @@ class HttpClient: raise e return await adecode_json(r) - async def delete(self, path: str) -> None: + async def delete(self, path: str, *, json: Optional[Any] = None) -> None: """Make a DELETE request.""" - r = await self.client.delete(path) + r = await self.client.request("DELETE", path, json=json) try: r.raise_for_status() except httpx.HTTPStatusError as e: @@ -1879,6 +1884,205 @@ class CronClient: return await self.http.post("/runs/crons/search", json=payload) +class StoreClient: + def __init__(self, http: HttpClient) -> None: + self.http = http + + async def put_item( + self, namespace: Sequence[str], /, key: str, value: dict[str, Any] + ) -> None: + """Store or update an item. + + Args: + namespace: A list of strings representing the namespace path. + key: The unique identifier for the item within the namespace. + value: A dictionary containing the item's data. + + Returns: + None + + Example Usage: + + await client.store.put_item( + ["documents", "user123"], + key="item456", + value={"title": "My Document", "content": "Hello World"} + ) + """ + for label in namespace: + if "." in label: + raise ValueError( + f"Invalid namespace label '{label}'. Namespace labels cannot contain periods ('.')." + ) + payload = { + "namespace": namespace, + "key": key, + "value": value, + } + await self.http.put("/store/items", json=payload) + + async def get_item(self, namespace: Sequence[str], /, key: str) -> Item: + """Retrieve a single item. + + Args: + key: The unique identifier for the item. + namespace: Optional list of strings representing the namespace path. + + Returns: + Item: The retrieved item. + + Example Usage: + + item = await client.store.get_item( + ["documents", "user123"], + key="item456", + ) + print(item) + + ---------------------------------------------------------------- + + { + 'namespace': ['documents', 'user123'], + 'key': 'item456', + 'value': {'title': 'My Document', 'content': 'Hello World'}, + 'created_at': '2024-07-30T12:00:00Z', + 'updated_at': '2024-07-30T12:00:00Z' + } + """ + for label in namespace: + if "." in label: + raise ValueError( + f"Invalid namespace label '{label}'. Namespace labels cannot contain periods ('.')." + ) + return await self.http.get( + "/store/items", params={"namespace": ".".join(namespace), "key": key} + ) + + async def delete_item(self, namespace: Sequence[str], /, key: str) -> None: + """Delete an item. + + Args: + key: The unique identifier for the item. + namespace: Optional list of strings representing the namespace path. + + Returns: + None + + Example Usage: + + await client.store.delete_item( + ["documents", "user123"], + key="item456", + ) + """ + await self.http.delete( + "/store/items", json={"namespace": namespace, "key": key} + ) + + async def search_items( + self, + namespace_prefix: Sequence[str], + /, + filter: Optional[dict[str, Any]] = None, + limit: int = 10, + offset: int = 0, + ) -> SearchItemsResponse: + """Search for items within a namespace prefix. + + Args: + namespace_prefix: List of strings representing the namespace prefix. + filter: Optional dictionary of key-value pairs to filter results. + limit: Maximum number of items to return (default is 10). + offset: Number of items to skip before returning results (default is 0). + + Returns: + List[Item]: A list of items matching the search criteria. + + Example Usage: + + items = await client.store.search_items( + ["documents"], + filter={"author": "John Doe"}, + limit=5, + offset=0 + ) + print(items) + + ---------------------------------------------------------------- + + { + "items": [ + { + "namespace": ["documents", "user123"], + "key": "item789", + "value": { + "title": "Another Document", + "author": "John Doe" + }, + "created_at": "2024-07-30T12:00:00Z", + "updated_at": "2024-07-30T12:00:00Z" + }, + # ... additional items ... + ] + } + """ + payload = { + "namespace_prefix": namespace_prefix, + "filter": filter, + "limit": limit, + "offset": offset, + } + + return await self.http.post("/store/items/search", json=_provided_vals(payload)) + + async def list_namespaces( + self, + prefix: Optional[List[str]] = None, + suffix: Optional[List[str]] = None, + max_depth: Optional[int] = None, + limit: int = 100, + offset: int = 0, + ) -> ListNamespaceResponse: + """List namespaces with optional match conditions. + + Args: + prefix: Optional list of strings representing the prefix to filter namespaces. + suffix: Optional list of strings representing the suffix to filter namespaces. + max_depth: Optional integer specifying the maximum depth of namespaces to return. + limit: Maximum number of namespaces to return (default is 100). + offset: Number of namespaces to skip before returning results (default is 0). + + Returns: + List[List[str]]: A list of namespaces matching the criteria. + + Example Usage: + + namespaces = await client.store.list_namespaces( + prefix=["documents"], + max_depth=3, + limit=10, + offset=0 + ) + print(namespaces) + + ---------------------------------------------------------------- + + [ + ["documents", "user123", "reports"], + ["documents", "user456", "invoices"], + ... + ] + """ + payload = { + "prefix": prefix, + "suffix": suffix, + "max_depth": max_depth, + "limit": limit, + "offset": offset, + } + return await self.http.post("/store/namespaces", json=_provided_vals(payload)) + + def get_sync_client( *, url: Optional[str] = None, @@ -1918,6 +2122,7 @@ class SyncLangGraphClient: self.threads = SyncThreadsClient(self.http) self.runs = SyncRunsClient(self.http) self.crons = SyncCronClient(self.http) + self.store = SyncStoreClient(self.http) class SyncHttpClient: @@ -1986,9 +2191,9 @@ class SyncHttpClient: raise e return decode_json(r) - def delete(self, path: str) -> None: + def delete(self, path: str, *, json: Optional[Any] = None) -> None: """Make a DELETE request.""" - r = self.client.delete(path) + r = self.client.request("DELETE", path, json=json) try: r.raise_for_status() except httpx.HTTPStatusError as e: @@ -3636,3 +3841,204 @@ class SyncCronClient: } payload = {k: v for k, v in payload.items() if v is not None} return self.http.post("/runs/crons/search", json=payload) + + +class SyncStoreClient: + def __init__(self, http: SyncHttpClient) -> None: + self.http = http + + def put_item( + self, namespace: Sequence[str], /, key: str, value: dict[str, Any] + ) -> None: + """Store or update an item. + + Args: + namespace: A list of strings representing the namespace path. + key: The unique identifier for the item within the namespace. + value: A dictionary containing the item's data. + + Returns: + None + + Example Usage: + + client.store.put_item( + ["documents", "user123"], + key="item456", + value={"title": "My Document", "content": "Hello World"} + ) + """ + for label in namespace: + if "." in label: + raise ValueError( + f"Invalid namespace label '{label}'. Namespace labels cannot contain periods ('.')." + ) + payload = { + "namespace": namespace, + "key": key, + "value": value, + } + self.http.put("/store/items", json=payload) + + def get_item(self, namespace: Sequence[str], /, key: str) -> Item: + """Retrieve a single item. + + Args: + key: The unique identifier for the item. + namespace: Optional list of strings representing the namespace path. + + Returns: + Item: The retrieved item. + + Example Usage: + + item = client.store.get_item( + ["documents", "user123"], + key="item456", + ) + print(item) + + ---------------------------------------------------------------- + + { + 'namespace': ['documents', 'user123'], + 'key': 'item456', + 'value': {'title': 'My Document', 'content': 'Hello World'}, + 'created_at': '2024-07-30T12:00:00Z', + 'updated_at': '2024-07-30T12:00:00Z' + } + """ + for label in namespace: + if "." in label: + raise ValueError( + f"Invalid namespace label '{label}'. Namespace labels cannot contain periods ('.')." + ) + + return self.http.get( + "/store/items", params={"key": key, "namespace": ".".join(namespace)} + ) + + def delete_item(self, namespace: Sequence[str], /, key: str) -> None: + """Delete an item. + + Args: + key: The unique identifier for the item. + namespace: Optional list of strings representing the namespace path. + + Returns: + None + + Example Usage: + + client.store.delete_item( + ["documents", "user123"], + key="item456", + ) + """ + self.http.delete("/store/items", json={"key": key, "namespace": namespace}) + + def search_items( + self, + namespace_prefix: Sequence[str], + /, + filter: Optional[dict[str, Any]] = None, + limit: int = 10, + offset: int = 0, + ) -> SearchItemsResponse: + """Search for items within a namespace prefix. + + Args: + namespace_prefix: List of strings representing the namespace prefix. + filter: Optional dictionary of key-value pairs to filter results. + limit: Maximum number of items to return (default is 10). + offset: Number of items to skip before returning results (default is 0). + + Returns: + List[Item]: A list of items matching the search criteria. + + Example Usage: + + items = client.store.search_items( + ["documents"], + filter={"author": "John Doe"}, + limit=5, + offset=0 + ) + print(items) + + ---------------------------------------------------------------- + + { + "items": [ + { + "namespace": ["documents", "user123"], + "key": "item789", + "value": { + "title": "Another Document", + "author": "John Doe" + }, + "created_at": "2024-07-30T12:00:00Z", + "updated_at": "2024-07-30T12:00:00Z" + }, + # ... additional items ... + ] + } + """ + payload = { + "namespace_prefix": namespace_prefix, + "filter": filter, + "limit": limit, + "offset": offset, + } + return self.http.post("/store/items/search", json=_provided_vals(payload)) + + def list_namespaces( + self, + prefix: Optional[List[str]] = None, + suffix: Optional[List[str]] = None, + max_depth: Optional[int] = None, + limit: int = 100, + offset: int = 0, + ) -> ListNamespaceResponse: + """List namespaces with optional match conditions. + + Args: + prefix: Optional list of strings representing the prefix to filter namespaces. + suffix: Optional list of strings representing the suffix to filter namespaces. + max_depth: Optional integer specifying the maximum depth of namespaces to return. + limit: Maximum number of namespaces to return (default is 100). + offset: Number of namespaces to skip before returning results (default is 0). + + Returns: + List[List[str]]: A list of namespaces matching the criteria. + + Example Usage: + + namespaces = client.store.list_namespaces( + prefix=["documents"], + max_depth=3, + limit=10, + offset=0 + ) + print(namespaces) + + ---------------------------------------------------------------- + + [ + ["documents", "user123", "reports"], + ["documents", "user456", "invoices"], + ... + ] + """ + payload = { + "prefix": prefix, + "suffix": suffix, + "max_depth": max_depth, + "limit": limit, + "offset": offset, + } + return self.http.post("/store/namespaces", json=_provided_vals(payload)) + + +def _provided_vals(d: dict): + return {k: v for k, v in d.items() if v is not None} diff --git a/libs/sdk-py/langgraph_sdk/schema.py b/libs/sdk-py/langgraph_sdk/schema.py index da1db204c..2ea166aef 100644 --- a/libs/sdk-py/langgraph_sdk/schema.py +++ b/libs/sdk-py/langgraph_sdk/schema.py @@ -197,6 +197,30 @@ class RunCreate(TypedDict): multitask_strategy: Optional[MultitaskStrategy] +class Item(TypedDict): + namespace: list[str] + """The namespace of the item.""" + key: str + """The unique identifier of the item within its namespace. + + In general, keys are not globally unique. + """ + value: dict[str, Any] + """The value stored in the item. This is the document itself.""" + created_at: datetime + """The timestamp when the item was created.""" + updated_at: datetime + """The timestamp when the item was last updated.""" + + +class ListNamespaceResponse(TypedDict): + namespaces: list[list[str]] + + +class SearchItemsResponse(TypedDict): + items: list[Item] + + class StreamPart(NamedTuple): event: str data: dict