diff --git a/test/_util.py b/test/_util.py index db24edf..a5575be 100644 --- a/test/_util.py +++ b/test/_util.py @@ -1,5 +1,7 @@ from __future__ import annotations +import asyncio +from collections.abc import AsyncIterator, Iterator from typing import TYPE_CHECKING from connectrpc._compression import IdentityCompression @@ -7,9 +9,15 @@ 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: @@ -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 diff --git a/test/test_client.py b/test/test_client.py index ee6c47e..5f08248 100644 --- a/test/test_client.py +++ b/test/test_client.py @@ -18,6 +18,7 @@ from connectrpc.client import ResponseMetadata from connectrpc.protocol import ProtocolType +from ._util import call from .connectrpc.example.haberdasher_connect import ( Haberdasher, HaberdasherASGIApplication, @@ -61,77 +62,52 @@ ] -@pytest.mark.parametrize( - ("headers", "trailers", "response_headers", "response_trailers"), _headers_cases -) -def test_headers_sync(headers, trailers, response_headers, response_trailers) -> None: - class HeadersHaberdasherSync(HaberdasherSync): - def __init__( - self, headers: list[tuple[str, str]], trailers: list[tuple[str, str]] - ) -> None: - self.headers = headers - self.trailers = trailers - - def make_hat(self, _request, ctx): - for key, value in self.headers: - ctx.response_headers.add(key, value) - for key, value in self.trailers: - ctx.response_trailers.add(key, value) - return Hat() - - transport = WSGITransport( - HaberdasherWSGIApplication(HeadersHaberdasherSync(headers, trailers)) - ) - - client = HaberdasherClientSync( - "http://localhost", http_client=SyncClient(transport=transport) - ) - - with ResponseMetadata() as resp: - assert resp.http_status is None - assert list(resp.headers.allitems()) == [] - assert list(resp.trailers.allitems()) == [] - client.make_hat(Size(inches=10)) - - assert resp.http_status == 200 - assert list(resp.headers.allitems()) == response_headers - assert list(resp.trailers.allitems()) == response_trailers - - @pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["async", "sync"]) @pytest.mark.parametrize( ("headers", "trailers", "response_headers", "response_trailers"), _headers_cases ) -async def test_headers_async( - headers, trailers, response_headers, response_trailers +async def test_headers( + mode, headers, trailers, response_headers, response_trailers ) -> None: - class HeadersHaberdasher(Haberdasher): - def __init__( - self, headers: list[tuple[str, str]], trailers: list[tuple[str, str]] - ) -> None: - self.headers = headers - self.trailers = trailers - - async def make_hat(self, _request, ctx): - for key, value in self.headers: - ctx.response_headers.add(key, value) - for key, value in self.trailers: - ctx.response_trailers.add(key, value) - return Hat() - - transport = ASGITransport( - HaberdasherASGIApplication(HeadersHaberdasher(headers, trailers)) - ) - - client = HaberdasherClient( - "http://localhost", http_client=Client(transport=transport) - ) + if mode == "async": + + class HeadersHaberdasher(Haberdasher): + async def make_hat(self, _request, ctx): + for key, value in headers: + ctx.response_headers.add(key, value) + for key, value in trailers: + ctx.response_trailers.add(key, value) + return Hat() + + client = HaberdasherClient( + "http://localhost", + http_client=Client( + ASGITransport(HaberdasherASGIApplication(HeadersHaberdasher())) + ), + ) + else: + + class HeadersHaberdasherSync(HaberdasherSync): + def make_hat(self, _request, ctx): + for key, value in headers: + ctx.response_headers.add(key, value) + for key, value in trailers: + ctx.response_trailers.add(key, value) + return Hat() + + client = HaberdasherClientSync( + "http://localhost", + http_client=SyncClient( + WSGITransport(HaberdasherWSGIApplication(HeadersHaberdasherSync())) + ), + ) with ResponseMetadata() as resp: assert resp.http_status is None assert list(resp.headers.allitems()) == [] assert list(resp.trailers.allitems()) == [] - await client.make_hat(Size(inches=10)) + await call(client.make_hat, Size(inches=10)) assert resp.http_status == 200 assert list(resp.headers.allitems()) == response_headers diff --git a/test/test_compression.py b/test/test_compression.py index e7c0021..cb71755 100644 --- a/test/test_compression.py +++ b/test/test_compression.py @@ -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, @@ -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"), [ @@ -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 @@ -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")] diff --git a/test/test_interceptor.py b/test/test_interceptor.py index 97ce54f..43b55d3 100644 --- a/test/test_interceptor.py +++ b/test/test_interceptor.py @@ -13,6 +13,7 @@ from connectrpc.code import Code from connectrpc.errors import ConnectError +from ._util import call from .connectrpc.example.haberdasher_connect import ( Haberdasher, HaberdasherASGIApplication, @@ -63,347 +64,178 @@ def server_interceptor(): return RequestInterceptor() -@pytest_asyncio.fixture -async def client_async( - client_interceptor: RequestInterceptor, server_interceptor: RequestInterceptor -): - class SimpleHaberdasher(Haberdasher): - async def make_hat(self, request, _ctx): - if request.inches < 0: +class SimpleHaberdasher(Haberdasher): + async def make_hat(self, request, _ctx): + if request.inches < 0: + raise ConnectError(Code.INVALID_ARGUMENT, "Size must be non-negative") + return Hat(size=request.inches, color="green") + + async def make_flexible_hat(self, request, _ctx): + size = 0 + async for s in request: + if s.inches < 0: raise ConnectError(Code.INVALID_ARGUMENT, "Size must be non-negative") - return Hat(size=request.inches, color="green") - - async def make_flexible_hat(self, request, _ctx): - size = 0 - async for s in request: - if s.inches < 0: - raise ConnectError( - Code.INVALID_ARGUMENT, "Size must be non-negative" - ) - size += s.inches - return Hat(size=size, color="red") + size += s.inches + return Hat(size=size, color="red") - async def make_similar_hats(self, request, _ctx): - if request.inches < 0: + async def make_similar_hats(self, request, _ctx): + if request.inches < 0: + raise ConnectError(Code.INVALID_ARGUMENT, "Size must be non-negative") + yield Hat(size=request.inches, color="orange") + yield Hat(size=request.inches, color="blue") + + async def make_various_hats(self, request, _ctx): + colors = itertools.cycle(("black", "white", "gold")) + async for s in request: + if s.inches < 0: raise ConnectError(Code.INVALID_ARGUMENT, "Size must be non-negative") - yield Hat(size=request.inches, color="orange") - yield Hat(size=request.inches, color="blue") - - async def make_various_hats(self, request, _ctx): - colors = itertools.cycle(("black", "white", "gold")) - async for s in request: - if s.inches < 0: - raise ConnectError( - Code.INVALID_ARGUMENT, "Size must be non-negative" - ) - yield Hat(size=s.inches, color=next(colors)) - - app = HaberdasherASGIApplication( - SimpleHaberdasher(), interceptors=(server_interceptor,) - ) - transport = ASGITransport(app) - async with HaberdasherClient( - "http://localhost", - interceptors=(client_interceptor,), - http_client=Client(transport=transport), - ) as client: - yield client - - -@pytest.mark.asyncio -async def test_intercept_unary_async( - client_async: HaberdasherClient, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - result = await client_async.make_hat(Size(inches=10)) - assert result == Hat(size=10, color="green") - assert client_interceptor.result == ["Hello MakeHat and goodbye"] - assert server_interceptor.result == ["Hello MakeHat and goodbye"] - - -@pytest.mark.asyncio -async def test_intercept_unary_async_error( - client_async: HaberdasherClient, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - with pytest.raises(ConnectError): - await client_async.make_hat(Size(inches=-10)) - assert client_interceptor.result == [ - "Hello MakeHat and goodbye with error Size must be non-negative" - ] - assert server_interceptor.result == [ - "Hello MakeHat and goodbye with error Size must be non-negative" - ] - + yield Hat(size=s.inches, color=next(colors)) -@pytest.mark.asyncio -async def test_intercept_client_stream_async( - client_async: HaberdasherClient, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - async def requests(): - yield Size(inches=10) - yield Size(inches=20) - - result = await client_async.make_flexible_hat(requests()) - assert result == Hat(size=30, color="red") - assert client_interceptor.result == ["Hello MakeFlexibleHat and goodbye"] - assert server_interceptor.result == ["Hello MakeFlexibleHat and goodbye"] - - -@pytest.mark.asyncio -async def test_intercept_client_stream_async_error( - client_async: HaberdasherClient, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - async def requests(): - yield Size(inches=-10) - yield Size(inches=20) - with pytest.raises(ConnectError): - await client_async.make_flexible_hat(requests()) - assert client_interceptor.result == [ - "Hello MakeFlexibleHat and goodbye with error Size must be non-negative" - ] - assert server_interceptor.result == [ - "Hello MakeFlexibleHat and goodbye with error Size must be non-negative" - ] - - -@pytest.mark.asyncio -async def test_intercept_server_stream_async( - client_async: HaberdasherClient, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - result = [r async for r in client_async.make_similar_hats(Size(inches=15))] - - assert result == [Hat(size=15, color="orange"), Hat(size=15, color="blue")] - assert client_interceptor.result == ["Hello MakeSimilarHats and goodbye"] - assert server_interceptor.result == ["Hello MakeSimilarHats and goodbye"] - - -@pytest.mark.asyncio -async def test_intercept_server_stream_async_error( - client_async: HaberdasherClient, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - with pytest.raises(ConnectError): - async for _ in client_async.make_similar_hats(Size(inches=-15)): - pass - - assert client_interceptor.result == [ - "Hello MakeSimilarHats and goodbye with error Size must be non-negative" - ] - assert server_interceptor.result == [ - "Hello MakeSimilarHats and goodbye with error Size must be non-negative" - ] - - -@pytest.mark.asyncio -async def test_intercept_bidi_stream_async( - client_async: HaberdasherClient, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - async def requests(): - yield Size(inches=25) - yield Size(inches=35) - yield Size(inches=45) - - result = [r async for r in client_async.make_various_hats(requests())] - - assert result == [ - Hat(size=25, color="black"), - Hat(size=35, color="white"), - Hat(size=45, color="gold"), - ] - assert client_interceptor.result == ["Hello MakeVariousHats and goodbye"] - assert server_interceptor.result == ["Hello MakeVariousHats and goodbye"] - - -@pytest.mark.asyncio -async def test_intercept_bidi_stream_async_error( - client_async: HaberdasherClient, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - async def requests(): - yield Size(inches=-25) - yield Size(inches=35) - yield Size(inches=45) - - with pytest.raises(ConnectError): - async for _ in client_async.make_various_hats(requests()): - pass - - assert client_interceptor.result == [ - "Hello MakeVariousHats and goodbye with error Size must be non-negative" - ] - assert server_interceptor.result == [ - "Hello MakeVariousHats and goodbye with error Size must be non-negative" - ] - - -@pytest.fixture -def client_sync( - client_interceptor: RequestInterceptor, server_interceptor: RequestInterceptor -): - class SimpleHaberdasherSync(HaberdasherSync): - def make_hat(self, request, _ctx): - if request.inches < 0: +class SimpleHaberdasherSync(HaberdasherSync): + def make_hat(self, request, _ctx): + if request.inches < 0: + raise ConnectError(Code.INVALID_ARGUMENT, "Size must be non-negative") + return Hat(size=request.inches, color="green") + + def make_flexible_hat(self, request, _ctx): + size = 0 + for s in request: + if s.inches < 0: raise ConnectError(Code.INVALID_ARGUMENT, "Size must be non-negative") - return Hat(size=request.inches, color="green") - - def make_flexible_hat(self, request, _ctx): - size = 0 - for s in request: - if s.inches < 0: - raise ConnectError( - Code.INVALID_ARGUMENT, "Size must be non-negative" - ) - size += s.inches - return Hat(size=size, color="red") + size += s.inches + return Hat(size=size, color="red") - def make_similar_hats(self, request, _ctx): - if request.inches < 0: + def make_similar_hats(self, request, _ctx): + if request.inches < 0: + raise ConnectError(Code.INVALID_ARGUMENT, "Size must be non-negative") + yield Hat(size=request.inches, color="orange") + yield Hat(size=request.inches, color="blue") + + def make_various_hats(self, request, _ctx): + colors = itertools.cycle(("black", "white", "gold")) + requests = [*request] + for s in requests: + if s.inches < 0: raise ConnectError(Code.INVALID_ARGUMENT, "Size must be non-negative") - yield Hat(size=request.inches, color="orange") - yield Hat(size=request.inches, color="blue") - - def make_various_hats(self, request, _ctx): - colors = itertools.cycle(("black", "white", "gold")) - requests = [*request] - for s in requests: - if s.inches < 0: - raise ConnectError( - Code.INVALID_ARGUMENT, "Size must be non-negative" - ) - yield Hat(size=s.inches, color=next(colors)) - - app = HaberdasherWSGIApplication( - SimpleHaberdasherSync(), interceptors=(server_interceptor,) - ) - transport = WSGITransport(app) - with HaberdasherClientSync( - "http://localhost", - interceptors=(client_interceptor,), - http_client=SyncClient(transport), - ) as client: - yield client + yield Hat(size=s.inches, color=next(colors)) -def test_intercept_unary_sync( - client_sync: HaberdasherClientSync, +@pytest_asyncio.fixture(params=["async", "sync"]) +async def client( + request: pytest.FixtureRequest, client_interceptor: RequestInterceptor, server_interceptor: RequestInterceptor, -) -> None: - result = client_sync.make_hat(Size(inches=10)) - assert result == Hat(size=10, color="green") - assert client_interceptor.result == ["Hello MakeHat and goodbye"] - assert server_interceptor.result == ["Hello MakeHat and goodbye"] - - -def test_intercept_unary_sync_error( - client_sync: HaberdasherClientSync, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - with pytest.raises(ConnectError): - client_sync.make_hat(Size(inches=-10)) - assert client_interceptor.result == [ - "Hello MakeHat and goodbye with error Size must be non-negative" - ] - assert server_interceptor.result == [ - "Hello MakeHat and goodbye with error Size must be non-negative" - ] - - -def test_intercept_client_stream_sync( - client_sync: HaberdasherClientSync, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - def requests(): - yield Size(inches=10) - yield Size(inches=20) - - result = client_sync.make_flexible_hat(requests()) - assert result == Hat(size=30, color="red") - assert client_interceptor.result == ["Hello MakeFlexibleHat and goodbye"] - assert server_interceptor.result == ["Hello MakeFlexibleHat and goodbye"] - - -def test_intercept_client_stream_sync_error( - client_sync: HaberdasherClientSync, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - def requests(): - yield Size(inches=-10) - yield Size(inches=20) - - with pytest.raises(ConnectError): - client_sync.make_flexible_hat(requests()) - assert client_interceptor.result == [ - "Hello MakeFlexibleHat and goodbye with error Size must be non-negative" - ] - assert server_interceptor.result == [ - "Hello MakeFlexibleHat and goodbye with error Size must be non-negative" - ] +): + if request.param == "async": + app = HaberdasherASGIApplication( + SimpleHaberdasher(), interceptors=(server_interceptor,) + ) + async with HaberdasherClient( + "http://localhost", + http_client=Client(ASGITransport(app)), + interceptors=(client_interceptor,), + ) as client: + yield client + else: + app = HaberdasherWSGIApplication( + SimpleHaberdasherSync(), interceptors=(server_interceptor,) + ) + with HaberdasherClientSync( + "http://localhost", + http_client=SyncClient(WSGITransport(app)), + interceptors=(client_interceptor,), + ) as client: + yield client -def test_intercept_server_stream_sync( - client_sync: HaberdasherClientSync, +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("method", "request_", "expected", "name"), + [ + pytest.param( + "make_hat", + Size(inches=10), + Hat(size=10, color="green"), + "MakeHat", + id="unary", + ), + pytest.param( + "make_flexible_hat", + [Size(inches=10), Size(inches=20)], + Hat(size=30, color="red"), + "MakeFlexibleHat", + id="client_stream", + ), + pytest.param( + "make_similar_hats", + Size(inches=15), + [Hat(size=15, color="orange"), Hat(size=15, color="blue")], + "MakeSimilarHats", + id="server_stream", + ), + pytest.param( + "make_various_hats", + [Size(inches=25), Size(inches=35), Size(inches=45)], + [ + Hat(size=25, color="black"), + Hat(size=35, color="white"), + Hat(size=45, color="gold"), + ], + "MakeVariousHats", + id="bidi_stream", + ), + ], +) +async def test_intercept( + client: HaberdasherClient | HaberdasherClientSync, client_interceptor: RequestInterceptor, server_interceptor: RequestInterceptor, + method: str, + request_: Size | list[Size], + expected: Hat | list[Hat], + name: str, ) -> None: - result = list(client_sync.make_similar_hats(Size(inches=15))) - - assert result == [Hat(size=15, color="orange"), Hat(size=15, color="blue")] - assert client_interceptor.result == ["Hello MakeSimilarHats and goodbye"] - assert server_interceptor.result == ["Hello MakeSimilarHats and goodbye"] + assert await call(getattr(client, method), request_) == expected + assert client_interceptor.result == [f"Hello {name} and goodbye"] + assert server_interceptor.result == [f"Hello {name} and goodbye"] -def test_intercept_server_stream_sync_error( - client_sync: HaberdasherClientSync, +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("method", "request_", "name"), + [ + pytest.param("make_hat", Size(inches=-10), "MakeHat", id="unary"), + pytest.param( + "make_flexible_hat", + [Size(inches=-10), Size(inches=20)], + "MakeFlexibleHat", + id="client_stream", + ), + pytest.param( + "make_similar_hats", Size(inches=-15), "MakeSimilarHats", id="server_stream" + ), + pytest.param( + "make_various_hats", + [Size(inches=-25), Size(inches=35), Size(inches=45)], + "MakeVariousHats", + id="bidi_stream", + ), + ], +) +async def test_intercept_error( + client: HaberdasherClient | HaberdasherClientSync, client_interceptor: RequestInterceptor, server_interceptor: RequestInterceptor, + method: str, + request_: Size | list[Size], + name: str, ) -> None: with pytest.raises(ConnectError): - list(client_sync.make_similar_hats(Size(inches=-15))) - assert client_interceptor.result == [ - "Hello MakeSimilarHats and goodbye with error Size must be non-negative" - ] - assert server_interceptor.result == [ - "Hello MakeSimilarHats and goodbye with error Size must be non-negative" - ] - - -def test_intercept_bidi_stream_sync( - client_sync: HaberdasherClientSync, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - def requests(): - yield Size(inches=25) - yield Size(inches=35) - yield Size(inches=45) - - result = list(client_sync.make_various_hats(requests())) - - assert result == [ - Hat(size=25, color="black"), - Hat(size=35, color="white"), - Hat(size=45, color="gold"), - ] - assert client_interceptor.result == ["Hello MakeVariousHats and goodbye"] - assert server_interceptor.result == ["Hello MakeVariousHats and goodbye"] + await call(getattr(client, method), request_) + expected = f"Hello {name} and goodbye with error Size must be non-negative" + assert client_interceptor.result == [expected] + assert server_interceptor.result == [expected] class _CountingHaberdasher(Haberdasher): @@ -584,7 +416,7 @@ async def make_hat(self, request, _ctx): interceptors=(EventMetadataInterceptor(), EventUnaryInterceptor()), ) async with HaberdasherClient( - "http://localhost", http_client=Client(transport=ASGITransport(app)) + "http://localhost", http_client=Client(ASGITransport(app)) ) as client: await client.make_hat(Size(inches=10)) @@ -741,7 +573,7 @@ async def make_similar_hats(self, request, _ctx): SimpleHaberdasher(), interceptors=(_ResponseMetadataInterceptor(),) ) async with HaberdasherClient( - "http://localhost", http_client=Client(transport=ASGITransport(app)) + "http://localhost", http_client=Client(ASGITransport(app)) ) as client: with ResponseMetadata() as resp: await client.make_hat(Size(inches=10)) @@ -779,24 +611,3 @@ def make_similar_hats(self, request, _ctx): for _ in client.make_similar_hats(Size(inches=10)): pass assert resp.trailers.get("x-interceptor-trailer") == "ran" - - -def test_intercept_bidi_stream_sync_error( - client_sync: HaberdasherClientSync, - client_interceptor: RequestInterceptor, - server_interceptor: RequestInterceptor, -) -> None: - def requests(): - yield Size(inches=-25) - yield Size(inches=35) - yield Size(inches=45) - - with pytest.raises(ConnectError): - list(client_sync.make_various_hats(requests())) - - assert client_interceptor.result == [ - "Hello MakeVariousHats and goodbye with error Size must be non-negative" - ] - assert server_interceptor.result == [ - "Hello MakeVariousHats and goodbye with error Size must be non-negative" - ] diff --git a/test/test_roundtrip.py b/test/test_roundtrip.py index 9a09820..8c91f48 100644 --- a/test/test_roundtrip.py +++ b/test/test_roundtrip.py @@ -15,7 +15,7 @@ from connectrpc.errors import ConnectError from connectrpc.server import DEFAULT_READ_MAX_BYTES -from ._util import resolve_compression +from ._util import call, resolve_compression from .connectrpc.example.haberdasher_connect import ( Haberdasher, HaberdasherASGIApplication, @@ -32,50 +32,47 @@ from asgiref.typing import HTTPDisconnectEvent, HTTPRequestEvent, HTTPScope +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["async", "sync"]) @pytest.mark.parametrize("proto_json", [False, True]) @pytest.mark.parametrize("compression_name", ["gzip", "br", "zstd", "identity"]) -def test_roundtrip_sync(proto_json: bool, compression_name: str) -> None: - class RoundtripHaberdasherSync(HaberdasherSync): - def make_hat(self, request, _ctx): - return Hat(size=request.inches, color="green") - +async def test_roundtrip(mode: str, proto_json: bool, compression_name: str) -> None: compression = resolve_compression(compression_name) - app = HaberdasherWSGIApplication( - RoundtripHaberdasherSync(), compressions=[compression] - ) - with HaberdasherClientSync( - "http://localhost", - http_client=SyncClient(WSGITransport(app=app)), - codec=proto_json_codec() if proto_json else None, - send_compression=compression, - accept_compression=[compression], - ) as client: - response = client.make_hat(request=Size(inches=10)) - assert response.size == 10 - assert response.color == "green" + codec = proto_json_codec() if proto_json else None + if mode == "async": + class RoundtripHaberdasher(Haberdasher): + async def make_hat(self, request, _ctx): + return Hat(size=request.inches, color="green") -@pytest.mark.parametrize("proto_json", [False, True]) -@pytest.mark.parametrize("compression_name", ["gzip", "br", "zstd", "identity"]) -@pytest.mark.asyncio -async def test_roundtrip_async(proto_json: bool, compression_name: str) -> None: - class DetailsHaberdasher(Haberdasher): - async def make_hat(self, request, _ctx): - return Hat(size=request.inches, color="green") + app = HaberdasherASGIApplication( + RoundtripHaberdasher(), compressions=[compression] + ) + client = HaberdasherClient( + "http://localhost", + http_client=Client(ASGITransport(app)), + codec=codec, + send_compression=compression, + accept_compression=[compression], + ) + else: - compression = resolve_compression(compression_name) - app = HaberdasherASGIApplication(DetailsHaberdasher(), compressions=[compression]) - transport = ASGITransport(app) - async with HaberdasherClient( - "http://localhost", - http_client=Client(transport), - codec=proto_json_codec() if proto_json else None, - send_compression=compression, - accept_compression=[compression], - ) as client: - response = await client.make_hat(request=Size(inches=10)) - assert response.size == 10 - assert response.color == "green" + class RoundtripHaberdasherSync(HaberdasherSync): + def make_hat(self, request, _ctx): + return Hat(size=request.inches, color="green") + + app = HaberdasherWSGIApplication( + RoundtripHaberdasherSync(), compressions=[compression] + ) + client = HaberdasherClientSync( + "http://localhost", + http_client=SyncClient(WSGITransport(app)), + codec=codec, + send_compression=compression, + accept_compression=[compression], + ) + response = await call(client.make_hat, Size(inches=10)) + assert response == Hat(size=10, color="green") # A request and a response containing a field the receiver's schema doesn't have, @@ -215,44 +212,43 @@ async def app(_scope, _receive, send): await client.make_hat(request=Size(inches=10)) -def test_roundtrip_sync_connect_get_empty_request() -> None: - class RoundtripHaberdasherSync(HaberdasherSync): - def make_hat(self, request, _ctx): - return Hat(size=request.inches, color="green") - +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["async", "sync"]) +async def test_roundtrip_connect_get_empty_request(mode: str) -> None: compression = resolve_compression("identity") - app = HaberdasherWSGIApplication( - RoundtripHaberdasherSync(), compressions=[compression] - ) - with HaberdasherClientSync( - "http://localhost", - http_client=SyncClient(WSGITransport(app=app)), - send_compression=compression, - accept_compression=[compression], - ) as client: - response = client.make_hat(request=Size(), use_get=True) - assert response.size == 0 - assert response.color == "green" + if mode == "async": + class RoundtripHaberdasher(Haberdasher): + async def make_hat(self, request, _ctx): + return Hat(size=request.inches, color="green") -@pytest.mark.asyncio -async def test_roundtrip_async_connect_get_empty_request() -> None: - class RoundtripHaberdasher(Haberdasher): - async def make_hat(self, request, _ctx): - return Hat(size=request.inches, color="green") + app = HaberdasherASGIApplication( + RoundtripHaberdasher(), compressions=[compression] + ) + client = HaberdasherClient( + "http://localhost", + http_client=Client(ASGITransport(app)), + send_compression=compression, + accept_compression=[compression], + ) + response = await client.make_hat(Size(), use_get=True) + else: - compression = resolve_compression("identity") - app = HaberdasherASGIApplication(RoundtripHaberdasher(), compressions=[compression]) - transport = ASGITransport(app) - async with HaberdasherClient( - "http://localhost", - http_client=Client(transport=transport), - send_compression=compression, - accept_compression=[compression], - ) as client: - response = await client.make_hat(request=Size(), use_get=True) - assert response.size == 0 - assert response.color == "green" + class RoundtripHaberdasherSync(HaberdasherSync): + def make_hat(self, request, _ctx): + return Hat(size=request.inches, color="green") + + app = HaberdasherWSGIApplication( + RoundtripHaberdasherSync(), compressions=[compression] + ) + client = HaberdasherClientSync( + "http://localhost", + http_client=SyncClient(WSGITransport(app)), + send_compression=compression, + accept_compression=[compression], + ) + response = client.make_hat(Size(), use_get=True) + assert response == Hat(size=0, color="green") @pytest.mark.parametrize("proto_json", [False, True]) @@ -469,49 +465,39 @@ async def request_stream(): assert len(responses) == 1 +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["async", "sync"]) @pytest.mark.parametrize("large", [False, True]) -def test_message_limit_unary_error_sync(large: bool) -> None: +async def test_message_limit_unary_error(mode: str, large: bool) -> None: message = "x" * 100 if large else "small" + if mode == "async": - class FailingHaberdasher(HaberdasherSync): - def make_hat(self, _request, _ctx): - raise ConnectError(Code.FAILED_PRECONDITION, message) + class FailingHaberdasher(Haberdasher): + async def make_hat(self, _request, _ctx): + raise ConnectError(Code.FAILED_PRECONDITION, message) - app = HaberdasherWSGIApplication(FailingHaberdasher()) - with ( - HaberdasherClientSync( + client = HaberdasherClient( "http://localhost", - http_client=SyncClient(WSGITransport(app)), + http_client=Client( + ASGITransport(HaberdasherASGIApplication(FailingHaberdasher())) + ), read_max_bytes=100, - ) as client, - pytest.raises(ConnectError) as exc_info, - ): - client.make_hat(request=Size()) - if large: - assert exc_info.value.code == Code.RESOURCE_EXHAUSTED - assert exc_info.value.message == "message is larger than configured max 100" + ) else: - assert exc_info.value.code == Code.FAILED_PRECONDITION - assert exc_info.value.message == message - - -@pytest.mark.parametrize("large", [False, True]) -@pytest.mark.asyncio -async def test_message_limit_unary_error_async(large: bool) -> None: - message = "x" * 100 if large else "small" - class FailingHaberdasher(Haberdasher): - async def make_hat(self, _request, _ctx): - raise ConnectError(Code.FAILED_PRECONDITION, message) + class FailingHaberdasherSync(HaberdasherSync): + def make_hat(self, _request, _ctx): + raise ConnectError(Code.FAILED_PRECONDITION, message) - app = HaberdasherASGIApplication(FailingHaberdasher()) - async with HaberdasherClient( - "http://localhost", - http_client=Client(transport=ASGITransport(app)), - read_max_bytes=100, - ) as client: - with pytest.raises(ConnectError) as exc_info: - await client.make_hat(request=Size()) + client = HaberdasherClientSync( + "http://localhost", + http_client=SyncClient( + WSGITransport(HaberdasherWSGIApplication(FailingHaberdasherSync())) + ), + read_max_bytes=100, + ) + with pytest.raises(ConnectError) as exc_info: + await call(client.make_hat, Size()) if large: assert exc_info.value.code == Code.RESOURCE_EXHAUSTED assert exc_info.value.message == "message is larger than configured max 100" @@ -524,65 +510,52 @@ async def make_hat(self, _request, _ctx): _BIG_DESCRIPTION_LENGTH = DEFAULT_READ_MAX_BYTES + 1 - 5 +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["async", "sync"]) @pytest.mark.parametrize("unlimited", [False, True]) -def test_message_limit_default_sync(unlimited: bool) -> None: - class EchoSizeHaberdasher(HaberdasherSync): - def make_hat(self, request, _ctx): - return Hat(size=len(request.description)) - +async def test_message_limit_default(mode: str, unlimited: bool) -> None: big_size = Size(description="X" * _BIG_DESCRIPTION_LENGTH) assert len(big_size.to_binary()) == DEFAULT_READ_MAX_BYTES + 1 # We specifically want to test not setting vs setting to None - app = ( - HaberdasherWSGIApplication(EchoSizeHaberdasher(), read_max_bytes=None) - if unlimited - else HaberdasherWSGIApplication(EchoSizeHaberdasher()) - ) - with HaberdasherClientSync( - "http://localhost", http_client=SyncClient(WSGITransport(app)) - ) as client: - if unlimited: - response = client.make_hat(request=big_size) - assert response.size == _BIG_DESCRIPTION_LENGTH - else: - with pytest.raises(ConnectError) as exc_info: - client.make_hat(request=big_size) - assert exc_info.value.code == Code.RESOURCE_EXHAUSTED - assert ( - exc_info.value.message - == f"message is larger than configured max {DEFAULT_READ_MAX_BYTES}" - ) + if mode == "async": + class EchoSizeHaberdasher(Haberdasher): + async def make_hat(self, request, _ctx): + return Hat(size=len(request.description)) -@pytest.mark.parametrize("unlimited", [False, True]) -@pytest.mark.asyncio -async def test_message_limit_default_async(unlimited: bool) -> None: - class EchoSizeHaberdasher(Haberdasher): - async def make_hat(self, request, _ctx): - return Hat(size=len(request.description)) + app = ( + HaberdasherASGIApplication(EchoSizeHaberdasher(), read_max_bytes=None) + if unlimited + else HaberdasherASGIApplication(EchoSizeHaberdasher()) + ) + client = HaberdasherClient( + "http://localhost", http_client=Client(ASGITransport(app)) + ) + else: - big_size = Size(description="X" * _BIG_DESCRIPTION_LENGTH) - assert len(big_size.to_binary()) == DEFAULT_READ_MAX_BYTES + 1 - # We specifically want to test not setting vs setting to None - app = ( - HaberdasherASGIApplication(EchoSizeHaberdasher(), read_max_bytes=None) - if unlimited - else HaberdasherASGIApplication(EchoSizeHaberdasher()) - ) - async with HaberdasherClient( - "http://localhost", http_client=Client(transport=ASGITransport(app)) - ) as client: - if unlimited: - response = await client.make_hat(request=big_size) - assert response.size == _BIG_DESCRIPTION_LENGTH - else: - with pytest.raises(ConnectError) as exc_info: - await client.make_hat(request=big_size) - assert exc_info.value.code == Code.RESOURCE_EXHAUSTED - assert ( - exc_info.value.message - == f"message is larger than configured max {DEFAULT_READ_MAX_BYTES}" - ) + class EchoSizeHaberdasherSync(HaberdasherSync): + def make_hat(self, request, _ctx): + return Hat(size=len(request.description)) + + app = ( + HaberdasherWSGIApplication(EchoSizeHaberdasherSync(), read_max_bytes=None) + if unlimited + else HaberdasherWSGIApplication(EchoSizeHaberdasherSync()) + ) + client = HaberdasherClientSync( + "http://localhost", http_client=SyncClient(WSGITransport(app)) + ) + if unlimited: + response = await call(client.make_hat, big_size) + assert response == Hat(size=_BIG_DESCRIPTION_LENGTH) + else: + with pytest.raises(ConnectError) as exc_info: + await call(client.make_hat, big_size) + assert exc_info.value.code == Code.RESOURCE_EXHAUSTED + assert ( + exc_info.value.message + == f"message is larger than configured max {DEFAULT_READ_MAX_BYTES}" + ) @pytest.mark.asyncio