1from __future__ import annotations
2
3import logging
4import os
5import shutil
6import sys
7import tempfile
8from enum import IntEnum, global_enum
9from io import BufferedRandom, BytesIO
10from numbers import Number
11from typing import TYPE_CHECKING, cast
12
13from .decoders import Base64Decoder, QuotedPrintableDecoder
14from .exceptions import FileError, FormParserError, MultipartParseError, QuerystringParseError
15
16if TYPE_CHECKING:
17 from collections.abc import Callable
18 from typing import Any, Literal, Protocol, TypeAlias, TypedDict
19
20 class SupportsRead(Protocol):
21 def read(self, __n: int) -> bytes: ...
22
23 class QuerystringCallbacks(TypedDict, total=False):
24 on_field_start: Callable[[], None]
25 on_field_name: Callable[[bytes, int, int], None]
26 on_field_data: Callable[[bytes, int, int], None]
27 on_field_end: Callable[[], None]
28 on_end: Callable[[], None]
29
30 class OctetStreamCallbacks(TypedDict, total=False):
31 on_start: Callable[[], None]
32 on_data: Callable[[bytes, int, int], None]
33 on_end: Callable[[], None]
34
35 class MultipartCallbacks(TypedDict, total=False):
36 on_part_begin: Callable[[], None]
37 on_part_data: Callable[[bytes, int, int], None]
38 on_part_end: Callable[[], None]
39 on_header_begin: Callable[[], None]
40 on_header_field: Callable[[bytes, int, int], None]
41 on_header_value: Callable[[bytes, int, int], None]
42 on_header_end: Callable[[], None]
43 on_headers_finished: Callable[[], None]
44 on_end: Callable[[], None]
45
46 class FileConfig(TypedDict, total=False):
47 UPLOAD_DIR: str | bytes | None
48 UPLOAD_DELETE_TMP: bool
49 UPLOAD_KEEP_FILENAME: bool
50 UPLOAD_KEEP_EXTENSIONS: bool
51 MAX_MEMORY_FILE_SIZE: int
52
53 class FormParserConfig(FileConfig):
54 UPLOAD_ERROR_ON_BAD_CTE: bool
55 MAX_BODY_SIZE: float
56 MAX_HEADER_COUNT: int
57 MAX_HEADER_SIZE: int
58
59 CallbackName: TypeAlias = Literal[
60 "start",
61 "data",
62 "end",
63 "field_start",
64 "field_name",
65 "field_data",
66 "field_end",
67 "part_begin",
68 "part_data",
69 "part_end",
70 "header_begin",
71 "header_field",
72 "header_value",
73 "header_end",
74 "headers_finished",
75 ]
76
77# Unique missing object.
78_missing = object()
79
80
81def _noop_event() -> None:
82 pass
83
84
85def _noop_data(_data: bytes, _start: int, _end: int) -> None:
86 pass
87
88
89@global_enum
90class QuerystringState(IntEnum):
91 """Querystring parser states.
92
93 These are used to keep track of the state of the parser, and are used to determine
94 what to do when new data is encountered.
95 """
96
97 BEFORE_FIELD = 0
98 FIELD_NAME = 1
99 FIELD_DATA = 2
100
101
102@global_enum
103class MultipartState(IntEnum):
104 """Multipart parser states.
105
106 These are used to keep track of the state of the parser, and are used to determine
107 what to do when new data is encountered.
108 """
109
110 START = 0
111 START_BOUNDARY = 1
112 HEADER_FIELD_START = 2
113 HEADER_FIELD = 3
114 HEADER_VALUE_START = 4
115 HEADER_VALUE = 5
116 HEADER_VALUE_ALMOST_DONE = 6
117 HEADERS_ALMOST_DONE = 7
118 PART_DATA_START = 8
119 PART_DATA = 9
120 PART_DATA_END = 10
121 END_BOUNDARY = 11
122 END = 12
123
124
125if TYPE_CHECKING:
126 BEFORE_FIELD = QuerystringState.BEFORE_FIELD
127 FIELD_NAME = QuerystringState.FIELD_NAME
128 FIELD_DATA = QuerystringState.FIELD_DATA
129
130 START = MultipartState.START
131 START_BOUNDARY = MultipartState.START_BOUNDARY
132 HEADER_FIELD_START = MultipartState.HEADER_FIELD_START
133 HEADER_FIELD = MultipartState.HEADER_FIELD
134 HEADER_VALUE_START = MultipartState.HEADER_VALUE_START
135 HEADER_VALUE = MultipartState.HEADER_VALUE
136 HEADER_VALUE_ALMOST_DONE = MultipartState.HEADER_VALUE_ALMOST_DONE
137 HEADERS_ALMOST_DONE = MultipartState.HEADERS_ALMOST_DONE
138 PART_DATA_START = MultipartState.PART_DATA_START
139 PART_DATA = MultipartState.PART_DATA
140 PART_DATA_END = MultipartState.PART_DATA_END
141 END_BOUNDARY = MultipartState.END_BOUNDARY
142 END = MultipartState.END
143
144# Flags for the multipart parser.
145FLAG_PART_BOUNDARY = 1
146FLAG_LAST_BOUNDARY = 2
147
148# Get constants. Since iterating over a str on Python 2 gives you a 1-length
149# string, but iterating over a bytes object on Python 3 gives you an integer,
150# we need to save these constants.
151CR = b"\r"[0]
152LF = b"\n"[0]
153COLON = b":"[0]
154SPACE = b" "[0]
155HYPHEN = b"-"[0]
156AMPERSAND = b"&"[0]
157LOWER_A = b"a"[0]
158LOWER_Z = b"z"[0]
159NULL = b"\x00"[0]
160
161# fmt: off
162# Mask for ASCII characters that can be http tokens.
163# Per RFC7230 - 3.2.6, this is all alpha-numeric characters
164# and these: !#$%&'*+-.^_`|~
165TOKEN_CHARS = (
166 b"ABCDEFGHIJKLMNOPQRSTUVWXYZ"
167 b"abcdefghijklmnopqrstuvwxyz"
168 b"0123456789"
169 b"!#$%&'*+-.^_`|~")
170TOKEN_CHARS_SET = frozenset(TOKEN_CHARS)
171# fmt: on
172
173DEFAULT_MAX_HEADER_COUNT = 8
174"""Default maximum number of headers allowed per multipart part."""
175
176DEFAULT_MAX_HEADER_SIZE = 4096 + 128
177"""Default maximum size of a single multipart header line, including syntax overhead."""
178
179MAX_BOUNDARY_LENGTH = 256
180"""Maximum allowed length of a multipart boundary.
181
182[RFC 2046 §5.1.1](https://datatracker.ietf.org/doc/html/rfc2046#section-5.1.1)
183recommends boundaries be at most 70 bytes. 256 bytes is generous headroom over
184every HTTP client.
185"""
186
187
188def _parseparam(s: str) -> list[str]:
189 # Vendored from the standard library's
190 # [`email.message._parseparam`](https://github.com/python/cpython/blob/v3.14.2/Lib/email/message.py#L73-L96)
191 # to split a header into its `;`-separated parts without treating a `;` inside a double-quoted string as a
192 # separator - and without the RFC 2231 decoding that `email.message.Message.get_params` would apply on top.
193 s = ";" + s
194 plist: list[str] = []
195 start = 0
196 while s.find(";", start) == start:
197 start += 1
198 end = s.find(";", start)
199 ind, diff = start, 0
200 while end > 0:
201 diff += s.count('"', ind, end) - s.count('\\"', ind, end)
202 if diff % 2 == 0:
203 break
204 end, ind = ind, s.find(";", end + 1)
205 if end < 0:
206 end = len(s)
207 i = s.find("=", start, end)
208 if i == -1:
209 f = s[start:end]
210 else:
211 f = s[start:i].rstrip().lower() + "=" + s[i + 1 : end].lstrip()
212 plist.append(f.strip())
213 start = end
214 return plist
215
216
217def parse_options_header(value: str | bytes | None) -> tuple[bytes, dict[bytes, bytes]]:
218 """Parses a Content-Type header into a value in the following format: (content_type, {parameters})."""
219 if not value:
220 return (b"", {})
221
222 # If we are passed bytes, we assume that it conforms to WSGI, encoding in latin-1.
223 if isinstance(value, bytes): # pragma: no cover
224 value = value.decode("latin-1")
225
226 # For types
227 assert isinstance(value, str), "Value should be a string by now"
228
229 # If we have no options, return the string as-is.
230 if ";" not in value:
231 return (value.lower().strip().encode("latin-1"), {})
232
233 ctype, *segments = _parseparam(value)
234 options: dict[bytes, bytes] = {}
235 for segment in segments:
236 key, _, val = segment.partition("=")
237 # [RFC 7578 §4.2](https://datatracker.ietf.org/doc/html/rfc7578#section-4.2)
238 # forbids the RFC 5987/2231 extended syntax (`key*=`, `key*0`, ...) in
239 # multipart/form-data, so we ignore those parameters and keep the plain
240 # `key` authoritative.
241 if "*" in key:
242 continue
243 if len(val) >= 2 and val[0] == '"' and val[-1] == '"':
244 val = val[1:-1].replace("\\\\", "\\").replace('\\"', '"')
245 # Work around an IE6 bug where the full file path is sent instead of
246 # just the filename.
247 if key == "filename" and (val[1:3] == ":\\" or val[:2] == "\\\\"):
248 val = val.split("\\")[-1]
249 options[key.encode("latin-1")] = val.encode("latin-1")
250 return ctype.encode("latin-1"), options
251
252
253class Field:
254 """A Field object represents a (parsed) form field. It represents a single
255 field with a corresponding name and value.
256
257 The name that a :class:`Field` will be instantiated with is the same name
258 that would be found in the following HTML::
259
260 <input name="name_goes_here" type="text"/>
261
262 This class defines two methods, :meth:`on_data` and :meth:`on_end`, that
263 will be called when data is written to the Field, and when the Field is
264 finalized, respectively.
265
266 Args:
267 name: The name of the form field.
268 content_type: The value of the Content-Type header for this field.
269 """
270
271 def __init__(self, name: bytes | None, *, content_type: str | None = None) -> None:
272 self._name = name
273 self._value: list[bytes] = []
274 self._content_type = content_type
275
276 # We cache the joined version of _value for speed.
277 self._cache = _missing
278
279 @classmethod
280 def from_value(cls, name: bytes, value: bytes | None) -> Field:
281 """Create an instance of a :class:`Field`, and set the corresponding
282 value - either None or an actual value. This method will also
283 finalize the Field itself.
284
285 Args:
286 name: the name of the form field.
287 value: the value of the form field - either a bytestring or None.
288
289 Returns:
290 A new instance of a [`Field`][python_multipart.Field].
291 """
292
293 f = cls(name)
294 if value is None:
295 f.set_none()
296 else:
297 f.write(value)
298 f.finalize()
299 return f
300
301 def write(self, data: bytes) -> int:
302 """Write some data into the form field.
303
304 Args:
305 data: The data to write to the field.
306
307 Returns:
308 The number of bytes written.
309 """
310 return self.on_data(data)
311
312 def on_data(self, data: bytes) -> int:
313 """This method is a callback that will be called whenever data is
314 written to the Field.
315
316 Args:
317 data: The data to write to the field.
318
319 Returns:
320 The number of bytes written.
321 """
322 self._value.append(data)
323 self._cache = _missing
324 return len(data)
325
326 def on_end(self) -> None:
327 """This method is called whenever the Field is finalized."""
328 if self._cache is _missing:
329 self._cache = b"".join(self._value)
330
331 def finalize(self) -> None:
332 """Finalize the form field."""
333 self.on_end()
334
335 def close(self) -> None:
336 """Close the Field object. This will free any underlying cache."""
337 # Free our value array.
338 if self._cache is _missing:
339 self._cache = b"".join(self._value)
340
341 del self._value
342
343 def set_none(self) -> None:
344 """Some fields in a querystring can possibly have a value of None - for
345 example, the string "foo&bar=&baz=asdf" will have a field with the
346 name "foo" and value None, one with name "bar" and value "", and one
347 with name "baz" and value "asdf". Since the write() interface doesn't
348 support writing None, this function will set the field value to None.
349 """
350 self._cache = None
351
352 @property
353 def field_name(self) -> bytes | None:
354 """This property returns the name of the field."""
355 return self._name
356
357 @property
358 def value(self) -> bytes | None:
359 """This property returns the value of the form field."""
360 if self._cache is _missing:
361 self._cache = b"".join(self._value)
362
363 assert isinstance(self._cache, bytes) or self._cache is None
364 return self._cache
365
366 @property
367 def content_type(self) -> str | None:
368 """This property returns the content_type value of the field."""
369 return self._content_type
370
371 def __eq__(self, other: object) -> bool:
372 if isinstance(other, Field):
373 return self.field_name == other.field_name and self.value == other.value
374 else:
375 return NotImplemented
376
377 def __repr__(self) -> str:
378 if self.value is not None and len(self.value) > 97:
379 # We get the repr, and then insert three dots before the final
380 # quote.
381 v = repr(self.value[:97])[:-1] + "...'"
382 else:
383 v = repr(self.value)
384
385 return f"{self.__class__.__name__}(field_name={self.field_name!r}, value={v})"
386
387
388class File:
389 """This class represents an uploaded file. It handles writing file data to
390 either an in-memory file or a temporary file on-disk, if the optional
391 threshold is passed.
392
393 There are some options that can be passed to the File to change behavior
394 of the class. Valid options are as follows:
395
396 | Name | Type | Default | Description |
397 |-----------------------|-------|---------|-------------|
398 | UPLOAD_DIR | `str` | None | The directory to store uploaded files in. If this is None, a temporary file will be created in the system's standard location. |
399 | UPLOAD_DELETE_TMP | `bool`| True | Delete automatically created TMP file |
400 | UPLOAD_KEEP_FILENAME | `bool`| False | Whether or not to keep the filename of the uploaded file. If True, then the filename is reduced to its basename (directory components are stripped) and the file is saved with that name. Otherwise, a temporary name will be used. |
401 | UPLOAD_KEEP_EXTENSIONS| `bool`| False | Whether or not to keep the uploaded file's extension. If False, the file will be saved with the default temporary extension (usually ".tmp"). Otherwise, the file's extension will be maintained. Note that this will properly combine with the UPLOAD_KEEP_FILENAME setting. |
402 | MAX_MEMORY_FILE_SIZE | `int` | 1 MiB | The maximum number of bytes of a File to keep in memory. By default, the contents of a File are kept into memory until a certain limit is reached, after which the contents of the File are written to a temporary file. This behavior can be disabled by setting this value to an appropriately large value (or, for example, infinity, such as `float('inf')`. |
403
404 Args:
405 file_name: The name of the file that this [`File`][python_multipart.File] represents.
406 field_name: The name of the form field that this file was uploaded with. This can be None, if, for example,
407 the file was uploaded with Content-Type application/octet-stream.
408 config: The configuration for this File. See above for valid configuration keys and their corresponding values.
409 content_type: The value of the Content-Type header.
410 """ # noqa: E501
411
412 def __init__(
413 self,
414 file_name: bytes | None,
415 field_name: bytes | None = None,
416 config: FileConfig = {},
417 *,
418 content_type: str | None = None,
419 ) -> None:
420 # Save configuration, set other variables default.
421 self.logger = logging.getLogger(__name__)
422 self._config = config
423 self._in_memory = True
424 self._bytes_written = 0
425 self._fileobj: BytesIO | BufferedRandom = BytesIO()
426
427 # Save the provided field/file name and content type.
428 self._field_name = field_name
429 self._file_name = file_name
430 self._content_type = content_type
431
432 # Our actual file name is None by default, since, depending on our
433 # config, we may not actually use the provided name.
434 self._actual_file_name: bytes | None = None
435
436 # Split the extension from the filename.
437 if file_name is not None:
438 # Extract just the basename to avoid directory traversal
439 basename = os.path.basename(file_name)
440 base, ext = os.path.splitext(basename)
441 self._file_base = base
442 self._ext = ext
443
444 @property
445 def field_name(self) -> bytes | None:
446 """The form field associated with this file. May be None if there isn't
447 one, for example when we have an application/octet-stream upload.
448 """
449 return self._field_name
450
451 @property
452 def file_name(self) -> bytes | None:
453 """The file name given in the upload request."""
454 return self._file_name
455
456 @property
457 def actual_file_name(self) -> bytes | None:
458 """The file name that this file is saved as. Will be None if it's not
459 currently saved on disk.
460 """
461 return self._actual_file_name
462
463 @property
464 def file_object(self) -> BytesIO | BufferedRandom:
465 """The file object that we're currently writing to. Note that this
466 will either be an instance of a :class:`io.BytesIO`, or a regular file
467 object.
468 """
469 return self._fileobj
470
471 @property
472 def size(self) -> int:
473 """The total size of this file, counted as the number of bytes that
474 currently have been written to the file.
475 """
476 return self._bytes_written
477
478 @property
479 def in_memory(self) -> bool:
480 """A boolean representing whether or not this file object is currently
481 stored in-memory or on-disk.
482 """
483 return self._in_memory
484
485 @property
486 def content_type(self) -> str | None:
487 """The Content-Type value for this part, if it was set."""
488 return self._content_type
489
490 def flush_to_disk(self) -> None:
491 """If the file is already on-disk, do nothing. Otherwise, copy from
492 the in-memory buffer to a disk file, and then reassign our internal
493 file object to this new disk file.
494
495 Note that if you attempt to flush a file that is already on-disk, a
496 warning will be logged to this module's logger.
497 """
498 if not self._in_memory:
499 self.logger.warning("Trying to flush to disk when we're not in memory")
500 return
501
502 # Go back to the start of our file.
503 self._fileobj.seek(0)
504
505 # Open a new file.
506 new_file = self._get_disk_file()
507
508 # Copy the file objects.
509 shutil.copyfileobj(self._fileobj, new_file)
510
511 # Seek to the new position in our new file.
512 new_file.seek(self._bytes_written)
513
514 # Reassign the fileobject.
515 old_fileobj = self._fileobj
516 self._fileobj = new_file
517
518 # We're no longer in memory.
519 self._in_memory = False
520
521 # Close the old file object.
522 old_fileobj.close()
523
524 def _get_disk_file(self) -> BufferedRandom:
525 """This function is responsible for getting a file object on-disk for us."""
526 self.logger.info("Opening a file on disk")
527
528 file_dir = self._config.get("UPLOAD_DIR")
529 keep_filename = self._config.get("UPLOAD_KEEP_FILENAME", False)
530 keep_extensions = self._config.get("UPLOAD_KEEP_EXTENSIONS", False)
531 delete_tmp = self._config.get("UPLOAD_DELETE_TMP", True)
532 tmp_file: None | BufferedRandom = None
533
534 # If we have a directory and are to keep the filename...
535 if file_dir is not None and keep_filename:
536 self.logger.info("Saving with filename in: %r", file_dir)
537
538 # Build our filename.
539 # TODO: what happens if we don't have a filename?
540 fname = self._file_base + self._ext if keep_extensions else self._file_base
541
542 path = os.path.join(file_dir, fname) # type: ignore[arg-type]
543 try:
544 self.logger.info("Opening file: %r", path)
545 tmp_file = open(path, "w+b")
546 except OSError:
547 tmp_file = None
548
549 self.logger.exception("Error opening temporary file")
550 raise FileError("Error opening temporary file: %r" % path)
551 else:
552 # Build options array.
553 # Note that on Python 3, tempfile doesn't support byte names. We
554 # encode our paths using the default filesystem encoding.
555 suffix = self._ext.decode(sys.getfilesystemencoding()) if keep_extensions else None
556
557 if file_dir is None:
558 dir = None
559 elif isinstance(file_dir, bytes):
560 dir = file_dir.decode(sys.getfilesystemencoding())
561 else:
562 dir = file_dir # pragma: no cover
563
564 # Create a temporary (named) file with the appropriate settings.
565 self.logger.info(
566 "Creating a temporary file with options: %r", {"suffix": suffix, "delete": delete_tmp, "dir": dir}
567 )
568 try:
569 tmp_file = cast(BufferedRandom, tempfile.NamedTemporaryFile(suffix=suffix, delete=delete_tmp, dir=dir))
570 except OSError:
571 self.logger.exception("Error creating named temporary file")
572 raise FileError("Error creating named temporary file")
573
574 assert tmp_file is not None
575 # Encode filename as bytes.
576 if isinstance(tmp_file.name, str):
577 fname = tmp_file.name.encode(sys.getfilesystemencoding())
578 else:
579 fname = cast(bytes, tmp_file.name) # pragma: no cover
580
581 self._actual_file_name = fname
582 return tmp_file
583
584 def write(self, data: bytes) -> int:
585 """Write some data to the File.
586
587 :param data: a bytestring
588 """
589 return self.on_data(data)
590
591 def on_data(self, data: bytes) -> int:
592 """This method is a callback that will be called whenever data is
593 written to the File.
594
595 Args:
596 data: The data to write to the file.
597
598 Returns:
599 The number of bytes written.
600 """
601 bwritten = self._fileobj.write(data)
602
603 # If the bytes written isn't the same as the length, just return.
604 if bwritten != len(data):
605 self.logger.warning("bwritten != len(data) (%d != %d)", bwritten, len(data))
606 return bwritten
607
608 # Keep track of how many bytes we've written.
609 self._bytes_written += bwritten
610
611 # If we're in-memory and are over our limit, we create a file.
612 max_memory_file_size = self._config.get("MAX_MEMORY_FILE_SIZE")
613 if self._in_memory and max_memory_file_size is not None and (self._bytes_written > max_memory_file_size):
614 self.logger.info("Flushing to disk")
615 self.flush_to_disk()
616
617 # Return the number of bytes written.
618 return bwritten
619
620 def on_end(self) -> None:
621 """This method is called whenever the Field is finalized."""
622 # Flush the underlying file object
623 self._fileobj.flush()
624
625 def finalize(self) -> None:
626 """Finalize the form file. This will not close the underlying file,
627 but simply signal that we are finished writing to the File.
628 """
629 self.on_end()
630
631 def close(self) -> None:
632 """Close the File object. This will actually close the underlying
633 file object (whether it's a :class:`io.BytesIO` or an actual file
634 object).
635 """
636 self._fileobj.close()
637
638 def __repr__(self) -> str:
639 return f"{self.__class__.__name__}(file_name={self.file_name!r}, field_name={self.field_name!r})"
640
641
642class BaseParser:
643 """This class is the base class for all parsers. It contains the logic for
644 calling and adding callbacks.
645
646 A callback can be one of two different forms. "Notification callbacks" are
647 callbacks that are called when something happens - for example, when a new
648 part of a multipart message is encountered by the parser. "Data callbacks"
649 are called when we get some sort of data - for example, part of the body of
650 a multipart chunk. Notification callbacks are called with no parameters,
651 whereas data callbacks are called with three, as follows::
652
653 data_callback(data, start, end)
654
655 The "data" parameter is a bytestring (i.e. "foo" on Python 2, or b"foo" on
656 Python 3). "start" and "end" are integer indexes into the "data" string
657 that represent the data of interest. Thus, in a data callback, the slice
658 `data[start:end]` represents the data that the callback is "interested in".
659 The callback is not passed a copy of the data, since copying severely hurts
660 performance.
661 """
662
663 def __init__(self) -> None:
664 self.logger = logging.getLogger(__name__)
665 self.callbacks: QuerystringCallbacks | OctetStreamCallbacks | MultipartCallbacks = {}
666
667 def callback(
668 self, name: CallbackName, data: bytes | None = None, start: int | None = None, end: int | None = None
669 ) -> None:
670 """This function calls a provided callback with some data. If the
671 callback is not set, will do nothing.
672
673 Args:
674 name: The name of the callback to call (as a string).
675 data: Data to pass to the callback. If None, then it is assumed that the callback is a notification
676 callback, and no parameters are given.
677 end: An integer that is passed to the data callback.
678 start: An integer that is passed to the data callback.
679 """
680 func = self.callbacks.get("on_" + name)
681 if func is None:
682 return
683 func = cast("Callable[..., Any]", func)
684 # Depending on whether we're given a buffer...
685 if data is not None:
686 # Don't do anything if we have start == end.
687 if start is not None and start == end:
688 return
689 func(data, start, end)
690 else:
691 func()
692
693 def set_callback(self, name: CallbackName, new_func: Callable[..., Any] | None) -> None:
694 """Update the function for a callback. Removes from the callbacks dict
695 if new_func is None.
696
697 :param name: The name of the callback to call (as a string).
698
699 :param new_func: The new function for the callback. If None, then the
700 callback will be removed (with no error if it does not
701 exist).
702 """
703 if new_func is None:
704 self.callbacks.pop("on_" + name, None) # type: ignore[misc]
705 else:
706 self.callbacks["on_" + name] = new_func # type: ignore[literal-required]
707
708 def close(self) -> None:
709 pass # pragma: no cover
710
711 def finalize(self) -> None:
712 pass # pragma: no cover
713
714 def __repr__(self) -> str:
715 return "%s()" % self.__class__.__name__
716
717
718class OctetStreamParser(BaseParser):
719 """This parser parses an octet-stream request body and calls callbacks when
720 incoming data is received. Callbacks are as follows:
721
722 | Callback Name | Parameters | Description |
723 |----------------|-----------------|-----------------------------------------------------|
724 | on_start | None | Called when the first data is parsed. |
725 | on_data | data, start, end| Called for each data chunk that is parsed. |
726 | on_end | None | Called when the parser is finished parsing all data.|
727
728 Args:
729 callbacks: A dictionary of callbacks. See the documentation for [`BaseParser`][python_multipart.BaseParser].
730 max_size: The maximum size of body to parse. Defaults to infinity - i.e. unbounded.
731 """
732
733 def __init__(self, callbacks: OctetStreamCallbacks = {}, max_size: float = float("inf")):
734 super().__init__()
735 self.callbacks = callbacks
736 self._started = False
737
738 if not isinstance(max_size, Number) or max_size < 1:
739 raise ValueError("max_size must be a positive number, not %r" % max_size)
740 self.max_size: int | float = max_size
741 self._current_size = 0
742
743 def write(self, data: bytes) -> int:
744 """Write some data to the parser, which will perform size verification,
745 and then pass the data to the underlying callback.
746
747 Args:
748 data: The data to write to the parser.
749
750 Returns:
751 The number of bytes written.
752 """
753 if not self._started:
754 self.callback("start")
755 self._started = True
756
757 # Truncate data length.
758 data_len = len(data)
759 if (self._current_size + data_len) > self.max_size:
760 # We truncate the length of data that we are to process.
761 new_size = int(self.max_size - self._current_size)
762 self.logger.warning(
763 "Current size is %d (max %d), so truncating data length from %d to %d",
764 self._current_size,
765 self.max_size,
766 data_len,
767 new_size,
768 )
769 data_len = new_size
770
771 # Increment size, then callback, in case there's an exception.
772 self._current_size += data_len
773 self.callback("data", data, 0, data_len)
774 return data_len
775
776 def finalize(self) -> None:
777 """Finalize this parser, which signals to that we are finished parsing,
778 and sends the on_end callback.
779 """
780 self.callback("end")
781
782 def __repr__(self) -> str:
783 return "%s()" % self.__class__.__name__
784
785
786class QuerystringParser(BaseParser):
787 """This is a streaming querystring parser. It will consume data, and call
788 the callbacks given when it has data.
789
790 | Callback Name | Parameters | Description |
791 |----------------|-----------------|-----------------------------------------------------|
792 | on_field_start | None | Called when a new field is encountered. |
793 | on_field_name | data, start, end| Called when a portion of a field's name is encountered. |
794 | on_field_data | data, start, end| Called when a portion of a field's data is encountered. |
795 | on_field_end | None | Called when the end of a field is encountered. |
796 | on_end | None | Called when the parser is finished parsing all data.|
797
798 Args:
799 callbacks: A dictionary of callbacks. See the documentation for [`BaseParser`][python_multipart.BaseParser].
800 strict_parsing: Whether or not to parse the body strictly. Defaults to False. If this is set to True, then the
801 behavior of the parser changes as the following: if a field has a value with an equal sign
802 (e.g. "foo=bar", or "foo="), it is always included. If a field has no equals sign (e.g. "...&name&..."),
803 it will be treated as an error if 'strict_parsing' is True, otherwise included. If an error is encountered,
804 then a [`QuerystringParseError`][python_multipart.exceptions.QuerystringParseError] will be raised.
805 max_size: The maximum size of body to parse. Defaults to infinity - i.e. unbounded.
806 """ # noqa: E501
807
808 state: QuerystringState
809
810 def __init__(
811 self, callbacks: QuerystringCallbacks = {}, strict_parsing: bool = False, max_size: float = float("inf")
812 ) -> None:
813 super().__init__()
814 self.state = BEFORE_FIELD
815 self._found_sep = False
816
817 self.callbacks = callbacks
818
819 # Max-size stuff
820 if not isinstance(max_size, Number) or max_size < 1:
821 raise ValueError("max_size must be a positive number, not %r" % max_size)
822 self.max_size: int | float = max_size
823 self._current_size = 0
824
825 # Should parsing be strict?
826 self.strict_parsing = strict_parsing
827
828 def write(self, data: bytes) -> int:
829 """Write some data to the parser, which will perform size verification,
830 parse into either a field name or value, and then pass the
831 corresponding data to the underlying callback. If an error is
832 encountered while parsing, a QuerystringParseError will be raised. The
833 "offset" attribute of the raised exception will be set to the offset in
834 the input data chunk (NOT the overall stream) that caused the error.
835
836 Args:
837 data: The data to write to the parser.
838
839 Returns:
840 The number of bytes written.
841 """
842 # Handle sizing.
843 data_len = len(data)
844 if (self._current_size + data_len) > self.max_size:
845 # We truncate the length of data that we are to process.
846 new_size = int(self.max_size - self._current_size)
847 self.logger.warning(
848 "Current size is %d (max %d), so truncating data length from %d to %d",
849 self._current_size,
850 self.max_size,
851 data_len,
852 new_size,
853 )
854 data_len = new_size
855
856 l = 0
857 try:
858 l = self._internal_write(data, data_len)
859 finally:
860 self._current_size += l
861
862 return l
863
864 def _internal_write(self, data: bytes, length: int) -> int:
865 state = self.state
866 strict_parsing = self.strict_parsing
867 found_sep = self._found_sep
868 callbacks = cast("QuerystringCallbacks", self.callbacks)
869 on_field_start = callbacks.get("on_field_start")
870 on_field_name = callbacks.get("on_field_name")
871 on_field_data = callbacks.get("on_field_data")
872 on_field_end = callbacks.get("on_field_end")
873 if on_field_start is None:
874 on_field_start = _noop_event
875 if on_field_name is None:
876 on_field_name = _noop_data
877 if on_field_data is None:
878 on_field_data = _noop_data
879 if on_field_end is None:
880 on_field_end = _noop_event
881
882 i = 0
883 while i < length:
884 ch = data[i]
885
886 # Depending on our state...
887 if state == BEFORE_FIELD:
888 # If the 'found_sep' flag is set, we've already encountered
889 # and skipped a single separator. If so, we check our strict
890 # parsing flag and decide what to do. Otherwise, we haven't
891 # yet reached a separator, and thus, if we do, we need to skip
892 # it as it will be the boundary between fields that's supposed
893 # to be there.
894 if ch == AMPERSAND:
895 if found_sep:
896 # If we're parsing strictly, we disallow blank chunks.
897 if strict_parsing:
898 raise QuerystringParseError("Skipping duplicate ampersand at %d" % i, offset=i)
899 else:
900 self.logger.debug("Skipping duplicate ampersand at %d", i)
901 else:
902 # This case is when we're skipping the (first)
903 # separator between fields, so we just set our flag
904 # and continue on.
905 found_sep = True
906 else:
907 # Emit a field-start event, and go to that state. Also,
908 # reset the "found_sep" flag, for the next time we get to
909 # this state.
910 on_field_start()
911 i -= 1
912 state = FIELD_NAME
913 found_sep = False
914
915 elif state == FIELD_NAME:
916 # Try and find a separator - we ensure that, if we do, we only
917 # look for the equal sign before it.
918 sep_pos = data.find(b"&", i, length)
919
920 # See if we can find an equals sign in the remaining data. If
921 # so, we can immediately emit the field name and jump to the
922 # data state.
923 if sep_pos != -1:
924 equals_pos = data.find(b"=", i, sep_pos)
925 else:
926 equals_pos = data.find(b"=", i, length)
927
928 if equals_pos != -1:
929 # Emit this name.
930 if i != equals_pos:
931 on_field_name(data, i, equals_pos)
932
933 # Jump i to this position. Note that it will then have 1
934 # added to it below, which means the next iteration of this
935 # loop will inspect the character after the equals sign.
936 i = equals_pos
937 state = FIELD_DATA
938 else:
939 # No equals sign found.
940 if not strict_parsing:
941 # See also comments in the QuerystringState.FIELD_DATA case below.
942 # If we found the separator, we emit the name and just
943 # end - there's no data callback at all (not even with
944 # a blank value).
945 if sep_pos != -1:
946 if i != sep_pos:
947 on_field_name(data, i, sep_pos)
948 on_field_end()
949
950 i = sep_pos - 1
951 state = BEFORE_FIELD
952 else:
953 # Otherwise, no separator in this block, so the
954 # rest of this chunk must be a name.
955 if i != length:
956 on_field_name(data, i, length)
957 i = length
958
959 else:
960 # We're parsing strictly. If we find a separator,
961 # this is an error - we require an equals sign.
962 if sep_pos != -1:
963 raise QuerystringParseError(
964 "When strict_parsing is True, we require an "
965 "equals sign in all field chunks. Did not "
966 "find one in the chunk that starts at %d" % (i,),
967 offset=i,
968 )
969
970 # No separator in the rest of this chunk, so it's just
971 # a field name.
972 if i != length:
973 on_field_name(data, i, length)
974 i = length
975
976 elif state == FIELD_DATA:
977 # Try finding an ampersand after this position.
978 sep_pos = data.find(b"&", i, length)
979
980 # If we found it, callback this bit as data and then go back
981 # to expecting to find a field.
982 if sep_pos != -1:
983 if i != sep_pos:
984 on_field_data(data, i, sep_pos)
985 on_field_end()
986
987 # Note that we go to the separator, which brings us to the
988 # "before field" state. This allows us to properly emit
989 # "field_start" events only when we actually have data for
990 # a field of some sort.
991 i = sep_pos - 1
992 state = BEFORE_FIELD
993
994 # Otherwise, emit the rest as data and finish.
995 else:
996 if i != length:
997 on_field_data(data, i, length)
998 i = length
999
1000 else: # pragma: no cover (error case)
1001 msg = "Reached an unknown state %d at %d" % (state, i)
1002 self.logger.warning(msg)
1003 raise QuerystringParseError(msg, offset=i)
1004
1005 i += 1
1006
1007 self.state = state
1008 self._found_sep = found_sep
1009 return length
1010
1011 def finalize(self) -> None:
1012 """Finalize this parser, which signals to that we are finished parsing,
1013 if we're still in the middle of a field, an on_field_end callback, and
1014 then the on_end callback.
1015 """
1016 callbacks = cast("QuerystringCallbacks", self.callbacks)
1017 # If we're currently in the middle of a field, we finish it.
1018 if self.state in (FIELD_DATA, FIELD_NAME):
1019 on_field_end = callbacks.get("on_field_end")
1020 if on_field_end is None:
1021 on_field_end = _noop_event
1022 on_field_end()
1023 on_end = callbacks.get("on_end")
1024 if on_end is None:
1025 on_end = _noop_event
1026 on_end()
1027
1028 def __repr__(self) -> str:
1029 return "{}(strict_parsing={!r}, max_size={!r})".format(
1030 self.__class__.__name__, self.strict_parsing, self.max_size
1031 )
1032
1033
1034class MultipartParser(BaseParser):
1035 """This class is a streaming multipart/form-data parser.
1036
1037 | Callback Name | Parameters | Description |
1038 |--------------------|-----------------|-------------|
1039 | on_part_begin | None | Called when a new part of the multipart message is encountered. |
1040 | on_part_data | data, start, end| Called when a portion of a part's data is encountered. |
1041 | on_part_end | None | Called when the end of a part is reached. |
1042 | on_header_begin | None | Called when we've found a new header in a part of a multipart message |
1043 | on_header_field | data, start, end| Called each time an additional portion of a header is read (i.e. the part of the header that is before the colon; the "Foo" in "Foo: Bar"). |
1044 | on_header_value | data, start, end| Called when we get data for a header. |
1045 | on_header_end | None | Called when the current header is finished - i.e. we've reached the newline at the end of the header. |
1046 | on_headers_finished| None | Called when all headers are finished, and before the part data starts. |
1047 | on_end | None | Called when the parser is finished parsing all data. |
1048
1049 Args:
1050 boundary: The multipart boundary. This is required, and must match what is given in the HTTP request - usually in the Content-Type header.
1051 callbacks: A dictionary of callbacks. See the documentation for [`BaseParser`][python_multipart.BaseParser].
1052 max_size: The maximum size of body to parse. Defaults to infinity - i.e. unbounded.
1053 max_header_count: The maximum number of headers allowed per part.
1054 max_header_size: The maximum size of a single header line (excluding the trailing CRLF).
1055 """ # noqa: E501
1056
1057 def __init__(
1058 self,
1059 boundary: bytes | str,
1060 callbacks: MultipartCallbacks = {},
1061 max_size: float = float("inf"),
1062 *,
1063 max_header_count: int = DEFAULT_MAX_HEADER_COUNT,
1064 max_header_size: int = DEFAULT_MAX_HEADER_SIZE,
1065 ) -> None:
1066 # Initialize parser state.
1067 super().__init__()
1068 self.state = START
1069 self.index = self.flags = 0
1070
1071 self.callbacks = callbacks
1072
1073 if not isinstance(max_size, Number) or max_size < 1:
1074 raise ValueError("max_size must be a positive number, not %r" % max_size)
1075 self.max_size = max_size
1076 self._current_size = 0
1077
1078 self.max_header_count = max_header_count
1079 self._current_header_count = 0
1080
1081 self.max_header_size = max_header_size
1082 self._current_header_size = 0
1083
1084 # Setup marks. These are used to track the state of data received.
1085 self.marks: dict[str, int] = {}
1086
1087 # Save our boundary.
1088 if isinstance(boundary, str): # pragma: no cover
1089 boundary = boundary.encode("latin-1")
1090 if len(boundary) > MAX_BOUNDARY_LENGTH:
1091 raise FormParserError(f"Boundary length {len(boundary)} exceeds maximum of {MAX_BOUNDARY_LENGTH}")
1092 self.boundary = b"\r\n--" + boundary
1093
1094 def write(self, data: bytes) -> int:
1095 """Write some data to the parser, which will perform size verification,
1096 and then parse the data into the appropriate location (e.g. header,
1097 data, etc.), and pass this on to the underlying callback. If an error
1098 is encountered, a MultipartParseError will be raised. The "offset"
1099 attribute on the raised exception will be set to the offset of the byte
1100 in the input chunk that caused the error.
1101
1102 Args:
1103 data: The data to write to the parser.
1104
1105 Returns:
1106 The number of bytes written.
1107 """
1108 # Handle sizing.
1109 data_len = len(data)
1110 if (self._current_size + data_len) > self.max_size:
1111 # We truncate the length of data that we are to process.
1112 new_size = int(self.max_size - self._current_size)
1113 self.logger.warning(
1114 "Current size is %d (max %d), so truncating data length from %d to %d",
1115 self._current_size,
1116 self.max_size,
1117 data_len,
1118 new_size,
1119 )
1120 data_len = new_size
1121
1122 l = 0
1123 try:
1124 l = self._internal_write(data, data_len)
1125 finally:
1126 self._current_size += l
1127
1128 return l
1129
1130 def _internal_write(self, data: bytes, length: int) -> int:
1131 # Get values from locals.
1132 boundary = self.boundary
1133 boundary_length = len(boundary)
1134
1135 # Get our state, flags and index. These are persisted between calls to
1136 # this function.
1137 state = self.state
1138 index = self.index
1139 flags = self.flags
1140 current_header_count = self._current_header_count
1141 current_header_size = self._current_header_size
1142
1143 # Our index defaults to 0.
1144 i = 0
1145
1146 def advance_header_size(amount: int = 1) -> None:
1147 nonlocal current_header_size
1148 current_header_size += amount
1149 if current_header_size > self.max_header_size:
1150 raise MultipartParseError("Maximum header size exceeded", offset=i)
1151
1152 # Set a mark.
1153 def set_mark(name: str) -> None:
1154 self.marks[name] = i
1155
1156 # Remove a mark.
1157 def delete_mark(name: str, reset: bool = False) -> None:
1158 self.marks.pop(name, None)
1159
1160 # Helper function that makes calling a callback with data easier. The
1161 # 'remaining' parameter will callback from the marked value until the
1162 # end of the buffer, and reset the mark, instead of deleting it. This
1163 # is used at the end of the function to call our callbacks with any
1164 # remaining data in this chunk.
1165 def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> None:
1166 marked_index = self.marks.get(name)
1167 if marked_index is None:
1168 return
1169
1170 # Otherwise, we call it from the mark to the current byte we're
1171 # processing.
1172 if end_i <= marked_index:
1173 # There is no additional data to send.
1174 pass
1175 elif marked_index >= 0:
1176 # We are emitting data from the local buffer.
1177 self.callback(name, data, marked_index, end_i)
1178 else:
1179 # Some of the data comes from a partial boundary match.
1180 # and requires look-behind.
1181 # We need to use self.flags (and not flags) because we care about
1182 # the state when we entered the loop.
1183 lookbehind_len = -marked_index
1184 if lookbehind_len <= boundary_length:
1185 self.callback(name, boundary, 0, lookbehind_len)
1186 elif self.flags & FLAG_PART_BOUNDARY:
1187 lookback = boundary + b"\r\n"
1188 self.callback(name, lookback, 0, lookbehind_len)
1189 elif self.flags & FLAG_LAST_BOUNDARY:
1190 lookback = boundary + b"--\r\n"
1191 self.callback(name, lookback, 0, lookbehind_len)
1192 else: # pragma: no cover (error case)
1193 self.logger.warning("Look-back buffer error")
1194
1195 if end_i > 0:
1196 self.callback(name, data, 0, end_i)
1197 # If we're getting remaining data, we have got all the data we
1198 # can be certain is not a boundary, leaving only a partial boundary match.
1199 if remaining:
1200 self.marks[name] = end_i - length
1201 else:
1202 self.marks.pop(name, None)
1203
1204 # For each byte...
1205 while i < length:
1206 c = data[i]
1207
1208 if state == START:
1209 # Skip leading newlines
1210 if c == CR or c == LF:
1211 i = data.find(b"-", i)
1212 if i == -1:
1213 # No boundary candidate in this chunk, so ignore the content after the leading CR/LF.
1214 i = length
1215 break
1216 continue
1217
1218 # index is used as in index into our boundary. Set to 0.
1219 index = 0
1220
1221 # Move to the next state, but decrement i so that we re-process
1222 # this character.
1223 state = START_BOUNDARY
1224 i -= 1
1225
1226 elif state == START_BOUNDARY:
1227 if index == 0 and data.startswith(boundary[2:], i, length):
1228 index = boundary_length - 2
1229 i += index
1230 continue
1231
1232 # Check to ensure that the last 2 characters in our boundary
1233 # are CRLF.
1234 if index == boundary_length - 2:
1235 if c == HYPHEN:
1236 # Potential empty message.
1237 state = END_BOUNDARY
1238 elif c != CR:
1239 # Error!
1240 msg = "Did not find CR at end of boundary (%d)" % (i,)
1241 self.logger.warning(msg)
1242 raise MultipartParseError(msg, offset=i)
1243
1244 index += 1
1245
1246 elif index == boundary_length - 1:
1247 if c != LF:
1248 msg = "Did not find LF at end of boundary (%d)" % (i,)
1249 self.logger.warning(msg)
1250 raise MultipartParseError(msg, offset=i)
1251
1252 # The index is now used for indexing into our boundary.
1253 index = 0
1254
1255 # Callback for the start of a part.
1256 self.callback("part_begin")
1257 current_header_count = 0
1258 current_header_size = 0
1259
1260 # Move to the next character and state.
1261 state = HEADER_FIELD_START
1262
1263 else:
1264 # Check to ensure our boundary matches
1265 if c != boundary[index + 2]:
1266 msg = "Expected boundary character %r, got %r at index %d" % (boundary[index + 2], c, index + 2)
1267 self.logger.warning(msg)
1268 raise MultipartParseError(msg, offset=i)
1269
1270 # Increment index into boundary and continue.
1271 index += 1
1272
1273 elif state == HEADER_FIELD_START:
1274 # Mark the start of a header field here, reset the index, and
1275 # continue parsing our header field.
1276 index = 0
1277
1278 if c != CR:
1279 current_header_count += 1
1280 if current_header_count > self.max_header_count:
1281 raise MultipartParseError("Maximum header count exceeded", offset=i)
1282 current_header_size = 0
1283
1284 # Set a mark of our header field.
1285 set_mark("header_field")
1286
1287 # Notify that we're starting a header if the next character is
1288 # not a CR; a CR at the beginning of the header will cause us
1289 # to stop parsing headers in the MultipartState.HEADER_FIELD state,
1290 # below.
1291 if c != CR:
1292 self.callback("header_begin")
1293
1294 # Move to parsing header fields.
1295 state = HEADER_FIELD
1296 i -= 1
1297
1298 elif state == HEADER_FIELD:
1299 # If we've reached a CR at the beginning of a header, it means
1300 # that we've reached the second of 2 newlines, and so there are
1301 # no more headers to parse.
1302 if c == CR and index == 0:
1303 delete_mark("header_field")
1304 state = HEADERS_ALMOST_DONE
1305 i += 1
1306 continue
1307
1308 # The field name runs until the colon; jump straight to it and
1309 # validate the whole span at once instead of byte by byte.
1310 colon = data.find(b":", i, length)
1311 end = colon if colon != -1 else length
1312
1313 # Enforce the size limit before slicing and validating, so an oversized header
1314 # name fails fast instead of copying and scanning a potentially huge span.
1315 advance_header_size(end - i if colon == -1 else end - i + 1)
1316
1317 field = data[i:end]
1318 if field.translate(None, TOKEN_CHARS):
1319 bad = next(b for b in field if b not in TOKEN_CHARS_SET)
1320 bad_i = i + field.index(bad)
1321 msg = "Found invalid character %r in header at %d" % (bad, bad_i)
1322 self.logger.warning(msg)
1323 raise MultipartParseError(msg, offset=bad_i)
1324
1325 index += end - i
1326 if colon == -1:
1327 # Field name continues into the next chunk.
1328 i = length
1329 else:
1330 # A 0-length header is an error.
1331 if index == 0:
1332 msg = "Found 0-length header at %d" % (i,)
1333 self.logger.warning(msg)
1334 raise MultipartParseError(msg, offset=i)
1335
1336 # Call our callback with the header field.
1337 i = colon
1338 data_callback("header_field", i)
1339
1340 # Move to parsing the header value.
1341 state = HEADER_VALUE_START
1342
1343 elif state == HEADER_VALUE_START:
1344 # Skip leading spaces.
1345 if c == SPACE:
1346 advance_header_size()
1347 i += 1
1348 continue
1349
1350 # Mark the start of the header value.
1351 set_mark("header_value")
1352
1353 # Move to the header-value state, reprocessing this character.
1354 state = HEADER_VALUE
1355 i -= 1
1356
1357 elif state == HEADER_VALUE:
1358 # The value runs until the terminating CR; jump straight to it
1359 # instead of inspecting every byte.
1360 cr = data.find(b"\r", i, length)
1361 end = cr if cr != -1 else length
1362 advance_header_size(end - i)
1363 if cr != -1:
1364 i = cr
1365 data_callback("header_value", i)
1366 self.callback("header_end")
1367 current_header_size = 0
1368 state = HEADER_VALUE_ALMOST_DONE
1369 else:
1370 i = length
1371
1372 elif state == HEADER_VALUE_ALMOST_DONE:
1373 # The last character should be a LF. If not, it's an error.
1374 if c != LF:
1375 msg = f"Did not find LF character at end of header (found {c!r})"
1376 self.logger.warning(msg)
1377 raise MultipartParseError(msg, offset=i)
1378
1379 # Move back to the start of another header. Note that if that
1380 # state detects ANOTHER newline, it'll trigger the end of our
1381 # headers.
1382 state = HEADER_FIELD_START
1383
1384 elif state == HEADERS_ALMOST_DONE:
1385 # We're almost done our headers. This is reached when we parse
1386 # a CR at the beginning of a header, so our next character
1387 # should be a LF, or it's an error.
1388 if c != LF:
1389 msg = f"Did not find LF at end of headers (found {c!r})"
1390 self.logger.warning(msg)
1391 raise MultipartParseError(msg, offset=i)
1392
1393 self.callback("headers_finished")
1394 state = PART_DATA_START
1395
1396 elif state == PART_DATA_START:
1397 # Mark the start of our part data.
1398 set_mark("part_data")
1399
1400 # Start processing part data, including this character.
1401 state = PART_DATA
1402 i -= 1
1403
1404 elif state == PART_DATA:
1405 # We're processing our part data right now. During this, we
1406 # need to efficiently search for our boundary, since any data
1407 # on any number of lines can be a part of the current data.
1408
1409 # Save the current value of our index. We use this in case we
1410 # find part of a boundary, but it doesn't match fully.
1411 prev_index = index
1412
1413 # If our index is 0, we're starting a new part, so start our
1414 # search.
1415 if index == 0:
1416 # The most common case is likely to be that the whole
1417 # boundary is present in the buffer.
1418 # Calling `find` is much faster than iterating here.
1419 i0 = data.find(boundary, i, length)
1420 if i0 >= 0:
1421 # We matched the whole boundary string.
1422 index = boundary_length - 1
1423 i = i0 + boundary_length - 1
1424 c = data[i]
1425 else:
1426 # No whole boundary, but the tail may hold a partial one
1427 # that completes in the next chunk. Boundary starts with
1428 # CR, which an RFC boundary contains nowhere else, so the
1429 # last CR in the tail is the only candidate prefix start.
1430 k = data.rfind(boundary[:1], max(i, length - boundary_length + 1), length)
1431 if k != -1 and boundary.startswith(data[k:length]):
1432 index = length - k
1433 # Carry the partial via index; the end-of-chunk flush
1434 # emits the data before it and re-marks the lookbehind.
1435 i = length
1436 continue
1437
1438 # Now, we have a couple of cases here. If our index is before
1439 # the end of the boundary...
1440 if index < boundary_length:
1441 # If the character matches...
1442 if boundary[index] == c:
1443 # The current character matches, so continue!
1444 index += 1
1445 else:
1446 index = 0
1447
1448 # Our index is equal to the length of our boundary!
1449 elif index == boundary_length:
1450 # First we increment it.
1451 index += 1
1452
1453 # Now, if we've reached a newline, we need to set this as
1454 # the potential end of our boundary.
1455 if c == CR:
1456 flags |= FLAG_PART_BOUNDARY
1457
1458 # Otherwise, if this is a hyphen, we might be at the last
1459 # of all boundaries.
1460 elif c == HYPHEN:
1461 flags |= FLAG_LAST_BOUNDARY
1462
1463 # Otherwise, we reset our index, since this isn't either a
1464 # newline or a hyphen.
1465 else:
1466 index = 0
1467
1468 # Our index is right after the part boundary, which should be
1469 # a LF.
1470 elif index == boundary_length + 1:
1471 # If we're at a part boundary (i.e. we've seen a CR
1472 # character already)...
1473 if flags & FLAG_PART_BOUNDARY:
1474 # We need a LF character next.
1475 if c == LF:
1476 # Unset the part boundary flag.
1477 flags &= ~FLAG_PART_BOUNDARY
1478
1479 # We have identified a boundary, callback for any data before it.
1480 data_callback("part_data", i - index)
1481 # Callback indicating that we've reached the end of
1482 # a part, and are starting a new one.
1483 self.callback("part_end")
1484 self.callback("part_begin")
1485 current_header_count = 0
1486 current_header_size = 0
1487
1488 # Move to parsing new headers.
1489 index = 0
1490 state = HEADER_FIELD_START
1491 i += 1
1492 continue
1493
1494 # We didn't find an LF character, so no match. Reset
1495 # our index and clear our flag.
1496 index = 0
1497 flags &= ~FLAG_PART_BOUNDARY
1498
1499 # Otherwise, if we're at the last boundary (i.e. we've
1500 # seen a hyphen already)...
1501 elif flags & FLAG_LAST_BOUNDARY:
1502 # We need a second hyphen here.
1503 if c == HYPHEN:
1504 # We have identified a boundary, callback for any data before it.
1505 data_callback("part_data", i - index)
1506 # Callback to end the current part, and then the
1507 # message.
1508 self.callback("part_end")
1509 self.callback("end")
1510 state = END
1511 else:
1512 # No match, so reset index.
1513 index = 0
1514
1515 # Otherwise, our index is 0. If the previous index is not, it
1516 # means we reset something, and we need to take the data we
1517 # thought was part of our boundary and send it along as actual
1518 # data.
1519 if index == 0 and prev_index > 0:
1520 # Overwrite our previous index.
1521 prev_index = 0
1522
1523 # Re-consider the current character, since this could be
1524 # the start of the boundary itself.
1525 i -= 1
1526
1527 elif state == END_BOUNDARY:
1528 if index == boundary_length - 1:
1529 if c != HYPHEN:
1530 msg = "Did not find - at end of boundary (%d)" % (i,)
1531 self.logger.warning(msg)
1532 raise MultipartParseError(msg, offset=i)
1533 index += 1
1534 self.callback("end")
1535 state = END
1536
1537 elif state == END:
1538 # Silently discard any epilogue data (RFC 2046 section 5.1.1 allows a CRLF and optional
1539 # epilogue after the closing boundary). Django and Werkzeug do the same.
1540 i = length
1541 break
1542
1543 else: # pragma: no cover (error case)
1544 # We got into a strange state somehow! Just stop processing.
1545 msg = "Reached an unknown state %d at %d" % (state, i)
1546 self.logger.warning(msg)
1547 raise MultipartParseError(msg, offset=i)
1548
1549 # Move to the next byte.
1550 i += 1
1551
1552 # We call our callbacks with any remaining data. Note that we pass
1553 # the 'remaining' flag, which sets the mark back to 0 instead of
1554 # deleting it, if it's found. This is because, if the mark is found
1555 # at this point, we assume that there's data for one of these things
1556 # that has been parsed, but not yet emitted. And, as such, it implies
1557 # that we haven't yet reached the end of this 'thing'. So, by setting
1558 # the mark to 0, we cause any data callbacks that take place in future
1559 # calls to this function to start from the beginning of that buffer.
1560 data_callback("header_field", length, True)
1561 data_callback("header_value", length, True)
1562 data_callback("part_data", length - index, True)
1563
1564 # Save values to locals.
1565 self.state = state
1566 self.index = index
1567 self.flags = flags
1568 self._current_header_count = current_header_count
1569 self._current_header_size = current_header_size
1570
1571 # Return our data length to indicate no errors, and that we processed
1572 # all of it.
1573 return length
1574
1575 def finalize(self) -> None:
1576 """Finalize this parser, which signals to that we are finished parsing.
1577
1578 Note: It does not currently, but in the future, it will verify that we
1579 are in the final state of the parser (i.e. the end of the multipart
1580 message is well-formed), and, if not, throw an error.
1581 """
1582 # TODO: verify that we're in the state MultipartState.END, otherwise throw an
1583 # error or otherwise state that we're not finished parsing.
1584 pass
1585
1586 def __repr__(self) -> str:
1587 return f"{self.__class__.__name__}(boundary={self.boundary!r})"
1588
1589
1590class FormParser:
1591 """This class is the all-in-one form parser. Given all the information
1592 necessary to parse a form, it will instantiate the correct parser, create
1593 the proper :class:`Field` and :class:`File` classes to store the data that
1594 is parsed, and call the two given callbacks with each field and file as
1595 they become available.
1596
1597 Args:
1598 content_type: The Content-Type of the incoming request. This is used to select the appropriate parser.
1599 on_field: The callback to call when a field has been parsed and is ready for usage. See above for parameters.
1600 on_file: The callback to call when a file has been parsed and is ready for usage. See above for parameters.
1601 on_end: An optional callback to call when all fields and files in a request has been parsed. Can be None.
1602 boundary: If the request is a multipart/form-data request, this should be the boundary of the request, as given
1603 in the Content-Type header, as a bytestring.
1604 file_name: If the request is of type application/octet-stream, then the body of the request will not contain any
1605 information about the uploaded file. In such cases, you can provide the file name of the uploaded file
1606 manually.
1607 config: Configuration to use for this FormParser. The default values are taken from the DEFAULT_CONFIG value,
1608 and then any keys present in this dictionary will overwrite the default values.
1609 """
1610
1611 #: This is the default configuration for our form parser.
1612 #: Note: all file sizes should be in bytes.
1613 DEFAULT_CONFIG: FormParserConfig = {
1614 "MAX_BODY_SIZE": float("inf"),
1615 "MAX_HEADER_COUNT": DEFAULT_MAX_HEADER_COUNT,
1616 "MAX_HEADER_SIZE": DEFAULT_MAX_HEADER_SIZE,
1617 "MAX_MEMORY_FILE_SIZE": 1 * 1024 * 1024,
1618 "UPLOAD_DIR": None,
1619 "UPLOAD_DELETE_TMP": True,
1620 "UPLOAD_KEEP_FILENAME": False,
1621 "UPLOAD_KEEP_EXTENSIONS": False,
1622 # Error on invalid Content-Transfer-Encoding?
1623 "UPLOAD_ERROR_ON_BAD_CTE": False,
1624 }
1625
1626 def __init__(
1627 self,
1628 content_type: str,
1629 on_field: Callable[[Field], None] | None,
1630 on_file: Callable[[File], None] | None,
1631 on_end: Callable[[], None] | None = None,
1632 boundary: bytes | str | None = None,
1633 file_name: bytes | None = None,
1634 config: dict[Any, Any] = {},
1635 ) -> None:
1636 self.logger = logging.getLogger(__name__)
1637
1638 # Save variables.
1639 self.content_type = content_type
1640 self.boundary = boundary
1641 self.bytes_received = 0
1642 self.parser = None
1643
1644 # Save callbacks.
1645 self.on_field = on_field
1646 self.on_file = on_file
1647 self.on_end = on_end
1648
1649 # Set configuration options.
1650 self.config: FormParserConfig = self.DEFAULT_CONFIG.copy()
1651 self.config.update(config) # type: ignore[typeddict-item]
1652
1653 parser: OctetStreamParser | MultipartParser | QuerystringParser | None = None
1654
1655 # Depending on the Content-Type, we instantiate the correct parser.
1656 if content_type == "application/octet-stream":
1657 file: File | None = None
1658
1659 def on_start() -> None:
1660 nonlocal file
1661 file = File(file_name, None, config=self.config)
1662
1663 def on_data(data: bytes, start: int, end: int) -> None:
1664 nonlocal file
1665 assert file is not None
1666 file.write(data[start:end])
1667
1668 def _on_end() -> None:
1669 nonlocal file
1670 assert file is not None
1671 # Finalize the file itself.
1672 file.finalize()
1673
1674 # Call our callback.
1675 if on_file:
1676 on_file(file)
1677
1678 # Call the on-end callback.
1679 if self.on_end is not None:
1680 self.on_end()
1681
1682 # Instantiate an octet-stream parser
1683 parser = OctetStreamParser(
1684 callbacks={"on_start": on_start, "on_data": on_data, "on_end": _on_end},
1685 max_size=self.config["MAX_BODY_SIZE"],
1686 )
1687
1688 elif content_type == "application/x-www-form-urlencoded" or content_type == "application/x-url-encoded":
1689 name_buffer: list[bytes] = []
1690
1691 f: Field | None = None
1692
1693 def on_field_start() -> None:
1694 pass
1695
1696 def on_field_name(data: bytes, start: int, end: int) -> None:
1697 name_buffer.append(data[start:end])
1698
1699 def on_field_data(data: bytes, start: int, end: int) -> None:
1700 nonlocal f
1701 if f is None:
1702 f = Field(b"".join(name_buffer))
1703 del name_buffer[:]
1704 f.write(data[start:end])
1705
1706 def on_field_end() -> None:
1707 nonlocal f
1708 # Finalize and call callback.
1709 if f is None:
1710 # If we get here, it's because there was no field data.
1711 # We create a field, set it to None, and then continue.
1712 f = Field(b"".join(name_buffer))
1713 del name_buffer[:]
1714 f.set_none()
1715
1716 f.finalize()
1717 if on_field:
1718 on_field(f)
1719 f = None
1720
1721 def _on_end() -> None:
1722 if self.on_end is not None:
1723 self.on_end()
1724
1725 # Instantiate parser.
1726 parser = QuerystringParser(
1727 callbacks={
1728 "on_field_start": on_field_start,
1729 "on_field_name": on_field_name,
1730 "on_field_data": on_field_data,
1731 "on_field_end": on_field_end,
1732 "on_end": _on_end,
1733 },
1734 max_size=self.config["MAX_BODY_SIZE"],
1735 )
1736
1737 elif content_type == "multipart/form-data":
1738 if boundary is None:
1739 self.logger.error("No boundary given")
1740 raise FormParserError("No boundary given")
1741
1742 header_name: list[bytes] = []
1743 header_value: list[bytes] = []
1744 headers: dict[bytes, bytes] = {}
1745
1746 f_multi: File | Field | None = None
1747 writer: File | Field | Base64Decoder | QuotedPrintableDecoder | None = None
1748 is_file = False
1749
1750 def on_part_begin() -> None:
1751 # Reset headers in case this isn't the first part.
1752 nonlocal headers
1753 headers = {}
1754
1755 def on_part_data(data: bytes, start: int, end: int) -> None:
1756 nonlocal writer
1757 assert writer is not None
1758 writer.write(data[start:end])
1759 # TODO: check for error here.
1760
1761 def on_part_end() -> None:
1762 nonlocal f_multi, is_file
1763 assert f_multi is not None
1764 f_multi.finalize()
1765 if is_file:
1766 if on_file:
1767 assert isinstance(f_multi, File)
1768 on_file(f_multi)
1769 else:
1770 if on_field:
1771 assert isinstance(f_multi, Field)
1772 on_field(f_multi)
1773
1774 def on_header_field(data: bytes, start: int, end: int) -> None:
1775 header_name.append(data[start:end])
1776
1777 def on_header_value(data: bytes, start: int, end: int) -> None:
1778 header_value.append(data[start:end])
1779
1780 def on_header_end() -> None:
1781 headers[b"".join(header_name).lower()] = b"".join(header_value)
1782 del header_name[:]
1783 del header_value[:]
1784
1785 def on_headers_finished() -> None:
1786 nonlocal is_file, f_multi, writer
1787 # Reset the 'is file' flag.
1788 is_file = False
1789
1790 # Parse the content-disposition header.
1791 content_disp = headers.get(b"content-disposition")
1792 disp, options = parse_options_header(content_disp)
1793
1794 # Get the field and filename.
1795 field_name = options.get(b"name")
1796 file_name = options.get(b"filename")
1797 # RFC 7578 §4.2: each part MUST have a Content-Disposition header with a "name" parameter.
1798 if field_name is None:
1799 raise FormParserError(f'Field name not found in Content-Disposition: "{content_disp!r}"')
1800
1801 # Create the proper class.
1802 content_type_b = headers.get(b"content-type")
1803 content_type = content_type_b.decode("latin-1") if content_type_b is not None else None
1804 if file_name is None:
1805 f_multi = Field(field_name, content_type=content_type)
1806 else:
1807 f_multi = File(file_name, field_name, config=self.config, content_type=content_type)
1808 is_file = True
1809
1810 # Parse the given Content-Transfer-Encoding to determine what
1811 # we need to do with the incoming data.
1812 # TODO: check that we properly handle 8bit / 7bit encoding.
1813 # RFC 2045 section 6.1: Content-Transfer-Encoding values are case-insensitive.
1814 # https://www.rfc-editor.org/rfc/rfc2045#section-6.1
1815 transfer_encoding = headers.get(b"content-transfer-encoding", b"7bit").lower()
1816
1817 if transfer_encoding in (b"binary", b"8bit", b"7bit"):
1818 writer = f_multi
1819
1820 elif transfer_encoding == b"base64":
1821 writer = Base64Decoder(f_multi)
1822
1823 elif transfer_encoding == b"quoted-printable":
1824 writer = QuotedPrintableDecoder(f_multi)
1825
1826 else:
1827 self.logger.warning("Unknown Content-Transfer-Encoding: %r", transfer_encoding)
1828 if self.config["UPLOAD_ERROR_ON_BAD_CTE"]:
1829 raise FormParserError(f'Unknown Content-Transfer-Encoding "{transfer_encoding!r}"')
1830 else:
1831 # If we aren't erroring, then we just treat this as an
1832 # unencoded Content-Transfer-Encoding.
1833 writer = f_multi
1834
1835 def _on_end() -> None:
1836 nonlocal writer
1837 if writer is not None:
1838 writer.finalize()
1839 if self.on_end is not None:
1840 self.on_end()
1841
1842 # Instantiate a multipart parser.
1843 parser = MultipartParser(
1844 boundary,
1845 callbacks={
1846 "on_part_begin": on_part_begin,
1847 "on_part_data": on_part_data,
1848 "on_part_end": on_part_end,
1849 "on_header_field": on_header_field,
1850 "on_header_value": on_header_value,
1851 "on_header_end": on_header_end,
1852 "on_headers_finished": on_headers_finished,
1853 "on_end": _on_end,
1854 },
1855 max_size=self.config["MAX_BODY_SIZE"],
1856 max_header_count=self.config["MAX_HEADER_COUNT"],
1857 max_header_size=self.config["MAX_HEADER_SIZE"],
1858 )
1859
1860 else:
1861 self.logger.warning("Unknown Content-Type: %r", content_type)
1862 raise FormParserError(f"Unknown Content-Type: {content_type}")
1863
1864 self.parser = parser
1865
1866 def write(self, data: bytes) -> int:
1867 """Write some data. The parser will forward this to the appropriate
1868 underlying parser.
1869
1870 Args:
1871 data: The data to write.
1872
1873 Returns:
1874 The number of bytes processed.
1875 """
1876 self.bytes_received += len(data)
1877 # TODO: check the parser's return value for errors?
1878 assert self.parser is not None
1879 return self.parser.write(data)
1880
1881 def finalize(self) -> None:
1882 """Finalize the parser."""
1883 if self.parser is not None and hasattr(self.parser, "finalize"):
1884 self.parser.finalize()
1885
1886 def close(self) -> None:
1887 """Close the parser."""
1888 if self.parser is not None and hasattr(self.parser, "close"):
1889 self.parser.close()
1890
1891 def __repr__(self) -> str:
1892 return f"{self.__class__.__name__}(content_type={self.content_type!r}, parser={self.parser!r})"
1893
1894
1895def create_form_parser(
1896 headers: dict[str, bytes],
1897 on_field: Callable[[Field], None] | None,
1898 on_file: Callable[[File], None] | None,
1899 config: dict[Any, Any] = {},
1900) -> FormParser:
1901 """This function is a helper function to aid in creating a FormParser
1902 instances. Given a dictionary-like headers object, it will determine
1903 the correct information needed, instantiate a FormParser with the
1904 appropriate values and given callbacks, and then return the corresponding
1905 parser.
1906
1907 Args:
1908 headers: A dictionary-like object of HTTP headers. The only required header is Content-Type.
1909 on_field: Callback to call with each parsed field.
1910 on_file: Callback to call with each parsed file.
1911 config: Configuration variables to pass to the FormParser.
1912 """
1913 content_type: str | bytes | None = headers.get("Content-Type")
1914 if content_type is None:
1915 logging.getLogger(__name__).warning("No Content-Type header given")
1916 raise ValueError("No Content-Type header given!")
1917
1918 # Boundaries are optional (the FormParser will raise if one is needed
1919 # but not given).
1920 content_type, params = parse_options_header(content_type)
1921 boundary = params.get(b"boundary")
1922
1923 # We need content_type to be a string, not a bytes object.
1924 content_type = content_type.decode("latin-1")
1925
1926 # Instantiate a form parser.
1927 form_parser = FormParser(content_type, on_field, on_file, boundary=boundary, config=config)
1928
1929 # Return our parser.
1930 return form_parser
1931
1932
1933def parse_form(
1934 headers: dict[str, bytes],
1935 input_stream: SupportsRead,
1936 on_field: Callable[[Field], None] | None,
1937 on_file: Callable[[File], None] | None,
1938 chunk_size: int = 1048576,
1939) -> None:
1940 """This function is useful if you just want to parse a request body,
1941 without too much work. Pass it a dictionary-like object of the request's
1942 headers, and a file-like object for the input stream, along with two
1943 callbacks that will get called whenever a field or file is parsed.
1944
1945 Args:
1946 headers: A dictionary-like object of HTTP headers. The only required header is Content-Type.
1947 input_stream: A file-like object that represents the request body. The read() method must return bytestrings.
1948 on_field: Callback to call with each parsed field.
1949 on_file: Callback to call with each parsed file.
1950 chunk_size: The maximum size to read from the input stream and write to the parser at one time.
1951 Defaults to 1 MiB.
1952 """
1953 if chunk_size < 1:
1954 raise ValueError(f"chunk_size must be a positive number, not {chunk_size!r}")
1955
1956 # Create our form parser.
1957 parser = create_form_parser(headers, on_field, on_file)
1958
1959 # Read chunks of 1MiB and write to the parser, but never read more than
1960 # the given Content-Length, if any.
1961 content_length: int | float | bytes | None = headers.get("Content-Length")
1962 if content_length is not None:
1963 content_length = int(content_length)
1964 if content_length < 0:
1965 raise ValueError("Content-Length must be non-negative")
1966 else:
1967 content_length = float("inf")
1968 bytes_read = 0
1969
1970 while True:
1971 # Read only up to the Content-Length given.
1972 max_readable = int(min(content_length - bytes_read, chunk_size))
1973 buff = input_stream.read(max_readable)
1974
1975 # Write to the parser and update our length.
1976 parser.write(buff)
1977 bytes_read += len(buff)
1978
1979 # If we get a buffer that's smaller than the size requested, or if we
1980 # have read up to our content length, we're done.
1981 if len(buff) != max_readable or bytes_read == content_length:
1982 break
1983
1984 # Tell our parser that we're done writing data.
1985 parser.finalize()