This commit is contained in:
William Fu-Hinthorn
2024-12-05 18:27:47 -08:00
parent b7e441d781
commit 4406828c39
4 changed files with 313 additions and 14 deletions
@@ -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))
+195 -12
View File
@@ -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