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

192 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.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)