From a02f71abc1b0b62b554c21ea9f83292ca3b10232 Mon Sep 17 00:00:00 2001 From: i2y <6240399+i2y@users.noreply.github.com> Date: Sun, 27 Sep 2026 11:34:50 +0900 Subject: [PATCH] Resolve request compression names in one place Each server path looked up the request compression on its own, and the copies drifted apart: #365 fixed an unknown Connect stream compression being treated as identity, Connect unary POST lowercasing the name, and empty encodings being accepted on some paths only. Resolve the name in resolve_request_compression instead, so every path treats an empty name as identity and matches names exactly. No behavior change. Co-Authored-By: Claude Signed-off-by: i2y <6240399+i2y@users.noreply.github.com> --- src/connectrpc/_compression.py | 11 +++++++++++ src/connectrpc/_protocol_connect.py | 12 +++++++----- src/connectrpc/_protocol_grpc.py | 11 ++++++++--- src/connectrpc/_server_async.py | 9 +++++---- src/connectrpc/_server_sync.py | 13 +++++++++---- test/test_compression.py | 16 ++++++++++++++++ 6 files changed, 56 insertions(+), 16 deletions(-) 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"), [