diff --git a/CHANGELOG.md b/CHANGELOG.md index aeb17fe..f4d5849 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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) diff --git a/python_multipart/multipart.py b/python_multipart/multipart.py index 49cdd8e..5f20891 100644 --- a/python_multipart/multipart.py +++ b/python_multipart/multipart.py @@ -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) @@ -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, @@ -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 @@ -1153,7 +1182,7 @@ 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. @@ -1161,18 +1190,18 @@ def data_callback(name: CallbackName, end_i: int, remaining: bool = False) -> No # 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: @@ -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 @@ -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 @@ -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 @@ -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: @@ -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: @@ -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 @@ -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. @@ -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: @@ -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 diff --git a/tests/test_multipart.py b/tests/test_multipart.py index 949d70c..4c83334 100644 --- a/tests/test_multipart.py +++ b/tests/test_multipart.py @@ -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)