Skip to content
Merged
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
15 changes: 15 additions & 0 deletions src/connectrpc/_compression.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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}")
11 changes: 5 additions & 6 deletions src/connectrpc/_protocol_connect.py
Original file line number Diff line number Diff line change
Expand Up @@ -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, ""
)
Expand Down
2 changes: 1 addition & 1 deletion src/connectrpc/_protocol_grpc.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
25 changes: 12 additions & 13 deletions src/connectrpc/_server_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand Down
28 changes: 12 additions & 16 deletions src/connectrpc/_server_sync.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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,
Expand Down
Loading
Loading