codekingpro/portable-devtools
114k
1"""Utilities for batching operations in a background task."""2 3from __future__ import annotations4 5import asyncio6import functools7import weakref8from collections.abc import Callable, Iterable9from typing import Any, Literal, TypeVar10 11from langgraph.store.base import (12 NOT_PROVIDED,13 BaseStore,14 GetOp,15 Item,16 ListNamespacesOp,17 MatchCondition,18 NamespacePath,19 NotProvided,20 Op,21 PutOp,22 Result,23 SearchItem,24 SearchOp,25 _ensure_refresh,26 _ensure_ttl,27 _validate_namespace,28)29 30F = TypeVar("F", bound=Callable)31 32 33def _check_loop(func: F) -> F:34 @functools.wraps(func)35 def wrapper(store: AsyncBatchedBaseStore, *args: Any, **kwargs: Any) -> Any:36 method_name: str = func.__name__37 try:38 current_loop = asyncio.get_running_loop()39 if current_loop is store._loop:40 replacement_str = (41 f"Specifically, replace `store.{method_name}(...)` with `await store.a{method_name}(...)"42 if method_name43 else "For example, replace `store.get(...)` with `await store.aget(...)`"44 )45 raise asyncio.InvalidStateError(46 f"Synchronous calls to {store.__class__.__name__} detected in the main event loop. "47 "This can lead to deadlocks or performance issues. "48 "Please use the asynchronous interface for main thread operations. "49 f"{replacement_str} "50 )51 except RuntimeError:52 pass53 return func(store, *args, **kwargs)54 55 return wrapper56 57 58class AsyncBatchedBaseStore(BaseStore):59 """Efficiently batch operations in a background task."""60 61 __slots__ = ("_loop", "_aqueue", "_task")62 63 def __init__(self) -> None:64 super().__init__()65 self._loop = asyncio.get_running_loop()66 self._aqueue: asyncio.Queue[tuple[asyncio.Future, Op]] = asyncio.Queue()67 self._task: asyncio.Task | None = None68 self._ensure_task()69 70 def __del__(self) -> None:71 try:72 if self._task:73 self._task.cancel()74 except RuntimeError:75 pass76 77 def _ensure_task(self) -> None:78 """Ensure the background processing loop is running."""79 if self._task is None or self._task.done():80 self._task = self._loop.create_task(_run(self._aqueue, weakref.ref(self)))81 82 async def aget(83 self,84 namespace: tuple[str, ...],85 key: str,86 *,87 refresh_ttl: bool | None = None,88 ) -> Item | None:89 self._ensure_task()90 fut = self._loop.create_future()91 self._aqueue.put_nowait(92 (93 fut,94 GetOp(95 namespace,96 key,97 refresh_ttl=_ensure_refresh(self.ttl_config, refresh_ttl),98 ),99 )100 )101 return await fut102 103 async def asearch(104 self,105 namespace_prefix: tuple[str, ...],106 /,107 *,108 query: str | None = None,109 filter: dict[str, Any] | None = None,110 limit: int = 10,111 offset: int = 0,112 refresh_ttl: bool | None = None,113 ) -> list[SearchItem]:114 self._ensure_task()115 fut = self._loop.create_future()116 self._aqueue.put_nowait(117 (118 fut,119 SearchOp(120 namespace_prefix,121 filter,122 limit,123 offset,124 query,125 refresh_ttl=_ensure_refresh(self.ttl_config, refresh_ttl),126 ),127 )128 )129 return await fut130 131 async def aput(132 self,133 namespace: tuple[str, ...],134 key: str,135 value: dict[str, Any],136 index: Literal[False] | list[str] | None = None,137 *,138 ttl: float | None | NotProvided = NOT_PROVIDED,139 ) -> None:140 self._ensure_task()141 _validate_namespace(namespace)142 fut = self._loop.create_future()143 self._aqueue.put_nowait(144 (145 fut,146 PutOp(147 namespace, key, value, index, ttl=_ensure_ttl(self.ttl_config, ttl)148 ),149 )150 )151 return await fut152 153 async def adelete(154 self,155 namespace: tuple[str, ...],156 key: str,157 ) -> None:158 self._ensure_task()159 fut = self._loop.create_future()160 self._aqueue.put_nowait((fut, PutOp(namespace, key, None)))161 return await fut162 163 async def alist_namespaces(164 self,165 *,166 prefix: NamespacePath | None = None,167 suffix: NamespacePath | None = None,168 max_depth: int | None = None,169 limit: int = 100,170 offset: int = 0,171 ) -> list[tuple[str, ...]]:172 self._ensure_task()173 fut = self._loop.create_future()174 match_conditions = []175 if prefix:176 match_conditions.append(MatchCondition(match_type="prefix", path=prefix))177 if suffix:178 match_conditions.append(MatchCondition(match_type="suffix", path=suffix))179 180 op = ListNamespacesOp(181 match_conditions=tuple(match_conditions),182 max_depth=max_depth,183 limit=limit,184 offset=offset,185 )186 self._aqueue.put_nowait((fut, op))187 return await fut188 189 @_check_loop190 def batch(self, ops: Iterable[Op]) -> list[Result]:191 return asyncio.run_coroutine_threadsafe(self.abatch(ops), self._loop).result()192 193 @_check_loop194 def get(195 self,196 namespace: tuple[str, ...],197 key: str,198 *,199 refresh_ttl: bool | None = None,200 ) -> Item | None:201 return asyncio.run_coroutine_threadsafe(202 self.aget(namespace, key=key, refresh_ttl=refresh_ttl), self._loop203 ).result()204 205 @_check_loop206 def search(207 self,208 namespace_prefix: tuple[str, ...],209 /,210 *,211 query: str | None = None,212 filter: dict[str, Any] | None = None,213 limit: int = 10,214 offset: int = 0,215 refresh_ttl: bool | None = None,216 ) -> list[SearchItem]:217 return asyncio.run_coroutine_threadsafe(218 self.asearch(219 namespace_prefix,220 query=query,221 filter=filter,222 limit=limit,223 offset=offset,224 refresh_ttl=refresh_ttl,225 ),226 self._loop,227 ).result()228 229 @_check_loop230 def put(231 self,232 namespace: tuple[str, ...],233 key: str,234 value: dict[str, Any],235 index: Literal[False] | list[str] | None = None,236 *,237 ttl: float | None | NotProvided = NOT_PROVIDED,238 ) -> None:239 _validate_namespace(namespace)240 asyncio.run_coroutine_threadsafe(241 self.aput(242 namespace,243 key=key,244 value=value,245 index=index,246 ttl=_ensure_ttl(self.ttl_config, ttl),247 ),248 self._loop,249 ).result()250 251 @_check_loop252 def delete(253 self,254 namespace: tuple[str, ...],255 key: str,256 ) -> None:257 asyncio.run_coroutine_threadsafe(258 self.adelete(namespace, key=key), self._loop259 ).result()260 261 @_check_loop262 def list_namespaces(263 self,264 *,265 prefix: NamespacePath | None = None,266 suffix: NamespacePath | None = None,267 max_depth: int | None = None,268 limit: int = 100,269 offset: int = 0,270 ) -> list[tuple[str, ...]]:271 return asyncio.run_coroutine_threadsafe(272 self.alist_namespaces(273 prefix=prefix,274 suffix=suffix,275 max_depth=max_depth,276 limit=limit,277 offset=offset,278 ),279 self._loop,280 ).result()281 282 283def _dedupe_ops(values: list[Op]) -> tuple[list[int] | None, list[Op]]:284 """Dedupe operations while preserving order for results.285 286 Args:287 values: List of operations to dedupe288 289 Returns:290 Tuple of (listen indices, deduped operations)291 where listen indices map deduped operation results back to original positions292 """293 if len(values) <= 1:294 return None, list(values)295 296 dedupped: list[Op] = []297 listen: list[int] = []298 puts: dict[tuple[tuple[str, ...], str], int] = {}299 300 for op in values:301 if isinstance(op, (GetOp, SearchOp, ListNamespacesOp)):302 try:303 listen.append(dedupped.index(op))304 except ValueError:305 listen.append(len(dedupped))306 dedupped.append(op)307 elif isinstance(op, PutOp):308 putkey = (op.namespace, op.key)309 if putkey in puts:310 # Overwrite previous put311 ix = puts[putkey]312 dedupped[ix] = op313 listen.append(ix)314 else:315 puts[putkey] = len(dedupped)316 listen.append(len(dedupped))317 dedupped.append(op)318 319 else: # Any new ops will be treated regularly320 listen.append(len(dedupped))321 dedupped.append(op)322 323 return listen, dedupped324 325 326async def _run(327 aqueue: asyncio.Queue[tuple[asyncio.Future, Op]],328 store: weakref.ReferenceType[BaseStore],329) -> None:330 while item := await aqueue.get():331 # don't run batch if the future is done (e.g. cancelled)332 if item[0].done():333 continue334 # check if store is still alive335 if s := store():336 try:337 # accumulate operations scheduled in same tick338 items = [item]339 try:340 while item := aqueue.get_nowait():341 # don't insert if the future is done (e.g. cancelled)342 if item[0].done():343 continue344 items.append(item)345 except asyncio.QueueEmpty:346 pass347 # get the operations to run348 futs = [item[0] for item in items]349 values = [item[1] for item in items]350 # action each operation351 try:352 listen, dedupped = _dedupe_ops(values)353 results = await s.abatch(dedupped)354 if listen is not None:355 results = [results[ix] for ix in listen]356 357 # set the results of each operation358 for fut, result in zip(futs, results, strict=False):359 # guard against future being done (e.g. cancelled)360 if not fut.done():361 fut.set_result(result)362 except Exception as e:363 for fut in futs:364 # guard against future being done (e.g. cancelled)365 if not fut.done():366 fut.set_exception(e)367 finally:368 # remove strong ref to store369 del s370 else:371 break372 