Implement new Store interface (#1834)

This commit is contained in:
Nuno Campos
2024-09-29 17:06:42 -07:00
committed by GitHub
parent ab47768bfd
commit b1a4243055
25 changed files with 1594 additions and 192 deletions
@@ -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:
+410
View File
@@ -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]
+88
View File
@@ -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
+119
View File
@@ -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}")
+8
View File
@@ -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()
+261
View File
@@ -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))
+11 -2
View File
@@ -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,
+14
View File
@@ -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)
+2 -2
View File
@@ -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
-21
View File
@@ -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
-65
View File
@@ -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)
-25
View File
@@ -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)
+24 -4
View File
@@ -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)
+10 -1
View File
@@ -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,
)
== {}
)
+61 -2
View File
@@ -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
+71 -14
View File
@@ -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
+43
View File
@@ -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)
-38
View File
@@ -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": {
+410 -4
View File
@@ -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}
+24
View File
@@ -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