cors.py 7.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179
  1. from __future__ import annotations
  2. import functools
  3. import re
  4. from collections.abc import Sequence
  5. from starlette.datastructures import Headers, MutableHeaders
  6. from starlette.responses import PlainTextResponse, Response
  7. from starlette.types import ASGIApp, Message, Receive, Scope, Send
  8. ALL_METHODS = ("DELETE", "GET", "HEAD", "OPTIONS", "PATCH", "POST", "PUT")
  9. SAFELISTED_HEADERS = {"Accept", "Accept-Language", "Content-Language", "Content-Type"}
  10. class CORSMiddleware:
  11. def __init__(
  12. self,
  13. app: ASGIApp,
  14. allow_origins: Sequence[str] = (),
  15. allow_methods: Sequence[str] = ("GET",),
  16. allow_headers: Sequence[str] = (),
  17. allow_credentials: bool = False,
  18. allow_origin_regex: str | None = None,
  19. allow_private_network: bool = False,
  20. expose_headers: Sequence[str] = (),
  21. max_age: int = 600,
  22. ) -> None:
  23. if "*" in allow_methods:
  24. allow_methods = ALL_METHODS
  25. compiled_allow_origin_regex = None
  26. if allow_origin_regex is not None:
  27. compiled_allow_origin_regex = re.compile(allow_origin_regex)
  28. allow_all_origins = "*" in allow_origins
  29. allow_all_headers = "*" in allow_headers
  30. preflight_explicit_allow_origin = not allow_all_origins or allow_credentials
  31. simple_headers: dict[str, str] = {}
  32. if allow_all_origins:
  33. simple_headers["Access-Control-Allow-Origin"] = "*"
  34. if allow_credentials:
  35. simple_headers["Access-Control-Allow-Credentials"] = "true"
  36. if expose_headers:
  37. simple_headers["Access-Control-Expose-Headers"] = ", ".join(expose_headers)
  38. preflight_headers: dict[str, str] = {}
  39. if preflight_explicit_allow_origin:
  40. # The origin value will be set in preflight_response() if it is allowed.
  41. preflight_headers["Vary"] = "Origin"
  42. else:
  43. preflight_headers["Access-Control-Allow-Origin"] = "*"
  44. preflight_headers.update(
  45. {
  46. "Access-Control-Allow-Methods": ", ".join(allow_methods),
  47. "Access-Control-Max-Age": str(max_age),
  48. }
  49. )
  50. allow_headers = sorted(SAFELISTED_HEADERS | set(allow_headers))
  51. if allow_headers and not allow_all_headers:
  52. preflight_headers["Access-Control-Allow-Headers"] = ", ".join(allow_headers)
  53. if allow_credentials:
  54. preflight_headers["Access-Control-Allow-Credentials"] = "true"
  55. self.app = app
  56. self.allow_origins = allow_origins
  57. self.allow_methods = allow_methods
  58. self.allow_headers = [h.lower() for h in allow_headers]
  59. self.allow_all_origins = allow_all_origins
  60. self.allow_all_headers = allow_all_headers
  61. self.allow_credentials = allow_credentials
  62. self.preflight_explicit_allow_origin = preflight_explicit_allow_origin
  63. self.allow_origin_regex = compiled_allow_origin_regex
  64. self.allow_private_network = allow_private_network
  65. self.simple_headers = simple_headers
  66. self.preflight_headers = preflight_headers
  67. async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
  68. if scope["type"] != "http": # pragma: no cover
  69. await self.app(scope, receive, send)
  70. return
  71. method = scope["method"]
  72. headers = Headers(scope=scope)
  73. origin = headers.get("origin")
  74. if origin is None:
  75. await self.app(scope, receive, send)
  76. return
  77. if method == "OPTIONS" and "access-control-request-method" in headers:
  78. response = self.preflight_response(request_headers=headers)
  79. await response(scope, receive, send)
  80. return
  81. await self.simple_response(scope, receive, send, request_headers=headers)
  82. def is_allowed_origin(self, origin: str) -> bool:
  83. if self.allow_all_origins:
  84. return True
  85. if self.allow_origin_regex is not None and self.allow_origin_regex.fullmatch(origin):
  86. return True
  87. return origin in self.allow_origins
  88. def preflight_response(self, request_headers: Headers) -> Response:
  89. requested_origin = request_headers["origin"]
  90. requested_method = request_headers["access-control-request-method"]
  91. requested_headers = request_headers.get("access-control-request-headers")
  92. requested_private_network = request_headers.get("access-control-request-private-network")
  93. headers = dict(self.preflight_headers)
  94. failures: list[str] = []
  95. if self.is_allowed_origin(origin=requested_origin):
  96. if self.preflight_explicit_allow_origin:
  97. # The "else" case is already accounted for in self.preflight_headers
  98. # and the value would be "*".
  99. headers["Access-Control-Allow-Origin"] = requested_origin
  100. else:
  101. failures.append("origin")
  102. if requested_method not in self.allow_methods:
  103. failures.append("method")
  104. # If we allow all headers, then we have to mirror back any requested
  105. # headers in the response.
  106. if self.allow_all_headers and requested_headers is not None:
  107. headers["Access-Control-Allow-Headers"] = requested_headers
  108. elif requested_headers is not None:
  109. for header in [h.lower() for h in requested_headers.split(",")]:
  110. if header.strip() not in self.allow_headers:
  111. failures.append("headers")
  112. break
  113. if requested_private_network is not None:
  114. if self.allow_private_network:
  115. headers["Access-Control-Allow-Private-Network"] = "true"
  116. else:
  117. failures.append("private-network")
  118. # We don't strictly need to use 400 responses here, since its up to
  119. # the browser to enforce the CORS policy, but its more informative
  120. # if we do.
  121. if failures:
  122. failure_text = "Disallowed CORS " + ", ".join(failures)
  123. return PlainTextResponse(failure_text, status_code=400, headers=headers)
  124. return PlainTextResponse("OK", status_code=200, headers=headers)
  125. async def simple_response(self, scope: Scope, receive: Receive, send: Send, request_headers: Headers) -> None:
  126. send = functools.partial(self.send, send=send, request_headers=request_headers)
  127. await self.app(scope, receive, send)
  128. async def send(self, message: Message, send: Send, request_headers: Headers) -> None:
  129. if message["type"] != "http.response.start":
  130. await send(message)
  131. return
  132. message.setdefault("headers", [])
  133. headers = MutableHeaders(scope=message)
  134. headers.update(self.simple_headers)
  135. origin = request_headers["Origin"]
  136. # If credentials are allowed, then we must respond with the specific origin instead of '*'.
  137. if self.allow_all_origins and self.allow_credentials:
  138. self.allow_explicit_origin(headers, origin)
  139. # If we only allow specific origins, then we have to mirror back the Origin header in the response.
  140. elif not self.allow_all_origins and self.is_allowed_origin(origin=origin):
  141. self.allow_explicit_origin(headers, origin)
  142. await send(message)
  143. @staticmethod
  144. def allow_explicit_origin(headers: MutableHeaders, origin: str) -> None:
  145. headers["Access-Control-Allow-Origin"] = origin
  146. headers.add_vary_header("Origin")