_utils.py 2.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111
  1. from __future__ import annotations
  2. import functools
  3. import sys
  4. from collections.abc import AsyncGenerator, Awaitable, Callable, Generator
  5. from contextlib import AbstractAsyncContextManager, asynccontextmanager
  6. from typing import Any, Generic, Protocol, TypeVar, overload
  7. import anyio.abc
  8. from starlette.types import Scope
  9. if sys.version_info >= (3, 13): # pragma: no cover
  10. from inspect import iscoroutinefunction
  11. from typing import TypeIs
  12. else: # pragma: no cover
  13. from asyncio import iscoroutinefunction
  14. from typing_extensions import TypeIs
  15. if sys.version_info < (3, 11): # pragma: no cover
  16. try:
  17. from exceptiongroup import BaseExceptionGroup
  18. except ImportError:
  19. class BaseExceptionGroup(BaseException): # type: ignore[no-redef]
  20. pass
  21. T = TypeVar("T")
  22. AwaitableCallable = Callable[..., Awaitable[T]]
  23. @overload
  24. def is_async_callable(obj: AwaitableCallable[T]) -> TypeIs[AwaitableCallable[T]]: ...
  25. @overload
  26. def is_async_callable(obj: Any) -> TypeIs[AwaitableCallable[Any]]: ...
  27. def is_async_callable(obj: Any) -> Any:
  28. while isinstance(obj, functools.partial):
  29. obj = obj.func
  30. return iscoroutinefunction(obj) or (callable(obj) and iscoroutinefunction(obj.__call__))
  31. T_co = TypeVar("T_co", covariant=True)
  32. class AwaitableOrContextManager(
  33. Awaitable[T_co], AbstractAsyncContextManager[T_co], Protocol[T_co]
  34. ): ... # pragma: no branch
  35. class SupportsAsyncClose(Protocol):
  36. async def close(self) -> None: ... # pragma: no cover
  37. SupportsAsyncCloseType = TypeVar("SupportsAsyncCloseType", bound=SupportsAsyncClose, covariant=False)
  38. class AwaitableOrContextManagerWrapper(Generic[SupportsAsyncCloseType]):
  39. __slots__ = ("aw", "entered")
  40. def __init__(self, aw: Awaitable[SupportsAsyncCloseType]) -> None:
  41. self.aw = aw
  42. def __await__(self) -> Generator[Any, None, SupportsAsyncCloseType]:
  43. return self.aw.__await__()
  44. async def __aenter__(self) -> SupportsAsyncCloseType:
  45. self.entered = await self.aw
  46. return self.entered
  47. async def __aexit__(self, *args: Any) -> None | bool:
  48. await self.entered.close()
  49. return None
  50. @asynccontextmanager
  51. async def create_collapsing_task_group() -> AsyncGenerator[anyio.abc.TaskGroup, None]:
  52. try:
  53. async with anyio.create_task_group() as tg:
  54. yield tg
  55. except BaseExceptionGroup as excs:
  56. if len(excs.exceptions) != 1:
  57. raise
  58. exc = excs.exceptions[0]
  59. context = None if exc.__suppress_context__ else exc.__context__
  60. raise exc from exc.__cause__ or context
  61. def get_route_path(scope: Scope) -> str:
  62. path: str = scope["path"]
  63. root_path = scope.get("root_path", "")
  64. if not root_path:
  65. return path
  66. if not path.startswith(root_path):
  67. return path
  68. if path == root_path:
  69. return ""
  70. if path[len(root_path)] == "/":
  71. return path[len(root_path) :]
  72. return path