From f219e0447755606634f579299dc6f1654af35cba Mon Sep 17 00:00:00 2001 From: HughhhhCoder Date: Mon, 31 Aug 2026 20:57:49 +0800 Subject: [PATCH 1/4] fix(client): retry only replayable request content --- src/openai/_base_client.py | 60 ++++++++++++++++++--- tests/test_client.py | 105 ++++++++++++++++++++++++++++++++++++- 2 files changed, 158 insertions(+), 7 deletions(-) diff --git a/src/openai/_base_client.py b/src/openai/_base_client.py index f195d04816..e50cf69526 100644 --- a/src/openai/_base_client.py +++ b/src/openai/_base_client.py @@ -110,6 +110,52 @@ log: logging.Logger = logging.getLogger(__name__) log.addFilter(SensitiveHeadersFilter()) + +class _RequestContentReplay: + def __init__(self, content: object) -> None: + self._content = content + self._position: int | None = None + self._replayable = True + + if content is None or isinstance(content, (bytes, bytearray)): + return + + if callable(getattr(content, "read", None)): + seekable = getattr(content, "seekable", None) + tell = getattr(content, "tell", None) + try: + if not callable(seekable) or not seekable() or not callable(tell): + self._replayable = False + return + position = tell() + except (OSError, ValueError): + self._replayable = False + return + + if isinstance(position, int): + self._position = position + else: + self._replayable = False + return + + self._replayable = not isinstance(content, (Iterator, AsyncIterator)) + + def rewind(self) -> bool: + if not self._replayable: + return False + if self._position is None: + return True + + seek = getattr(self._content, "seek", None) + if not callable(seek): + return False + try: + seek(self._position) + except (OSError, ValueError): + return False + return True + + # TODO: make base page type vars covariant SyncPageT = TypeVar("SyncPageT", bound="BaseSyncPage[Any]") AsyncPageT = TypeVar("AsyncPageT", bound="BaseAsyncPage[Any]") @@ -1047,6 +1093,7 @@ def request( response: httpx2.Response | None = None max_retries = input_options.get_max_retries(self.max_retries) + content_replay = _RequestContentReplay(input_options.content) retries_taken = 0 for retries_taken in range(max_retries + 1): @@ -1081,7 +1128,7 @@ def request( except timeout_exceptions() as err: log.debug("Encountered a timeout exception: %s", type(err).__name__) - if remaining_retries > 0: + if remaining_retries > 0 and content_replay.rewind(): self._sleep_for_retry( retries_taken=retries_taken, max_retries=max_retries, @@ -1098,7 +1145,7 @@ def request( except Exception as err: log.debug("Encountered exception: %s", type(err).__name__) - if remaining_retries > 0: + if remaining_retries > 0 and content_replay.rewind(): self._sleep_for_retry( retries_taken=retries_taken, max_retries=max_retries, @@ -1122,7 +1169,7 @@ def request( except status_exceptions() as err: # thrown on 4xx and 5xx status code log.debug("Encountered an HTTP status error: %i", response.status_code) - if remaining_retries > 0 and self._should_retry(err.response): + if remaining_retries > 0 and self._should_retry(err.response) and content_replay.rewind(): err.response.close() self._sleep_for_retry( retries_taken=retries_taken, @@ -1671,6 +1718,7 @@ async def request( response: httpx2.Response | None = None max_retries = input_options.get_max_retries(self.max_retries) + content_replay = _RequestContentReplay(input_options.content) retries_taken = 0 for retries_taken in range(max_retries + 1): @@ -1704,7 +1752,7 @@ async def request( except timeout_exceptions() as err: log.debug("Encountered a timeout exception: %s", type(err).__name__) - if remaining_retries > 0: + if remaining_retries > 0 and content_replay.rewind(): await self._sleep_for_retry( retries_taken=retries_taken, max_retries=max_retries, @@ -1721,7 +1769,7 @@ async def request( except Exception as err: log.debug("Encountered exception: %s", type(err).__name__) - if remaining_retries > 0: + if remaining_retries > 0 and content_replay.rewind(): await self._sleep_for_retry( retries_taken=retries_taken, max_retries=max_retries, @@ -1745,7 +1793,7 @@ async def request( except status_exceptions() as err: # thrown on 4xx and 5xx status code log.debug("Encountered an HTTP status error: %i", response.status_code) - if remaining_retries > 0 and self._should_retry(err.response): + if remaining_retries > 0 and self._should_retry(err.response) and content_replay.rewind(): await err.response.aclose() await self._sleep_for_retry( retries_taken=retries_taken, diff --git a/tests/test_client.py b/tests/test_client.py index d82c39e616..35ae2289f0 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -1,6 +1,7 @@ from __future__ import annotations import gc +import io import os import sys import json @@ -23,7 +24,7 @@ from openai._utils import asyncify from openai._models import BaseModel, FinalRequestOptions from openai._streaming import Stream, AsyncStream -from openai._exceptions import APIStatusError, APITimeoutError, APIResponseValidationError +from openai._exceptions import APIStatusError, APITimeoutError, APIConnectionError, APIResponseValidationError from openai._base_client import ( DEFAULT_TIMEOUT, HTTPX_DEFAULT_TIMEOUT, @@ -807,6 +808,70 @@ def mock_handler(request: httpx2.Request) -> httpx2.Response: assert response.content == file_content assert counter.value == 1 + @pytest.mark.parametrize("failure_mode", ["status", "timeout", "connection"]) + def test_binary_content_retry_does_not_reuse_iterator( + self, failure_mode: Literal["status", "timeout", "connection"] + ) -> None: + file_content = b"Hello, this is a test file." + request_bodies: list[bytes] = [] + + def mock_handler(request: httpx2.Request) -> httpx2.Response: + request_bodies.append(request.read()) + if len(request_bodies) > 1: + return httpx2.Response(200) + if failure_mode == "timeout": + raise httpx2.ReadTimeout("timed out", request=request) + if failure_mode == "connection": + raise httpx2.ConnectError("connection failed", request=request) + return httpx2.Response(500, json={"error": {}}) + + expected_error = { + "status": APIStatusError, + "timeout": APITimeoutError, + "connection": APIConnectionError, + }[failure_mode] + + with OpenAI( + base_url=base_url, + api_key=api_key, + max_retries=1, + http_client=httpx2.Client(transport=MockTransport(handler=mock_handler)), + ) as client: + with pytest.raises(expected_error): + client.post( + "/upload", + content=_make_sync_iterator([file_content]), + cast_to=httpx2.Response, + ) + + assert request_bodies == [file_content] + + @mock.patch("openai._base_client.BaseClient._calculate_retry_timeout", _low_retry_timeout) + def test_binary_content_retry_rewinds_seekable_stream(self) -> None: + file_content = b"Hello, this is a test file." + request_bodies: list[bytes] = [] + + def mock_handler(request: httpx2.Request) -> httpx2.Response: + request_bodies.append(request.read()) + return httpx2.Response(500 if len(request_bodies) == 1 else 200, json={"error": {}}) + + with OpenAI( + base_url=base_url, + api_key=api_key, + max_retries=1, + http_client=httpx2.Client(transport=MockTransport(handler=mock_handler)), + ) as client: + content = io.BytesIO(b"prefix" + file_content) + content.seek(len(b"prefix")) + response = client.post( + "/upload", + content=content, + cast_to=httpx2.Response, + ) + + assert response.status_code == 200 + assert request_bodies == [file_content, file_content] + @pytest.mark.respx2(base_url=base_url) def test_binary_content_upload_with_body_is_deprecated(self, respx2_mock: MockRouter, client: OpenAI) -> None: respx2_mock.post("/upload").mock(side_effect=mirror_request_content) @@ -2109,6 +2174,44 @@ async def mock_handler(request: httpx2.Request) -> httpx2.Response: assert response.content == file_content assert counter.value == 1 + @pytest.mark.parametrize("failure_mode", ["status", "timeout", "connection"]) + async def test_binary_content_retry_does_not_reuse_asynciterator( + self, failure_mode: Literal["status", "timeout", "connection"] + ) -> None: + file_content = b"Hello, this is a test file." + request_bodies: list[bytes] = [] + + async def mock_handler(request: httpx2.Request) -> httpx2.Response: + request_bodies.append(await request.aread()) + if len(request_bodies) > 1: + return httpx2.Response(200) + if failure_mode == "timeout": + raise httpx2.ReadTimeout("timed out", request=request) + if failure_mode == "connection": + raise httpx2.ConnectError("connection failed", request=request) + return httpx2.Response(500, json={"error": {}}) + + expected_error = { + "status": APIStatusError, + "timeout": APITimeoutError, + "connection": APIConnectionError, + }[failure_mode] + + async with AsyncOpenAI( + base_url=base_url, + api_key=api_key, + max_retries=1, + http_client=httpx2.AsyncClient(transport=MockTransport(handler=mock_handler)), + ) as client: + with pytest.raises(expected_error): + await client.post( + "/upload", + content=_make_async_iterator([file_content]), + cast_to=httpx2.Response, + ) + + assert request_bodies == [file_content] + @pytest.mark.respx2(base_url=base_url) async def test_binary_content_upload_with_body_is_deprecated( self, respx2_mock: MockRouter, async_client: AsyncOpenAI From 53256c27654918936c2681edd8db2a75c1f017a8 Mon Sep 17 00:00:00 2001 From: HughhhhCoder Date: Tue, 1 Sep 2026 10:44:11 +0800 Subject: [PATCH 2/4] fix(client): treat opaque iterables as non-replayable --- src/openai/_base_client.py | 5 ++- tests/test_client.py | 65 ++++++++++++++++++++++++++++++++++---- 2 files changed, 62 insertions(+), 8 deletions(-) diff --git a/src/openai/_base_client.py b/src/openai/_base_client.py index e50cf69526..36115106d2 100644 --- a/src/openai/_base_client.py +++ b/src/openai/_base_client.py @@ -138,7 +138,10 @@ def __init__(self, content: object) -> None: self._replayable = False return - self._replayable = not isinstance(content, (Iterator, AsyncIterator)) + # The iterable protocols do not guarantee a fresh iterator for each iteration. An object + # can return the same stored generator from __iter__ or __aiter__ without being an iterator + # itself, so only retry concrete containers whose repeatability is known here. + self._replayable = type(content) in (list, tuple) def rewind(self) -> bool: if not self._replayable: diff --git a/tests/test_client.py b/tests/test_client.py index 35ae2289f0..0c70ed8796 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -9,7 +9,7 @@ import inspect import dataclasses import tracemalloc -from typing import Any, Union, TypeVar, Callable, Iterable, Iterator, Optional, Coroutine, cast +from typing import Any, Union, TypeVar, Callable, Iterable, Iterator, Optional, Coroutine, AsyncIterable, cast from unittest import mock from typing_extensions import Literal, AsyncIterator, override @@ -114,6 +114,24 @@ async def _make_async_iterator(iterable: Iterable[T], counter: Optional[Counter] yield item +class _OneShotIterable(Iterable[T]): + def __init__(self, iterable: Iterable[T]) -> None: + self._iterator = iter(iterable) + + @override + def __iter__(self) -> Iterator[T]: + return self._iterator + + +class _OneShotAsyncIterable(AsyncIterable[T]): + def __init__(self, iterable: Iterable[T]) -> None: + self._iterator = _make_async_iterator(iterable) + + @override + def __aiter__(self) -> AsyncIterator[T]: + return self._iterator + + def _get_open_connections(client: OpenAI | AsyncOpenAI) -> int: transport = client._client._transport if isinstance(transport, httpx2.HTTPTransport) or isinstance(transport, httpx2.AsyncHTTPTransport): @@ -808,9 +826,12 @@ def mock_handler(request: httpx2.Request) -> httpx2.Response: assert response.content == file_content assert counter.value == 1 + @pytest.mark.parametrize("content_factory", [_make_sync_iterator, _OneShotIterable]) @pytest.mark.parametrize("failure_mode", ["status", "timeout", "connection"]) - def test_binary_content_retry_does_not_reuse_iterator( - self, failure_mode: Literal["status", "timeout", "connection"] + def test_binary_content_retry_does_not_reuse_one_shot_iterable( + self, + content_factory: Callable[[Iterable[bytes]], Iterable[bytes]], + failure_mode: Literal["status", "timeout", "connection"], ) -> None: file_content = b"Hello, this is a test file." request_bodies: list[bytes] = [] @@ -840,12 +861,39 @@ def mock_handler(request: httpx2.Request) -> httpx2.Response: with pytest.raises(expected_error): client.post( "/upload", - content=_make_sync_iterator([file_content]), + content=content_factory([file_content]), cast_to=httpx2.Response, ) assert request_bodies == [file_content] + @pytest.mark.parametrize("content_type", [list, tuple]) + @mock.patch("openai._base_client.BaseClient._calculate_retry_timeout", _low_retry_timeout) + def test_binary_content_retry_reuses_known_repeatable_iterable( + self, content_type: type[list[bytes]] | type[tuple[bytes, ...]] + ) -> None: + file_content = b"Hello, this is a test file." + request_bodies: list[bytes] = [] + + def mock_handler(request: httpx2.Request) -> httpx2.Response: + request_bodies.append(request.read()) + return httpx2.Response(500 if len(request_bodies) == 1 else 200, json={"error": {}}) + + with OpenAI( + base_url=base_url, + api_key=api_key, + max_retries=1, + http_client=httpx2.Client(transport=MockTransport(handler=mock_handler)), + ) as client: + response = client.post( + "/upload", + content=content_type([file_content]), + cast_to=httpx2.Response, + ) + + assert response.status_code == 200 + assert request_bodies == [file_content, file_content] + @mock.patch("openai._base_client.BaseClient._calculate_retry_timeout", _low_retry_timeout) def test_binary_content_retry_rewinds_seekable_stream(self) -> None: file_content = b"Hello, this is a test file." @@ -2174,9 +2222,12 @@ async def mock_handler(request: httpx2.Request) -> httpx2.Response: assert response.content == file_content assert counter.value == 1 + @pytest.mark.parametrize("content_factory", [_make_async_iterator, _OneShotAsyncIterable]) @pytest.mark.parametrize("failure_mode", ["status", "timeout", "connection"]) - async def test_binary_content_retry_does_not_reuse_asynciterator( - self, failure_mode: Literal["status", "timeout", "connection"] + async def test_binary_content_retry_does_not_reuse_one_shot_asynciterable( + self, + content_factory: Callable[[Iterable[bytes]], AsyncIterable[bytes]], + failure_mode: Literal["status", "timeout", "connection"], ) -> None: file_content = b"Hello, this is a test file." request_bodies: list[bytes] = [] @@ -2206,7 +2257,7 @@ async def mock_handler(request: httpx2.Request) -> httpx2.Response: with pytest.raises(expected_error): await client.post( "/upload", - content=_make_async_iterator([file_content]), + content=content_factory([file_content]), cast_to=httpx2.Response, ) From 06817547ee2a0766aa412b58be10563332fe087a Mon Sep 17 00:00:00 2001 From: HughhhhCoder Date: Thu, 3 Sep 2026 17:18:19 +0800 Subject: [PATCH 3/4] fix(client): honor multipart file replayability --- src/openai/_base_client.py | 42 ++++++++++++--- tests/test_client.py | 101 +++++++++++++++++++++++++++++++++++++ 2 files changed, 135 insertions(+), 8 deletions(-) diff --git a/src/openai/_base_client.py b/src/openai/_base_client.py index 36115106d2..4e9bfa43e2 100644 --- a/src/openai/_base_client.py +++ b/src/openai/_base_client.py @@ -159,6 +159,18 @@ def rewind(self) -> bool: return True +def _iter_file_contents(files: HttpxRequestFiles | None) -> Iterator[object]: + if files is None: + return + + entries = files.items() if isinstance(files, Mapping) else files + for _, file in entries: + if isinstance(file, tuple) and len(file) > 1: + yield file[1] + else: + yield file + + # TODO: make base page type vars covariant SyncPageT = TypeVar("SyncPageT", bound="BaseSyncPage[Any]") AsyncPageT = TypeVar("AsyncPageT", bound="BaseAsyncPage[Any]") @@ -1096,7 +1108,10 @@ def request( response: httpx2.Response | None = None max_retries = input_options.get_max_retries(self.max_retries) - content_replay = _RequestContentReplay(input_options.content) + request_body_replays = [ + _RequestContentReplay(input_options.content), + *(_RequestContentReplay(content) for content in _iter_file_contents(input_options.files)), + ] retries_taken = 0 for retries_taken in range(max_retries + 1): @@ -1131,7 +1146,7 @@ def request( except timeout_exceptions() as err: log.debug("Encountered a timeout exception: %s", type(err).__name__) - if remaining_retries > 0 and content_replay.rewind(): + if remaining_retries > 0 and all(replay.rewind() for replay in request_body_replays): self._sleep_for_retry( retries_taken=retries_taken, max_retries=max_retries, @@ -1148,7 +1163,7 @@ def request( except Exception as err: log.debug("Encountered exception: %s", type(err).__name__) - if remaining_retries > 0 and content_replay.rewind(): + if remaining_retries > 0 and all(replay.rewind() for replay in request_body_replays): self._sleep_for_retry( retries_taken=retries_taken, max_retries=max_retries, @@ -1172,7 +1187,11 @@ def request( except status_exceptions() as err: # thrown on 4xx and 5xx status code log.debug("Encountered an HTTP status error: %i", response.status_code) - if remaining_retries > 0 and self._should_retry(err.response) and content_replay.rewind(): + if ( + remaining_retries > 0 + and self._should_retry(err.response) + and all(replay.rewind() for replay in request_body_replays) + ): err.response.close() self._sleep_for_retry( retries_taken=retries_taken, @@ -1721,7 +1740,10 @@ async def request( response: httpx2.Response | None = None max_retries = input_options.get_max_retries(self.max_retries) - content_replay = _RequestContentReplay(input_options.content) + request_body_replays = [ + _RequestContentReplay(input_options.content), + *(_RequestContentReplay(content) for content in _iter_file_contents(input_options.files)), + ] retries_taken = 0 for retries_taken in range(max_retries + 1): @@ -1755,7 +1777,7 @@ async def request( except timeout_exceptions() as err: log.debug("Encountered a timeout exception: %s", type(err).__name__) - if remaining_retries > 0 and content_replay.rewind(): + if remaining_retries > 0 and all(replay.rewind() for replay in request_body_replays): await self._sleep_for_retry( retries_taken=retries_taken, max_retries=max_retries, @@ -1772,7 +1794,7 @@ async def request( except Exception as err: log.debug("Encountered exception: %s", type(err).__name__) - if remaining_retries > 0 and content_replay.rewind(): + if remaining_retries > 0 and all(replay.rewind() for replay in request_body_replays): await self._sleep_for_retry( retries_taken=retries_taken, max_retries=max_retries, @@ -1796,7 +1818,11 @@ async def request( except status_exceptions() as err: # thrown on 4xx and 5xx status code log.debug("Encountered an HTTP status error: %i", response.status_code) - if remaining_retries > 0 and self._should_retry(err.response) and content_replay.rewind(): + if ( + remaining_retries > 0 + and self._should_retry(err.response) + and all(replay.rewind() for replay in request_body_replays) + ): await err.response.aclose() await self._sleep_for_retry( retries_taken=retries_taken, diff --git a/tests/test_client.py b/tests/test_client.py index 0c70ed8796..ee840d0a35 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -132,6 +132,11 @@ def __aiter__(self) -> AsyncIterator[T]: return self._iterator +class _NonSeekableBytesIO(io.BytesIO): + def seekable(self) -> bool: + return False + + def _get_open_connections(client: OpenAI | AsyncOpenAI) -> int: transport = client._client._transport if isinstance(transport, httpx2.HTTPTransport) or isinstance(transport, httpx2.AsyncHTTPTransport): @@ -894,6 +899,54 @@ def mock_handler(request: httpx2.Request) -> httpx2.Response: assert response.status_code == 200 assert request_bodies == [file_content, file_content] + def test_multipart_retry_does_not_reuse_non_seekable_file(self) -> None: + file_content = b"Hello, this multipart file must not be replayed." + request_bodies: list[bytes] = [] + + def mock_handler(request: httpx2.Request) -> httpx2.Response: + request_bodies.append(request.read()) + return httpx2.Response(500, json={"error": {}}) + + with OpenAI( + base_url=base_url, + api_key=api_key, + max_retries=1, + http_client=httpx2.Client(transport=MockTransport(handler=mock_handler)), + ) as client: + with pytest.raises(APIStatusError): + client.post( + "/upload", + files={"file": ("upload.txt", _NonSeekableBytesIO(file_content), "text/plain")}, + cast_to=httpx2.Response, + ) + + assert len(request_bodies) == 1 + assert file_content in request_bodies[0] + + def test_multipart_retry_rewinds_seekable_file(self) -> None: + file_content = b"Hello, this multipart file can be replayed." + request_bodies: list[bytes] = [] + + def mock_handler(request: httpx2.Request) -> httpx2.Response: + request_bodies.append(request.read()) + return httpx2.Response(500 if len(request_bodies) == 1 else 200) + + with OpenAI( + base_url=base_url, + api_key=api_key, + max_retries=1, + http_client=httpx2.Client(transport=MockTransport(handler=mock_handler)), + ) as client: + response = client.post( + "/upload", + files={"file": ("upload.txt", io.BytesIO(file_content), "text/plain")}, + cast_to=httpx2.Response, + ) + + assert response.status_code == 200 + assert len(request_bodies) == 2 + assert all(file_content in body for body in request_bodies) + @mock.patch("openai._base_client.BaseClient._calculate_retry_timeout", _low_retry_timeout) def test_binary_content_retry_rewinds_seekable_stream(self) -> None: file_content = b"Hello, this is a test file." @@ -2263,6 +2316,54 @@ async def mock_handler(request: httpx2.Request) -> httpx2.Response: assert request_bodies == [file_content] + async def test_multipart_retry_does_not_reuse_non_seekable_file(self) -> None: + file_content = b"Hello, this multipart file must not be replayed." + request_bodies: list[bytes] = [] + + async def mock_handler(request: httpx2.Request) -> httpx2.Response: + request_bodies.append(await request.aread()) + return httpx2.Response(500, json={"error": {}}) + + async with AsyncOpenAI( + base_url=base_url, + api_key=api_key, + max_retries=1, + http_client=httpx2.AsyncClient(transport=MockTransport(handler=mock_handler)), + ) as client: + with pytest.raises(APIStatusError): + await client.post( + "/upload", + files={"file": ("upload.txt", _NonSeekableBytesIO(file_content), "text/plain")}, + cast_to=httpx2.Response, + ) + + assert len(request_bodies) == 1 + assert file_content in request_bodies[0] + + async def test_multipart_retry_rewinds_seekable_file(self) -> None: + file_content = b"Hello, this multipart file can be replayed." + request_bodies: list[bytes] = [] + + async def mock_handler(request: httpx2.Request) -> httpx2.Response: + request_bodies.append(await request.aread()) + return httpx2.Response(500 if len(request_bodies) == 1 else 200) + + async with AsyncOpenAI( + base_url=base_url, + api_key=api_key, + max_retries=1, + http_client=httpx2.AsyncClient(transport=MockTransport(handler=mock_handler)), + ) as client: + response = await client.post( + "/upload", + files={"file": ("upload.txt", io.BytesIO(file_content), "text/plain")}, + cast_to=httpx2.Response, + ) + + assert response.status_code == 200 + assert len(request_bodies) == 2 + assert all(file_content in body for body in request_bodies) + @pytest.mark.respx2(base_url=base_url) async def test_binary_content_upload_with_body_is_deprecated( self, respx2_mock: MockRouter, async_client: AsyncOpenAI From 50fb12243260d095d5bcfa4c64a3c9387f860665 Mon Sep 17 00:00:00 2001 From: HughhhhCoder Date: Wed, 16 Sep 2026 14:05:38 +0800 Subject: [PATCH 4/4] fix(client): track prepared request bodies --- src/openai/_base_client.py | 16 ++++++------ tests/test_client.py | 50 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 58 insertions(+), 8 deletions(-) diff --git a/src/openai/_base_client.py b/src/openai/_base_client.py index 25ff9340be..3a48945800 100644 --- a/src/openai/_base_client.py +++ b/src/openai/_base_client.py @@ -1118,15 +1118,15 @@ def request( response: httpx2.Response | None = None max_retries = input_options.get_max_retries(self.max_retries) self._validate_max_retries(max_retries) - request_body_replays = [ - _RequestContentReplay(input_options.content), - *(_RequestContentReplay(content) for content in _iter_file_contents(input_options.files)), - ] retries_taken = 0 for retries_taken in range(max_retries + 1): options = model_copy(input_options) options = self._prepare_options(options) + request_body_replays = [ + _RequestContentReplay(options.content), + *(_RequestContentReplay(content) for content in _iter_file_contents(options.files)), + ] remaining_retries = max_retries - retries_taken request = self._build_request(options, retries_taken=retries_taken) @@ -1751,15 +1751,15 @@ async def request( response: httpx2.Response | None = None max_retries = input_options.get_max_retries(self.max_retries) self._validate_max_retries(max_retries) - request_body_replays = [ - _RequestContentReplay(input_options.content), - *(_RequestContentReplay(content) for content in _iter_file_contents(input_options.files)), - ] retries_taken = 0 for retries_taken in range(max_retries + 1): options = model_copy(input_options) options = await self._prepare_options(options) + request_body_replays = [ + _RequestContentReplay(options.content), + *(_RequestContentReplay(content) for content in _iter_file_contents(options.files)), + ] remaining_retries = max_retries - retries_taken request = self._build_request(options, retries_taken=retries_taken) diff --git a/tests/test_client.py b/tests/test_client.py index b1ea6b2068..576e72b3eb 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -935,6 +935,31 @@ def mock_handler(request: httpx2.Request) -> httpx2.Response: assert response.status_code == 200 assert request_bodies == [file_content, file_content] + def test_binary_content_retry_checks_prepared_options(self) -> None: + file_content = b"Hello, this prepared body must not be replayed." + prepared_content = _OneShotIterable([file_content]) + request_bodies: list[bytes] = [] + + def prepare_options(options: FinalRequestOptions) -> FinalRequestOptions: + options.content = prepared_content + return options + + def mock_handler(request: httpx2.Request) -> httpx2.Response: + request_bodies.append(request.read()) + return httpx2.Response(500 if len(request_bodies) == 1 else 200, json={"error": {}}) + + with OpenAI( + base_url=base_url, + api_key=api_key, + max_retries=1, + http_client=httpx2.Client(transport=MockTransport(handler=mock_handler)), + ) as client: + with mock.patch.object(client, "_prepare_options", side_effect=prepare_options): + with pytest.raises(APIStatusError): + client.post("/upload", content=file_content, cast_to=httpx2.Response) + + assert request_bodies == [file_content] + def test_multipart_retry_does_not_reuse_non_seekable_file(self) -> None: file_content = b"Hello, this multipart file must not be replayed." request_bodies: list[bytes] = [] @@ -2405,6 +2430,31 @@ async def mock_handler(request: httpx2.Request) -> httpx2.Response: assert request_bodies == [file_content] + async def test_binary_content_retry_checks_prepared_options(self) -> None: + file_content = b"Hello, this prepared body must not be replayed." + prepared_content = _OneShotAsyncIterable([file_content]) + request_bodies: list[bytes] = [] + + async def prepare_options(options: FinalRequestOptions) -> FinalRequestOptions: + options.content = prepared_content + return options + + async def mock_handler(request: httpx2.Request) -> httpx2.Response: + request_bodies.append(await request.aread()) + return httpx2.Response(500 if len(request_bodies) == 1 else 200, json={"error": {}}) + + async with AsyncOpenAI( + base_url=base_url, + api_key=api_key, + max_retries=1, + http_client=httpx2.AsyncClient(transport=MockTransport(handler=mock_handler)), + ) as client: + with mock.patch.object(client, "_prepare_options", side_effect=prepare_options): + with pytest.raises(APIStatusError): + await client.post("/upload", content=file_content, cast_to=httpx2.Response) + + assert request_bodies == [file_content] + async def test_multipart_retry_does_not_reuse_non_seekable_file(self) -> None: file_content = b"Hello, this multipart file must not be replayed." request_bodies: list[bytes] = []