diff --git a/docs/errors.md b/docs/errors.md index e17fdca..cb7778b 100644 --- a/docs/errors.md +++ b/docs/errors.md @@ -186,7 +186,7 @@ Unlike `DecodeError`, this error is raised before the request is sent. ## `ResponseTooLargeError` -Both clients accept `max_response_body_bytes: int | None = None`. By default there is no limit. When it is set, a response body larger than the cap raises `ResponseTooLargeError` instead of being returned, whatever the status: a `200` trips it as easily as a `500`. The cap counts decoded bytes, after decompression. It applies to `send()` and the verb methods, and to the error body that `stream()` reads before raising a `StatusError`. Bytes you read yourself while iterating a `stream()` are never capped. With a cap set and `follow_redirects=True`, httpware follows the redirects itself and caps only the final response. It closes each intermediate redirect response without reading its body, so the responses in `response.history` have no content. Client `auth` is sent to the first URL only. `httpx2` keeps its `Authorization` header on a redirect within the same origin or from `http` to `https` on the same host, and drops it otherwise. +Both clients accept `max_response_body_bytes: int | None = None`. By default there is no limit. When it is set, a response body larger than the cap raises `ResponseTooLargeError` instead of being returned, whatever the status: a `200` trips it as easily as a `500`. The cap counts decoded bytes, after decompression. It applies to `send()` and the verb methods, and to the error body that `stream()` reads before raising a `StatusError`. Bytes you read yourself while iterating a `stream()` are never capped, except with an `auth` that sets `requires_response_body`: that auth needs every body, so `stream()` reads the final one under the cap before yielding it. With a cap set, httpware runs the client's `auth` flow and, with `follow_redirects=True`, follows the redirects itself, and caps only the final response. It closes each intermediate response, such as a redirect or a `DigestAuth` challenge, without reading its body, so the responses in `response.history` have no content. An `auth` that sets `requires_response_body`, like a token refresh that parses the token endpoint's JSON, gets each response read under the cap instead. Client `auth` is sent to the first URL only. `httpx2` keeps its `Authorization` header on a redirect within the same origin or from `http` to `https` on the same host, and drops it otherwise. `ResponseTooLargeError` carries: diff --git a/src/httpware/client.py b/src/httpware/client.py index 57f90b9..b141a89 100644 --- a/src/httpware/client.py +++ b/src/httpware/client.py @@ -40,6 +40,7 @@ "event_hooks": "event_hooks=... is not supported; use middleware=... instead.", } _TOO_MANY_REDIRECTS_MESSAGE = "Exceeded maximum allowed redirects." +_NO_AUTH = httpx2.Auth() _BASE_URL_QUERY_MESSAGE = ( "base_url must not contain a query string: httpx2 appends request paths after it, " "producing malformed URLs. Pass the query as params=... instead." @@ -133,58 +134,153 @@ def _select_httpx2_options( return forwarded -async def _send_following_redirects_async(client: httpx2.AsyncClient, request: httpx2.Request) -> httpx2.Response: - """Send `request` streaming, following redirects hop by hop without reading intermediate bodies.""" - history: list[httpx2.Response] = [] - response = await client.send(request, stream=True, follow_redirects=False) - while client.follow_redirects and response.next_request is not None: +def _request_auth(client: httpx2.Client | httpx2.AsyncClient, request: httpx2.Request) -> httpx2.Auth: + """Return the auth httpx2 applies to `request`: the client's, else Basic from URL credentials, else none.""" + if client.auth is not None: + return client.auth + if request.url.username or request.url.password: + return httpx2.BasicAuth(request.url.username, request.url.password) + return _NO_AUTH + + +async def _read_capped_and_close_async(streaming: httpx2.Response, cap: int) -> httpx2.Response: + """Buffer `streaming` under `cap` via `_read_capped_async`, closing it either way.""" + try: + return await _read_capped_async(streaming, cap, streaming.request) + finally: + await streaming.aclose() + + +async def _send_redirect_hops_async( + client: httpx2.AsyncClient, + request: httpx2.Request, + prior_history: list[httpx2.Response], +) -> httpx2.Response: + """Send `request` streaming, following redirects hop by hop without reading intermediate bodies. + + Histories and the `max_redirects` count include `prior_history`, as in httpx2. + """ + hops = list(prior_history) + while True: + if len(hops) > client.max_redirects: + raise httpx2.TooManyRedirects(_TOO_MANY_REDIRECTS_MESSAGE, request=request) + response = await client.send(request, stream=True, follow_redirects=False, auth=_NO_AUTH) + response.history = list(hops) + if not client.follow_redirects or response.next_request is None: + return response await response.aclose() - history.append(response) - if len(history) > client.max_redirects: - raise httpx2.TooManyRedirects(_TOO_MANY_REDIRECTS_MESSAGE, request=response.next_request) - response = await client.send(response.next_request, stream=True, follow_redirects=False, auth=None) - response.history = history - return response + hops.append(response) + request = response.next_request + + +async def _send_capped_async(client: httpx2.AsyncClient, request: httpx2.Request, cap: int) -> httpx2.Response: + """Send `request` streaming, driving the client's auth flow and redirects without reading intermediate bodies. + + An auth that sets `requires_response_body` gets each response buffered under `cap` instead. + """ + auth = _request_auth(client, request) + flow = auth.async_auth_flow(request) + history: list[httpx2.Response] = [] + try: + request = await anext(flow) + while True: + response = await _send_redirect_hops_async(client, request, history) + if auth.requires_response_body: + response = await _read_capped_and_close_async(response, cap) + try: + next_request = await flow.asend(response) + except StopAsyncIteration: + return response + except BaseException: + await response.aclose() + raise + await response.aclose() + response.history = list(history) + history.append(response) + request = next_request + finally: + await flow.aclose() @contextlib.asynccontextmanager -async def _stream_following_redirects_async( +async def _stream_capped_async( client: httpx2.AsyncClient, method: str, url: httpx2.URL | str, kwargs: dict[str, typing.Any], + cap: int, ) -> AsyncIterator[httpx2.Response]: - """Async mirror of `httpx2.AsyncClient.stream` that follows redirects via `_send_following_redirects_async`.""" - response = await _send_following_redirects_async(client, client.build_request(method, url, **kwargs)) + """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) try: yield response finally: await response.aclose() -def _send_following_redirects(client: httpx2.Client, request: httpx2.Request) -> httpx2.Response: - """Sync mirror of `_send_following_redirects_async`.""" - history: list[httpx2.Response] = [] - response = client.send(request, stream=True, follow_redirects=False) - while client.follow_redirects and response.next_request is not None: +def _read_capped_and_close(streaming: httpx2.Response, cap: int) -> httpx2.Response: + """Sync mirror of `_read_capped_and_close_async`.""" + try: + return _read_capped(streaming, cap, streaming.request) + finally: + streaming.close() + + +def _send_redirect_hops( + client: httpx2.Client, + request: httpx2.Request, + prior_history: list[httpx2.Response], +) -> httpx2.Response: + """Sync mirror of `_send_redirect_hops_async`.""" + hops = list(prior_history) + while True: + if len(hops) > client.max_redirects: + raise httpx2.TooManyRedirects(_TOO_MANY_REDIRECTS_MESSAGE, request=request) + response = client.send(request, stream=True, follow_redirects=False, auth=_NO_AUTH) + response.history = list(hops) + if not client.follow_redirects or response.next_request is None: + return response response.close() - history.append(response) - if len(history) > client.max_redirects: - raise httpx2.TooManyRedirects(_TOO_MANY_REDIRECTS_MESSAGE, request=response.next_request) - response = client.send(response.next_request, stream=True, follow_redirects=False, auth=None) - response.history = history - return response + hops.append(response) + request = response.next_request + + +def _send_capped(client: httpx2.Client, request: httpx2.Request, cap: int) -> httpx2.Response: + """Sync mirror of `_send_capped_async`.""" + auth = _request_auth(client, request) + flow = auth.sync_auth_flow(request) + history: list[httpx2.Response] = [] + try: + request = next(flow) + while True: + response = _send_redirect_hops(client, request, history) + if auth.requires_response_body: + response = _read_capped_and_close(response, cap) + try: + next_request = flow.send(response) + except StopIteration: + return response + except BaseException: + response.close() + raise + response.close() + response.history = list(history) + history.append(response) + request = next_request + finally: + flow.close() @contextlib.contextmanager -def _stream_following_redirects( +def _stream_capped( client: httpx2.Client, method: str, url: httpx2.URL | str, kwargs: dict[str, typing.Any], + cap: int, ) -> Iterator[httpx2.Response]: - """Sync mirror of `_stream_following_redirects_async`.""" - response = _send_following_redirects(client, client.build_request(method, url, **kwargs)) + """Sync mirror of `_stream_capped_async`.""" + response = _send_capped(client, client.build_request(method, url, **kwargs), cap) try: yield response finally: @@ -284,11 +380,8 @@ async def _terminal(self, request: httpx2.Request) -> httpx2.Response: if cap is None: response = await self._httpx2_client.send(request) else: - streaming = await _send_following_redirects_async(self._httpx2_client, request) - try: - response = await _read_capped_async(streaming, cap, streaming.request) - finally: - await streaming.aclose() + streaming = await _send_capped_async(self._httpx2_client, request, cap) + response = await _read_capped_and_close_async(streaming, cap) except RuntimeError as exc: if self._httpx2_client.is_closed: raise TransportError(str(exc)) from exc @@ -1140,7 +1233,7 @@ async def stream( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwa opened = ( self._httpx2_client.stream(method, merged_url, **kwargs) if cap is None - else _stream_following_redirects_async(self._httpx2_client, method, merged_url, kwargs) + else _stream_capped_async(self._httpx2_client, method, merged_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 @@ -1225,11 +1318,8 @@ def _terminal(self, request: httpx2.Request) -> httpx2.Response: if cap is None: response = self._httpx2_client.send(request) else: - streaming = _send_following_redirects(self._httpx2_client, request) - try: - response = _read_capped(streaming, cap, streaming.request) - finally: - streaming.close() + streaming = _send_capped(self._httpx2_client, request, cap) + response = _read_capped_and_close(streaming, cap) except RuntimeError as exc: if self._httpx2_client.is_closed: raise TransportError(str(exc)) from exc @@ -2102,7 +2192,7 @@ def stream( # noqa: PLR0913 — mirrors httpx2 per-method signatures; kwargs-fo opened = ( self._httpx2_client.stream(method, merged_url, **kwargs) if cap is None - else _stream_following_redirects(self._httpx2_client, method, merged_url, kwargs) + else _stream_capped(self._httpx2_client, method, merged_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 diff --git a/tests/test_client_body_cap_auth.py b/tests/test_client_body_cap_auth.py new file mode 100644 index 0000000..151fb51 --- /dev/null +++ b/tests/test_client_body_cap_auth.py @@ -0,0 +1,309 @@ +"""max_response_body_bytes with a multi-step auth flow: intermediate auth responses stay under the cap.""" + +import typing +from collections.abc import AsyncIterator, Callable, Generator, Iterator +from http import HTTPStatus + +import httpx2 +import pytest + +from httpware import AsyncClient, Client +from httpware.errors import ResponseTooLargeError, TransportError + + +_CHALLENGE = 'Digest realm="api", nonce="abc", qop="auth"' + + +class _TokenRefreshAuth(httpx2.Auth): + requires_response_body = True + + def auth_flow(self, request: httpx2.Request) -> Generator[httpx2.Request, httpx2.Response]: + response = yield request + if response.status_code == HTTPStatus.UNAUTHORIZED: + token_response = yield httpx2.Request("POST", "https://example.test/token") + request.headers["authorization"] = f"Bearer {token_response.json()['token']}" + yield request + + +def _token_endpoint(token_padding: int) -> httpx2.MockTransport: + def handler(request: httpx2.Request) -> httpx2.Response: + if request.url.path == "/token": + return httpx2.Response(HTTPStatus.OK, json={"token": "fresh", "padding": "x" * token_padding}) + if request.headers.get("authorization") == "Bearer fresh": + return httpx2.Response(HTTPStatus.OK, content=b"done") + return httpx2.Response(HTTPStatus.UNAUTHORIZED) + + return httpx2.MockTransport(handler) + + +def _digest_challenge(body: Callable[[], AsyncIterator[bytes] | Iterator[bytes]]) -> httpx2.MockTransport: + def handler(request: httpx2.Request) -> httpx2.Response: + if request.headers.get("authorization", "").startswith("Digest "): + return httpx2.Response(HTTPStatus.OK, content=b"done") + return httpx2.Response(HTTPStatus.UNAUTHORIZED, headers={"www-authenticate": _CHALLENGE}, content=body()) + + return httpx2.MockTransport(handler) + + +def _huge_challenge_body(pulled: list[bytes]) -> httpx2.MockTransport: + async def huge_body() -> AsyncIterator[bytes]: + for _ in range(100): + pulled.append(b"x" * 1024) + yield pulled[-1] + + return _digest_challenge(huge_body) + + +def _huge_challenge_body_sync(pulled: list[bytes]) -> httpx2.MockTransport: + def huge_body() -> Iterator[bytes]: + for _ in range(100): + pulled.append(b"x" * 1024) + yield pulled[-1] + + return _digest_challenge(huge_body) + + +@pytest.mark.parametrize(("cap", "challenge_read"), [(None, True), (1024, False)]) +async def test_async_reads_a_digest_challenge_body_only_without_a_cap( + cap: int | None, + challenge_read: bool, +) -> None: + pulled: list[bytes] = [] + async with AsyncClient( + transport=_huge_challenge_body(pulled), + auth=httpx2.DigestAuth("u", "p"), + max_response_body_bytes=cap, + ) as client: + response = await client.get("https://example.test/") + assert response.content == b"done" + assert bool(pulled) is challenge_read + + +@pytest.mark.parametrize(("cap", "challenge_read"), [(None, True), (1024, False)]) +def test_sync_reads_a_digest_challenge_body_only_without_a_cap( + cap: int | None, + challenge_read: bool, +) -> None: + pulled: list[bytes] = [] + with Client( + transport=_huge_challenge_body_sync(pulled), + auth=httpx2.DigestAuth("u", "p"), + max_response_body_bytes=cap, + ) as client: + response = client.get("https://example.test/") + assert response.content == b"done" + assert bool(pulled) is challenge_read + + +async def test_async_auth_flow_reading_bodies_gets_them_under_the_cap() -> None: + async with AsyncClient( + transport=_token_endpoint(token_padding=0), + auth=_TokenRefreshAuth(), + max_response_body_bytes=1024, + ) as client: + response = await client.get("https://example.test/") + assert response.content == b"done" + + +async def test_async_auth_flow_reading_bodies_rejects_one_over_the_cap() -> None: + async with AsyncClient( + transport=_token_endpoint(token_padding=2048), + auth=_TokenRefreshAuth(), + max_response_body_bytes=1024, + ) as client: + with pytest.raises(ResponseTooLargeError): + await client.get("https://example.test/") + + +def test_sync_auth_flow_reading_bodies_gets_them_under_the_cap() -> None: + with Client( + transport=_token_endpoint(token_padding=0), + auth=_TokenRefreshAuth(), + max_response_body_bytes=1024, + ) as client: + response = client.get("https://example.test/") + assert response.content == b"done" + + +def test_sync_auth_flow_reading_bodies_rejects_one_over_the_cap() -> None: + with ( + Client( + transport=_token_endpoint(token_padding=2048), + auth=_TokenRefreshAuth(), + max_response_body_bytes=1024, + ) as client, + pytest.raises(ResponseTooLargeError), + ): + client.get("https://example.test/") + + +def _echo_authorization(request: httpx2.Request) -> httpx2.Response: + return httpx2.Response(HTTPStatus.OK, content=request.headers.get("authorization", "").encode()) + + +@pytest.mark.parametrize("cap", [None, 1024]) +async def test_async_url_credentials_authenticate_with_or_without_a_cap(cap: int | None) -> None: + async with AsyncClient(transport=httpx2.MockTransport(_echo_authorization), max_response_body_bytes=cap) as client: + response = await client.get("https://u:p@example.test/") + assert response.content == b"Basic dTpw" + + +@pytest.mark.parametrize("cap", [None, 1024]) +def test_sync_url_credentials_authenticate_with_or_without_a_cap(cap: int | None) -> None: + with Client(transport=httpx2.MockTransport(_echo_authorization), max_response_body_bytes=cap) as client: + response = client.get("https://u:p@example.test/") + assert response.content == b"Basic dTpw" + + +def _digest_behind_redirect(request: httpx2.Request) -> httpx2.Response: + if request.url.path == "/start": + return httpx2.Response(HTTPStatus.FOUND, headers={"location": "/final"}) + if request.headers.get("authorization", "").startswith("Digest "): + return httpx2.Response(HTTPStatus.OK, content=b"done") + return httpx2.Response(HTTPStatus.UNAUTHORIZED, headers={"www-authenticate": _CHALLENGE}) + + +def _history_shape(response: httpx2.Response) -> list[tuple[int, str, list[typing.Any]]]: + return [(hop.status_code, hop.url.path, _history_shape(hop)) for hop in response.history] + + +def _malformed_challenge(request: httpx2.Request) -> httpx2.Response: # noqa: ARG001 + return httpx2.Response(HTTPStatus.UNAUTHORIZED, headers={"www-authenticate": 'Digest realm="api"'}) + + +@pytest.mark.parametrize("cap", [None, 1024]) +async def test_async_digest_after_a_redirect_is_the_same_with_or_without_a_cap(cap: int | None) -> None: + async with AsyncClient( + transport=httpx2.MockTransport(_digest_behind_redirect), + auth=httpx2.DigestAuth("u", "p"), + follow_redirects=True, + max_response_body_bytes=cap, + ) as client: + response = await client.get("https://example.test/start") + assert response.content == b"done" + assert _history_shape(response) == [ + (HTTPStatus.UNAUTHORIZED, "/final", []), + (HTTPStatus.FOUND, "/start", [(HTTPStatus.UNAUTHORIZED, "/final", [])]), + ] + + +@pytest.mark.parametrize("cap", [None, 1024]) +def test_sync_digest_after_a_redirect_is_the_same_with_or_without_a_cap(cap: int | None) -> None: + with Client( + transport=httpx2.MockTransport(_digest_behind_redirect), + auth=httpx2.DigestAuth("u", "p"), + follow_redirects=True, + max_response_body_bytes=cap, + ) as client: + response = client.get("https://example.test/start") + assert response.content == b"done" + assert _history_shape(response) == [ + (HTTPStatus.UNAUTHORIZED, "/final", []), + (HTTPStatus.FOUND, "/start", [(HTTPStatus.UNAUTHORIZED, "/final", [])]), + ] + + +@pytest.mark.parametrize("cap", [None, 1024]) +async def test_async_auth_flow_error_is_the_same_with_or_without_a_cap(cap: int | None) -> None: + async with AsyncClient( + transport=httpx2.MockTransport(_malformed_challenge), + auth=httpx2.DigestAuth("u", "p"), + max_response_body_bytes=cap, + ) as client: + with pytest.raises(TransportError, match="Malformed Digest") as caught: + await client.get("https://example.test/") + assert type(caught.value) is TransportError + + +@pytest.mark.parametrize("cap", [None, 1024]) +def test_sync_auth_flow_error_is_the_same_with_or_without_a_cap(cap: int | None) -> None: + with ( + Client( + transport=httpx2.MockTransport(_malformed_challenge), + auth=httpx2.DigestAuth("u", "p"), + max_response_body_bytes=cap, + ) as client, + pytest.raises(TransportError, match="Malformed Digest") as caught, + ): + client.get("https://example.test/") + assert type(caught.value) is TransportError + + +@pytest.mark.parametrize("cap", [None, 1024]) +async def test_async_auth_steps_count_toward_max_redirects_with_or_without_a_cap(cap: int | None) -> None: + async with AsyncClient( + transport=httpx2.MockTransport(_digest_behind_redirect), + auth=httpx2.DigestAuth("u", "p"), + max_redirects=0, + max_response_body_bytes=cap, + ) as client: + with pytest.raises(TransportError, match="Exceeded maximum allowed redirects"): + await client.get("https://example.test/final") + + +@pytest.mark.parametrize("cap", [None, 1024]) +def test_sync_auth_steps_count_toward_max_redirects_with_or_without_a_cap(cap: int | None) -> None: + with ( + Client( + transport=httpx2.MockTransport(_digest_behind_redirect), + auth=httpx2.DigestAuth("u", "p"), + max_redirects=0, + max_response_body_bytes=cap, + ) as client, + pytest.raises(TransportError, match="Exceeded maximum allowed redirects"), + ): + client.get("https://example.test/final") + + +async def test_async_stream_never_reads_a_digest_challenge_body_under_a_cap() -> None: + pulled: list[bytes] = [] + async with ( + AsyncClient( + transport=_huge_challenge_body(pulled), + auth=httpx2.DigestAuth("u", "p"), + max_response_body_bytes=1024, + ) as client, + client.stream("GET", "https://example.test/") as response, + ): + body = await response.aread() + assert body == b"done" + assert pulled == [] + + +def test_sync_stream_never_reads_a_digest_challenge_body_under_a_cap() -> None: + pulled: list[bytes] = [] + with ( + Client( + transport=_huge_challenge_body_sync(pulled), + auth=httpx2.DigestAuth("u", "p"), + max_response_body_bytes=1024, + ) as client, + client.stream("GET", "https://example.test/") as response, + ): + body = response.read() + assert body == b"done" + assert pulled == [] + + +async def test_async_stream_rejects_an_auth_read_body_over_the_cap() -> None: + async with AsyncClient( + transport=_token_endpoint(token_padding=2048), + auth=_TokenRefreshAuth(), + max_response_body_bytes=1024, + ) as client: + with pytest.raises(ResponseTooLargeError): + async with client.stream("GET", "https://example.test/"): + pytest.fail("unreachable") # pragma: no cover — stream() raises on enter + + +def test_sync_stream_rejects_an_auth_read_body_over_the_cap() -> None: + with ( + Client( + transport=_token_endpoint(token_padding=2048), + auth=_TokenRefreshAuth(), + max_response_body_bytes=1024, + ) as client, + pytest.raises(ResponseTooLargeError), + client.stream("GET", "https://example.test/"), + ): + pytest.fail("unreachable") # pragma: no cover — stream() raises on enter