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
89 changes: 83 additions & 6 deletions src/openai/_base_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]")
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand Down
Loading