diff --git a/litellm/utils.py b/litellm/utils.py index a2dac2fcf8d..ff8dd321495 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1661,6 +1661,14 @@ def post_call_processing( raise e +def _is_litellm_router_call(kwargs: Mapping[str, object], *, is_async: bool) -> bool: + """Router completion uses metadata. Async generic calls retry with litellm_metadata; sync generic calls need SDK retries.""" + metadata_buckets: Final = ( + (kwargs.get("metadata"), kwargs.get("litellm_metadata")) if is_async else (kwargs.get("metadata"),) + ) + return any(isinstance(bucket, Mapping) and "model_group" in bucket for bucket in metadata_buckets) + + def client(original_function): from litellm.litellm_core_utils.core_helpers import max_retries_per_request_hit @@ -1903,11 +1911,9 @@ def client(original_function): litellm.num_retries = None # set retries to None to prevent infinite loops context_window_fallback_dict: Final = kwargs.get("context_window_fallback_dict", {}) - _is_litellm_router_call = "model_group" in ( - kwargs.get("metadata") or {} - ) # check if call from litellm.router/proxy + is_completion_litellm_router_call: Final = _is_litellm_router_call(kwargs, is_async=False) if ( - num_retries and not _is_litellm_router_call + num_retries and not is_completion_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) @@ -1920,7 +1926,7 @@ def client(original_function): isinstance(e, litellm.exceptions.ContextWindowExceededError) and context_window_fallback_dict and model in context_window_fallback_dict - and not _is_litellm_router_call + and not is_completion_litellm_router_call ): if len(args) > 0: args[0] = context_window_fallback_dict[model] @@ -1939,11 +1945,9 @@ def client(original_function): 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") or {} - ) # check if call from litellm.router/proxy + is_responses_litellm_router_call: Final = _is_litellm_router_call(kwargs, is_async=False) if ( - num_retries and not _is_litellm_router_call + num_retries and not is_responses_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) @@ -2218,12 +2222,10 @@ def client(original_function): if call_type == CallTypes.acompletion.value: context_window_fallback_dict: Final = kwargs.get("context_window_fallback_dict", {}) - _is_litellm_router_call = "model_group" in ( - kwargs.get("metadata") or {} - ) # check if call from litellm.router/proxy + is_acompletion_litellm_router_call: Final = _is_litellm_router_call(kwargs, is_async=True) if ( - num_retries and not _is_litellm_router_call + num_retries and not is_acompletion_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 @@ -2242,7 +2244,7 @@ def client(original_function): isinstance(e, litellm.exceptions.ContextWindowExceededError) and context_window_fallback_dict and model in context_window_fallback_dict - and not _is_litellm_router_call + and not is_acompletion_litellm_router_call ): if len(args) > 0: args[0] = context_window_fallback_dict[model] @@ -2251,12 +2253,10 @@ def client(original_function): result = await original_function(*args, **kwargs) return result elif call_type == CallTypes.aresponses.value: - _is_litellm_router_call = "model_group" in ( - kwargs.get("metadata") or {} - ) # check if call from litellm.router/proxy + is_aresponses_litellm_router_call: Final = _is_litellm_router_call(kwargs, is_async=True) if ( - num_retries and not _is_litellm_router_call + num_retries and not is_aresponses_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 diff --git a/tests/integration/observability/test_s3_v2_streaming_failure_dedupe.py b/tests/integration/observability/test_s3_v2_streaming_failure_dedupe.py index ee347592d37..003181c72d8 100644 --- a/tests/integration/observability/test_s3_v2_streaming_failure_dedupe.py +++ b/tests/integration/observability/test_s3_v2_streaming_failure_dedupe.py @@ -87,16 +87,18 @@ def _responses_stream_id(response: httpx.Response) -> str: return response_id -def _register_models(scenario: Scenario, upstream_url: str) -> tuple[str, str]: +def _register_models(scenario: Scenario, upstream_url: str, num_retries: int = 0) -> tuple[str, str]: openai_model: Final = scenario.model( model="openai/gpt-4o-mini", api_base=upstream_url + "/v1", api_key="synthetic-provider-key", + num_retries=num_retries, ) anthropic_model: Final = scenario.model( model="anthropic/claude-sonnet-4-5-20250929", api_base=upstream_url, api_key="synthetic-provider-key", + num_retries=num_retries, ) return openai_model, anthropic_model @@ -165,7 +167,7 @@ def test_retried_failure_uploads_one_s3_object( marker: Final = f"s3-a-{surface}-{uuid.uuid4().hex}" sink: Final = RecordingS3Sink() with wire_server(_failure_provider) as upstream, wire_server(sink.respond) as bucket: - config: Final = s3_config(tmp_path, bucket.url, {}, settings={"num_retries": 2}) + config: Final = s3_config(tmp_path, bucket.url, {}) with ( owned_proxy( gateway, @@ -176,7 +178,7 @@ def test_retried_failure_uploads_one_s3_object( ) as candidate, candidate.scenario() as scenario, ): - openai_model, anthropic_model = _register_models(scenario, upstream.url) + openai_model, anthropic_model = _register_models(scenario, upstream.url, num_retries=2) key: Final = scenario.key(models=[openai_model, anthropic_model]) response: Final = _surface_request( candidate, @@ -316,7 +318,7 @@ def test_failure_burst_through_sink_outage_lands_each_request_once( ) sink: Final = RecordingS3Sink() with wire_server(_failure_provider) as upstream, wire_server(sink.respond) as bucket: - config: Final = s3_config(tmp_path, bucket.url, {}, settings={"num_retries": 2}) + config: Final = s3_config(tmp_path, bucket.url, {}) with ( owned_proxy( gateway, @@ -327,7 +329,7 @@ def test_failure_burst_through_sink_outage_lands_each_request_once( ) as candidate, candidate.scenario() as scenario, ): - openai_model, anthropic_model = _register_models(scenario, upstream.url) + openai_model, anthropic_model = _register_models(scenario, upstream.url, num_retries=2) key: Final = scenario.key(models=[openai_model, anthropic_model]) sink.fail_until = float("inf") diff --git a/tests/integration/providers/test_bedrock_stream_timeout_wire.py b/tests/integration/providers/test_bedrock_stream_timeout_wire.py index e7a76052558..631232a0420 100644 --- a/tests/integration/providers/test_bedrock_stream_timeout_wire.py +++ b/tests/integration/providers/test_bedrock_stream_timeout_wire.py @@ -490,7 +490,7 @@ def test_c7_the_bedrock_passthrough_stream_is_relayed_verbatim(gateway: Gateway) def test_c8_the_bedrock_passthrough_stream_already_retries_at_the_deployment_timeout(gateway: Gateway) -> None: with _peer("stall") as wire, gateway.scenario() as scenario: - model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS, num_retries=2) with httpx.Client(base_url=_proxy_url(gateway), timeout=_RETRY_WINDOW, trust_env=False) as client: response: Final = client.post( f"/bedrock/model/{model}/converse-stream", json=_PASSTHROUGH_BODY, headers=_auth(gateway) diff --git a/tests/integration/providers/test_gemini_thinking_replay_wire.py b/tests/integration/providers/test_gemini_thinking_replay_wire.py index 571274d20e2..055aa2fee4e 100644 --- a/tests/integration/providers/test_gemini_thinking_replay_wire.py +++ b/tests/integration/providers/test_gemini_thinking_replay_wire.py @@ -12,6 +12,7 @@ import anthropic import httpx import openai import pytest +from _pytest.mark.structures import ParameterSet from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives.asymmetric import rsa from integration._support.client import Gateway, Scenario, eventually @@ -485,7 +486,7 @@ def _only_request(wire: Wire, provider: Provider, stream: bool) -> Mapping[str, return _JSON_OBJECT.validate_json(received[0].body) -def _happy_cells() -> tuple[pytest.ParameterSet, ...]: +def _happy_cells() -> tuple[ParameterSet, ...]: return tuple( pytest.param( provider, endpoint, stream, client, id=f"{provider}-{endpoint}-{'stream' if stream else 'sync'}-{client}" diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index a72aec07766..47f893d0697 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -69,6 +69,7 @@ from litellm.utils import ( TextCompletionStreamWrapper, _check_provider_match, _get_potential_model_names, + _is_litellm_router_call, _is_streaming_request, _run_success_deployment_hook_on_converted_chat_stream, _snapshot_exception_for_hook, @@ -4755,6 +4756,158 @@ async def test_wrapper_async_logs_converted_responses_stream_with_standard_loggi assert success_kwargs["stream"] is True +@pytest.mark.parametrize( + "kwargs, is_async, expected", + [ + ({"metadata": {"model_group": "g"}}, False, True), + ({"metadata": {"model_group": "g"}}, True, True), + ({"litellm_metadata": {"model_group": "g"}}, False, False), + ({"litellm_metadata": {"model_group": "g"}}, True, True), + ({}, False, False), + ({}, True, False), + ({"metadata": None}, False, False), + ({"metadata": None}, True, False), + ], +) +def test_is_litellm_router_call_is_async_aware( + kwargs: Mapping[str, object], is_async: bool, expected: bool +) -> None: + assert _is_litellm_router_call(kwargs, is_async=is_async) is expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", [False, True], ids=["non_streaming", "streaming"]) +async def test_router_aresponses_does_not_run_sdk_retries( + monkeypatch: pytest.MonkeyPatch, stream: bool +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + model_list: Final = [ + { + "model_name": "responses-retry", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "sk-test", + "api_base": "https://responses-retry.local/v1", + "num_retries": 2, + }, + } + ] + router: Final = litellm.Router( + model_list=model_list, num_retries=0, retry_after=0, disable_cooldowns=True + ) + + try: + with respx.mock(assert_all_called=True) as respx_mock: + upstream: Final = respx_mock.post("https://responses-retry.local/v1/responses").mock( + return_value=httpx.Response( + 503, + headers={"retry-after": "0"}, + json={"error": {"message": "model is down", "type": "server_error"}}, + ) + ) + with pytest.raises(litellm.ServiceUnavailableError): + await router.aresponses(model="responses-retry", input="hi", stream=stream) + + assert upstream.call_count == 3 + finally: + router.discard() + litellm.in_memory_llm_clients_cache.flush_cache() + + +def test_router_responses_keeps_sdk_retries_for_sync_router_call( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + model_list: Final = [ + { + "model_name": "responses-retry", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "sk-test", + "api_base": "https://responses-retry.local/v1", + "num_retries": 2, + }, + } + ] + router: Final = litellm.Router( + model_list=model_list, num_retries=0, retry_after=0, disable_cooldowns=True + ) + + try: + with respx.mock(assert_all_called=True) as respx_mock: + upstream: Final = respx_mock.post("https://responses-retry.local/v1/responses").mock( + return_value=httpx.Response( + 503, + headers={"retry-after": "0"}, + json={"error": {"message": "model is down", "type": "server_error"}}, + ) + ) + with pytest.raises(litellm.ServiceUnavailableError): + router.responses(model="responses-retry", input="hi") + + assert upstream.call_count == 3 + finally: + router.discard() + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.asyncio +async def test_aresponses_uses_sdk_retries_without_router(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + + try: + with respx.mock(assert_all_called=True) as respx_mock: + upstream: Final = respx_mock.post("https://responses-direct.local/v1/responses").mock( + return_value=httpx.Response( + 503, + headers={"retry-after": "0"}, + json={"error": {"message": "model is down", "type": "server_error"}}, + ) + ) + with pytest.raises(litellm.ServiceUnavailableError): + await litellm.aresponses( + model="openai/gpt-4o-mini", + input="hi", + api_base="https://responses-direct.local/v1", + api_key="k", + num_retries=2, + ) + + assert upstream.call_count == 3 + finally: + litellm.in_memory_llm_clients_cache.flush_cache() + + +def test_responses_uses_sdk_retries_without_router(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + + try: + with respx.mock(assert_all_called=True) as respx_mock: + upstream: Final = respx_mock.post("https://responses-sync.local/v1/responses").mock( + return_value=httpx.Response( + 503, + headers={"retry-after": "0"}, + json={"error": {"message": "model is down", "type": "server_error"}}, + ) + ) + with pytest.raises(litellm.ServiceUnavailableError): + litellm.responses( + model="openai/gpt-4o-mini", + input="hi", + api_base="https://responses-sync.local/v1", + api_key="k", + num_retries=2, + ) + + assert upstream.call_count == 3 + finally: + litellm.in_memory_llm_clients_cache.flush_cache() + + @pytest.mark.asyncio async def test_wrapper_async_replays_cached_converted_chat_stream_as_stream( monkeypatch: pytest.MonkeyPatch,