diff --git a/libs/checkpoint/langgraph/store/base/batch.py b/libs/checkpoint/langgraph/store/base/batch.py index 6cfc11419..1bf2ac79e 100644 --- a/libs/checkpoint/langgraph/store/base/batch.py +++ b/libs/checkpoint/langgraph/store/base/batch.py @@ -1,7 +1,8 @@ import asyncio import functools import weakref -from typing import Any, Callable, Iterable, Literal, Optional, TypeVar, Union +from collections.abc import Iterable +from typing import Any, Callable, Literal, Optional, TypeVar, Union from langgraph.store.base import ( BaseStore, @@ -54,7 +55,7 @@ class AsyncBatchedBaseStore(BaseStore): def __init__(self) -> None: super().__init__() self._loop = asyncio.get_running_loop() - self._aqueue: dict[asyncio.Future, Op] = {} + self._aqueue: asyncio.Queue[tuple[asyncio.Future, Op]] = asyncio.Queue() self._task = self._loop.create_task(_run(self._aqueue, weakref.ref(self))) def __del__(self) -> None: @@ -65,8 +66,9 @@ class AsyncBatchedBaseStore(BaseStore): namespace: tuple[str, ...], key: str, ) -> Optional[Item]: + assert not self._task.done() fut = self._loop.create_future() - self._aqueue[fut] = GetOp(namespace, key) + self._aqueue.put_nowait((fut, GetOp(namespace, key))) return await fut async def asearch( @@ -79,8 +81,11 @@ class AsyncBatchedBaseStore(BaseStore): limit: int = 10, offset: int = 0, ) -> list[SearchItem]: + assert not self._task.done() fut = self._loop.create_future() - self._aqueue[fut] = SearchOp(namespace_prefix, filter, limit, offset, query) + self._aqueue.put_nowait( + (fut, SearchOp(namespace_prefix, filter, limit, offset, query)) + ) return await fut async def aput( @@ -90,9 +95,10 @@ class AsyncBatchedBaseStore(BaseStore): value: dict[str, Any], index: Optional[Union[Literal[False], list[str]]] = None, ) -> None: + assert not self._task.done() _validate_namespace(namespace) fut = self._loop.create_future() - self._aqueue[fut] = PutOp(namespace, key, value, index) + self._aqueue.put_nowait((fut, PutOp(namespace, key, value, index))) return await fut async def adelete( @@ -100,8 +106,9 @@ class AsyncBatchedBaseStore(BaseStore): namespace: tuple[str, ...], key: str, ) -> None: + assert not self._task.done() fut = self._loop.create_future() - self._aqueue[fut] = PutOp(namespace, key, None) + self._aqueue.put_nowait((fut, PutOp(namespace, key, None))) return await fut async def alist_namespaces( @@ -113,6 +120,7 @@ class AsyncBatchedBaseStore(BaseStore): limit: int = 100, offset: int = 0, ) -> list[tuple[str, ...]]: + assert not self._task.done() fut = self._loop.create_future() match_conditions = [] if prefix: @@ -126,7 +134,7 @@ class AsyncBatchedBaseStore(BaseStore): limit=limit, offset=offset, ) - self._aqueue[fut] = op + self._aqueue.put_nowait((fut, op)) return await fut @_check_loop @@ -250,34 +258,38 @@ def _dedupe_ops(values: list[Op]) -> tuple[Optional[list[int]], list[Op]]: async def _run( - aqueue: dict[asyncio.Future, Op], + aqueue: asyncio.Queue[tuple[asyncio.Future, Op]], store: weakref.ReferenceType[BaseStore], ) -> None: - while True: - await asyncio.sleep(0) - if not aqueue: - continue + while item := await aqueue.get(): + # check if store is still alive if s := store(): - # get the operations to run - taken = aqueue.copy() - # action each operation try: - values = list(taken.values()) - listen, dedupped = _dedupe_ops(values) - results = await s.abatch(dedupped) - if listen is not None: - results = [results[ix] for ix in listen] + # accumulate operations scheduled in same tick + items = [item] + try: + while item := aqueue.get_nowait(): + items.append(item) + except asyncio.QueueEmpty: + pass + # get the operations to run + futs = [item[0] for item in items] + values = [item[1] for item in items] + # action each operation + try: + listen, dedupped = _dedupe_ops(values) + results = await s.abatch(dedupped) + 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 aqueue[fut] + # set the results of each operation + for fut, result in zip(futs, results): + fut.set_result(result) + except Exception as e: + for fut in futs: + fut.set_exception(e) + finally: + # remove strong ref to store + del s else: break - # remove strong ref to store - del s