Fix busy loop in AsyncBatchedBaseStore (#3445)

- while loop w asyncio.sleep(0) takes up cpu
This commit is contained in:
Nuno Campos
2025-02-14 12:12:28 -08:00
committed by GitHub
+43 -31
View File
@@ -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