Merge pull request #19074 from BerriAI/litellm_19046-bug-retry-policies-are-not-applied-on-responses-calls

Add retry policy support to responses API
This commit is contained in:
Sameer Kankute 2026-01-14 17:56:11 +05:30 • committed by GitHub
commit 4aadc0d41f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 213 additions and 0 deletions

View file

@ -4227,6 +4227,71 @@ async def acompletion_with_retries(*args, **kwargs):
return await retryer(original_function, *args, **kwargs)
def responses_with_retries(*args, **kwargs):
"""
Executes a litellm.responses() with retries
"""
try:
import tenacity
except Exception as e:
raise Exception(
f"tenacity import failed please run `pip install tenacity`. Error{e}"
)
from litellm.responses.main import responses
num_retries = kwargs.pop("num_retries", 3)
# reset retries in .responses()
kwargs["max_retries"] = 0
kwargs["num_retries"] = 0
retry_strategy: Literal["exponential_backoff_retry", "constant_retry"] = kwargs.pop(
"retry_strategy", "constant_retry"
) # type: ignore
original_function = kwargs.pop("original_function", responses)
if retry_strategy == "exponential_backoff_retry":
retryer = tenacity.Retrying(
wait=tenacity.wait_exponential(multiplier=1, max=10),
stop=tenacity.stop_after_attempt(num_retries),
reraise=True,
)
else:
retryer = tenacity.Retrying(
stop=tenacity.stop_after_attempt(num_retries), reraise=True
)
return retryer(original_function, *args, **kwargs)
async def aresponses_with_retries(*args, **kwargs):
"""
Executes a litellm.aresponses() with retries
"""
try:
import tenacity
except Exception as e:
raise Exception(
f"tenacity import failed please run `pip install tenacity`. Error{e}"
)
from litellm.responses.main import aresponses
num_retries = kwargs.pop("num_retries", 3)
kwargs["max_retries"] = 0
kwargs["num_retries"] = 0
retry_strategy = kwargs.pop("retry_strategy", "constant_retry")
original_function = kwargs.pop("original_function", aresponses)
if retry_strategy == "exponential_backoff_retry":
retryer = tenacity.AsyncRetrying(
wait=tenacity.wait_exponential(multiplier=1, max=10),
stop=tenacity.stop_after_attempt(num_retries),
reraise=True,
)
else:
retryer = tenacity.AsyncRetrying(
stop=tenacity.stop_after_attempt(num_retries), reraise=True
)
return await retryer(original_function, *args, **kwargs)
### EMBEDDING ENDPOINTS ####################
@client
async def aembedding(*args, **kwargs) -> EmbeddingResponse:

View file

