Skip to content
Open
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
10 changes: 6 additions & 4 deletions httpx_retries/transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
133 changes: 133 additions & 0 deletions tests/test_transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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
Loading