diff --git a/src/h2/stream.py b/src/h2/stream.py index 249f73e0c..9b9cd544c 100644 --- a/src/h2/stream.py +++ b/src/h2/stream.py @@ -906,7 +906,9 @@ def send_headers(self, self._authority = authority_from_headers(bytes_headers) # store request method for _initialize_content_length - self.request_method = extract_method_header(bytes_headers) + method = extract_method_header(bytes_headers) + if method is not None: + self.request_method = method return frames diff --git a/tests/test_head_request.py b/tests/test_head_request.py index e32d1525f..5e7bd9046 100644 --- a/tests/test_head_request.py +++ b/tests/test_head_request.py @@ -45,6 +45,35 @@ def test_non_zero_content_and_no_body(self, frame_factory, headers) -> None: assert event.stream_id == 1 assert event.headers == self.example_response_headers + @pytest.mark.parametrize("headers", [EXAMPLE_REQUEST_HEADERS, EXAMPLE_REQUEST_HEADERS_BYTES]) + def test_non_zero_content_and_no_body_after_trailers(self, frame_factory, headers) -> None: + c = h2.connection.H2Connection() + c.initiate_connection() + c.send_headers(1, headers) + c.send_headers(1, [(b"x-checksum", b"abc123")], end_stream=True) + + f = frame_factory.build_headers_frame( + [ + (b":status", b"200"), + (b"server", b"fake-serv/0.1.0"), + (b"content-length", b"1234"), + ], + ) + events = c.receive_data(f.serialize()) + + assert len(events) == 1 + event = events[0] + + assert isinstance(event, h2.events.ResponseReceived) + assert event.stream_id == 1 + + data = frame_factory.build_data_frame(b"", flags=["END_STREAM"]) + events = c.receive_data(data.serialize()) + + assert len(events) == 2 + assert isinstance(events[0], h2.events.DataReceived) + assert isinstance(events[1], h2.events.StreamEnded) + @pytest.mark.parametrize("headers", [EXAMPLE_REQUEST_HEADERS, EXAMPLE_REQUEST_HEADERS_BYTES]) def test_reject_non_zero_content_and_body(self, frame_factory, headers) -> None: c = h2.connection.H2Connection()