mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-13 13:17:52 +02:00
test
This commit is contained in:
@@ -19,7 +19,11 @@ from langgraph.store.base import (
|
||||
Result,
|
||||
SearchOp,
|
||||
)
|
||||
from langgraph.store.base.batch import AsyncBatchedBaseStore
|
||||
from langgraph.store.base.batch import (
|
||||
BatchedBaseStore,
|
||||
AsyncBatchedBaseStore,
|
||||
SyncBatchedBaseStore,
|
||||
)
|
||||
from langgraph.store.postgres.base import (
|
||||
_PLACEHOLDER,
|
||||
BasePostgresStore,
|
||||
@@ -156,8 +160,15 @@ class AsyncPostgresStore(AsyncBatchedBaseStore, BasePostgresStore[_ainternal.Con
|
||||
|
||||
return results
|
||||
|
||||
# def batch(self, ops: Iterable[Op]) -> list[Result]:
|
||||
# return asyncio.run_coroutine_threadsafe(self.abatch(ops), self.loop).result()
|
||||
def batch(self, ops: Iterable[Op]) -> list[Result]:
|
||||
return asyncio.run_coroutine_threadsafe(self.abatch(ops), self.loop).result()
|
||||
futures = []
|
||||
for op in ops:
|
||||
fut = self._loop.create_future()
|
||||
self._aqueue[fut] = op
|
||||
futures.append(fut)
|
||||
return [fut.result() for fut in asyncio.as_completed(futures)]
|
||||
|
||||
@classmethod
|
||||
@asynccontextmanager
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# type: ignore
|
||||
import asyncio
|
||||
import itertools
|
||||
import sys
|
||||
import uuid
|
||||
@@ -63,6 +64,106 @@ async def store(request) -> AsyncIterator[AsyncPostgresStore]:
|
||||
await conn.execute(f"DROP DATABASE {database}")
|
||||
|
||||
|
||||
async def test_large_batches(store: AsyncPostgresStore) -> None:
|
||||
N = 100
|
||||
M = 10
|
||||
coros = []
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
with ThreadPoolExecutor(max_workers=10) as executor:
|
||||
for m in range(M):
|
||||
ops = []
|
||||
for i in range(N):
|
||||
for i in range(N):
|
||||
coros.append(
|
||||
executor.submit(
|
||||
store.put,
|
||||
("test", "foo", "bar", "baz", str(m % 2)),
|
||||
f"key{i}",
|
||||
value={"foo": "bar" + str(i)},
|
||||
)
|
||||
)
|
||||
coros.append(
|
||||
executor.submit(
|
||||
store.get,
|
||||
("test", "foo", "bar", "baz", str(m % 2)),
|
||||
f"key{i}",
|
||||
)
|
||||
)
|
||||
coros.append(
|
||||
executor.submit(
|
||||
store.list_namespaces,
|
||||
prefix=None,
|
||||
max_depth=m + 1,
|
||||
)
|
||||
)
|
||||
coros.append(
|
||||
executor.submit(
|
||||
store.search,
|
||||
("test",),
|
||||
)
|
||||
)
|
||||
coros.append(
|
||||
executor.submit(
|
||||
store.put,
|
||||
("test", "foo", "bar", "baz", str(m % 2)),
|
||||
f"key{i}",
|
||||
value={"foo": "bar" + str(i)},
|
||||
)
|
||||
)
|
||||
coros.append(
|
||||
executor.submit(
|
||||
store.put,
|
||||
("test", "foo", "bar", "baz", str(m % 2)),
|
||||
f"key{i}",
|
||||
None,
|
||||
)
|
||||
)
|
||||
# ops.extend(
|
||||
# [
|
||||
# PutOp(
|
||||
# namespace=("test", "foo", "bar", "baz", str(m % 2)),
|
||||
# key=f"key{i}", # {m}",
|
||||
# value=None,
|
||||
# ),
|
||||
# GetOp(namespace=("test",), key=f"key{i}{m}"),
|
||||
# ListNamespacesOp(
|
||||
# match_conditions=None, max_depth=i + 1, limit=m, offset=0
|
||||
# ),
|
||||
# SearchOp(
|
||||
# namespace_prefix=("test",),
|
||||
# filter=None,
|
||||
# limit=10,
|
||||
# offset=0,
|
||||
# ),
|
||||
# ]
|
||||
#
|
||||
# ops.extend(
|
||||
# [
|
||||
# # PutOp(
|
||||
# # namespace=("test", "foo", "bar", "baz", str(m % 2)),
|
||||
# # key=f"key{i}", # {m}",
|
||||
# # value={"data": f"value{i}{m}"},
|
||||
# # ),
|
||||
# # GetOp(namespace=("test",), key=f"key{i}{m}"),
|
||||
# # ListNamespacesOp(
|
||||
# # match_conditions=None, max_depth=i + 1, limit=m, offset=0
|
||||
# # ),
|
||||
# # SearchOp(
|
||||
# # namespace_prefix=("test",),
|
||||
# # filter={"data": f"value{i}{m}"},
|
||||
# # limit=10,
|
||||
# # offset=0,
|
||||
# # ),
|
||||
# ]
|
||||
# )
|
||||
|
||||
# coros.extend(ops)
|
||||
# executor.map(store.batch, [[op] for op in coros])
|
||||
|
||||
# await asyncio.gather(*coros) # *[store.abatch(ops) for ops in coros])
|
||||
|
||||
|
||||
async def test_abatch_order(store: AsyncPostgresStore) -> None:
|
||||
# Setup test data
|
||||
await store.aput(("test", "foo"), "key1", {"data": "value1"})
|
||||
|
||||
@@ -808,6 +808,8 @@ class BaseStore(ABC):
|
||||
# [("a", "b", "c"), ("a", "b", "d"), ("a", "b", "f")]
|
||||
```
|
||||
"""
|
||||
if max_depth is not None and max_depth <= 0:
|
||||
raise ValueError("If provided, max_depth must be greater than 0")
|
||||
match_conditions = []
|
||||
if prefix:
|
||||
match_conditions.append(MatchCondition(match_type="prefix", path=prefix))
|
||||
@@ -1004,6 +1006,8 @@ class BaseStore(ABC):
|
||||
# Returns: [("a", "b", "c"), ("a", "b", "d"), ("a", "b", "f")]
|
||||
```
|
||||
"""
|
||||
if max_depth is not None and max_depth <= 0:
|
||||
raise ValueError("If provided, max_depth must be greater than 0")
|
||||
match_conditions = []
|
||||
if prefix:
|
||||
match_conditions.append(MatchCondition(match_type="prefix", path=prefix))
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
import asyncio
|
||||
import threading
|
||||
import time
|
||||
import weakref
|
||||
from typing import Any, Literal, Optional, Union
|
||||
from concurrent.futures import Future
|
||||
from typing import Any, Iterable, Literal, Optional, Union
|
||||
|
||||
from langgraph.store.base import (
|
||||
BaseStore,
|
||||
@@ -17,19 +20,9 @@ from langgraph.store.base import (
|
||||
)
|
||||
|
||||
|
||||
class AsyncBatchedBaseStore(BaseStore):
|
||||
class AsyncBatchedBaseStoreMixin:
|
||||
"""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, ...],
|
||||
@@ -99,6 +92,29 @@ class AsyncBatchedBaseStore(BaseStore):
|
||||
self._aqueue[fut] = op
|
||||
return await fut
|
||||
|
||||
# def batch(self, ops: Iterable[Op]) -> list[Result]:
|
||||
# futures = []
|
||||
# for op in ops:
|
||||
# fut = self._loop.create_future()
|
||||
# self._aqueue[fut] = op
|
||||
# futures.append(fut)
|
||||
# return [fut.result() for fut in asyncio.as_completed(futures)]
|
||||
|
||||
|
||||
class AsyncBatchedBaseStore(AsyncBatchedBaseStoreMixin, BaseStore):
|
||||
"""Efficiently batch operations in a background task."""
|
||||
|
||||
__slots__ = ("_loop", "_aqueue", "_task")
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
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()
|
||||
|
||||
|
||||
def _dedupe_ops(values: list[Op]) -> tuple[Optional[list[int]], list[Op]]:
|
||||
"""Dedupe operations while preserving order for results.
|
||||
@@ -174,3 +190,170 @@ async def _run(
|
||||
break
|
||||
# remove strong ref to store
|
||||
del s
|
||||
|
||||
|
||||
class SyncBatchedBaseStoreMixin(BaseStore):
|
||||
"""Efficiently batch operations in a background thread."""
|
||||
|
||||
def get(
|
||||
self,
|
||||
namespace: tuple[str, ...],
|
||||
key: str,
|
||||
) -> Optional[Item]:
|
||||
fut = Future()
|
||||
self._queue[fut] = GetOp(namespace, key)
|
||||
return fut.result()
|
||||
|
||||
def search(
|
||||
self,
|
||||
namespace_prefix: tuple[str, ...],
|
||||
/,
|
||||
*,
|
||||
query: Optional[str] = None,
|
||||
filter: Optional[dict[str, Any]] = None,
|
||||
limit: int = 10,
|
||||
offset: int = 0,
|
||||
) -> list[SearchItem]:
|
||||
fut = Future()
|
||||
self._queue[fut] = SearchOp(namespace_prefix, filter, limit, offset, query)
|
||||
return fut.result()
|
||||
|
||||
def put(
|
||||
self,
|
||||
namespace: tuple[str, ...],
|
||||
key: str,
|
||||
value: dict[str, Any],
|
||||
index: Optional[Union[Literal[False], list[str]]] = None,
|
||||
) -> None:
|
||||
_validate_namespace(namespace)
|
||||
fut = Future()
|
||||
self._queue[fut] = PutOp(namespace, key, value, index)
|
||||
return fut.result()
|
||||
|
||||
def delete(
|
||||
self,
|
||||
namespace: tuple[str, ...],
|
||||
key: str,
|
||||
) -> None:
|
||||
fut = Future()
|
||||
self._queue[fut] = PutOp(namespace, key, None)
|
||||
return fut.result()
|
||||
|
||||
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, ...]]:
|
||||
fut = Future()
|
||||
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,
|
||||
)
|
||||
self._queue[fut] = op
|
||||
return fut.result()
|
||||
|
||||
|
||||
class SyncBatchedBaseStore(SyncBatchedBaseStoreMixin, BaseStore):
|
||||
"""Efficiently batch operations in a background thread."""
|
||||
|
||||
__slots__ = ("_sync_queue", "_sync_thread")
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self._sync_queue: dict[Future, Op] = {}
|
||||
self._sync_thread = threading.Thread(
|
||||
target=_sync_run,
|
||||
args=(self._sync_queue, weakref.ref(self)),
|
||||
daemon=True,
|
||||
)
|
||||
self._sync_thread.start()
|
||||
|
||||
def __del__(self) -> None:
|
||||
# Signal the thread to stop
|
||||
if self._sync_thread.is_alive():
|
||||
empty_future = Future()
|
||||
self._sync_queue[empty_future] = None # type: ignore
|
||||
self._sync_thread.join()
|
||||
|
||||
|
||||
class BatchedBaseStore(
|
||||
AsyncBatchedBaseStoreMixin, SyncBatchedBaseStoreMixin, BaseStore
|
||||
):
|
||||
__slots__ = (
|
||||
"_sync_queue",
|
||||
"_sync_thread",
|
||||
"_task",
|
||||
"_loop",
|
||||
"_aqueue",
|
||||
)
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
# Setup async processing
|
||||
self._loop = asyncio.get_running_loop()
|
||||
self._aqueue: dict[asyncio.Future, Op] = {}
|
||||
self._task = self._loop.create_task(_run(self._aqueue, weakref.ref(self)))
|
||||
|
||||
self._sync_queue: dict[Future, Op] = {}
|
||||
self._sync_thread = threading.Thread(
|
||||
target=_sync_run,
|
||||
args=(self._sync_queue, weakref.ref(self)),
|
||||
daemon=True,
|
||||
)
|
||||
self._sync_thread.start()
|
||||
|
||||
def __del__(self) -> None:
|
||||
# Signal the thread to stop
|
||||
if self._sync_thread.is_alive():
|
||||
empty_future = Future()
|
||||
self._sync_queue[empty_future] = None # type: ignore
|
||||
self._sync_thread.join()
|
||||
|
||||
# Signal the thread to stop
|
||||
if self._task is not None:
|
||||
self._task.cancel()
|
||||
|
||||
|
||||
def _sync_run(queue: dict[Future, Op], store: weakref.ReferenceType[BaseStore]) -> None:
|
||||
while True:
|
||||
time.sleep(0.001) # Yield to other threads
|
||||
if not queue:
|
||||
continue
|
||||
if s := store():
|
||||
# get the operations to run
|
||||
taken = queue.copy()
|
||||
# action each operation
|
||||
try:
|
||||
values = list(taken.values())
|
||||
if None in values: # Exit signal
|
||||
break
|
||||
listen, dedupped = _dedupe_ops(values)
|
||||
results = s.batch(dedupped) # Note: Using sync batch here
|
||||
if listen is not None:
|
||||
results = [results[ix] for ix in listen]
|
||||
|
||||
# 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 queue[fut]
|
||||
else:
|
||||
break
|
||||
# remove strong ref to store
|
||||
del s
|
||||
|
||||
Reference in New Issue
Block a user