| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973974975976977978979980981982983984985986987988989990991992993994995996997998999100010011002100310041005100610071008100910101011101210131014101510161017101810191020102110221023102410251026102710281029103010311032103310341035103610371038103910401041104210431044104510461047104810491050105110521053105410551056105710581059106010611062106310641065106610671068106910701071107210731074107510761077107810791080108110821083108410851086108710881089109010911092109310941095109610971098109911001101110211031104110511061107110811091110111111121113111411151116111711181119112011211122112311241125112611271128112911301131113211331134113511361137113811391140114111421143114411451146114711481149115011511152115311541155115611571158115911601161116211631164116511661167116811691170117111721173117411751176117711781179118011811182118311841185118611871188118911901191119211931194119511961197119811991200120112021203120412051206120712081209121012111212121312141215121612171218121912201221122212231224122512261227122812291230123112321233123412351236123712381239124012411242124312441245124612471248124912501251125212531254125512561257125812591260126112621263126412651266126712681269127012711272127312741275127612771278127912801281128212831284128512861287128812891290129112921293129412951296129712981299130013011302130313041305130613071308130913101311131213131314131513161317131813191320132113221323132413251326132713281329133013311332133313341335133613371338133913401341134213431344134513461347134813491350135113521353135413551356135713581359136013611362136313641365136613671368136913701371137213731374137513761377137813791380138113821383138413851386138713881389139013911392139313941395139613971398139914001401140214031404140514061407140814091410141114121413141414151416141714181419142014211422142314241425142614271428142914301431143214331434143514361437143814391440144114421443144414451446144714481449145014511452145314541455145614571458145914601461146214631464146514661467146814691470147114721473147414751476147714781479148014811482148314841485148614871488148914901491149214931494149514961497149814991500150115021503150415051506150715081509151015111512151315141515151615171518151915201521152215231524152515261527152815291530153115321533153415351536153715381539154015411542154315441545154615471548154915501551155215531554155515561557155815591560156115621563156415651566156715681569157015711572157315741575157615771578157915801581158215831584158515861587158815891590159115921593159415951596159715981599160016011602160316041605160616071608160916101611161216131614161516161617161816191620162116221623162416251626162716281629163016311632163316341635163616371638163916401641164216431644164516461647164816491650165116521653165416551656165716581659166016611662166316641665166616671668166916701671167216731674167516761677167816791680168116821683168416851686168716881689169016911692169316941695169616971698169917001701170217031704170517061707170817091710171117121713171417151716171717181719172017211722172317241725172617271728172917301731173217331734173517361737173817391740174117421743174417451746174717481749175017511752175317541755175617571758175917601761176217631764176517661767176817691770177117721773177417751776177717781779178017811782178317841785178617871788178917901791179217931794179517961797179817991800180118021803180418051806180718081809181018111812181318141815181618171818181918201821182218231824182518261827182818291830183118321833183418351836183718381839184018411842184318441845184618471848184918501851185218531854185518561857185818591860186118621863186418651866186718681869187018711872187318741875187618771878187918801881188218831884188518861887188818891890189118921893189418951896189718981899190019011902190319041905190619071908190919101911191219131914191519161917191819191920192119221923192419251926192719281929193019311932193319341935193619371938193919401941194219431944194519461947194819491950195119521953195419551956195719581959196019611962196319641965196619671968196919701971197219731974197519761977197819791980198119821983198419851986198719881989199019911992199319941995199619971998199920002001200220032004200520062007200820092010201120122013201420152016201720182019202020212022202320242025202620272028202920302031203220332034203520362037203820392040204120422043204420452046204720482049205020512052205320542055205620572058205920602061206220632064206520662067206820692070207120722073207420752076207720782079208020812082208320842085208620872088208920902091209220932094209520962097209820992100210121022103210421052106210721082109211021112112211321142115211621172118211921202121212221232124212521262127212821292130213121322133213421352136213721382139214021412142214321442145214621472148214921502151215221532154215521562157215821592160216121622163216421652166216721682169217021712172217321742175217621772178217921802181218221832184218521862187218821892190219121922193219421952196219721982199220022012202220322042205220622072208220922102211221222132214221522162217221822192220222122222223222422252226222722282229223022312232223322342235223622372238223922402241224222432244224522462247224822492250225122522253225422552256225722582259226022612262226322642265226622672268226922702271227222732274227522762277227822792280228122822283228422852286228722882289229022912292229322942295229622972298229923002301230223032304230523062307230823092310231123122313231423152316231723182319232023212322232323242325232623272328232923302331233223332334233523362337233823392340234123422343234423452346234723482349235023512352235323542355235623572358235923602361236223632364236523662367236823692370237123722373237423752376237723782379238023812382238323842385238623872388238923902391239223932394239523962397239823992400240124022403240424052406240724082409241024112412241324142415241624172418241924202421242224232424242524262427242824292430243124322433243424352436243724382439244024412442244324442445244624472448244924502451245224532454245524562457245824592460246124622463246424652466246724682469247024712472247324742475247624772478247924802481248224832484248524862487248824892490249124922493249424952496249724982499250025012502250325042505250625072508250925102511251225132514251525162517251825192520252125222523252425252526252725282529253025312532253325342535253625372538253925402541254225432544254525462547254825492550255125522553255425552556255725582559256025612562256325642565256625672568256925702571257225732574257525762577257825792580258125822583258425852586258725882589259025912592259325942595259625972598259926002601260226032604260526062607260826092610261126122613261426152616261726182619262026212622262326242625262626272628262926302631263226332634263526362637263826392640264126422643264426452646264726482649265026512652265326542655265626572658265926602661266226632664266526662667266826692670267126722673267426752676267726782679268026812682268326842685268626872688268926902691269226932694269526962697269826992700270127022703270427052706270727082709271027112712271327142715271627172718271927202721272227232724272527262727272827292730273127322733273427352736273727382739274027412742274327442745274627472748274927502751275227532754275527562757275827592760276127622763276427652766276727682769277027712772277327742775277627772778277927802781278227832784278527862787278827892790279127922793279427952796279727982799280028012802280328042805280628072808280928102811281228132814281528162817281828192820282128222823282428252826282728282829283028312832283328342835283628372838283928402841284228432844284528462847284828492850285128522853285428552856285728582859286028612862286328642865286628672868286928702871287228732874287528762877287828792880288128822883288428852886288728882889289028912892289328942895289628972898289929002901290229032904290529062907290829092910291129122913291429152916291729182919292029212922292329242925292629272928292929302931293229332934293529362937293829392940294129422943294429452946294729482949295029512952295329542955295629572958295929602961296229632964296529662967296829692970297129722973297429752976297729782979298029812982298329842985298629872988298929902991299229932994299529962997299829993000300130023003300430053006300730083009301030113012301330143015301630173018301930203021302230233024302530263027302830293030303130323033303430353036303730383039304030413042304330443045304630473048304930503051305230533054305530563057305830593060306130623063306430653066306730683069307030713072307330743075307630773078307930803081308230833084308530863087308830893090309130923093309430953096309730983099310031013102310331043105310631073108310931103111311231133114311531163117311831193120312131223123312431253126312731283129313031313132313331343135313631373138313931403141314231433144314531463147314831493150315131523153315431553156315731583159316031613162316331643165316631673168316931703171317231733174317531763177317831793180318131823183318431853186318731883189319031913192319331943195319631973198319932003201 |
- from __future__ import annotations
- import array
- import asyncio
- import concurrent.futures
- import math
- import os
- import socket
- import sys
- import threading
- import weakref
- from asyncio import (
- AbstractEventLoop,
- CancelledError,
- all_tasks,
- create_task,
- current_task,
- get_running_loop,
- sleep,
- )
- from asyncio.base_events import _run_until_complete_cb # type: ignore[attr-defined]
- from collections import OrderedDict, deque
- from collections.abc import (
- AsyncGenerator,
- AsyncIterator,
- Awaitable,
- Callable,
- Collection,
- Coroutine,
- Iterable,
- Sequence,
- )
- from concurrent.futures import Future
- from contextlib import AbstractContextManager, suppress
- from contextvars import Context, copy_context
- from dataclasses import dataclass, field
- from functools import partial, wraps
- from inspect import (
- CORO_RUNNING,
- CORO_SUSPENDED,
- getcoroutinestate,
- )
- from io import IOBase
- from os import PathLike
- from queue import Queue
- from signal import Signals
- from socket import AddressFamily, SocketKind
- from threading import Thread
- from types import CodeType, TracebackType
- from typing import (
- IO,
- TYPE_CHECKING,
- Any,
- Literal,
- ParamSpec,
- TypeVar,
- cast,
- )
- from weakref import WeakKeyDictionary
- from .. import (
- CapacityLimiterStatistics,
- EventStatistics,
- LockStatistics,
- TaskInfo,
- abc,
- )
- from .._core._eventloop import (
- claim_worker_thread,
- set_current_async_library,
- threadlocals,
- )
- from .._core._exceptions import (
- BrokenResourceError,
- BusyResourceError,
- ClosedResourceError,
- EndOfStream,
- RunFinishedError,
- WouldBlock,
- )
- from .._core._sockets import convert_ipv6_sockaddr
- from .._core._streams import create_memory_object_stream
- from .._core._synchronization import (
- CapacityLimiter as BaseCapacityLimiter,
- )
- from .._core._synchronization import Event as BaseEvent
- from .._core._synchronization import Lock as BaseLock
- from .._core._synchronization import (
- ResourceGuard,
- SemaphoreStatistics,
- )
- from .._core._synchronization import Semaphore as BaseSemaphore
- from .._core._tasks import CancelScope as BaseCancelScope
- from .._core._tasks import TaskHandle
- from ..abc import (
- AsyncBackend,
- IPSockAddrType,
- SocketListener,
- UDPPacketType,
- UNIXDatagramPacketType,
- )
- from ..abc._tasks import call_for_coroutine, get_callable_name, get_coro_name
- from ..lowlevel import RunVar, _run_vars
- if TYPE_CHECKING:
- from _typeshed import FileDescriptorLike
- from ..abc._eventloop import StrOrBytesPath
- from ..streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream
- else:
- FileDescriptorLike = object
- if sys.version_info >= (3, 11):
- from asyncio import Runner
- from typing import Self, TypeVarTuple, Unpack
- else:
- import contextvars
- import enum
- import signal
- from asyncio import coroutines, events, exceptions, tasks
- from exceptiongroup import BaseExceptionGroup
- from typing_extensions import Self, TypeVarTuple, Unpack
- class _State(enum.Enum):
- CREATED = "created"
- INITIALIZED = "initialized"
- CLOSED = "closed"
- class Runner:
- # Copied from CPython 3.11
- def __init__(
- self,
- *,
- debug: bool | None = None,
- loop_factory: Callable[[], AbstractEventLoop] | None = None,
- ):
- self._state = _State.CREATED
- self._debug = debug
- self._loop_factory = loop_factory
- self._loop: AbstractEventLoop | None = None
- self._context = None
- self._interrupt_count = 0
- self._set_event_loop = False
- def __enter__(self) -> Self:
- self._lazy_init()
- return self
- def __exit__(
- self,
- exc_type: type[BaseException] | None,
- exc_val: BaseException | None,
- exc_tb: TracebackType | None,
- ) -> None:
- self.close()
- def close(self) -> None:
- """Shutdown and close event loop."""
- loop = self._loop
- if self._state is not _State.INITIALIZED or loop is None:
- return
- try:
- _cancel_all_tasks(loop)
- loop.run_until_complete(loop.shutdown_asyncgens())
- if hasattr(loop, "shutdown_default_executor"):
- loop.run_until_complete(loop.shutdown_default_executor())
- else:
- loop.run_until_complete(_shutdown_default_executor(loop))
- finally:
- if self._set_event_loop:
- events.set_event_loop(None)
- loop.close()
- self._loop = None
- self._state = _State.CLOSED
- def get_loop(self) -> AbstractEventLoop:
- """Return embedded event loop."""
- self._lazy_init()
- return self._loop
- def run(self, coro: Coroutine[T_Retval], *, context=None) -> T_Retval:
- """Run a coroutine inside the embedded event loop."""
- if not coroutines.iscoroutine(coro):
- raise ValueError(f"a coroutine was expected, got {coro!r}")
- if events._get_running_loop() is not None:
- # fail fast with short traceback
- raise RuntimeError(
- "Runner.run() cannot be called from a running event loop"
- )
- self._lazy_init()
- if context is None:
- context = self._context
- task = context.run(self._loop.create_task, coro)
- if (
- threading.current_thread() is threading.main_thread()
- and signal.getsignal(signal.SIGINT) is signal.default_int_handler
- ):
- sigint_handler = partial(self._on_sigint, main_task=task)
- try:
- signal.signal(signal.SIGINT, sigint_handler)
- except ValueError:
- # `signal.signal` may throw if `threading.main_thread` does
- # not support signals (e.g. embedded interpreter with signals
- # not registered - see gh-91880)
- sigint_handler = None
- else:
- sigint_handler = None
- self._interrupt_count = 0
- try:
- return self._loop.run_until_complete(task)
- except exceptions.CancelledError:
- if self._interrupt_count > 0:
- uncancel = getattr(task, "uncancel", None)
- if uncancel is not None and uncancel() == 0:
- raise KeyboardInterrupt # noqa: B904
- raise # CancelledError
- finally:
- if (
- sigint_handler is not None
- and signal.getsignal(signal.SIGINT) is sigint_handler
- ):
- signal.signal(signal.SIGINT, signal.default_int_handler)
- def _lazy_init(self) -> None:
- if self._state is _State.CLOSED:
- raise RuntimeError("Runner is closed")
- if self._state is _State.INITIALIZED:
- return
- if self._loop_factory is None:
- self._loop = events.new_event_loop()
- if not self._set_event_loop:
- # Call set_event_loop only once to avoid calling
- # attach_loop multiple times on child watchers
- events.set_event_loop(self._loop)
- self._set_event_loop = True
- else:
- self._loop = self._loop_factory()
- if self._debug is not None:
- self._loop.set_debug(self._debug)
- self._context = contextvars.copy_context()
- self._state = _State.INITIALIZED
- def _on_sigint(self, signum, frame, main_task: asyncio.Task) -> None:
- self._interrupt_count += 1
- if self._interrupt_count == 1 and not main_task.done():
- main_task.cancel()
- # wakeup loop if it is blocked by select() with long timeout
- self._loop.call_soon_threadsafe(lambda: None)
- return
- raise KeyboardInterrupt()
- def _cancel_all_tasks(loop: AbstractEventLoop) -> None:
- to_cancel = tasks.all_tasks(loop)
- if not to_cancel:
- return
- for task in to_cancel:
- task.cancel()
- loop.run_until_complete(tasks.gather(*to_cancel, return_exceptions=True))
- for task in to_cancel:
- if task.cancelled():
- continue
- if task.exception() is not None:
- loop.call_exception_handler(
- {
- "message": "unhandled exception during asyncio.run() shutdown",
- "exception": task.exception(),
- "task": task,
- }
- )
- async def _shutdown_default_executor(loop: AbstractEventLoop) -> None:
- """Schedule the shutdown of the default executor."""
- def _do_shutdown(future: asyncio.futures.Future) -> None:
- try:
- loop._default_executor.shutdown(wait=True) # type: ignore[attr-defined]
- loop.call_soon_threadsafe(future.set_result, None)
- except Exception as ex:
- loop.call_soon_threadsafe(future.set_exception, ex)
- loop._executor_shutdown_called = True
- if loop._default_executor is None:
- return
- future = loop.create_future()
- thread = threading.Thread(target=_do_shutdown, args=(future,))
- thread.start()
- try:
- await future
- finally:
- thread.join()
- T_Retval = TypeVar("T_Retval")
- T_co = TypeVar("T_co", covariant=True)
- T_contra = TypeVar("T_contra", contravariant=True)
- PosArgsT = TypeVarTuple("PosArgsT")
- P = ParamSpec("P")
- _root_task: RunVar[asyncio.Task[Any] | None] = RunVar("_root_task")
- def find_root_task() -> asyncio.Task:
- root_task = _root_task.get(None)
- if root_task is not None and not root_task.done():
- return root_task
- # Look for a task that has been started via run_until_complete()
- for task in all_tasks():
- if task._callbacks and not task.done():
- for cb, context in task._callbacks:
- if (
- cb is _run_until_complete_cb
- or getattr(cb, "__module__", None) == "uvloop.loop"
- ):
- _root_task.set(task)
- def _unset(t: asyncio.Task[Any]) -> None:
- if vars := _run_vars.get(t.get_loop()):
- vars.pop(_root_task, None)
- # Register a callback to break the task -> loop -> _run_var[loop][_root_task] -> task cycle
- # Also run it in its own context to not create another reference.
- # We can't use RunVar.reset() here since these are called synchronously
- # and thus lowlevel.current_token() (which RunVar.reset() depends on) fails.
- task.add_done_callback(_unset, context=context)
- return task
- # Look up the topmost task in the AnyIO task tree, if possible
- task = cast(asyncio.Task, current_task())
- state = _task_states.get(task)
- if state:
- cancel_scope = state.cancel_scope
- while cancel_scope and cancel_scope._parent_scope is not None:
- cancel_scope = cancel_scope._parent_scope
- if cancel_scope is not None:
- return cast(asyncio.Task, cancel_scope._host_task)
- return task
- def _task_started(task: asyncio.Task) -> bool:
- """Return ``True`` if the task has been started and has not finished."""
- # The task coro should never be None here, as we never add finished tasks to the
- # task list
- coro = task.get_coro()
- assert coro is not None
- return getcoroutinestate(coro) in (CORO_RUNNING, CORO_SUSPENDED)
- #
- # Timeouts and cancellation
- #
- def is_anyio_cancellation(exc: CancelledError) -> bool:
- # Sometimes third party frameworks catch a CancelledError and raise a new one, so as
- # a workaround we have to look at the previous ones in __context__ too for a
- # matching cancel message
- while True:
- if (
- exc.args
- and isinstance(exc.args[0], str)
- and exc.args[0].startswith("Cancelled via cancel scope ")
- ):
- return True
- if isinstance(exc.__context__, CancelledError):
- exc = exc.__context__
- continue
- return False
- class CancelScope(BaseCancelScope):
- __slots__ = (
- "_active",
- "_cancel_called",
- "_cancel_handle",
- "_cancel_reason",
- "_cancelled_caught",
- "_child_scopes",
- "_deadline",
- "_host_task",
- "_parent_scope",
- "_pending_uncancellations",
- "_shield",
- "_tasks",
- "_timeout_handle",
- )
- def __new__(cls, *, deadline: float = math.inf, shield: bool = False) -> Self:
- return object.__new__(cls)
- def __init__(self, deadline: float = math.inf, shield: bool = False):
- self._deadline = deadline
- self._shield = shield
- self._parent_scope: CancelScope | None = None
- self._child_scopes: set[CancelScope] = set()
- self._cancel_called = False
- self._cancel_reason: str | None = None
- self._cancelled_caught = False
- self._active = False
- self._timeout_handle: asyncio.TimerHandle | None = None
- self._cancel_handle: asyncio.Handle | None = None
- self._tasks: set[asyncio.Task] = set()
- self._host_task: asyncio.Task | None = None
- if sys.version_info >= (3, 11):
- self._pending_uncancellations: int | None = 0
- else:
- self._pending_uncancellations = None
- def __enter__(self) -> Self:
- if self._active:
- raise RuntimeError(
- "Each CancelScope may only be used for a single 'with' block"
- )
- self._host_task = host_task = cast(asyncio.Task, current_task())
- self._tasks.add(host_task)
- try:
- task_state = _task_states[host_task]
- except KeyError:
- task_state = TaskState(None, self)
- _task_states[host_task] = task_state
- else:
- self._parent_scope = task_state.cancel_scope
- task_state.cancel_scope = self
- if self._parent_scope is not None:
- # If using an eager task factory, the parent scope may not even contain
- # the host task
- self._parent_scope._child_scopes.add(self)
- self._parent_scope._tasks.discard(host_task)
- self._timeout()
- self._active = True
- # Start cancelling the host task if the scope was cancelled before entering
- if self._cancel_called:
- self._deliver_cancellation(self)
- return self
- def __exit__(
- self,
- exc_type: type[BaseException] | None,
- exc_val: BaseException | None,
- exc_tb: TracebackType | None,
- ) -> bool:
- del exc_tb
- if not self._active:
- raise RuntimeError("This cancel scope is not active")
- if current_task() is not self._host_task:
- raise RuntimeError(
- "Attempted to exit cancel scope in a different task than it was "
- "entered in"
- )
- assert self._host_task is not None
- host_task_state = _task_states.get(self._host_task)
- if host_task_state is None or host_task_state.cancel_scope is not self:
- raise RuntimeError(
- "Attempted to exit a cancel scope that isn't the current tasks's "
- "current cancel scope"
- )
- try:
- self._active = False
- if self._timeout_handle:
- self._timeout_handle.cancel()
- self._timeout_handle = None
- self._tasks.remove(self._host_task)
- if self._parent_scope is not None:
- self._parent_scope._child_scopes.remove(self)
- self._parent_scope._tasks.add(self._host_task)
- host_task_state.cancel_scope = self._parent_scope
- # Restart the cancellation effort in the closest visible, cancelled parent
- # scope if necessary
- self._restart_cancellation_in_parent()
- # We only swallow the exception iff it was an AnyIO CancelledError, either
- # directly as exc_val or inside an exception group and there are no cancelled
- # parent cancel scopes visible to us here
- if self._cancel_called and not self._parent_cancellation_is_visible_to_us:
- # For each level-cancel() call made on the host task, call uncancel()
- while self._pending_uncancellations:
- self._host_task.uncancel()
- self._pending_uncancellations -= 1
- # Update cancelled_caught and check for exceptions we must not swallow
- if isinstance(exc_val, BaseExceptionGroup):
- cancelleds_caught, remaining = exc_val.split(
- lambda exc: (
- isinstance(exc, CancelledError)
- and is_anyio_cancellation(exc)
- )
- )
- if cancelleds_caught is None:
- return False
- self._cancelled_caught = True
- if remaining is None:
- return True
- context = remaining.__context__
- try:
- # Preserve __cause__ and __suppress_context__ by avoiding `raise
- # ... from ...`
- raise remaining
- finally:
- # Preserve __context__
- remaining.__context__ = context
- del context
- else:
- if isinstance(exc_val, CancelledError) and is_anyio_cancellation(
- exc_val
- ):
- self._cancelled_caught = True
- return True
- else:
- return False
- else:
- if self._pending_uncancellations:
- assert self._parent_scope is not None
- assert self._parent_scope._pending_uncancellations is not None
- self._parent_scope._pending_uncancellations += (
- self._pending_uncancellations
- )
- self._pending_uncancellations = 0
- return False
- finally:
- self._host_task = None
- del exc_val
- @property
- def _effectively_cancelled(self) -> bool:
- cancel_scope: CancelScope | None = self
- while cancel_scope is not None:
- if cancel_scope._cancel_called:
- return True
- if cancel_scope.shield:
- return False
- cancel_scope = cancel_scope._parent_scope
- return False
- @property
- def _parent_cancellation_is_visible_to_us(self) -> bool:
- return (
- self._parent_scope is not None
- and not self.shield
- and self._parent_scope._effectively_cancelled
- )
- def _timeout(self) -> None:
- if self._deadline != math.inf:
- loop = get_running_loop()
- if loop.time() >= self._deadline:
- self.cancel("deadline exceeded")
- else:
- self._timeout_handle = loop.call_at(self._deadline, self._timeout)
- def _deliver_cancellation(self, origin: CancelScope) -> bool:
- """
- Deliver cancellation to directly contained tasks and nested cancel scopes.
- Schedule another run at the end if we still have tasks eligible for
- cancellation.
- :param origin: the cancel scope that originated the cancellation
- :return: ``True`` if the delivery needs to be retried on the next cycle
- """
- should_retry = False
- current = current_task()
- for task in self._tasks:
- # Always skip tasks that are already done (see issue #1111)
- if task.done():
- continue
- should_retry = True
- if task._must_cancel: # type: ignore[attr-defined]
- continue
- # The task is eligible for cancellation if it has started
- if task is not current and (task is self._host_task or _task_started(task)):
- waiter = task._fut_waiter # type: ignore[attr-defined]
- if not isinstance(waiter, asyncio.Future) or not waiter.done():
- task.cancel(origin._cancel_reason)
- if (
- task is origin._host_task
- and origin._pending_uncancellations is not None
- ):
- origin._pending_uncancellations += 1
- # Deliver cancellation to child scopes that aren't shielded or running their own
- # cancellation callbacks
- for scope in self._child_scopes:
- if not scope._shield and not scope.cancel_called:
- should_retry = scope._deliver_cancellation(origin) or should_retry
- # Schedule another callback if there are still tasks left
- if origin is self:
- if should_retry:
- self._cancel_handle = get_running_loop().call_soon(
- self._deliver_cancellation, origin
- )
- else:
- self._cancel_handle = None
- return should_retry
- def _restart_cancellation_in_parent(self) -> None:
- """
- Restart the cancellation effort in the closest directly cancelled parent scope.
- """
- scope = self._parent_scope
- while scope is not None:
- if scope._cancel_called:
- if scope._cancel_handle is None:
- scope._deliver_cancellation(scope)
- break
- # No point in looking beyond any shielded scope
- if scope._shield:
- break
- scope = scope._parent_scope
- def _reparent(self, new_parent: CancelScope) -> None:
- """
- Move this active scope from its current parent to ``new_parent``.
- Used by :meth:`TaskGroup.start` to move a task that has just called
- ``task_status.started()`` into the target task group's cancel scope.
- """
- if self._parent_scope is new_parent:
- return
- if self._parent_scope is not None:
- self._parent_scope._child_scopes.discard(self)
- self._parent_scope = new_parent
- new_parent._child_scopes.add(self)
- # If the new parent (or an ancestor) is already cancelled, (re)start the
- # delivery loop to ensure we're cancelled at next checkpoint like Trio.
- self._restart_cancellation_in_parent()
- def cancel(self, reason: str | None = None) -> None:
- if not self._cancel_called:
- if self._timeout_handle:
- self._timeout_handle.cancel()
- self._timeout_handle = None
- self._cancel_called = True
- self._cancel_reason = f"Cancelled via cancel scope {id(self):x}"
- if task := current_task():
- self._cancel_reason += f" by {task}"
- if reason:
- self._cancel_reason += f"; reason: {reason}"
- if self._host_task is not None:
- self._deliver_cancellation(self)
- @property
- def deadline(self) -> float:
- return self._deadline
- @deadline.setter
- def deadline(self, value: float) -> None:
- self._deadline = float(value)
- if self._timeout_handle is not None:
- self._timeout_handle.cancel()
- self._timeout_handle = None
- if self._active and not self._cancel_called:
- self._timeout()
- @property
- def cancel_called(self) -> bool:
- return self._cancel_called
- @property
- def cancelled_caught(self) -> bool:
- return self._cancelled_caught
- @property
- def shield(self) -> bool:
- return self._shield
- @shield.setter
- def shield(self, value: bool) -> None:
- if self._shield != value:
- self._shield = value
- if not value:
- self._restart_cancellation_in_parent()
- #
- # Task states
- #
- class TaskState:
- """
- Encapsulates auxiliary task information that cannot be added to the Task instance
- itself because there are no guarantees about its implementation.
- """
- __slots__ = "__weakref__", "cancel_scope", "parent_id"
- def __init__(self, parent_id: int | None, cancel_scope: CancelScope | None):
- self.parent_id = parent_id
- self.cancel_scope = cancel_scope
- _task_states: WeakKeyDictionary[asyncio.Task, TaskState] = WeakKeyDictionary()
- #
- # Task groups
- #
- class _AsyncioTaskStatus(abc.TaskStatus[T_contra]):
- def __init__(
- self,
- future: asyncio.Future,
- parent_id: int,
- target_scope: CancelScope,
- spawn_scope: CancelScope,
- ):
- self._future = future
- self._parent_id = parent_id
- # The eventual parent scope for this spawn_scope
- # (after task_status.started() has been called)
- self._target_scope = target_scope
- # The task's own cancel scope, also held by its TaskHandle
- self._spawn_scope = spawn_scope
- def started(self, value: T_contra | None = None) -> None:
- task = cast(asyncio.Task, current_task())
- _task_states[task].parent_id = self._parent_id
- if self._future.done():
- if not self._future.cancelled():
- raise RuntimeError("called 'started' twice on the same task status")
- else:
- # Caller of start() was cancelled, nothing to reparent
- return
- self._future.set_result(value)
- self._spawn_scope._reparent(self._target_scope)
- if sys.version_info >= (3, 12):
- _eager_task_factory_code: CodeType | None = asyncio.eager_task_factory.__code__
- else:
- _eager_task_factory_code = None
- class TaskGroup(abc.TaskGroup):
- def __init__(self) -> None:
- self.cancel_scope: CancelScope = CancelScope()
- self._entered = False
- self._exceptions: list[BaseException] = []
- self._tasks: set[asyncio.Task] = set()
- self._on_completed_fut: asyncio.Future[None] | None = None
- async def __aenter__(self) -> Self:
- if self._entered:
- raise RuntimeError("TaskGroup cannot be entered more than once")
- self._entered = True
- self.cancel_scope.__enter__()
- return self
- async def __aexit__(
- self,
- exc_type: type[BaseException] | None,
- exc_val: BaseException | None,
- exc_tb: TracebackType | None,
- ) -> bool:
- try:
- if exc_val is not None:
- self.cancel_scope.cancel()
- if not isinstance(exc_val, CancelledError):
- self._exceptions.append(exc_val)
- loop = get_running_loop()
- try:
- if self._tasks:
- with CancelScope() as wait_scope:
- while self._tasks:
- self._on_completed_fut = loop.create_future()
- try:
- await self._on_completed_fut
- except CancelledError as exc:
- # Shield the scope against further cancellation attempts,
- # as they're not productive (#695)
- wait_scope.shield = True
- self.cancel_scope.cancel()
- # Set exc_val from the cancellation exception if it was
- # previously unset. However, we should not replace a native
- # cancellation exception with one raise by a cancel scope.
- if exc_val is None or (
- isinstance(exc_val, CancelledError)
- and not is_anyio_cancellation(exc)
- ):
- exc_val = exc
- self._on_completed_fut = None
- else:
- # If there are no child tasks to wait on, run at least one checkpoint
- # anyway
- await AsyncIOBackend.cancel_shielded_checkpoint()
- if self._exceptions:
- # The exception that got us here should already have been
- # added to self._exceptions so it's ok to break exception
- # chaining and avoid adding a "During handling of above..."
- # for each nesting level.
- raise BaseExceptionGroup(
- "unhandled errors in a TaskGroup", self._exceptions
- ) from None
- elif exc_val:
- raise exc_val
- except BaseException as exc:
- if self.cancel_scope.__exit__(type(exc), exc, exc.__traceback__):
- return True
- raise
- return self.cancel_scope.__exit__(exc_type, exc_val, exc_tb)
- finally:
- del exc_val, exc_tb, self._exceptions
- def _spawn(
- self,
- coro: Coroutine[Any, Any, T_co],
- name: object,
- task_status: _AsyncioTaskStatus | None = None,
- ) -> TaskHandle[T_co]:
- task_status_future = task_status._future if task_status is not None else None
- def task_done(_task: asyncio.Task) -> None:
- if sys.version_info >= (3, 14) and self.cancel_scope._host_task is not None:
- asyncio.future_discard_from_awaited_by(
- _task, self.cancel_scope._host_task
- )
- task_state = _task_states[_task]
- assert task_state.cancel_scope is not None
- assert _task in task_state.cancel_scope._tasks
- task_state.cancel_scope._tasks.remove(_task)
- self._tasks.remove(task)
- del _task_states[_task]
- if self._on_completed_fut is not None and not self._tasks:
- try:
- self._on_completed_fut.set_result(None)
- except asyncio.InvalidStateError:
- pass
- try:
- exc = _task.exception()
- except CancelledError as e:
- while isinstance(e.__context__, CancelledError):
- e = e.__context__
- exc = e
- if exc is not None:
- # The future can only be in the cancelled state if the host task was
- # cancelled, so return immediately instead of adding one more
- # CancelledError to the exceptions list
- if task_status_future is not None and task_status_future.cancelled():
- return
- if task_status_future is None or task_status_future.done():
- if not isinstance(exc, CancelledError):
- self._exceptions.append(exc)
- if not self.cancel_scope._effectively_cancelled:
- self.cancel_scope.cancel()
- else:
- task_status_future.set_exception(exc)
- elif task_status_future is not None and not task_status_future.done():
- task_status_future.set_exception(
- RuntimeError("Child exited without calling task_status.started()")
- )
- if task_status_future is not None:
- parent_id = id(current_task())
- caller_state = _task_states.get(cast(asyncio.Task, current_task()))
- if caller_state is not None and caller_state.cancel_scope is not None:
- initial_scope = caller_state.cancel_scope
- else:
- # The caller is an unmanaged task (no task state)
- initial_scope = self.cancel_scope
- else:
- parent_id = id(self.cancel_scope._host_task)
- initial_scope = self.cancel_scope
- spawn_scope = task_status._spawn_scope if task_status is not None else None
- handle = TaskHandle(coro, name, cancel_scope=spawn_scope)
- loop = asyncio.get_running_loop()
- wrapper_coro = handle._run_coro()
- try:
- if (
- (factory := loop.get_task_factory())
- and getattr(factory, "__code__", None) is _eager_task_factory_code
- and (closure := getattr(factory, "__closure__", None))
- ):
- custom_task_constructor = closure[0].cell_contents
- task = custom_task_constructor(
- wrapper_coro, loop=loop, name=handle.name
- )
- else:
- task = loop.create_task(wrapper_coro, name=handle.name)
- except BaseException:
- with suppress(BaseException):
- wrapper_coro.close()
- with suppress(BaseException):
- coro.close()
- raise
- # Make the spawned task inherit the initial cancel scope
- _task_states[task] = TaskState(parent_id=parent_id, cancel_scope=initial_scope)
- initial_scope._tasks.add(task)
- self._tasks.add(task)
- if sys.version_info >= (3, 14) and self.cancel_scope._host_task is not None:
- asyncio.future_add_to_awaited_by(task, self.cancel_scope._host_task)
- task.add_done_callback(task_done)
- return handle
- def create_task(
- self,
- coro: Coroutine[Any, Any, T_co],
- *,
- name: object = None,
- context: Context | None = None,
- ) -> TaskHandle[T_co]:
- if not isinstance(coro, Coroutine):
- raise TypeError(f"expected a coroutine, got {coro.__class__.__qualname__}")
- if not self._entered or not self.cancel_scope._active:
- coro.close()
- raise RuntimeError(
- "This task group is not active; no new tasks can be started."
- )
- final_name = get_coro_name(coro, name)
- if context is not None:
- return context.run(self._spawn, coro, name=final_name)
- else:
- return self._spawn(coro, name=final_name)
- async def start(
- self,
- func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]],
- *args: Unpack[PosArgsT],
- name: object = None,
- return_handle: Literal[False, True] = False,
- ) -> Any:
- if not self._entered or not self.cancel_scope._active:
- raise RuntimeError(
- "This task group is not active; no new tasks can be started."
- )
- # Until task_status.started() is called, a task spawned via start() belongs to
- # the *caller's* cancel scope, not the target group's, so cancelling the group
- # does not cancel a task that hasn't reported startup yet. The
- # task_status.started() call moves the task's cancel scope to this task group's
- # cancel scope.
- #
- # The caller may be an unmanaged task (no _task_states entry), in which case
- # fall back to the group's own scope.
- future: asyncio.Future = asyncio.Future()
- final_name = get_callable_name(func, name)
- task_status: _AsyncioTaskStatus[Any] = _AsyncioTaskStatus(
- future, id(self.cancel_scope._host_task), self.cancel_scope, CancelScope()
- )
- coro = call_for_coroutine(func, args, task_status=task_status)
- handle = self._spawn(coro, final_name, task_status)
- # If the task raises an exception after sending a start value without a switch
- # point between, the task group is cancelled and this method never proceeds to
- # process the completed future. That's why we have to have a shielded cancel
- # scope here.
- try:
- await future
- except BaseException:
- if handle.status is TaskHandle.Status.PENDING:
- # Cancel the task and wait for it to exit before returning
- handle.cancel()
- with CancelScope(shield=True):
- await handle.wait()
- raise
- if return_handle:
- handle._start_value = future.result()
- return handle
- else:
- return future.result()
- #
- # Threads
- #
- _Retval_Queue_Type = tuple[T_Retval | None, BaseException | None]
- class WorkerThread(Thread):
- MAX_IDLE_TIME = 10 # seconds
- def __init__(
- self,
- root_task: asyncio.Task,
- workers: set[WorkerThread],
- idle_workers: deque[WorkerThread],
- ):
- kwargs: dict[str, Any] = {}
- if sys.version_info >= (3, 14):
- kwargs["context"] = Context()
- super().__init__(name="AnyIO worker thread", **kwargs)
- self.root_task = root_task
- self.workers = workers
- self.idle_workers = idle_workers
- self.loop = root_task._loop
- self.queue: Queue[
- tuple[Context, Callable, tuple, asyncio.Future, CancelScope] | None
- ] = Queue(2)
- self.idle_since = AsyncIOBackend.current_time()
- self.stopping = False
- def _report_result(
- self, future: asyncio.Future, result: Any, exc: BaseException | None
- ) -> None:
- self.idle_since = AsyncIOBackend.current_time()
- if not self.stopping:
- self.idle_workers.append(self)
- if not future.cancelled():
- if exc is not None:
- if isinstance(exc, StopIteration):
- new_exc = RuntimeError("coroutine raised StopIteration")
- new_exc.__cause__ = exc
- exc = new_exc
- future.set_exception(exc)
- else:
- future.set_result(result)
- def run(self) -> None:
- with claim_worker_thread(AsyncIOBackend, self.loop):
- while True:
- item = self.queue.get()
- if item is None:
- # Shutdown command received
- return
- context, func, args, future, cancel_scope = item
- if not future.cancelled():
- result = None
- exception: BaseException | None = None
- threadlocals.current_cancel_scope = cancel_scope
- try:
- result = context.run(func, *args)
- except BaseException as exc:
- exception = exc
- finally:
- del threadlocals.current_cancel_scope
- try:
- self.loop.call_soon_threadsafe(
- self._report_result, future, result, exception
- )
- except RuntimeError:
- if not self.loop.is_closed():
- raise
- del result, exception
- self.queue.task_done()
- del item, context, func, args, future, cancel_scope
- def stop(self, f: asyncio.Task | None = None) -> None:
- self.stopping = True
- self.queue.put_nowait(None)
- self.workers.discard(self)
- try:
- self.idle_workers.remove(self)
- except ValueError:
- pass
- _threadpool_idle_workers: RunVar[deque[WorkerThread]] = RunVar(
- "_threadpool_idle_workers"
- )
- _threadpool_workers: RunVar[set[WorkerThread]] = RunVar("_threadpool_workers")
- #
- # Subprocesses
- #
- @dataclass(eq=False)
- class StreamReaderWrapper(abc.ByteReceiveStream):
- _stream: asyncio.StreamReader
- async def receive(self, max_bytes: int = 65536) -> bytes:
- if max_bytes < 1:
- raise ValueError("max_bytes must be a positive integer")
- data = await self._stream.read(max_bytes)
- if data:
- return data
- else:
- raise EndOfStream
- async def aclose(self) -> None:
- self._stream.set_exception(ClosedResourceError())
- await AsyncIOBackend.checkpoint()
- @dataclass(eq=False)
- class StreamWriterWrapper(abc.ByteSendStream):
- _stream: asyncio.StreamWriter
- _closed: bool = field(init=False, default=False)
- async def send(self, item: bytes) -> None:
- await AsyncIOBackend.checkpoint_if_cancelled()
- stream_paused = self._stream._protocol._paused # type: ignore[attr-defined]
- try:
- self._stream.write(item)
- await self._stream.drain()
- except (ConnectionResetError, BrokenPipeError, RuntimeError) as exc:
- # If closed by us and/or the peer:
- # * on stdlib, drain() raises ConnectionResetError or BrokenPipeError
- # * on uvloop and Winloop, write() eventually starts raising RuntimeError
- if self._closed:
- raise ClosedResourceError from exc
- elif self._stream.is_closing():
- raise BrokenResourceError from exc
- raise
- if not stream_paused:
- await AsyncIOBackend.cancel_shielded_checkpoint()
- async def aclose(self) -> None:
- self._closed = True
- self._stream.close()
- await AsyncIOBackend.checkpoint()
- @dataclass(eq=False)
- class Process(abc.Process):
- _process: asyncio.subprocess.Process
- _stdin: StreamWriterWrapper | None
- _stdout: StreamReaderWrapper | None
- _stderr: StreamReaderWrapper | None
- _exited: asyncio.Event
- _transport: asyncio.SubprocessTransport
- async def aclose(self) -> None:
- with CancelScope(shield=True) as scope:
- # We need to close the underlying pipe_transports as well to allow a
- # process blocking on full buffers to receive SIGPIPE and exit.
- if self._stdin:
- await self._stdin.aclose()
- if pipe := self._transport.get_pipe_transport(0):
- pipe.close()
- if self._stdout:
- await self._stdout.aclose()
- if pipe := self._transport.get_pipe_transport(1):
- pipe.close()
- if self._stderr:
- await self._stderr.aclose()
- if pipe := self._transport.get_pipe_transport(2):
- pipe.close()
- scope.shield = False
- try:
- await self.wait()
- except BaseException:
- scope.shield = True
- # Closing the transport on asyncio also handles sending kill
- self._transport.close()
- await self.wait()
- raise
- async def wait(self) -> int:
- await self._exited.wait()
- assert self._process.returncode is not None
- return self._process.returncode
- def terminate(self) -> None:
- self._process.terminate()
- def kill(self) -> None:
- self._process.kill()
- def send_signal(self, signal: int) -> None:
- self._process.send_signal(signal)
- @property
- def pid(self) -> int:
- return self._process.pid
- @property
- def returncode(self) -> int | None:
- return self._process.returncode
- @property
- def stdin(self) -> abc.ByteSendStream | None:
- return self._stdin
- @property
- def stdout(self) -> abc.ByteReceiveStream | None:
- return self._stdout
- @property
- def stderr(self) -> abc.ByteReceiveStream | None:
- return self._stderr
- def _forcibly_shutdown_process_pool_on_exit(
- workers: set[Process], _task: object
- ) -> None:
- """
- Forcibly shuts down worker processes belonging to this event loop."""
- child_watcher: asyncio.AbstractChildWatcher | None = None # type: ignore[name-defined]
- if sys.version_info < (3, 12):
- try:
- child_watcher = asyncio.get_event_loop_policy().get_child_watcher()
- except NotImplementedError:
- pass
- # Close as much as possible (w/o async/await) to avoid warnings
- for process in workers.copy():
- if process.returncode is not None:
- continue
- process._stdin._stream._transport.close() # type: ignore[union-attr]
- process._stdout._stream._transport.close() # type: ignore[union-attr]
- process._stderr._stream._transport.close() # type: ignore[union-attr]
- process.kill()
- if child_watcher:
- child_watcher.remove_child_handler(process.pid)
- async def _shutdown_process_pool_on_exit(workers: set[abc.Process]) -> None:
- """
- Shuts down worker processes belonging to this event loop.
- NOTE: this only works when the event loop was started using asyncio.run() or
- anyio.run().
- """
- process: abc.Process
- try:
- await sleep(math.inf)
- except asyncio.CancelledError:
- workers = workers.copy()
- for process in workers:
- if process.returncode is None:
- process.kill()
- for process in workers:
- await process.aclose()
- #
- # Sockets and networking
- #
- class StreamProtocol(asyncio.Protocol):
- read_queue: deque[bytes]
- read_event: asyncio.Event
- write_event: asyncio.Event
- exception: Exception | None = None
- is_at_eof: bool = False
- def connection_made(self, transport: asyncio.BaseTransport) -> None:
- self.read_queue = deque()
- self.read_event = asyncio.Event()
- self.write_event = asyncio.Event()
- self.write_event.set()
- cast(asyncio.Transport, transport).set_write_buffer_limits(0)
- def connection_lost(self, exc: Exception | None) -> None:
- if exc:
- self.exception = exc
- self.read_event.set()
- self.write_event.set()
- def data_received(self, data: bytes) -> None:
- # ProactorEventloop sometimes sends bytearray instead of bytes
- self.read_queue.append(bytes(data))
- self.read_event.set()
- def eof_received(self) -> bool | None:
- self.is_at_eof = True
- self.read_event.set()
- return True
- def pause_writing(self) -> None:
- self.write_event = asyncio.Event()
- def resume_writing(self) -> None:
- self.write_event.set()
- class DatagramProtocol(asyncio.DatagramProtocol):
- read_queue: deque[tuple[bytes, IPSockAddrType]]
- read_event: asyncio.Event
- write_event: asyncio.Event
- closed_event: asyncio.Event
- exception: Exception | None = None
- def connection_made(self, transport: asyncio.BaseTransport) -> None:
- self.read_queue = deque(maxlen=100) # arbitrary value
- self.read_event = asyncio.Event()
- self.write_event = asyncio.Event()
- self.closed_event = asyncio.Event()
- self.write_event.set()
- def connection_lost(self, exc: Exception | None) -> None:
- self.read_event.set()
- self.write_event.set()
- self.closed_event.set()
- def datagram_received(self, data: bytes, addr: IPSockAddrType) -> None:
- addr = convert_ipv6_sockaddr(addr)
- self.read_queue.append((data, addr))
- self.read_event.set()
- def error_received(self, exc: Exception) -> None:
- self.exception = exc
- def pause_writing(self) -> None:
- self.write_event.clear()
- def resume_writing(self) -> None:
- self.write_event.set()
- class SocketStream(abc.SocketStream):
- def __init__(self, transport: asyncio.Transport, protocol: StreamProtocol):
- self._transport = transport
- self._protocol = protocol
- self._receive_guard = ResourceGuard("reading from")
- self._send_guard = ResourceGuard("writing to")
- self._closed = False
- @property
- def _raw_socket(self) -> socket.socket:
- return self._transport.get_extra_info("socket")
- async def receive(self, max_bytes: int = 65536) -> bytes:
- if max_bytes < 1:
- raise ValueError("max_bytes must be a positive integer")
- with self._receive_guard:
- if (
- not self._protocol.read_event.is_set()
- and not self._transport.is_closing()
- and not self._protocol.is_at_eof
- ):
- self._transport.resume_reading()
- await self._protocol.read_event.wait()
- self._transport.pause_reading()
- else:
- await AsyncIOBackend.checkpoint()
- try:
- chunk = self._protocol.read_queue.popleft()
- except IndexError:
- if self._closed:
- raise ClosedResourceError from None
- elif self._protocol.exception:
- raise BrokenResourceError from self._protocol.exception
- else:
- raise EndOfStream from None
- if len(chunk) > max_bytes:
- # Split the oversized chunk
- chunk, leftover = chunk[:max_bytes], chunk[max_bytes:]
- self._protocol.read_queue.appendleft(leftover)
- # If the read queue is empty, clear the flag so that the next call will
- # block until data is available
- if not self._protocol.read_queue:
- self._protocol.read_event.clear()
- return chunk
- async def send(self, item: bytes) -> None:
- with self._send_guard:
- await AsyncIOBackend.checkpoint()
- if self._closed:
- raise ClosedResourceError
- elif self._protocol.exception is not None:
- raise BrokenResourceError from self._protocol.exception
- try:
- self._transport.write(item)
- except RuntimeError as exc:
- if self._transport.is_closing():
- raise BrokenResourceError from exc
- else:
- raise
- await self._protocol.write_event.wait()
- async def send_eof(self) -> None:
- try:
- self._transport.write_eof()
- except OSError:
- pass
- async def aclose(self) -> None:
- self._closed = True
- if not self._transport.is_closing():
- try:
- self._transport.write_eof()
- except OSError:
- pass
- self._transport.close()
- await sleep(0)
- self._transport.abort()
- class _RawSocketMixin:
- _receive_future: asyncio.Future | None = None
- _send_future: asyncio.Future | None = None
- _closing = False
- def __init__(self, raw_socket: socket.socket):
- self.__raw_socket = raw_socket
- self._receive_guard = ResourceGuard("reading from")
- self._send_guard = ResourceGuard("writing to")
- @property
- def _raw_socket(self) -> socket.socket:
- return self.__raw_socket
- def _wait_until_readable(self, loop: asyncio.AbstractEventLoop) -> asyncio.Future:
- def callback(f: object) -> None:
- del self._receive_future
- loop.remove_reader(self.__raw_socket)
- f = self._receive_future = asyncio.Future()
- loop.add_reader(self.__raw_socket, f.set_result, None)
- f.add_done_callback(callback)
- return f
- def _wait_until_writable(self, loop: asyncio.AbstractEventLoop) -> asyncio.Future:
- def callback(f: object) -> None:
- del self._send_future
- loop.remove_writer(self.__raw_socket)
- f = self._send_future = asyncio.Future()
- loop.add_writer(self.__raw_socket, f.set_result, None)
- f.add_done_callback(callback)
- return f
- async def aclose(self) -> None:
- if not self._closing:
- self._closing = True
- if self.__raw_socket.fileno() != -1:
- self.__raw_socket.close()
- if self._receive_future and not self._receive_future.done():
- self._receive_future.set_result(None)
- if self._send_future and not self._send_future.done():
- self._send_future.set_result(None)
- class UNIXSocketStream(_RawSocketMixin, abc.UNIXSocketStream):
- async def send_eof(self) -> None:
- with self._send_guard:
- self._raw_socket.shutdown(socket.SHUT_WR)
- async def receive(self, max_bytes: int = 65536) -> bytes:
- if max_bytes < 1:
- raise ValueError("max_bytes must be a positive integer")
- loop = get_running_loop()
- await AsyncIOBackend.checkpoint()
- with self._receive_guard:
- while True:
- try:
- data = self._raw_socket.recv(max_bytes)
- except BlockingIOError:
- await self._wait_until_readable(loop)
- except OSError as exc:
- if self._closing:
- raise ClosedResourceError from None
- else:
- raise BrokenResourceError from exc
- else:
- if not data:
- raise EndOfStream
- return data
- async def send(self, item: bytes) -> None:
- loop = get_running_loop()
- await AsyncIOBackend.checkpoint()
- with self._send_guard:
- view = memoryview(item)
- while view:
- try:
- bytes_sent = self._raw_socket.send(view)
- except BlockingIOError:
- await self._wait_until_writable(loop)
- except OSError as exc:
- if self._closing:
- raise ClosedResourceError from None
- else:
- raise BrokenResourceError from exc
- else:
- view = view[bytes_sent:]
- async def receive_fds(self, msglen: int, maxfds: int) -> tuple[bytes, list[int]]:
- if not isinstance(msglen, int) or msglen < 0:
- raise ValueError("msglen must be a non-negative integer")
- if not isinstance(maxfds, int) or maxfds < 1:
- raise ValueError("maxfds must be a positive integer")
- loop = get_running_loop()
- fds = array.array("i")
- await AsyncIOBackend.checkpoint()
- with self._receive_guard:
- while True:
- try:
- message, ancdata, _flags, _addr = self._raw_socket.recvmsg(
- msglen, socket.CMSG_LEN(maxfds * fds.itemsize)
- )
- except BlockingIOError:
- await self._wait_until_readable(loop)
- except OSError as exc:
- if self._closing:
- raise ClosedResourceError from None
- else:
- raise BrokenResourceError from exc
- else:
- if not message and not ancdata:
- raise EndOfStream
- break
- for cmsg_level, cmsg_type, cmsg_data in ancdata:
- if cmsg_level != socket.SOL_SOCKET or cmsg_type != socket.SCM_RIGHTS:
- raise RuntimeError(
- f"Received unexpected ancillary data; message = {message!r}, "
- f"cmsg_level = {cmsg_level}, cmsg_type = {cmsg_type}"
- )
- fds.frombytes(cmsg_data[: len(cmsg_data) - (len(cmsg_data) % fds.itemsize)])
- return message, list(fds)
- async def send_fds(self, message: bytes, fds: Collection[int | IOBase]) -> None:
- if not message:
- raise ValueError("message must not be empty")
- if not fds:
- raise ValueError("fds must not be empty")
- loop = get_running_loop()
- filenos: list[int] = []
- for fd in fds:
- if isinstance(fd, int):
- filenos.append(fd)
- elif isinstance(fd, IOBase):
- filenos.append(fd.fileno())
- fdarray = array.array("i", filenos)
- await AsyncIOBackend.checkpoint()
- with self._send_guard:
- while True:
- try:
- # The ignore can be removed after mypy picks up
- # https://github.com/python/typeshed/pull/5545
- self._raw_socket.sendmsg(
- [message], [(socket.SOL_SOCKET, socket.SCM_RIGHTS, fdarray)]
- )
- break
- except BlockingIOError:
- await self._wait_until_writable(loop)
- except OSError as exc:
- if self._closing:
- raise ClosedResourceError from None
- else:
- raise BrokenResourceError from exc
- class TCPSocketListener(abc.SocketListener):
- _accept_scope: CancelScope | None = None
- _closed = False
- def __init__(self, raw_socket: socket.socket):
- self.__raw_socket = raw_socket
- self._loop = cast(asyncio.BaseEventLoop, get_running_loop())
- self._accept_guard = ResourceGuard("accepting connections from")
- @property
- def _raw_socket(self) -> socket.socket:
- return self.__raw_socket
- async def accept(self) -> abc.SocketStream:
- if self._closed:
- raise ClosedResourceError
- with self._accept_guard:
- await AsyncIOBackend.checkpoint()
- with CancelScope() as self._accept_scope:
- try:
- client_sock, _addr = await self._loop.sock_accept(self._raw_socket)
- except asyncio.CancelledError:
- # Workaround for https://bugs.python.org/issue41317
- try:
- self._loop.remove_reader(self._raw_socket)
- except (ValueError, NotImplementedError):
- pass
- if self._closed:
- raise ClosedResourceError from None
- raise
- finally:
- self._accept_scope = None
- client_sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
- transport, protocol = await self._loop.connect_accepted_socket(
- StreamProtocol, client_sock
- )
- return SocketStream(transport, protocol)
- async def aclose(self) -> None:
- if self._closed:
- return
- self._closed = True
- if self._accept_scope:
- # Workaround for https://bugs.python.org/issue41317
- try:
- self._loop.remove_reader(self._raw_socket)
- except (ValueError, NotImplementedError):
- pass
- self._accept_scope.cancel()
- await sleep(0)
- self._raw_socket.close()
- class UNIXSocketListener(abc.SocketListener):
- def __init__(self, raw_socket: socket.socket):
- self.__raw_socket = raw_socket
- self._loop = get_running_loop()
- self._accept_guard = ResourceGuard("accepting connections from")
- self._closed = False
- async def accept(self) -> abc.SocketStream:
- await AsyncIOBackend.checkpoint()
- with self._accept_guard:
- while True:
- try:
- client_sock, _ = self.__raw_socket.accept()
- client_sock.setblocking(False)
- return UNIXSocketStream(client_sock)
- except BlockingIOError:
- f: asyncio.Future = asyncio.Future()
- self._loop.add_reader(self.__raw_socket, f.set_result, None)
- f.add_done_callback(
- lambda _: self._loop.remove_reader(self.__raw_socket)
- )
- await f
- except OSError as exc:
- if self._closed:
- raise ClosedResourceError from None
- else:
- raise BrokenResourceError from exc
- async def aclose(self) -> None:
- self._closed = True
- self.__raw_socket.close()
- @property
- def _raw_socket(self) -> socket.socket:
- return self.__raw_socket
- class UDPSocket(abc.UDPSocket):
- def __init__(
- self, transport: asyncio.DatagramTransport, protocol: DatagramProtocol
- ):
- self._transport = transport
- self._protocol = protocol
- self._receive_guard = ResourceGuard("reading from")
- self._send_guard = ResourceGuard("writing to")
- self._closed = False
- @property
- def _raw_socket(self) -> socket.socket:
- return self._transport.get_extra_info("socket")
- async def aclose(self) -> None:
- self._closed = True
- if not self._transport.is_closing():
- self._transport.close()
- await self._protocol.closed_event.wait()
- async def receive(self) -> tuple[bytes, IPSockAddrType]:
- with self._receive_guard:
- await AsyncIOBackend.checkpoint()
- # If the buffer is empty, ask for more data
- if not self._protocol.read_queue and not self._transport.is_closing():
- self._protocol.read_event.clear()
- await self._protocol.read_event.wait()
- try:
- return self._protocol.read_queue.popleft()
- except IndexError:
- if self._closed:
- raise ClosedResourceError from None
- else:
- raise BrokenResourceError from None
- async def send(self, item: UDPPacketType) -> None:
- with self._send_guard:
- await AsyncIOBackend.checkpoint()
- await self._protocol.write_event.wait()
- if self._closed:
- raise ClosedResourceError
- elif self._transport.is_closing():
- raise BrokenResourceError
- else:
- self._transport.sendto(*item)
- class ConnectedUDPSocket(abc.ConnectedUDPSocket):
- def __init__(
- self, transport: asyncio.DatagramTransport, protocol: DatagramProtocol
- ):
- self._transport = transport
- self._protocol = protocol
- self._receive_guard = ResourceGuard("reading from")
- self._send_guard = ResourceGuard("writing to")
- self._closed = False
- @property
- def _raw_socket(self) -> socket.socket:
- return self._transport.get_extra_info("socket")
- async def aclose(self) -> None:
- self._closed = True
- if not self._transport.is_closing():
- self._transport.close()
- await self._protocol.closed_event.wait()
- async def receive(self) -> bytes:
- with self._receive_guard:
- await AsyncIOBackend.checkpoint()
- # If the buffer is empty, ask for more data
- if not self._protocol.read_queue and not self._transport.is_closing():
- self._protocol.read_event.clear()
- await self._protocol.read_event.wait()
- try:
- packet = self._protocol.read_queue.popleft()
- except IndexError:
- if self._closed:
- raise ClosedResourceError from None
- else:
- raise BrokenResourceError from None
- return packet[0]
- async def send(self, item: bytes) -> None:
- with self._send_guard:
- await AsyncIOBackend.checkpoint()
- await self._protocol.write_event.wait()
- if self._closed:
- raise ClosedResourceError
- elif self._transport.is_closing():
- raise BrokenResourceError
- else:
- self._transport.sendto(item)
- class UNIXDatagramSocket(_RawSocketMixin, abc.UNIXDatagramSocket):
- async def receive(self) -> UNIXDatagramPacketType:
- loop = get_running_loop()
- await AsyncIOBackend.checkpoint()
- with self._receive_guard:
- while True:
- try:
- data = self._raw_socket.recvfrom(65536)
- except BlockingIOError:
- await self._wait_until_readable(loop)
- except OSError as exc:
- if self._closing:
- raise ClosedResourceError from None
- else:
- raise BrokenResourceError from exc
- else:
- return data
- async def send(self, item: UNIXDatagramPacketType) -> None:
- loop = get_running_loop()
- await AsyncIOBackend.checkpoint()
- with self._send_guard:
- while True:
- try:
- self._raw_socket.sendto(*item)
- except BlockingIOError:
- await self._wait_until_writable(loop)
- except OSError as exc:
- if self._closing:
- raise ClosedResourceError from None
- else:
- raise BrokenResourceError from exc
- else:
- return
- class ConnectedUNIXDatagramSocket(_RawSocketMixin, abc.ConnectedUNIXDatagramSocket):
- async def receive(self) -> bytes:
- loop = get_running_loop()
- await AsyncIOBackend.checkpoint()
- with self._receive_guard:
- while True:
- try:
- data = self._raw_socket.recv(65536)
- except BlockingIOError:
- await self._wait_until_readable(loop)
- except OSError as exc:
- if self._closing:
- raise ClosedResourceError from None
- else:
- raise BrokenResourceError from exc
- else:
- return data
- async def send(self, item: bytes) -> None:
- loop = get_running_loop()
- await AsyncIOBackend.checkpoint()
- with self._send_guard:
- while True:
- try:
- self._raw_socket.send(item)
- except BlockingIOError:
- await self._wait_until_writable(loop)
- except OSError as exc:
- if self._closing:
- raise ClosedResourceError from None
- else:
- raise BrokenResourceError from exc
- else:
- return
- _read_events: RunVar[dict[int, asyncio.Future[bool]]] = RunVar("read_events")
- _write_events: RunVar[dict[int, asyncio.Future[bool]]] = RunVar("write_events")
- #
- # Synchronization
- #
- class Event(BaseEvent):
- __slots__ = ("_event",)
- def __new__(cls) -> Self:
- return object.__new__(cls)
- def __init__(self) -> None:
- self._event = asyncio.Event()
- def set(self) -> None:
- self._event.set()
- def is_set(self) -> bool:
- return self._event.is_set()
- async def wait(self) -> None:
- if self.is_set():
- await AsyncIOBackend.checkpoint()
- else:
- await self._event.wait()
- def statistics(self) -> EventStatistics:
- return EventStatistics(len(self._event._waiters))
- class Lock(BaseLock):
- __slots__ = "_fast_acquire", "_owner_task", "_waiters"
- def __new__(cls, *, fast_acquire: bool = False) -> Self:
- return object.__new__(cls)
- def __init__(self, *, fast_acquire: bool = False) -> None:
- self._fast_acquire = fast_acquire
- self._owner_task: asyncio.Task | None = None
- self._waiters: deque[tuple[asyncio.Task, asyncio.Future]] = deque()
- async def acquire(self) -> None:
- task = cast(asyncio.Task, current_task())
- if self._owner_task is None and not self._waiters:
- await AsyncIOBackend.checkpoint_if_cancelled()
- self._owner_task = task
- # Unless on the "fast path", yield control of the event loop so that other
- # tasks can run too
- if not self._fast_acquire:
- try:
- await AsyncIOBackend.cancel_shielded_checkpoint()
- except CancelledError:
- self.release()
- raise
- return
- if self._owner_task == task:
- raise RuntimeError("Attempted to acquire an already held Lock")
- fut: asyncio.Future[None] = asyncio.Future()
- item = task, fut
- self._waiters.append(item)
- try:
- await fut
- except CancelledError:
- if fut.cancelled():
- try:
- self._waiters.remove(item)
- except ValueError:
- pass
- else:
- self.release()
- raise
- def acquire_nowait(self) -> None:
- task = cast(asyncio.Task, current_task())
- if self._owner_task is None and not self._waiters:
- self._owner_task = task
- return
- if self._owner_task is task:
- raise RuntimeError("Attempted to acquire an already held Lock")
- raise WouldBlock
- def locked(self) -> bool:
- return self._owner_task is not None
- def release(self) -> None:
- if self._owner_task != current_task():
- raise RuntimeError("The current task is not holding this lock")
- # A cancelled waiter that already received ownership removes itself from
- # _waiters before calling release(); any cancelled waiter still queued here
- # was cancelled before being woken, so drop it.
- while self._waiters:
- task, fut = self._waiters.popleft()
- if fut.cancelled():
- continue
- self._owner_task = task
- fut.set_result(None)
- return
- self._owner_task = None
- def statistics(self) -> LockStatistics:
- task_info = AsyncIOTaskInfo(self._owner_task) if self._owner_task else None
- return LockStatistics(self.locked(), task_info, len(self._waiters))
- class Semaphore(BaseSemaphore):
- __slots__ = "_fast_acquire", "_max_value", "_value", "_waiters"
- def __new__(
- cls,
- initial_value: int,
- *,
- max_value: int | None = None,
- fast_acquire: bool = False,
- ) -> Self:
- return object.__new__(cls)
- def __init__(
- self,
- initial_value: int,
- *,
- max_value: int | None = None,
- fast_acquire: bool = False,
- ):
- super().__init__(initial_value, max_value=max_value)
- self._value = initial_value
- self._max_value = max_value
- self._fast_acquire = fast_acquire
- self._waiters: deque[asyncio.Future[None]] = deque()
- async def acquire(self) -> None:
- if self._value > 0 and not self._waiters:
- await AsyncIOBackend.checkpoint_if_cancelled()
- self._value -= 1
- # Unless on the "fast path", yield control of the event loop so that other
- # tasks can run too
- if not self._fast_acquire:
- try:
- await AsyncIOBackend.cancel_shielded_checkpoint()
- except CancelledError:
- self.release()
- raise
- return
- fut: asyncio.Future[None] = asyncio.Future()
- self._waiters.append(fut)
- try:
- await fut
- except CancelledError:
- if fut.cancelled():
- try:
- self._waiters.remove(fut)
- except ValueError:
- pass
- else:
- self.release()
- raise
- def acquire_nowait(self) -> None:
- if self._value == 0:
- raise WouldBlock
- self._value -= 1
- def release(self) -> None:
- if self._max_value is not None and self._value == self._max_value:
- raise ValueError("semaphore released too many times")
- while self._waiters:
- fut = self._waiters.popleft()
- if fut.cancelled():
- continue
- fut.set_result(None)
- return
- self._value += 1
- @property
- def value(self) -> int:
- return self._value
- @property
- def max_value(self) -> int | None:
- return self._max_value
- def statistics(self) -> SemaphoreStatistics:
- return SemaphoreStatistics(len(self._waiters))
- class CapacityLimiter(BaseCapacityLimiter):
- __slots__ = "_borrowers", "_total_tokens", "_wait_queue"
- def __new__(cls, total_tokens: float) -> Self:
- return object.__new__(cls)
- def __init__(self, total_tokens: float):
- self._total_tokens: float = 0
- self._borrowers: set[Any] = set()
- self._wait_queue: OrderedDict[Any, asyncio.Event] = OrderedDict()
- self.total_tokens = total_tokens
- async def __aenter__(self) -> None:
- await self.acquire()
- async def __aexit__(
- self,
- exc_type: type[BaseException] | None,
- exc_val: BaseException | None,
- exc_tb: TracebackType | None,
- ) -> None:
- self.release()
- @property
- def total_tokens(self) -> float:
- return self._total_tokens
- @total_tokens.setter
- def total_tokens(self, value: float) -> None:
- if not isinstance(value, int) and not math.isinf(value):
- raise TypeError("total_tokens must be an int or math.inf")
- if value < 0:
- raise ValueError("total_tokens must be >= 0")
- self._total_tokens = value
- # Notify waiting tasks that they have acquired the limiter while
- # there is spare capacity.
- while self._wait_queue and len(self._borrowers) < self._total_tokens:
- borrower, event = self._wait_queue.popitem(last=False)
- self._borrowers.add(borrower)
- event.set()
- @property
- def borrowed_tokens(self) -> int:
- return len(self._borrowers)
- @property
- def available_tokens(self) -> float:
- return self._total_tokens - len(self._borrowers)
- def _notify_next_waiter(self) -> None:
- """Hand a free token to the next task in line, if any."""
- if self._wait_queue and len(self._borrowers) < self._total_tokens:
- borrower, event = self._wait_queue.popitem(last=False)
- self._borrowers.add(borrower)
- event.set()
- def acquire_nowait(self) -> None:
- self.acquire_on_behalf_of_nowait(current_task())
- def acquire_on_behalf_of_nowait(self, borrower: object) -> None:
- if borrower in self._borrowers:
- raise RuntimeError(
- "this borrower is already holding one of this CapacityLimiter's tokens"
- )
- if self._wait_queue or len(self._borrowers) >= self._total_tokens:
- raise WouldBlock
- self._borrowers.add(borrower)
- async def acquire(self) -> None:
- return await self.acquire_on_behalf_of(current_task())
- async def acquire_on_behalf_of(self, borrower: object) -> None:
- await AsyncIOBackend.checkpoint_if_cancelled()
- try:
- self.acquire_on_behalf_of_nowait(borrower)
- except WouldBlock:
- event = asyncio.Event()
- self._wait_queue[borrower] = event
- try:
- await event.wait()
- except BaseException:
- self._wait_queue.pop(borrower, None)
- if event.is_set():
- self._borrowers.discard(borrower)
- self._notify_next_waiter()
- raise
- else:
- try:
- await AsyncIOBackend.cancel_shielded_checkpoint()
- except BaseException:
- self.release()
- raise
- def release(self) -> None:
- self.release_on_behalf_of(current_task())
- def release_on_behalf_of(self, borrower: object) -> None:
- try:
- self._borrowers.remove(borrower)
- except KeyError:
- raise RuntimeError(
- "this borrower isn't holding any of this CapacityLimiter's tokens"
- ) from None
- self._notify_next_waiter()
- def statistics(self) -> CapacityLimiterStatistics:
- return CapacityLimiterStatistics(
- self.borrowed_tokens,
- self.total_tokens,
- tuple(self._borrowers),
- len(self._wait_queue),
- )
- _default_thread_limiter: RunVar[CapacityLimiter] = RunVar("_default_thread_limiter")
- #
- # Operating system signals
- #
- class _SignalReceiver:
- def __init__(self, signals: tuple[Signals, ...]):
- self._signals = signals
- self._loop = get_running_loop()
- self._signal_queue: deque[Signals] = deque()
- self._future: asyncio.Future = asyncio.Future()
- self._handled_signals: set[Signals] = set()
- def _deliver(self, signum: Signals) -> None:
- self._signal_queue.append(signum)
- if not self._future.done():
- self._future.set_result(None)
- def __enter__(self) -> Self:
- for sig in set(self._signals):
- self._loop.add_signal_handler(sig, self._deliver, sig)
- self._handled_signals.add(sig)
- return self
- def __exit__(
- self,
- exc_type: type[BaseException] | None,
- exc_val: BaseException | None,
- exc_tb: TracebackType | None,
- ) -> None:
- for sig in self._handled_signals:
- self._loop.remove_signal_handler(sig)
- def __aiter__(self) -> _SignalReceiver:
- return self
- async def __anext__(self) -> Signals:
- await AsyncIOBackend.checkpoint()
- if not self._signal_queue:
- self._future = asyncio.Future()
- await self._future
- return self._signal_queue.popleft()
- #
- # Testing and debugging
- #
- class AsyncIOTaskInfo(TaskInfo):
- def __init__(self, task: asyncio.Task):
- task_state = _task_states.get(task)
- if task_state is None:
- parent_id = None
- else:
- parent_id = task_state.parent_id
- coro = task.get_coro()
- assert coro is not None, "created TaskInfo from a completed Task"
- super().__init__(id(task), parent_id, task.get_name(), coro)
- self._task = weakref.ref(task)
- def has_pending_cancellation(self) -> bool:
- if not (task := self._task()):
- # If the task isn't around anymore, it won't have a pending cancellation
- return False
- if task._must_cancel or ( # type: ignore[attr-defined]
- isinstance(task._fut_waiter, asyncio.Future) # type: ignore[attr-defined]
- and task._fut_waiter.cancelled() # type: ignore[attr-defined]
- ):
- return True
- if task_state := _task_states.get(task):
- if cancel_scope := task_state.cancel_scope:
- return cancel_scope._effectively_cancelled
- return False
- class TestRunner(abc.TestRunner):
- _send_stream: MemoryObjectSendStream[tuple[Awaitable[Any], asyncio.Future[Any]]]
- def __init__(
- self,
- *,
- debug: bool | None = None,
- use_uvloop: bool = False,
- loop_factory: Callable[[], AbstractEventLoop] | None = None,
- ) -> None:
- if use_uvloop and loop_factory is None:
- if sys.platform != "win32":
- import uvloop
- loop_factory = uvloop.new_event_loop
- else:
- import winloop
- loop_factory = winloop.new_event_loop
- self._runner = Runner(debug=debug, loop_factory=loop_factory)
- self._exceptions: list[BaseException] = []
- self._runner_task: asyncio.Task | None = None
- def __enter__(self) -> Self:
- self._runner.__enter__()
- self.get_loop().set_exception_handler(self._exception_handler)
- return self
- def __exit__(
- self,
- exc_type: type[BaseException] | None,
- exc_val: BaseException | None,
- exc_tb: TracebackType | None,
- ) -> None:
- self._runner.__exit__(exc_type, exc_val, exc_tb)
- def get_loop(self) -> AbstractEventLoop:
- return self._runner.get_loop()
- def is_running(self) -> bool:
- try:
- asyncio.get_running_loop()
- return True
- except RuntimeError:
- return False
- def _exception_handler(
- self, loop: asyncio.AbstractEventLoop, context: dict[str, Any]
- ) -> None:
- if isinstance(context.get("exception"), Exception):
- self._exceptions.append(context["exception"])
- else:
- loop.default_exception_handler(context)
- def _raise_async_exceptions(self) -> None:
- # Re-raise any exceptions raised in asynchronous callbacks
- if self._exceptions:
- exceptions, self._exceptions = self._exceptions, []
- if len(exceptions) == 1:
- raise exceptions[0]
- elif exceptions:
- raise BaseExceptionGroup(
- "Multiple exceptions occurred in asynchronous callbacks", exceptions
- )
- async def _run_tests_and_fixtures(
- self,
- receive_stream: MemoryObjectReceiveStream[
- tuple[Awaitable[T_Retval], asyncio.Future[T_Retval]]
- ],
- ) -> None:
- from _pytest.outcomes import OutcomeException
- with receive_stream, self._send_stream:
- async for coro, future in receive_stream:
- try:
- retval = await coro
- except CancelledError as exc:
- if not future.cancelled():
- future.cancel(*exc.args)
- raise
- except BaseException as exc:
- if not future.cancelled():
- future.set_exception(exc)
- if not isinstance(exc, (Exception, OutcomeException)):
- raise
- else:
- if not future.cancelled():
- future.set_result(retval)
- async def _call_in_runner_task(
- self,
- func: Callable[P, Awaitable[T_Retval]],
- /,
- *args: P.args,
- **kwargs: P.kwargs,
- ) -> T_Retval:
- if not self._runner_task:
- self._send_stream, receive_stream = create_memory_object_stream[
- tuple[Awaitable[Any], asyncio.Future]
- ](1)
- self._runner_task = self.get_loop().create_task(
- self._run_tests_and_fixtures(receive_stream)
- )
- coro = func(*args, **kwargs)
- future: asyncio.Future[T_Retval] = self.get_loop().create_future()
- self._send_stream.send_nowait((coro, future))
- return await future
- def run_asyncgen_fixture(
- self,
- fixture_func: Callable[..., AsyncGenerator[T_Retval, Any]],
- kwargs: dict[str, Any],
- ) -> Iterable[T_Retval]:
- asyncgen = fixture_func(**kwargs)
- fixturevalue: T_Retval = self.get_loop().run_until_complete(
- self._call_in_runner_task(asyncgen.asend, None)
- )
- self._raise_async_exceptions()
- yield fixturevalue
- try:
- self.get_loop().run_until_complete(
- self._call_in_runner_task(asyncgen.asend, None)
- )
- except StopAsyncIteration:
- self._raise_async_exceptions()
- else:
- self.get_loop().run_until_complete(asyncgen.aclose())
- raise RuntimeError("Async generator fixture did not stop")
- def run_fixture(
- self,
- fixture_func: Callable[..., Coroutine[Any, Any, T_Retval]],
- kwargs: dict[str, Any],
- ) -> T_Retval:
- retval = self.get_loop().run_until_complete(
- self._call_in_runner_task(fixture_func, **kwargs)
- )
- self._raise_async_exceptions()
- return retval
- def run_test(
- self, test_func: Callable[..., Coroutine[Any, Any, Any]], kwargs: dict[str, Any]
- ) -> None:
- from _pytest.outcomes import OutcomeException
- try:
- self.get_loop().run_until_complete(
- self._call_in_runner_task(test_func, **kwargs)
- )
- except Exception as exc:
- self._exceptions.append(exc)
- except OutcomeException:
- raise
- except BaseException:
- # A BaseException (e.g. KeyboardInterrupt, SystemExit) interrupted the event loop before
- # the test completed. Cancel _runner_task so it does not resume when the event
- # loop is re-entered during async generator fixture teardown.
- if self._runner_task is not None and not self._runner_task.done():
- self._runner_task.cancel()
- self._send_stream.close()
- try:
- self.get_loop().run_until_complete(self._runner_task)
- except CancelledError:
- pass
- finally:
- self._runner_task = None
- raise
- self._raise_async_exceptions()
- class _ProcessStreamProtocol(asyncio.subprocess.SubprocessStreamProtocol):
- """
- A subprocess protocol that allows us to be notified of ``process_exited``
- asyncio's own ``Process.wait()`` only resolves once every pipe transport has
- disconnected so to get same semantics as on trio and uvloop we need this.
- """
- def __init__(self) -> None:
- # Match the standard factory for asyncio.create_process
- super().__init__(limit=2**16, loop=asyncio.get_running_loop())
- self.exited = asyncio.Event()
- def process_exited(self) -> None:
- super().process_exited()
- self.exited.set()
- class AsyncIOBackend(AsyncBackend):
- @classmethod
- def run(
- cls,
- func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval]],
- args: tuple[Unpack[PosArgsT]],
- kwargs: dict[str, Any],
- options: dict[str, Any],
- ) -> T_Retval:
- @wraps(func)
- async def wrapper() -> T_Retval:
- task = cast(asyncio.Task, current_task())
- task.set_name(get_callable_name(func))
- _task_states[task] = TaskState(None, None)
- try:
- return await func(*args)
- finally:
- del _task_states[task]
- debug = options.get("debug", None)
- loop_factory = options.get("loop_factory", None)
- if loop_factory is None and options.get("use_uvloop", False):
- if sys.platform != "win32":
- import uvloop
- loop_factory = uvloop.new_event_loop
- else:
- import winloop
- loop_factory = winloop.new_event_loop
- with Runner(debug=debug, loop_factory=loop_factory) as runner:
- return runner.run(wrapper())
- @classmethod
- def current_token(cls) -> object:
- return get_running_loop()
- @classmethod
- def current_time(cls) -> float:
- return get_running_loop().time()
- @classmethod
- def cancelled_exception_class(cls) -> type[BaseException]:
- return CancelledError
- @classmethod
- async def checkpoint(cls) -> None:
- await sleep(0)
- @classmethod
- async def checkpoint_if_cancelled(cls) -> None:
- task = current_task()
- if task is None:
- return
- try:
- cancel_scope = _task_states[task].cancel_scope
- except KeyError:
- return
- while cancel_scope:
- if cancel_scope.cancel_called:
- await sleep(0)
- elif cancel_scope.shield:
- break
- else:
- cancel_scope = cancel_scope._parent_scope
- @classmethod
- async def cancel_shielded_checkpoint(cls) -> None:
- with CancelScope(shield=True):
- await sleep(0)
- @classmethod
- async def sleep(cls, delay: float) -> None:
- await sleep(delay)
- @classmethod
- def create_cancel_scope(
- cls, *, deadline: float = math.inf, shield: bool = False
- ) -> CancelScope:
- return CancelScope(deadline=deadline, shield=shield)
- @classmethod
- def current_effective_deadline(cls) -> float:
- if (task := current_task()) is None:
- return math.inf
- try:
- cancel_scope = _task_states[task].cancel_scope
- except KeyError:
- return math.inf
- deadline = math.inf
- while cancel_scope:
- deadline = min(deadline, cancel_scope.deadline)
- if cancel_scope._cancel_called:
- deadline = -math.inf
- break
- elif cancel_scope.shield:
- break
- else:
- cancel_scope = cancel_scope._parent_scope
- return deadline
- @classmethod
- def create_task_group(cls) -> abc.TaskGroup:
- return TaskGroup()
- @classmethod
- def create_event(cls) -> BaseEvent:
- return Event()
- @classmethod
- def create_lock(cls, *, fast_acquire: bool) -> BaseLock:
- return Lock(fast_acquire=fast_acquire)
- @classmethod
- def create_semaphore(
- cls,
- initial_value: int,
- *,
- max_value: int | None = None,
- fast_acquire: bool = False,
- ) -> BaseSemaphore:
- return Semaphore(initial_value, max_value=max_value, fast_acquire=fast_acquire)
- @classmethod
- def create_capacity_limiter(cls, total_tokens: float) -> BaseCapacityLimiter:
- return CapacityLimiter(total_tokens)
- @classmethod
- async def run_sync_in_worker_thread( # type: ignore[return]
- cls,
- func: Callable[[Unpack[PosArgsT]], T_Retval],
- args: tuple[Unpack[PosArgsT]],
- abandon_on_cancel: bool = False,
- limiter: BaseCapacityLimiter | None = None,
- ) -> T_Retval:
- await cls.checkpoint()
- # If this is the first run in this event loop thread, set up the necessary
- # variables
- try:
- idle_workers = _threadpool_idle_workers.get()
- workers = _threadpool_workers.get()
- except LookupError:
- idle_workers = deque()
- workers = set()
- _threadpool_idle_workers.set(idle_workers)
- _threadpool_workers.set(workers)
- async with limiter or cls.current_default_thread_limiter():
- with CancelScope(shield=not abandon_on_cancel) as scope:
- future = asyncio.Future[T_Retval]()
- root_task = find_root_task()
- if not idle_workers:
- worker = WorkerThread(root_task, workers, idle_workers)
- worker.start()
- workers.add(worker)
- root_task.add_done_callback(worker.stop, context=Context())
- else:
- worker = idle_workers.pop()
- # Prune any other workers that have been idle for MAX_IDLE_TIME
- # seconds or longer
- now = cls.current_time()
- while idle_workers:
- if (
- now - idle_workers[0].idle_since
- < WorkerThread.MAX_IDLE_TIME
- ):
- break
- expired_worker = idle_workers.popleft()
- expired_worker.root_task.remove_done_callback(
- expired_worker.stop
- )
- expired_worker.stop()
- context = copy_context()
- context.run(set_current_async_library, None)
- if abandon_on_cancel or scope._parent_scope is None:
- worker_scope = scope
- else:
- worker_scope = scope._parent_scope
- worker.queue.put_nowait((context, func, args, future, worker_scope))
- return await future
- @classmethod
- def check_cancelled(cls) -> None:
- scope: CancelScope | None = threadlocals.current_cancel_scope
- while scope is not None:
- if scope.cancel_called:
- raise CancelledError(f"Cancelled via cancel scope {id(scope):x}")
- if scope.shield:
- return
- scope = scope._parent_scope
- @classmethod
- def run_async_from_thread(
- cls,
- func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]],
- args: tuple[Unpack[PosArgsT]],
- token: object,
- ) -> T_co:
- async def task_wrapper() -> T_co:
- __tracebackhide__ = True
- if scope is not None:
- task = cast(asyncio.Task, current_task())
- _task_states[task] = TaskState(None, scope)
- scope._tasks.add(task)
- try:
- return await func(*args)
- except CancelledError as exc:
- raise concurrent.futures.CancelledError(str(exc)) from None
- finally:
- if scope is not None:
- scope._tasks.discard(task)
- loop = cast(
- "AbstractEventLoop", token or threadlocals.current_token.native_token
- )
- if loop.is_closed():
- raise RunFinishedError
- context = copy_context()
- context.run(set_current_async_library, "asyncio")
- scope = getattr(threadlocals, "current_cancel_scope", None)
- f: concurrent.futures.Future[T_co] = context.run(
- asyncio.run_coroutine_threadsafe, task_wrapper(), loop=loop
- )
- return f.result()
- @classmethod
- def run_sync_from_thread(
- cls,
- func: Callable[[Unpack[PosArgsT]], T_Retval],
- args: tuple[Unpack[PosArgsT]],
- token: object,
- ) -> T_Retval:
- @wraps(func)
- def wrapper() -> None:
- try:
- set_current_async_library("asyncio")
- f.set_result(func(*args))
- except BaseException as exc:
- f.set_exception(exc)
- if not isinstance(exc, Exception):
- raise
- loop = cast(
- "AbstractEventLoop", token or threadlocals.current_token.native_token
- )
- if loop.is_closed():
- raise RunFinishedError
- f: concurrent.futures.Future[T_Retval] = Future()
- loop.call_soon_threadsafe(wrapper)
- return f.result()
- @classmethod
- async def open_process(
- cls,
- command: StrOrBytesPath | Sequence[StrOrBytesPath],
- *,
- stdin: int | IO[Any] | None,
- stdout: int | IO[Any] | None,
- stderr: int | IO[Any] | None,
- **kwargs: Any,
- ) -> Process:
- await cls.checkpoint()
- if isinstance(command, PathLike):
- command = os.fspath(command)
- # Use loop.subprocess_shell()/subprocess_exec() rather than their
- # asyncio.create_subprocess_*() counterparts to get access to
- # transport/protocol.
- loop = asyncio.get_running_loop()
- if isinstance(command, (str, bytes)):
- transport, protocol = await loop.subprocess_shell(
- _ProcessStreamProtocol,
- command,
- stdin=stdin,
- stdout=stdout,
- stderr=stderr,
- **kwargs,
- )
- else:
- transport, protocol = await loop.subprocess_exec(
- _ProcessStreamProtocol,
- *command,
- stdin=stdin,
- stdout=stdout,
- stderr=stderr,
- **kwargs,
- )
- process = asyncio.subprocess.Process(transport, protocol, loop)
- stdin_stream = StreamWriterWrapper(process.stdin) if process.stdin else None
- stdout_stream = StreamReaderWrapper(process.stdout) if process.stdout else None
- stderr_stream = StreamReaderWrapper(process.stderr) if process.stderr else None
- return Process(
- process,
- stdin_stream,
- stdout_stream,
- stderr_stream,
- protocol.exited,
- transport,
- )
- @classmethod
- def setup_process_pool_exit_at_shutdown(cls, workers: set[abc.Process]) -> None:
- create_task(
- _shutdown_process_pool_on_exit(workers),
- name="AnyIO process pool shutdown task",
- )
- find_root_task().add_done_callback(
- partial(_forcibly_shutdown_process_pool_on_exit, workers) # type:ignore[arg-type]
- )
- @classmethod
- async def connect_tcp(
- cls, host: str, port: int, local_address: IPSockAddrType | None = None
- ) -> abc.SocketStream:
- transport, protocol = cast(
- tuple[asyncio.Transport, StreamProtocol],
- await get_running_loop().create_connection(
- StreamProtocol, host, port, local_addr=local_address
- ),
- )
- transport.pause_reading()
- return SocketStream(transport, protocol)
- @classmethod
- async def connect_unix(cls, path: str | bytes) -> abc.UNIXSocketStream:
- await cls.checkpoint()
- loop = get_running_loop()
- raw_socket = socket.socket(socket.AF_UNIX)
- raw_socket.setblocking(False)
- while True:
- try:
- raw_socket.connect(path)
- except BlockingIOError:
- f: asyncio.Future = asyncio.Future()
- loop.add_writer(raw_socket, f.set_result, None)
- f.add_done_callback(lambda _: loop.remove_writer(raw_socket))
- await f
- except BaseException:
- raw_socket.close()
- raise
- else:
- return UNIXSocketStream(raw_socket)
- @classmethod
- def create_tcp_listener(cls, sock: socket.socket) -> SocketListener:
- return TCPSocketListener(sock)
- @classmethod
- def create_unix_listener(cls, sock: socket.socket) -> SocketListener:
- return UNIXSocketListener(sock)
- @classmethod
- async def create_udp_socket(
- cls,
- family: AddressFamily,
- local_address: IPSockAddrType | None,
- remote_address: IPSockAddrType | None,
- reuse_port: bool,
- ) -> UDPSocket | ConnectedUDPSocket:
- transport, protocol = await get_running_loop().create_datagram_endpoint(
- DatagramProtocol,
- local_addr=local_address,
- remote_addr=remote_address,
- family=family,
- reuse_port=reuse_port,
- )
- if protocol.exception:
- transport.close()
- raise protocol.exception
- if not remote_address:
- return UDPSocket(transport, protocol)
- else:
- return ConnectedUDPSocket(transport, protocol)
- @classmethod
- async def create_unix_datagram_socket( # type: ignore[override]
- cls, raw_socket: socket.socket, remote_path: str | bytes | None
- ) -> abc.UNIXDatagramSocket | abc.ConnectedUNIXDatagramSocket:
- await cls.checkpoint()
- loop = get_running_loop()
- if remote_path:
- while True:
- try:
- raw_socket.connect(remote_path)
- except BlockingIOError:
- f: asyncio.Future = asyncio.Future()
- loop.add_writer(raw_socket, f.set_result, None)
- f.add_done_callback(lambda _: loop.remove_writer(raw_socket))
- await f
- except BaseException:
- raw_socket.close()
- raise
- else:
- return ConnectedUNIXDatagramSocket(raw_socket)
- else:
- return UNIXDatagramSocket(raw_socket)
- @classmethod
- async def getaddrinfo(
- cls,
- host: bytes | str | None,
- port: str | int | None,
- *,
- family: int | AddressFamily = 0,
- type: int | SocketKind = 0,
- proto: int = 0,
- flags: int = 0,
- ) -> Sequence[
- tuple[
- AddressFamily,
- SocketKind,
- int,
- str,
- tuple[str, int] | tuple[str, int, int, int] | tuple[int, bytes],
- ]
- ]:
- return await get_running_loop().getaddrinfo(
- host, port, family=family, type=type, proto=proto, flags=flags
- )
- @classmethod
- async def getnameinfo(
- cls, sockaddr: IPSockAddrType, flags: int = 0
- ) -> tuple[str, str]:
- return await get_running_loop().getnameinfo(sockaddr, flags)
- @classmethod
- async def wait_readable(cls, obj: FileDescriptorLike) -> None:
- try:
- read_events = _read_events.get()
- except LookupError:
- read_events = {}
- _read_events.set(read_events)
- fd = obj if isinstance(obj, int) else obj.fileno()
- if read_events.get(fd):
- raise BusyResourceError("reading from")
- loop = get_running_loop()
- fut: asyncio.Future[bool] = loop.create_future()
- def cb() -> None:
- try:
- del read_events[fd]
- except KeyError:
- pass
- else:
- remove_reader(fd)
- try:
- fut.set_result(True)
- except asyncio.InvalidStateError:
- pass
- try:
- loop.add_reader(fd, cb)
- except NotImplementedError:
- from anyio._core._asyncio_selector_thread import get_selector
- selector = get_selector()
- selector.add_reader(fd, cb)
- remove_reader = selector.remove_reader
- else:
- remove_reader = loop.remove_reader
- read_events[fd] = fut
- try:
- success = await fut
- finally:
- try:
- del read_events[fd]
- except KeyError:
- pass
- else:
- remove_reader(fd)
- if not success:
- raise ClosedResourceError
- @classmethod
- async def wait_writable(cls, obj: FileDescriptorLike) -> None:
- try:
- write_events = _write_events.get()
- except LookupError:
- write_events = {}
- _write_events.set(write_events)
- fd = obj if isinstance(obj, int) else obj.fileno()
- if write_events.get(fd):
- raise BusyResourceError("writing to")
- loop = get_running_loop()
- fut: asyncio.Future[bool] = loop.create_future()
- def cb() -> None:
- try:
- del write_events[fd]
- except KeyError:
- pass
- else:
- remove_writer(fd)
- try:
- fut.set_result(True)
- except asyncio.InvalidStateError:
- pass
- try:
- loop.add_writer(fd, cb)
- except NotImplementedError:
- from anyio._core._asyncio_selector_thread import get_selector
- selector = get_selector()
- selector.add_writer(fd, cb)
- remove_writer = selector.remove_writer
- else:
- remove_writer = loop.remove_writer
- write_events[fd] = fut
- try:
- success = await fut
- finally:
- try:
- del write_events[fd]
- except KeyError:
- pass
- else:
- remove_writer(fd)
- if not success:
- raise ClosedResourceError
- @classmethod
- def notify_closing(cls, obj: FileDescriptorLike) -> None:
- fd = obj if isinstance(obj, int) else obj.fileno()
- loop = get_running_loop()
- try:
- write_events = _write_events.get()
- except LookupError:
- pass
- else:
- try:
- fut = write_events.pop(fd)
- except KeyError:
- pass
- else:
- try:
- fut.set_result(False)
- except asyncio.InvalidStateError:
- pass
- try:
- loop.remove_writer(fd)
- except NotImplementedError:
- from anyio._core._asyncio_selector_thread import get_selector
- get_selector().remove_writer(fd)
- try:
- read_events = _read_events.get()
- except LookupError:
- pass
- else:
- try:
- fut = read_events.pop(fd)
- except KeyError:
- pass
- else:
- try:
- fut.set_result(False)
- except asyncio.InvalidStateError:
- pass
- try:
- loop.remove_reader(fd)
- except NotImplementedError:
- from anyio._core._asyncio_selector_thread import get_selector
- get_selector().remove_reader(fd)
- @classmethod
- async def wrap_listener_socket(cls, sock: socket.socket) -> SocketListener:
- if hasattr(socket, "AF_UNIX") and sock.family == socket.AF_UNIX:
- return UNIXSocketListener(sock)
- return TCPSocketListener(sock)
- @classmethod
- async def wrap_stream_socket(cls, sock: socket.socket) -> SocketStream:
- transport, protocol = await get_running_loop().create_connection(
- StreamProtocol, sock=sock
- )
- return SocketStream(transport, protocol)
- @classmethod
- async def wrap_unix_stream_socket(cls, sock: socket.socket) -> UNIXSocketStream:
- return UNIXSocketStream(sock)
- @classmethod
- async def wrap_udp_socket(cls, sock: socket.socket) -> UDPSocket:
- transport, protocol = await get_running_loop().create_datagram_endpoint(
- DatagramProtocol, sock=sock
- )
- return UDPSocket(transport, protocol)
- @classmethod
- async def wrap_connected_udp_socket(cls, sock: socket.socket) -> ConnectedUDPSocket:
- transport, protocol = await get_running_loop().create_datagram_endpoint(
- DatagramProtocol, sock=sock
- )
- return ConnectedUDPSocket(transport, protocol)
- @classmethod
- async def wrap_unix_datagram_socket(cls, sock: socket.socket) -> UNIXDatagramSocket:
- return UNIXDatagramSocket(sock)
- @classmethod
- async def wrap_connected_unix_datagram_socket(
- cls, sock: socket.socket
- ) -> ConnectedUNIXDatagramSocket:
- return ConnectedUNIXDatagramSocket(sock)
- @classmethod
- def current_default_thread_limiter(cls) -> CapacityLimiter:
- try:
- return _default_thread_limiter.get()
- except LookupError:
- limiter = CapacityLimiter(40)
- _default_thread_limiter.set(limiter)
- return limiter
- @classmethod
- def open_signal_receiver(
- cls, *signals: Signals
- ) -> AbstractContextManager[AsyncIterator[Signals]]:
- return _SignalReceiver(signals)
- @classmethod
- def get_current_task(cls) -> TaskInfo:
- return AsyncIOTaskInfo(current_task()) # type: ignore[arg-type]
- @classmethod
- def get_running_tasks(cls) -> Sequence[TaskInfo]:
- return [AsyncIOTaskInfo(task) for task in all_tasks() if not task.done()]
- @classmethod
- async def wait_all_tasks_blocked(cls) -> None:
- await cls.checkpoint()
- this_task = current_task()
- while True:
- for task in all_tasks():
- if task is this_task:
- continue
- waiter = task._fut_waiter # type: ignore[attr-defined]
- if waiter is None or waiter.done():
- await sleep(0.1)
- break
- else:
- return
- @classmethod
- def create_test_runner(cls, options: dict[str, Any]) -> TestRunner:
- return TestRunner(**options)
- backend_class = AsyncIOBackend
|