diff --git a/simvue/api/request.py b/simvue/api/request.py index 89ab9555..a4514a09 100644 --- a/simvue/api/request.py +++ b/simvue/api/request.py @@ -14,6 +14,7 @@ import requests from tenacity import ( + RetryCallState, retry, retry_if_exception_type, stop_after_attempt, @@ -46,6 +47,17 @@ class RetryableHTTPError(Exception): pass +def _rewind_request_streams(retry_state: RetryCallState) -> None: + """Rewind file-like request bodies before retrying.""" + files = retry_state.kwargs.get("files") or {} + streams = (*files.values(), retry_state.kwargs.get("data")) + + for value in streams: + stream = value[1] if isinstance(value, tuple) else value + if callable(seek := getattr(stream, "seek", None)): + seek(0) + + @retry( wait=wait_exponential(multiplier=RETRY_MULTIPLIER, min=RETRY_MIN, max=RETRY_MAX), stop=stop_after_attempt(RETRY_STOP), @@ -56,6 +68,7 @@ class RetryableHTTPError(Exception): requests.exceptions.ConnectionError, ), ), + before_sleep=_rewind_request_streams, reraise=True, ) def post( @@ -138,6 +151,7 @@ def post( ), ), stop=stop_after_attempt(RETRY_STOP), + before_sleep=_rewind_request_streams, reraise=True, ) def put( diff --git a/tests/unit/test_request.py b/tests/unit/test_request.py new file mode 100644 index 00000000..019513b5 --- /dev/null +++ b/tests/unit/test_request.py @@ -0,0 +1,65 @@ +import io + +import pytest +from tenacity import wait_none + +from simvue.api import request + +_PAYLOAD = b"stable-file-contents" + + +class _Response: + def __init__(self, status_code: int) -> None: + self.status_code = status_code + + +@pytest.mark.local +@pytest.mark.parametrize("tuple_form", (False, True), ids=("file-object", "tuple")) +def test_post_retry_rewinds_file_stream(mocker, tuple_form: bool) -> None: + payloads = [] + stream = io.BytesIO(_PAYLOAD) + file_value = ("artifact.bin", stream) if tuple_form else stream + + def fake_post(*args, **kwargs): + file = kwargs["files"]["file"] + file = file[1] if isinstance(file, tuple) else file + payloads.append(file.read()) + return _Response(503 if len(payloads) == 1 else 200) + + mocker.patch.object(request.requests, "post", side_effect=fake_post) + + response = request.post.retry_with(wait=wait_none())( + "https://example.invalid", + headers={}, + params={}, + data={}, + is_json=False, + files={"file": file_value}, + timeout=1, + ) + + assert response.status_code == 200 + assert payloads == [_PAYLOAD, _PAYLOAD] + + +@pytest.mark.local +def test_put_retry_rewinds_data_stream(mocker) -> None: + payloads = [] + stream = io.BytesIO(_PAYLOAD) + + def fake_put(*args, **kwargs): + payloads.append(kwargs["data"].read()) + return _Response(503 if len(payloads) == 1 else 200) + + mocker.patch.object(request.requests, "put", side_effect=fake_put) + + response = request.put.retry_with(wait=wait_none())( + "https://example.invalid", + headers={}, + data=stream, + is_json=False, + timeout=1, + ) + + assert response.status_code == 200 + assert payloads == [_PAYLOAD, _PAYLOAD]