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