formparsers.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297
  1. from __future__ import annotations
  2. from collections.abc import AsyncGenerator
  3. from dataclasses import dataclass, field
  4. from enum import Enum
  5. from tempfile import SpooledTemporaryFile
  6. from typing import TYPE_CHECKING
  7. from urllib.parse import unquote_plus
  8. from starlette.datastructures import FormData, Headers, UploadFile
  9. if TYPE_CHECKING:
  10. import python_multipart as multipart
  11. from python_multipart.multipart import MultipartCallbacks, QuerystringCallbacks, parse_options_header
  12. else:
  13. try:
  14. try:
  15. import python_multipart as multipart
  16. from python_multipart.multipart import parse_options_header
  17. except ModuleNotFoundError: # pragma: no cover
  18. import multipart
  19. from multipart.multipart import parse_options_header
  20. except ModuleNotFoundError: # pragma: no cover
  21. multipart = None
  22. parse_options_header = None
  23. class FormMessage(Enum):
  24. FIELD_START = 1
  25. FIELD_NAME = 2
  26. FIELD_DATA = 3
  27. FIELD_END = 4
  28. END = 5
  29. @dataclass
  30. class MultipartPart:
  31. content_disposition: bytes | None = None
  32. field_name: str = ""
  33. data: bytearray = field(default_factory=bytearray)
  34. file: UploadFile | None = None
  35. item_headers: list[tuple[bytes, bytes]] = field(default_factory=list)
  36. def _user_safe_decode(src: bytes | bytearray, codec: str) -> str:
  37. try:
  38. return src.decode(codec)
  39. except (UnicodeDecodeError, LookupError):
  40. return src.decode("latin-1")
  41. class MultiPartException(Exception):
  42. def __init__(self, message: str) -> None:
  43. self.message = message
  44. class FormParser:
  45. def __init__(
  46. self,
  47. headers: Headers,
  48. stream: AsyncGenerator[bytes, None],
  49. *,
  50. max_fields: int | float = 1000,
  51. max_part_size: int = 1024 * 1024, # 1MB
  52. ) -> None:
  53. assert multipart is not None, "The `python-multipart` library must be installed to use form parsing."
  54. self.headers = headers
  55. self.stream = stream
  56. self.max_fields = max_fields
  57. self.max_part_size = max_part_size
  58. self.messages: list[tuple[FormMessage, bytes]] = []
  59. self._current_field_size = 0
  60. self._current_fields = 0
  61. def on_field_start(self) -> None:
  62. self._current_field_size = 0
  63. message = (FormMessage.FIELD_START, b"")
  64. self.messages.append(message)
  65. def on_field_name(self, data: bytes, start: int, end: int) -> None:
  66. self._current_field_size += end - start
  67. if self._current_field_size > self.max_part_size:
  68. raise MultiPartException(f"Field exceeded maximum size of {int(self.max_part_size / 1024)}KB.")
  69. message = (FormMessage.FIELD_NAME, data[start:end])
  70. self.messages.append(message)
  71. def on_field_data(self, data: bytes, start: int, end: int) -> None:
  72. self._current_field_size += end - start
  73. if self._current_field_size > self.max_part_size:
  74. raise MultiPartException(f"Field exceeded maximum size of {int(self.max_part_size / 1024)}KB.")
  75. message = (FormMessage.FIELD_DATA, data[start:end])
  76. self.messages.append(message)
  77. def on_field_end(self) -> None:
  78. self._current_fields += 1
  79. if self._current_fields > self.max_fields:
  80. raise MultiPartException(f"Too many fields. Maximum number of fields is {self.max_fields}.")
  81. message = (FormMessage.FIELD_END, b"")
  82. self.messages.append(message)
  83. def on_end(self) -> None:
  84. message = (FormMessage.END, b"")
  85. self.messages.append(message)
  86. async def parse(self) -> FormData:
  87. # Callbacks dictionary.
  88. callbacks: QuerystringCallbacks = {
  89. "on_field_start": self.on_field_start,
  90. "on_field_name": self.on_field_name,
  91. "on_field_data": self.on_field_data,
  92. "on_field_end": self.on_field_end,
  93. "on_end": self.on_end,
  94. }
  95. # Create the parser.
  96. parser = multipart.QuerystringParser(callbacks)
  97. field_name = bytearray()
  98. field_value = bytearray()
  99. items: list[tuple[str, str | UploadFile]] = []
  100. # Feed the parser with data from the request.
  101. async for chunk in self.stream:
  102. if chunk:
  103. parser.write(chunk)
  104. else:
  105. parser.finalize()
  106. messages = list(self.messages)
  107. self.messages.clear()
  108. for message_type, message_bytes in messages:
  109. if message_type == FormMessage.FIELD_START:
  110. field_name = bytearray()
  111. field_value = bytearray()
  112. elif message_type == FormMessage.FIELD_NAME:
  113. field_name.extend(message_bytes)
  114. elif message_type == FormMessage.FIELD_DATA:
  115. field_value.extend(message_bytes)
  116. elif message_type == FormMessage.FIELD_END:
  117. name = unquote_plus(field_name.decode("latin-1"))
  118. value = unquote_plus(field_value.decode("latin-1"))
  119. items.append((name, value))
  120. return FormData(items)
  121. class MultiPartParser:
  122. spool_max_size = 1024 * 1024 # 1MB
  123. """The maximum size of the spooled temporary file used to store file data."""
  124. max_part_size = 1024 * 1024 # 1MB
  125. """The maximum size of a part in the multipart request."""
  126. def __init__(
  127. self,
  128. headers: Headers,
  129. stream: AsyncGenerator[bytes, None],
  130. *,
  131. max_files: int | float = 1000,
  132. max_fields: int | float = 1000,
  133. max_part_size: int = 1024 * 1024, # 1MB
  134. ) -> None:
  135. assert multipart is not None, "The `python-multipart` library must be installed to use form parsing."
  136. self.headers = headers
  137. self.stream = stream
  138. self.max_files = max_files
  139. self.max_fields = max_fields
  140. self.items: list[tuple[str, str | UploadFile]] = []
  141. self._current_files = 0
  142. self._current_fields = 0
  143. self._current_partial_header_name: bytes = b""
  144. self._current_partial_header_value: bytes = b""
  145. self._current_part = MultipartPart()
  146. self._charset = ""
  147. self._file_parts_to_write: list[tuple[MultipartPart, bytes]] = []
  148. self._file_parts_to_finish: list[MultipartPart] = []
  149. self._files_to_close_on_error: list[SpooledTemporaryFile[bytes]] = []
  150. self.max_part_size = max_part_size
  151. def on_part_begin(self) -> None:
  152. self._current_part = MultipartPart()
  153. def on_part_data(self, data: bytes, start: int, end: int) -> None:
  154. message_bytes = data[start:end]
  155. if self._current_part.file is None:
  156. if len(self._current_part.data) + len(message_bytes) > self.max_part_size:
  157. raise MultiPartException(f"Part exceeded maximum size of {int(self.max_part_size / 1024)}KB.")
  158. self._current_part.data.extend(message_bytes)
  159. else:
  160. self._file_parts_to_write.append((self._current_part, message_bytes))
  161. def on_part_end(self) -> None:
  162. if self._current_part.file is None:
  163. self.items.append(
  164. (
  165. self._current_part.field_name,
  166. _user_safe_decode(self._current_part.data, self._charset),
  167. )
  168. )
  169. else:
  170. self._file_parts_to_finish.append(self._current_part)
  171. # The file can be added to the items right now even though it's not
  172. # finished yet, because it will be finished in the `parse()` method, before
  173. # self.items is used in the return value.
  174. self.items.append((self._current_part.field_name, self._current_part.file))
  175. def on_header_field(self, data: bytes, start: int, end: int) -> None:
  176. self._current_partial_header_name += data[start:end]
  177. def on_header_value(self, data: bytes, start: int, end: int) -> None:
  178. self._current_partial_header_value += data[start:end]
  179. def on_header_end(self) -> None:
  180. field = self._current_partial_header_name.lower()
  181. if field == b"content-disposition":
  182. self._current_part.content_disposition = self._current_partial_header_value
  183. self._current_part.item_headers.append((field, self._current_partial_header_value))
  184. self._current_partial_header_name = b""
  185. self._current_partial_header_value = b""
  186. def on_headers_finished(self) -> None:
  187. disposition, options = parse_options_header(self._current_part.content_disposition)
  188. try:
  189. self._current_part.field_name = _user_safe_decode(options[b"name"], self._charset)
  190. except KeyError:
  191. raise MultiPartException('The Content-Disposition header field "name" must be provided.')
  192. if b"filename" in options:
  193. self._current_files += 1
  194. if self._current_files > self.max_files:
  195. raise MultiPartException(f"Too many files. Maximum number of files is {self.max_files}.")
  196. filename = _user_safe_decode(options[b"filename"], self._charset)
  197. tempfile = SpooledTemporaryFile(max_size=self.spool_max_size)
  198. self._files_to_close_on_error.append(tempfile)
  199. self._current_part.file = UploadFile(
  200. file=tempfile, # type: ignore[arg-type]
  201. size=0,
  202. filename=filename,
  203. headers=Headers(raw=self._current_part.item_headers),
  204. )
  205. else:
  206. self._current_fields += 1
  207. if self._current_fields > self.max_fields:
  208. raise MultiPartException(f"Too many fields. Maximum number of fields is {self.max_fields}.")
  209. self._current_part.file = None
  210. def on_end(self) -> None:
  211. pass
  212. async def parse(self) -> FormData:
  213. # Parse the Content-Type header to get the multipart boundary.
  214. _, params = parse_options_header(self.headers["Content-Type"])
  215. charset = params.get(b"charset", "utf-8")
  216. if isinstance(charset, bytes):
  217. charset = charset.decode("latin-1")
  218. self._charset = charset
  219. try:
  220. boundary = params[b"boundary"]
  221. except KeyError:
  222. raise MultiPartException("Missing boundary in multipart.")
  223. # Callbacks dictionary.
  224. callbacks: MultipartCallbacks = {
  225. "on_part_begin": self.on_part_begin,
  226. "on_part_data": self.on_part_data,
  227. "on_part_end": self.on_part_end,
  228. "on_header_field": self.on_header_field,
  229. "on_header_value": self.on_header_value,
  230. "on_header_end": self.on_header_end,
  231. "on_headers_finished": self.on_headers_finished,
  232. "on_end": self.on_end,
  233. }
  234. # Create the parser.
  235. parser = multipart.MultipartParser(boundary, callbacks)
  236. try:
  237. # Feed the parser with data from the request.
  238. async for chunk in self.stream:
  239. parser.write(chunk)
  240. # Write file data, it needs to use await with the UploadFile methods
  241. # that call the corresponding file methods *in a threadpool*,
  242. # otherwise, if they were called directly in the callback methods above
  243. # (regular, non-async functions), that would block the event loop in
  244. # the main thread.
  245. for part, data in self._file_parts_to_write:
  246. assert part.file # for type checkers
  247. await part.file.write(data)
  248. for part in self._file_parts_to_finish:
  249. assert part.file # for type checkers
  250. await part.file.seek(0)
  251. self._file_parts_to_write.clear()
  252. self._file_parts_to_finish.clear()
  253. parser.finalize()
  254. except BaseException:
  255. # Close all the files if parsing or reading the request stream fails.
  256. for file in self._files_to_close_on_error:
  257. file.close()
  258. raise
  259. return FormData(self.items)