_compat.py 2.8 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091
  1. from __future__ import annotations
  2. import asyncio
  3. import sys
  4. from collections.abc import Callable, Coroutine
  5. from typing import Any, TypeVar
  6. __all__ = ["asyncio_run", "iscoroutinefunction"]
  7. if sys.version_info >= (3, 14):
  8. from inspect import iscoroutinefunction
  9. else:
  10. from asyncio import iscoroutinefunction
  11. _T = TypeVar("_T")
  12. if sys.version_info >= (3, 12):
  13. asyncio_run = asyncio.run
  14. elif sys.version_info >= (3, 11):
  15. def asyncio_run(
  16. main: Coroutine[Any, Any, _T],
  17. *,
  18. debug: bool = False,
  19. loop_factory: Callable[[], asyncio.AbstractEventLoop] | None = None,
  20. ) -> _T:
  21. # asyncio.run from Python 3.12
  22. # https://docs.python.org/3/license.html#psf-license
  23. with asyncio.Runner(debug=debug, loop_factory=loop_factory) as runner:
  24. return runner.run(main)
  25. else:
  26. # modified version of asyncio.run from Python 3.10 to add loop_factory kwarg
  27. # https://docs.python.org/3/license.html#psf-license
  28. def asyncio_run(
  29. main: Coroutine[Any, Any, _T],
  30. *,
  31. debug: bool = False,
  32. loop_factory: Callable[[], asyncio.AbstractEventLoop] | None = None,
  33. ) -> _T:
  34. try:
  35. asyncio.get_running_loop()
  36. except RuntimeError:
  37. pass
  38. else:
  39. raise RuntimeError("asyncio.run() cannot be called from a running event loop")
  40. if not asyncio.iscoroutine(main):
  41. raise ValueError(f"a coroutine was expected, got {main!r}")
  42. if loop_factory is None:
  43. loop = asyncio.new_event_loop()
  44. else:
  45. loop = loop_factory()
  46. try:
  47. if loop_factory is None:
  48. asyncio.set_event_loop(loop)
  49. if debug is not None:
  50. loop.set_debug(debug)
  51. return loop.run_until_complete(main)
  52. finally:
  53. try:
  54. _cancel_all_tasks(loop)
  55. loop.run_until_complete(loop.shutdown_asyncgens())
  56. loop.run_until_complete(loop.shutdown_default_executor())
  57. finally:
  58. if loop_factory is None:
  59. asyncio.set_event_loop(None)
  60. loop.close()
  61. def _cancel_all_tasks(loop: asyncio.AbstractEventLoop) -> None:
  62. to_cancel = asyncio.all_tasks(loop)
  63. if not to_cancel:
  64. return
  65. for task in to_cancel:
  66. task.cancel()
  67. loop.run_until_complete(asyncio.gather(*to_cancel, return_exceptions=True))
  68. for task in to_cancel:
  69. if task.cancelled():
  70. continue
  71. if task.exception() is not None:
  72. loop.call_exception_handler(
  73. {
  74. "message": "unhandled exception during asyncio.run() shutdown",
  75. "exception": task.exception(),
  76. "task": task,
  77. }
  78. )