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
67 changes: 44 additions & 23 deletions src/httpware/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -202,6 +202,18 @@ async def _send_capped_async(client: httpx2.AsyncClient, request: httpx2.Request
await flow.aclose()


@contextlib.asynccontextmanager
async def _stream_async(
client: httpx2.AsyncClient, method: str, url: httpx2.URL | str, kwargs: dict[str, typing.Any]
) -> AsyncIterator[httpx2.Response]:
"""Mirror of `httpx2.AsyncClient.stream` that builds via `_build_request`."""
response = await client.send(_build_request(client, method, url, kwargs), stream=True)
try:
yield response
finally:
await response.aclose()


@contextlib.asynccontextmanager
async def _stream_capped_async(
client: httpx2.AsyncClient,
Expand All @@ -211,7 +223,7 @@ async def _stream_capped_async(
cap: int,
) -> AsyncIterator[httpx2.Response]:
"""Async mirror of `httpx2.AsyncClient.stream` that sends via `_send_capped_async`."""
response = await _send_capped_async(client, client.build_request(method, url, **kwargs), cap)
response = await _send_capped_async(client, _build_request(client, method, url, kwargs), cap)
try:
yield response
finally:
Expand Down Expand Up @@ -271,6 +283,18 @@ def _send_capped(client: httpx2.Client, request: httpx2.Request, cap: int) -> ht
flow.close()


@contextlib.contextmanager
def _stream(
client: httpx2.Client, method: str, url: httpx2.URL | str, kwargs: dict[str, typing.Any]
) -> Iterator[httpx2.Response]:
"""Sync mirror of `_stream_async`."""
response = client.send(_build_request(client, method, url, kwargs), stream=True)
try:
yield response
finally:
response.close()


@contextlib.contextmanager
def _stream_capped(
client: httpx2.Client,
Expand All @@ -280,7 +304,7 @@ def _stream_capped(
cap: int,
) -> Iterator[httpx2.Response]:
"""Sync mirror of `_stream_capped_async`."""
response = _send_capped(client, client.build_request(method, url, **kwargs), cap)
response = _send_capped(client, _build_request(client, method, url, kwargs), cap)
try:
yield response
finally:
Expand Down Expand Up @@ -322,14 +346,17 @@ 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."""
def _build_request(
client: httpx2.Client | httpx2.AsyncClient, method: str, url: httpx2.URL | str, kwargs: dict[str, typing.Any]
) -> httpx2.Request:
"""Build via `client`, keeping the URL's own query bytes ahead of any params, which httpx2 would replace."""
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)
if not parsed.query:
return client.build_request(method, url, **kwargs)
request = client.build_request(method, parsed.copy_with(query=None), **kwargs)
built_query = request.url.query
request.url = request.url.copy_with(query=parsed.query + b"&" + built_query if built_query else parsed.query)
return request


class AsyncClient:
Expand Down Expand Up @@ -430,8 +457,7 @@ async def send_with_response(

def build_request(self, method: str, url: str, **kwargs: typing.Any) -> httpx2.Request:
"""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)
return _build_request(self._httpx2_client, method, url, kwargs)

def _prepare_request( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwargs-forwarding complexity is structural
self,
Expand All @@ -448,7 +474,6 @@ 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 @@ -460,7 +485,7 @@ def _prepare_request( # noqa: PLR0913 — mirrors httpx2 per-method signatures;
data=data,
files=files,
)
request = self._httpx2_client.build_request(method, merged_url, **kwargs)
request = _build_request(self._httpx2_client, method, 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 @@ -1216,7 +1241,6 @@ 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 @@ -1231,9 +1255,9 @@ async def stream( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwa

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

def build_request(self, method: str, url: str, **kwargs: typing.Any) -> httpx2.Request:
"""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)
return _build_request(self._httpx2_client, method, url, kwargs)

def _prepare_request( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwargs-forwarding complexity is structural
self,
Expand All @@ -1410,7 +1433,6 @@ 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 @@ -1422,7 +1444,7 @@ def _prepare_request( # noqa: PLR0913 — mirrors httpx2 per-method signatures;
data=data,
files=files,
)
request = self._httpx2_client.build_request(method, merged_url, **kwargs)
request = _build_request(self._httpx2_client, method, 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 @@ -2175,7 +2197,6 @@ 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 @@ -2190,9 +2211,9 @@ def stream( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwargs-fo

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

from http import HTTPStatus

Expand Down Expand Up @@ -30,14 +30,15 @@ def handler(request: httpx2.Request) -> httpx2.Response:

_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&b=1", {"b": "2"}, {}, "a=1&b=1&b=2", id="same-key-appended"),
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", None, {"c": "3"}, "a=1&c=3", id="client-params-after-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"
"https://example.test/x?a=1&c=1", {"b": "2"}, {"c": "3"}, "a=1&c=1&c=3&b=2", id="url-and-client-same-key-kept"
),
pytest.param("https://example.test/x", {"b": "2"}, {"c": "3"}, "c=3&b=2", id="no-url-query-unchanged"),
pytest.param("https://example.test/x?q=a%20b&flag", {"p": "1"}, {}, "q=a%20b&flag&p=1", id="url-query-bytes-kept"),
]


Expand Down Expand Up @@ -98,7 +99,7 @@ def test_sync_stream_merges_url_query(
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"
assert str(request.url) == "https://example.test/api/items?a=1&c=3&b=2"


def test_url_query_without_params_is_left_verbatim() -> None:
Expand Down
Loading