codekingpro/portable-devtools
114k
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 