Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
95 changes: 58 additions & 37 deletions python_multipart/multipart.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
import shutil
import sys
import tempfile
from enum import IntEnum
from enum import IntEnum, global_enum
from io import BufferedRandom, BytesIO
from numbers import Number
from typing import TYPE_CHECKING, cast
Expand Down Expand Up @@ -86,6 +86,7 @@ def _noop_data(_data: bytes, _start: int, _end: int) -> None:
pass


@global_enum
class QuerystringState(IntEnum):
"""Querystring parser states.

Expand All @@ -98,6 +99,7 @@ class QuerystringState(IntEnum):
FIELD_DATA = 2


@global_enum
class MultipartState(IntEnum):
"""Multipart parser states.

Expand All @@ -120,6 +122,25 @@ class MultipartState(IntEnum):
END = 12


if TYPE_CHECKING:
BEFORE_FIELD = QuerystringState.BEFORE_FIELD
FIELD_NAME = QuerystringState.FIELD_NAME
FIELD_DATA = QuerystringState.FIELD_DATA

START = MultipartState.START
START_BOUNDARY = MultipartState.START_BOUNDARY
HEADER_FIELD_START = MultipartState.HEADER_FIELD_START
HEADER_FIELD = MultipartState.HEADER_FIELD
HEADER_VALUE_START = MultipartState.HEADER_VALUE_START
HEADER_VALUE = MultipartState.HEADER_VALUE
HEADER_VALUE_ALMOST_DONE = MultipartState.HEADER_VALUE_ALMOST_DONE
HEADERS_ALMOST_DONE = MultipartState.HEADERS_ALMOST_DONE
PART_DATA_START = MultipartState.PART_DATA_START
PART_DATA = MultipartState.PART_DATA
PART_DATA_END = MultipartState.PART_DATA_END
END_BOUNDARY = MultipartState.END_BOUNDARY
END = MultipartState.END

Comment thread
Kludex marked this conversation as resolved.
# Flags for the multipart parser.
FLAG_PART_BOUNDARY = 1
FLAG_LAST_BOUNDARY = 2
Expand Down Expand Up @@ -790,7 +811,7 @@ def __init__(
self, callbacks: QuerystringCallbacks = {}, strict_parsing: bool = False, max_size: float = float("inf")
) -> None:
super().__init__()
self.state = QuerystringState.BEFORE_FIELD
self.state = BEFORE_FIELD
self._found_sep = False

self.callbacks = callbacks
Expand Down Expand Up @@ -863,7 +884,7 @@ def _internal_write(self, data: bytes, length: int) -> int:
ch = data[i]

# Depending on our state...
if state == QuerystringState.BEFORE_FIELD:
if state == BEFORE_FIELD:
# If the 'found_sep' flag is set, we've already encountered
# and skipped a single separator. If so, we check our strict
# parsing flag and decide what to do. Otherwise, we haven't
Expand All @@ -888,10 +909,10 @@ def _internal_write(self, data: bytes, length: int) -> int:
# this state.
on_field_start()
i -= 1
state = QuerystringState.FIELD_NAME
state = FIELD_NAME
found_sep = False

elif state == QuerystringState.FIELD_NAME:
elif state == FIELD_NAME:
# Try and find a separator - we ensure that, if we do, we only
# look for the equal sign before it.
sep_pos = data.find(b"&", i, length)
Expand All @@ -913,7 +934,7 @@ def _internal_write(self, data: bytes, length: int) -> int:
# added to it below, which means the next iteration of this
# loop will inspect the character after the equals sign.
i = equals_pos
state = QuerystringState.FIELD_DATA
state = FIELD_DATA
else:
# No equals sign found.
if not strict_parsing:
Expand All @@ -927,7 +948,7 @@ def _internal_write(self, data: bytes, length: int) -> int:
on_field_end()

i = sep_pos - 1
state = QuerystringState.BEFORE_FIELD
state = BEFORE_FIELD
else:
# Otherwise, no separator in this block, so the
# rest of this chunk must be a name.
Expand All @@ -952,7 +973,7 @@ def _internal_write(self, data: bytes, length: int) -> int:
on_field_name(data, i, length)
i = length

