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
32 changes: 24 additions & 8 deletions src/httpware/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -145,6 +145,16 @@ def _assemble_request_kwargs( # noqa: PLR0913 — 9 per-request kwargs from htt
return kwargs


def _merge_url_query(
url: httpx2.URL | str, params: typing.Any | None, client_params: httpx2.QueryParams
) -> tuple[httpx2.URL | str, typing.Any | None]:
"""Fold the URL's own query into `params`; httpx2 would otherwise replace it."""
parsed = httpx2.URL(url)
if not parsed.query or (params is None and not client_params):
return url, params
return parsed.copy_with(query=None), parsed.params.merge(params)


class AsyncClient:
"""Async HTTP client: thin wrapper around httpx2 with typed decoding and middleware."""

Expand Down Expand Up @@ -261,8 +271,9 @@ async def send_with_response(
return response, bound.decode(response)

def build_request(self, method: str, url: str, **kwargs: typing.Any) -> httpx2.Request:
"""Delegate request construction to the wrapped httpx2.AsyncClient."""
return self._httpx2_client.build_request(method, url, **kwargs)
"""Delegate request construction to the wrapped httpx2.AsyncClient, keeping the URL's own query."""
merged_url, params = _merge_url_query(url, kwargs.pop("params", None), self._httpx2_client.params)
return self._httpx2_client.build_request(method, merged_url, params=params, **kwargs)

def _prepare_request( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwargs-forwarding complexity is structural
self,
Expand All @@ -279,6 +290,7 @@ def _prepare_request( # noqa: PLR0913 — mirrors httpx2 per-method signatures;
data: typing.Any | None = None,
files: typing.Any | None = None,
) -> httpx2.Request:
merged_url, params = _merge_url_query(url, params, self._httpx2_client.params)
kwargs = _assemble_request_kwargs(
params=params,
headers=headers,
Expand All @@ -290,7 +302,7 @@ def _prepare_request( # noqa: PLR0913 — mirrors httpx2 per-method signatures;
data=data,
files=files,
)
request = self._httpx2_client.build_request(method, url, **kwargs)
request = self._httpx2_client.build_request(method, merged_url, **kwargs)
if _is_streaming_body_async(content) or _is_streaming_body_async(data) or _is_streaming_body_async(files):
request.extensions[STREAMING_BODY_MARKER] = True
return request
Expand Down Expand Up @@ -1046,6 +1058,7 @@ async def stream( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwa
Maps httpx2 exceptions raised during the request OR body consumption to
httpware exceptions via _httpx2_exception_mapper.
"""
merged_url, params = _merge_url_query(url, params, self._httpx2_client.params)
kwargs = _assemble_request_kwargs(
params=params,
headers=headers,
Expand All @@ -1058,7 +1071,7 @@ async def stream( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwa
files=files,
)

async with _httpx2_exception_mapper(), self._httpx2_client.stream(method, url, **kwargs) as response:
async with _httpx2_exception_mapper(), self._httpx2_client.stream(method, merged_url, **kwargs) as response:
if HTTPStatus.BAD_REQUEST <= response.status_code < 600: # noqa: PLR2004 — 600 is the synthetic upper bound for 5xx
cap = self._max_response_body_bytes
if cap is None:
Expand Down Expand Up @@ -1234,8 +1247,9 @@ def send_with_response(
return response, bound.decode(response)

def build_request(self, method: str, url: str, **kwargs: typing.Any) -> httpx2.Request:
"""Delegate request construction to the wrapped httpx2.Client."""
return self._httpx2_client.build_request(method, url, **kwargs)
"""Delegate request construction to the wrapped httpx2.Client, keeping the URL's own query."""
merged_url, params = _merge_url_query(url, kwargs.pop("params", None), self._httpx2_client.params)
return self._httpx2_client.build_request(method, merged_url, params=params, **kwargs)

def _prepare_request( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwargs-forwarding complexity is structural
self,
Expand All @@ -1252,6 +1266,7 @@ def _prepare_request( # noqa: PLR0913 — mirrors httpx2 per-method signatures;
data: typing.Any | None = None,
files: typing.Any | None = None,
) -> httpx2.Request:
merged_url, params = _merge_url_query(url, params, self._httpx2_client.params)
kwargs = _assemble_request_kwargs(
params=params,
headers=headers,
Expand All @@ -1263,7 +1278,7 @@ def _prepare_request( # noqa: PLR0913 — mirrors httpx2 per-method signatures;
data=data,
files=files,
)
request = self._httpx2_client.build_request(method, url, **kwargs)
request = self._httpx2_client.build_request(method, merged_url, **kwargs)
if _is_streaming_body_sync(content) or _is_streaming_body_sync(data) or _is_streaming_body_sync(files):
request.extensions[STREAMING_BODY_MARKER] = True
return request
Expand Down Expand Up @@ -2016,6 +2031,7 @@ def stream( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwargs-fo
Maps httpx2 exceptions raised during the request OR body consumption to
httpware exceptions via _httpx2_exception_mapper_sync.
"""
merged_url, params = _merge_url_query(url, params, self._httpx2_client.params)
kwargs = _assemble_request_kwargs(
params=params,
headers=headers,
Expand All @@ -2028,7 +2044,7 @@ def stream( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwargs-fo
files=files,
)

with _httpx2_exception_mapper_sync(), self._httpx2_client.stream(method, url, **kwargs) as response:
with _httpx2_exception_mapper_sync(), self._httpx2_client.stream(method, merged_url, **kwargs) as response:
if HTTPStatus.BAD_REQUEST <= response.status_code < 600: # noqa: PLR2004 — 600 is the synthetic upper bound for 5xx
cap = self._max_response_body_bytes
if cap is None:
Expand Down
107 changes: 107 additions & 0 deletions tests/test_url_query_merge.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,107 @@
"""The URL's own query string survives per-request and client-level `params`."""

from http import HTTPStatus

import httpx2
import pytest

from httpware import AsyncClient, Client


def _ok(request: httpx2.Request) -> httpx2.Response:
return httpx2.Response(HTTPStatus.OK, request=request)


def _async_client(captured: list[httpx2.Request], params: dict[str, str] | None = None) -> AsyncClient:
def handler(request: httpx2.Request) -> httpx2.Response:
captured.append(request)
return _ok(request)

return AsyncClient(httpx2_client=httpx2.AsyncClient(transport=httpx2.MockTransport(handler), params=params))


def _sync_client(captured: list[httpx2.Request], params: dict[str, str] | None = None) -> Client:
def handler(request: httpx2.Request) -> httpx2.Response:
captured.append(request)
return _ok(request)

return Client(httpx2_client=httpx2.Client(transport=httpx2.MockTransport(handler), params=params))


_CASES = [
pytest.param("https://example.test/x?a=1", {"b": "2"}, {}, "a=1&b=2", id="url-query-plus-params"),
pytest.param("https://example.test/x?a=1&b=1", {"b": "2"}, {}, "a=1&b=2", id="params-override-same-key"),
pytest.param("https://example.test/x?a=1&a=2", {"b": "3"}, {}, "a=1&a=2&b=3", id="repeated-url-keys-kept"),
pytest.param("https://example.test/x?a=1", {}, {}, "a=1", id="empty-params-keeps-url-query"),
pytest.param("https://example.test/x?a=1", None, {"c": "3"}, "c=3&a=1", id="client-params-keep-url-query"),
pytest.param(
"https://example.test/x?a=1&c=1", {"b": "2"}, {"c": "3"}, "c=1&a=1&b=2", id="url-query-overrides-client"
),
pytest.param("https://example.test/x", {"b": "2"}, {"c": "3"}, "c=3&b=2", id="no-url-query-unchanged"),
]


@pytest.mark.parametrize(("url", "params", "client_params", "expected_query"), _CASES)
def test_async_build_request_merges_url_query(
url: str, params: dict[str, str] | None, client_params: dict[str, str], expected_query: str
) -> None:
client = _async_client([], params=client_params)
assert client.build_request("GET", url, params=params).url.query == expected_query.encode()


@pytest.mark.parametrize(("url", "params", "client_params", "expected_query"), _CASES)
def test_sync_build_request_merges_url_query(
url: str, params: dict[str, str] | None, client_params: dict[str, str], expected_query: str
) -> None:
client = _sync_client([], params=client_params)
assert client.build_request("GET", url, params=params).url.query == expected_query.encode()


@pytest.mark.parametrize(("url", "params", "client_params", "expected_query"), _CASES)
async def test_async_get_merges_url_query(
url: str, params: dict[str, str] | None, client_params: dict[str, str], expected_query: str
) -> None:
captured: list[httpx2.Request] = []
await _async_client(captured, params=client_params).get(url, params=params)
assert captured[0].url.query == expected_query.encode()


@pytest.mark.parametrize(("url", "params", "client_params", "expected_query"), _CASES)
def test_sync_get_merges_url_query(
url: str, params: dict[str, str] | None, client_params: dict[str, str], expected_query: str
) -> None:
captured: list[httpx2.Request] = []
_sync_client(captured, params=client_params).get(url, params=params)
assert captured[0].url.query == expected_query.encode()


@pytest.mark.parametrize(("url", "params", "client_params", "expected_query"), _CASES)
async def test_async_stream_merges_url_query(
url: str, params: dict[str, str] | None, client_params: dict[str, str], expected_query: str
) -> None:
captured: list[httpx2.Request] = []
async with _async_client(captured, params=client_params).stream("GET", url, params=params):
pass
assert captured[0].url.query == expected_query.encode()


@pytest.mark.parametrize(("url", "params", "client_params", "expected_query"), _CASES)
def test_sync_stream_merges_url_query(
url: str, params: dict[str, str] | None, client_params: dict[str, str], expected_query: str
) -> None:
captured: list[httpx2.Request] = []
with _sync_client(captured, params=client_params).stream("GET", url, params=params):
pass
assert captured[0].url.query == expected_query.encode()


def test_owned_client_relative_url_merges_with_base_url() -> None:
client = Client(base_url="https://example.test/api", params={"c": "3"})
request = client.build_request("GET", "items?a=1", params={"b": "2"})
assert str(request.url) == "https://example.test/api/items?c=3&a=1&b=2"


def test_url_query_without_params_is_left_verbatim() -> None:
client = _sync_client([])
url = "https://example.test/x?cursor=a%2Fb+c"
assert str(client.build_request("GET", url).url) == url
Loading