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

87 statements  

1from __future__ import annotations 

2 

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 

11 

12import anyio.abc 

13 

14from starlette.types import Scope 

15 

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 

21 

22 from typing_extensions import TypeIs 

23 

24if sys.version_info < (3, 11): # pragma: no cover 

25 from exceptiongroup import BaseExceptionGroup 

26 

27 

28T = TypeVar("T") 

29AwaitableCallable = Callable[..., Awaitable[T]] 

30 

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) 

37 

38 

39@dataclass(frozen=True, slots=True) 

40class ParsedHost: 

41 host: str 

42 port: str | None 

43 

44 @property 

45 def authority(self) -> str: 

46 return self.host if self.port is None else f"{self.host}:{self.port}" 

47 

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 

54 

55 

56@overload 

57def is_async_callable(obj: AwaitableCallable[T]) -> TypeIs[AwaitableCallable[T]]: ... 

58 

59 

60@overload 

61def is_async_callable(obj: Any) -> TypeIs[AwaitableCallable[Any]]: ... 

62 

63 

64def is_async_callable(obj: Any) -> Any: 

65 while isinstance(obj, functools.partial): 

66 obj = obj.func 

67 

68 return iscoroutinefunction(obj) or (callable(obj) and iscoroutinefunction(obj.__call__)) 

69 

70 

71T_co = TypeVar("T_co", covariant=True) 

72 

73 

74class AwaitableOrContextManager( 

75 Awaitable[T_co], AbstractAsyncContextManager[T_co], Protocol[T_co] 

76): ... # pragma: no branch 

77 

78 

79class SupportsAsyncClose(Protocol): 

80 async def close(self) -> None: ... # pragma: no cover 

81 

82 

83SupportsAsyncCloseType = TypeVar("SupportsAsyncCloseType", bound=SupportsAsyncClose, covariant=False) 

84 

85 

86class AwaitableOrContextManagerWrapper(Generic[SupportsAsyncCloseType]): 

87 __slots__ = ("aw", "entered") 

88 

89 def __init__(self, aw: Awaitable[SupportsAsyncCloseType]) -> None: 

90 self.aw = aw 

91 

92 def __await__(self) -> Generator[Any, None, SupportsAsyncCloseType]: 

93 return self.aw.__await__() 

94 

95 async def __aenter__(self) -> SupportsAsyncCloseType: 

96 self.entered = await self.aw 

97 return self.entered 

98 

99 async def __aexit__(self, *args: Any) -> None | bool: 

100 await self.entered.close() 

101 return None 

102 

103 

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 

112 

113 exc = excs.exceptions[0] 

114 context = None if exc.__suppress_context__ else exc.__context__ 

115 raise exc from exc.__cause__ or context 

116 

117 

118def parse_host_header(host_header: str | None) -> ParsedHost | None: 

119 """Parse `host_header` into its host and port components. 

120 

121 The host preserves brackets around IP literals. Invalid headers produce `None`. 

122 """ 

123 if host_header is None: 

124 return None 

125 

126 match = _HOST_RE.fullmatch(host_header) 

127 if match is None: 

128 return None 

129 

130 ipv6 = match["ipv6"] 

131 if ipv6 is not None: 

132 try: 

133 IPv6Address(ipv6) 

134 except AddressValueError: 

135 return None 

136 

137 return ParsedHost(match["host"], match["port"]) 

138 

139 

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 

145 

146 if not path.startswith(root_path): 

147 return path 

148 

149 if path == root_path: 

150 return "" 

151 

152 if path[len(root_path)] == "/": 

153 return path[len(root_path) :] 

154 

155 return path