body_limit.py 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132
  1. from __future__ import annotations
  2. from typing import cast
  3. from starlette.datastructures import Headers
  4. from starlette.exceptions import HTTPException
  5. from starlette.responses import PlainTextResponse
  6. from starlette.types import ASGIApp, Message, Receive, Scope, Send
  7. MAX_BODY_SIZE_SCOPE_KEY = "starlette.max_body_size"
  8. _BODY_LIMIT_RESPONDER_SCOPE_KEY = "starlette._body_limit_responder"
  9. class _Missing:
  10. __slots__ = ()
  11. _MISSING = _Missing()
  12. class _RequestBodyTooLarge(HTTPException):
  13. def __init__(self) -> None:
  14. super().__init__(status_code=413, detail="Content Too Large")
  15. class _RequestBodyLimitResponseSent(Exception):
  16. pass
  17. class RequestBodyLimitMiddleware:
  18. """Limit the total size of an HTTP request body."""
  19. def __init__(self, app: ASGIApp, max_body_size: int) -> None:
  20. self.app = app
  21. self.max_body_size = max_body_size
  22. async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
  23. if scope["type"] != "http":
  24. return await self.app(scope, receive, send)
  25. responder = RequestBodyLimitResponder(self.app, self.max_body_size)
  26. await responder(scope, receive, send)
  27. class RequestBodyLimitResponder:
  28. def __init__(self, app: ASGIApp, max_body_size: int) -> None:
  29. self.app = app
  30. self.max_body_size = max_body_size
  31. self._scope: Scope | None = None
  32. self._receive: Receive | None = None
  33. self._send: Send | None = None
  34. self.content_length: int | None = None
  35. self.total_size = 0
  36. self.response_started = False
  37. @property
  38. def scope(self) -> Scope:
  39. assert self._scope is not None
  40. return self._scope
  41. @property
  42. def receive(self) -> Receive:
  43. assert self._receive is not None
  44. return self._receive
  45. @property
  46. def send(self) -> Send:
  47. assert self._send is not None
  48. return self._send
  49. async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
  50. previous_scope_limit = cast(int | _Missing, scope.get(MAX_BODY_SIZE_SCOPE_KEY, _MISSING))
  51. scope[MAX_BODY_SIZE_SCOPE_KEY] = self.max_body_size
  52. active_responder = cast(RequestBodyLimitResponder | None, scope.get(_BODY_LIMIT_RESPONDER_SCOPE_KEY))
  53. if active_responder is not None:
  54. active_responder.max_body_size = self.max_body_size
  55. if active_responder.total_size > active_responder.max_body_size:
  56. raise _RequestBodyTooLarge
  57. return await self.app(scope, receive, send)
  58. self._scope = scope
  59. self._receive = receive
  60. self._send = send
  61. self.content_length = _get_content_length(scope)
  62. scope[_BODY_LIMIT_RESPONDER_SCOPE_KEY] = self
  63. try:
  64. await self.app(scope, self.receive_with_limit, self.send_with_limit)
  65. except _RequestBodyTooLarge:
  66. if self.response_started:
  67. raise
  68. response = PlainTextResponse("Content Too Large", status_code=413)
  69. await response(scope, receive, send)
  70. except _RequestBodyLimitResponseSent:
  71. pass
  72. finally:
  73. scope.pop(_BODY_LIMIT_RESPONDER_SCOPE_KEY, None)
  74. if isinstance(previous_scope_limit, _Missing):
  75. scope.pop(MAX_BODY_SIZE_SCOPE_KEY, None)
  76. else:
  77. scope[MAX_BODY_SIZE_SCOPE_KEY] = previous_scope_limit
  78. async def receive_with_limit(self) -> Message:
  79. if self.content_length is not None and self.content_length > self.max_body_size:
  80. raise _RequestBodyTooLarge
  81. message = await self.receive()
  82. if message["type"] == "http.request":
  83. self.total_size += len(message.get("body", b""))
  84. if self.total_size > self.max_body_size:
  85. raise _RequestBodyTooLarge
  86. return message
  87. async def send_with_limit(self, message: Message) -> None:
  88. if message["type"] == "http.response.start":
  89. self.response_started = True
  90. if self.content_length is not None and self.content_length > self.max_body_size:
  91. response = PlainTextResponse("Content Too Large", status_code=413)
  92. await response(self.scope, self.receive, self.send)
  93. raise _RequestBodyLimitResponseSent
  94. await self.send(message)
  95. def _get_content_length(scope: Scope) -> int | None:
  96. content_length = Headers(scope=scope).get("content-length")
  97. if content_length is None:
  98. return None
  99. try:
  100. return int(content_length)
  101. except ValueError:
  102. return None