formparsers.py 11 KB

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