From c322e53ee06179e7ce9069df4e37c9ac35f326b0 Mon Sep 17 00:00:00 2001 From: bensynapse <118375461+bensynapse@users.noreply.github.com> Date: Sat, 3 Oct 2026 20:06:05 +0300 Subject: [PATCH] Close response streams on terminal validation errors --- httpx_retries/transport.py | 10 +-- tests/test_transport.py | 133 +++++++++++++++++++++++++++++++++++++ 2 files changed, 139 insertions(+), 4 deletions(-) diff --git a/httpx_retries/transport.py b/httpx_retries/transport.py index 4abd729..0d13c6f 100644 --- a/httpx_retries/transport.py +++ b/httpx_retries/transport.py @@ -45,8 +45,9 @@ def _retry_operation( response.request = request try: retry.validate_response(response) - except Exception as e: - if retry.is_exhausted() or not retry.is_retryable_exception(e): + except BaseException as e: + if not isinstance(e, Exception) or retry.is_exhausted() or not retry.is_retryable_exception(e): + response.close() raise continue response.extensions["retry"] = retry @@ -90,8 +91,9 @@ async def _retry_operation_async( await retry.validate_response(response) else: retry.validate_response(response) - except Exception as e: - if retry.is_exhausted() or not retry.is_retryable_exception(e): + except BaseException as e: + if not isinstance(e, Exception) or retry.is_exhausted() or not retry.is_retryable_exception(e): + await response.aclose() raise continue response.extensions["retry"] = retry diff --git a/tests/test_transport.py b/tests/test_transport.py index f1998b8..4609840 100644 --- a/tests/test_transport.py +++ b/tests/test_transport.py @@ -103,6 +103,29 @@ async def handle_async_request(self, request: Request) -> Response: return create_response(request, 200) +class TrackingByteStream(httpx.SyncByteStream): + def __init__(self) -> None: + self.close_count = 0 + + def __iter__(self) -> Generator[bytes, None, None]: + yield b"response body" + + def close(self) -> None: + self.close_count += 1 + + +class TrackingAsyncByteStream(httpx.AsyncByteStream): + def __init__(self) -> None: + self.close_count = 0 + + async def __aiter__(self) -> AsyncGenerator[bytes, None]: + yield b"response body" + + async def aclose(self) -> None: + await asyncio.sleep(0) + self.close_count += 1 + + def test_successful_request(mock_responses: MockResponse) -> None: mock_sleep, _ = mock_responses transport = RetryTransport() @@ -924,3 +947,113 @@ def validate(response: httpx.Response) -> None: assert response.status_code == 200 assert validated == [200] + + +@pytest.mark.parametrize("error", [ValueError("bad response"), KeyboardInterrupt()]) +def test_validate_response_terminal_error_closes_stream(error: BaseException) -> None: + stream = TrackingByteStream() + response = Response(200, stream=stream) + send = Mock(return_value=response) + + def validate(actual_response: Response) -> None: + assert actual_response is response + raise error + + transport = RetryTransport(transport=httpx.MockTransport(send), retry=Retry(validate_response=validate)) + with httpx.Client(transport=transport) as client: + with pytest.raises(type(error)) as exc_info: + client.get("https://example.com") + + assert exc_info.value is error + assert response.is_closed + assert stream.close_count == 1 + send.assert_called_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("async_validator", [False, True]) +async def test_async_validate_response_terminal_error_closes_stream(async_validator: bool) -> None: + stream = TrackingAsyncByteStream() + response = Response(200, stream=stream) + send = Mock(return_value=response) + error = ValueError("bad response") + + def validate(actual_response: Response) -> None: + assert actual_response is response + raise error + + async def avalidate(actual_response: Response) -> None: + validate(actual_response) + + retry = Retry(validate_response=avalidate if async_validator else validate) + transport = RetryTransport(transport=httpx.MockTransport(send), retry=retry) + async with httpx.AsyncClient(transport=transport) as client: + with pytest.raises(ValueError) as exc_info: + await client.get("https://example.com") + + assert exc_info.value is error + assert response.is_closed + assert stream.close_count == 1 + send.assert_called_once() + + +@pytest.mark.asyncio +async def test_async_validate_response_cancellation_closes_stream() -> None: + stream = TrackingAsyncByteStream() + response = Response(200, stream=stream) + send = Mock(return_value=response) + validating = asyncio.Event() + + async def validate(actual_response: Response) -> None: + assert actual_response is response + validating.set() + await asyncio.Event().wait() + + transport = RetryTransport(transport=httpx.MockTransport(send), retry=Retry(validate_response=validate)) + async with httpx.AsyncClient(transport=transport) as client: + task = asyncio.create_task(client.get("https://example.com")) + await asyncio.wait_for(validating.wait(), timeout=5) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert response.is_closed + assert stream.close_count == 1 + send.assert_called_once() + + +def test_validate_response_success_keeps_stream_open() -> None: + stream = TrackingByteStream() + response = Response(200, stream=stream) + validate = Mock() + transport = RetryTransport( + transport=httpx.MockTransport(lambda _: response), retry=Retry(validate_response=validate) + ) + with httpx.Client(transport=transport) as client: + result = client.send(client.build_request("GET", "https://example.com"), stream=True) + + validate.assert_called_once_with(response) + assert result is response + assert not response.is_closed + assert stream.close_count == 0 + assert response.read() == b"response body" + assert stream.close_count == 1 + + +@pytest.mark.asyncio +async def test_async_validate_response_success_keeps_stream_open() -> None: + stream = TrackingAsyncByteStream() + response = Response(200, stream=stream) + validate = AsyncMock() + transport = RetryTransport( + transport=httpx.MockTransport(lambda _: response), retry=Retry(validate_response=validate) + ) + async with httpx.AsyncClient(transport=transport) as client: + result = await client.send(client.build_request("GET", "https://example.com"), stream=True) + + validate.assert_awaited_once_with(response) + assert result is response + assert not response.is_closed + assert stream.close_count == 0 + assert await response.aread() == b"response body" + assert stream.close_count == 1