mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 12:35:08 +02:00
Implement new Store interface (#1834)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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]
|
||||
@@ -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
|
||||
@@ -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}")
|
||||
@@ -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()
|
||||
|
||||
@@ -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 == []
|
||||
@@ -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))
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
== {}
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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"],),
|
||||
]
|
||||
@@ -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(
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user