Skip to content

Commit ee237cb

Browse files
feat(client): support unlimited retries and configurable backoff
1 parent 5e33cc6 commit ee237cb

11 files changed

Lines changed: 342 additions & 67 deletions

File tree

‎README.md‎

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -776,6 +776,26 @@ client.with_options(max_retries=5).chat.completions.create(
776776
)
777777
```
778778

779+
`max_retries` accepts non-negative integers, or `math.inf` to keep retrying eligible
780+
failures until success or cancellation. `0` disables retries. Other values raise
781+
an error before a request is sent.
782+
783+
To customize exponential backoff, set `backoff_factor` (the initial delay, default
784+
`0.5` seconds) and `max_backoff` (the delay cap, default `8.0` seconds). Both must
785+
be finite, non-negative numbers. Delays are reduced by up to 25% jitter. A valid
786+
server `Retry-After` takes precedence over these settings; server delays above two
787+
minutes are surfaced without retrying.
788+
789+
```python
790+
import math
791+
792+
client = OpenAI(max_retries=math.inf, backoff_factor=1.0, max_backoff=30.0)
793+
# The same options can be set with client.with_options(...).
794+
```
795+
796+
Transport failures are retried. Application exceptions raised by custom transports
797+
or hooks propagate unchanged, including task-executor cancellation signals.
798+
779799
## Timeouts
780800

781801
By default requests time out after 10 minutes. You can configure this with a `timeout` option,

‎src/openai/__init__.py‎

Lines changed: 34 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -134,6 +134,7 @@
134134

135135
import httpx2 as _httpx
136136

137+
from ._constants import MAX_RETRY_DELAY, INITIAL_RETRY_DELAY
137138
from ._base_client import DEFAULT_TIMEOUT, DEFAULT_MAX_RETRIES
138139

139140
api_key: str | None = None
@@ -150,7 +151,11 @@
150151

151152
timeout: float | Timeout | None = DEFAULT_TIMEOUT
152153

153-
max_retries: int = DEFAULT_MAX_RETRIES
154+
max_retries: int | float = DEFAULT_MAX_RETRIES
155+
156+
backoff_factor: float = INITIAL_RETRY_DELAY
157+
158+
max_backoff: float = MAX_RETRY_DELAY
154159

155160
default_headers: _t.Mapping[str, str] | None = None
156161

@@ -260,15 +265,35 @@ def timeout(self, value: float | Timeout | None) -> None: # type: ignore
260265

261266
@property # type: ignore
262267
@override
263-
def max_retries(self) -> int:
268+
def max_retries(self) -> int | float:
264269
return max_retries
265270

266271
@max_retries.setter # type: ignore
267-
def max_retries(self, value: int) -> None: # type: ignore
272+
def max_retries(self, value: int | float) -> None: # type: ignore
268273
global max_retries
269274

270275
max_retries = value
271276

277+
@property # type: ignore
278+
@override
279+
def backoff_factor(self) -> float:
280+
return backoff_factor
281+
282+
@backoff_factor.setter # type: ignore
283+
def backoff_factor(self, value: float) -> None: # type: ignore
284+
global backoff_factor
285+
backoff_factor = value
286+
287+
@property # type: ignore
288+
@override
289+
def max_backoff(self) -> float:
290+
return max_backoff
291+
292+
@max_backoff.setter # type: ignore
293+
def max_backoff(self, value: float) -> None: # type: ignore
294+
global max_backoff
295+
max_backoff = value
296+
272297
@property # type: ignore
273298
@override
274299
def _custom_headers(self) -> _t.Mapping[str, str] | None:
@@ -395,6 +420,8 @@ def _load_client() -> OpenAI: # type: ignore[reportUnusedFunction]
395420
base_url=base_url,
396421
timeout=timeout,
397422
max_retries=max_retries,
423+
backoff_factor=backoff_factor,
424+
max_backoff=max_backoff,
398425
default_headers=default_headers,
399426
default_query=default_query,
400427
http_client=http_client,
@@ -411,6 +438,8 @@ def _load_client() -> OpenAI: # type: ignore[reportUnusedFunction]
411438
base_url=base_url,
412439
timeout=timeout,
413440
max_retries=max_retries,
441+
backoff_factor=backoff_factor,
442+
max_backoff=max_backoff,
414443
default_headers=default_headers,
415444
default_query=default_query,
416445
http_client=http_client,
@@ -426,6 +455,8 @@ def _load_client() -> OpenAI: # type: ignore[reportUnusedFunction]
426455
base_url=base_url,
427456
timeout=timeout,
428457
max_retries=max_retries,
458+
backoff_factor=backoff_factor,
459+
max_backoff=max_backoff,
429460
default_headers=default_headers,
430461
default_query=default_query,
431462
http_client=http_client,

‎src/openai/_base_client.py‎

Lines changed: 44 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@
3131
cast,
3232
overload,
3333
)
34+
from itertools import count
3435
from typing_extensions import Unpack, Literal, override, get_origin
3536

3637
import anyio
@@ -385,7 +386,7 @@ class BaseClient(Generic[_HttpxClientT, _DefaultStreamT]):
385386
_client: _HttpxClientT
386387
_version: str
387388
_base_url: URL
388-
max_retries: int
389+
max_retries: int | float
389390
timeout: Union[float, Timeout, None]
390391
_strict_response_validation: bool
391392
_idempotency_header: str | None
@@ -397,25 +398,26 @@ def __init__(
397398
version: str,
398399
base_url: str | URL,
399400
_strict_response_validation: bool,
400-
max_retries: int = DEFAULT_MAX_RETRIES,
401+
max_retries: int | float = DEFAULT_MAX_RETRIES,
402+
backoff_factor: float = INITIAL_RETRY_DELAY,
403+
max_backoff: float = MAX_RETRY_DELAY,
401404
timeout: float | Timeout | None = DEFAULT_TIMEOUT,
402405
custom_headers: Mapping[str, str] | None = None,
403406
custom_query: Mapping[str, object] | None = None,
404407
) -> None:
405408
self._version = version
406409
self._base_url = self._enforce_trailing_slash(normalize_httpx_url(base_url))
407410
self.max_retries = max_retries
411+
self.backoff_factor = backoff_factor
412+
self.max_backoff = max_backoff
408413
self.timeout = timeout
409414
self._custom_headers = custom_headers or {}
410415
self._custom_query = custom_query or {}
411416
self._strict_response_validation = _strict_response_validation
412417
self._idempotency_header = None
413418
self._platform: Platform | None = None
414419

415-
if max_retries is None: # pyright: ignore[reportUnnecessaryComparison]
416-
raise TypeError(
417-
"max_retries cannot be None. If you want to disable retries, pass `0`; if you want unlimited retries, pass `math.inf` or a very high number; if you want the default behavior, pass `openai.DEFAULT_MAX_RETRIES`"
418-
)
420+
self._validate_retry_options(max_retries)
419421

420422
def _enforce_trailing_slash(self, url: URL) -> URL:
421423
if url.raw_path.endswith(b"/"):
@@ -789,24 +791,32 @@ def _parse_retry_after_header(self, response_headers: Optional[httpx2.Headers] =
789791

790792
return float(retry_date - time.time())
791793

794+
def _validate_retry_options(self, value: object) -> None:
795+
if value is None:
796+
raise TypeError("max_retries cannot be None. Use 0 to disable retries or math.inf for unlimited retries.")
797+
if not isinstance(value, (int, float)):
798+
raise TypeError("max_retries must be a non-negative integer or math.inf")
799+
if not (isinstance(value, int) and value >= 0) and value != math.inf:
800+
raise ValueError("max_retries must be a non-negative integer or math.inf")
801+
for name, value in (("backoff_factor", self.backoff_factor), ("max_backoff", self.max_backoff)):
802+
if not math.isfinite(value) or value < 0:
803+
raise ValueError(f"{name} must be a finite, non-negative number")
804+
792805
def _calculate_retry_timeout(
793806
self,
794-
remaining_retries: int,
795-
options: FinalRequestOptions,
807+
retries_taken: int,
796808
response_headers: Optional[httpx2.Headers] = None,
797809
) -> float:
798-
max_retries = options.get_max_retries(self.max_retries)
799-
800810
# Honor server-directed delays up to two minutes.
801811
retry_after = self._parse_retry_after_header(response_headers)
802812
if retry_after is not None and math.isfinite(retry_after) and 0 < retry_after <= MAX_RETRY_AFTER_DELAY:
803813
return retry_after
804814

805815
# Also cap retry count to 1000 to avoid any potential overflows with `pow`
806-
nb_retries = min(max_retries - remaining_retries, 1000)
816+
nb_retries = min(retries_taken, 1000)
807817

808818
# Apply exponential backoff, but not more than the max.
809-
sleep_seconds = min(INITIAL_RETRY_DELAY * pow(2.0, nb_retries), MAX_RETRY_DELAY)
819+
sleep_seconds = min(self.backoff_factor * pow(2.0, nb_retries), self.max_backoff)
810820

811821
# Reduce the calculated timeout by a random range between 0-25%
812822
jitter = 1 - 0.25 * random()
@@ -901,7 +911,9 @@ def __init__(
901911
*,
902912
version: str,
903913
base_url: str | URL,
904-
max_retries: int = DEFAULT_MAX_RETRIES,
914+
max_retries: int | float = DEFAULT_MAX_RETRIES,
915+
backoff_factor: float = INITIAL_RETRY_DELAY,
916+
max_backoff: float = MAX_RETRY_DELAY,
905917
timeout: float | Timeout | None | NotGiven = not_given,
906918
http_client: httpx2.Client | None = None,
907919
custom_headers: Mapping[str, str] | None = None,
@@ -938,6 +950,8 @@ def __init__(
938950
timeout=cast(Timeout, timeout),
939951
base_url=base_url,
940952
max_retries=max_retries,
953+
backoff_factor=backoff_factor,
954+
max_backoff=max_backoff,
941955
custom_query=custom_query,
942956
custom_headers=custom_headers,
943957
_strict_response_validation=_strict_response_validation,
@@ -1048,9 +1062,10 @@ def request(
10481062

10491063
response: httpx2.Response | None = None
10501064
max_retries = input_options.get_max_retries(self.max_retries)
1065+
self._validate_retry_options(max_retries)
10511066

10521067
retries_taken = 0
1053-
for retries_taken in range(max_retries + 1):
1068+
for retries_taken in count():
10541069
options = model_copy(input_options)
10551070
options = self._prepare_options(options)
10561071

@@ -1086,7 +1101,6 @@ def request(
10861101
self._sleep_for_retry(
10871102
retries_taken=retries_taken,
10881103
max_retries=max_retries,
1089-
options=input_options,
10901104
response=None,
10911105
)
10921106
continue
@@ -1103,7 +1117,6 @@ def request(
11031117
self._sleep_for_retry(
11041118
retries_taken=retries_taken,
11051119
max_retries=max_retries,
1106-
options=input_options,
11071120
response=None,
11081121
)
11091122
continue
@@ -1128,7 +1141,6 @@ def request(
11281141
self._sleep_for_retry(
11291142
retries_taken=retries_taken,
11301143
max_retries=max_retries,
1131-
options=input_options,
11321144
response=response,
11331145
)
11341146
continue
@@ -1154,16 +1166,16 @@ def request(
11541166
)
11551167

11561168
def _sleep_for_retry(
1157-
self, *, retries_taken: int, max_retries: int, options: FinalRequestOptions, response: httpx2.Response | None
1169+
self, *, retries_taken: int, max_retries: int | float, response: httpx2.Response | None
11581170
) -> None:
11591171
remaining_retries = max_retries - retries_taken
11601172
if remaining_retries == 1:
11611173
log.debug("1 retry left")
11621174
else:
1163-
log.debug("%i retries left", remaining_retries)
1175+
log.debug("%s retries left", remaining_retries)
11641176

1165-
timeout = self._calculate_retry_timeout(remaining_retries, options, response.headers if response else None)
1166-
log.info("Retrying request in %f seconds", timeout)
1177+
timeout = self._calculate_retry_timeout(retries_taken, response.headers if response else None)
1178+
log.info("Retrying request in %f seconds (retry %i of %s)", timeout, retries_taken + 1, max_retries)
11671179

11681180
time.sleep(timeout)
11691181

@@ -1524,7 +1536,9 @@ def __init__(
15241536
version: str,
15251537
base_url: str | URL,
15261538
_strict_response_validation: bool,
1527-
max_retries: int = DEFAULT_MAX_RETRIES,
1539+
max_retries: int | float = DEFAULT_MAX_RETRIES,
1540+
backoff_factor: float = INITIAL_RETRY_DELAY,
1541+
max_backoff: float = MAX_RETRY_DELAY,
15281542
timeout: float | Timeout | None | NotGiven = not_given,
15291543
http_client: httpx2.AsyncClient | None = None,
15301544
custom_headers: Mapping[str, str] | None = None,
@@ -1560,6 +1574,8 @@ def __init__(
15601574
# cast to a valid type because mypy doesn't understand our type narrowing
15611575
timeout=cast(Timeout, timeout),
15621576
max_retries=max_retries,
1577+
backoff_factor=backoff_factor,
1578+
max_backoff=max_backoff,
15631579
custom_query=custom_query,
15641580
custom_headers=custom_headers,
15651581
_strict_response_validation=_strict_response_validation,
@@ -1672,9 +1688,10 @@ async def request(
16721688

16731689
response: httpx2.Response | None = None
16741690
max_retries = input_options.get_max_retries(self.max_retries)
1691+
self._validate_retry_options(max_retries)
16751692

16761693
retries_taken = 0
1677-
for retries_taken in range(max_retries + 1):
1694+
for retries_taken in count():
16781695
options = model_copy(input_options)
16791696
options = await self._prepare_options(options)
16801697

@@ -1709,7 +1726,6 @@ async def request(
17091726
await self._sleep_for_retry(
17101727
retries_taken=retries_taken,
17111728
max_retries=max_retries,
1712-
options=input_options,
17131729
response=None,
17141730
)
17151731
continue
@@ -1726,7 +1742,6 @@ async def request(
17261742
await self._sleep_for_retry(
17271743
retries_taken=retries_taken,
17281744
max_retries=max_retries,
1729-
options=input_options,
17301745
response=None,
17311746
)
17321747
continue
@@ -1751,7 +1766,6 @@ async def request(
17511766
await self._sleep_for_retry(
17521767
retries_taken=retries_taken,
17531768
max_retries=max_retries,
1754-
options=input_options,
17551769
response=response,
17561770
)
17571771
continue
@@ -1777,16 +1791,16 @@ async def request(
17771791
)
17781792

17791793
async def _sleep_for_retry(
1780-
self, *, retries_taken: int, max_retries: int, options: FinalRequestOptions, response: httpx2.Response | None
1794+
self, *, retries_taken: int, max_retries: int | float, response: httpx2.Response | None
17811795
) -> None:
17821796
remaining_retries = max_retries - retries_taken
17831797
if remaining_retries == 1:
17841798
log.debug("1 retry left")
17851799
else:
1786-
log.debug("%i retries left", remaining_retries)
1800+
log.debug("%s retries left", remaining_retries)
17871801

1788-
timeout = self._calculate_retry_timeout(remaining_retries, options, response.headers if response else None)
1789-
log.info("Retrying request in %f seconds", timeout)
1802+
timeout = self._calculate_retry_timeout(retries_taken, response.headers if response else None)
1803+
log.info("Retrying request in %f seconds (retry %i of %s)", timeout, retries_taken + 1, max_retries)
17901804

17911805
await anyio.sleep(timeout)
17921806

0 commit comments

Comments
 (0)