_asyncio.py 102 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495969798991001011021031041051061071081091101111121131141151161171181191201211221231241251261271281291301311321331341351361371381391401411421431441451461471481491501511521531541551561571581591601611621631641651661671681691701711721731741751761771781791801811821831841851861871881891901911921931941951961971981992002012022032042052062072082092102112122132142152162172182192202212222232242252262272282292302312322332342352362372382392402412422432442452462472482492502512522532542552562572582592602612622632642652662672682692702712722732742752762772782792802812822832842852862872882892902912922932942952962972982993003013023033043053063073083093103113123133143153163173183193203213223233243253263273283293303313323333343353363373383393403413423433443453463473483493503513523533543553563573583593603613623633643653663673683693703713723733743753763773783793803813823833843853863873883893903913923933943953963973983994004014024034044054064074084094104114124134144154164174184194204214224234244254264274284294304314324334344354364374384394404414424434444454464474484494504514524534544554564574584594604614624634644654664674684694704714724734744754764774784794804814824834844854864874884894904914924934944954964974984995005015025035045055065075085095105115125135145155165175185195205215225235245255265275285295305315325335345355365375385395405415425435445455465475485495505515525535545555565575585595605615625635645655665675685695705715725735745755765775785795805815825835845855865875885895905915925935945955965975985996006016026036046056066076086096106116126136146156166176186196206216226236246256266276286296306316326336346356366376386396406416426436446456466476486496506516526536546556566576586596606616626636646656666676686696706716726736746756766776786796806816826836846856866876886896906916926936946956966976986997007017027037047057067077087097107117127137147157167177187197207217227237247257267277287297307317327337347357367377387397407417427437447457467477487497507517527537547557567577587597607617627637647657667677687697707717727737747757767777787797807817827837847857867877887897907917927937947957967977987998008018028038048058068078088098108118128138148158168178188198208218228238248258268278288298308318328338348358368378388398408418428438448458468478488498508518528538548558568578588598608618628638648658668678688698708718728738748758768778788798808818828838848858868878888898908918928938948958968978988999009019029039049059069079089099109119129139149159169179189199209219229239249259269279289299309319329339349359369379389399409419429439449459469479489499509519529539549559569579589599609619629639649659669679689699709719729739749759769779789799809819829839849859869879889899909919929939949959969979989991000100110021003100410051006100710081009101010111012101310141015101610171018101910201021102210231024102510261027102810291030103110321033103410351036103710381039104010411042104310441045104610471048104910501051105210531054105510561057105810591060106110621063106410651066106710681069107010711072107310741075107610771078107910801081108210831084108510861087108810891090109110921093109410951096109710981099110011011102110311041105110611071108110911101111111211131114111511161117111811191120112111221123112411251126112711281129113011311132113311341135113611371138113911401141114211431144114511461147114811491150115111521153115411551156115711581159116011611162116311641165116611671168116911701171117211731174117511761177117811791180118111821183118411851186118711881189119011911192119311941195119611971198119912001201120212031204120512061207120812091210121112121213121412151216121712181219122012211222122312241225122612271228122912301231123212331234123512361237123812391240124112421243124412451246124712481249125012511252125312541255125612571258125912601261126212631264126512661267126812691270127112721273127412751276127712781279128012811282128312841285128612871288128912901291129212931294129512961297129812991300130113021303130413051306130713081309131013111312131313141315131613171318131913201321132213231324132513261327132813291330133113321333133413351336133713381339134013411342134313441345134613471348134913501351135213531354135513561357135813591360136113621363136413651366136713681369137013711372137313741375137613771378137913801381138213831384138513861387138813891390139113921393139413951396139713981399140014011402140314041405140614071408140914101411141214131414141514161417141814191420142114221423142414251426142714281429143014311432143314341435143614371438143914401441144214431444144514461447144814491450145114521453145414551456145714581459146014611462146314641465146614671468146914701471147214731474147514761477147814791480148114821483148414851486148714881489149014911492149314941495149614971498149915001501150215031504150515061507150815091510151115121513151415151516151715181519152015211522152315241525152615271528152915301531153215331534153515361537153815391540154115421543154415451546154715481549155015511552155315541555155615571558155915601561156215631564156515661567156815691570157115721573157415751576157715781579158015811582158315841585158615871588158915901591159215931594159515961597159815991600160116021603160416051606160716081609161016111612161316141615161616171618161916201621162216231624162516261627162816291630163116321633163416351636163716381639164016411642164316441645164616471648164916501651165216531654165516561657165816591660166116621663166416651666166716681669167016711672167316741675167616771678167916801681168216831684168516861687168816891690169116921693169416951696169716981699170017011702170317041705170617071708170917101711171217131714171517161717171817191720172117221723172417251726172717281729173017311732173317341735173617371738173917401741174217431744174517461747174817491750175117521753175417551756175717581759176017611762176317641765176617671768176917701771177217731774177517761777177817791780178117821783178417851786178717881789179017911792179317941795179617971798179918001801180218031804180518061807180818091810181118121813181418151816181718181819182018211822182318241825182618271828182918301831183218331834183518361837183818391840184118421843184418451846184718481849185018511852185318541855185618571858185918601861186218631864186518661867186818691870187118721873187418751876187718781879188018811882188318841885188618871888188918901891189218931894189518961897189818991900190119021903190419051906190719081909191019111912191319141915191619171918191919201921192219231924192519261927192819291930193119321933193419351936193719381939194019411942194319441945194619471948194919501951195219531954195519561957195819591960196119621963196419651966196719681969197019711972197319741975197619771978197919801981198219831984198519861987198819891990199119921993199419951996199719981999200020012002200320042005200620072008200920102011201220132014201520162017201820192020202120222023202420252026202720282029203020312032203320342035203620372038203920402041204220432044204520462047204820492050205120522053205420552056205720582059206020612062206320642065206620672068206920702071207220732074207520762077207820792080208120822083208420852086208720882089209020912092209320942095209620972098209921002101210221032104210521062107210821092110211121122113211421152116211721182119212021212122212321242125212621272128212921302131213221332134213521362137213821392140214121422143214421452146214721482149215021512152215321542155215621572158215921602161216221632164216521662167216821692170217121722173217421752176217721782179218021812182218321842185218621872188218921902191219221932194219521962197219821992200220122022203220422052206220722082209221022112212221322142215221622172218221922202221222222232224222522262227222822292230223122322233223422352236223722382239224022412242224322442245224622472248224922502251225222532254225522562257225822592260226122622263226422652266226722682269227022712272227322742275227622772278227922802281228222832284228522862287228822892290229122922293229422952296229722982299230023012302230323042305230623072308230923102311231223132314231523162317231823192320232123222323232423252326232723282329233023312332233323342335233623372338233923402341234223432344234523462347234823492350235123522353235423552356235723582359236023612362236323642365236623672368236923702371237223732374237523762377237823792380238123822383238423852386238723882389239023912392239323942395239623972398239924002401240224032404240524062407240824092410241124122413241424152416241724182419242024212422242324242425242624272428242924302431243224332434243524362437243824392440244124422443244424452446244724482449245024512452245324542455245624572458245924602461246224632464246524662467246824692470247124722473247424752476247724782479248024812482248324842485248624872488248924902491249224932494249524962497249824992500250125022503250425052506250725082509251025112512251325142515251625172518251925202521252225232524252525262527252825292530253125322533253425352536253725382539254025412542254325442545254625472548254925502551255225532554255525562557255825592560256125622563256425652566256725682569257025712572257325742575257625772578257925802581258225832584258525862587258825892590259125922593259425952596259725982599260026012602260326042605260626072608260926102611261226132614261526162617261826192620262126222623262426252626262726282629263026312632263326342635263626372638263926402641264226432644264526462647264826492650265126522653265426552656265726582659266026612662266326642665266626672668266926702671267226732674267526762677267826792680268126822683268426852686268726882689269026912692269326942695269626972698269927002701270227032704270527062707270827092710271127122713271427152716271727182719272027212722272327242725272627272728272927302731273227332734273527362737273827392740274127422743274427452746274727482749275027512752275327542755275627572758275927602761276227632764276527662767276827692770277127722773277427752776277727782779278027812782278327842785278627872788278927902791279227932794279527962797279827992800280128022803280428052806280728082809281028112812281328142815281628172818281928202821282228232824282528262827282828292830283128322833283428352836283728382839284028412842284328442845284628472848284928502851285228532854285528562857285828592860286128622863286428652866286728682869287028712872287328742875287628772878287928802881288228832884288528862887288828892890289128922893289428952896289728982899290029012902290329042905290629072908290929102911291229132914291529162917291829192920292129222923292429252926292729282929293029312932293329342935293629372938293929402941294229432944294529462947294829492950295129522953295429552956295729582959296029612962296329642965296629672968296929702971297229732974297529762977297829792980298129822983298429852986298729882989299029912992299329942995299629972998299930003001300230033004300530063007300830093010301130123013301430153016301730183019302030213022302330243025302630273028302930303031303230333034303530363037303830393040304130423043304430453046304730483049305030513052305330543055305630573058305930603061306230633064306530663067306830693070307130723073307430753076307730783079308030813082308330843085308630873088308930903091309230933094309530963097309830993100310131023103310431053106310731083109311031113112311331143115311631173118311931203121312231233124312531263127312831293130313131323133313431353136
  1. from __future__ import annotations
  2. import array
  3. import asyncio
  4. import concurrent.futures
  5. import contextvars
  6. import math
  7. import os
  8. import socket
  9. import sys
  10. import threading
  11. import weakref
  12. from asyncio import (
  13. AbstractEventLoop,
  14. CancelledError,
  15. all_tasks,
  16. create_task,
  17. current_task,
  18. get_running_loop,
  19. sleep,
  20. )
  21. from asyncio.base_events import _run_until_complete_cb # type: ignore[attr-defined]
  22. from collections import OrderedDict, deque
  23. from collections.abc import (
  24. AsyncGenerator,
  25. AsyncIterator,
  26. Awaitable,
  27. Callable,
  28. Collection,
  29. Coroutine,
  30. Iterable,
  31. Sequence,
  32. )
  33. from concurrent.futures import Future
  34. from contextlib import AbstractContextManager
  35. from contextvars import Context, copy_context
  36. from dataclasses import dataclass, field
  37. from functools import partial, wraps
  38. from inspect import (
  39. CORO_RUNNING,
  40. CORO_SUSPENDED,
  41. getcoroutinestate,
  42. )
  43. from io import IOBase
  44. from os import PathLike
  45. from queue import Queue
  46. from signal import Signals
  47. from socket import AddressFamily, SocketKind
  48. from threading import Thread
  49. from types import CodeType, TracebackType
  50. from typing import (
  51. IO,
  52. TYPE_CHECKING,
  53. Any,
  54. Literal,
  55. ParamSpec,
  56. TypeVar,
  57. cast,
  58. )
  59. from weakref import WeakKeyDictionary
  60. from .. import (
  61. CapacityLimiterStatistics,
  62. EventStatistics,
  63. LockStatistics,
  64. TaskInfo,
  65. abc,
  66. )
  67. from .._core._eventloop import (
  68. claim_worker_thread,
  69. set_current_async_library,
  70. threadlocals,
  71. )
  72. from .._core._exceptions import (
  73. BrokenResourceError,
  74. BusyResourceError,
  75. ClosedResourceError,
  76. EndOfStream,
  77. RunFinishedError,
  78. WouldBlock,
  79. )
  80. from .._core._sockets import convert_ipv6_sockaddr
  81. from .._core._streams import create_memory_object_stream
  82. from .._core._synchronization import (
  83. CapacityLimiter as BaseCapacityLimiter,
  84. )
  85. from .._core._synchronization import Event as BaseEvent
  86. from .._core._synchronization import Lock as BaseLock
  87. from .._core._synchronization import (
  88. ResourceGuard,
  89. SemaphoreStatistics,
  90. )
  91. from .._core._synchronization import Semaphore as BaseSemaphore
  92. from .._core._tasks import CancelScope as BaseCancelScope
  93. from .._core._tasks import TaskHandle
  94. from ..abc import (
  95. AsyncBackend,
  96. IPSockAddrType,
  97. SocketListener,
  98. UDPPacketType,
  99. UNIXDatagramPacketType,
  100. )
  101. from ..abc._eventloop import StrOrBytesPath
  102. from ..abc._tasks import call_for_coroutine, get_callable_name
  103. from ..lowlevel import RunVar
  104. from ..streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream
  105. if TYPE_CHECKING:
  106. from _typeshed import FileDescriptorLike
  107. else:
  108. FileDescriptorLike = object
  109. if sys.version_info >= (3, 11):
  110. from asyncio import Runner
  111. from typing import TypeVarTuple, Unpack
  112. else:
  113. import contextvars
  114. import enum
  115. import signal
  116. from asyncio import coroutines, events, exceptions, tasks
  117. from exceptiongroup import BaseExceptionGroup
  118. from typing_extensions import TypeVarTuple, Unpack
  119. class _State(enum.Enum):
  120. CREATED = "created"
  121. INITIALIZED = "initialized"
  122. CLOSED = "closed"
  123. class Runner:
  124. # Copied from CPython 3.11
  125. def __init__(
  126. self,
  127. *,
  128. debug: bool | None = None,
  129. loop_factory: Callable[[], AbstractEventLoop] | None = None,
  130. ):
  131. self._state = _State.CREATED
  132. self._debug = debug
  133. self._loop_factory = loop_factory
  134. self._loop: AbstractEventLoop | None = None
  135. self._context = None
  136. self._interrupt_count = 0
  137. self._set_event_loop = False
  138. def __enter__(self) -> Runner:
  139. self._lazy_init()
  140. return self
  141. def __exit__(
  142. self,
  143. exc_type: type[BaseException] | None,
  144. exc_val: BaseException | None,
  145. exc_tb: TracebackType | None,
  146. ) -> None:
  147. self.close()
  148. def close(self) -> None:
  149. """Shutdown and close event loop."""
  150. loop = self._loop
  151. if self._state is not _State.INITIALIZED or loop is None:
  152. return
  153. try:
  154. _cancel_all_tasks(loop)
  155. loop.run_until_complete(loop.shutdown_asyncgens())
  156. if hasattr(loop, "shutdown_default_executor"):
  157. loop.run_until_complete(loop.shutdown_default_executor())
  158. else:
  159. loop.run_until_complete(_shutdown_default_executor(loop))
  160. finally:
  161. if self._set_event_loop:
  162. events.set_event_loop(None)
  163. loop.close()
  164. self._loop = None
  165. self._state = _State.CLOSED
  166. def get_loop(self) -> AbstractEventLoop:
  167. """Return embedded event loop."""
  168. self._lazy_init()
  169. return self._loop
  170. def run(self, coro: Coroutine[T_Retval], *, context=None) -> T_Retval:
  171. """Run a coroutine inside the embedded event loop."""
  172. if not coroutines.iscoroutine(coro):
  173. raise ValueError(f"a coroutine was expected, got {coro!r}")
  174. if events._get_running_loop() is not None:
  175. # fail fast with short traceback
  176. raise RuntimeError(
  177. "Runner.run() cannot be called from a running event loop"
  178. )
  179. self._lazy_init()
  180. if context is None:
  181. context = self._context
  182. task = context.run(self._loop.create_task, coro)
  183. if (
  184. threading.current_thread() is threading.main_thread()
  185. and signal.getsignal(signal.SIGINT) is signal.default_int_handler
  186. ):
  187. sigint_handler = partial(self._on_sigint, main_task=task)
  188. try:
  189. signal.signal(signal.SIGINT, sigint_handler)
  190. except ValueError:
  191. # `signal.signal` may throw if `threading.main_thread` does
  192. # not support signals (e.g. embedded interpreter with signals
  193. # not registered - see gh-91880)
  194. sigint_handler = None
  195. else:
  196. sigint_handler = None
  197. self._interrupt_count = 0
  198. try:
  199. return self._loop.run_until_complete(task)
  200. except exceptions.CancelledError:
  201. if self._interrupt_count > 0:
  202. uncancel = getattr(task, "uncancel", None)
  203. if uncancel is not None and uncancel() == 0:
  204. raise KeyboardInterrupt # noqa: B904
  205. raise # CancelledError
  206. finally:
  207. if (
  208. sigint_handler is not None
  209. and signal.getsignal(signal.SIGINT) is sigint_handler
  210. ):
  211. signal.signal(signal.SIGINT, signal.default_int_handler)
  212. def _lazy_init(self) -> None:
  213. if self._state is _State.CLOSED:
  214. raise RuntimeError("Runner is closed")
  215. if self._state is _State.INITIALIZED:
  216. return
  217. if self._loop_factory is None:
  218. self._loop = events.new_event_loop()
  219. if not self._set_event_loop:
  220. # Call set_event_loop only once to avoid calling
  221. # attach_loop multiple times on child watchers
  222. events.set_event_loop(self._loop)
  223. self._set_event_loop = True
  224. else:
  225. self._loop = self._loop_factory()
  226. if self._debug is not None:
  227. self._loop.set_debug(self._debug)
  228. self._context = contextvars.copy_context()
  229. self._state = _State.INITIALIZED
  230. def _on_sigint(self, signum, frame, main_task: asyncio.Task) -> None:
  231. self._interrupt_count += 1
  232. if self._interrupt_count == 1 and not main_task.done():
  233. main_task.cancel()
  234. # wakeup loop if it is blocked by select() with long timeout
  235. self._loop.call_soon_threadsafe(lambda: None)
  236. return
  237. raise KeyboardInterrupt()
  238. def _cancel_all_tasks(loop: AbstractEventLoop) -> None:
  239. to_cancel = tasks.all_tasks(loop)
  240. if not to_cancel:
  241. return
  242. for task in to_cancel:
  243. task.cancel()
  244. loop.run_until_complete(tasks.gather(*to_cancel, return_exceptions=True))
  245. for task in to_cancel:
  246. if task.cancelled():
  247. continue
  248. if task.exception() is not None:
  249. loop.call_exception_handler(
  250. {
  251. "message": "unhandled exception during asyncio.run() shutdown",
  252. "exception": task.exception(),
  253. "task": task,
  254. }
  255. )
  256. async def _shutdown_default_executor(loop: AbstractEventLoop) -> None:
  257. """Schedule the shutdown of the default executor."""
  258. def _do_shutdown(future: asyncio.futures.Future) -> None:
  259. try:
  260. loop._default_executor.shutdown(wait=True) # type: ignore[attr-defined]
  261. loop.call_soon_threadsafe(future.set_result, None)
  262. except Exception as ex:
  263. loop.call_soon_threadsafe(future.set_exception, ex)
  264. loop._executor_shutdown_called = True
  265. if loop._default_executor is None:
  266. return
  267. future = loop.create_future()
  268. thread = threading.Thread(target=_do_shutdown, args=(future,))
  269. thread.start()
  270. try:
  271. await future
  272. finally:
  273. thread.join()
  274. T_Retval = TypeVar("T_Retval")
  275. T_co = TypeVar("T_co", covariant=True)
  276. T_contra = TypeVar("T_contra", contravariant=True)
  277. PosArgsT = TypeVarTuple("PosArgsT")
  278. P = ParamSpec("P")
  279. _root_task: RunVar[asyncio.Task | None] = RunVar("_root_task")
  280. def find_root_task() -> asyncio.Task:
  281. root_task = _root_task.get(None)
  282. if root_task is not None and not root_task.done():
  283. return root_task
  284. # Look for a task that has been started via run_until_complete()
  285. for task in all_tasks():
  286. if task._callbacks and not task.done():
  287. callbacks = [cb for cb, context in task._callbacks]
  288. for cb in callbacks:
  289. if (
  290. cb is _run_until_complete_cb
  291. or getattr(cb, "__module__", None) == "uvloop.loop"
  292. ):
  293. _root_task.set(task)
  294. return task
  295. # Look up the topmost task in the AnyIO task tree, if possible
  296. task = cast(asyncio.Task, current_task())
  297. state = _task_states.get(task)
  298. if state:
  299. cancel_scope = state.cancel_scope
  300. while cancel_scope and cancel_scope._parent_scope is not None:
  301. cancel_scope = cancel_scope._parent_scope
  302. if cancel_scope is not None:
  303. return cast(asyncio.Task, cancel_scope._host_task)
  304. return task
  305. #
  306. # Event loop
  307. #
  308. _run_vars: WeakKeyDictionary[asyncio.AbstractEventLoop, Any] = WeakKeyDictionary()
  309. def _task_started(task: asyncio.Task) -> bool:
  310. """Return ``True`` if the task has been started and has not finished."""
  311. # The task coro should never be None here, as we never add finished tasks to the
  312. # task list
  313. coro = task.get_coro()
  314. assert coro is not None
  315. return getcoroutinestate(coro) in (CORO_RUNNING, CORO_SUSPENDED)
  316. #
  317. # Timeouts and cancellation
  318. #
  319. def is_anyio_cancellation(exc: CancelledError) -> bool:
  320. # Sometimes third party frameworks catch a CancelledError and raise a new one, so as
  321. # a workaround we have to look at the previous ones in __context__ too for a
  322. # matching cancel message
  323. while True:
  324. if (
  325. exc.args
  326. and isinstance(exc.args[0], str)
  327. and exc.args[0].startswith("Cancelled via cancel scope ")
  328. ):
  329. return True
  330. if isinstance(exc.__context__, CancelledError):
  331. exc = exc.__context__
  332. continue
  333. return False
  334. class CancelScope(BaseCancelScope):
  335. __slots__ = (
  336. "_active",
  337. "_cancel_called",
  338. "_cancel_handle",
  339. "_cancel_reason",
  340. "_cancelled_caught",
  341. "_child_scopes",
  342. "_deadline",
  343. "_host_task",
  344. "_parent_scope",
  345. "_pending_uncancellations",
  346. "_shield",
  347. "_tasks",
  348. "_timeout_handle",
  349. )
  350. def __new__(
  351. cls, *, deadline: float = math.inf, shield: bool = False
  352. ) -> CancelScope:
  353. return object.__new__(cls)
  354. def __init__(self, deadline: float = math.inf, shield: bool = False):
  355. self._deadline = deadline
  356. self._shield = shield
  357. self._parent_scope: CancelScope | None = None
  358. self._child_scopes: set[CancelScope] = set()
  359. self._cancel_called = False
  360. self._cancel_reason: str | None = None
  361. self._cancelled_caught = False
  362. self._active = False
  363. self._timeout_handle: asyncio.TimerHandle | None = None
  364. self._cancel_handle: asyncio.Handle | None = None
  365. self._tasks: set[asyncio.Task] = set()
  366. self._host_task: asyncio.Task | None = None
  367. if sys.version_info >= (3, 11):
  368. self._pending_uncancellations: int | None = 0
  369. else:
  370. self._pending_uncancellations = None
  371. def __enter__(self) -> CancelScope:
  372. if self._active:
  373. raise RuntimeError(
  374. "Each CancelScope may only be used for a single 'with' block"
  375. )
  376. self._host_task = host_task = cast(asyncio.Task, current_task())
  377. self._tasks.add(host_task)
  378. try:
  379. task_state = _task_states[host_task]
  380. except KeyError:
  381. task_state = TaskState(None, self)
  382. _task_states[host_task] = task_state
  383. else:
  384. self._parent_scope = task_state.cancel_scope
  385. task_state.cancel_scope = self
  386. if self._parent_scope is not None:
  387. # If using an eager task factory, the parent scope may not even contain
  388. # the host task
  389. self._parent_scope._child_scopes.add(self)
  390. self._parent_scope._tasks.discard(host_task)
  391. self._timeout()
  392. self._active = True
  393. # Start cancelling the host task if the scope was cancelled before entering
  394. if self._cancel_called:
  395. self._deliver_cancellation(self)
  396. return self
  397. def __exit__(
  398. self,
  399. exc_type: type[BaseException] | None,
  400. exc_val: BaseException | None,
  401. exc_tb: TracebackType | None,
  402. ) -> bool:
  403. del exc_tb
  404. if not self._active:
  405. raise RuntimeError("This cancel scope is not active")
  406. if current_task() is not self._host_task:
  407. raise RuntimeError(
  408. "Attempted to exit cancel scope in a different task than it was "
  409. "entered in"
  410. )
  411. assert self._host_task is not None
  412. host_task_state = _task_states.get(self._host_task)
  413. if host_task_state is None or host_task_state.cancel_scope is not self:
  414. raise RuntimeError(
  415. "Attempted to exit a cancel scope that isn't the current tasks's "
  416. "current cancel scope"
  417. )
  418. try:
  419. self._active = False
  420. if self._timeout_handle:
  421. self._timeout_handle.cancel()
  422. self._timeout_handle = None
  423. self._tasks.remove(self._host_task)
  424. if self._parent_scope is not None:
  425. self._parent_scope._child_scopes.remove(self)
  426. self._parent_scope._tasks.add(self._host_task)
  427. host_task_state.cancel_scope = self._parent_scope
  428. # Restart the cancellation effort in the closest visible, cancelled parent
  429. # scope if necessary
  430. self._restart_cancellation_in_parent()
  431. # We only swallow the exception iff it was an AnyIO CancelledError, either
  432. # directly as exc_val or inside an exception group and there are no cancelled
  433. # parent cancel scopes visible to us here
  434. if self._cancel_called and not self._parent_cancellation_is_visible_to_us:
  435. # For each level-cancel() call made on the host task, call uncancel()
  436. while self._pending_uncancellations:
  437. self._host_task.uncancel()
  438. self._pending_uncancellations -= 1
  439. # Update cancelled_caught and check for exceptions we must not swallow
  440. if isinstance(exc_val, BaseExceptionGroup):
  441. cancelleds_caught, remaining = exc_val.split(
  442. lambda exc: (
  443. isinstance(exc, CancelledError)
  444. and is_anyio_cancellation(exc)
  445. )
  446. )
  447. if cancelleds_caught is None:
  448. return False
  449. self._cancelled_caught = True
  450. if remaining is None:
  451. return True
  452. context = remaining.__context__
  453. try:
  454. # Preserve __cause__ and __suppress_context__ by avoiding `raise
  455. # ... from ...`
  456. raise remaining
  457. finally:
  458. # Preserve __context__
  459. remaining.__context__ = context
  460. del context
  461. else:
  462. if isinstance(exc_val, CancelledError) and is_anyio_cancellation(
  463. exc_val
  464. ):
  465. self._cancelled_caught = True
  466. return True
  467. else:
  468. return False
  469. else:
  470. if self._pending_uncancellations:
  471. assert self._parent_scope is not None
  472. assert self._parent_scope._pending_uncancellations is not None
  473. self._parent_scope._pending_uncancellations += (
  474. self._pending_uncancellations
  475. )
  476. self._pending_uncancellations = 0
  477. return False
  478. finally:
  479. self._host_task = None
  480. del exc_val
  481. @property
  482. def _effectively_cancelled(self) -> bool:
  483. cancel_scope: CancelScope | None = self
  484. while cancel_scope is not None:
  485. if cancel_scope._cancel_called:
  486. return True
  487. if cancel_scope.shield:
  488. return False
  489. cancel_scope = cancel_scope._parent_scope
  490. return False
  491. @property
  492. def _parent_cancellation_is_visible_to_us(self) -> bool:
  493. return (
  494. self._parent_scope is not None
  495. and not self.shield
  496. and self._parent_scope._effectively_cancelled
  497. )
  498. def _timeout(self) -> None:
  499. if self._deadline != math.inf:
  500. loop = get_running_loop()
  501. if loop.time() >= self._deadline:
  502. self.cancel("deadline exceeded")
  503. else:
  504. self._timeout_handle = loop.call_at(self._deadline, self._timeout)
  505. def _deliver_cancellation(self, origin: CancelScope) -> bool:
  506. """
  507. Deliver cancellation to directly contained tasks and nested cancel scopes.
  508. Schedule another run at the end if we still have tasks eligible for
  509. cancellation.
  510. :param origin: the cancel scope that originated the cancellation
  511. :return: ``True`` if the delivery needs to be retried on the next cycle
  512. """
  513. should_retry = False
  514. current = current_task()
  515. for task in self._tasks:
  516. # Always skip tasks that are already done (see issue #1111)
  517. if task.done():
  518. continue
  519. should_retry = True
  520. if task._must_cancel: # type: ignore[attr-defined]
  521. continue
  522. # The task is eligible for cancellation if it has started
  523. if task is not current and (task is self._host_task or _task_started(task)):
  524. waiter = task._fut_waiter # type: ignore[attr-defined]
  525. if not isinstance(waiter, asyncio.Future) or not waiter.done():
  526. task.cancel(origin._cancel_reason)
  527. if (
  528. task is origin._host_task
  529. and origin._pending_uncancellations is not None
  530. ):
  531. origin._pending_uncancellations += 1
  532. # Deliver cancellation to child scopes that aren't shielded or running their own
  533. # cancellation callbacks
  534. for scope in self._child_scopes:
  535. if not scope._shield and not scope.cancel_called:
  536. should_retry = scope._deliver_cancellation(origin) or should_retry
  537. # Schedule another callback if there are still tasks left
  538. if origin is self:
  539. if should_retry:
  540. self._cancel_handle = get_running_loop().call_soon(
  541. self._deliver_cancellation, origin
  542. )
  543. else:
  544. self._cancel_handle = None
  545. return should_retry
  546. def _restart_cancellation_in_parent(self) -> None:
  547. """
  548. Restart the cancellation effort in the closest directly cancelled parent scope.
  549. """
  550. scope = self._parent_scope
  551. while scope is not None:
  552. if scope._cancel_called:
  553. if scope._cancel_handle is None:
  554. scope._deliver_cancellation(scope)
  555. break
  556. # No point in looking beyond any shielded scope
  557. if scope._shield:
  558. break
  559. scope = scope._parent_scope
  560. def cancel(self, reason: str | None = None) -> None:
  561. if not self._cancel_called:
  562. if self._timeout_handle:
  563. self._timeout_handle.cancel()
  564. self._timeout_handle = None
  565. self._cancel_called = True
  566. self._cancel_reason = f"Cancelled via cancel scope {id(self):x}"
  567. if task := current_task():
  568. self._cancel_reason += f" by {task}"
  569. if reason:
  570. self._cancel_reason += f"; reason: {reason}"
  571. if self._host_task is not None:
  572. self._deliver_cancellation(self)
  573. @property
  574. def deadline(self) -> float:
  575. return self._deadline
  576. @deadline.setter
  577. def deadline(self, value: float) -> None:
  578. self._deadline = float(value)
  579. if self._timeout_handle is not None:
  580. self._timeout_handle.cancel()
  581. self._timeout_handle = None
  582. if self._active and not self._cancel_called:
  583. self._timeout()
  584. @property
  585. def cancel_called(self) -> bool:
  586. return self._cancel_called
  587. @property
  588. def cancelled_caught(self) -> bool:
  589. return self._cancelled_caught
  590. @property
  591. def shield(self) -> bool:
  592. return self._shield
  593. @shield.setter
  594. def shield(self, value: bool) -> None:
  595. if self._shield != value:
  596. self._shield = value
  597. if not value:
  598. self._restart_cancellation_in_parent()
  599. #
  600. # Task states
  601. #
  602. class TaskState:
  603. """
  604. Encapsulates auxiliary task information that cannot be added to the Task instance
  605. itself because there are no guarantees about its implementation.
  606. """
  607. __slots__ = "parent_id", "cancel_scope", "__weakref__"
  608. def __init__(self, parent_id: int | None, cancel_scope: CancelScope | None):
  609. self.parent_id = parent_id
  610. self.cancel_scope = cancel_scope
  611. _task_states: WeakKeyDictionary[asyncio.Task, TaskState] = WeakKeyDictionary()
  612. #
  613. # Task groups
  614. #
  615. class _AsyncioTaskStatus(abc.TaskStatus):
  616. def __init__(self, future: asyncio.Future, parent_id: int):
  617. self._future = future
  618. self._parent_id = parent_id
  619. def started(self, value: T_contra | None = None) -> None:
  620. try:
  621. self._future.set_result(value)
  622. except asyncio.InvalidStateError:
  623. if not self._future.cancelled():
  624. raise RuntimeError(
  625. "called 'started' twice on the same task status"
  626. ) from None
  627. task = cast(asyncio.Task, current_task())
  628. _task_states[task].parent_id = self._parent_id
  629. if sys.version_info >= (3, 12):
  630. _eager_task_factory_code: CodeType | None = asyncio.eager_task_factory.__code__
  631. else:
  632. _eager_task_factory_code = None
  633. class TaskGroup(abc.TaskGroup):
  634. def __init__(self) -> None:
  635. self.cancel_scope: CancelScope = CancelScope()
  636. self._entered = False
  637. self._exceptions: list[BaseException] = []
  638. self._tasks: set[asyncio.Task] = set()
  639. self._on_completed_fut: asyncio.Future[None] | None = None
  640. async def __aenter__(self) -> TaskGroup:
  641. if self._entered:
  642. raise RuntimeError("TaskGroup cannot be entered more than once")
  643. self._entered = True
  644. self.cancel_scope.__enter__()
  645. return self
  646. async def __aexit__(
  647. self,
  648. exc_type: type[BaseException] | None,
  649. exc_val: BaseException | None,
  650. exc_tb: TracebackType | None,
  651. ) -> bool:
  652. try:
  653. if exc_val is not None:
  654. self.cancel_scope.cancel()
  655. if not isinstance(exc_val, CancelledError):
  656. self._exceptions.append(exc_val)
  657. loop = get_running_loop()
  658. try:
  659. if self._tasks:
  660. with CancelScope() as wait_scope:
  661. while self._tasks:
  662. self._on_completed_fut = loop.create_future()
  663. try:
  664. await self._on_completed_fut
  665. except CancelledError as exc:
  666. # Shield the scope against further cancellation attempts,
  667. # as they're not productive (#695)
  668. wait_scope.shield = True
  669. self.cancel_scope.cancel()
  670. # Set exc_val from the cancellation exception if it was
  671. # previously unset. However, we should not replace a native
  672. # cancellation exception with one raise by a cancel scope.
  673. if exc_val is None or (
  674. isinstance(exc_val, CancelledError)
  675. and not is_anyio_cancellation(exc)
  676. ):
  677. exc_val = exc
  678. self._on_completed_fut = None
  679. else:
  680. # If there are no child tasks to wait on, run at least one checkpoint
  681. # anyway
  682. await AsyncIOBackend.cancel_shielded_checkpoint()
  683. if self._exceptions:
  684. # The exception that got us here should already have been
  685. # added to self._exceptions so it's ok to break exception
  686. # chaining and avoid adding a "During handling of above..."
  687. # for each nesting level.
  688. raise BaseExceptionGroup(
  689. "unhandled errors in a TaskGroup", self._exceptions
  690. ) from None
  691. elif exc_val:
  692. raise exc_val
  693. except BaseException as exc:
  694. if self.cancel_scope.__exit__(type(exc), exc, exc.__traceback__):
  695. return True
  696. raise
  697. return self.cancel_scope.__exit__(exc_type, exc_val, exc_tb)
  698. finally:
  699. del exc_val, exc_tb, self._exceptions
  700. def _spawn(
  701. self,
  702. coro: Coroutine[Any, Any, T_co],
  703. name: object,
  704. task_status_future: asyncio.Future | None = None,
  705. ) -> TaskHandle[T_co]:
  706. def task_done(_task: asyncio.Task) -> None:
  707. if sys.version_info >= (3, 14) and self.cancel_scope._host_task is not None:
  708. asyncio.future_discard_from_awaited_by(
  709. _task, self.cancel_scope._host_task
  710. )
  711. task_state = _task_states[_task]
  712. assert task_state.cancel_scope is not None
  713. assert _task in task_state.cancel_scope._tasks
  714. task_state.cancel_scope._tasks.remove(_task)
  715. self._tasks.remove(task)
  716. del _task_states[_task]
  717. if self._on_completed_fut is not None and not self._tasks:
  718. try:
  719. self._on_completed_fut.set_result(None)
  720. except asyncio.InvalidStateError:
  721. pass
  722. try:
  723. exc = _task.exception()
  724. except CancelledError as e:
  725. while isinstance(e.__context__, CancelledError):
  726. e = e.__context__
  727. exc = e
  728. if exc is not None:
  729. # The future can only be in the cancelled state if the host task was
  730. # cancelled, so return immediately instead of adding one more
  731. # CancelledError to the exceptions list
  732. if task_status_future is not None and task_status_future.cancelled():
  733. return
  734. if task_status_future is None or task_status_future.done():
  735. if not isinstance(exc, CancelledError):
  736. self._exceptions.append(exc)
  737. if not self.cancel_scope._effectively_cancelled:
  738. self.cancel_scope.cancel()
  739. else:
  740. task_status_future.set_exception(exc)
  741. elif task_status_future is not None and not task_status_future.done():
  742. task_status_future.set_exception(
  743. RuntimeError("Child exited without calling task_status.started()")
  744. )
  745. if task_status_future:
  746. parent_id = id(current_task())
  747. else:
  748. parent_id = id(self.cancel_scope._host_task)
  749. handle = TaskHandle(coro, name)
  750. loop = asyncio.get_running_loop()
  751. wrapper_coro = handle._run_coro()
  752. if (
  753. (factory := loop.get_task_factory())
  754. and getattr(factory, "__code__", None) is _eager_task_factory_code
  755. and (closure := getattr(factory, "__closure__", None))
  756. ):
  757. custom_task_constructor = closure[0].cell_contents
  758. task = custom_task_constructor(wrapper_coro, loop=loop, name=handle.name)
  759. else:
  760. task = loop.create_task(wrapper_coro, name=handle.name)
  761. # Make the spawned task inherit the task group's cancel scope
  762. _task_states[task] = TaskState(
  763. parent_id=parent_id, cancel_scope=self.cancel_scope
  764. )
  765. self.cancel_scope._tasks.add(task)
  766. self._tasks.add(task)
  767. if sys.version_info >= (3, 14) and self.cancel_scope._host_task is not None:
  768. asyncio.future_add_to_awaited_by(task, self.cancel_scope._host_task)
  769. task.add_done_callback(task_done)
  770. return handle
  771. def create_task(
  772. self,
  773. coro: Coroutine[Any, Any, T_co],
  774. *,
  775. name: object = None,
  776. context: Context | None = None,
  777. ) -> TaskHandle[T_co]:
  778. if not isinstance(coro, Coroutine):
  779. raise TypeError(f"expected a coroutine, got {coro.__class__.__qualname__}")
  780. if not self._entered or not self.cancel_scope._active:
  781. coro.close()
  782. raise RuntimeError(
  783. "This task group is not active; no new tasks can be started."
  784. )
  785. if context is not None:
  786. return context.run(self._spawn, coro, name=name)
  787. else:
  788. return self._spawn(coro, name=name)
  789. async def start(
  790. self,
  791. func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]],
  792. *args: Unpack[PosArgsT],
  793. name: object = None,
  794. return_handle: Literal[False] | Literal[True] = False,
  795. ) -> Any:
  796. if not self._entered or not self.cancel_scope._active:
  797. raise RuntimeError(
  798. "This task group is not active; no new tasks can be started."
  799. )
  800. future: asyncio.Future = asyncio.Future()
  801. final_name = get_callable_name(func, name)
  802. task_status = _AsyncioTaskStatus(future, id(self.cancel_scope._host_task))
  803. coro = call_for_coroutine(func, args, task_status=task_status)
  804. handle = self._spawn(coro, final_name, future)
  805. # If the task raises an exception after sending a start value without a switch
  806. # point between, the task group is cancelled and this method never proceeds to
  807. # process the completed future. That's why we have to have a shielded cancel
  808. # scope here.
  809. try:
  810. await future
  811. except BaseException:
  812. if handle.status is TaskHandle.Status.PENDING:
  813. # Cancel the task and wait for it to exit before returning
  814. handle.cancel()
  815. with CancelScope(shield=True):
  816. await handle.wait()
  817. raise
  818. if return_handle:
  819. handle._start_value = future.result()
  820. return handle
  821. else:
  822. return future.result()
  823. #
  824. # Threads
  825. #
  826. _Retval_Queue_Type = tuple[T_Retval | None, BaseException | None]
  827. class WorkerThread(Thread):
  828. MAX_IDLE_TIME = 10 # seconds
  829. def __init__(
  830. self,
  831. root_task: asyncio.Task,
  832. workers: set[WorkerThread],
  833. idle_workers: deque[WorkerThread],
  834. ):
  835. super().__init__(name="AnyIO worker thread")
  836. self.root_task = root_task
  837. self.workers = workers
  838. self.idle_workers = idle_workers
  839. self.loop = root_task._loop
  840. self.queue: Queue[
  841. tuple[Context, Callable, tuple, asyncio.Future, CancelScope] | None
  842. ] = Queue(2)
  843. self.idle_since = AsyncIOBackend.current_time()
  844. self.stopping = False
  845. def _report_result(
  846. self, future: asyncio.Future, result: Any, exc: BaseException | None
  847. ) -> None:
  848. self.idle_since = AsyncIOBackend.current_time()
  849. if not self.stopping:
  850. self.idle_workers.append(self)
  851. if not future.cancelled():
  852. if exc is not None:
  853. if isinstance(exc, StopIteration):
  854. new_exc = RuntimeError("coroutine raised StopIteration")
  855. new_exc.__cause__ = exc
  856. exc = new_exc
  857. future.set_exception(exc)
  858. else:
  859. future.set_result(result)
  860. def run(self) -> None:
  861. with claim_worker_thread(AsyncIOBackend, self.loop):
  862. while True:
  863. item = self.queue.get()
  864. if item is None:
  865. # Shutdown command received
  866. return
  867. context, func, args, future, cancel_scope = item
  868. if not future.cancelled():
  869. result = None
  870. exception: BaseException | None = None
  871. threadlocals.current_cancel_scope = cancel_scope
  872. try:
  873. result = context.run(func, *args)
  874. except BaseException as exc:
  875. exception = exc
  876. finally:
  877. del threadlocals.current_cancel_scope
  878. if not self.loop.is_closed():
  879. self.loop.call_soon_threadsafe(
  880. self._report_result, future, result, exception
  881. )
  882. del result, exception
  883. self.queue.task_done()
  884. del item, context, func, args, future, cancel_scope
  885. def stop(self, f: asyncio.Task | None = None) -> None:
  886. self.stopping = True
  887. self.queue.put_nowait(None)
  888. self.workers.discard(self)
  889. try:
  890. self.idle_workers.remove(self)
  891. except ValueError:
  892. pass
  893. _threadpool_idle_workers: RunVar[deque[WorkerThread]] = RunVar(
  894. "_threadpool_idle_workers"
  895. )
  896. _threadpool_workers: RunVar[set[WorkerThread]] = RunVar("_threadpool_workers")
  897. #
  898. # Subprocesses
  899. #
  900. @dataclass(eq=False)
  901. class StreamReaderWrapper(abc.ByteReceiveStream):
  902. _stream: asyncio.StreamReader
  903. async def receive(self, max_bytes: int = 65536) -> bytes:
  904. if max_bytes < 1:
  905. raise ValueError("max_bytes must be a positive integer")
  906. data = await self._stream.read(max_bytes)
  907. if data:
  908. return data
  909. else:
  910. raise EndOfStream
  911. async def aclose(self) -> None:
  912. self._stream.set_exception(ClosedResourceError())
  913. await AsyncIOBackend.checkpoint()
  914. @dataclass(eq=False)
  915. class StreamWriterWrapper(abc.ByteSendStream):
  916. _stream: asyncio.StreamWriter
  917. _closed: bool = field(init=False, default=False)
  918. async def send(self, item: bytes) -> None:
  919. await AsyncIOBackend.checkpoint_if_cancelled()
  920. stream_paused = self._stream._protocol._paused # type: ignore[attr-defined]
  921. try:
  922. self._stream.write(item)
  923. await self._stream.drain()
  924. except (ConnectionResetError, BrokenPipeError, RuntimeError) as exc:
  925. # If closed by us and/or the peer:
  926. # * on stdlib, drain() raises ConnectionResetError or BrokenPipeError
  927. # * on uvloop and Winloop, write() eventually starts raising RuntimeError
  928. if self._closed:
  929. raise ClosedResourceError from exc
  930. elif self._stream.is_closing():
  931. raise BrokenResourceError from exc
  932. raise
  933. if not stream_paused:
  934. await AsyncIOBackend.cancel_shielded_checkpoint()
  935. async def aclose(self) -> None:
  936. self._closed = True
  937. self._stream.close()
  938. await AsyncIOBackend.checkpoint()
  939. @dataclass(eq=False)
  940. class Process(abc.Process):
  941. _process: asyncio.subprocess.Process
  942. _stdin: StreamWriterWrapper | None
  943. _stdout: StreamReaderWrapper | None
  944. _stderr: StreamReaderWrapper | None
  945. _exited: asyncio.Event
  946. _transport: asyncio.SubprocessTransport
  947. async def aclose(self) -> None:
  948. with CancelScope(shield=True) as scope:
  949. # We need to close the underlying pipe_transports as well to allow a
  950. # process blocking on full buffers to receive SIGPIPE and exit.
  951. if self._stdin:
  952. await self._stdin.aclose()
  953. if pipe := self._transport.get_pipe_transport(0):
  954. pipe.close()
  955. if self._stdout:
  956. await self._stdout.aclose()
  957. if pipe := self._transport.get_pipe_transport(1):
  958. pipe.close()
  959. if self._stderr:
  960. await self._stderr.aclose()
  961. if pipe := self._transport.get_pipe_transport(2):
  962. pipe.close()
  963. scope.shield = False
  964. try:
  965. await self.wait()
  966. except BaseException:
  967. scope.shield = True
  968. # Closing the transport on asyncio also handles sending kill
  969. self._transport.close()
  970. await self.wait()
  971. raise
  972. async def wait(self) -> int:
  973. await self._exited.wait()
  974. assert self._process.returncode is not None
  975. return self._process.returncode
  976. def terminate(self) -> None:
  977. self._process.terminate()
  978. def kill(self) -> None:
  979. self._process.kill()
  980. def send_signal(self, signal: int) -> None:
  981. self._process.send_signal(signal)
  982. @property
  983. def pid(self) -> int:
  984. return self._process.pid
  985. @property
  986. def returncode(self) -> int | None:
  987. return self._process.returncode
  988. @property
  989. def stdin(self) -> abc.ByteSendStream | None:
  990. return self._stdin
  991. @property
  992. def stdout(self) -> abc.ByteReceiveStream | None:
  993. return self._stdout
  994. @property
  995. def stderr(self) -> abc.ByteReceiveStream | None:
  996. return self._stderr
  997. def _forcibly_shutdown_process_pool_on_exit(
  998. workers: set[Process], _task: object
  999. ) -> None:
  1000. """
  1001. Forcibly shuts down worker processes belonging to this event loop."""
  1002. child_watcher: asyncio.AbstractChildWatcher | None = None # type: ignore[name-defined]
  1003. if sys.version_info < (3, 12):
  1004. try:
  1005. child_watcher = asyncio.get_event_loop_policy().get_child_watcher()
  1006. except NotImplementedError:
  1007. pass
  1008. # Close as much as possible (w/o async/await) to avoid warnings
  1009. for process in workers.copy():
  1010. if process.returncode is not None:
  1011. continue
  1012. process._stdin._stream._transport.close() # type: ignore[union-attr]
  1013. process._stdout._stream._transport.close() # type: ignore[union-attr]
  1014. process._stderr._stream._transport.close() # type: ignore[union-attr]
  1015. process.kill()
  1016. if child_watcher:
  1017. child_watcher.remove_child_handler(process.pid)
  1018. async def _shutdown_process_pool_on_exit(workers: set[abc.Process]) -> None:
  1019. """
  1020. Shuts down worker processes belonging to this event loop.
  1021. NOTE: this only works when the event loop was started using asyncio.run() or
  1022. anyio.run().
  1023. """
  1024. process: abc.Process
  1025. try:
  1026. await sleep(math.inf)
  1027. except asyncio.CancelledError:
  1028. workers = workers.copy()
  1029. for process in workers:
  1030. if process.returncode is None:
  1031. process.kill()
  1032. for process in workers:
  1033. await process.aclose()
  1034. #
  1035. # Sockets and networking
  1036. #
  1037. class StreamProtocol(asyncio.Protocol):
  1038. read_queue: deque[bytes]
  1039. read_event: asyncio.Event
  1040. write_event: asyncio.Event
  1041. exception: Exception | None = None
  1042. is_at_eof: bool = False
  1043. def connection_made(self, transport: asyncio.BaseTransport) -> None:
  1044. self.read_queue = deque()
  1045. self.read_event = asyncio.Event()
  1046. self.write_event = asyncio.Event()
  1047. self.write_event.set()
  1048. cast(asyncio.Transport, transport).set_write_buffer_limits(0)
  1049. def connection_lost(self, exc: Exception | None) -> None:
  1050. if exc:
  1051. self.exception = exc
  1052. self.read_event.set()
  1053. self.write_event.set()
  1054. def data_received(self, data: bytes) -> None:
  1055. # ProactorEventloop sometimes sends bytearray instead of bytes
  1056. self.read_queue.append(bytes(data))
  1057. self.read_event.set()
  1058. def eof_received(self) -> bool | None:
  1059. self.is_at_eof = True
  1060. self.read_event.set()
  1061. return True
  1062. def pause_writing(self) -> None:
  1063. self.write_event = asyncio.Event()
  1064. def resume_writing(self) -> None:
  1065. self.write_event.set()
  1066. class DatagramProtocol(asyncio.DatagramProtocol):
  1067. read_queue: deque[tuple[bytes, IPSockAddrType]]
  1068. read_event: asyncio.Event
  1069. write_event: asyncio.Event
  1070. closed_event: asyncio.Event
  1071. exception: Exception | None = None
  1072. def connection_made(self, transport: asyncio.BaseTransport) -> None:
  1073. self.read_queue = deque(maxlen=100) # arbitrary value
  1074. self.read_event = asyncio.Event()
  1075. self.write_event = asyncio.Event()
  1076. self.closed_event = asyncio.Event()
  1077. self.write_event.set()
  1078. def connection_lost(self, exc: Exception | None) -> None:
  1079. self.read_event.set()
  1080. self.write_event.set()
  1081. self.closed_event.set()
  1082. def datagram_received(self, data: bytes, addr: IPSockAddrType) -> None:
  1083. addr = convert_ipv6_sockaddr(addr)
  1084. self.read_queue.append((data, addr))
  1085. self.read_event.set()
  1086. def error_received(self, exc: Exception) -> None:
  1087. self.exception = exc
  1088. def pause_writing(self) -> None:
  1089. self.write_event.clear()
  1090. def resume_writing(self) -> None:
  1091. self.write_event.set()
  1092. class SocketStream(abc.SocketStream):
  1093. def __init__(self, transport: asyncio.Transport, protocol: StreamProtocol):
  1094. self._transport = transport
  1095. self._protocol = protocol
  1096. self._receive_guard = ResourceGuard("reading from")
  1097. self._send_guard = ResourceGuard("writing to")
  1098. self._closed = False
  1099. @property
  1100. def _raw_socket(self) -> socket.socket:
  1101. return self._transport.get_extra_info("socket")
  1102. async def receive(self, max_bytes: int = 65536) -> bytes:
  1103. if max_bytes < 1:
  1104. raise ValueError("max_bytes must be a positive integer")
  1105. with self._receive_guard:
  1106. if (
  1107. not self._protocol.read_event.is_set()
  1108. and not self._transport.is_closing()
  1109. and not self._protocol.is_at_eof
  1110. ):
  1111. self._transport.resume_reading()
  1112. await self._protocol.read_event.wait()
  1113. self._transport.pause_reading()
  1114. else:
  1115. await AsyncIOBackend.checkpoint()
  1116. try:
  1117. chunk = self._protocol.read_queue.popleft()
  1118. except IndexError:
  1119. if self._closed:
  1120. raise ClosedResourceError from None
  1121. elif self._protocol.exception:
  1122. raise BrokenResourceError from self._protocol.exception
  1123. else:
  1124. raise EndOfStream from None
  1125. if len(chunk) > max_bytes:
  1126. # Split the oversized chunk
  1127. chunk, leftover = chunk[:max_bytes], chunk[max_bytes:]
  1128. self._protocol.read_queue.appendleft(leftover)
  1129. # If the read queue is empty, clear the flag so that the next call will
  1130. # block until data is available
  1131. if not self._protocol.read_queue:
  1132. self._protocol.read_event.clear()
  1133. return chunk
  1134. async def send(self, item: bytes) -> None:
  1135. with self._send_guard:
  1136. await AsyncIOBackend.checkpoint()
  1137. if self._closed:
  1138. raise ClosedResourceError
  1139. elif self._protocol.exception is not None:
  1140. raise BrokenResourceError from self._protocol.exception
  1141. try:
  1142. self._transport.write(item)
  1143. except RuntimeError as exc:
  1144. if self._transport.is_closing():
  1145. raise BrokenResourceError from exc
  1146. else:
  1147. raise
  1148. await self._protocol.write_event.wait()
  1149. async def send_eof(self) -> None:
  1150. try:
  1151. self._transport.write_eof()
  1152. except OSError:
  1153. pass
  1154. async def aclose(self) -> None:
  1155. self._closed = True
  1156. if not self._transport.is_closing():
  1157. try:
  1158. self._transport.write_eof()
  1159. except OSError:
  1160. pass
  1161. self._transport.close()
  1162. await sleep(0)
  1163. self._transport.abort()
  1164. class _RawSocketMixin:
  1165. _receive_future: asyncio.Future | None = None
  1166. _send_future: asyncio.Future | None = None
  1167. _closing = False
  1168. def __init__(self, raw_socket: socket.socket):
  1169. self.__raw_socket = raw_socket
  1170. self._receive_guard = ResourceGuard("reading from")
  1171. self._send_guard = ResourceGuard("writing to")
  1172. @property
  1173. def _raw_socket(self) -> socket.socket:
  1174. return self.__raw_socket
  1175. def _wait_until_readable(self, loop: asyncio.AbstractEventLoop) -> asyncio.Future:
  1176. def callback(f: object) -> None:
  1177. del self._receive_future
  1178. loop.remove_reader(self.__raw_socket)
  1179. f = self._receive_future = asyncio.Future()
  1180. loop.add_reader(self.__raw_socket, f.set_result, None)
  1181. f.add_done_callback(callback)
  1182. return f
  1183. def _wait_until_writable(self, loop: asyncio.AbstractEventLoop) -> asyncio.Future:
  1184. def callback(f: object) -> None:
  1185. del self._send_future
  1186. loop.remove_writer(self.__raw_socket)
  1187. f = self._send_future = asyncio.Future()
  1188. loop.add_writer(self.__raw_socket, f.set_result, None)
  1189. f.add_done_callback(callback)
  1190. return f
  1191. async def aclose(self) -> None:
  1192. if not self._closing:
  1193. self._closing = True
  1194. if self.__raw_socket.fileno() != -1:
  1195. self.__raw_socket.close()
  1196. if self._receive_future:
  1197. self._receive_future.set_result(None)
  1198. if self._send_future:
  1199. self._send_future.set_result(None)
  1200. class UNIXSocketStream(_RawSocketMixin, abc.UNIXSocketStream):
  1201. async def send_eof(self) -> None:
  1202. with self._send_guard:
  1203. self._raw_socket.shutdown(socket.SHUT_WR)
  1204. async def receive(self, max_bytes: int = 65536) -> bytes:
  1205. if max_bytes < 1:
  1206. raise ValueError("max_bytes must be a positive integer")
  1207. loop = get_running_loop()
  1208. await AsyncIOBackend.checkpoint()
  1209. with self._receive_guard:
  1210. while True:
  1211. try:
  1212. data = self._raw_socket.recv(max_bytes)
  1213. except BlockingIOError:
  1214. await self._wait_until_readable(loop)
  1215. except OSError as exc:
  1216. if self._closing:
  1217. raise ClosedResourceError from None
  1218. else:
  1219. raise BrokenResourceError from exc
  1220. else:
  1221. if not data:
  1222. raise EndOfStream
  1223. return data
  1224. async def send(self, item: bytes) -> None:
  1225. loop = get_running_loop()
  1226. await AsyncIOBackend.checkpoint()
  1227. with self._send_guard:
  1228. view = memoryview(item)
  1229. while view:
  1230. try:
  1231. bytes_sent = self._raw_socket.send(view)
  1232. except BlockingIOError:
  1233. await self._wait_until_writable(loop)
  1234. except OSError as exc:
  1235. if self._closing:
  1236. raise ClosedResourceError from None
  1237. else:
  1238. raise BrokenResourceError from exc
  1239. else:
  1240. view = view[bytes_sent:]
  1241. async def receive_fds(self, msglen: int, maxfds: int) -> tuple[bytes, list[int]]:
  1242. if not isinstance(msglen, int) or msglen < 0:
  1243. raise ValueError("msglen must be a non-negative integer")
  1244. if not isinstance(maxfds, int) or maxfds < 1:
  1245. raise ValueError("maxfds must be a positive integer")
  1246. loop = get_running_loop()
  1247. fds = array.array("i")
  1248. await AsyncIOBackend.checkpoint()
  1249. with self._receive_guard:
  1250. while True:
  1251. try:
  1252. message, ancdata, flags, addr = self._raw_socket.recvmsg(
  1253. msglen, socket.CMSG_LEN(maxfds * fds.itemsize)
  1254. )
  1255. except BlockingIOError:
  1256. await self._wait_until_readable(loop)
  1257. except OSError as exc:
  1258. if self._closing:
  1259. raise ClosedResourceError from None
  1260. else:
  1261. raise BrokenResourceError from exc
  1262. else:
  1263. if not message and not ancdata:
  1264. raise EndOfStream
  1265. break
  1266. for cmsg_level, cmsg_type, cmsg_data in ancdata:
  1267. if cmsg_level != socket.SOL_SOCKET or cmsg_type != socket.SCM_RIGHTS:
  1268. raise RuntimeError(
  1269. f"Received unexpected ancillary data; message = {message!r}, "
  1270. f"cmsg_level = {cmsg_level}, cmsg_type = {cmsg_type}"
  1271. )
  1272. fds.frombytes(cmsg_data[: len(cmsg_data) - (len(cmsg_data) % fds.itemsize)])
  1273. return message, list(fds)
  1274. async def send_fds(self, message: bytes, fds: Collection[int | IOBase]) -> None:
  1275. if not message:
  1276. raise ValueError("message must not be empty")
  1277. if not fds:
  1278. raise ValueError("fds must not be empty")
  1279. loop = get_running_loop()
  1280. filenos: list[int] = []
  1281. for fd in fds:
  1282. if isinstance(fd, int):
  1283. filenos.append(fd)
  1284. elif isinstance(fd, IOBase):
  1285. filenos.append(fd.fileno())
  1286. fdarray = array.array("i", filenos)
  1287. await AsyncIOBackend.checkpoint()
  1288. with self._send_guard:
  1289. while True:
  1290. try:
  1291. # The ignore can be removed after mypy picks up
  1292. # https://github.com/python/typeshed/pull/5545
  1293. self._raw_socket.sendmsg(
  1294. [message], [(socket.SOL_SOCKET, socket.SCM_RIGHTS, fdarray)]
  1295. )
  1296. break
  1297. except BlockingIOError:
  1298. await self._wait_until_writable(loop)
  1299. except OSError as exc:
  1300. if self._closing:
  1301. raise ClosedResourceError from None
  1302. else:
  1303. raise BrokenResourceError from exc
  1304. class TCPSocketListener(abc.SocketListener):
  1305. _accept_scope: CancelScope | None = None
  1306. _closed = False
  1307. def __init__(self, raw_socket: socket.socket):
  1308. self.__raw_socket = raw_socket
  1309. self._loop = cast(asyncio.BaseEventLoop, get_running_loop())
  1310. self._accept_guard = ResourceGuard("accepting connections from")
  1311. @property
  1312. def _raw_socket(self) -> socket.socket:
  1313. return self.__raw_socket
  1314. async def accept(self) -> abc.SocketStream:
  1315. if self._closed:
  1316. raise ClosedResourceError
  1317. with self._accept_guard:
  1318. await AsyncIOBackend.checkpoint()
  1319. with CancelScope() as self._accept_scope:
  1320. try:
  1321. client_sock, _addr = await self._loop.sock_accept(self._raw_socket)
  1322. except asyncio.CancelledError:
  1323. # Workaround for https://bugs.python.org/issue41317
  1324. try:
  1325. self._loop.remove_reader(self._raw_socket)
  1326. except (ValueError, NotImplementedError):
  1327. pass
  1328. if self._closed:
  1329. raise ClosedResourceError from None
  1330. raise
  1331. finally:
  1332. self._accept_scope = None
  1333. client_sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
  1334. transport, protocol = await self._loop.connect_accepted_socket(
  1335. StreamProtocol, client_sock
  1336. )
  1337. return SocketStream(transport, protocol)
  1338. async def aclose(self) -> None:
  1339. if self._closed:
  1340. return
  1341. self._closed = True
  1342. if self._accept_scope:
  1343. # Workaround for https://bugs.python.org/issue41317
  1344. try:
  1345. self._loop.remove_reader(self._raw_socket)
  1346. except (ValueError, NotImplementedError):
  1347. pass
  1348. self._accept_scope.cancel()
  1349. await sleep(0)
  1350. self._raw_socket.close()
  1351. class UNIXSocketListener(abc.SocketListener):
  1352. def __init__(self, raw_socket: socket.socket):
  1353. self.__raw_socket = raw_socket
  1354. self._loop = get_running_loop()
  1355. self._accept_guard = ResourceGuard("accepting connections from")
  1356. self._closed = False
  1357. async def accept(self) -> abc.SocketStream:
  1358. await AsyncIOBackend.checkpoint()
  1359. with self._accept_guard:
  1360. while True:
  1361. try:
  1362. client_sock, _ = self.__raw_socket.accept()
  1363. client_sock.setblocking(False)
  1364. return UNIXSocketStream(client_sock)
  1365. except BlockingIOError:
  1366. f: asyncio.Future = asyncio.Future()
  1367. self._loop.add_reader(self.__raw_socket, f.set_result, None)
  1368. f.add_done_callback(
  1369. lambda _: self._loop.remove_reader(self.__raw_socket)
  1370. )
  1371. await f
  1372. except OSError as exc:
  1373. if self._closed:
  1374. raise ClosedResourceError from None
  1375. else:
  1376. raise BrokenResourceError from exc
  1377. async def aclose(self) -> None:
  1378. self._closed = True
  1379. self.__raw_socket.close()
  1380. @property
  1381. def _raw_socket(self) -> socket.socket:
  1382. return self.__raw_socket
  1383. class UDPSocket(abc.UDPSocket):
  1384. def __init__(
  1385. self, transport: asyncio.DatagramTransport, protocol: DatagramProtocol
  1386. ):
  1387. self._transport = transport
  1388. self._protocol = protocol
  1389. self._receive_guard = ResourceGuard("reading from")
  1390. self._send_guard = ResourceGuard("writing to")
  1391. self._closed = False
  1392. @property
  1393. def _raw_socket(self) -> socket.socket:
  1394. return self._transport.get_extra_info("socket")
  1395. async def aclose(self) -> None:
  1396. self._closed = True
  1397. if not self._transport.is_closing():
  1398. self._transport.close()
  1399. await self._protocol.closed_event.wait()
  1400. async def receive(self) -> tuple[bytes, IPSockAddrType]:
  1401. with self._receive_guard:
  1402. await AsyncIOBackend.checkpoint()
  1403. # If the buffer is empty, ask for more data
  1404. if not self._protocol.read_queue and not self._transport.is_closing():
  1405. self._protocol.read_event.clear()
  1406. await self._protocol.read_event.wait()
  1407. try:
  1408. return self._protocol.read_queue.popleft()
  1409. except IndexError:
  1410. if self._closed:
  1411. raise ClosedResourceError from None
  1412. else:
  1413. raise BrokenResourceError from None
  1414. async def send(self, item: UDPPacketType) -> None:
  1415. with self._send_guard:
  1416. await AsyncIOBackend.checkpoint()
  1417. await self._protocol.write_event.wait()
  1418. if self._closed:
  1419. raise ClosedResourceError
  1420. elif self._transport.is_closing():
  1421. raise BrokenResourceError
  1422. else:
  1423. self._transport.sendto(*item)
  1424. class ConnectedUDPSocket(abc.ConnectedUDPSocket):
  1425. def __init__(
  1426. self, transport: asyncio.DatagramTransport, protocol: DatagramProtocol
  1427. ):
  1428. self._transport = transport
  1429. self._protocol = protocol
  1430. self._receive_guard = ResourceGuard("reading from")
  1431. self._send_guard = ResourceGuard("writing to")
  1432. self._closed = False
  1433. @property
  1434. def _raw_socket(self) -> socket.socket:
  1435. return self._transport.get_extra_info("socket")
  1436. async def aclose(self) -> None:
  1437. self._closed = True
  1438. if not self._transport.is_closing():
  1439. self._transport.close()
  1440. await self._protocol.closed_event.wait()
  1441. async def receive(self) -> bytes:
  1442. with self._receive_guard:
  1443. await AsyncIOBackend.checkpoint()
  1444. # If the buffer is empty, ask for more data
  1445. if not self._protocol.read_queue and not self._transport.is_closing():
  1446. self._protocol.read_event.clear()
  1447. await self._protocol.read_event.wait()
  1448. try:
  1449. packet = self._protocol.read_queue.popleft()
  1450. except IndexError:
  1451. if self._closed:
  1452. raise ClosedResourceError from None
  1453. else:
  1454. raise BrokenResourceError from None
  1455. return packet[0]
  1456. async def send(self, item: bytes) -> None:
  1457. with self._send_guard:
  1458. await AsyncIOBackend.checkpoint()
  1459. await self._protocol.write_event.wait()
  1460. if self._closed:
  1461. raise ClosedResourceError
  1462. elif self._transport.is_closing():
  1463. raise BrokenResourceError
  1464. else:
  1465. self._transport.sendto(item)
  1466. class UNIXDatagramSocket(_RawSocketMixin, abc.UNIXDatagramSocket):
  1467. async def receive(self) -> UNIXDatagramPacketType:
  1468. loop = get_running_loop()
  1469. await AsyncIOBackend.checkpoint()
  1470. with self._receive_guard:
  1471. while True:
  1472. try:
  1473. data = self._raw_socket.recvfrom(65536)
  1474. except BlockingIOError:
  1475. await self._wait_until_readable(loop)
  1476. except OSError as exc:
  1477. if self._closing:
  1478. raise ClosedResourceError from None
  1479. else:
  1480. raise BrokenResourceError from exc
  1481. else:
  1482. return data
  1483. async def send(self, item: UNIXDatagramPacketType) -> None:
  1484. loop = get_running_loop()
  1485. await AsyncIOBackend.checkpoint()
  1486. with self._send_guard:
  1487. while True:
  1488. try:
  1489. self._raw_socket.sendto(*item)
  1490. except BlockingIOError:
  1491. await self._wait_until_writable(loop)
  1492. except OSError as exc:
  1493. if self._closing:
  1494. raise ClosedResourceError from None
  1495. else:
  1496. raise BrokenResourceError from exc
  1497. else:
  1498. return
  1499. class ConnectedUNIXDatagramSocket(_RawSocketMixin, abc.ConnectedUNIXDatagramSocket):
  1500. async def receive(self) -> bytes:
  1501. loop = get_running_loop()
  1502. await AsyncIOBackend.checkpoint()
  1503. with self._receive_guard:
  1504. while True:
  1505. try:
  1506. data = self._raw_socket.recv(65536)
  1507. except BlockingIOError:
  1508. await self._wait_until_readable(loop)
  1509. except OSError as exc:
  1510. if self._closing:
  1511. raise ClosedResourceError from None
  1512. else:
  1513. raise BrokenResourceError from exc
  1514. else:
  1515. return data
  1516. async def send(self, item: bytes) -> None:
  1517. loop = get_running_loop()
  1518. await AsyncIOBackend.checkpoint()
  1519. with self._send_guard:
  1520. while True:
  1521. try:
  1522. self._raw_socket.send(item)
  1523. except BlockingIOError:
  1524. await self._wait_until_writable(loop)
  1525. except OSError as exc:
  1526. if self._closing:
  1527. raise ClosedResourceError from None
  1528. else:
  1529. raise BrokenResourceError from exc
  1530. else:
  1531. return
  1532. _read_events: RunVar[dict[int, asyncio.Future[bool]]] = RunVar("read_events")
  1533. _write_events: RunVar[dict[int, asyncio.Future[bool]]] = RunVar("write_events")
  1534. #
  1535. # Synchronization
  1536. #
  1537. class Event(BaseEvent):
  1538. __slots__ = ("_event",)
  1539. def __new__(cls) -> Event:
  1540. return object.__new__(cls)
  1541. def __init__(self) -> None:
  1542. self._event = asyncio.Event()
  1543. def set(self) -> None:
  1544. self._event.set()
  1545. def is_set(self) -> bool:
  1546. return self._event.is_set()
  1547. async def wait(self) -> None:
  1548. if self.is_set():
  1549. await AsyncIOBackend.checkpoint()
  1550. else:
  1551. await self._event.wait()
  1552. def statistics(self) -> EventStatistics:
  1553. return EventStatistics(len(self._event._waiters))
  1554. class Lock(BaseLock):
  1555. __slots__ = "_fast_acquire", "_owner_task", "_waiters"
  1556. def __new__(cls, *, fast_acquire: bool = False) -> Lock:
  1557. return object.__new__(cls)
  1558. def __init__(self, *, fast_acquire: bool = False) -> None:
  1559. self._fast_acquire = fast_acquire
  1560. self._owner_task: asyncio.Task | None = None
  1561. self._waiters: deque[tuple[asyncio.Task, asyncio.Future]] = deque()
  1562. async def acquire(self) -> None:
  1563. task = cast(asyncio.Task, current_task())
  1564. if self._owner_task is None and not self._waiters:
  1565. await AsyncIOBackend.checkpoint_if_cancelled()
  1566. self._owner_task = task
  1567. # Unless on the "fast path", yield control of the event loop so that other
  1568. # tasks can run too
  1569. if not self._fast_acquire:
  1570. try:
  1571. await AsyncIOBackend.cancel_shielded_checkpoint()
  1572. except CancelledError:
  1573. self.release()
  1574. raise
  1575. return
  1576. if self._owner_task == task:
  1577. raise RuntimeError("Attempted to acquire an already held Lock")
  1578. fut: asyncio.Future[None] = asyncio.Future()
  1579. item = task, fut
  1580. self._waiters.append(item)
  1581. try:
  1582. await fut
  1583. except CancelledError:
  1584. if fut.cancelled():
  1585. try:
  1586. self._waiters.remove(item)
  1587. except ValueError:
  1588. pass
  1589. else:
  1590. self.release()
  1591. raise
  1592. def acquire_nowait(self) -> None:
  1593. task = cast(asyncio.Task, current_task())
  1594. if self._owner_task is None and not self._waiters:
  1595. self._owner_task = task
  1596. return
  1597. if self._owner_task is task:
  1598. raise RuntimeError("Attempted to acquire an already held Lock")
  1599. raise WouldBlock
  1600. def locked(self) -> bool:
  1601. return self._owner_task is not None
  1602. def release(self) -> None:
  1603. if self._owner_task != current_task():
  1604. raise RuntimeError("The current task is not holding this lock")
  1605. # A cancelled waiter that already received ownership removes itself from
  1606. # _waiters before calling release(); any cancelled waiter still queued here
  1607. # was cancelled before being woken, so drop it.
  1608. while self._waiters:
  1609. task, fut = self._waiters.popleft()
  1610. if fut.cancelled():
  1611. continue
  1612. self._owner_task = task
  1613. fut.set_result(None)
  1614. return
  1615. self._owner_task = None
  1616. def statistics(self) -> LockStatistics:
  1617. task_info = AsyncIOTaskInfo(self._owner_task) if self._owner_task else None
  1618. return LockStatistics(self.locked(), task_info, len(self._waiters))
  1619. class Semaphore(BaseSemaphore):
  1620. __slots__ = "_value", "_max_value", "_fast_acquire", "_waiters"
  1621. def __new__(
  1622. cls,
  1623. initial_value: int,
  1624. *,
  1625. max_value: int | None = None,
  1626. fast_acquire: bool = False,
  1627. ) -> Semaphore:
  1628. return object.__new__(cls)
  1629. def __init__(
  1630. self,
  1631. initial_value: int,
  1632. *,
  1633. max_value: int | None = None,
  1634. fast_acquire: bool = False,
  1635. ):
  1636. super().__init__(initial_value, max_value=max_value)
  1637. self._value = initial_value
  1638. self._max_value = max_value
  1639. self._fast_acquire = fast_acquire
  1640. self._waiters: deque[asyncio.Future[None]] = deque()
  1641. async def acquire(self) -> None:
  1642. if self._value > 0 and not self._waiters:
  1643. await AsyncIOBackend.checkpoint_if_cancelled()
  1644. self._value -= 1
  1645. # Unless on the "fast path", yield control of the event loop so that other
  1646. # tasks can run too
  1647. if not self._fast_acquire:
  1648. try:
  1649. await AsyncIOBackend.cancel_shielded_checkpoint()
  1650. except CancelledError:
  1651. self.release()
  1652. raise
  1653. return
  1654. fut: asyncio.Future[None] = asyncio.Future()
  1655. self._waiters.append(fut)
  1656. try:
  1657. await fut
  1658. except CancelledError:
  1659. if fut.cancelled():
  1660. try:
  1661. self._waiters.remove(fut)
  1662. except ValueError:
  1663. pass
  1664. else:
  1665. self.release()
  1666. raise
  1667. def acquire_nowait(self) -> None:
  1668. if self._value == 0:
  1669. raise WouldBlock
  1670. self._value -= 1
  1671. def release(self) -> None:
  1672. if self._max_value is not None and self._value == self._max_value:
  1673. raise ValueError("semaphore released too many times")
  1674. while self._waiters:
  1675. fut = self._waiters.popleft()
  1676. if fut.cancelled():
  1677. continue
  1678. fut.set_result(None)
  1679. return
  1680. self._value += 1
  1681. @property
  1682. def value(self) -> int:
  1683. return self._value
  1684. @property
  1685. def max_value(self) -> int | None:
  1686. return self._max_value
  1687. def statistics(self) -> SemaphoreStatistics:
  1688. return SemaphoreStatistics(len(self._waiters))
  1689. class CapacityLimiter(BaseCapacityLimiter):
  1690. __slots__ = "_total_tokens", "_borrowers", "_wait_queue"
  1691. def __new__(cls, total_tokens: float) -> CapacityLimiter:
  1692. return object.__new__(cls)
  1693. def __init__(self, total_tokens: float):
  1694. self._total_tokens: float = 0
  1695. self._borrowers: set[Any] = set()
  1696. self._wait_queue: OrderedDict[Any, asyncio.Event] = OrderedDict()
  1697. self.total_tokens = total_tokens
  1698. async def __aenter__(self) -> None:
  1699. await self.acquire()
  1700. async def __aexit__(
  1701. self,
  1702. exc_type: type[BaseException] | None,
  1703. exc_val: BaseException | None,
  1704. exc_tb: TracebackType | None,
  1705. ) -> None:
  1706. self.release()
  1707. @property
  1708. def total_tokens(self) -> float:
  1709. return self._total_tokens
  1710. @total_tokens.setter
  1711. def total_tokens(self, value: float) -> None:
  1712. if not isinstance(value, int) and not math.isinf(value):
  1713. raise TypeError("total_tokens must be an int or math.inf")
  1714. if value < 0:
  1715. raise ValueError("total_tokens must be >= 0")
  1716. waiters_to_notify = max(value - self._total_tokens, 0)
  1717. self._total_tokens = value
  1718. # Notify waiting tasks that they have acquired the limiter
  1719. while self._wait_queue and waiters_to_notify:
  1720. borrower, event = self._wait_queue.popitem(last=False)
  1721. self._borrowers.add(borrower)
  1722. event.set()
  1723. waiters_to_notify -= 1
  1724. @property
  1725. def borrowed_tokens(self) -> int:
  1726. return len(self._borrowers)
  1727. @property
  1728. def available_tokens(self) -> float:
  1729. return self._total_tokens - len(self._borrowers)
  1730. def _notify_next_waiter(self) -> None:
  1731. """Hand a free token to the next task in line, if any."""
  1732. if self._wait_queue and len(self._borrowers) < self._total_tokens:
  1733. borrower, event = self._wait_queue.popitem(last=False)
  1734. self._borrowers.add(borrower)
  1735. event.set()
  1736. def acquire_nowait(self) -> None:
  1737. self.acquire_on_behalf_of_nowait(current_task())
  1738. def acquire_on_behalf_of_nowait(self, borrower: object) -> None:
  1739. if borrower in self._borrowers:
  1740. raise RuntimeError(
  1741. "this borrower is already holding one of this CapacityLimiter's tokens"
  1742. )
  1743. if self._wait_queue or len(self._borrowers) >= self._total_tokens:
  1744. raise WouldBlock
  1745. self._borrowers.add(borrower)
  1746. async def acquire(self) -> None:
  1747. return await self.acquire_on_behalf_of(current_task())
  1748. async def acquire_on_behalf_of(self, borrower: object) -> None:
  1749. await AsyncIOBackend.checkpoint_if_cancelled()
  1750. try:
  1751. self.acquire_on_behalf_of_nowait(borrower)
  1752. except WouldBlock:
  1753. event = asyncio.Event()
  1754. self._wait_queue[borrower] = event
  1755. try:
  1756. await event.wait()
  1757. except BaseException:
  1758. self._wait_queue.pop(borrower, None)
  1759. if event.is_set():
  1760. self._borrowers.discard(borrower)
  1761. self._notify_next_waiter()
  1762. raise
  1763. else:
  1764. try:
  1765. await AsyncIOBackend.cancel_shielded_checkpoint()
  1766. except BaseException:
  1767. self.release()
  1768. raise
  1769. def release(self) -> None:
  1770. self.release_on_behalf_of(current_task())
  1771. def release_on_behalf_of(self, borrower: object) -> None:
  1772. try:
  1773. self._borrowers.remove(borrower)
  1774. except KeyError:
  1775. raise RuntimeError(
  1776. "this borrower isn't holding any of this CapacityLimiter's tokens"
  1777. ) from None
  1778. self._notify_next_waiter()
  1779. def statistics(self) -> CapacityLimiterStatistics:
  1780. return CapacityLimiterStatistics(
  1781. self.borrowed_tokens,
  1782. self.total_tokens,
  1783. tuple(self._borrowers),
  1784. len(self._wait_queue),
  1785. )
  1786. _default_thread_limiter: RunVar[CapacityLimiter] = RunVar("_default_thread_limiter")
  1787. #
  1788. # Operating system signals
  1789. #
  1790. class _SignalReceiver:
  1791. def __init__(self, signals: tuple[Signals, ...]):
  1792. self._signals = signals
  1793. self._loop = get_running_loop()
  1794. self._signal_queue: deque[Signals] = deque()
  1795. self._future: asyncio.Future = asyncio.Future()
  1796. self._handled_signals: set[Signals] = set()
  1797. def _deliver(self, signum: Signals) -> None:
  1798. self._signal_queue.append(signum)
  1799. if not self._future.done():
  1800. self._future.set_result(None)
  1801. def __enter__(self) -> _SignalReceiver:
  1802. for sig in set(self._signals):
  1803. self._loop.add_signal_handler(sig, self._deliver, sig)
  1804. self._handled_signals.add(sig)
  1805. return self
  1806. def __exit__(
  1807. self,
  1808. exc_type: type[BaseException] | None,
  1809. exc_val: BaseException | None,
  1810. exc_tb: TracebackType | None,
  1811. ) -> None:
  1812. for sig in self._handled_signals:
  1813. self._loop.remove_signal_handler(sig)
  1814. def __aiter__(self) -> _SignalReceiver:
  1815. return self
  1816. async def __anext__(self) -> Signals:
  1817. await AsyncIOBackend.checkpoint()
  1818. if not self._signal_queue:
  1819. self._future = asyncio.Future()
  1820. await self._future
  1821. return self._signal_queue.popleft()
  1822. #
  1823. # Testing and debugging
  1824. #
  1825. class AsyncIOTaskInfo(TaskInfo):
  1826. def __init__(self, task: asyncio.Task):
  1827. task_state = _task_states.get(task)
  1828. if task_state is None:
  1829. parent_id = None
  1830. else:
  1831. parent_id = task_state.parent_id
  1832. coro = task.get_coro()
  1833. assert coro is not None, "created TaskInfo from a completed Task"
  1834. super().__init__(id(task), parent_id, task.get_name(), coro)
  1835. self._task = weakref.ref(task)
  1836. def has_pending_cancellation(self) -> bool:
  1837. if not (task := self._task()):
  1838. # If the task isn't around anymore, it won't have a pending cancellation
  1839. return False
  1840. if task._must_cancel: # type: ignore[attr-defined]
  1841. return True
  1842. elif (
  1843. isinstance(task._fut_waiter, asyncio.Future) # type: ignore[attr-defined]
  1844. and task._fut_waiter.cancelled() # type: ignore[attr-defined]
  1845. ):
  1846. return True
  1847. if task_state := _task_states.get(task):
  1848. if cancel_scope := task_state.cancel_scope:
  1849. return cancel_scope._effectively_cancelled
  1850. return False
  1851. class TestRunner(abc.TestRunner):
  1852. _send_stream: MemoryObjectSendStream[tuple[Awaitable[Any], asyncio.Future[Any]]]
  1853. def __init__(
  1854. self,
  1855. *,
  1856. debug: bool | None = None,
  1857. use_uvloop: bool = False,
  1858. loop_factory: Callable[[], AbstractEventLoop] | None = None,
  1859. ) -> None:
  1860. if use_uvloop and loop_factory is None:
  1861. if sys.platform != "win32":
  1862. import uvloop
  1863. loop_factory = uvloop.new_event_loop
  1864. else:
  1865. import winloop
  1866. loop_factory = winloop.new_event_loop
  1867. self._runner = Runner(debug=debug, loop_factory=loop_factory)
  1868. self._exceptions: list[BaseException] = []
  1869. self._runner_task: asyncio.Task | None = None
  1870. def __enter__(self) -> TestRunner:
  1871. self._runner.__enter__()
  1872. self.get_loop().set_exception_handler(self._exception_handler)
  1873. return self
  1874. def __exit__(
  1875. self,
  1876. exc_type: type[BaseException] | None,
  1877. exc_val: BaseException | None,
  1878. exc_tb: TracebackType | None,
  1879. ) -> None:
  1880. self._runner.__exit__(exc_type, exc_val, exc_tb)
  1881. def get_loop(self) -> AbstractEventLoop:
  1882. return self._runner.get_loop()
  1883. def is_running(self) -> bool:
  1884. try:
  1885. asyncio.get_running_loop()
  1886. return True
  1887. except RuntimeError:
  1888. return False
  1889. def _exception_handler(
  1890. self, loop: asyncio.AbstractEventLoop, context: dict[str, Any]
  1891. ) -> None:
  1892. if isinstance(context.get("exception"), Exception):
  1893. self._exceptions.append(context["exception"])
  1894. else:
  1895. loop.default_exception_handler(context)
  1896. def _raise_async_exceptions(self) -> None:
  1897. # Re-raise any exceptions raised in asynchronous callbacks
  1898. if self._exceptions:
  1899. exceptions, self._exceptions = self._exceptions, []
  1900. if len(exceptions) == 1:
  1901. raise exceptions[0]
  1902. elif exceptions:
  1903. raise BaseExceptionGroup(
  1904. "Multiple exceptions occurred in asynchronous callbacks", exceptions
  1905. )
  1906. async def _run_tests_and_fixtures(
  1907. self,
  1908. receive_stream: MemoryObjectReceiveStream[
  1909. tuple[Awaitable[T_Retval], asyncio.Future[T_Retval]]
  1910. ],
  1911. ) -> None:
  1912. from _pytest.outcomes import OutcomeException
  1913. with receive_stream, self._send_stream:
  1914. async for coro, future in receive_stream:
  1915. try:
  1916. retval = await coro
  1917. except CancelledError as exc:
  1918. if not future.cancelled():
  1919. future.cancel(*exc.args)
  1920. raise
  1921. except BaseException as exc:
  1922. if not future.cancelled():
  1923. future.set_exception(exc)
  1924. if not isinstance(exc, (Exception, OutcomeException)):
  1925. raise
  1926. else:
  1927. if not future.cancelled():
  1928. future.set_result(retval)
  1929. async def _call_in_runner_task(
  1930. self,
  1931. func: Callable[P, Awaitable[T_Retval]],
  1932. /,
  1933. *args: P.args,
  1934. **kwargs: P.kwargs,
  1935. ) -> T_Retval:
  1936. if not self._runner_task:
  1937. self._send_stream, receive_stream = create_memory_object_stream[
  1938. tuple[Awaitable[Any], asyncio.Future]
  1939. ](1)
  1940. self._runner_task = self.get_loop().create_task(
  1941. self._run_tests_and_fixtures(receive_stream)
  1942. )
  1943. coro = func(*args, **kwargs)
  1944. future: asyncio.Future[T_Retval] = self.get_loop().create_future()
  1945. self._send_stream.send_nowait((coro, future))
  1946. return await future
  1947. def run_asyncgen_fixture(
  1948. self,
  1949. fixture_func: Callable[..., AsyncGenerator[T_Retval, Any]],
  1950. kwargs: dict[str, Any],
  1951. ) -> Iterable[T_Retval]:
  1952. asyncgen = fixture_func(**kwargs)
  1953. fixturevalue: T_Retval = self.get_loop().run_until_complete(
  1954. self._call_in_runner_task(asyncgen.asend, None)
  1955. )
  1956. self._raise_async_exceptions()
  1957. yield fixturevalue
  1958. try:
  1959. self.get_loop().run_until_complete(
  1960. self._call_in_runner_task(asyncgen.asend, None)
  1961. )
  1962. except StopAsyncIteration:
  1963. self._raise_async_exceptions()
  1964. else:
  1965. self.get_loop().run_until_complete(asyncgen.aclose())
  1966. raise RuntimeError("Async generator fixture did not stop")
  1967. def run_fixture(
  1968. self,
  1969. fixture_func: Callable[..., Coroutine[Any, Any, T_Retval]],
  1970. kwargs: dict[str, Any],
  1971. ) -> T_Retval:
  1972. retval = self.get_loop().run_until_complete(
  1973. self._call_in_runner_task(fixture_func, **kwargs)
  1974. )
  1975. self._raise_async_exceptions()
  1976. return retval
  1977. def run_test(
  1978. self, test_func: Callable[..., Coroutine[Any, Any, Any]], kwargs: dict[str, Any]
  1979. ) -> None:
  1980. from _pytest.outcomes import OutcomeException
  1981. try:
  1982. self.get_loop().run_until_complete(
  1983. self._call_in_runner_task(test_func, **kwargs)
  1984. )
  1985. except Exception as exc:
  1986. self._exceptions.append(exc)
  1987. except OutcomeException:
  1988. raise
  1989. except BaseException:
  1990. # A BaseException (e.g. KeyboardInterrupt, SystemExit) interrupted the event loop before
  1991. # the test completed. Cancel _runner_task so it does not resume when the event
  1992. # loop is re-entered during async generator fixture teardown.
  1993. if self._runner_task is not None and not self._runner_task.done():
  1994. self._runner_task.cancel()
  1995. self._send_stream.close()
  1996. try:
  1997. self.get_loop().run_until_complete(self._runner_task)
  1998. except CancelledError:
  1999. pass
  2000. finally:
  2001. self._runner_task = None
  2002. raise
  2003. self._raise_async_exceptions()
  2004. class _ProcessStreamProtocol(asyncio.subprocess.SubprocessStreamProtocol):
  2005. """
  2006. A subprocess protocol that allows us to be notified of ``process_exited``
  2007. asyncio's own ``Process.wait()`` only resolves once every pipe transport has
  2008. disconnected so to get same semantics as on trio and uvloop we need this.
  2009. """
  2010. def __init__(self) -> None:
  2011. # Match the standard factory for asyncio.create_process
  2012. super().__init__(limit=2**16, loop=asyncio.get_running_loop())
  2013. self.exited = asyncio.Event()
  2014. def process_exited(self) -> None:
  2015. super().process_exited()
  2016. self.exited.set()
  2017. class AsyncIOBackend(AsyncBackend):
  2018. @classmethod
  2019. def run(
  2020. cls,
  2021. func: Callable[[Unpack[PosArgsT]], Awaitable[T_Retval]],
  2022. args: tuple[Unpack[PosArgsT]],
  2023. kwargs: dict[str, Any],
  2024. options: dict[str, Any],
  2025. ) -> T_Retval:
  2026. @wraps(func)
  2027. async def wrapper() -> T_Retval:
  2028. task = cast(asyncio.Task, current_task())
  2029. task.set_name(get_callable_name(func))
  2030. _task_states[task] = TaskState(None, None)
  2031. try:
  2032. return await func(*args)
  2033. finally:
  2034. del _task_states[task]
  2035. debug = options.get("debug", None)
  2036. loop_factory = options.get("loop_factory", None)
  2037. if loop_factory is None and options.get("use_uvloop", False):
  2038. if sys.platform != "win32":
  2039. import uvloop
  2040. loop_factory = uvloop.new_event_loop
  2041. else:
  2042. import winloop
  2043. loop_factory = winloop.new_event_loop
  2044. with Runner(debug=debug, loop_factory=loop_factory) as runner:
  2045. return runner.run(wrapper())
  2046. @classmethod
  2047. def current_token(cls) -> object:
  2048. return get_running_loop()
  2049. @classmethod
  2050. def current_time(cls) -> float:
  2051. return get_running_loop().time()
  2052. @classmethod
  2053. def cancelled_exception_class(cls) -> type[BaseException]:
  2054. return CancelledError
  2055. @classmethod
  2056. async def checkpoint(cls) -> None:
  2057. await sleep(0)
  2058. @classmethod
  2059. async def checkpoint_if_cancelled(cls) -> None:
  2060. task = current_task()
  2061. if task is None:
  2062. return
  2063. try:
  2064. cancel_scope = _task_states[task].cancel_scope
  2065. except KeyError:
  2066. return
  2067. while cancel_scope:
  2068. if cancel_scope.cancel_called:
  2069. await sleep(0)
  2070. elif cancel_scope.shield:
  2071. break
  2072. else:
  2073. cancel_scope = cancel_scope._parent_scope
  2074. @classmethod
  2075. async def cancel_shielded_checkpoint(cls) -> None:
  2076. with CancelScope(shield=True):
  2077. await sleep(0)
  2078. @classmethod
  2079. async def sleep(cls, delay: float) -> None:
  2080. await sleep(delay)
  2081. @classmethod
  2082. def create_cancel_scope(
  2083. cls, *, deadline: float = math.inf, shield: bool = False
  2084. ) -> CancelScope:
  2085. return CancelScope(deadline=deadline, shield=shield)
  2086. @classmethod
  2087. def current_effective_deadline(cls) -> float:
  2088. if (task := current_task()) is None:
  2089. return math.inf
  2090. try:
  2091. cancel_scope = _task_states[task].cancel_scope
  2092. except KeyError:
  2093. return math.inf
  2094. deadline = math.inf
  2095. while cancel_scope:
  2096. deadline = min(deadline, cancel_scope.deadline)
  2097. if cancel_scope._cancel_called:
  2098. deadline = -math.inf
  2099. break
  2100. elif cancel_scope.shield:
  2101. break
  2102. else:
  2103. cancel_scope = cancel_scope._parent_scope
  2104. return deadline
  2105. @classmethod
  2106. def create_task_group(cls) -> abc.TaskGroup:
  2107. return TaskGroup()
  2108. @classmethod
  2109. def create_event(cls) -> abc.Event:
  2110. return Event()
  2111. @classmethod
  2112. def create_lock(cls, *, fast_acquire: bool) -> abc.Lock:
  2113. return Lock(fast_acquire=fast_acquire)
  2114. @classmethod
  2115. def create_semaphore(
  2116. cls,
  2117. initial_value: int,
  2118. *,
  2119. max_value: int | None = None,
  2120. fast_acquire: bool = False,
  2121. ) -> abc.Semaphore:
  2122. return Semaphore(initial_value, max_value=max_value, fast_acquire=fast_acquire)
  2123. @classmethod
  2124. def create_capacity_limiter(cls, total_tokens: float) -> abc.CapacityLimiter:
  2125. return CapacityLimiter(total_tokens)
  2126. @classmethod
  2127. async def run_sync_in_worker_thread( # type: ignore[return]
  2128. cls,
  2129. func: Callable[[Unpack[PosArgsT]], T_Retval],
  2130. args: tuple[Unpack[PosArgsT]],
  2131. abandon_on_cancel: bool = False,
  2132. limiter: abc.CapacityLimiter | None = None,
  2133. ) -> T_Retval:
  2134. await cls.checkpoint()
  2135. # If this is the first run in this event loop thread, set up the necessary
  2136. # variables
  2137. try:
  2138. idle_workers = _threadpool_idle_workers.get()
  2139. workers = _threadpool_workers.get()
  2140. except LookupError:
  2141. idle_workers = deque()
  2142. workers = set()
  2143. _threadpool_idle_workers.set(idle_workers)
  2144. _threadpool_workers.set(workers)
  2145. async with limiter or cls.current_default_thread_limiter():
  2146. with CancelScope(shield=not abandon_on_cancel) as scope:
  2147. future = asyncio.Future[T_Retval]()
  2148. root_task = find_root_task()
  2149. if not idle_workers:
  2150. worker = WorkerThread(root_task, workers, idle_workers)
  2151. worker.start()
  2152. workers.add(worker)
  2153. root_task.add_done_callback(
  2154. worker.stop, context=contextvars.Context()
  2155. )
  2156. else:
  2157. worker = idle_workers.pop()
  2158. # Prune any other workers that have been idle for MAX_IDLE_TIME
  2159. # seconds or longer
  2160. now = cls.current_time()
  2161. while idle_workers:
  2162. if (
  2163. now - idle_workers[0].idle_since
  2164. < WorkerThread.MAX_IDLE_TIME
  2165. ):
  2166. break
  2167. expired_worker = idle_workers.popleft()
  2168. expired_worker.root_task.remove_done_callback(
  2169. expired_worker.stop
  2170. )
  2171. expired_worker.stop()
  2172. context = copy_context()
  2173. context.run(set_current_async_library, None)
  2174. if abandon_on_cancel or scope._parent_scope is None:
  2175. worker_scope = scope
  2176. else:
  2177. worker_scope = scope._parent_scope
  2178. worker.queue.put_nowait((context, func, args, future, worker_scope))
  2179. return await future
  2180. @classmethod
  2181. def check_cancelled(cls) -> None:
  2182. scope: CancelScope | None = threadlocals.current_cancel_scope
  2183. while scope is not None:
  2184. if scope.cancel_called:
  2185. raise CancelledError(f"Cancelled via cancel scope {id(scope):x}")
  2186. if scope.shield:
  2187. return
  2188. scope = scope._parent_scope
  2189. @classmethod
  2190. def run_async_from_thread(
  2191. cls,
  2192. func: Callable[[Unpack[PosArgsT]], Coroutine[Any, Any, T_co]],
  2193. args: tuple[Unpack[PosArgsT]],
  2194. token: object,
  2195. ) -> T_co:
  2196. async def task_wrapper() -> T_co:
  2197. __tracebackhide__ = True
  2198. if scope is not None:
  2199. task = cast(asyncio.Task, current_task())
  2200. _task_states[task] = TaskState(None, scope)
  2201. scope._tasks.add(task)
  2202. try:
  2203. return await func(*args)
  2204. except CancelledError as exc:
  2205. raise concurrent.futures.CancelledError(str(exc)) from None
  2206. finally:
  2207. if scope is not None:
  2208. scope._tasks.discard(task)
  2209. loop = cast(
  2210. "AbstractEventLoop", token or threadlocals.current_token.native_token
  2211. )
  2212. if loop.is_closed():
  2213. raise RunFinishedError
  2214. context = copy_context()
  2215. context.run(set_current_async_library, "asyncio")
  2216. scope = getattr(threadlocals, "current_cancel_scope", None)
  2217. f: concurrent.futures.Future[T_co] = context.run(
  2218. asyncio.run_coroutine_threadsafe, task_wrapper(), loop=loop
  2219. )
  2220. return f.result()
  2221. @classmethod
  2222. def run_sync_from_thread(
  2223. cls,
  2224. func: Callable[[Unpack[PosArgsT]], T_Retval],
  2225. args: tuple[Unpack[PosArgsT]],
  2226. token: object,
  2227. ) -> T_Retval:
  2228. @wraps(func)
  2229. def wrapper() -> None:
  2230. try:
  2231. set_current_async_library("asyncio")
  2232. f.set_result(func(*args))
  2233. except BaseException as exc:
  2234. f.set_exception(exc)
  2235. if not isinstance(exc, Exception):
  2236. raise
  2237. loop = cast(
  2238. "AbstractEventLoop", token or threadlocals.current_token.native_token
  2239. )
  2240. if loop.is_closed():
  2241. raise RunFinishedError
  2242. f: concurrent.futures.Future[T_Retval] = Future()
  2243. loop.call_soon_threadsafe(wrapper)
  2244. return f.result()
  2245. @classmethod
  2246. async def open_process(
  2247. cls,
  2248. command: StrOrBytesPath | Sequence[StrOrBytesPath],
  2249. *,
  2250. stdin: int | IO[Any] | None,
  2251. stdout: int | IO[Any] | None,
  2252. stderr: int | IO[Any] | None,
  2253. **kwargs: Any,
  2254. ) -> Process:
  2255. await cls.checkpoint()
  2256. if isinstance(command, PathLike):
  2257. command = os.fspath(command)
  2258. # Use loop.subprocess_shell()/subprocess_exec() rather than their
  2259. # asyncio.create_subprocess_*() counterparts to get access to
  2260. # transport/protocol.
  2261. loop = asyncio.get_running_loop()
  2262. if isinstance(command, (str, bytes)):
  2263. transport, protocol = await loop.subprocess_shell(
  2264. _ProcessStreamProtocol,
  2265. command,
  2266. stdin=stdin,
  2267. stdout=stdout,
  2268. stderr=stderr,
  2269. **kwargs,
  2270. )
  2271. else:
  2272. transport, protocol = await loop.subprocess_exec(
  2273. _ProcessStreamProtocol,
  2274. *command,
  2275. stdin=stdin,
  2276. stdout=stdout,
  2277. stderr=stderr,
  2278. **kwargs,
  2279. )
  2280. process = asyncio.subprocess.Process(transport, protocol, loop)
  2281. stdin_stream = StreamWriterWrapper(process.stdin) if process.stdin else None
  2282. stdout_stream = StreamReaderWrapper(process.stdout) if process.stdout else None
  2283. stderr_stream = StreamReaderWrapper(process.stderr) if process.stderr else None
  2284. return Process(
  2285. process,
  2286. stdin_stream,
  2287. stdout_stream,
  2288. stderr_stream,
  2289. protocol.exited,
  2290. transport,
  2291. )
  2292. @classmethod
  2293. def setup_process_pool_exit_at_shutdown(cls, workers: set[abc.Process]) -> None:
  2294. create_task(
  2295. _shutdown_process_pool_on_exit(workers),
  2296. name="AnyIO process pool shutdown task",
  2297. )
  2298. find_root_task().add_done_callback(
  2299. partial(_forcibly_shutdown_process_pool_on_exit, workers) # type:ignore[arg-type]
  2300. )
  2301. @classmethod
  2302. async def connect_tcp(
  2303. cls, host: str, port: int, local_address: IPSockAddrType | None = None
  2304. ) -> abc.SocketStream:
  2305. transport, protocol = cast(
  2306. tuple[asyncio.Transport, StreamProtocol],
  2307. await get_running_loop().create_connection(
  2308. StreamProtocol, host, port, local_addr=local_address
  2309. ),
  2310. )
  2311. transport.pause_reading()
  2312. return SocketStream(transport, protocol)
  2313. @classmethod
  2314. async def connect_unix(cls, path: str | bytes) -> abc.UNIXSocketStream:
  2315. await cls.checkpoint()
  2316. loop = get_running_loop()
  2317. raw_socket = socket.socket(socket.AF_UNIX)
  2318. raw_socket.setblocking(False)
  2319. while True:
  2320. try:
  2321. raw_socket.connect(path)
  2322. except BlockingIOError:
  2323. f: asyncio.Future = asyncio.Future()
  2324. loop.add_writer(raw_socket, f.set_result, None)
  2325. f.add_done_callback(lambda _: loop.remove_writer(raw_socket))
  2326. await f
  2327. except BaseException:
  2328. raw_socket.close()
  2329. raise
  2330. else:
  2331. return UNIXSocketStream(raw_socket)
  2332. @classmethod
  2333. def create_tcp_listener(cls, sock: socket.socket) -> SocketListener:
  2334. return TCPSocketListener(sock)
  2335. @classmethod
  2336. def create_unix_listener(cls, sock: socket.socket) -> SocketListener:
  2337. return UNIXSocketListener(sock)
  2338. @classmethod
  2339. async def create_udp_socket(
  2340. cls,
  2341. family: AddressFamily,
  2342. local_address: IPSockAddrType | None,
  2343. remote_address: IPSockAddrType | None,
  2344. reuse_port: bool,
  2345. ) -> UDPSocket | ConnectedUDPSocket:
  2346. transport, protocol = await get_running_loop().create_datagram_endpoint(
  2347. DatagramProtocol,
  2348. local_addr=local_address,
  2349. remote_addr=remote_address,
  2350. family=family,
  2351. reuse_port=reuse_port,
  2352. )
  2353. if protocol.exception:
  2354. transport.close()
  2355. raise protocol.exception
  2356. if not remote_address:
  2357. return UDPSocket(transport, protocol)
  2358. else:
  2359. return ConnectedUDPSocket(transport, protocol)
  2360. @classmethod
  2361. async def create_unix_datagram_socket( # type: ignore[override]
  2362. cls, raw_socket: socket.socket, remote_path: str | bytes | None
  2363. ) -> abc.UNIXDatagramSocket | abc.ConnectedUNIXDatagramSocket:
  2364. await cls.checkpoint()
  2365. loop = get_running_loop()
  2366. if remote_path:
  2367. while True:
  2368. try:
  2369. raw_socket.connect(remote_path)
  2370. except BlockingIOError:
  2371. f: asyncio.Future = asyncio.Future()
  2372. loop.add_writer(raw_socket, f.set_result, None)
  2373. f.add_done_callback(lambda _: loop.remove_writer(raw_socket))
  2374. await f
  2375. except BaseException:
  2376. raw_socket.close()
  2377. raise
  2378. else:
  2379. return ConnectedUNIXDatagramSocket(raw_socket)
  2380. else:
  2381. return UNIXDatagramSocket(raw_socket)
  2382. @classmethod
  2383. async def getaddrinfo(
  2384. cls,
  2385. host: bytes | str | None,
  2386. port: str | int | None,
  2387. *,
  2388. family: int | AddressFamily = 0,
  2389. type: int | SocketKind = 0,
  2390. proto: int = 0,
  2391. flags: int = 0,
  2392. ) -> Sequence[
  2393. tuple[
  2394. AddressFamily,
  2395. SocketKind,
  2396. int,
  2397. str,
  2398. tuple[str, int] | tuple[str, int, int, int] | tuple[int, bytes],
  2399. ]
  2400. ]:
  2401. return await get_running_loop().getaddrinfo(
  2402. host, port, family=family, type=type, proto=proto, flags=flags
  2403. )
  2404. @classmethod
  2405. async def getnameinfo(
  2406. cls, sockaddr: IPSockAddrType, flags: int = 0
  2407. ) -> tuple[str, str]:
  2408. return await get_running_loop().getnameinfo(sockaddr, flags)
  2409. @classmethod
  2410. async def wait_readable(cls, obj: FileDescriptorLike) -> None:
  2411. try:
  2412. read_events = _read_events.get()
  2413. except LookupError:
  2414. read_events = {}
  2415. _read_events.set(read_events)
  2416. fd = obj if isinstance(obj, int) else obj.fileno()
  2417. if read_events.get(fd):
  2418. raise BusyResourceError("reading from")
  2419. loop = get_running_loop()
  2420. fut: asyncio.Future[bool] = loop.create_future()
  2421. def cb() -> None:
  2422. try:
  2423. del read_events[fd]
  2424. except KeyError:
  2425. pass
  2426. else:
  2427. remove_reader(fd)
  2428. try:
  2429. fut.set_result(True)
  2430. except asyncio.InvalidStateError:
  2431. pass
  2432. try:
  2433. loop.add_reader(fd, cb)
  2434. except NotImplementedError:
  2435. from anyio._core._asyncio_selector_thread import get_selector
  2436. selector = get_selector()
  2437. selector.add_reader(fd, cb)
  2438. remove_reader = selector.remove_reader
  2439. else:
  2440. remove_reader = loop.remove_reader
  2441. read_events[fd] = fut
  2442. try:
  2443. success = await fut
  2444. finally:
  2445. try:
  2446. del read_events[fd]
  2447. except KeyError:
  2448. pass
  2449. else:
  2450. remove_reader(fd)
  2451. if not success:
  2452. raise ClosedResourceError
  2453. @classmethod
  2454. async def wait_writable(cls, obj: FileDescriptorLike) -> None:
  2455. try:
  2456. write_events = _write_events.get()
  2457. except LookupError:
  2458. write_events = {}
  2459. _write_events.set(write_events)
  2460. fd = obj if isinstance(obj, int) else obj.fileno()
  2461. if write_events.get(fd):
  2462. raise BusyResourceError("writing to")
  2463. loop = get_running_loop()
  2464. fut: asyncio.Future[bool] = loop.create_future()
  2465. def cb() -> None:
  2466. try:
  2467. del write_events[fd]
  2468. except KeyError:
  2469. pass
  2470. else:
  2471. remove_writer(fd)
  2472. try:
  2473. fut.set_result(True)
  2474. except asyncio.InvalidStateError:
  2475. pass
  2476. try:
  2477. loop.add_writer(fd, cb)
  2478. except NotImplementedError:
  2479. from anyio._core._asyncio_selector_thread import get_selector
  2480. selector = get_selector()
  2481. selector.add_writer(fd, cb)
  2482. remove_writer = selector.remove_writer
  2483. else:
  2484. remove_writer = loop.remove_writer
  2485. write_events[fd] = fut
  2486. try:
  2487. success = await fut
  2488. finally:
  2489. try:
  2490. del write_events[fd]
  2491. except KeyError:
  2492. pass
  2493. else:
  2494. remove_writer(fd)
  2495. if not success:
  2496. raise ClosedResourceError
  2497. @classmethod
  2498. def notify_closing(cls, obj: FileDescriptorLike) -> None:
  2499. fd = obj if isinstance(obj, int) else obj.fileno()
  2500. loop = get_running_loop()
  2501. try:
  2502. write_events = _write_events.get()
  2503. except LookupError:
  2504. pass
  2505. else:
  2506. try:
  2507. fut = write_events.pop(fd)
  2508. except KeyError:
  2509. pass
  2510. else:
  2511. try:
  2512. fut.set_result(False)
  2513. except asyncio.InvalidStateError:
  2514. pass
  2515. try:
  2516. loop.remove_writer(fd)
  2517. except NotImplementedError:
  2518. from anyio._core._asyncio_selector_thread import get_selector
  2519. get_selector().remove_writer(fd)
  2520. try:
  2521. read_events = _read_events.get()
  2522. except LookupError:
  2523. pass
  2524. else:
  2525. try:
  2526. fut = read_events.pop(fd)
  2527. except KeyError:
  2528. pass
  2529. else:
  2530. try:
  2531. fut.set_result(False)
  2532. except asyncio.InvalidStateError:
  2533. pass
  2534. try:
  2535. loop.remove_reader(fd)
  2536. except NotImplementedError:
  2537. from anyio._core._asyncio_selector_thread import get_selector
  2538. get_selector().remove_reader(fd)
  2539. @classmethod
  2540. async def wrap_listener_socket(cls, sock: socket.socket) -> SocketListener:
  2541. if hasattr(socket, "AF_UNIX") and sock.family == socket.AF_UNIX:
  2542. return UNIXSocketListener(sock)
  2543. return TCPSocketListener(sock)
  2544. @classmethod
  2545. async def wrap_stream_socket(cls, sock: socket.socket) -> SocketStream:
  2546. transport, protocol = await get_running_loop().create_connection(
  2547. StreamProtocol, sock=sock
  2548. )
  2549. return SocketStream(transport, protocol)
  2550. @classmethod
  2551. async def wrap_unix_stream_socket(cls, sock: socket.socket) -> UNIXSocketStream:
  2552. return UNIXSocketStream(sock)
  2553. @classmethod
  2554. async def wrap_udp_socket(cls, sock: socket.socket) -> UDPSocket:
  2555. transport, protocol = await get_running_loop().create_datagram_endpoint(
  2556. DatagramProtocol, sock=sock
  2557. )
  2558. return UDPSocket(transport, protocol)
  2559. @classmethod
  2560. async def wrap_connected_udp_socket(cls, sock: socket.socket) -> ConnectedUDPSocket:
  2561. transport, protocol = await get_running_loop().create_datagram_endpoint(
  2562. DatagramProtocol, sock=sock
  2563. )
  2564. return ConnectedUDPSocket(transport, protocol)
  2565. @classmethod
  2566. async def wrap_unix_datagram_socket(cls, sock: socket.socket) -> UNIXDatagramSocket:
  2567. return UNIXDatagramSocket(sock)
  2568. @classmethod
  2569. async def wrap_connected_unix_datagram_socket(
  2570. cls, sock: socket.socket
  2571. ) -> ConnectedUNIXDatagramSocket:
  2572. return ConnectedUNIXDatagramSocket(sock)
  2573. @classmethod
  2574. def current_default_thread_limiter(cls) -> CapacityLimiter:
  2575. try:
  2576. return _default_thread_limiter.get()
  2577. except LookupError:
  2578. limiter = CapacityLimiter(40)
  2579. _default_thread_limiter.set(limiter)
  2580. return limiter
  2581. @classmethod
  2582. def open_signal_receiver(
  2583. cls, *signals: Signals
  2584. ) -> AbstractContextManager[AsyncIterator[Signals]]:
  2585. return _SignalReceiver(signals)
  2586. @classmethod
  2587. def get_current_task(cls) -> TaskInfo:
  2588. return AsyncIOTaskInfo(current_task()) # type: ignore[arg-type]
  2589. @classmethod
  2590. def get_running_tasks(cls) -> Sequence[TaskInfo]:
  2591. return [AsyncIOTaskInfo(task) for task in all_tasks() if not task.done()]
  2592. @classmethod
  2593. async def wait_all_tasks_blocked(cls) -> None:
  2594. await cls.checkpoint()
  2595. this_task = current_task()
  2596. while True:
  2597. for task in all_tasks():
  2598. if task is this_task:
  2599. continue
  2600. waiter = task._fut_waiter # type: ignore[attr-defined]
  2601. if waiter is None or waiter.done():
  2602. await sleep(0.1)
  2603. break
  2604. else:
  2605. return
  2606. @classmethod
  2607. def create_test_runner(cls, options: dict[str, Any]) -> TestRunner:
  2608. return TestRunner(**options)
  2609. backend_class = AsyncIOBackend