diff --git a/src/connectrpc/_compression.py b/src/connectrpc/_compression.py index f139fa50..6475d084 100644 --- a/src/connectrpc/_compression.py +++ b/src/connectrpc/_compression.py @@ -5,7 +5,9 @@ from connectrpc.compression.gzip import GzipCompression from ._shared import message_too_large_error +from .code import Code from .compression import Compression +from .errors import ConnectError if TYPE_CHECKING: from collections.abc import Iterable @@ -53,3 +55,16 @@ def negotiate_compression( if compression: return compression return _identity + + +def unknown_compression_error( + name: str, compressions: dict[str, Compression] +) -> ConnectError: + # identity is always accepted and is not a compression, so it is not listed. + supported = [n for n in compressions if n != "identity"] + detail = ( + f"supported encodings are {', '.join(supported)}" + if supported + else "compression is not supported" + ) + return ConnectError(Code.UNIMPLEMENTED, f"unknown compression: '{name}': {detail}") diff --git a/src/connectrpc/_protocol_connect.py b/src/connectrpc/_protocol_connect.py index 06e607c4..cb609152 100644 --- a/src/connectrpc/_protocol_connect.py +++ b/src/connectrpc/_protocol_connect.py @@ -143,13 +143,12 @@ def codec_name_from_content_type(self, content_type: str, *, stream: bool) -> st def negotiate_stream_compression( self, headers: Headers, compressions: dict[str, Compression] - ) -> tuple[Compression, Compression]: - req_compression_name = headers.get( - CONNECT_STREAMING_HEADER_COMPRESSION, "identity" - ) - req_compression = ( - compressions.get(req_compression_name) or IdentityCompression() + ) -> tuple[Compression | None, Compression]: + # An empty header means identity too, as in connect-go. + req_compression_name = ( + headers.get(CONNECT_STREAMING_HEADER_COMPRESSION) or "identity" ) + req_compression = compressions.get(req_compression_name) accept_compression = headers.get( CONNECT_STREAMING_HEADER_ACCEPT_COMPRESSION, "" ) diff --git a/src/connectrpc/_protocol_grpc.py b/src/connectrpc/_protocol_grpc.py index 5a01fa5b..7ac4140b 100644 --- a/src/connectrpc/_protocol_grpc.py +++ b/src/connectrpc/_protocol_grpc.py @@ -96,7 +96,7 @@ def codec_name_from_content_type(self, content_type: str, *, stream: bool) -> st def negotiate_stream_compression( self, headers: Headers, compressions: dict[str, Compression] ) -> tuple[Compression | None, Compression]: - req_compression_name = headers.get(GRPC_HEADER_COMPRESSION, "identity") + req_compression_name = headers.get(GRPC_HEADER_COMPRESSION) or "identity" req_compression = compressions.get(req_compression_name) accept_compression = headers.get(GRPC_HEADER_ACCEPT_COMPRESSION, "") resp_compression = negotiate_compression(accept_compression, compressions) diff --git a/src/connectrpc/_server_async.py b/src/connectrpc/_server_async.py index fa202ea7..bef10199 100644 --- a/src/connectrpc/_server_async.py +++ b/src/connectrpc/_server_async.py @@ -12,7 +12,11 @@ from urllib.parse import parse_qs from ._codec import Codec, get_default_codecs -from ._compression import negotiate_compression, resolve_compressions +from ._compression import ( + negotiate_compression, + resolve_compressions, + unknown_compression_error, +) from ._envelope import EnvelopeReader from ._interceptor_async import ( BidiStreamInterceptor, @@ -338,13 +342,10 @@ async def _read_get_request( message = message.encode("utf-8") # Handle compression - compression_name = params.get("compression", ["identity"])[0] + compression_name = params.get("compression", [""])[0] or "identity" compression = self._compressions.get(compression_name) if not compression: - raise ConnectError( - Code.UNIMPLEMENTED, - f"unknown compression: '{compression_name}': supported encodings are {', '.join(self._compressions.keys())}", - ) + raise unknown_compression_error(compression_name, self._compressions) # Decompress and decode message if message: # Don't decompress empty messages @@ -375,13 +376,10 @@ async def _read_post_request( req_body = b"".join(chunks) # Handle compression if specified - compression_name = headers.get("content-encoding", "identity").lower() + compression_name = headers.get("content-encoding") or "identity" compression = self._compressions.get(compression_name) if not compression: - raise ConnectError( - Code.UNIMPLEMENTED, - f"unknown compression: '{compression_name}': supported encodings are {', '.join(self._compressions.keys())}", - ) + raise unknown_compression_error(compression_name, self._compressions) if req_body: # Don't decompress empty body req_body = compression.decompress(req_body, self._read_max_bytes) @@ -410,8 +408,9 @@ async def _handle_stream( try: await metadata_run.start() if not req_compression: - raise ConnectError( - Code.UNIMPLEMENTED, "Unrecognized request compression" + raise unknown_compression_error( + headers.get(protocol.compression_header_name(), ""), + self._compressions, ) request_stream = _request_stream( receive, diff --git a/src/connectrpc/_server_sync.py b/src/connectrpc/_server_sync.py index 8698c014..2f203ccc 100644 --- a/src/connectrpc/_server_sync.py +++ b/src/connectrpc/_server_sync.py @@ -11,7 +11,11 @@ from . import _server_shared from ._codec import Codec, get_default_codecs -from ._compression import negotiate_compression, resolve_compressions +from ._compression import ( + negotiate_compression, + resolve_compressions, + unknown_compression_error, +) from ._envelope import EnvelopeReader, EnvelopeWriter from ._interceptor_sync import ( BidiStreamInterceptorSync, @@ -402,13 +406,10 @@ def _handle_post_request( req_body = b"".join(chunks) # Handle compression if specified - compression_name = environ.get("HTTP_CONTENT_ENCODING", "identity").lower() + compression_name = environ.get("HTTP_CONTENT_ENCODING") or "identity" compression = self._compressions.get(compression_name) if not compression: - raise ConnectError( - Code.UNIMPLEMENTED, - f"unknown compression: '{compression_name}': supported encodings are {', '.join(self._compressions.keys())}", - ) + raise unknown_compression_error(compression_name, self._compressions) try: req_body = compression.decompress(req_body, self._read_max_bytes) except ConnectError: @@ -460,16 +461,10 @@ def _handle_get_request( message = message.encode("utf-8") # Handle compression if specified - if "compression" in params: - compression_name = params["compression"][0] - else: - compression_name = "identity" + compression_name = params.get("compression", [""])[0] or "identity" compression = self._compressions.get(compression_name) if not compression: - raise ConnectError( - Code.UNIMPLEMENTED, - f"unknown compression: '{compression_name}': supported encodings are {', '.join(self._compressions.keys())}", - ) + raise unknown_compression_error(compression_name, self._compressions) message = compression.decompress(message, self._read_max_bytes) codec_name = params.get("encoding", ("",))[0] @@ -528,8 +523,9 @@ def _handle_stream( try: metadata_run.start() if not req_compression: - raise ConnectError( - Code.UNIMPLEMENTED, "Unrecognized request compression" + raise unknown_compression_error( + headers.get(protocol.compression_header_name(), ""), + self._compressions, ) request_stream = _request_stream( request_body, diff --git a/test/test_compression.py b/test/test_compression.py index eab5bccf..f3d35890 100644 --- a/test/test_compression.py +++ b/test/test_compression.py @@ -1,19 +1,28 @@ from __future__ import annotations from typing import TYPE_CHECKING +from urllib.parse import urlencode import brotli as brotli_lib import pytest from pyqwest import Client, SyncClient from pyqwest.testing import ASGITransport, WSGITransport -from connectrpc._compression import IdentityCompression +from connectrpc._compression import ( + IdentityCompression, + resolve_compressions, + unknown_compression_error, +) +from connectrpc._protocol_connect import ConnectServerProtocol +from connectrpc._protocol_grpc import GRPCServerProtocol from connectrpc.client import ResponseMetadata from connectrpc.code import Code from connectrpc.compression.brotli import BrotliCompression from connectrpc.compression.gzip import GzipCompression from connectrpc.compression.zstd import ZstdCompression from connectrpc.errors import ConnectError +from connectrpc.protocol import ProtocolType +from connectrpc.request import Headers from ._util import resolve_compression from .connectrpc.example.haberdasher_connect import ( @@ -27,6 +36,7 @@ from .connectrpc.example.haberdasher_pb import Hat, Size if TYPE_CHECKING: + from connectrpc._protocol import ServerProtocol from connectrpc.compression import Compression @@ -99,6 +109,251 @@ def make_hat(self, _request, _ctx): assert meta.headers.get("content-encoding") == encoding +_protocols = [ProtocolType.CONNECT, ProtocolType.GRPC, ProtocolType.GRPC_WEB] +_streams = [pytest.param(False, id="unary"), pytest.param(True, id="stream")] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("protocol", _protocols) +@pytest.mark.parametrize("stream", _streams) +async def test_unknown_request_compression_async( + protocol: ProtocolType, stream: bool +) -> None: + class SimpleHaberdasher(Haberdasher): + async def make_hat(self, _request, _ctx): + return Hat(size=10, color="blue") + + async def make_similar_hats(self, _request, _ctx): + yield Hat(size=10, color="blue") + + # The server only supports the default gzip. + app = HaberdasherASGIApplication(SimpleHaberdasher()) + client = HaberdasherClient( + "http://localhost", + protocol=protocol, + http_client=Client(ASGITransport(app)), + send_compression=ZstdCompression(), + ) + with pytest.raises(ConnectError) as exc_info: + if stream: + async for _ in client.make_similar_hats(Size(inches=10)): + pass + else: + await client.make_hat(Size(inches=10)) + assert exc_info.value.code == Code.UNIMPLEMENTED + assert ( + exc_info.value.message + == "unknown compression: 'zstd': supported encodings are gzip" + ) + + +@pytest.mark.parametrize("protocol", _protocols) +@pytest.mark.parametrize("stream", _streams) +def test_unknown_request_compression_sync(protocol: ProtocolType, stream: bool) -> None: + class SimpleHaberdasher(HaberdasherSync): + def make_hat(self, _request, _ctx): + return Hat(size=10, color="blue") + + def make_similar_hats(self, _request, _ctx): + yield Hat(size=10, color="blue") + + # The server only supports the default gzip. + app = HaberdasherWSGIApplication(SimpleHaberdasher()) + client = HaberdasherClientSync( + "http://localhost", + protocol=protocol, + http_client=SyncClient(WSGITransport(app)), + send_compression=ZstdCompression(), + ) + with pytest.raises(ConnectError) as exc_info: + if stream: + for _ in client.make_similar_hats(Size(inches=10)): + pass + else: + client.make_hat(Size(inches=10)) + assert exc_info.value.code == Code.UNIMPLEMENTED + assert ( + exc_info.value.message + == "unknown compression: 'zstd': supported encodings are gzip" + ) + + +@pytest.mark.parametrize( + ("compressions", "detail"), + [ + pytest.param(None, "supported encodings are gzip", id="default"), + pytest.param((), "compression is not supported", id="none"), + pytest.param( + (ZstdCompression(), GzipCompression()), + "supported encodings are zstd, gzip", + id="multiple", + ), + ], +) +def test_unknown_compression_error( + compressions: tuple[Compression, ...] | None, detail: str +) -> None: + error = unknown_compression_error("foo", resolve_compressions(compressions)) + assert error.code == Code.UNIMPLEMENTED + assert error.message == f"unknown compression: 'foo': {detail}" + + +class _XorCompression: + """A toy compression registered under a name that is not all lowercase.""" + + def name(self) -> str: + return "Xor" + + def compress(self, data: bytes | bytearray | memoryview) -> bytes: + return bytes(b ^ 0x5A for b in data) + + def decompress( + self, + data: bytes | bytearray | memoryview, + read_max_bytes: int | None = None, # noqa: ARG002 + ) -> bytes: + return bytes(b ^ 0x5A for b in data) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", _streams) +async def test_mixed_case_request_compression_async(stream: bool) -> None: + class SimpleHaberdasher(Haberdasher): + async def make_hat(self, request, _ctx): + return Hat(size=request.inches, color="blue") + + async def make_similar_hats(self, request, _ctx): + yield Hat(size=request.inches, color="blue") + + app = HaberdasherASGIApplication( + SimpleHaberdasher(), compressions=[_XorCompression()] + ) + client = HaberdasherClient( + "http://localhost", + http_client=Client(ASGITransport(app)), + send_compression=_XorCompression(), + ) + if stream: + hats = [hat async for hat in client.make_similar_hats(Size(inches=10))] + else: + hats = [await client.make_hat(Size(inches=10))] + assert hats == [Hat(size=10, color="blue")] + + +@pytest.mark.parametrize("stream", _streams) +def test_mixed_case_request_compression_sync(stream: bool) -> None: + class SimpleHaberdasher(HaberdasherSync): + def make_hat(self, request, _ctx): + return Hat(size=request.inches, color="blue") + + def make_similar_hats(self, request, _ctx): + yield Hat(size=request.inches, color="blue") + + app = HaberdasherWSGIApplication( + SimpleHaberdasher(), compressions=[_XorCompression()] + ) + client = HaberdasherClientSync( + "http://localhost", + http_client=SyncClient(WSGITransport(app)), + send_compression=_XorCompression(), + ) + if stream: + hats = list(client.make_similar_hats(Size(inches=10))) + else: + hats = [client.make_hat(Size(inches=10))] + assert hats == [Hat(size=10, color="blue")] + + +@pytest.mark.parametrize( + ("protocol", "header_name"), + [ + pytest.param(ConnectServerProtocol(), "connect-content-encoding", id="connect"), + pytest.param(GRPCServerProtocol(), "grpc-encoding", id="grpc"), + ], +) +@pytest.mark.parametrize( + ("header", "expected"), + [ + pytest.param(None, "identity", id="absent"), + pytest.param("", "identity", id="empty"), + pytest.param("identity", "identity", id="identity"), + pytest.param("gzip", "gzip", id="gzip"), + pytest.param("zstd", None, id="unknown"), + ], +) +def test_stream_request_compression( + protocol: ServerProtocol, header_name: str, header: str | None, expected: str | None +) -> None: + headers = Headers() + if header is not None: + headers[header_name] = header + compression, _ = protocol.negotiate_stream_compression( + headers, resolve_compressions(None) + ) + assert (compression.name() if compression else None) == expected + + +# An empty compression means identity, as in connect-go. +_empty_compression_requests = [ + pytest.param( + "POST", + "", + {"content-type": "application/json", "content-encoding": ""}, + b'{"inches": 10}', + id="post", + ), + pytest.param( + "GET", + "?" + + urlencode( + {"encoding": "json", "compression": "", "message": '{"inches": 10}'} + ), + {}, + b"", + id="get", + ), +] + + +@pytest.mark.parametrize( + ("method", "query", "headers", "body"), _empty_compression_requests +) +def test_empty_request_compression_sync(method, query, headers, body) -> None: + class SimpleHaberdasher(HaberdasherSync): + def make_hat(self, request, _ctx): + return Hat(size=request.inches, color="blue") + + transport = WSGITransport(HaberdasherWSGIApplication(SimpleHaberdasher())) + res = SyncClient(transport).execute( + method=method, + url=f"http://localhost/connectrpc.example.Haberdasher/MakeHat{query}", + headers=headers, + content=body, + ) + assert res.status == 200 + assert res.json() == {"size": 10, "color": "blue"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("method", "query", "headers", "body"), _empty_compression_requests +) +async def test_empty_request_compression_async(method, query, headers, body) -> None: + class SimpleHaberdasher(Haberdasher): + async def make_hat(self, request, _ctx): + return Hat(size=request.inches, color="blue") + + transport = ASGITransport(HaberdasherASGIApplication(SimpleHaberdasher())) + res = await Client(transport).execute( + method=method, + url=f"http://localhost/connectrpc.example.Haberdasher/MakeHat{query}", + headers=headers, + content=body, + ) + assert res.status == 200 + assert res.json() == {"size": 10, "color": "blue"} + + class TestIdentityCompression: def test_name(self): assert IdentityCompression().name() == "identity"