routing.py 29 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757
  1. from __future__ import annotations
  2. import contextlib
  3. import functools
  4. import inspect
  5. import re
  6. import traceback
  7. import types
  8. import warnings
  9. from collections.abc import Awaitable, Callable, Collection, Generator, Sequence
  10. from contextlib import AbstractAsyncContextManager, AbstractContextManager, asynccontextmanager
  11. from enum import Enum
  12. from re import Pattern
  13. from typing import Any, TypeVar
  14. from starlette._exception_handler import wrap_app_handling_exceptions
  15. from starlette._utils import get_route_path, is_async_callable
  16. from starlette.concurrency import run_in_threadpool
  17. from starlette.convertors import CONVERTOR_TYPES, Convertor
  18. from starlette.datastructures import URL, Headers, URLPath
  19. from starlette.exceptions import HTTPException, StarletteDeprecationWarning
  20. from starlette.middleware import Middleware
  21. from starlette.middleware.body_limit import RequestBodyLimitMiddleware
  22. from starlette.requests import Request
  23. from starlette.responses import PlainTextResponse, RedirectResponse, Response
  24. from starlette.types import ASGIApp, Lifespan, Receive, Scope, Send
  25. from starlette.websockets import WebSocket, WebSocketClose
  26. class NoMatchFound(Exception):
  27. """
  28. Raised by `.url_for(name, **path_params)` and `.url_path_for(name, **path_params)`
  29. if no matching route exists.
  30. """
  31. def __init__(self, name: str, path_params: dict[str, Any]) -> None:
  32. params = ", ".join(list(path_params.keys()))
  33. super().__init__(f'No route exists for name "{name}" and params "{params}".')
  34. class Match(Enum):
  35. NONE = 0
  36. PARTIAL = 1
  37. FULL = 2
  38. def request_response(
  39. func: Callable[[Request], Awaitable[Response] | Response],
  40. ) -> ASGIApp:
  41. """
  42. Takes a function or coroutine `func(request) -> response`,
  43. and returns an ASGI application.
  44. """
  45. f: Callable[[Request], Awaitable[Response]] = (
  46. func if is_async_callable(func) else functools.partial(run_in_threadpool, func) # type: ignore[assignment, call-arg]
  47. )
  48. async def app(scope: Scope, receive: Receive, send: Send) -> None:
  49. request = Request(scope, receive, send)
  50. async def app(scope: Scope, receive: Receive, send: Send) -> None:
  51. response = await f(request)
  52. await response(scope, receive, send)
  53. await wrap_app_handling_exceptions(app, request)(scope, receive, send)
  54. return app
  55. def websocket_session(
  56. func: Callable[[WebSocket], Awaitable[None]],
  57. ) -> ASGIApp:
  58. """
  59. Takes a coroutine `func(session)`, and returns an ASGI application.
  60. """
  61. # assert asyncio.iscoroutinefunction(func), "WebSocket endpoints must be async"
  62. async def app(scope: Scope, receive: Receive, send: Send) -> None:
  63. session = WebSocket(scope, receive=receive, send=send)
  64. async def app(scope: Scope, receive: Receive, send: Send) -> None:
  65. await func(session)
  66. await wrap_app_handling_exceptions(app, session)(scope, receive, send)
  67. return app
  68. def get_name(endpoint: Callable[..., Any]) -> str:
  69. return getattr(endpoint, "__name__", endpoint.__class__.__name__)
  70. def replace_params(
  71. path: str,
  72. param_convertors: dict[str, Convertor[Any]],
  73. path_params: dict[str, str],
  74. ) -> tuple[str, dict[str, str]]:
  75. for key, value in list(path_params.items()):
  76. if "{" + key + "}" in path:
  77. convertor = param_convertors[key]
  78. value = convertor.to_string(value)
  79. path = path.replace("{" + key + "}", value)
  80. path_params.pop(key)
  81. return path, path_params
  82. # Match parameters in URL paths, eg. '{param}', and '{param:int}'
  83. PARAM_REGEX = re.compile("{([a-zA-Z_][a-zA-Z0-9_]*)(:[a-zA-Z_][a-zA-Z0-9_]*)?}")
  84. def compile_path(
  85. path: str,
  86. ) -> tuple[Pattern[str], str, dict[str, Convertor[Any]]]:
  87. """
  88. Given a path string, like: "/{username:str}",
  89. or a host string, like: "{subdomain}.mydomain.org", return a three-tuple
  90. of (regex, format, {param_name:convertor}).
  91. regex: "/(?P<username>[^/]+)"
  92. format: "/{username}"
  93. convertors: {"username": StringConvertor()}
  94. """
  95. is_host = not path.startswith("/")
  96. path_regex = "^"
  97. path_format = ""
  98. duplicated_params: set[str] = set()
  99. idx = 0
  100. param_convertors = {}
  101. for match in PARAM_REGEX.finditer(path):
  102. param_name, convertor_type = match.groups("str")
  103. convertor_type = convertor_type.lstrip(":")
  104. assert convertor_type in CONVERTOR_TYPES, f"Unknown path convertor '{convertor_type}'"
  105. convertor = CONVERTOR_TYPES[convertor_type]
  106. path_regex += re.escape(path[idx : match.start()])
  107. path_regex += f"(?P<{param_name}>{convertor.regex})"
  108. path_format += path[idx : match.start()]
  109. path_format += "{%s}" % param_name
  110. if param_name in param_convertors:
  111. duplicated_params.add(param_name)
  112. param_convertors[param_name] = convertor
  113. idx = match.end()
  114. if duplicated_params:
  115. names = ", ".join(sorted(duplicated_params))
  116. ending = "s" if len(duplicated_params) > 1 else ""
  117. raise ValueError(f"Duplicated param name{ending} {names} at path {path}")
  118. if is_host:
  119. # Align with `Host.matches()` behavior, which ignores port.
  120. hostname = path[idx:].split(":")[0]
  121. path_regex += re.escape(hostname) + "$"
  122. else:
  123. path_regex += re.escape(path[idx:]) + "$"
  124. path_format += path[idx:]
  125. return re.compile(path_regex), path_format, param_convertors
  126. class BaseRoute:
  127. def matches(self, scope: Scope) -> tuple[Match, Scope]:
  128. raise NotImplementedError() # pragma: no cover
  129. def url_path_for(self, name: str, /, **path_params: Any) -> URLPath:
  130. raise NotImplementedError() # pragma: no cover
  131. async def handle(self, scope: Scope, receive: Receive, send: Send) -> None:
  132. raise NotImplementedError() # pragma: no cover
  133. async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
  134. """
  135. A route may be used in isolation as a stand-alone ASGI app.
  136. This is a somewhat contrived case, as they'll almost always be used
  137. within a Router, but could be useful for some tooling and minimal apps.
  138. """
  139. match, child_scope = self.matches(scope)
  140. if match == Match.NONE:
  141. if scope["type"] == "http":
  142. response = PlainTextResponse("Not Found", status_code=404)
  143. await response(scope, receive, send)
  144. elif scope["type"] == "websocket": # pragma: no branch
  145. websocket_close = WebSocketClose()
  146. await websocket_close(scope, receive, send)
  147. return
  148. scope.update(child_scope)
  149. await self.handle(scope, receive, send)
  150. class Route(BaseRoute):
  151. def __init__(
  152. self,
  153. path: str,
  154. endpoint: Callable[..., Any],
  155. *,
  156. methods: Collection[str] | None = None,
  157. name: str | None = None,
  158. include_in_schema: bool = True,
  159. middleware: Sequence[Middleware] | None = None,
  160. max_body_size: int | None = None,
  161. ) -> None:
  162. assert path.startswith("/"), "Routed paths must start with '/'"
  163. self.path = path
  164. self.endpoint = endpoint
  165. self.name = get_name(endpoint) if name is None else name
  166. self.include_in_schema = include_in_schema
  167. endpoint_handler = endpoint
  168. while isinstance(endpoint_handler, functools.partial):
  169. endpoint_handler = endpoint_handler.func
  170. if inspect.isfunction(endpoint_handler) or inspect.ismethod(endpoint_handler):
  171. # Endpoint is function or method. Treat it as `func(request) -> response`.
  172. self.app = request_response(endpoint)
  173. if methods is None:
  174. methods = ["GET"]
  175. else:
  176. # Endpoint is a class. Treat it as ASGI.
  177. self.app = endpoint
  178. if middleware is not None:
  179. for cls, args, kwargs in reversed(middleware):
  180. self.app = cls(self.app, *args, **kwargs)
  181. if max_body_size is not None:
  182. self.app = RequestBodyLimitMiddleware(self.app, max_body_size=max_body_size)
  183. if methods is None:
  184. self.methods = None
  185. else:
  186. self.methods = {method.upper() for method in methods}
  187. if "GET" in self.methods:
  188. self.methods.add("HEAD")
  189. self.path_regex, self.path_format, self.param_convertors = compile_path(path)
  190. def matches(self, scope: Scope) -> tuple[Match, Scope]:
  191. path_params: dict[str, Any]
  192. if scope["type"] == "http":
  193. route_path = get_route_path(scope)
  194. match = self.path_regex.match(route_path)
  195. if match:
  196. matched_params = match.groupdict()
  197. for key, value in matched_params.items():
  198. matched_params[key] = self.param_convertors[key].convert(value)
  199. path_params = dict(scope.get("path_params", {}))
  200. path_params.update(matched_params)
  201. child_scope = {"endpoint": self.endpoint, "path_params": path_params}
  202. if self.methods and scope["method"] not in self.methods:
  203. return Match.PARTIAL, child_scope
  204. else:
  205. return Match.FULL, child_scope
  206. return Match.NONE, {}
  207. def url_path_for(self, name: str, /, **path_params: Any) -> URLPath:
  208. seen_params = set(path_params.keys())
  209. expected_params = set(self.param_convertors.keys())
  210. if name != self.name or seen_params != expected_params:
  211. raise NoMatchFound(name, path_params)
  212. path, remaining_params = replace_params(self.path_format, self.param_convertors, path_params)
  213. assert not remaining_params
  214. return URLPath(path=path, protocol="http")
  215. async def handle(self, scope: Scope, receive: Receive, send: Send) -> None:
  216. if self.methods and scope["method"] not in self.methods:
  217. headers = {"Allow": ", ".join(self.methods)}
  218. if "app" in scope:
  219. raise HTTPException(status_code=405, headers=headers)
  220. else:
  221. response = PlainTextResponse("Method Not Allowed", status_code=405, headers=headers)
  222. await response(scope, receive, send)
  223. else:
  224. await self.app(scope, receive, send)
  225. def __eq__(self, other: Any) -> bool:
  226. return (
  227. isinstance(other, Route)
  228. and self.path == other.path
  229. and self.endpoint == other.endpoint
  230. and self.methods == other.methods
  231. )
  232. def __repr__(self) -> str:
  233. class_name = self.__class__.__name__
  234. methods = sorted(self.methods or [])
  235. path, name = self.path, self.name
  236. return f"{class_name}(path={path!r}, name={name!r}, methods={methods!r})"
  237. class WebSocketRoute(BaseRoute):
  238. def __init__(
  239. self,
  240. path: str,
  241. endpoint: Callable[..., Any],
  242. *,
  243. name: str | None = None,
  244. middleware: Sequence[Middleware] | None = None,
  245. ) -> None:
  246. assert path.startswith("/"), "Routed paths must start with '/'"
  247. self.path = path
  248. self.endpoint = endpoint
  249. self.name = get_name(endpoint) if name is None else name
  250. endpoint_handler = endpoint
  251. while isinstance(endpoint_handler, functools.partial):
  252. endpoint_handler = endpoint_handler.func
  253. if inspect.isfunction(endpoint_handler) or inspect.ismethod(endpoint_handler):
  254. # Endpoint is function or method. Treat it as `func(websocket)`.
  255. self.app = websocket_session(endpoint)
  256. else:
  257. # Endpoint is a class. Treat it as ASGI.
  258. self.app = endpoint
  259. if middleware is not None:
  260. for cls, args, kwargs in reversed(middleware):
  261. self.app = cls(self.app, *args, **kwargs)
  262. self.path_regex, self.path_format, self.param_convertors = compile_path(path)
  263. def matches(self, scope: Scope) -> tuple[Match, Scope]:
  264. path_params: dict[str, Any]
  265. if scope["type"] == "websocket":
  266. route_path = get_route_path(scope)
  267. match = self.path_regex.match(route_path)
  268. if match:
  269. matched_params = match.groupdict()
  270. for key, value in matched_params.items():
  271. matched_params[key] = self.param_convertors[key].convert(value)
  272. path_params = dict(scope.get("path_params", {}))
  273. path_params.update(matched_params)
  274. child_scope = {"endpoint": self.endpoint, "path_params": path_params}
  275. return Match.FULL, child_scope
  276. return Match.NONE, {}
  277. def url_path_for(self, name: str, /, **path_params: Any) -> URLPath:
  278. seen_params = set(path_params.keys())
  279. expected_params = set(self.param_convertors.keys())
  280. if name != self.name or seen_params != expected_params:
  281. raise NoMatchFound(name, path_params)
  282. path, remaining_params = replace_params(self.path_format, self.param_convertors, path_params)
  283. assert not remaining_params
  284. return URLPath(path=path, protocol="websocket")
  285. async def handle(self, scope: Scope, receive: Receive, send: Send) -> None:
  286. await self.app(scope, receive, send)
  287. def __eq__(self, other: Any) -> bool:
  288. return isinstance(other, WebSocketRoute) and self.path == other.path and self.endpoint == other.endpoint
  289. def __repr__(self) -> str:
  290. return f"{self.__class__.__name__}(path={self.path!r}, name={self.name!r})"
  291. class Mount(BaseRoute):
  292. def __init__(
  293. self,
  294. path: str,
  295. app: ASGIApp | None = None,
  296. routes: Sequence[BaseRoute] | None = None,
  297. name: str | None = None,
  298. *,
  299. middleware: Sequence[Middleware] | None = None,
  300. max_body_size: int | None = None,
  301. ) -> None:
  302. assert path == "" or path.startswith("/"), "Routed paths must start with '/'"
  303. assert app is not None or routes is not None, "Either 'app=...', or 'routes=' must be specified"
  304. self.path = path.rstrip("/")
  305. if app is not None:
  306. self._base_app: ASGIApp = app
  307. else:
  308. self._base_app = Router(routes=routes)
  309. self.app = self._base_app
  310. if middleware is not None:
  311. for cls, args, kwargs in reversed(middleware):
  312. self.app = cls(self.app, *args, **kwargs)
  313. if max_body_size is not None:
  314. self.app = RequestBodyLimitMiddleware(self.app, max_body_size=max_body_size)
  315. self.name = name
  316. self.path_regex, self.path_format, self.param_convertors = compile_path(self.path + "/{path:path}")
  317. @property
  318. def routes(self) -> list[BaseRoute]:
  319. return getattr(self._base_app, "routes", [])
  320. def matches(self, scope: Scope) -> tuple[Match, Scope]:
  321. path_params: dict[str, Any]
  322. if scope["type"] in ("http", "websocket"): # pragma: no branch
  323. root_path = scope.get("root_path", "")
  324. route_path = get_route_path(scope)
  325. match = self.path_regex.match(route_path)
  326. if match:
  327. matched_params = match.groupdict()
  328. for key, value in matched_params.items():
  329. matched_params[key] = self.param_convertors[key].convert(value)
  330. remaining_path = "/" + matched_params.pop("path")
  331. matched_path = route_path[: -len(remaining_path)]
  332. path_params = dict(scope.get("path_params", {}))
  333. path_params.update(matched_params)
  334. child_scope = {
  335. "path_params": path_params,
  336. # app_root_path will only be set at the top level scope,
  337. # initialized with the (optional) value of a root_path
  338. # set above/before Starlette. And even though any
  339. # mount will have its own child scope with its own respective
  340. # root_path, the app_root_path will always be available in all
  341. # the child scopes with the same top level value because it's
  342. # set only once here with a default, any other child scope will
  343. # just inherit that app_root_path default value stored in the
  344. # scope. All this is needed to support Request.url_for(), as it
  345. # uses the app_root_path to build the URL path.
  346. "app_root_path": scope.get("app_root_path", root_path),
  347. "root_path": root_path + matched_path,
  348. "endpoint": self.app,
  349. }
  350. return Match.FULL, child_scope
  351. return Match.NONE, {}
  352. def url_path_for(self, name: str, /, **path_params: Any) -> URLPath:
  353. if self.name is not None and name == self.name and "path" in path_params:
  354. # 'name' matches "<mount_name>".
  355. path_params["path"] = path_params["path"].lstrip("/")
  356. path, remaining_params = replace_params(self.path_format, self.param_convertors, path_params)
  357. if not remaining_params:
  358. return URLPath(path=path)
  359. elif self.name is None or name.startswith(self.name + ":"):
  360. if self.name is None:
  361. # No mount name.
  362. remaining_name = name
  363. else:
  364. # 'name' matches "<mount_name>:<child_name>".
  365. remaining_name = name[len(self.name) + 1 :]
  366. path_kwarg = path_params.get("path")
  367. path_params["path"] = ""
  368. path_prefix, remaining_params = replace_params(self.path_format, self.param_convertors, path_params)
  369. if path_kwarg is not None:
  370. remaining_params["path"] = path_kwarg
  371. for route in self.routes or []:
  372. try:
  373. url = route.url_path_for(remaining_name, **remaining_params)
  374. return URLPath(path=path_prefix.rstrip("/") + str(url), protocol=url.protocol)
  375. except NoMatchFound:
  376. pass
  377. raise NoMatchFound(name, path_params)
  378. async def handle(self, scope: Scope, receive: Receive, send: Send) -> None:
  379. await self.app(scope, receive, send)
  380. def __eq__(self, other: Any) -> bool:
  381. return isinstance(other, Mount) and self.path == other.path and self.app == other.app
  382. def __repr__(self) -> str:
  383. class_name = self.__class__.__name__
  384. name = self.name or ""
  385. return f"{class_name}(path={self.path!r}, name={name!r}, app={self.app!r})"
  386. class Host(BaseRoute):
  387. def __init__(self, host: str, app: ASGIApp, name: str | None = None) -> None:
  388. assert not host.startswith("/"), "Host must not start with '/'"
  389. self.host = host
  390. self.app = app
  391. self.name = name
  392. self.host_regex, self.host_format, self.param_convertors = compile_path(host)
  393. @property
  394. def routes(self) -> list[BaseRoute]:
  395. return getattr(self.app, "routes", [])
  396. def matches(self, scope: Scope) -> tuple[Match, Scope]:
  397. if scope["type"] in ("http", "websocket"): # pragma:no branch
  398. headers = Headers(scope=scope)
  399. host = headers.get("host", "").split(":")[0]
  400. match = self.host_regex.match(host)
  401. if match:
  402. matched_params = match.groupdict()
  403. for key, value in matched_params.items():
  404. matched_params[key] = self.param_convertors[key].convert(value)
  405. path_params = dict(scope.get("path_params", {}))
  406. path_params.update(matched_params)
  407. child_scope = {"path_params": path_params, "endpoint": self.app}
  408. return Match.FULL, child_scope
  409. return Match.NONE, {}
  410. def url_path_for(self, name: str, /, **path_params: Any) -> URLPath:
  411. if self.name is not None and name == self.name and "path" in path_params:
  412. # 'name' matches "<mount_name>".
  413. path = path_params.pop("path")
  414. host, remaining_params = replace_params(self.host_format, self.param_convertors, path_params)
  415. if not remaining_params:
  416. return URLPath(path=path, host=host)
  417. elif self.name is None or name.startswith(self.name + ":"):
  418. if self.name is None:
  419. # No mount name.
  420. remaining_name = name
  421. else:
  422. # 'name' matches "<mount_name>:<child_name>".
  423. remaining_name = name[len(self.name) + 1 :]
  424. host, remaining_params = replace_params(self.host_format, self.param_convertors, path_params)
  425. for route in self.routes or []:
  426. try:
  427. url = route.url_path_for(remaining_name, **remaining_params)
  428. return URLPath(path=str(url), protocol=url.protocol, host=host)
  429. except NoMatchFound:
  430. pass
  431. raise NoMatchFound(name, path_params)
  432. async def handle(self, scope: Scope, receive: Receive, send: Send) -> None:
  433. await self.app(scope, receive, send)
  434. def __eq__(self, other: Any) -> bool:
  435. return isinstance(other, Host) and self.host == other.host and self.app == other.app
  436. def __repr__(self) -> str:
  437. class_name = self.__class__.__name__
  438. name = self.name or ""
  439. return f"{class_name}(host={self.host!r}, name={name!r}, app={self.app!r})"
  440. _T = TypeVar("_T")
  441. class _AsyncLiftContextManager(AbstractAsyncContextManager[_T]):
  442. def __init__(self, cm: AbstractContextManager[_T]):
  443. self._cm = cm
  444. async def __aenter__(self) -> _T:
  445. return self._cm.__enter__()
  446. async def __aexit__(
  447. self,
  448. exc_type: type[BaseException] | None,
  449. exc_value: BaseException | None,
  450. traceback: types.TracebackType | None,
  451. ) -> bool | None:
  452. return self._cm.__exit__(exc_type, exc_value, traceback)
  453. def _wrap_gen_lifespan_context(
  454. lifespan_context: Callable[[Any], Generator[Any, Any, Any]],
  455. ) -> Callable[[Any], AbstractAsyncContextManager[Any]]:
  456. cmgr = contextlib.contextmanager(lifespan_context)
  457. @functools.wraps(cmgr)
  458. def wrapper(app: Any) -> _AsyncLiftContextManager[Any]:
  459. return _AsyncLiftContextManager(cmgr(app))
  460. return wrapper
  461. class _DefaultLifespan:
  462. def __init__(self, router: Router):
  463. self._router = router
  464. async def __aenter__(self) -> None:
  465. pass
  466. async def __aexit__(self, *exc_info: object) -> None:
  467. pass
  468. def __call__(self: _T, app: object) -> _T:
  469. return self
  470. class Router:
  471. def __init__(
  472. self,
  473. routes: Sequence[BaseRoute] | None = None,
  474. redirect_slashes: bool = True,
  475. default: ASGIApp | None = None,
  476. # the generic to Lifespan[AppType] is the type of the top level application
  477. # which the router cannot know statically, so we use Any
  478. lifespan: Lifespan[Any] | None = None,
  479. *,
  480. middleware: Sequence[Middleware] | None = None,
  481. max_body_size: int | None = None,
  482. ) -> None:
  483. self.routes = [] if routes is None else list(routes)
  484. self.redirect_slashes = redirect_slashes
  485. self.default = self.not_found if default is None else default
  486. if lifespan is None:
  487. self.lifespan_context: Lifespan[Any] = _DefaultLifespan(self)
  488. elif inspect.isasyncgenfunction(lifespan):
  489. warnings.warn(
  490. "async generator function lifespans are deprecated, "
  491. "use an @contextlib.asynccontextmanager function instead",
  492. StarletteDeprecationWarning,
  493. )
  494. self.lifespan_context = asynccontextmanager(lifespan)
  495. elif inspect.isgeneratorfunction(lifespan):
  496. warnings.warn(
  497. "generator function lifespans are deprecated, use an @contextlib.asynccontextmanager function instead",
  498. StarletteDeprecationWarning,
  499. )
  500. self.lifespan_context = _wrap_gen_lifespan_context(lifespan)
  501. else:
  502. self.lifespan_context = lifespan
  503. self.middleware_stack = self.app
  504. if middleware:
  505. for cls, args, kwargs in reversed(middleware):
  506. self.middleware_stack = cls(self.middleware_stack, *args, **kwargs)
  507. if max_body_size is not None:
  508. self.middleware_stack = RequestBodyLimitMiddleware(self.middleware_stack, max_body_size=max_body_size)
  509. async def not_found(self, scope: Scope, receive: Receive, send: Send) -> None:
  510. if scope["type"] == "websocket":
  511. websocket_close = WebSocketClose()
  512. await websocket_close(scope, receive, send)
  513. return
  514. # If we're running inside a starlette application then raise an
  515. # exception, so that the configurable exception handler can deal with
  516. # returning the response. For plain ASGI apps, just return the response.
  517. if "app" in scope:
  518. raise HTTPException(status_code=404)
  519. else:
  520. response = PlainTextResponse("Not Found", status_code=404)
  521. await response(scope, receive, send)
  522. def url_path_for(self, name: str, /, **path_params: Any) -> URLPath:
  523. for route in self.routes:
  524. try:
  525. return route.url_path_for(name, **path_params)
  526. except NoMatchFound:
  527. pass
  528. raise NoMatchFound(name, path_params)
  529. async def lifespan(self, scope: Scope, receive: Receive, send: Send) -> None:
  530. """
  531. Handle ASGI lifespan messages, which allows us to manage application
  532. startup and shutdown events.
  533. """
  534. started = False
  535. app: Any = scope.get("app")
  536. await receive()
  537. try:
  538. async with self.lifespan_context(app) as maybe_state:
  539. if maybe_state is not None:
  540. if "state" not in scope:
  541. raise RuntimeError('The server does not support "state" in the lifespan scope.')
  542. scope["state"].update(maybe_state)
  543. await send({"type": "lifespan.startup.complete"})
  544. started = True
  545. await receive()
  546. except BaseException:
  547. exc_text = traceback.format_exc()
  548. if started:
  549. await send({"type": "lifespan.shutdown.failed", "message": exc_text})
  550. else:
  551. await send({"type": "lifespan.startup.failed", "message": exc_text})
  552. raise
  553. else:
  554. await send({"type": "lifespan.shutdown.complete"})
  555. async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
  556. """
  557. The main entry point to the Router class.
  558. """
  559. await self.middleware_stack(scope, receive, send)
  560. async def app(self, scope: Scope, receive: Receive, send: Send) -> None:
  561. assert scope["type"] in ("http", "websocket", "lifespan")
  562. if "router" not in scope:
  563. scope["router"] = self
  564. if scope["type"] == "lifespan":
  565. await self.lifespan(scope, receive, send)
  566. return
  567. partial = None
  568. for route in self.routes:
  569. # Determine if any route matches the incoming scope,
  570. # and hand over to the matching route if found.
  571. match, child_scope = route.matches(scope)
  572. if match == Match.FULL:
  573. scope.update(child_scope)
  574. await route.handle(scope, receive, send)
  575. return
  576. elif match == Match.PARTIAL and partial is None:
  577. partial = route
  578. partial_scope = child_scope
  579. if partial is not None:
  580. #  Handle partial matches. These are cases where an endpoint is
  581. # able to handle the request, but is not a preferred option.
  582. # We use this in particular to deal with "405 Method Not Allowed".
  583. scope.update(partial_scope)
  584. await partial.handle(scope, receive, send)
  585. return
  586. route_path = get_route_path(scope)
  587. if scope["type"] == "http" and self.redirect_slashes and route_path != "/":
  588. redirect_scope = dict(scope)
  589. if route_path.endswith("/"):
  590. redirect_scope["path"] = redirect_scope["path"].rstrip("/")
  591. else:
  592. redirect_scope["path"] = redirect_scope["path"] + "/"
  593. for route in self.routes:
  594. match, child_scope = route.matches(redirect_scope)
  595. if match != Match.NONE:
  596. redirect_url = URL(scope=redirect_scope)
  597. response = RedirectResponse(url=str(redirect_url))
  598. await response(scope, receive, send)
  599. return
  600. await self.default(scope, receive, send)
  601. def __eq__(self, other: Any) -> bool:
  602. return isinstance(other, Router) and self.routes == other.routes
  603. def mount(self, path: str, app: ASGIApp, name: str | None = None) -> None: # pragma: no cover
  604. route = Mount(path, app=app, name=name)
  605. self.routes.append(route)
  606. def host(self, host: str, app: ASGIApp, name: str | None = None) -> None: # pragma: no cover
  607. route = Host(host, app=app, name=name)
  608. self.routes.append(route)
  609. def add_route(
  610. self,
  611. path: str,
  612. endpoint: Callable[[Request], Awaitable[Response] | Response],
  613. methods: Collection[str] | None = None,
  614. name: str | None = None,
  615. include_in_schema: bool = True,
  616. ) -> None: # pragma: no cover
  617. route = Route(
  618. path,
  619. endpoint=endpoint,
  620. methods=methods,
  621. name=name,
  622. include_in_schema=include_in_schema,
  623. )
  624. self.routes.append(route)
  625. def add_websocket_route(
  626. self,
  627. path: str,
  628. endpoint: Callable[[WebSocket], Awaitable[None]],
  629. name: str | None = None,
  630. ) -> None: # pragma: no cover
  631. route = WebSocketRoute(path, endpoint=endpoint, name=name)
  632. self.routes.append(route)