Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_executor.py224 linesDownload Raw Back to pregel
1from __future__ import annotations2 3import asyncio4import concurrent.futures5import time6from collections.abc import Awaitable, Callable, Coroutine7from contextlib import AbstractAsyncContextManager, AbstractContextManager, ExitStack8from contextvars import copy_context9from types import TracebackType10from typing import (11    Protocol,12    TypeVar,13    cast,14)15 16from langchain_core.runnables import RunnableConfig17from langchain_core.runnables.config import get_executor_for_config18from typing_extensions import ParamSpec19 20from langgraph._internal._future import CONTEXT_NOT_SUPPORTED, run_coroutine_threadsafe21from langgraph.errors import GraphBubbleUp22 23P = ParamSpec("P")24T = TypeVar("T")25 26 27class Submit(Protocol[P, T]):28    def __call__(  # type: ignore[valid-type]29        self,30        fn: Callable[P, T],31        *args: P.args,32        __name__: str | None = None,33        __cancel_on_exit__: bool = False,34        __reraise_on_exit__: bool = True,35        __next_tick__: bool = False,36        **kwargs: P.kwargs,37    ) -> concurrent.futures.Future[T]: ...38 39 40class BackgroundExecutor(AbstractContextManager):41    """A context manager that runs sync tasks in the background.42    Uses a thread pool executor to delegate tasks to separate threads.43    On exit,44    - cancels any (not yet started) tasks with `__cancel_on_exit__=True`45    - waits for all tasks to finish46    - re-raises the first exception from tasks with `__reraise_on_exit__=True`"""47 48    def __init__(self, config: RunnableConfig) -> None:49        self.stack = ExitStack()50        self.executor = self.stack.enter_context(get_executor_for_config(config))51        # mapping of Future to (__cancel_on_exit__, __reraise_on_exit__) flags52        self.tasks: dict[concurrent.futures.Future, tuple[bool, bool]] = {}53 54    def submit(  # type: ignore[valid-type]55        self,56        fn: Callable[P, T],57        *args: P.args,58        __name__: str | None = None,  # currently not used in sync version59        __cancel_on_exit__: bool = False,  # for sync, can cancel only if not started60        __reraise_on_exit__: bool = True,61        __next_tick__: bool = False,62        **kwargs: P.kwargs,63    ) -> concurrent.futures.Future[T]:64        ctx = copy_context()65        if __next_tick__:66            task = cast(67                concurrent.futures.Future[T],68                self.executor.submit(next_tick, ctx.run, fn, *args, **kwargs),  # type: ignore[arg-type]69            )70        else:71            task = self.executor.submit(ctx.run, fn, *args, **kwargs)72        self.tasks[task] = (__cancel_on_exit__, __reraise_on_exit__)73        # add a callback to remove the task from the tasks dict when it's done74        task.add_done_callback(self.done)75        return task76 77    def done(self, task: concurrent.futures.Future) -> None:78        """Remove the task from the tasks dict when it's done."""79        try:80            task.result()81        except GraphBubbleUp:82            # This exception is an interruption signal, not an error83            # so we don't want to re-raise it on exit84            self.tasks.pop(task)85        except BaseException:86            pass87        else:88            self.tasks.pop(task)89 90    def __enter__(self) -> Submit:91        return self.submit92 93    def __exit__(94        self,95        exc_type: type[BaseException] | None,96        exc_value: BaseException | None,97        traceback: TracebackType | None,98    ) -> bool | None:99        # copy the tasks as done() callback may modify the dict100        tasks = self.tasks.copy()101        # cancel all tasks that should be cancelled102        for task, (cancel, _) in tasks.items():103            if cancel:104                task.cancel()105        # wait for all tasks to finish106        if pending := {t for t in tasks if not t.done()}:107            concurrent.futures.wait(pending)108        # shutdown the executor109        self.stack.__exit__(exc_type, exc_value, traceback)110        # if there's already an exception being raised, don't raise another one111        if exc_type is None:112            # re-raise the first exception that occurred in a task113            for task, (_, reraise) in tasks.items():114                if not reraise:115                    continue116                try:117                    task.result()118                except concurrent.futures.CancelledError:119                    pass120 121 122class AsyncBackgroundExecutor(AbstractAsyncContextManager):123    """A context manager that runs async tasks in the background.124    Uses the current event loop to delegate tasks to asyncio tasks.125    On exit,126    - cancels any tasks with `__cancel_on_exit__=True`127    - waits for all tasks to finish128    - re-raises the first exception from tasks with `__reraise_on_exit__=True`129      ignoring CancelledError"""130 131    def __init__(self, config: RunnableConfig) -> None:132        self.tasks: dict[asyncio.Future, tuple[bool, bool]] = {}133        self.sentinel = object()134        self.loop = asyncio.get_running_loop()135        if max_concurrency := config.get("max_concurrency"):136            self.semaphore: asyncio.Semaphore | None = asyncio.Semaphore(137                max_concurrency138            )139        else:140            self.semaphore = None141 142    def submit(  # type: ignore[valid-type]143        self,144        fn: Callable[P, Awaitable[T]],145        *args: P.args,146        __name__: str | None = None,147        __cancel_on_exit__: bool = False,148        __reraise_on_exit__: bool = True,149        __next_tick__: bool = False,  # noop in async (always True)150        **kwargs: P.kwargs,151    ) -> asyncio.Future[T]:152        coro = cast(Coroutine[None, None, T], fn(*args, **kwargs))153        if self.semaphore:154            coro = gated(self.semaphore, coro)155        if CONTEXT_NOT_SUPPORTED:156            task = run_coroutine_threadsafe(157                coro, self.loop, name=__name__, lazy=__next_tick__158            )159        else:160            task = run_coroutine_threadsafe(161                coro,162                self.loop,163                name=__name__,164                context=copy_context(),165                lazy=__next_tick__,166            )167        self.tasks[task] = (__cancel_on_exit__, __reraise_on_exit__)168        task.add_done_callback(self.done)169        return task170 171    def done(self, task: asyncio.Future) -> None:172        try:173            if exc := task.exception():174                # This exception is an interruption signal, not an error175                # so we don't want to re-raise it on exit176                if isinstance(exc, GraphBubbleUp):177                    self.tasks.pop(task)178            else:179                self.tasks.pop(task)180        except asyncio.CancelledError:181            self.tasks.pop(task)182 183    async def __aenter__(self) -> Submit:184        return self.submit185 186    async def __aexit__(187        self,188        exc_type: type[BaseException] | None,189        exc_value: BaseException | None,190        traceback: TracebackType | None,191    ) -> None:192        # copy the tasks as done() callback may modify the dict193        tasks = self.tasks.copy()194        # cancel all tasks that should be cancelled195        for task, (cancel, _) in tasks.items():196            if cancel:197                task.cancel(self.sentinel)198        # wait for all tasks to finish199        if tasks:200            await asyncio.wait(tasks)201        # if there's already an exception being raised, don't raise another one202        if exc_type is None:203            # re-raise the first exception that occurred in a task204            for task, (_, reraise) in tasks.items():205                if not reraise:206                    continue207                try:208                    if exc := task.exception():209                        raise exc210                except asyncio.CancelledError:211                    pass212 213 214async def gated(semaphore: asyncio.Semaphore, coro: Coroutine[None, None, T]) -> T:215    """A coroutine that waits for a semaphore before running another coroutine."""216    async with semaphore:217        return await coro218 219 220def next_tick(fn: Callable[P, T], *args: P.args, **kwargs: P.kwargs) -> T:221    """A function that yields control to other threads before running another function."""222    time.sleep(0)223    return fn(*args, **kwargs)224 
codekingpro/portable-devtools · Team Ai