mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
commit
4aadc0d41f
3 changed files with 213 additions and 0 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue