Skip to content
Closed
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
1 change: 1 addition & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

## Unreleased

* Speed up multipart callback dispatch [#317](https://github.com/Kludex/python-multipart/pull/317).
* Speed up querystring callback dispatch [#316](https://github.com/Kludex/python-multipart/pull/316).

## 0.0.32 (2026-06-04)
Expand Down
73 changes: 51 additions & 22 deletions python_multipart/multipart.py
Original file line number Diff line number Diff line change
Expand Up @@ -1048,6 +1048,24 @@ def __init__(
self.index = self.flags = 0

self.callbacks = callbacks
on_part_begin = callbacks.get("on_part_begin")
on_part_data = callbacks.get("on_part_data")
on_part_end = callbacks.get("on_part_end")
on_header_begin = callbacks.get("on_header_begin")
on_header_field = callbacks.get("on_header_field")
on_header_value = callbacks.get("on_header_value")
on_header_end = callbacks.get("on_header_end")
on_headers_finished = callbacks.get("on_headers_finished")
on_end = callbacks.get("on_end")
self._on_part_begin = _noop_event if on_part_begin is None else on_part_begin
self._on_part_data = _noop_data if on_part_data is None else on_part_data
self._on_part_end = _noop_event if on_part_end is None else on_part_end
self._on_header_begin = _noop_event if on_header_begin is None else on_header_begin
self._on_header_field = _noop_data if on_header_field is None else on_header_field
self._on_header_value = _noop_data if on_header_value is None else on_header_value
self._on_header_end = _noop_event if on_header_end is None else on_header_end
self._on_headers_finished = _noop_event if on_headers_finished is None else on_headers_finished
self._on_end = _noop_event if on_end is None else on_end

if not isinstance(max_size, Number) or max_size < 1:
raise ValueError("max_size must be a positive number, not %r" % max_size)
Expand All @@ -1070,6 +1088,15 @@ def __init__(
raise FormParserError(f"Boundary length {len(boundary)} exceeds maximum of {MAX_BOUNDARY_LENGTH}")
self.boundary = b"\r\n--" + boundary

def set_callback(self, name: CallbackName, new_func: Callable[..., Any] | None) -> None:
super().set_callback(name, new_func)
if name in ("part_data", "header_field", "header_value"):
callback = _noop_data if new_func is None else new_func
setattr(self, "_on_" + name, callback)
elif name in ("part_begin", "part_end", "header_begin", "header_end", "headers_finished", "end"):
callback = _noop_event if new_func is None else new_func
setattr(self, "_on_" + name, callback)

def write(self, data: bytes) -> int:
"""Write some data to the parser, which will perform size verification,
and then parse the data into the appropriate location (e.g. header,
Expand Down Expand Up @@ -1141,7 +1168,9 @@ def delete_mark(name: str, reset: bool = False) -> None:
# end of the buffer, and reset the mark, instead of deleting it. This
# is used at the end of the function to call our callbacks with any
# remaining data in this chunk.
def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> None:
def data_callback(
name: CallbackName, callback: Callable[[bytes, int, int], None], end_i: int, remaining: bool = False
) -> None:
marked_index = self.marks.get(name)
if marked_index is None:
return
Expand All @@ -1153,26 +1182,26 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
pass
elif marked_index >= 0:
# We are emitting data from the local buffer.
self.callback(name, data, marked_index, end_i)
callback(data, marked_index, end_i)
else:
# Some of the data comes from a partial boundary match.
# and requires look-behind.
# We need to use self.flags (and not flags) because we care about
# the state when we entered the loop.
lookbehind_len = -marked_index
if lookbehind_len <= boundary_length:
self.callback(name, boundary, 0, lookbehind_len)
callback(boundary, 0, lookbehind_len)
elif self.flags & FLAG_PART_BOUNDARY:
lookback = boundary + b"\r\n"
self.callback(name, lookback, 0, lookbehind_len)
callback(lookback, 0, lookbehind_len)
elif self.flags & FLAG_LAST_BOUNDARY:
lookback = boundary + b"--\r\n"
self.callback(name, lookback, 0, lookbehind_len)
callback(lookback, 0, lookbehind_len)
else: # pragma: no cover (error case)
self.logger.warning("Look-back buffer error")

if end_i > 0:
self.callback(name, data, 0, end_i)
callback(data, 0, end_i)
# If we're getting remaining data, we have got all the data we
# can be certain is not a boundary, leaving only a partial boundary match.
if remaining:
Expand Down Expand Up @@ -1227,7 +1256,7 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
index = 0

# Callback for the start of a part.
self.callback("part_begin")
self._on_part_begin()
current_header_count = 0
current_header_size = 0

Expand Down Expand Up @@ -1263,7 +1292,7 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
# to stop parsing headers in the MultipartState.HEADER_FIELD state,
# below.
if c != CR:
self.callback("header_begin")
self._on_header_begin()

# Move to parsing header fields.
state = MultipartState.HEADER_FIELD
Expand Down Expand Up @@ -1309,7 +1338,7 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No

# Call our callback with the header field.
i = colon
data_callback("header_field", i)
data_callback("header_field", self._on_header_field, i)

# Move to parsing the header value.
state = MultipartState.HEADER_VALUE_START
Expand All @@ -1336,8 +1365,8 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
advance_header_size(end - i)
if cr != -1:
i = cr
data_callback("header_value", i)
self.callback("header_end")
data_callback("header_value", self._on_header_value, i)
self._on_header_end()
current_header_size = 0
state = MultipartState.HEADER_VALUE_ALMOST_DONE
else:
Expand All @@ -1364,7 +1393,7 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
self.logger.warning(msg)
raise MultipartParseError(msg, offset=i)

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

elif state == MultipartState.PART_DATA_START:
Expand Down Expand Up @@ -1451,11 +1480,11 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
flags &= ~FLAG_PART_BOUNDARY

# We have identified a boundary, callback for any data before it.
data_callback("part_data", i - index)
data_callback("part_data", self._on_part_data, i - index)
# Callback indicating that we've reached the end of
# a part, and are starting a new one.
self.callback("part_end")
self.callback("part_begin")
self._on_part_end()
self._on_part_begin()
current_header_count = 0
current_header_size = 0

Expand All @@ -1476,11 +1505,11 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
# We need a second hyphen here.
if c == HYPHEN:
# We have identified a boundary, callback for any data before it.
data_callback("part_data", i - index)
data_callback("part_data", self._on_part_data, i - index)
# Callback to end the current part, and then the
# message.
self.callback("part_end")
self.callback("end")
self._on_part_end()
self._on_end()
state = MultipartState.END
else:
# No match, so reset index.
Expand All @@ -1505,7 +1534,7 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
self.logger.warning(msg)
raise MultipartParseError(msg, offset=i)
index += 1
self.callback("end")
self._on_end()
state = MultipartState.END

elif state == MultipartState.END:
Expand All @@ -1531,9 +1560,9 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No
# that we haven't yet reached the end of this 'thing'. So, by setting
# the mark to 0, we cause any data callbacks that take place in future
# calls to this function to start from the beginning of that buffer.
data_callback("header_field", length, True)
data_callback("header_value", length, True)
data_callback("part_data", length - index, True)
data_callback("header_field", self._on_header_field, length, True)
data_callback("header_value", self._on_header_value, length, True)
data_callback("part_data", self._on_part_data, length - index, True)

# Save values to locals.
self.state = state
Expand Down
39 changes: 39 additions & 0 deletions tests/test_multipart.py
Original file line number Diff line number Diff line change
Expand Up @@ -1592,6 +1592,45 @@ def test_invalid_max_size_multipart(self) -> None:
with self.assertRaises(ValueError):
MultipartParser(b"bound", max_size="foo") # type: ignore[arg-type]

def test_multipart_set_callback(self) -> None:
header_begins = 0
header_fields: list[bytes] = []

def on_header_begin() -> None:
nonlocal header_begins
header_begins += 1

def on_header_field(data: bytes, start: int, end: int) -> None:
header_fields.append(data[start:end])

parser = MultipartParser(b"boundary")
parser.set_callback("header_begin", on_header_begin)
parser.set_callback("header_field", on_header_field)
parser.set_callback("part_data", None)
data = b"--boundary\r\nX: y\r\n\r\nbody\r\n--boundary--\r\n"
parser.write(data)

self.assertEqual(header_begins, 1)
self.assertEqual(header_fields, [b"X"])

def test_multipart_none_callbacks(self) -> None:
callbacks: Any = {
"on_part_begin": None,
"on_part_data": None,
"on_part_end": None,
"on_header_begin": None,
"on_header_field": None,
"on_header_value": None,
"on_header_end": None,
"on_headers_finished": None,
"on_end": None,
}
parser = MultipartParser(b"boundary", callbacks)
data = b"--boundary\r\nX: y\r\n\r\nbody\r\n--boundary--\r\n"

self.assertEqual(parser.write(data), len(data))
parser.finalize()

def test_boundary_too_long(self) -> None:
with self.assertRaisesRegex(FormParserError, "Boundary length 257 exceeds maximum of 256"):
MultipartParser(b"x" * 257)
Expand Down
Loading