gzip.py 8.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222
  1. from __future__ import annotations
  2. import zlib
  3. from typing import NoReturn
  4. import anyio.lowlevel
  5. import anyio.to_thread
  6. from starlette.datastructures import Headers, MutableHeaders
  7. from starlette.types import ASGIApp, Message, Receive, Scope, Send
  8. # TODO(v2): We should rename `DEFAULT_EXCLUDED_CONTENT_TYPES` to `DEFAULT_EXCLUDE_CONTENT_TYPES`.
  9. DEFAULT_EXCLUDED_CONTENT_TYPES = (
  10. "application/gzip",
  11. "application/x-gzip",
  12. "application/zip",
  13. "audio/*",
  14. "font/woff",
  15. "font/woff2",
  16. "image/avif",
  17. "image/gif",
  18. "image/jpeg",
  19. "image/png",
  20. "image/webp",
  21. "text/event-stream",
  22. "video/*",
  23. )
  24. _gzip_capacity_limiter: anyio.lowlevel.RunVar[anyio.CapacityLimiter] = anyio.lowlevel.RunVar("_gzip_capacity_limiter")
  25. def _get_gzip_capacity_limiter() -> anyio.CapacityLimiter:
  26. """Return the capacity limiter used for worker-thread GZip compression."""
  27. try:
  28. return _gzip_capacity_limiter.get()
  29. except LookupError:
  30. # Keep gzip compression isolated from AnyIO's default worker-thread
  31. # capacity limiter while matching its default concurrency.
  32. limiter = anyio.CapacityLimiter(40)
  33. _gzip_capacity_limiter.set(limiter)
  34. return limiter
  35. class GZipMiddleware:
  36. def __init__(
  37. self,
  38. app: ASGIApp,
  39. minimum_size: int = 500,
  40. compresslevel: int = 9,
  41. thread_minimum_size: int = 128 * 1024, # 128 KiB
  42. *,
  43. exclude_content_types: tuple[str, ...] = DEFAULT_EXCLUDED_CONTENT_TYPES,
  44. ) -> None:
  45. self.app = app
  46. self.minimum_size = minimum_size
  47. self.compresslevel = compresslevel
  48. self.thread_minimum_size = thread_minimum_size
  49. self.exclude_content_types = _normalize_content_types(exclude_content_types)
  50. async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
  51. if scope["type"] != "http": # pragma: no cover
  52. await self.app(scope, receive, send)
  53. return
  54. headers = Headers(scope=scope)
  55. responder: ASGIApp
  56. if "gzip" in headers.get("Accept-Encoding", ""):
  57. responder = GZipResponder(
  58. self.app,
  59. self.minimum_size,
  60. compresslevel=self.compresslevel,
  61. thread_minimum_size=self.thread_minimum_size,
  62. exclude_content_types=self.exclude_content_types,
  63. )
  64. else:
  65. responder = IdentityResponder(self.app, self.minimum_size, exclude_content_types=self.exclude_content_types)
  66. await responder(scope, receive, send)
  67. class IdentityResponder:
  68. content_encoding: str
  69. def __init__(
  70. self,
  71. app: ASGIApp,
  72. minimum_size: int,
  73. *,
  74. exclude_content_types: tuple[str, ...] = DEFAULT_EXCLUDED_CONTENT_TYPES,
  75. ) -> None:
  76. self.app = app
  77. self.minimum_size = minimum_size
  78. self.exclude_content_types = _normalize_content_types(exclude_content_types)
  79. self.send: Send = unattached_send
  80. self.initial_message: Message = {}
  81. self.started = False
  82. self.content_encoding_set = False
  83. self.content_type_is_excluded = False
  84. self.partial_response = False
  85. async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
  86. self.send = send
  87. await self.app(scope, receive, self.send_with_compression)
  88. async def send_with_compression(self, message: Message) -> None:
  89. message_type = message["type"]
  90. if message_type == "http.response.start":
  91. # Don't send the initial message until we've determined how to
  92. # modify the outgoing headers correctly.
  93. self.initial_message = message
  94. headers = Headers(raw=self.initial_message["headers"])
  95. self.content_encoding_set = "content-encoding" in headers
  96. self.partial_response = message["status"] == 206
  97. media_type = headers.get("content-type", "").partition(";")[0].strip().lower()
  98. media_types = {media_type, media_type.partition("/")[0] + "/*"}
  99. self.content_type_is_excluded = not media_types.isdisjoint(self.exclude_content_types)
  100. elif message_type == "http.response.body" and (
  101. self.content_encoding_set or self.partial_response or self.content_type_is_excluded
  102. ):
  103. if not self.started:
  104. self.started = True
  105. await self.send(self.initial_message)
  106. await self.send(message)
  107. elif message_type == "http.response.body" and not self.started:
  108. self.started = True
  109. body = message.get("body", b"")
  110. more_body = message.get("more_body", False)
  111. if len(body) < self.minimum_size and not more_body:
  112. # Don't apply compression to small outgoing responses.
  113. await self.send(self.initial_message)
  114. await self.send(message)
  115. elif not more_body:
  116. # Standard response.
  117. body = await self.apply_compression(body, more_body=False)
  118. headers = MutableHeaders(raw=self.initial_message["headers"])
  119. headers.add_vary_header("Accept-Encoding")
  120. if body != message["body"]:
  121. headers["Content-Encoding"] = self.content_encoding
  122. headers["Content-Length"] = str(len(body))
  123. message["body"] = body
  124. await self.send(self.initial_message)
  125. await self.send(message)
  126. else:
  127. # Initial body in streaming response.
  128. body = await self.apply_compression(body, more_body=True)
  129. headers = MutableHeaders(raw=self.initial_message["headers"])
  130. headers.add_vary_header("Accept-Encoding")
  131. if body != message["body"]:
  132. headers["Content-Encoding"] = self.content_encoding
  133. del headers["Content-Length"]
  134. message["body"] = body
  135. await self.send(self.initial_message)
  136. await self.send(message)
  137. elif message_type == "http.response.body":
  138. # Remaining body in streaming response.
  139. body = message.get("body", b"")
  140. more_body = message.get("more_body", False)
  141. message["body"] = await self.apply_compression(body, more_body=more_body)
  142. await self.send(message)
  143. elif message_type == "http.response.pathsend": # pragma: no branch
  144. # Don't apply GZip to pathsend responses
  145. await self.send(self.initial_message)
  146. await self.send(message)
  147. async def apply_compression(self, body: bytes, *, more_body: bool) -> bytes:
  148. """Apply compression on the response body.
  149. If more_body is False, the compression stream is finalized. Compression
  150. resources are only allocated once a body is actually compressed.
  151. """
  152. return body
  153. class GZipResponder(IdentityResponder):
  154. content_encoding = "gzip"
  155. def __init__(
  156. self,
  157. app: ASGIApp,
  158. minimum_size: int,
  159. compresslevel: int = 9,
  160. *,
  161. thread_minimum_size: int = 128 * 1024, # 128 KiB
  162. exclude_content_types: tuple[str, ...] = DEFAULT_EXCLUDED_CONTENT_TYPES,
  163. ) -> None:
  164. super().__init__(app, minimum_size, exclude_content_types=exclude_content_types)
  165. self.compresslevel = compresslevel
  166. self.thread_minimum_size = thread_minimum_size
  167. self._compressor: zlib._Compress | None = None
  168. @property
  169. def compressor(self) -> zlib._Compress:
  170. if self._compressor is None:
  171. self._compressor = zlib.compressobj(self.compresslevel, zlib.DEFLATED, 16 + zlib.MAX_WBITS)
  172. return self._compressor
  173. async def apply_compression(self, body: bytes, *, more_body: bool) -> bytes:
  174. if len(body) >= self.thread_minimum_size:
  175. # Compressing large chunks inline would block the event loop.
  176. limiter = _get_gzip_capacity_limiter()
  177. return await anyio.to_thread.run_sync(self._compress_body, body, more_body, limiter=limiter)
  178. return self._compress_body(body, more_body)
  179. def _compress_body(self, body: bytes, more_body: bool) -> bytes:
  180. if more_body:
  181. return self.compressor.compress(body) + self.compressor.flush(zlib.Z_SYNC_FLUSH)
  182. return self.compressor.compress(body) + self.compressor.flush()
  183. async def unattached_send(message: Message) -> NoReturn:
  184. raise RuntimeError("send awaitable not set") # pragma: no cover
  185. def _normalize_content_types(content_types: tuple[str, ...]) -> tuple[str, ...]:
  186. return tuple(content_type.partition(";")[0].strip().lower() for content_type in content_types)