Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/starlette/_utils.py: 47%
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
1from __future__ import annotations
3import functools
4import re
5import sys
6from collections.abc import AsyncGenerator, Awaitable, Callable, Generator
7from contextlib import AbstractAsyncContextManager, asynccontextmanager
8from dataclasses import dataclass
9from ipaddress import AddressValueError, IPv6Address
10from typing import Any, Generic, Protocol, TypeVar, overload
12import anyio.abc
14from starlette.types import Scope
16if sys.version_info >= (3, 13): # pragma: no cover
17 from inspect import iscoroutinefunction
18 from typing import TypeIs
19else: # pragma: no cover
20 from asyncio import iscoroutinefunction
22 from typing_extensions import TypeIs
24if sys.version_info < (3, 11): # pragma: no cover
25 from exceptiongroup import BaseExceptionGroup
28T = TypeVar("T")
29AwaitableCallable = Callable[..., Awaitable[T]]
31# Reject characters that could make a Host header change the URL path or authority.
32_HOST_RE = re.compile(
33 r"^(?P<host>[a-z0-9._~%!$&'()*+,;=-]+|\[(?:(?P<ipv6>[a-f0-9]*:[a-f0-9.:]+)|"
34 r"(?-i:v)[a-f0-9]+\.[a-z0-9._~!$&'()*+,;=:-]+)\])(?::(?P<port>[0-9]+))?$",
35 re.IGNORECASE,
36)
39@dataclass(frozen=True, slots=True)
40class ParsedHost:
41 host: str
42 port: str | None
44 @property
45 def authority(self) -> str:
46 return self.host if self.port is None else f"{self.host}:{self.port}"
48 @property
49 def is_valid_port(self) -> bool:
50 if self.port is None:
51 return True
52 port = self.port.lstrip("0")
53 return len(port) <= 5 and int(port or "0") <= 65535
56@overload
57def is_async_callable(obj: AwaitableCallable[T]) -> TypeIs[AwaitableCallable[T]]: ...
60@overload
61def is_async_callable(obj: Any) -> TypeIs[AwaitableCallable[Any]]: ...
64def is_async_callable(obj: Any) -> Any:
65 while isinstance(obj, functools.partial):
66 obj = obj.func
68 return iscoroutinefunction(obj) or (callable(obj) and iscoroutinefunction(obj.__call__))
71T_co = TypeVar("T_co", covariant=True)
74class AwaitableOrContextManager(
75 Awaitable[T_co], AbstractAsyncContextManager[T_co], Protocol[T_co]
76): ... # pragma: no branch
79class SupportsAsyncClose(Protocol):
80 async def close(self) -> None: ... # pragma: no cover
83SupportsAsyncCloseType = TypeVar("SupportsAsyncCloseType", bound=SupportsAsyncClose, covariant=False)
86class AwaitableOrContextManagerWrapper(Generic[SupportsAsyncCloseType]):
87 __slots__ = ("aw", "entered")
89 def __init__(self, aw: Awaitable[SupportsAsyncCloseType]) -> None:
90 self.aw = aw
92 def __await__(self) -> Generator[Any, None, SupportsAsyncCloseType]:
93 return self.aw.__await__()
95 async def __aenter__(self) -> SupportsAsyncCloseType:
96 self.entered = await self.aw
97 return self.entered
99 async def __aexit__(self, *args: Any) -> None | bool:
100 await self.entered.close()
101 return None
104@asynccontextmanager
105async def create_collapsing_task_group() -> AsyncGenerator[anyio.abc.TaskGroup, None]:
106 try:
107 async with anyio.create_task_group() as tg:
108 yield tg
109 except BaseExceptionGroup as excs:
110 if len(excs.exceptions) != 1:
111 raise
113 exc = excs.exceptions[0]
114 context = None if exc.__suppress_context__ else exc.__context__
115 raise exc from exc.__cause__ or context
118def parse_host_header(host_header: str | None) -> ParsedHost | None:
119 """Parse `host_header` into its host and port components.
121 The host preserves brackets around IP literals. Invalid headers produce `None`.
122 """
123 if host_header is None:
124 return None
126 match = _HOST_RE.fullmatch(host_header)
127 if match is None:
128 return None
130 ipv6 = match["ipv6"]
131 if ipv6 is not None:
132 try:
133 IPv6Address(ipv6)
134 except AddressValueError:
135 return None
137 return ParsedHost(match["host"], match["port"])
140def get_route_path(scope: Scope) -> str:
141 path: str = scope["path"]
142 root_path = scope.get("root_path", "")
143 if not root_path:
144 return path
146 if not path.startswith(root_path):
147 return path
149 if path == root_path:
150 return ""
152 if path[len(root_path)] == "/":
153 return path[len(root_path) :]
155 return path