gzip.py 4.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108
  1. import gzip
  2. import io
  3. import typing
  4. from starlette.datastructures import Headers, MutableHeaders
  5. from starlette.types import ASGIApp, Message, Receive, Scope, Send
  6. class GZipMiddleware:
  7. def __init__(self, app: ASGIApp, minimum_size: int = 500, compresslevel: int = 9) -> None:
  8. self.app = app
  9. self.minimum_size = minimum_size
  10. self.compresslevel = compresslevel
  11. async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
  12. if scope["type"] == "http":
  13. headers = Headers(scope=scope)
  14. if "gzip" in headers.get("Accept-Encoding", ""):
  15. responder = GZipResponder(self.app, self.minimum_size, compresslevel=self.compresslevel)
  16. await responder(scope, receive, send)
  17. return
  18. await self.app(scope, receive, send)
  19. class GZipResponder:
  20. def __init__(self, app: ASGIApp, minimum_size: int, compresslevel: int = 9) -> None:
  21. self.app = app
  22. self.minimum_size = minimum_size
  23. self.send: Send = unattached_send
  24. self.initial_message: Message = {}
  25. self.started = False
  26. self.content_encoding_set = False
  27. self.gzip_buffer = io.BytesIO()
  28. self.gzip_file = gzip.GzipFile(mode="wb", fileobj=self.gzip_buffer, compresslevel=compresslevel)
  29. async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
  30. self.send = send
  31. with self.gzip_buffer, self.gzip_file:
  32. await self.app(scope, receive, self.send_with_gzip)
  33. async def send_with_gzip(self, message: Message) -> None:
  34. message_type = message["type"]
  35. if message_type == "http.response.start":
  36. # Don't send the initial message until we've determined how to
  37. # modify the outgoing headers correctly.
  38. self.initial_message = message
  39. headers = Headers(raw=self.initial_message["headers"])
  40. self.content_encoding_set = "content-encoding" in headers
  41. elif message_type == "http.response.body" and self.content_encoding_set:
  42. if not self.started:
  43. self.started = True
  44. await self.send(self.initial_message)
  45. await self.send(message)
  46. elif message_type == "http.response.body" and not self.started:
  47. self.started = True
  48. body = message.get("body", b"")
  49. more_body = message.get("more_body", False)
  50. if len(body) < self.minimum_size and not more_body:
  51. # Don't apply GZip to small outgoing responses.
  52. await self.send(self.initial_message)
  53. await self.send(message)
  54. elif not more_body:
  55. # Standard GZip response.
  56. self.gzip_file.write(body)
  57. self.gzip_file.close()
  58. body = self.gzip_buffer.getvalue()
  59. headers = MutableHeaders(raw=self.initial_message["headers"])
  60. headers["Content-Encoding"] = "gzip"
  61. headers["Content-Length"] = str(len(body))
  62. headers.add_vary_header("Accept-Encoding")
  63. message["body"] = body
  64. await self.send(self.initial_message)
  65. await self.send(message)
  66. else:
  67. # Initial body in streaming GZip response.
  68. headers = MutableHeaders(raw=self.initial_message["headers"])
  69. headers["Content-Encoding"] = "gzip"
  70. headers.add_vary_header("Accept-Encoding")
  71. del headers["Content-Length"]
  72. self.gzip_file.write(body)
  73. message["body"] = self.gzip_buffer.getvalue()
  74. self.gzip_buffer.seek(0)
  75. self.gzip_buffer.truncate()
  76. await self.send(self.initial_message)
  77. await self.send(message)
  78. elif message_type == "http.response.body":
  79. # Remaining body in streaming GZip response.
  80. body = message.get("body", b"")
  81. more_body = message.get("more_body", False)
  82. self.gzip_file.write(body)
  83. if not more_body:
  84. self.gzip_file.close()
  85. message["body"] = self.gzip_buffer.getvalue()
  86. self.gzip_buffer.seek(0)
  87. self.gzip_buffer.truncate()
  88. await self.send(message)
  89. async def unattached_send(message: Message) -> typing.NoReturn:
  90. raise RuntimeError("send awaitable not set") # pragma: no cover