diff --git a/src/openai/_base_client.py b/src/openai/_base_client.py index f195d04816..4e9bfa43e2 100644 --- a/src/openai/_base_client.py +++ b/src/openai/_base_client.py @@ -110,6 +110,67 @@ 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 + + # 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: + 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 + + +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]") @@ -1047,6 +1108,10 @@ def request( response: httpx2.Response | None = None max_retries = input_options.get_max_retries(self.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): @@ -1081,7 +1146,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 all(replay.rewind() for replay in request_body_replays): self._sleep_for_retry( retries_taken=retries_taken, max_retries=max_retries, @@ -1098,7 +1163,7 @@ def request( except Exception as err: log.debug("Encountered exception: %s", type(err).__name__) - if remaining_retries > 0: + 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, @@ -1122,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): + 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, @@ -1671,6 +1740,10 @@ async def request( response: httpx2.Response | None = None max_retries = input_options.get_max_retries(self.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): @@ -1704,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: + 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, @@ -1721,7 +1794,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 all(replay.rewind() for replay in request_body_replays): await self._sleep_for_retry( retries_taken=retries_taken, max_retries=max_retries, @@ -1745,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): + 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 d82c39e616..ee840d0a35 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 @@ -8,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 @@ -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, @@ -113,6 +114,29 @@ 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 + + +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): @@ -807,6 +831,148 @@ 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_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] = [] + + 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=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] + + 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." + 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 +2275,95 @@ 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_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] = [] + + 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=content_factory([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] = [] + + 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