diff --git a/venv/Lib/site-packages/anyio/__init__.py b/venv/Lib/site-packages/anyio/__init__.py new file mode 100644 index 0000000..2502c76 --- /dev/null +++ b/venv/Lib/site-packages/anyio/__init__.py @@ -0,0 +1,115 @@ +from __future__ import annotations + +from ._core._contextmanagers import AsyncContextManagerMixin as AsyncContextManagerMixin +from ._core._contextmanagers import ContextManagerMixin as ContextManagerMixin +from ._core._eventloop import current_time as current_time +from ._core._eventloop import get_all_backends as get_all_backends +from ._core._eventloop import get_available_backends as get_available_backends +from ._core._eventloop import get_cancelled_exc_class as get_cancelled_exc_class +from ._core._eventloop import run as run +from ._core._eventloop import sleep as sleep +from ._core._eventloop import sleep_forever as sleep_forever +from ._core._eventloop import sleep_until as sleep_until +from ._core._exceptions import BrokenResourceError as BrokenResourceError +from ._core._exceptions import BrokenWorkerInterpreter as BrokenWorkerInterpreter +from ._core._exceptions import BrokenWorkerProcess as BrokenWorkerProcess +from ._core._exceptions import BusyResourceError as BusyResourceError +from ._core._exceptions import ClosedResourceError as ClosedResourceError +from ._core._exceptions import ConnectionFailed as ConnectionFailed +from ._core._exceptions import DelimiterNotFound as DelimiterNotFound +from ._core._exceptions import EndOfStream as EndOfStream +from ._core._exceptions import IncompleteRead as IncompleteRead +from ._core._exceptions import NoEventLoopError as NoEventLoopError +from ._core._exceptions import RunFinishedError as RunFinishedError +from ._core._exceptions import TaskCancelled as TaskCancelled +from ._core._exceptions import TaskFailed as TaskFailed +from ._core._exceptions import TaskNotFinished as TaskNotFinished +from ._core._exceptions import TypedAttributeLookupError as TypedAttributeLookupError +from ._core._exceptions import WouldBlock as WouldBlock +from ._core._fileio import AsyncFile as AsyncFile +from ._core._fileio import Path as Path +from ._core._fileio import open_file as open_file +from ._core._fileio import wrap_file as wrap_file +from ._core._resources import aclose_forcefully as aclose_forcefully +from ._core._signals import open_signal_receiver as open_signal_receiver +from ._core._sockets import TCPConnectable as TCPConnectable +from ._core._sockets import UNIXConnectable as UNIXConnectable +from ._core._sockets import as_connectable as as_connectable +from ._core._sockets import connect_tcp as connect_tcp +from ._core._sockets import connect_unix as connect_unix +from ._core._sockets import create_connected_udp_socket as create_connected_udp_socket +from ._core._sockets import ( + create_connected_unix_datagram_socket as create_connected_unix_datagram_socket, +) +from ._core._sockets import create_tcp_listener as create_tcp_listener +from ._core._sockets import create_udp_socket as create_udp_socket +from ._core._sockets import create_unix_datagram_socket as create_unix_datagram_socket +from ._core._sockets import create_unix_listener as create_unix_listener +from ._core._sockets import getaddrinfo as getaddrinfo +from ._core._sockets import getnameinfo as getnameinfo +from ._core._sockets import notify_closing as notify_closing +from ._core._sockets import wait_readable as wait_readable +from ._core._sockets import wait_socket_readable as wait_socket_readable +from ._core._sockets import wait_socket_writable as wait_socket_writable +from ._core._sockets import wait_writable as wait_writable +from ._core._streams import create_memory_object_stream as create_memory_object_stream +from ._core._subprocesses import open_process as open_process +from ._core._subprocesses import run_process as run_process +from ._core._synchronization import CapacityLimiter as CapacityLimiter +from ._core._synchronization import ( + CapacityLimiterStatistics as CapacityLimiterStatistics, +) +from ._core._synchronization import Condition as Condition +from ._core._synchronization import ConditionStatistics as ConditionStatistics +from ._core._synchronization import Event as Event +from ._core._synchronization import EventStatistics as EventStatistics +from ._core._synchronization import Lock as Lock +from ._core._synchronization import LockStatistics as LockStatistics +from ._core._synchronization import ResourceGuard as ResourceGuard +from ._core._synchronization import Semaphore as Semaphore +from ._core._synchronization import SemaphoreStatistics as SemaphoreStatistics +from ._core._tasks import TASK_STATUS_IGNORED as TASK_STATUS_IGNORED +from ._core._tasks import CancelScope as CancelScope +from ._core._tasks import TaskHandle as TaskHandle +from ._core._tasks import create_task_group as create_task_group +from ._core._tasks import current_effective_deadline as current_effective_deadline +from ._core._tasks import fail_after as fail_after +from ._core._tasks import move_on_after as move_on_after +from ._core._tempfile import NamedTemporaryFile as NamedTemporaryFile +from ._core._tempfile import SpooledTemporaryFile as SpooledTemporaryFile +from ._core._tempfile import TemporaryDirectory as TemporaryDirectory +from ._core._tempfile import TemporaryFile as TemporaryFile +from ._core._tempfile import gettempdir as gettempdir +from ._core._tempfile import gettempdirb as gettempdirb +from ._core._tempfile import mkdtemp as mkdtemp +from ._core._tempfile import mkstemp as mkstemp +from ._core._testing import TaskInfo as TaskInfo +from ._core._testing import get_current_task as get_current_task +from ._core._testing import get_running_tasks as get_running_tasks +from ._core._testing import wait_all_tasks_blocked as wait_all_tasks_blocked +from ._core._typedattr import TypedAttributeProvider as TypedAttributeProvider +from ._core._typedattr import TypedAttributeSet as TypedAttributeSet +from ._core._typedattr import typed_attribute as typed_attribute + +# Re-export imports so they look like they live directly in this package +for __value in list(locals().values()): + if getattr(__value, "__module__", "").startswith("anyio."): + __value.__module__ = __name__ + + +del __value + + +def __getattr__(attr: str) -> type[BrokenWorkerInterpreter]: + """Support deprecated aliases.""" + if attr == "BrokenWorkerIntepreter": + import warnings + + warnings.warn( + "The 'BrokenWorkerIntepreter' alias is deprecated, use 'BrokenWorkerInterpreter' instead.", + DeprecationWarning, + stacklevel=2, + ) + return BrokenWorkerInterpreter + + raise AttributeError(f"module {__name__!r} has no attribute {attr!r}") diff --git a/venv/Lib/site-packages/anyio/from_thread.py b/venv/Lib/site-packages/anyio/from_thread.py new file mode 100644 index 0000000..8c7914c --- /dev/null +++ b/venv/Lib/site-packages/anyio/from_thread.py @@ -0,0 +1,582 @@ +from __future__ import annotations + +__all__ = ( + "BlockingPortal", + "BlockingPortalProvider", + "check_cancelled", + "run", + "run_sync", + "start_blocking_portal", +) + +import sys +from collections.abc import Awaitable, Callable, Coroutine, Generator +from concurrent.futures import Future +from contextlib import ( + AbstractAsyncContextManager, + AbstractContextManager, + contextmanager, +) +from dataclasses import dataclass, field +from functools import partial +from inspect import isawaitable +from threading import Lock, Thread, current_thread, get_ident +from types import TracebackType +from typing import ( + Any, + Generic, + TypeVar, + cast, + overload, +) + +from ._core._eventloop import ( + get_cancelled_exc_class, + threadlocals, +) +from ._core._eventloop import run as run_eventloop +from ._core._exceptions import NoEventLoopError +from ._core._synchronization import Event +from ._core._tasks import CancelScope, create_task_group +from .abc._tasks import TaskStatus +from .lowlevel import EventLoopToken, current_token + +if sys.version_info >= (3, 11): + from typing import TypeVarTuple, Unpack +else: + from typing_extensions import TypeVarTuple, Unpack + +T_Retval = TypeVar("T_Retval") +T_co = TypeVar("T_co", covariant=True) +PosArgsT = TypeVarTuple("PosArgsT") + + +def _token_or_error(token: EventLoopToken | None) -> EventLoopToken: + if token is not None: + return token + + try: + return threadlocals.current_token + except AttributeError: + raise NoEventLoopError( + "Not running inside an AnyIO worker thread, and no event loop token was " + "provided" + ) from None + + +def run( + func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]], + *args: Unpack[PosArgsT], + token: EventLoopToken | None = None, +) -> T_co: + """ + Call a coroutine function from a worker thread. + + :param func: a coroutine function + :param args: positional arguments for the callable + :param token: an event loop token to use to get back to the event loop thread + (required if calling this function from outside an AnyIO worker thread) + :return: the return value of the coroutine function + :raises MissingTokenError: if no token was provided and called from outside an + AnyIO worker thread + :raises RunFinishedError: if the event loop tied to ``token`` is no longer running + + .. versionchanged:: 4.11.0 + Added the ``token`` parameter. + + """ + explicit_token = token is not None + token = _token_or_error(token) + return token.backend_class.run_async_from_thread( + func, args, token=token.native_token if explicit_token else None + ) + + +def run_sync( + func: Callable[[Unpack[PosArgsT]], T_Retval], + *args: Unpack[PosArgsT], + token: EventLoopToken | None = None, +) -> T_Retval: + """ + Call a function in the event loop thread from a worker thread. + + :param func: a callable + :param args: positional arguments for the callable + :param token: an event loop token to use to get back to the event loop thread + (required if calling this function from outside an AnyIO worker thread) + :return: the return value of the callable + :raises MissingTokenError: if no token was provided and called from outside an + AnyIO worker thread + :raises RunFinishedError: if the event loop tied to ``token`` is no longer running + + .. versionchanged:: 4.11.0 + Added the ``token`` parameter. + + """ + explicit_token = token is not None + token = _token_or_error(token) + return token.backend_class.run_sync_from_thread( + func, args, token=token.native_token if explicit_token else None + ) + + +class _BlockingAsyncContextManager(Generic[T_co], AbstractContextManager): + _enter_future: Future[T_co] + _exit_future: Future[bool | None] + _exit_event: Event + _exit_exc_info: tuple[ + type[BaseException] | None, BaseException | None, TracebackType | None + ] = (None, None, None) + + def __init__( + self, async_cm: AbstractAsyncContextManager[T_co], portal: BlockingPortal + ): + self._async_cm = async_cm + self._portal = portal + + async def run_async_cm(self) -> bool | None: + try: + self._exit_event = Event() + value = await self._async_cm.__aenter__() + except BaseException as exc: + self._enter_future.set_exception(exc) + raise + else: + self._enter_future.set_result(value) + + try: + # Wait for the sync context manager to exit. + # This next statement can raise `get_cancelled_exc_class()` if + # something went wrong in a task group in this async context + # manager. + await self._exit_event.wait() + finally: + # In case of cancellation, it could be that we end up here before + # `_BlockingAsyncContextManager.__exit__` is called, and an + # `_exit_exc_info` has been set. + result = await self._async_cm.__aexit__(*self._exit_exc_info) + + return result + + def __enter__(self) -> T_co: + self._enter_future = Future() + self._exit_future = self._portal.start_task_soon(self.run_async_cm) + return self._enter_future.result() + + def __exit__( + self, + __exc_type: type[BaseException] | None, + __exc_value: BaseException | None, + __traceback: TracebackType | None, + ) -> bool | None: + self._exit_exc_info = __exc_type, __exc_value, __traceback + self._portal.call(self._exit_event.set) + return self._exit_future.result() + + +class _BlockingPortalTaskStatus(TaskStatus): + def __init__(self, future: Future): + self._future = future + + def started(self, value: object = None) -> None: + self._future.set_result(value) + + +class BlockingPortal: + """ + An object that lets external threads run code in an asynchronous event loop. + + :raises NoEventLoopError: if no supported asynchronous event loop is running in the + current thread + """ + + def __init__(self) -> None: + self._token = current_token() + self._event_loop_thread_id: int | None = get_ident() + self._stop_event = Event() + self._task_group = create_task_group() + + async def __aenter__(self) -> BlockingPortal: + await self._task_group.__aenter__() + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: TracebackType | None, + ) -> bool: + await self.stop() + return await self._task_group.__aexit__(exc_type, exc_val, exc_tb) + + def _check_running(self) -> None: + if self._event_loop_thread_id is None: + raise RuntimeError("This portal is not running") + if self._event_loop_thread_id == get_ident(): + raise RuntimeError( + "This method cannot be called from the event loop thread" + ) + + async def sleep_until_stopped(self) -> None: + """Sleep until :meth:`stop` is called.""" + await self._stop_event.wait() + + async def stop(self, cancel_remaining: bool = False) -> None: + """ + Signal the portal to shut down. + + This marks the portal as no longer accepting new calls and exits from + :meth:`sleep_until_stopped`. + + :param cancel_remaining: ``True`` to cancel all the remaining tasks, ``False`` + to let them finish before returning + + """ + self._event_loop_thread_id = None + self._stop_event.set() + if cancel_remaining: + self._task_group.cancel_scope.cancel("the blocking portal is shutting down") + + async def _call_func( + self, + func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval] | T_Retval], + args: tuple[Unpack[PosArgsT]], + kwargs: dict[str, Any], + future: Future[T_Retval], + ) -> None: + event_loop_thread_id = self._event_loop_thread_id + + def callback(f: Future[T_Retval]) -> None: + if f.cancelled(): + if event_loop_thread_id == get_ident(): + scope.cancel("the future was cancelled") + elif event_loop_thread_id is not None: + run_sync( + scope.cancel, "the future was cancelled", token=self._token + ) + + try: + retval_or_awaitable = func(*args, **kwargs) + if isawaitable(retval_or_awaitable): + with CancelScope() as scope: + future.add_done_callback(callback) + retval = await retval_or_awaitable + else: + retval = retval_or_awaitable + except get_cancelled_exc_class(): + future.cancel() + future.set_running_or_notify_cancel() + except BaseException as exc: + if not future.cancelled(): + future.set_exception(exc) + + # Let base exceptions fall through + if not isinstance(exc, Exception): + raise + else: + if not future.cancelled(): + future.set_result(retval) + finally: + scope = None # type: ignore[assignment] + + def _spawn_task_from_thread( + self, + func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval] | T_Retval], + args: tuple[Unpack[PosArgsT]], + kwargs: dict[str, Any], + name: object, + future: Future[T_Retval], + ) -> None: + """ + Spawn a new task using the given callable. + + :param func: a callable + :param args: positional arguments to be passed to the callable + :param kwargs: keyword arguments to be passed to the callable + :param name: name of the task (will be coerced to a string if not ``None``) + :param future: a future that will resolve to the return value of the callable, + or the exception raised during its execution + + """ + run_sync( + partial(self._task_group.start_soon, name=name), + self._call_func, + func, + args, + kwargs, + future, + token=self._token, + ) + + @overload + def call( + self, + func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval]], + *args: Unpack[PosArgsT], + ) -> T_Retval: ... + + @overload + def call( + self, func: Callable[[Unpack[PosArgsT]], T_Retval], *args: Unpack[PosArgsT] + ) -> T_Retval: ... + + def call( + self, + func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval] | T_Retval], + *args: Unpack[PosArgsT], + ) -> T_Retval: + """ + Call the given function in the event loop thread. + + If the callable returns a coroutine object, it is awaited on. + + :param func: any callable + :raises RuntimeError: if the portal is not running or if this method is called + from within the event loop thread + + """ + return cast(T_Retval, self.start_task_soon(func, *args).result()) + + @overload + def start_task_soon( + self, + func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval]], + *args: Unpack[PosArgsT], + name: object = None, + ) -> Future[T_Retval]: ... + + @overload + def start_task_soon( + self, + func: Callable[[Unpack[PosArgsT]], T_Retval], + *args: Unpack[PosArgsT], + name: object = None, + ) -> Future[T_Retval]: ... + + def start_task_soon( + self, + func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval] | T_Retval], + *args: Unpack[PosArgsT], + name: object = None, + ) -> Future[T_Retval]: + """ + Start a task in the portal's task group. + + The task will be run inside a cancel scope which can be cancelled by cancelling + the returned future. + + :param func: the target function + :param args: positional arguments passed to ``func`` + :param name: name of the task (will be coerced to a string if not ``None``) + :return: a future that resolves with the return value of the callable if the + task completes successfully, or with the exception raised in the task + :raises RuntimeError: if the portal is not running or if this method is called + from within the event loop thread + :rtype: concurrent.futures.Future[T_Retval] + + .. versionadded:: 3.0 + + """ + self._check_running() + f: Future[T_Retval] = Future() + self._spawn_task_from_thread(func, args, {}, name, f) + return f + + def start_task( + self, + func: Callable[..., Awaitable[T_Retval]], + *args: object, + name: object = None, + ) -> tuple[Future[T_Retval], Any]: + """ + Start a task in the portal's task group and wait until it signals for readiness. + + This method works the same way as :meth:`.abc.TaskGroup.start`. + + :param func: the target function + :param args: positional arguments passed to ``func`` + :param name: name of the task (will be coerced to a string if not ``None``) + :return: a tuple of (future, task_status_value) where the ``task_status_value`` + is the value passed to ``task_status.started()`` from within the target + function + :rtype: tuple[concurrent.futures.Future[T_Retval], Any] + + .. versionadded:: 3.0 + + """ + + def task_done(future: Future[T_Retval]) -> None: + if not task_status_future.done(): + if future.cancelled(): + task_status_future.cancel() + elif future.exception(): + task_status_future.set_exception(future.exception()) + else: + exc = RuntimeError( + "Task exited without calling task_status.started()" + ) + task_status_future.set_exception(exc) + + self._check_running() + task_status_future: Future = Future() + task_status = _BlockingPortalTaskStatus(task_status_future) + f: Future = Future() + f.add_done_callback(task_done) + self._spawn_task_from_thread(func, args, {"task_status": task_status}, name, f) + return f, task_status_future.result() + + def wrap_async_context_manager( + self, cm: AbstractAsyncContextManager[T_co] + ) -> AbstractContextManager[T_co]: + """ + Wrap an async context manager as a synchronous context manager via this portal. + + Spawns a task that will call both ``__aenter__()`` and ``__aexit__()``, stopping + in the middle until the synchronous context manager exits. + + :param cm: an asynchronous context manager + :return: a synchronous context manager + + .. versionadded:: 2.1 + + """ + return _BlockingAsyncContextManager(cm, self) + + +@dataclass +class BlockingPortalProvider: + """ + A manager for a blocking portal. Used as a context manager. The first thread to + enter this context manager causes a blocking portal to be started with the specific + parameters, and the last thread to exit causes the portal to be shut down. Thus, + there will be exactly one blocking portal running in this context as long as at + least one thread has entered this context manager. + + The parameters are the same as for :func:`~anyio.run`. + + :param backend: name of the backend + :param backend_options: backend options + + .. versionadded:: 4.4 + """ + + backend: str = "asyncio" + backend_options: dict[str, Any] | None = None + _lock: Lock = field(init=False, default_factory=Lock) + _leases: int = field(init=False, default=0) + _portal: BlockingPortal = field(init=False) + _portal_cm: AbstractContextManager[BlockingPortal] | None = field( + init=False, default=None + ) + + def __enter__(self) -> BlockingPortal: + with self._lock: + if self._portal_cm is None: + self._portal_cm = start_blocking_portal( + self.backend, self.backend_options + ) + self._portal = self._portal_cm.__enter__() + + self._leases += 1 + return self._portal + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: TracebackType | None, + ) -> None: + portal_cm: AbstractContextManager[BlockingPortal] | None = None + with self._lock: + assert self._portal_cm + assert self._leases > 0 + self._leases -= 1 + if not self._leases: + portal_cm = self._portal_cm + self._portal_cm = None + del self._portal + + if portal_cm: + portal_cm.__exit__(None, None, None) + + +@contextmanager +def start_blocking_portal( + backend: str = "asyncio", + backend_options: dict[str, Any] | None = None, + *, + name: str | None = None, +) -> Generator[BlockingPortal, Any, None]: + """ + Start a new event loop in a new thread and run a blocking portal in its main task. + + The parameters are the same as for :func:`~anyio.run`. + + :param backend: name of the backend + :param backend_options: backend options + :param name: name of the thread + :return: a context manager that yields a blocking portal + + .. versionchanged:: 3.0 + Usage as a context manager is now required. + + """ + + async def run_portal() -> None: + async with BlockingPortal() as portal_: + if name is None: + current_thread().name = f"{backend}-portal-{id(portal_):x}" + + future.set_result(portal_) + await portal_.sleep_until_stopped() + + def run_blocking_portal() -> None: + if future.set_running_or_notify_cancel(): + try: + run_eventloop( + run_portal, backend=backend, backend_options=backend_options + ) + except BaseException as exc: + if not future.done(): + future.set_exception(exc) + + future: Future[BlockingPortal] = Future() + thread = Thread(target=run_blocking_portal, daemon=True, name=name) + thread.start() + try: + cancel_remaining_tasks = False + portal = future.result() + try: + yield portal + except BaseException: + cancel_remaining_tasks = True + raise + finally: + try: + portal.call(portal.stop, cancel_remaining_tasks) + except RuntimeError: + pass + finally: + thread.join() + + +def check_cancelled() -> None: + """ + Check if the cancel scope of the host task's running the current worker thread has + been cancelled. + + If the host task's current cancel scope has indeed been cancelled, the + backend-specific cancellation exception will be raised. + + :raises RuntimeError: if the current thread was not spawned by + :func:`.to_thread.run_sync` + + """ + try: + token: EventLoopToken = threadlocals.current_token + except AttributeError: + raise NoEventLoopError( + "This function can only be called inside an AnyIO worker thread" + ) from None + + token.backend_class.check_cancelled() diff --git a/venv/Lib/site-packages/anyio/functools.py b/venv/Lib/site-packages/anyio/functools.py new file mode 100644 index 0000000..b0bdfb4 --- /dev/null +++ b/venv/Lib/site-packages/anyio/functools.py @@ -0,0 +1,400 @@ +from __future__ import annotations + +__all__ = ( + "AsyncCacheInfo", + "AsyncCacheParameters", + "AsyncLRUCacheWrapper", + "cache", + "lru_cache", + "reduce", +) + +import functools +from collections import OrderedDict +from collections.abc import ( + AsyncIterable, + Awaitable, + Callable, + Coroutine, + Hashable, + Iterable, +) +from functools import update_wrapper +from inspect import iscoroutinefunction +from typing import ( + Any, + Generic, + NamedTuple, + ParamSpec, + TypedDict, + TypeVar, + cast, + final, + overload, +) +from weakref import WeakKeyDictionary + +from ._core._eventloop import current_time +from ._core._synchronization import Lock +from .lowlevel import RunVar, checkpoint + +T = TypeVar("T") +S = TypeVar("S") +P = ParamSpec("P") +lru_cache_items: RunVar[ + WeakKeyDictionary[ + AsyncLRUCacheWrapper[Any, Any], + OrderedDict[ + Hashable, + tuple[_InitialMissingType, Lock, float | None] + | tuple[Any, None, float | None], + ], + ] +] = RunVar("lru_cache_items") + + +class _InitialMissingType: + pass + + +initial_missing: _InitialMissingType = _InitialMissingType() + + +class AsyncCacheInfo(NamedTuple): + hits: int + misses: int + maxsize: int | None + currsize: int + ttl: int | None + + +class AsyncCacheParameters(TypedDict): + maxsize: int | None + typed: bool + always_checkpoint: bool + ttl: int | None + + +class _LRUMethodWrapper(Generic[T]): + def __init__(self, wrapper: AsyncLRUCacheWrapper[..., T], instance: object): + self.__wrapper = wrapper + self.__instance = instance + + def cache_info(self) -> AsyncCacheInfo: + return self.__wrapper.cache_info() + + def cache_parameters(self) -> AsyncCacheParameters: + return self.__wrapper.cache_parameters() + + def cache_clear(self) -> None: + self.__wrapper.cache_clear() + + async def __call__(self, *args: Any, **kwargs: Any) -> T: + if self.__instance is None: + return await self.__wrapper(*args, **kwargs) + + return await self.__wrapper(self.__instance, *args, **kwargs) + + +@final +class AsyncLRUCacheWrapper(Generic[P, T]): + def __init__( + self, + func: Callable[P, Awaitable[T]], + maxsize: int | None, + typed: bool, + always_checkpoint: bool, + ttl: int | None, + ): + self.__wrapped__ = func + self._hits: int = 0 + self._misses: int = 0 + self._maxsize = max(maxsize, 0) if maxsize is not None else None + self._currsize: int = 0 + self._typed = typed + self._always_checkpoint = always_checkpoint + self._ttl = ttl + update_wrapper(self, func) + + def cache_info(self) -> AsyncCacheInfo: + return AsyncCacheInfo( + self._hits, self._misses, self._maxsize, self._currsize, self._ttl + ) + + def cache_parameters(self) -> AsyncCacheParameters: + return { + "maxsize": self._maxsize, + "typed": self._typed, + "always_checkpoint": self._always_checkpoint, + "ttl": self._ttl, + } + + def cache_clear(self) -> None: + if cache := lru_cache_items.get(None): + cache.pop(self, None) + self._hits = self._misses = self._currsize = 0 + + async def __call__(self, *args: P.args, **kwargs: P.kwargs) -> T: + # Easy case first: if maxsize == 0, no caching is done + if self._maxsize == 0: + value = await self.__wrapped__(*args, **kwargs) + self._misses += 1 + return value + + # The key is constructed as a flat tuple to avoid memory overhead + key: tuple[Any, ...] = args + if kwargs: + # initial_missing is used as a separator + key += (initial_missing,) + sum(kwargs.items(), ()) + + if self._typed: + key += tuple(type(arg) for arg in args) + if kwargs: + key += (initial_missing,) + tuple(type(val) for val in kwargs.values()) + + try: + cache = lru_cache_items.get() + except LookupError: + cache = WeakKeyDictionary() + lru_cache_items.set(cache) + + try: + cache_entry = cache[self] + except KeyError: + cache_entry = cache[self] = OrderedDict() + + cached_value: T | _InitialMissingType + try: + cached_value, lock, expires_at = cache_entry[key] + except KeyError: + # We're the first task to call this function + cached_value, lock, expires_at = ( + initial_missing, + Lock(fast_acquire=not self._always_checkpoint), + None, + ) + cache_entry[key] = cached_value, lock, expires_at + + if lock is None: + if expires_at is not None and current_time() >= expires_at: + self._currsize -= 1 + cached_value, lock, expires_at = ( + initial_missing, + Lock(fast_acquire=not self._always_checkpoint), + None, + ) + cache_entry[key] = cached_value, lock, expires_at + else: + # The value was already cached + self._hits += 1 + cache_entry.move_to_end(key) + if self._always_checkpoint: + await checkpoint() + + return cast(T, cached_value) + + async with lock: + # Check if another task filled the cache while we acquired the lock + if (cached_value := cache_entry[key][0]) is initial_missing: + self._misses += 1 + if self._maxsize is not None and self._currsize >= self._maxsize: + cache_entry.popitem(last=False) + else: + self._currsize += 1 + + value = await self.__wrapped__(*args, **kwargs) + expires_at = ( + current_time() + self._ttl if self._ttl is not None else None + ) + cache_entry[key] = value, None, expires_at + else: + # Another task filled the cache while we were waiting for the lock + self._hits += 1 + cache_entry.move_to_end(key) + value = cast(T, cached_value) + + return value + + def __get__( + self, instance: object, owner: type | None = None + ) -> _LRUMethodWrapper[T]: + wrapper = _LRUMethodWrapper(self, instance) + update_wrapper(wrapper, self.__wrapped__) + return wrapper + + +class _LRUCacheWrapper: + def __init__( + self, maxsize: int | None, typed: bool, always_checkpoint: bool, ttl: int | None + ): + self._maxsize = maxsize + self._typed = typed + self._always_checkpoint = always_checkpoint + self._ttl = ttl + + @overload + def __call__( # type: ignore[overload-overlap] + self, func: Callable[P, Coroutine[Any, Any, T]], / + ) -> AsyncLRUCacheWrapper[P, T]: ... + + @overload + def __call__( + self, func: Callable[..., T], / + ) -> functools._lru_cache_wrapper[T]: ... + + def __call__( + self, f: Callable[P, Coroutine[Any, Any, T]] | Callable[..., T], / + ) -> AsyncLRUCacheWrapper[P, T] | functools._lru_cache_wrapper[T]: + if iscoroutinefunction(f): + return AsyncLRUCacheWrapper( + f, self._maxsize, self._typed, self._always_checkpoint, self._ttl + ) + + return functools.lru_cache(maxsize=self._maxsize, typed=self._typed)(f) # type: ignore[arg-type] + + +@overload +def cache( # type: ignore[overload-overlap] + func: Callable[P, Coroutine[Any, Any, T]], / +) -> AsyncLRUCacheWrapper[P, T]: ... + + +@overload +def cache(func: Callable[..., T], /) -> functools._lru_cache_wrapper[T]: ... + + +def cache(func: Callable[..., Any] | Callable[P, Coroutine[Any, Any, Any]], /) -> Any: + """ + A convenient shortcut for :func:`lru_cache` with ``maxsize=None``. + + This is the asynchronous equivalent to :func:`functools.cache`. + + """ + return lru_cache(maxsize=None)(func) + + +@overload +def lru_cache( + *, + maxsize: int | None = ..., + typed: bool = ..., + always_checkpoint: bool = ..., + ttl: int | None = ..., +) -> _LRUCacheWrapper: ... + + +@overload +def lru_cache( # type: ignore[overload-overlap] + func: Callable[P, Coroutine[Any, Any, T]], / +) -> AsyncLRUCacheWrapper[P, T]: ... + + +@overload +def lru_cache(func: Callable[..., T], /) -> functools._lru_cache_wrapper[T]: ... + + +def lru_cache( + func: Callable[..., Coroutine[Any, Any, Any]] | Callable[..., Any] | None = None, + /, + *, + maxsize: int | None = 128, + typed: bool = False, + always_checkpoint: bool = False, + ttl: int | None = None, +) -> Any: + """ + An asynchronous version of :func:`functools.lru_cache`. + + If a synchronous function is passed, the standard library + :func:`functools.lru_cache` is applied instead. + + :param always_checkpoint: if ``True``, every call to the cached function will be + guaranteed to yield control to the event loop at least once + :param ttl: time in seconds after which to invalidate cache entries + + .. note:: Caches and locks are managed on a per-event loop basis. + + """ + if func is None: + return _LRUCacheWrapper(maxsize, typed, always_checkpoint, ttl) + + if not callable(func): + raise TypeError("the first argument must be callable") + + return _LRUCacheWrapper(maxsize, typed, always_checkpoint, ttl)(func) + + +@overload +async def reduce( + function: Callable[[T, S], Awaitable[T]], + iterable: Iterable[S] | AsyncIterable[S], + /, + initial: T, +) -> T: ... + + +@overload +async def reduce( + function: Callable[[T, T], Awaitable[T]], + iterable: Iterable[T] | AsyncIterable[T], + /, +) -> T: ... + + +async def reduce( # type: ignore[misc] + function: Callable[[T, T], Awaitable[T]] | Callable[[T, S], Awaitable[T]], + iterable: Iterable[T] | Iterable[S] | AsyncIterable[T] | AsyncIterable[S], + /, + initial: T | _InitialMissingType = initial_missing, +) -> T: + """ + Asynchronous version of :func:`functools.reduce`. + + :param function: a coroutine function that takes two arguments: the accumulated + value and the next element from the iterable + :param iterable: an iterable or async iterable + :param initial: the initial value (if missing, the first element of the iterable is + used as the initial value) + + """ + element: Any + function_called = False + if isinstance(iterable, AsyncIterable): + async_it = iterable.__aiter__() + if initial is initial_missing: + try: + value = cast(T, await async_it.__anext__()) + except StopAsyncIteration: + raise TypeError( + "reduce() of empty sequence with no initial value" + ) from None + else: + value = cast(T, initial) + + async for element in async_it: + value = await function(value, element) + function_called = True + elif isinstance(iterable, Iterable): + it = iter(iterable) + if initial is initial_missing: + try: + value = cast(T, next(it)) + except StopIteration: + raise TypeError( + "reduce() of empty sequence with no initial value" + ) from None + else: + value = cast(T, initial) + + for element in it: + value = await function(value, element) + function_called = True + else: + raise TypeError("reduce() argument 2 must be an iterable or async iterable") + + # Make sure there is at least one checkpoint, even if an empty iterable and an + # initial value were given + if not function_called: + await checkpoint() + + return value diff --git a/venv/Lib/site-packages/anyio/itertools.py b/venv/Lib/site-packages/anyio/itertools.py new file mode 100644 index 0000000..7e5248e --- /dev/null +++ b/venv/Lib/site-packages/anyio/itertools.py @@ -0,0 +1,626 @@ +from __future__ import annotations + +__all__ = ( + "accumulate", + "batched", + "Chain", + "combinations", + "combinations_with_replacement", + "compress", + "count", + "cycle", + "dropwhile", + "filterfalse", + "groupby", + "islice", + "pairwise", + "permutations", + "product", + "repeat", + "starmap", + "tee", + "takewhile", + "zip_longest", +) + +import itertools +import operator +import sys +from collections.abc import ( + AsyncGenerator, + AsyncIterable, + AsyncIterator, + Awaitable, + Callable, + Iterable, + Iterator, +) +from dataclasses import dataclass, field +from typing import Any, Generic, TypeVar, cast, overload + +from ._core._synchronization import Lock +from ._core._tasks import CancelScope +from .lowlevel import cancel_shielded_checkpoint, checkpoint, checkpoint_if_cancelled + +T = TypeVar("T") +R = TypeVar("R") +_tee_end = object() + + +@dataclass(eq=False) +class _IterableAsyncIterator(AsyncIterator[T]): + iterator: Iterator[T] + + async def __anext__(self) -> T: + await checkpoint_if_cancelled() + try: + result = next(self.iterator) + except StopIteration: + await cancel_shielded_checkpoint() + raise StopAsyncIteration from None + + await cancel_shielded_checkpoint() + return result + + +def _iterate(iterable: Iterable[T] | AsyncIterable[T]) -> AsyncIterator[T]: + if isinstance(iterable, AsyncIterator): + return iterable + + if isinstance(iterable, AsyncIterable): + return iterable.__aiter__() + + return _IterableAsyncIterator(iter(iterable)) + + +@dataclass(eq=False) +class _TeeLink(Generic[T]): + value: object | None = None + next: _TeeLink[T] | None = None + filled: bool = False + + +@dataclass(eq=False) +class _TeeState(Generic[T]): + iterator: AsyncIterator[T] + lock: Lock = field(default_factory=Lock) + + async def fill(self, link: _TeeLink[T]) -> bool: + if link.filled: + return False + + async with self.lock: + if link.filled: + return True + + link.value = await anext(self.iterator, _tee_end) + if link.value is not _tee_end: + link.next = _TeeLink() + + link.filled = True + return True + + +class _TeeAsyncIterator(AsyncIterator[T]): + _state: _TeeState[T] + _link: _TeeLink[T] + _element_yielded: bool + + def __init__( + self, iterable: Iterable[T] | AsyncIterable[T] | _TeeAsyncIterator[T] + ) -> None: + if isinstance(iterable, _TeeAsyncIterator): + self._state = iterable._state + self._link = iterable._link + else: + self._state = _TeeState(_iterate(iterable)) + self._link = _TeeLink() + + self._element_yielded = False + + async def __anext__(self) -> T: + had_yieldpoint = await self._state.fill(self._link) + if self._link.value is _tee_end: + if not self._element_yielded: + await checkpoint() + + raise StopAsyncIteration + + if not had_yieldpoint: + await checkpoint_if_cancelled() + + self._element_yielded = True + value = cast(T, self._link.value) + next_link = self._link.next + assert next_link is not None + self._link = next_link + if not had_yieldpoint: + await cancel_shielded_checkpoint() + + return value + + +async def _operator_add(x: T, y: T) -> T: + return operator.add(x, y) + + +async def accumulate( + iterable: Iterable[T] | AsyncIterable[T], + function: Callable[[T, T], Awaitable[T]] = _operator_add, + *, + initial: T | None = None, +) -> AsyncGenerator[T, None]: + iterator = _iterate(iterable) + if initial is None: + try: + total = await anext(iterator) + except StopAsyncIteration: + await checkpoint() + return + else: + await checkpoint_if_cancelled() + total = initial + await cancel_shielded_checkpoint() + + yield total + + async for element in iterator: + total = await function(total, element) + yield total + + +async def batched( + iterable: Iterable[T] | AsyncIterable[T], n: int, *, strict: bool = False +) -> AsyncGenerator[tuple[T, ...], None]: + if n < 1: + raise ValueError("n must be at least one") + + iterator = _iterate(iterable) + + while True: + batch: list[T] = [] + for _ in range(n): + try: + batch.append(await anext(iterator)) + except StopAsyncIteration: + if not batch: + await checkpoint() + return + if strict: + raise ValueError("batched(): incomplete batch") from None + + yield tuple(batch) + return + + yield tuple(batch) + + +class Chain: + def __call__( + self, *iterables: Iterable[T] | AsyncIterable[T] + ) -> AsyncGenerator[T, None]: + return self.from_iterable(iterables) + + async def from_iterable( + self, + iterables: ( + Iterable[Iterable[T] | AsyncIterable[T]] + | AsyncIterable[Iterable[T] | AsyncIterable[T]] + ), + ) -> AsyncGenerator[T, None]: + element_yielded = False + outer_iter = _iterate(iterables) + + try: + async for iterable in outer_iter: + async for element in _iterate(iterable): + element_yielded = True + yield element + finally: + aclose = getattr(outer_iter, "aclose", None) + if aclose is not None: + with CancelScope(shield=True): + await aclose() + + if not element_yielded: + await checkpoint() + + +chain: Chain = Chain() + + +async def combinations( + iterable: Iterable[T] | AsyncIterable[T], r: int +) -> AsyncGenerator[tuple[T, ...], None]: + pool: list[T] = [element async for element in _iterate(iterable)] + async for combination in _iterate(itertools.combinations(pool, r)): + yield combination + + +async def combinations_with_replacement( + iterable: Iterable[T] | AsyncIterable[T], r: int +) -> AsyncGenerator[tuple[T, ...], None]: + pool: list[T] = [element async for element in _iterate(iterable)] + async for combination in _iterate(itertools.combinations_with_replacement(pool, r)): + yield combination + + +async def compress( + data: Iterable[T] | AsyncIterable[T], + selectors: Iterable[object] | AsyncIterable[object], +) -> AsyncGenerator[T, None]: + data_iterator = _iterate(data) + selector_iterator = _iterate(selectors) + element_yielded = False + + while True: + try: + datum = await anext(data_iterator) + selector = await anext(selector_iterator) + except StopAsyncIteration: + if not element_yielded: + await checkpoint() + + return + + if selector: + element_yielded = True + yield datum + + +async def count(start: int = 0, step: int = 1) -> AsyncGenerator[int, None]: + n = start + while True: + await checkpoint_if_cancelled() + value = n + n += step + await cancel_shielded_checkpoint() + yield value + + +async def cycle( + iterable: Iterable[T] | AsyncIterable[T], +) -> AsyncGenerator[T, None]: + saved: list[T] = [] + async for element in _iterate(iterable): + saved.append(element) + yield element + + if not saved: + await checkpoint() + return + + while True: + for element in saved: + await checkpoint() + yield element + + +async def dropwhile( + predicate: Callable[[T], Awaitable[object]], + iterable: Iterable[T] | AsyncIterable[T], +) -> AsyncGenerator[T, None]: + element_yielded = False + dropping = True + + async for element in _iterate(iterable): + if dropping and await predicate(element): + continue + + dropping = False + element_yielded = True + yield element + + if not element_yielded: + await checkpoint() + + +async def filterfalse( + predicate: Callable[[T], Awaitable[object]], + iterable: Iterable[T] | AsyncIterable[T], +) -> AsyncGenerator[T, None]: + element_yielded = False + + async for element in _iterate(iterable): + if not await predicate(element): + element_yielded = True + yield element + + if not element_yielded: + await checkpoint() + + +@overload +def groupby( + iterable: Iterable[T] | AsyncIterable[T], +) -> AsyncGenerator[tuple[T, list[T]], None]: ... + + +@overload +def groupby( + iterable: Iterable[T] | AsyncIterable[T], + key: Callable[[T], Awaitable[R]], +) -> AsyncGenerator[tuple[R, list[T]], None]: ... + + +async def groupby( + iterable: Iterable[T] | AsyncIterable[T], + key: Callable[[T], Awaitable[object]] | None = None, +) -> AsyncGenerator[tuple[object, list[T]], None]: + iterator = _iterate(iterable) + try: + element = await anext(iterator) + except StopAsyncIteration: + await checkpoint() + return + + group_key = element if key is None else await key(element) + values = [element] + + async for element in iterator: + next_key = element if key is None else await key(element) + if next_key != group_key: + completed_group = group_key, values + group_key = next_key + values = [element] + yield completed_group + else: + values.append(element) + + yield group_key, values + + +@overload +def islice( + iterable: Iterable[T] | AsyncIterable[T], + stop: int | None, + /, +) -> AsyncGenerator[T, None]: ... + + +@overload +def islice( + iterable: Iterable[T] | AsyncIterable[T], + start: int | None, + stop: int | None, + step: int | None = 1, + /, +) -> AsyncGenerator[T, None]: ... + + +async def islice( + iterable: Iterable[T] | AsyncIterable[T], + *args: int | None, +) -> AsyncGenerator[T, None]: + if not args: + raise TypeError("islice expected at least 2 arguments, got 1") + if len(args) > 3: + raise TypeError(f"islice expected at most 4 arguments, got {len(args) + 1}") + + slice_args = slice(*args) + + start_message = ( + "Indices for islice() must be None or an integer: 0 <= x <= sys.maxsize." + ) + stop_message = ( + "Stop argument for islice() must be None or an integer: 0 <= x <= sys.maxsize." + ) + step_message = "Step for islice() must be a positive integer or None." + + def normalize_index(value: object, message: str) -> int: + try: + index = operator.index(cast(Any, value)) + except TypeError: + raise ValueError(message) from None + + if index < 0 or index > sys.maxsize: + raise ValueError(message) + + return index + + start = ( + 0 + if slice_args.start is None + else normalize_index(slice_args.start, start_message) + ) + stop = ( + None + if slice_args.stop is None + else normalize_index(slice_args.stop, stop_message) + ) + step = ( + 1 if slice_args.step is None else normalize_index(slice_args.step, step_message) + ) + + if step <= 0: + raise ValueError(step_message) + + if stop == 0 or start == stop: + await checkpoint() + return + + iterator = _iterate(iterable) + index = 0 + element_yielded = False + + while stop is None or index < stop: + try: + element = await anext(iterator) + except StopAsyncIteration: + if not element_yielded: + await checkpoint() + + return + + if index >= start and (index - start) % step == 0: + index += 1 + element_yielded = True + yield element + else: + index += 1 + + if not element_yielded: + await checkpoint() + + +async def pairwise( + iterable: Iterable[T] | AsyncIterable[T], +) -> AsyncGenerator[tuple[T, T], None]: + iterator = _iterate(iterable) + try: + previous = await anext(iterator) + except StopAsyncIteration: + await checkpoint() + return + + element_yielded = False + async for element in iterator: + element_yielded = True + pair = (previous, element) + previous = element + yield pair + + if not element_yielded: + await checkpoint() + + +async def permutations( + iterable: Iterable[T] | AsyncIterable[T], r: int | None = None +) -> AsyncGenerator[tuple[T, ...], None]: + pool: list[T] = [element async for element in _iterate(iterable)] + n = len(pool) + if r is None: + r = n + elif not isinstance(r, int): + raise TypeError("Expected int as r") + elif r < 0: + raise ValueError("r must be non-negative") + + async for permutation in _iterate(itertools.permutations(pool, r)): + yield permutation + + +async def product( + *iterables: Iterable[T] | AsyncIterable[T], repeat: int = 1 +) -> AsyncGenerator[tuple[T, ...], None]: + repeat = operator.index(repeat) + if repeat < 0: + raise ValueError("repeat argument cannot be negative") + + pools: list[tuple[T, ...]] = [] + for iterable in iterables: + pool: list[T] = [element async for element in _iterate(iterable)] + pools.append(tuple(pool)) + + async for value in _iterate(itertools.product(*pools, repeat=repeat)): + yield value + + +async def repeat(element: T, times: int | None = None) -> AsyncGenerator[T, None]: + if times is None: + while True: + await checkpoint() + yield element + + remaining = operator.index(cast(Any, times)) + if remaining <= 0: + await checkpoint() + return + + while remaining > 0: + await checkpoint_if_cancelled() + remaining -= 1 + await cancel_shielded_checkpoint() + yield element + + +async def starmap( + function: Callable[..., Awaitable[R]], + iterable: ( + Iterable[Iterable[object] | AsyncIterable[object]] + | AsyncIterable[Iterable[object] | AsyncIterable[object]] + ), +) -> AsyncGenerator[R, None]: + result_yielded = False + + async for args_iterable in _iterate(iterable): + args = [element async for element in _iterate(args_iterable)] + result_yielded = True + yield await function(*args) + + if not result_yielded: + await checkpoint() + + +def tee( + iterable: Iterable[T] | AsyncIterable[T], n: int = 2 +) -> tuple[AsyncIterator[T], ...]: + n = operator.index(cast(Any, n)) + if n < 0: + raise ValueError("n must be >= 0") + if n == 0: + return () + + iterator = _TeeAsyncIterator(iterable) + iterators: list[AsyncIterator[T]] = [iterator] + iterators.extend(_TeeAsyncIterator(iterator) for _ in range(n - 1)) + return tuple(iterators) + + +async def takewhile( + predicate: Callable[[T], Awaitable[object]], + iterable: Iterable[T] | AsyncIterable[T], +) -> AsyncGenerator[T, None]: + element_yielded = False + + async for element in _iterate(iterable): + if not await predicate(element): + if not element_yielded: + await checkpoint() + + return + + element_yielded = True + yield element + + if not element_yielded: + await checkpoint() + + +async def zip_longest( + *iterables: Iterable[object] | AsyncIterable[object], + fillvalue: object = None, +) -> AsyncGenerator[tuple[object, ...], None]: + iterators = [_iterate(iterable) for iterable in iterables] + num_active = len(iterators) + if not num_active: + await checkpoint() + return + + active = [True] * num_active + tuple_yielded = False + + while True: + values: list[object] = [] + for index, iterator in enumerate(iterators): + if not active[index]: + values.append(fillvalue) + continue + + try: + value = await anext(iterator) + except StopAsyncIteration: + active[index] = False + num_active -= 1 + if not num_active: + if not tuple_yielded: + await checkpoint() + + return + + value = fillvalue + + values.append(value) + + tuple_yielded = True + yield tuple(values) diff --git a/venv/Lib/site-packages/anyio/lowlevel.py b/venv/Lib/site-packages/anyio/lowlevel.py new file mode 100644 index 0000000..d045791 --- /dev/null +++ b/venv/Lib/site-packages/anyio/lowlevel.py @@ -0,0 +1,226 @@ +from __future__ import annotations + +__all__ = ( + "EventLoopToken", + "RunvarToken", + "RunVar", + "checkpoint", + "checkpoint_if_cancelled", + "cancel_shielded_checkpoint", + "current_token", +) + +import enum +from dataclasses import dataclass +from types import TracebackType +from typing import Any, Generic, Literal, TypeVar, final, overload +from weakref import WeakKeyDictionary + +from ._core._eventloop import get_async_backend +from .abc import AsyncBackend + +T = TypeVar("T") +D = TypeVar("D") + + +async def checkpoint() -> None: + """ + Check for cancellation and allow the scheduler to switch to another task. + + Equivalent to (but more efficient than):: + + await checkpoint_if_cancelled() + await cancel_shielded_checkpoint() + + .. versionadded:: 3.0 + + """ + await get_async_backend().checkpoint() + + +async def checkpoint_if_cancelled() -> None: + """ + Enter a checkpoint if the enclosing cancel scope has been cancelled. + + This does not allow the scheduler to switch to a different task. + + .. versionadded:: 3.0 + + """ + await get_async_backend().checkpoint_if_cancelled() + + +async def cancel_shielded_checkpoint() -> None: + """ + Allow the scheduler to switch to another task but without checking for cancellation. + + Equivalent to (but potentially more efficient than):: + + with CancelScope(shield=True): + await checkpoint() + + .. versionadded:: 3.0 + + """ + await get_async_backend().cancel_shielded_checkpoint() + + +@final +@dataclass(frozen=True, repr=False) +class EventLoopToken: + """ + An opaque object that holds a reference to an event loop. + + .. versionadded:: 4.11.0 + """ + + backend_class: type[AsyncBackend] + native_token: object + + +def current_token() -> EventLoopToken: + """ + Return a token object that can be used to call code in the current event loop from + another thread. + + :raises NoEventLoopError: if no supported asynchronous event loop is running in the + current thread + + .. versionadded:: 4.11.0 + + """ + backend_class = get_async_backend() + raw_token = backend_class.current_token() + return EventLoopToken(backend_class, raw_token) + + +_run_vars: WeakKeyDictionary[object, dict[RunVar[Any], Any]] = WeakKeyDictionary() + + +class _NoValueSet(enum.Enum): + NO_VALUE_SET = enum.auto() + + +class RunvarToken(Generic[T]): + """ + A token that can be used to restore a :class:`RunVar` to its previous value. + + Returned by :meth:`RunVar.set`. Can be used as a context manager to automatically + reset the variable on exit, or passed directly to :meth:`RunVar.reset`. + """ + + __slots__ = "_var", "_value", "_redeemed" + + def __init__(self, var: RunVar[T], value: T | Literal[_NoValueSet.NO_VALUE_SET]): + self._var = var + self._value: T | Literal[_NoValueSet.NO_VALUE_SET] = value + self._redeemed = False + + def __enter__(self) -> RunvarToken[T]: + return self + + def __exit__( + self, + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: TracebackType | None, + ) -> None: + self._var.reset(self) + + +class RunVar(Generic[T]): + """ + Like a :class:`~contextvars.ContextVar`, except scoped to the running event loop. + + Can be used as a context manager, Just like :class:`~contextvars.ContextVar`, that + will reset the variable to its previous value when the context block is exited. + """ + + __slots__ = "_name", "_default" + + NO_VALUE_SET: Literal[_NoValueSet.NO_VALUE_SET] = _NoValueSet.NO_VALUE_SET + + def __init__( + self, name: str, default: T | Literal[_NoValueSet.NO_VALUE_SET] = NO_VALUE_SET + ): + self._name = name + self._default = default + + @property + def _current_vars(self) -> dict[RunVar[T], T]: + native_token = current_token().native_token + try: + return _run_vars[native_token] + except KeyError: + run_vars = _run_vars[native_token] = {} + return run_vars + + @overload + def get(self, default: D) -> T | D: ... + + @overload + def get(self) -> T: ... + + def get( + self, default: D | Literal[_NoValueSet.NO_VALUE_SET] = NO_VALUE_SET + ) -> T | D: + """ + Return the current value of this run variable. + + :param default: a fallback value to return if no value has been set + :return: the current value, the provided default, or the variable's own default + :raises LookupError: if no value is set and no default is available + + """ + try: + return self._current_vars[self] + except KeyError: + if default is not RunVar.NO_VALUE_SET: + return default + elif self._default is not RunVar.NO_VALUE_SET: + return self._default + + raise LookupError( + f'Run variable "{self._name}" has no value and no default set' + ) + + def set(self, value: T) -> RunvarToken[T]: + """ + Set the value of this run variable for the current event loop. + + :param value: the new value + :return: a token that can be used to restore the previous value + + """ + current_vars = self._current_vars + token = RunvarToken(self, current_vars.get(self, RunVar.NO_VALUE_SET)) + current_vars[self] = value + return token + + def reset(self, token: RunvarToken[T]) -> None: + """ + Restore this run variable to the value it held before the matching :meth:`set`. + + :param token: the token returned by :meth:`set` + :raises ValueError: if the token belongs to a different :class:`RunVar` or the token + has already been used + + """ + if token._var is not self: + raise ValueError("This token does not belong to this RunVar") + + if token._redeemed: + raise ValueError("This token has already been used") + + if token._value is _NoValueSet.NO_VALUE_SET: + try: + del self._current_vars[self] + except KeyError: + pass + else: + self._current_vars[self] = token._value + + token._redeemed = True + + def __repr__(self) -> str: + return f""