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
19 changes: 14 additions & 5 deletions src/hypercorn/protocol/h2.py
Original file line number Diff line number Diff line change
Expand Up @@ -126,6 +126,10 @@ def __init__(
def idle(self) -> bool:
return len(self.streams) == 0 or all(stream.idle for stream in self.streams.values())

@property
def _max_requests_reached(self) -> bool:
return self.keep_alive_requests >= self.config.keep_alive_max_requests

async def initiate(
self, headers: list[tuple[bytes, bytes]] | None = None, settings: bytes | None = None
) -> None:
Expand Down Expand Up @@ -225,10 +229,12 @@ async def stream_send(self, event: StreamEvent) -> None:
idle = len(self.streams) == 0 or all(
stream.idle for stream in self.streams.values()
)
if idle and self.context.terminated.is_set():
if idle and (self.context.terminated.is_set() or self._max_requests_reached):
self.connection.close_connection()
await self._flush()
await self.send(Updated(idle=idle))
await self.send(Closed())
else:
await self.send(Updated(idle=idle))
elif isinstance(event, Request):
await self._create_server_push(event.stream_id, event.raw_path, event.headers)
except (
Expand All @@ -252,9 +258,12 @@ async def _handle_events(self, events: list[h2.events.Event]) -> None:
else:
await self._create_stream(event)
await self.send(Updated(idle=False))

if self.keep_alive_requests > self.config.keep_alive_max_requests:
self.connection.close_connection()
if self._max_requests_reached:
# Stop the client opening further streams; the GOAWAY follows
# once the streams in flight have completed, see stream_send.
self.connection.update_settings(
{h2.settings.SettingCodes.MAX_CONCURRENT_STREAMS: 0}
)
elif isinstance(event, h2.events.DataReceived):
await self.streams[event.stream_id].handle(
Body(stream_id=event.stream_id, data=event.data)
Expand Down
38 changes: 33 additions & 5 deletions tests/protocol/test_h2.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,11 +5,13 @@

import pytest
from h2.connection import H2Connection
from h2.events import ConnectionTerminated
from h2.events import ConnectionTerminated, RemoteSettingsChanged, StreamReset
from h2.settings import SettingCodes

from hypercorn.asyncio.worker_context import EventWrapper, WorkerContext
from hypercorn.config import Config
from hypercorn.events import Closed, RawData
from hypercorn.protocol.events import StreamClosed
from hypercorn.protocol.h2 import BUFFER_HIGH_WATER, BufferCompleteError, H2Protocol, StreamBuffer
from hypercorn.typing import ConnectionState

Expand Down Expand Up @@ -102,7 +104,7 @@ async def test_protocol_keep_alive_max_requests() -> None:
None,
AsyncMock(),
)
protocol.config.keep_alive_max_requests = 0
protocol.config.keep_alive_max_requests = 2
client = H2Connection()
client.initiate_connection()
headers = [
Expand All @@ -111,8 +113,34 @@ async def test_protocol_keep_alive_max_requests() -> None:
(":authority", "hypercorn"),
(":scheme", "https"),
]
client.send_headers(1, headers, end_stream=True)
# Three requests in flight when the limit is reached, the third one sent before
# the client could learn that the connection is winding down.
for stream_id in (1, 3, 5):
client.send_headers(stream_id, headers, end_stream=True)
await protocol.handle(RawData(data=client.data_to_send()))
protocol.send.assert_awaited() # type: ignore
events = client.receive_data(protocol.send.call_args_list[1].args[0].data) # type: ignore
assert isinstance(events[-1], ConnectionTerminated)
events = _received_events(client, protocol)
settings = [event for event in events if isinstance(event, RemoteSettingsChanged)]
assert settings[-1].changed_settings[SettingCodes.MAX_CONCURRENT_STREAMS].new_value == 0
# Nothing is closed while requests are in flight...
assert not any(isinstance(event, (ConnectionTerminated, StreamReset)) for event in events)
assert set(protocol.streams) == {1, 3, 5}

for stream_id in (1, 3):
await protocol.stream_send(StreamClosed(stream_id=stream_id))
assert not any(
isinstance(event, ConnectionTerminated) for event in _received_events(client, protocol)
)
# ...the GOAWAY follows once the last one has completed.
await protocol.stream_send(StreamClosed(stream_id=5))
assert protocol.send.call_args_list[-1] == call(Closed()) # type: ignore
assert isinstance(_received_events(client, protocol)[-1], ConnectionTerminated)


def _received_events(client: H2Connection, protocol: H2Protocol) -> list:
events = []
for call_ in protocol.send.call_args_list: # type: ignore
if isinstance(call_.args[0], RawData):
events.extend(client.receive_data(call_.args[0].data))
protocol.send.reset_mock() # type: ignore
return events