Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
batch.py372 linesDownload Raw Back to base
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 
codekingpro/portable-devtools · Team Ai