elif state == QuerystringState.FIELD_DATA:
elif state == FIELD_DATA:
# Try finding an ampersand after this position.
sep_pos = data.find(b"&", i, length)

Expand All @@ -968,7 +989,7 @@ def _internal_write(self, data: bytes, length: int) -> int:
# "field_start" events only when we actually have data for
# a field of some sort.
i = sep_pos - 1
state = QuerystringState.BEFORE_FIELD
state = BEFORE_FIELD

# Otherwise, emit the rest as data and finish.
else:
Expand All @@ -994,7 +1015,7 @@ def finalize(self) -> None:
"""
callbacks = cast("QuerystringCallbacks", self.callbacks)
# If we're currently in the middle of a field, we finish it.
if self.state in (QuerystringState.FIELD_DATA, QuerystringState.FIELD_NAME):
if self.state in (FIELD_DATA, FIELD_NAME):
on_field_end = callbacks.get("on_field_end")
if on_field_end is None:
on_field_end = _noop_event
Expand Down Expand Up @@ -1044,7 +1065,7 @@ def __init__(
) -> None:
# Initialize parser state.
super().__init__()
self.state = MultipartState.START
self.state = START
self.index = self.flags = 0

self.callbacks = callbacks
Expand Down Expand Up @@ -1184,7 +1205,7 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
while i < length:
c = data[i]

if state == MultipartState.START:
if state == START:
# Skip leading newlines
if c == CR or c == LF:
i = data.find(b"-", i)
Expand All @@ -1199,10 +1220,10 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No

# Move to the next state, but decrement i so that we re-process
# this character.
state = MultipartState.START_BOUNDARY
state = START_BOUNDARY
i -= 1

elif state == MultipartState.START_BOUNDARY:
elif state == START_BOUNDARY:
if index == 0 and data.startswith(boundary[2:], i, length):
index = boundary_length - 2
i += index
Expand All @@ -1213,7 +1234,7 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
if index == boundary_length - 2:
if c == HYPHEN:
# Potential empty message.
state = MultipartState.END_BOUNDARY
state = END_BOUNDARY
elif c != CR:
# Error!
msg = "Did not find CR at end of boundary (%d)" % (i,)
Expand All @@ -1237,7 +1258,7 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
current_header_size = 0

# Move to the next character and state.
state = MultipartState.HEADER_FIELD_START
state = HEADER_FIELD_START

else:
# Check to ensure our boundary matches
Expand All @@ -1249,7 +1270,7 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
# Increment index into boundary and continue.
index += 1

elif state == MultipartState.HEADER_FIELD_START:
elif state == HEADER_FIELD_START:
# Mark the start of a header field here, reset the index, and
# continue parsing our header field.
index = 0
Expand All @@ -1271,16 +1292,16 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
self.callback("header_begin")

# Move to parsing header fields.
state = MultipartState.HEADER_FIELD
state = HEADER_FIELD
i -= 1

elif state == MultipartState.HEADER_FIELD:
elif state == HEADER_FIELD:
# If we've reached a CR at the beginning of a header, it means
# that we've reached the second of 2 newlines, and so there are
# no more headers to parse.
if c == CR and index == 0:
delete_mark("header_field")
state = MultipartState.HEADERS_ALMOST_DONE
state = HEADERS_ALMOST_DONE
i += 1
continue

Expand Down Expand Up @@ -1317,9 +1338,9 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
data_callback("header_field", i)

# Move to parsing the header value.
state = MultipartState.HEADER_VALUE_START
state = HEADER_VALUE_START

elif state == MultipartState.HEADER_VALUE_START:
elif state == HEADER_VALUE_START:
# Skip leading spaces.
if c == SPACE:
advance_header_size()
Expand All @@ -1330,10 +1351,10 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
set_mark("header_value")

# Move to the header-value state, reprocessing this character.
state = MultipartState.HEADER_VALUE
state = HEADER_VALUE
i -= 1

elif state == MultipartState.HEADER_VALUE:
elif state == HEADER_VALUE:
# The value runs until the terminating CR; jump straight to it
# instead of inspecting every byte.
cr = data.find(b"\r", i, length)
Expand All @@ -1344,11 +1365,11 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
data_callback("header_value", i)
self.callback("header_end")
current_header_size = 0
state = MultipartState.HEADER_VALUE_ALMOST_DONE
state = HEADER_VALUE_ALMOST_DONE
else:
i = length

elif state == MultipartState.HEADER_VALUE_ALMOST_DONE:
elif state == HEADER_VALUE_ALMOST_DONE:
# The last character should be a LF. If not, it's an error.
if c != LF:
msg = f"Did not find LF character at end of header (found {c!r})"
Expand All @@ -1358,9 +1379,9 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
# Move back to the start of another header. Note that if that
# state detects ANOTHER newline, it'll trigger the end of our
# headers.
state = MultipartState.HEADER_FIELD_START
state = HEADER_FIELD_START

elif state == MultipartState.HEADERS_ALMOST_DONE:
elif state == HEADERS_ALMOST_DONE:
# We're almost done our headers. This is reached when we parse
# a CR at the beginning of a header, so our next character
# should be a LF, or it's an error.
Expand All @@ -1370,17 +1391,17 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
raise MultipartParseError(msg, offset=i)

self.callback("headers_finished")
state = MultipartState.PART_DATA_START
state = PART_DATA_START

elif state == MultipartState.PART_DATA_START:
elif state == PART_DATA_START:
# Mark the start of our part data.
set_mark("part_data")

# Start processing part data, including this character.
state = MultipartState.PART_DATA
state = PART_DATA
i -= 1

elif state == MultipartState.PART_DATA:
elif state == PART_DATA:
# We're processing our part data right now. During this, we
# need to efficiently search for our boundary, since any data
# on any number of lines can be a part of the current data.
Expand Down Expand Up @@ -1466,7 +1487,7 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No

# Move to parsing new headers.
index = 0
state = MultipartState.HEADER_FIELD_START
state = HEADER_FIELD_START
i += 1
continue

Expand All @@ -1486,7 +1507,7 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
# message.
self.callback("part_end")
self.callback("end")
state = MultipartState.END
state = END
else:
# No match, so reset index.
index = 0
Expand All @@ -1503,17 +1524,17 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
# the start of the boundary itself.
i -= 1

elif state == MultipartState.END_BOUNDARY:
elif state == END_BOUNDARY:
if index == boundary_length - 1:
if c != HYPHEN:
msg = "Did not find - at end of boundary (%d)" % (i,)
self.logger.warning(msg)
raise MultipartParseError(msg, offset=i)
index += 1
self.callback("end")
state = MultipartState.END
state = END

elif state == MultipartState.END:
elif state == END:
# Silently discard any epilogue data (RFC 2046 section 5.1.1 allows a CRLF and optional
# epilogue after the closing boundary). Django and Werkzeug do the same.
i = length
Expand Down
16 changes: 16 additions & 0 deletions tests/test_multipart.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,13 +21,17 @@
QuerystringParseError,
)
from python_multipart.multipart import (
END,
FIELD_DATA,
BaseParser,
Field,
File,
FormParser,
MultipartParser,
MultipartState,
OctetStreamParser,
QuerystringParser,
QuerystringState,
create_form_parser,
parse_form,
parse_options_header,
Expand Down Expand Up @@ -850,6 +854,18 @@ def test_content_transfer_encoding_is_case_insensitive(content_transfer_encoding
assert file.file_object.read() == b"Test"


def test_parser_states_remain_enum_members() -> None:
multipart_parser = MultipartParser(b"boundary")
multipart_parser.write(b"--boundary--\r\n")
multipart_parser.finalize()
assert multipart_parser.state is MultipartState.END is END

querystring_parser = QuerystringParser()
querystring_parser.write(b"field=value")
querystring_parser.finalize()
assert querystring_parser.state is QuerystringState.FIELD_DATA is FIELD_DATA


def test_multipart_opening_boundary_max_size() -> None:
data = b"--boundary\r\n\r\nvalue\r\n--boundary--"
max_size = 9
Expand Down
Loading