Skip to content
28 changes: 28 additions & 0 deletions test/_util.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,23 @@
from __future__ import annotations

import asyncio
from collections.abc import AsyncIterator, Iterator
from typing import TYPE_CHECKING

from connectrpc._compression import IdentityCompression
from connectrpc.compression.brotli import BrotliCompression
from connectrpc.compression.gzip import GzipCompression
from connectrpc.compression.zstd import ZstdCompression

from .connectrpc.example.haberdasher_connect import HaberdasherClientSync

if TYPE_CHECKING:
from types import MethodType

from connectrpc.compression import Compression

from .connectrpc.example.haberdasher_pb import Hat, Size


def resolve_compression(encoding: str) -> Compression:
match encoding:
Expand All @@ -24,3 +32,23 @@ def resolve_compression(encoding: str) -> Compression:
case _:
msg = f"unknown encoding '{encoding}'"
raise ValueError(msg)


async def call(method: MethodType, request: Size | list[Size]) -> Hat | list[Hat]:
"""Calls a client method, passing and returning streams as lists."""
if isinstance(method.__self__, HaberdasherClientSync):

def run() -> Hat | list[Hat]:
result = method(iter(request) if isinstance(request, list) else request)
return list(result) if isinstance(result, Iterator) else result

return await asyncio.to_thread(run)

async def stream(requests: list[Size]) -> AsyncIterator[Size]:
for r in requests:
yield r

result = method(stream(request) if isinstance(request, list) else request)
if isinstance(result, AsyncIterator):
return [r async for r in result]
return await result
221 changes: 73 additions & 148 deletions test/test_compression.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@
from connectrpc.protocol import ProtocolType
from connectrpc.request import Headers

from ._util import resolve_compression
from ._util import call, resolve_compression
from .connectrpc.example.haberdasher_connect import (
Haberdasher,
HaberdasherASGIApplication,
Expand All @@ -36,48 +36,65 @@
from .connectrpc.example.haberdasher_pb import Hat, Size

if TYPE_CHECKING:
from collections.abc import Iterable

from connectrpc._protocol import ServerProtocol
from connectrpc.compression import Compression


@pytest.mark.asyncio
@pytest.mark.parametrize(
("compressions", "encoding"),
[
pytest.param((), "identity", id="none"),
pytest.param(("gzip",), "gzip", id="gzip"),
pytest.param(("zstd",), "zstd", id="zstd"),
pytest.param(("br",), "br", id="br"),
pytest.param(("gzip", "br", "zstd"), "zstd", id="all"),
],
)
async def test_server_compressions_async(
compressions: tuple[str], encoding: str
) -> None:
class SimpleHaberdasher(Haberdasher):
async def make_hat(self, _request, _ctx):
return Hat(size=10, color="blue")

app = HaberdasherASGIApplication(
SimpleHaberdasher(), compressions=[resolve_compression(c) for c in compressions]
)
with ResponseMetadata() as meta:
client = HaberdasherClient(
class _BlueHaberdasher(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")


class _BlueHaberdasherSync(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")


@pytest.fixture(params=["async", "sync"])
def new_client(request: pytest.FixtureRequest):
"""Returns a factory for clients of a server with the given compressions."""

def new_client(
compressions: Iterable[Compression] | None = None,
*,
send_compression: Compression | None,
accept_compression: Iterable[Compression] | None = None,
protocol: ProtocolType = ProtocolType.CONNECT,
) -> HaberdasherClient | HaberdasherClientSync:
if request.param == "async":
app = HaberdasherASGIApplication(
_BlueHaberdasher(), compressions=compressions
)
return HaberdasherClient(
"http://localhost",
http_client=Client(ASGITransport(app)),
send_compression=send_compression,
accept_compression=accept_compression,
protocol=protocol,
)
app = HaberdasherWSGIApplication(
_BlueHaberdasherSync(), compressions=compressions
)
return HaberdasherClientSync(
"http://localhost",
http_client=Client(ASGITransport(app)),
accept_compression=(
ZstdCompression(),
GzipCompression(),
BrotliCompression(),
),
send_compression=None,
http_client=SyncClient(WSGITransport(app)),
send_compression=send_compression,
accept_compression=accept_compression,
protocol=protocol,
)
res = await client.make_hat(Size(inches=10))
assert res.size == 10
assert res.color == "blue"
assert meta.headers.get("content-encoding") == encoding

return new_client


@pytest.mark.asyncio
@pytest.mark.parametrize(
("compressions", "encoding"),
[
Expand All @@ -88,89 +105,37 @@ async def make_hat(self, _request, _ctx):
pytest.param(("gzip", "br", "zstd"), "zstd", id="all"),
],
)
def test_server_compressions_sync(compressions: tuple[str], encoding: str) -> None:
class SimpleHaberdasher(HaberdasherSync):
def make_hat(self, _request, _ctx):
return Hat(size=10, color="blue")

app = HaberdasherWSGIApplication(
SimpleHaberdasher(), compressions=[resolve_compression(c) for c in compressions]
)
client = HaberdasherClientSync(
"http://localhost",
http_client=SyncClient(WSGITransport(app)),
async def test_server_compressions(
new_client, compressions: tuple[str], encoding: str
) -> None:
client = new_client(
[resolve_compression(c) for c in compressions],
accept_compression=(ZstdCompression(), GzipCompression(), BrotliCompression()),
send_compression=None,
)
with ResponseMetadata() as meta:
res = client.make_hat(Size(inches=10))
assert res.size == 10
assert res.color == "blue"
res = await call(client.make_hat, Size(inches=10))
assert res == Hat(size=10, color="blue")
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")]
_methods = [
pytest.param("make_hat", id="unary"),
pytest.param("make_similar_hats", 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
@pytest.mark.parametrize("method", _methods)
async def test_unknown_request_compression(
new_client, protocol: ProtocolType, method: str
) -> 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(),
)
client = new_client(protocol=protocol, 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))
await call(getattr(client, method), Size(inches=10))
assert exc_info.value.code == Code.UNIMPLEMENTED
assert (
exc_info.value.message
Expand Down Expand Up @@ -216,51 +181,11 @@ def decompress(


@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))]
@pytest.mark.parametrize("method", _methods)
async def test_mixed_case_request_compression(new_client, method: str) -> None:
client = new_client([_XorCompression()], send_compression=_XorCompression())
res = await call(getattr(client, method), Size(inches=10))
hats = res if isinstance(res, list) else [res]
assert hats == [Hat(size=10, color="blue")]


Expand Down
Loading
Loading