diff --git a/src/connectrpc/_compression.py b/src/connectrpc/_compression.py index 6475d084..48475182 100644 --- a/src/connectrpc/_compression.py +++ b/src/connectrpc/_compression.py @@ -57,6 +57,17 @@ def negotiate_compression( return _identity +def resolve_request_compression( + name: str, compressions: dict[str, Compression] +) -> Compression | None: + """Return the compression for a request's encoding, or None if unsupported. + + Every request path resolves the name here: an empty name means identity, as + in connect-go, and names match exactly. + """ + return compressions.get(name or "identity") + + def unknown_compression_error( name: str, compressions: dict[str, Compression] ) -> ConnectError: diff --git a/src/connectrpc/_protocol_connect.py b/src/connectrpc/_protocol_connect.py index cb609152..7dd37fd4 100644 --- a/src/connectrpc/_protocol_connect.py +++ b/src/connectrpc/_protocol_connect.py @@ -6,7 +6,11 @@ from typing import TYPE_CHECKING, Any, TypeVar from ._codec import CODEC_NAME_JSON, Codec -from ._compression import IdentityCompression, negotiate_compression +from ._compression import ( + IdentityCompression, + negotiate_compression, + resolve_request_compression, +) from ._envelope import EnvelopeReader, EnvelopeWriter from ._protocol import ( ConnectWireError, @@ -144,11 +148,9 @@ 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]: - # An empty header means identity too, as in connect-go. - req_compression_name = ( - headers.get(CONNECT_STREAMING_HEADER_COMPRESSION) or "identity" + req_compression = resolve_request_compression( + headers.get(CONNECT_STREAMING_HEADER_COMPRESSION, ""), compressions ) - 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 7ac4140b..9df477e5 100644 --- a/src/connectrpc/_protocol_grpc.py +++ b/src/connectrpc/_protocol_grpc.py @@ -9,7 +9,11 @@ from pyqwest import Headers as HTTPHeaders -from ._compression import IdentityCompression, negotiate_compression +from ._compression import ( + IdentityCompression, + negotiate_compression, + resolve_request_compression, +) from ._envelope import EnvelopeReader, EnvelopeWriter from ._gen.google.rpc.status_pb import Status from ._protocol import ( @@ -96,8 +100,9 @@ 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) or "identity" - req_compression = compressions.get(req_compression_name) + req_compression = resolve_request_compression( + headers.get(GRPC_HEADER_COMPRESSION, ""), compressions + ) accept_compression = headers.get(GRPC_HEADER_ACCEPT_COMPRESSION, "") resp_compression = negotiate_compression(accept_compression, compressions) return req_compression, resp_compression diff --git a/src/connectrpc/_server_async.py b/src/connectrpc/_server_async.py index bef10199..87eb81e8 100644 --- a/src/connectrpc/_server_async.py +++ b/src/connectrpc/_server_async.py @@ -15,6 +15,7 @@ from ._compression import ( negotiate_compression, resolve_compressions, + resolve_request_compression, unknown_compression_error, ) from ._envelope import EnvelopeReader @@ -342,8 +343,8 @@ async def _read_get_request( message = message.encode("utf-8") # Handle compression - compression_name = params.get("compression", [""])[0] or "identity" - compression = self._compressions.get(compression_name) + compression_name = params.get("compression", [""])[0] + compression = resolve_request_compression(compression_name, self._compressions) if not compression: raise unknown_compression_error(compression_name, self._compressions) @@ -376,8 +377,8 @@ async def _read_post_request( req_body = b"".join(chunks) # Handle compression if specified - compression_name = headers.get("content-encoding") or "identity" - compression = self._compressions.get(compression_name) + compression_name = headers.get("content-encoding", "") + compression = resolve_request_compression(compression_name, self._compressions) if not compression: raise unknown_compression_error(compression_name, self._compressions) diff --git a/src/connectrpc/_server_sync.py b/src/connectrpc/_server_sync.py index 2f203ccc..6dcdcf45 100644 --- a/src/connectrpc/_server_sync.py +++ b/src/connectrpc/_server_sync.py @@ -14,6 +14,7 @@ from ._compression import ( negotiate_compression, resolve_compressions, + resolve_request_compression, unknown_compression_error, ) from ._envelope import EnvelopeReader, EnvelopeWriter @@ -406,8 +407,10 @@ def _handle_post_request( req_body = b"".join(chunks) # Handle compression if specified - compression_name = environ.get("HTTP_CONTENT_ENCODING") or "identity" - compression = self._compressions.get(compression_name) + compression_name = environ.get("HTTP_CONTENT_ENCODING", "") + compression = resolve_request_compression( + compression_name, self._compressions + ) if not compression: raise unknown_compression_error(compression_name, self._compressions) try: @@ -461,8 +464,10 @@ def _handle_get_request( message = message.encode("utf-8") # Handle compression if specified - compression_name = params.get("compression", [""])[0] or "identity" - compression = self._compressions.get(compression_name) + compression_name = params.get("compression", [""])[0] + compression = resolve_request_compression( + compression_name, self._compressions + ) if not compression: raise unknown_compression_error(compression_name, self._compressions) message = compression.decompress(message, self._read_max_bytes) diff --git a/test/test_compression.py b/test/test_compression.py index f3d35890..f66c6d96 100644 --- a/test/test_compression.py +++ b/test/test_compression.py @@ -11,6 +11,7 @@ from connectrpc._compression import ( IdentityCompression, resolve_compressions, + resolve_request_compression, unknown_compression_error, ) from connectrpc._protocol_connect import ConnectServerProtocol @@ -264,6 +265,21 @@ def make_similar_hats(self, request, _ctx): assert hats == [Hat(size=10, color="blue")] +@pytest.mark.parametrize( + ("name", "expected"), + [ + pytest.param("", "identity", id="empty"), + pytest.param("identity", "identity", id="identity"), + pytest.param("gzip", "gzip", id="gzip"), + pytest.param("GZIP", None, id="exact-match"), + pytest.param("zstd", None, id="unknown"), + ], +) +def test_resolve_request_compression(name: str, expected: str | None) -> None: + compression = resolve_request_compression(name, resolve_compressions(None)) + assert (compression.name() if compression else None) == expected + + @pytest.mark.parametrize( ("protocol", "header_name"), [