From 52d3c9dcfc11d16a914f34636b9733dec324ae60 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 14 Jan 2026 14:55:17 +0530 Subject: [PATCH] Add retry policy support to responses API --- litellm/main.py | 65 +++++++++++++ litellm/utils.py | 55 +++++++++++ .../test_completion_with_retries.py | 93 +++++++++++++++++++ 3 files changed, 213 insertions(+) diff --git a/litellm/main.py b/litellm/main.py index a1992752bdb..c1c4efd943f 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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: diff --git a/litellm/utils.py b/litellm/utils.py index ea29d22221d..3cf300802aa 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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 diff --git a/tests/local_testing/test_completion_with_retries.py b/tests/local_testing/test_completion_with_retries.py index 09dacdb651a..6eb3ad460e6 100644 --- a/tests/local_testing/test_completion_with_retries.py +++ b/tests/local_testing/test_completion_with_retries.py @@ -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