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
98 changes: 37 additions & 61 deletions test/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
Loading
Loading