Skip to content
Open
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
78 changes: 74 additions & 4 deletions src/hypercorn/protocol/h3.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,15 @@

from collections.abc import Awaitable, Callable

from aioquic.h3.connection import H3Connection
from aioquic.h3.events import DataReceived, HeadersReceived
from aioquic.h3.connection import (
H3Connection,
HeadersState,
FrameType,
FrameUnexpected,
ProtocolError,
encode_frame
)
from aioquic.h3.events import DataReceived, HeadersReceived , Headers
from aioquic.h3.exceptions import NoAvailablePushIDError
from aioquic.quic.connection import QuicConnection
from aioquic.quic.events import QuicEvent
Expand All @@ -26,7 +33,6 @@
from ..typing import AppWrapper, ConnectionState, TaskGroup, WorkerContext
from ..utils import filter_pseudo_headers


class H3Protocol:
def __init__(
self,
Expand All @@ -44,7 +50,7 @@ def __init__(
self.client = client
self.config = config
self.context = context
self.connection = H3Connection(quic)
self.connection = H3Wrapper(quic)
self.send = send
self.server = server
self.streams: dict[int, HTTPStream | WSStream] = {}
Expand All @@ -69,11 +75,13 @@ async def handle(self, quic_event: QuicEvent) -> None:

async def stream_send(self, event: StreamEvent) -> None:
if isinstance(event, (InformationalResponse, Response)):
is_informational = isinstance(event,InformationalResponse)
self.connection.send_headers(
event.stream_id,
[(b":status", b"%d" % event.status_code)]
+ event.headers
+ self.config.response_headers("h3"),
is_informational=is_informational
)
await self.send()
elif isinstance(event, (Body, Data)):
Expand Down Expand Up @@ -154,3 +162,65 @@ async def _create_server_push(
)
await self._create_stream(event)
await self.streams[event.stream_id].handle(EndBody(stream_id=event.stream_id))


class H3Wrapper(H3Connection):
"""
High level wrapper for H3Connection
"""
def __init__(self, quic, enable_webtransport = False):
super().__init__(quic, enable_webtransport)

def send_headers(
self,
stream_id: int,
headers: Headers,
end_stream: bool = False ,
is_informational: bool = False
) -> None:
"""
Send headers on the given stream.

.. aioquic_transmit::

:param stream_id: The stream ID on which to send the headers.
:param headers: The HTTP headers to send.
:param end_stream: Whether to end the stream.
:param is_informational: Wheather headers contains informational response headers
"""
if is_informational and end_stream:
raise ProtocolError("Informational headers (1xx) cannot end the stream.")

# check HEADERS frame is allowed
with self._get_or_create_stream(stream_id) as stream:
if stream.headers_send_state == HeadersState.AFTER_TRAILERS:
raise FrameUnexpected("HEADERS frame is not allowed in this state")

# Cannot send informational headers after final headers were already sent
if (
is_informational
and stream.headers_send_state != HeadersState.INITIAL
):
raise FrameUnexpected("Informational headers cannot be sent after final headers.")

if end_stream:
stream.finish_sending()
frame_data = self._encode_headers(stream_id, headers)
# log frame
if self._quic_logger is not None:
self._quic_logger.log_event(
category="http",
event="frame_created",
data=self._quic_logger.encode_http3_headers_frame(
length=len(frame_data), headers=headers, stream_id=stream_id
),
)
# update state and send headers
if not is_informational:
if stream.headers_send_state == HeadersState.INITIAL:
stream.headers_send_state = HeadersState.AFTER_HEADERS
else:
stream.headers_send_state = HeadersState.AFTER_TRAILERS
self._quic.send_stream_data(
stream_id, encode_frame(FrameType.HEADERS, frame_data), end_stream
)
64 changes: 64 additions & 0 deletions tests/protocol/test_h3.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
from unittest.mock import Mock

import pytest
from aioquic.h3.connection import (
FrameUnexpected,
HeadersState,
ProtocolError,
)

from hypercorn.protocol.h3 import H3Wrapper


@pytest.mark.parametrize("interim_count", [1, 2])
def test_informational_headers_do_not_advance_send_state(
interim_count: int,
) -> None:
quic = Mock()
connection = H3Wrapper(quic)
stream_id = 0

for _ in range(interim_count):
connection.send_headers(
stream_id,
[(b":status", b"103")],
is_informational=True,
)
assert connection._stream[stream_id].headers_send_state == HeadersState.INITIAL

connection.send_headers(stream_id, [(b":status", b"200")])
assert connection._stream[stream_id].headers_send_state == HeadersState.AFTER_HEADERS

# DATA and end-of-stream remain valid after the final response headers.
connection.send_data(stream_id, b"response body", end_stream=True)


def test_informational_headers_cannot_end_stream() -> None:
connection = H3Wrapper(Mock())

with pytest.raises((ProtocolError,FrameUnexpected), match="Informational headers.*cannot end the stream"):
connection.send_headers(
0,
[(b":status", b"103")],
end_stream=True,
is_informational=True,
)


def test_informational_headers_cannot_follow_final_headers() -> None:
connection = H3Wrapper(Mock())
stream_id = 0

connection.send_headers(stream_id, [(b":status", b"200")])

with pytest.raises(
(FrameUnexpected,ProtocolError),
match="Informational headers cannot be sent after final headers",
):
connection.send_headers(
stream_id,
[(b":status", b"103")],
is_informational=True,
)

assert connection._stream[stream_id].headers_send_state == HeadersState.AFTER_HEADERS