Skip to content
Draft
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
11 changes: 11 additions & 0 deletions src/connectrpc/_compression.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
12 changes: 7 additions & 5 deletions src/connectrpc/_protocol_connect.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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, ""
)
Expand Down
11 changes: 8 additions & 3 deletions src/connectrpc/_protocol_grpc.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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
Expand Down
9 changes: 5 additions & 4 deletions src/connectrpc/_server_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
from ._compression import (
negotiate_compression,
resolve_compressions,
resolve_request_compression,
unknown_compression_error,
)
from ._envelope import EnvelopeReader
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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)

Expand Down
13 changes: 9 additions & 4 deletions src/connectrpc/_server_sync.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from ._compression import (
negotiate_compression,
resolve_compressions,
resolve_request_compression,
unknown_compression_error,
)
from ._envelope import EnvelopeReader, EnvelopeWriter
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
16 changes: 16 additions & 0 deletions test/test_compression.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
from connectrpc._compression import (
IdentityCompression,
resolve_compressions,
resolve_request_compression,
unknown_compression_error,
)
from connectrpc._protocol_connect import ConnectServerProtocol
Expand Down Expand Up @@ -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"),
[
Expand Down
Loading