diff --git a/h11/_connection.py b/h11/_connection.py index e37d82a..215a794 100644 --- a/h11/_connection.py +++ b/h11/_connection.py @@ -10,9 +10,13 @@ overload, Tuple, Type, + TYPE_CHECKING, Union, ) +if TYPE_CHECKING: + from typing_extensions import Buffer + from ._events import ( ConnectionClosed, Data, @@ -361,7 +365,7 @@ def trailing_data(self) -> Tuple[bytes, bool]: """ return (bytes(self._receive_buffer), self._receive_buffer_closed) - def receive_data(self, data: bytes) -> None: + def receive_data(self, data: "Buffer") -> None: """Add data to our internal receive buffer. This does not actually do any processing on the data, just stores @@ -501,18 +505,15 @@ def next_event(self) -> Union[Event, Type[NEED_DATA], Type[PAUSED]]: raise @overload - def send(self, event: ConnectionClosed) -> None: - ... + def send(self, event: ConnectionClosed) -> None: ... @overload def send( self, event: Union[Request, InformationalResponse, Response, Data, EndOfMessage] - ) -> bytes: - ... + ) -> bytes: ... @overload - def send(self, event: Event) -> Optional[bytes]: - ... + def send(self, event: Event) -> Optional[bytes]: ... def send(self, event: Event) -> Optional[bytes]: """Convert a high-level event into bytes that can be sent to the peer, diff --git a/h11/_receivebuffer.py b/h11/_receivebuffer.py index e5c4e08..3bfe0b1 100644 --- a/h11/_receivebuffer.py +++ b/h11/_receivebuffer.py @@ -1,6 +1,9 @@ import re import sys -from typing import List, Optional, Union +from typing import List, Optional, TYPE_CHECKING + +if TYPE_CHECKING: + from typing_extensions import Buffer __all__ = ["ReceiveBuffer"] @@ -50,7 +53,7 @@ def __init__(self) -> None: self._next_line_search = 0 self._multiple_lines_search = 0 - def __iadd__(self, byteslike: Union[bytes, bytearray]) -> "ReceiveBuffer": + def __iadd__(self, byteslike: "Buffer") -> "ReceiveBuffer": self._data += byteslike return self diff --git a/h11/tests/test_connection.py b/h11/tests/test_connection.py index 01260dc..1899bc0 100644 --- a/h11/tests/test_connection.py +++ b/h11/tests/test_connection.py @@ -1,4 +1,4 @@ -from typing import Any, cast, Dict, List, Optional, Tuple, Type +from typing import Any, Callable, cast, Dict, List, Optional, Tuple, Type, Union import pytest @@ -238,6 +238,22 @@ def test_chunk_boundaries() -> None: assert conn.next_event() == EndOfMessage() +@pytest.mark.parametrize("data_wrapper", [bytearray, memoryview]) +def test_receive_data_accepts_byteslike_objects( + data_wrapper: Callable[[bytes], Union[bytearray, memoryview]], +) -> None: + conn = Connection(our_role=SERVER) + + conn.receive_data(data_wrapper(b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n")) + + assert conn.next_event() == Request( + method="GET", + target="/", + headers=[("Host", "example.com")], + ) + assert conn.next_event() == EndOfMessage() + + def test_client_talking_to_http10_server() -> None: c = Connection(CLIENT) c.send(Request(method="GET", target="/", headers=[("Host", "example.com")])) diff --git a/newsfragments/186.feature.rst b/newsfragments/186.feature.rst new file mode 100644 index 0000000..a3e4750 --- /dev/null +++ b/newsfragments/186.feature.rst @@ -0,0 +1,2 @@ +Allow ``Connection.receive_data()`` to accept ``bytearray`` and +``memoryview`` inputs in its public type hints.