@ -1640,6 +1640,37 @@ def client(original_function): # noqa: PLR0915
else:
kwargs["model"] = context_window_fallback_dict[model]
return original_function(*args, **kwargs)
elif call_type == CallTypes.responses.value:
num_retries = (
kwargs.get("num_retries", None) or litellm.num_retries or None
)
if kwargs.get("retry_policy", None):
get_num_retries_from_retry_policy = getattr(sys.modules[__name__], 'get_num_retries_from_retry_policy')
reset_retry_policy = getattr(sys.modules[__name__], 'reset_retry_policy')
num_retries = get_num_retries_from_retry_policy(
exception=e,
retry_policy=kwargs.get("retry_policy"),
)
kwargs["retry_policy"] = (
reset_retry_policy()
) # prevent infinite loops
litellm.num_retries = (
None # set retries to None to prevent infinite loops
)
_is_litellm_router_call = "model_group" in kwargs.get(
"metadata", {}
) # check if call from litellm.router/proxy
if (
num_retries and not _is_litellm_router_call
): # only enter this if call is not from litellm router/proxy. router has it's own logic for retrying
if (
isinstance(e, openai.APIError)
or isinstance(e, openai.Timeout)
or isinstance(e, openai.APIConnectionError)
):
kwargs["num_retries"] = num_retries
return litellm.responses_with_retries(*args, **kwargs)
traceback_exception = traceback.format_exc()
end_time = datetime.datetime.now()
@ -1899,12 +1930,36 @@ def client(original_function): # noqa: PLR0915
isinstance(e, litellm.exceptions.ContextWindowExceededError)
and context_window_fallback_dict
and model in context_window_fallback_dict
and not _is_litellm_router_call
):
if len(args) > 0:
args[0] = context_window_fallback_dict[model] # type: ignore
else:
kwargs["model"] = context_window_fallback_dict[model]
return await original_function(*args, **kwargs)
elif call_type == CallTypes.aresponses.value:
_is_litellm_router_call = "model_group" in kwargs.get(
"metadata", {}
) # check if call from litellm.router/proxy
if (
num_retries and not _is_litellm_router_call
): # only enter this if call is not from litellm router/proxy. router has it's own logic for retrying
try:
litellm.num_retries = (
None # set retries to None to prevent infinite loops
)
kwargs["num_retries"] = num_retries
kwargs["original_function"] = original_function
if isinstance(
e, openai.RateLimitError
): # rate limiting specific error
kwargs["retry_strategy"] = "exponential_backoff_retry"
elif isinstance(e, openai.APIError): # generic api error
kwargs["retry_strategy"] = "constant_retry"
return await litellm.aresponses_with_retries(*args, **kwargs)
except Exception:
pass
setattr(
e, "num_retries", num_retries

View file

@ -12,6 +12,8 @@ import pytest
import openai
import litellm
from litellm import completion_with_retries, completion, acompletion_with_retries
from litellm import responses_with_retries, aresponses_with_retries
from litellm.responses.main import responses, aresponses
from litellm import (
AuthenticationError,
BadRequestError,
@ -146,3 +148,94 @@ async def test_completion_with_retries(sync_mode):
mock_completion.assert_called_once()
assert mock_completion.call_args.kwargs["num_retries"] == 0
assert mock_completion.call_args.kwargs["max_retries"] == 0
# ==================== Responses API Retry Tests ====================
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_responses_with_retries(sync_mode):
"""
Test that responses() and aresponses() properly handle num_retries parameter.
If responses_with_retries is called with num_retries=3, and max_retries=0,
then litellm.responses should receive num_retries=0, max_retries=0
"""
from unittest.mock import patch, MagicMock, AsyncMock
if sync_mode:
target_function = "responses"
retry_function = responses_with_retries
else:
target_function = "aresponses"
retry_function = aresponses_with_retries
# Mock the responses/aresponses function
with patch("litellm.responses.main.responses" if sync_mode else "litellm.responses.main.aresponses") as mock_responses:
if sync_mode:
mock_responses.return_value = MagicMock()
retry_function(
model="gpt-4o",
input="Hello, what's the weather?",
num_retries=3,
original_function=mock_responses,
)
else:
mock_responses.return_value = AsyncMock()
await retry_function(
model="gpt-4o",
input="Hello, what's the weather?",
num_retries=3,
original_function=mock_responses,
)
mock_responses.assert_called_once()
assert mock_responses.call_args.kwargs["num_retries"] == 0
assert mock_responses.call_args.kwargs["max_retries"] == 0
@pytest.mark.asyncio
@pytest.mark.parametrize("sync_mode", [True, False])
async def test_responses_retry_on_auth_error(sync_mode):
"""
Test that responses API actually retries when encountering authentication errors.
This validates that the @client decorator properly handles responses/aresponses retries.
"""
from unittest.mock import patch
import openai
num_retries = 2
# Mock the responses/aresponses to raise an authentication error
if sync_mode:
with patch.object(litellm, "responses_with_retries") as mock_retry:
mock_retry.return_value = None
try:
responses(
model="gpt-4o",
input="Test input",
num_retries=num_retries,
api_key="sk-invalid-key-12345",
)
except Exception:
pass # Expected to fail with invalid key
# Check if retry function was called (means @client decorator triggered retry)
if mock_retry.called:
assert mock_retry.call_args.kwargs.get("num_retries") == num_retries
else:
with patch.object(litellm, "aresponses_with_retries") as mock_retry:
mock_retry.return_value = None
try:
await aresponses(
model="gpt-4o",
input="Test input",
num_retries=num_retries,
api_key="sk-invalid-key-12345",
)
except Exception:
pass # Expected to fail with invalid key
# Check if retry function was called (means @client decorator triggered retry)
if mock_retry.called:
assert mock_retry.call_args.kwargs.get("num_retries") == num_retries