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