diff --git a/src/h2/stream.py b/src/h2/stream.py index 249f73e0c..e3384eecb 100644 --- a/src/h2/stream.py +++ b/src/h2/stream.py @@ -1094,12 +1094,17 @@ def receive_headers(self, ).stream_ended = cast("StreamEnded", es_events[0]) events += es_events - self._initialize_content_length(headers) - if isinstance(headers_event, TrailersReceived) and not end_stream: msg = "Trailers must have END_STREAM set" raise ProtocolError(msg) + if isinstance(headers_event, TrailersReceived): + # Trailers are not part of the content, but the stream ends here, + # so this is the only point at which the body length can be policed. + self._track_content_length(0, end_stream=True) + else: + self._initialize_content_length(headers) + hdr_validation_flags = self._build_hdr_validation_flags(events) headers_event.headers = self._process_received_headers( headers, hdr_validation_flags, header_encoding, diff --git a/tests/test_basic_logic.py b/tests/test_basic_logic.py index d1bc0f1eb..d340c03f3 100644 --- a/tests/test_basic_logic.py +++ b/tests/test_basic_logic.py @@ -716,7 +716,7 @@ def test_can_receive_trailers(self, frame_factory) -> None: c.receive_data(f.serialize()) # Send in trailers. - trailers = [("content-length", "0")] + trailers = [("x-checksum", "0")] f = frame_factory.build_headers_frame( trailers, flags=["END_STREAM"], @@ -742,7 +742,7 @@ def test_reject_trailers_not_ending_stream(self, frame_factory) -> None: # Send in trailers. c.clear_outbound_data_buffer() - trailers = [("content-length", "0")] + trailers = [("x-checksum", "0")] f = frame_factory.build_headers_frame( trailers, flags=[], @@ -1646,7 +1646,7 @@ def test_can_receive_trailers(self, frame_factory) -> None: c.receive_data(f.serialize()) # Send in trailers. - trailers = [("content-length", "0")] + trailers = [("x-checksum", "0")] f = frame_factory.build_headers_frame( trailers, flags=["END_STREAM"], @@ -1671,7 +1671,7 @@ def test_reject_trailers_not_ending_stream(self, frame_factory) -> None: # Send in trailers. c.clear_outbound_data_buffer() - trailers = [("content-length", "0")] + trailers = [("x-checksum", "0")] f = frame_factory.build_headers_frame( trailers, flags=[], diff --git a/tests/test_invalid_content_lengths.py b/tests/test_invalid_content_lengths.py index 3927fb5e2..f2a1605e6 100644 --- a/tests/test_invalid_content_lengths.py +++ b/tests/test_invalid_content_lengths.py @@ -255,3 +255,119 @@ def test_insufficient_data_empty_frame(self, frame_factory, request_headers) -> error_code=h2.errors.ErrorCodes.PROTOCOL_ERROR, ) assert c.data_to_send() == expected_frame.serialize() + + +class TestContentLengthEnforcedAtTrailers: + """ + RFC 9113 ยง 8.1.1: a request or response is malformed if the value of a + content-length header field does not equal the sum of the DATA frame + payload lengths that form the content. The listed exemptions are 204, 304 + and HEAD, none of which is a trailers section, so a stream that ends with + trailers must still have its body length policed. + + A trailers section may not carry a content-length header field at all, so + it can never redefine the expected length either. + """ + + example_request_headers = [ + (":authority", "example.com"), + (":path", "/"), + (":scheme", "https"), + (":method", "POST"), + ("content-length", "15"), + ] + server_config = h2.config.H2Configuration(client_side=False) + + def _server(self, frame_factory, request_headers) -> h2.connection.H2Connection: + c = h2.connection.H2Connection(config=self.server_config) + c.initiate_connection() + c.receive_data(frame_factory.preamble()) + c.receive_data(frame_factory.build_headers_frame(headers=request_headers).serialize()) + return c + + @pytest.mark.parametrize("request_headers", [example_request_headers]) + def test_insufficient_data_ended_by_trailers(self, frame_factory, request_headers) -> None: + """ + Remote peers sending less data than content-length and then ending the + stream with trailers causes Protocol Errors. + """ + c = self._server(frame_factory, request_headers) + c.receive_data(frame_factory.build_data_frame(data=b"\x01"*13).serialize()) + c.clear_outbound_data_buffer() + + trailers = frame_factory.build_headers_frame( + headers=[("x-checksum", "0")], + flags=["END_STREAM"], + ) + with pytest.raises(h2.exceptions.InvalidBodyLengthError) as exp: + c.receive_data(trailers.serialize()) + + assert exp.value.expected_length == 15 + assert exp.value.actual_length == 13 + assert str(exp.value) == ( + "InvalidBodyLengthError: Expected 15 bytes, received 13" + ) + + expected_frame = frame_factory.build_goaway_frame( + last_stream_id=1, + error_code=h2.errors.ErrorCodes.PROTOCOL_ERROR, + ) + assert c.data_to_send() == expected_frame.serialize() + + def test_no_data_ended_by_trailers(self, frame_factory) -> None: + """ + Remote peers sending no data at all for a non-zero content-length and + then ending the stream with trailers causes Protocol Errors. + """ + c = self._server(frame_factory, self.example_request_headers) + c.clear_outbound_data_buffer() + + trailers = frame_factory.build_headers_frame( + headers=[("x-checksum", "0")], + flags=["END_STREAM"], + ) + with pytest.raises(h2.exceptions.InvalidBodyLengthError) as exp: + c.receive_data(trailers.serialize()) + + assert exp.value.expected_length == 15 + assert exp.value.actual_length == 0 + + def test_matching_body_ended_by_trailers_is_accepted(self, frame_factory) -> None: + """ + A trailers section that ends a stream whose body matches content-length + is still accepted, and emits TrailersReceived. + """ + c = self._server(frame_factory, self.example_request_headers) + c.receive_data(frame_factory.build_data_frame(data=b"\x01"*15).serialize()) + c.clear_outbound_data_buffer() + + trailers = frame_factory.build_headers_frame( + headers=[("x-checksum", "0")], + flags=["END_STREAM"], + ) + events = c.receive_data(trailers.serialize()) + + assert any(isinstance(e, h2.events.TrailersReceived) for e in events) + + def test_trailers_without_content_length_unchanged(self, frame_factory) -> None: + """ + A request with no content-length that ends with trailers is unaffected + by trailers-time validation. + """ + headers = [ + (":authority", "example.com"), + (":path", "/"), + (":scheme", "https"), + (":method", "POST"), + ] + c = self._server(frame_factory, headers) + c.receive_data(frame_factory.build_data_frame(data=b"\x01"*3).serialize()) + c.clear_outbound_data_buffer() + + trailers = frame_factory.build_headers_frame( + headers=[("x-checksum", "0")], + flags=["END_STREAM"], + ) + events = c.receive_data(trailers.serialize()) + + assert any(isinstance(e, h2.events.TrailersReceived) for e in events)