Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/starlette/formparsers.py: 27%

Shortcuts on this page

r m x   toggle line displays

j k   next/prev highlighted chunk

0   (zero) top of page

1   (one) first highlighted chunk

196 statements  

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)