diff --git a/tests/test_litellm/router_utils/test_cooldown_cache.py b/tests/test_litellm/router_utils/test_cooldown_cache.py index b6338ca69a0..a48402684b4 100644 --- a/tests/test_litellm/router_utils/test_cooldown_cache.py +++ b/tests/test_litellm/router_utils/test_cooldown_cache.py @@ -95,9 +95,7 @@ class TestCooldownCacheExceptionMasking: assert "magical kingdom" not in masked_exception # Should preserve the error type information at the beginning (first 50 chars) - assert masked_exception.startswith( - "litellm.proxy.proxy_server._handle_llm_api_excepti" - ) + assert masked_exception.startswith("litellm.proxy.proxy_server._handle_llm_api_excepti") def test_exception_with_api_keys_masked(self, cooldown_cache): """Test that API keys in exceptions are properly masked""" @@ -120,9 +118,7 @@ class TestCooldownCacheExceptionMasking: masked_exception = cooldown_data["exception_received"] # Should mask the sensitive content while preserving structure - assert masked_exception.startswith( - "Authentication failed with api_key=sk-12345678" - ) + assert masked_exception.startswith("Authentication failed with api_key=sk-12345678") assert "*" in masked_exception assert len(masked_exception) == len(exception_with_key) @@ -180,9 +176,7 @@ class TestCooldownCacheExceptionMasking: # Should successfully convert exception to string assert isinstance(cooldown_data["exception_received"], str) - assert ( - str(exc) == cooldown_data["exception_received"] - ) # Short exceptions not masked + assert str(exc) == cooldown_data["exception_received"] # Short exceptions not masked def test_masking_preserves_error_debugging_info(self, cooldown_cache): """Test that masking preserves essential debugging information""" @@ -209,9 +203,7 @@ class TestCooldownCacheExceptionMasking: masked_exception = cooldown_data["exception_received"] # Should preserve error type and initial debugging info (first 50 chars) - assert masked_exception.startswith( - "RateLimitError: Rate limit exceeded for model gpt-" - ) + assert masked_exception.startswith("RateLimitError: Rate limit exceeded for model gpt-") # Should mask the prompt content assert "Write a comprehensive analysis" not in masked_exception @@ -258,6 +250,130 @@ class TestCooldownCacheExceptionMasking: assert masked == expected +class TestCooldownCacheTTLCorrection: + def _make_cooldown_cache(self) -> CooldownCache: + in_memory = InMemoryCache() + dual_cache = DualCache(in_memory_cache=in_memory) + return CooldownCache(cache=dual_cache, default_cooldown_time=60.0) + + def test_expired_entry_evicted_and_not_returned(self): + """ + An entry with timestamp+cooldown_time in the past must be evicted from + in-memory cache and excluded from the active cooldown list. + """ + cc = self._make_cooldown_cache() + model_id = "expired-deployment" + key = CooldownCache.get_cooldown_cache_key(model_id) + + expired_value: CooldownCacheValue = { + "exception_received": "Rate limit", + "status_code": "429", + "timestamp": time.time() - 120.0, + "cooldown_time": 60.0, + } + cc.cache.in_memory_cache.set_cache(key, expired_value, ttl=600) + + active = cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None) + + assert active == [], "Expired cooldown entry must not appear in active cooldowns" + assert cc.cache.in_memory_cache.get_cache(key) is None, "Expired entry must be evicted from in-memory cache" + + def test_active_entry_is_returned(self): + """ + An entry whose cooldown window has not elapsed must appear in the active list. + """ + cc = self._make_cooldown_cache() + model_id = "active-deployment" + key = CooldownCache.get_cooldown_cache_key(model_id) + + active_value: CooldownCacheValue = { + "exception_received": "Rate limit", + "status_code": "429", + "timestamp": time.time(), + "cooldown_time": 60.0, + } + cc.cache.in_memory_cache.set_cache(key, active_value, ttl=60) + + active = cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None) + + assert len(active) == 1 + assert active[0][0] == model_id + + def test_ttl_corrected_when_in_memory_expiry_far_exceeds_remaining(self): + """ + When DualCache backfills from Redis using the default 600s TTL, the in-memory + TTL must be corrected to min(remaining, 60) seconds. + """ + cc = self._make_cooldown_cache() + model_id = "backfilled-deployment" + key = CooldownCache.get_cooldown_cache_key(model_id) + + remaining = 30.0 + value: CooldownCacheValue = { + "exception_received": "Rate limit", + "status_code": "429", + "timestamp": time.time() - (60.0 - remaining), + "cooldown_time": 60.0, + } + cc.cache.in_memory_cache.set_cache(key, value, ttl=600) + + before_expiry = cc.cache.in_memory_cache.ttl_dict.get(key) + assert before_expiry is not None + + cc.get_active_cooldowns(model_ids=[model_id], parent_otel_span=None) + + after_expiry = cc.cache.in_memory_cache.ttl_dict.get(key) + assert after_expiry is not None + corrected_remaining = after_expiry - time.time() + assert corrected_remaining <= 60.0, "Corrected TTL must not exceed 60s" + assert corrected_remaining > 0, "Corrected TTL must be positive (cooldown still active)" + + @pytest.mark.asyncio + async def test_async_expired_entry_evicted(self): + """ + Async path must also evict expired entries. + """ + cc = self._make_cooldown_cache() + model_id = "async-expired" + key = CooldownCache.get_cooldown_cache_key(model_id) + + expired_value: CooldownCacheValue = { + "exception_received": "Rate limit", + "status_code": "429", + "timestamp": time.time() - 120.0, + "cooldown_time": 60.0, + } + cc.cache.in_memory_cache.set_cache(key, expired_value, ttl=600) + + active = await cc.async_get_active_cooldowns(model_ids=[model_id], parent_otel_span=None) + + assert active == [], "Expired entry must not appear in async active cooldowns" + assert cc.cache.in_memory_cache.get_cache(key) is None + + @pytest.mark.asyncio + async def test_async_active_entry_is_returned(self): + """ + Async counterpart of test_active_entry_is_returned: an entry whose cooldown + window has not elapsed must appear in the async active list too. + """ + cc = self._make_cooldown_cache() + model_id = "async-active-deployment" + key = CooldownCache.get_cooldown_cache_key(model_id) + + active_value: CooldownCacheValue = { + "exception_received": "Rate limit", + "status_code": "429", + "timestamp": time.time(), + "cooldown_time": 60.0, + } + cc.cache.in_memory_cache.set_cache(key, active_value, ttl=60) + + active = await cc.async_get_active_cooldowns(model_ids=[model_id], parent_otel_span=None) + + assert len(active) == 1 + assert active[0][0] == model_id + + class TestCorrectedActiveCooldown: def _make_cooldown_cache(self) -> CooldownCache: in_memory = InMemoryCache() diff --git a/tests/test_litellm/router_utils/test_fallback_event_handlers.py b/tests/test_litellm/router_utils/test_fallback_event_handlers.py index 03ecc64d8d6..68395737469 100644 --- a/tests/test_litellm/router_utils/test_fallback_event_handlers.py +++ b/tests/test_litellm/router_utils/test_fallback_event_handlers.py @@ -4,6 +4,8 @@ from unittest.mock import MagicMock, patch import httpx import pytest +import litellm +from litellm.router_utils.cooldown_handlers import mark_advisor_orchestration_failure from litellm.router_utils.fallback_event_handlers import ( AttemptedFallbackTargets, _trigger_cooldown_for_failed_deployment, @@ -147,311 +149,6 @@ async def test_run_async_fallback_skips_original_model_group(): assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 1 -def test_trigger_cooldown_calls_set_cooldown_when_deployment_id_present(): - router = MagicMock() - router.cooldown_time = 60 - router.get_model_info.return_value = None - - exc = RuntimeError("upstream error") - exc.status_code = 429 - exc.failed_deployment_id = "deployment-abc" - - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set: - _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc) - - mock_set.assert_called_once() - _, call_kwargs = mock_set.call_args - assert call_kwargs["deployment"] == "deployment-abc" - assert call_kwargs["exception_status"] == 429 - - -def test_trigger_cooldown_skips_when_no_deployment_id(): - router = MagicMock() - - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set: - _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=RuntimeError("err")) - - mock_set.assert_not_called() - - -def test_trigger_cooldown_does_not_trust_caller_supplied_metadata_bucket(): - """A metadata bucket can't reliably be told apart from a caller-supplied one - without knowing the call's function_name, so a client with permission to set - metadata must not be able to get an arbitrary deployment cooled down by - forging a deployment_model_name marker.""" - router = MagicMock() - router.cooldown_time = 60 - router.get_model_info.return_value = None - - exc = RuntimeError("err") - kwargs = {"metadata": {"model_info": {"id": "attacker-chosen-deployment"}, "deployment_model_name": "gpt-4"}} - - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set: - _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs=kwargs, exception=exc) - - mock_set.assert_not_called() - - -def test_trigger_cooldown_increments_failure_counter_before_cooldown_check(): - """The fallback path must feed the same per-minute failure counter the - primary path uses, or repeated fallback failures never accumulate toward - the default percent-fail-rate cooldown threshold.""" - router = MagicMock() - router.cooldown_time = 60 - router.get_model_info.return_value = None - - exc = RuntimeError("err") - exc.failed_deployment_id = "deployment-abc" - - with ( - patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set, - patch( - "litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute" - ) as mock_increment, - ): - _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc) - - mock_increment.assert_called_once_with(litellm_router_instance=router, deployment_id="deployment-abc") - mock_set.assert_called_once() - - -def test_trigger_cooldown_uses_deployment_cooldown_time_when_present(): - router = MagicMock() - router.cooldown_time = 60 - router.get_model_info.return_value = {"model_info": {"cooldown_time": 30}} - - exc = RuntimeError("upstream error") - exc.status_code = 429 - exc.failed_deployment_id = "deployment-abc" - - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set: - _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc) - - _, call_kwargs = mock_set.call_args - assert call_kwargs["time_to_cooldown"] == 30 - - -def test_trigger_cooldown_falls_back_to_litellm_params_cooldown_time(): - """cooldown_time has pre-existing litellm_params support on the primary - failure path, so it must still be honored here when model_info doesn't set - it, unlike the new allowed_fails/allowed_fails_policy fields.""" - router = MagicMock() - router.cooldown_time = 60 - router.get_model_info.return_value = {"litellm_params": {"cooldown_time": 30}} - - exc = RuntimeError("upstream error") - exc.status_code = 429 - exc.failed_deployment_id = "deployment-abc" - - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set: - _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc) - - _, call_kwargs = mock_set.call_args - assert call_kwargs["time_to_cooldown"] == 30 - - -def test_trigger_cooldown_uses_response_header_when_no_deployment_config(): - """Precedence must match Router.deployment_callback_on_failure's primary path: - deployment config, then the response's Retry-After header, then the router - default. Without this, the fallback path always skips straight to the router - default whenever no deployment-level cooldown_time is configured.""" - router = MagicMock() - router.cooldown_time = 60 - router.get_model_info.return_value = {"model_info": {}} - - exc = RuntimeError("upstream error") - exc.status_code = 429 - exc.failed_deployment_id = "deployment-abc" - exc.litellm_response_headers = httpx.Headers({"retry-after": "45"}) - - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set: - _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc) - - _, call_kwargs = mock_set.call_args - assert call_kwargs["time_to_cooldown"] == 45 - - -def test_trigger_cooldown_silently_catches_exceptions(): - router = MagicMock() - router.cooldown_time = 60 - router.get_model_info.return_value = None - - exc = RuntimeError("upstream error") - exc.failed_deployment_id = "deployment-abc" - - with patch( - "litellm.router_utils.fallback_event_handlers._set_cooldown_deployments", - side_effect=RuntimeError("cooldown error"), - ): - _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc) - - -def test_trigger_cooldown_skips_request_scoped_404_on_generic_api_call(): - """A generic API call (files/batches/threads/rerank/...) forwards a caller-supplied - resource id, so a 404 there means "that id doesn't exist", not "this deployment is - unhealthy". Without this guard, a single bad id would 404 every deployment in the - fallback chain and cool all of them down from one request.""" - router = MagicMock() - router.cooldown_time = 60 - router.get_model_info.return_value = None - - exc = RuntimeError("not found") - exc.status_code = 404 - exc.failed_deployment_id = "deployment-abc" - - with ( - patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set, - patch( - "litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute" - ) as mock_increment, - ): - _trigger_cooldown_for_failed_deployment( - litellm_router=router, - kwargs={"original_generic_function": MagicMock()}, - exception=exc, - ) - - mock_set.assert_not_called() - mock_increment.assert_not_called() - - -def test_trigger_cooldown_still_cools_down_404_outside_generic_api_call(): - """The request-scoped-404 guard is scoped to generic API calls only: a 404 on a - regular completion fallback (no original_generic_function in kwargs) must still - cool down the deployment as before.""" - router = MagicMock() - router.cooldown_time = 60 - router.get_model_info.return_value = None - - exc = RuntimeError("not found") - exc.status_code = 404 - exc.failed_deployment_id = "deployment-abc" - - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set: - _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc) - - mock_set.assert_called_once() - - -def test_trigger_cooldown_skips_client_side_timeout_408(): - """The proxy's x-litellm-timeout header lets a caller set an arbitrarily short - timeout, which litellm.Timeout reports as status 408 regardless of the - deployment's actual health. Without this guard, a caller could force a 408 on - every deployment in the fallback chain from a single request.""" - router = MagicMock() - router.cooldown_time = 60 - router.get_model_info.return_value = None - - exc = RuntimeError("timeout") - exc.status_code = 408 - exc.failed_deployment_id = "deployment-abc" - - with ( - patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set, - patch( - "litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute" - ) as mock_increment, - ): - _trigger_cooldown_for_failed_deployment( - litellm_router=router, - kwargs={"client_side_timeout": True}, - exception=exc, - ) - - mock_set.assert_not_called() - mock_increment.assert_not_called() - - -def test_trigger_cooldown_still_cools_down_408_without_client_side_timeout_flag(): - """The client-side-timeout guard is scoped to caller-supplied timeouts only: a - 408 that did not come from x-litellm-timeout (no client_side_timeout in kwargs) - must still cool down the deployment as before.""" - router = MagicMock() - router.cooldown_time = 60 - router.get_model_info.return_value = None - - exc = RuntimeError("timeout") - exc.status_code = 408 - exc.failed_deployment_id = "deployment-abc" - - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set: - _trigger_cooldown_for_failed_deployment(litellm_router=router, kwargs={}, exception=exc) - - mock_set.assert_called_once() - - -@pytest.mark.asyncio -async def test_run_async_fallback_triggers_cooldown_when_logging_obj_has_logged(): - router = MagicMock() - router.cooldown_time = 60 - router.get_model_info.return_value = None - router.log_retry = MagicMock(side_effect=lambda kwargs, e: kwargs) - - exc = RuntimeError("fallback failed") - exc.failed_deployment_id = "dep-xyz" - - async def _always_fail(*args, **kwargs): - raise exc - - router.async_function_with_fallbacks = _always_fail - - logging_obj = MagicMock() - logging_obj.model_call_details = {"has_logged_async_failure": True} - - kwargs = { - "litellm_logging_obj": logging_obj, - } - - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set: - with pytest.raises(RuntimeError): - await run_async_fallback( - litellm_router=router, - fallback_model_group=["fallback-model"], - original_model_group="primary-model", - original_exception=RuntimeError("original"), - max_fallbacks=3, - fallback_depth=0, - **kwargs, - ) - - mock_set.assert_called_once() - - -@pytest.mark.asyncio -async def test_run_async_fallback_skips_cooldown_when_logging_obj_not_logged(): - router = MagicMock() - router.log_retry = MagicMock(side_effect=lambda kwargs, e: kwargs) - - exc = RuntimeError("fallback failed") - exc.failed_deployment_id = "dep-xyz" - - async def _always_fail(*args, **kwargs): - raise exc - - router.async_function_with_fallbacks = _always_fail - - logging_obj = MagicMock() - logging_obj.model_call_details = {"has_logged_async_failure": False} - - kwargs = { - "litellm_logging_obj": logging_obj, - } - - with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set: - with pytest.raises(RuntimeError): - await run_async_fallback( - litellm_router=router, - fallback_model_group=["fallback-model"], - original_model_group="primary-model", - original_exception=RuntimeError("original"), - max_fallbacks=3, - fallback_depth=0, - **kwargs, - ) - - mock_set.assert_not_called() - - class AttemptRecordingRouter: def __init__(self): self.attempted_model_groups = [] @@ -804,3 +501,300 @@ def test_get_fallback_model_group_does_not_mutate_fallbacks(): assert fallback_model_group == ["gpt-4o-mini"] assert fallbacks == [{"gpt-3.5-turbo": ["claude-3-haiku"]}, "gpt-4o-mini"] + + +class TestTriggerCooldownForFailedDeployment: + def test_calls_set_cooldown_deployments_with_stamped_deployment_id(self): + mock_router = MagicMock() + mock_router.cooldown_time = 60.0 + mock_router.get_model_info.return_value = None + + exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4") + exc.failed_deployment_id = "fallback-deployment" + + with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) + + mock_set_cooldown.assert_called_once() + call_kwargs = mock_set_cooldown.call_args[1] + assert call_kwargs["deployment"] == "fallback-deployment" + assert call_kwargs["original_exception"] is exc + + def test_does_not_trust_caller_supplied_metadata_bucket(self): + """A metadata bucket can't reliably be told apart from a caller-supplied + one without knowing this call's function_name, so a client with + permission to set metadata must not be able to get an arbitrary + deployment cooled down by forging a deployment_model_name marker.""" + mock_router = MagicMock() + mock_router.cooldown_time = 60.0 + mock_router.get_model_info.return_value = None + + exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4") + kwargs = { + "metadata": { + "model_info": {"id": "attacker-chosen-deployment"}, + "deployment_model_name": "gpt-4", + } + } + + with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs=kwargs, exception=exc) + + mock_set_cooldown.assert_not_called() + + def test_increments_failure_counter_before_cooldown_check(self): + """The fallback path must feed the same per-minute failure counter the + primary path uses, or repeated fallback failures never accumulate + toward the default percent-fail-rate cooldown threshold.""" + mock_router = MagicMock() + mock_router.cooldown_time = 60.0 + mock_router.get_model_info.return_value = None + + exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4") + exc.failed_deployment_id = "fallback-deployment" + + with ( + patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown, + patch( + "litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute" + ) as mock_increment, + ): + _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) + + mock_increment.assert_called_once_with( + litellm_router_instance=mock_router, deployment_id="fallback-deployment" + ) + mock_set_cooldown.assert_called_once() + + def test_no_op_when_deployment_id_missing(self): + mock_router = MagicMock() + + with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + _trigger_cooldown_for_failed_deployment( + litellm_router=mock_router, kwargs={}, exception=RuntimeError("no metadata") + ) + + mock_set_cooldown.assert_not_called() + + def test_skipped_for_advisor_orchestration_failure(self): + mock_router = MagicMock() + mock_router.cooldown_time = 60.0 + mock_router.get_model_info.return_value = None + + exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4") + exc.failed_deployment_id = "fallback-deployment" + mark_advisor_orchestration_failure(exc) + + with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) + + mock_set_cooldown.assert_not_called() + + def test_uses_deployment_litellm_params_cooldown_time_override(self): + mock_router = MagicMock() + mock_router.cooldown_time = 300.0 + mock_router.get_model_info.return_value = {"litellm_params": {"cooldown_time": 30.0}} + + exc = litellm.RateLimitError("Rate limit", "openai", "gpt-4") + exc.failed_deployment_id = "fallback-deployment" + + with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) + + call_kwargs = mock_set_cooldown.call_args[1] + assert call_kwargs["time_to_cooldown"] == 30.0 + + def test_uses_response_header_when_no_deployment_config(self): + """Precedence must match Router.deployment_callback_on_failure's primary + path: deployment config, then the response's Retry-After header, then the + router default.""" + mock_router = MagicMock() + mock_router.cooldown_time = 60.0 + mock_router.get_model_info.return_value = {"litellm_params": {}} + + exc = RuntimeError("upstream error") + exc.failed_deployment_id = "fallback-deployment" + exc.litellm_response_headers = httpx.Headers({"retry-after": "45"}) + + with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) + + call_kwargs = mock_set_cooldown.call_args[1] + assert call_kwargs["time_to_cooldown"] == 45 + + def test_silently_catches_exceptions(self): + mock_router = MagicMock() + mock_router.cooldown_time = 60.0 + mock_router.get_model_info.return_value = None + + exc = RuntimeError("upstream error") + exc.failed_deployment_id = "fallback-deployment" + + with patch( + "litellm.router_utils.fallback_event_handlers._set_cooldown_deployments", + side_effect=RuntimeError("cooldown error"), + ): + _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) + + def test_skips_request_scoped_404_on_generic_api_call(self): + """A generic API call (files/batches/threads/rerank/...) forwards a caller-supplied + resource id, so a 404 there means "that id doesn't exist", not "this deployment is + unhealthy". Without this guard, a single bad id would 404 every deployment in the + fallback chain and cool all of them down from one request.""" + mock_router = MagicMock() + mock_router.cooldown_time = 60.0 + mock_router.get_model_info.return_value = None + + exc = litellm.NotFoundError("not found", "openai", "gpt-4") + exc.failed_deployment_id = "fallback-deployment" + + with ( + patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown, + patch( + "litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute" + ) as mock_increment, + ): + _trigger_cooldown_for_failed_deployment( + litellm_router=mock_router, + kwargs={"original_generic_function": MagicMock()}, + exception=exc, + ) + + mock_set_cooldown.assert_not_called() + mock_increment.assert_not_called() + + def test_still_cools_down_404_outside_generic_api_call(self): + """The request-scoped-404 guard is scoped to generic API calls only: a 404 on a + regular completion fallback (no original_generic_function in kwargs) must still + cool down the deployment as before.""" + mock_router = MagicMock() + mock_router.cooldown_time = 60.0 + mock_router.get_model_info.return_value = None + + exc = litellm.NotFoundError("not found", "openai", "gpt-4") + exc.failed_deployment_id = "fallback-deployment" + + with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) + + mock_set_cooldown.assert_called_once() + + def test_skips_client_side_timeout_408(self): + """The proxy's x-litellm-timeout header lets a caller set an arbitrarily short + timeout, which litellm.Timeout reports as status 408 regardless of the + deployment's actual health. Without this guard, a caller could force a 408 on + every deployment in the fallback chain from a single request.""" + mock_router = MagicMock() + mock_router.cooldown_time = 60.0 + mock_router.get_model_info.return_value = None + + exc = litellm.Timeout(message="timeout", model="gpt-4", llm_provider="openai") + exc.failed_deployment_id = "fallback-deployment" + + with ( + patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown, + patch( + "litellm.router_utils.fallback_event_handlers.increment_deployment_failures_for_current_minute" + ) as mock_increment, + ): + _trigger_cooldown_for_failed_deployment( + litellm_router=mock_router, + kwargs={"client_side_timeout": True}, + exception=exc, + ) + + mock_set_cooldown.assert_not_called() + mock_increment.assert_not_called() + + def test_still_cools_down_408_without_client_side_timeout_flag(self): + """The client-side-timeout guard is scoped to caller-supplied timeouts only: a + 408 that did not come from x-litellm-timeout (no client_side_timeout in kwargs) + must still cool down the deployment as before.""" + mock_router = MagicMock() + mock_router.cooldown_time = 60.0 + mock_router.get_model_info.return_value = None + + exc = litellm.Timeout(message="timeout", model="gpt-4", llm_provider="openai") + exc.failed_deployment_id = "fallback-deployment" + + with patch("litellm.router_utils.fallback_event_handlers._set_cooldown_deployments") as mock_set_cooldown: + _trigger_cooldown_for_failed_deployment(litellm_router=mock_router, kwargs={}, exception=exc) + + mock_set_cooldown.assert_called_once() + + +class TestRunAsyncFallbackTriggersCooldown: + class RouterWithLoggingKwarg: + def __init__(self): + self.cooldown_time = 60.0 + + def log_retry(self, kwargs, e): + return kwargs + + def get_model_info(self, id): + return None + + async def async_function_with_fallbacks(self, *args, **kwargs): + raise RuntimeError("fallback model also failed") + + def _logging_obj(self, has_logged_async_failure: bool) -> MagicMock: + logging_obj = MagicMock() + logging_obj.model_call_details = {"has_logged_async_failure": has_logged_async_failure} + return logging_obj + + @pytest.mark.asyncio + async def test_triggers_cooldown_when_has_logged_async_failure_is_true(self): + with patch( + "litellm.router_utils.fallback_event_handlers._trigger_cooldown_for_failed_deployment" + ) as mock_trigger: + with pytest.raises(RuntimeError, match="fallback model also failed"): + await run_async_fallback( + litellm_router=self.RouterWithLoggingKwarg(), + fallback_model_group=["fallback-model"], + original_model_group="primary-model", + original_exception=RuntimeError("original request failed"), + max_fallbacks=3, + fallback_depth=0, + litellm_logging_obj=self._logging_obj(has_logged_async_failure=True), + ) + + mock_trigger.assert_called_once() + + @pytest.mark.asyncio + async def test_does_not_trigger_cooldown_when_has_logged_async_failure_is_false(self): + """This is the exact dead-code scenario the bug fix addresses: before it, + the normal failure callback runs for the first attempt in a fallback chain + (has_logged_async_failure is still False at that point), so no explicit + trigger is needed there.""" + with patch( + "litellm.router_utils.fallback_event_handlers._trigger_cooldown_for_failed_deployment" + ) as mock_trigger: + with pytest.raises(RuntimeError, match="fallback model also failed"): + await run_async_fallback( + litellm_router=self.RouterWithLoggingKwarg(), + fallback_model_group=["fallback-model"], + original_model_group="primary-model", + original_exception=RuntimeError("original request failed"), + max_fallbacks=3, + fallback_depth=0, + litellm_logging_obj=self._logging_obj(has_logged_async_failure=False), + ) + + mock_trigger.assert_not_called() + + @pytest.mark.asyncio + async def test_does_not_trigger_cooldown_when_no_logging_obj_present(self): + with patch( + "litellm.router_utils.fallback_event_handlers._trigger_cooldown_for_failed_deployment" + ) as mock_trigger: + with pytest.raises(RuntimeError, match="fallback model also failed"): + await run_async_fallback( + litellm_router=self.RouterWithLoggingKwarg(), + fallback_model_group=["fallback-model"], + original_model_group="primary-model", + original_exception=RuntimeError("original request failed"), + max_fallbacks=3, + fallback_depth=0, + ) + + mock_trigger.assert_not_called() diff --git a/tests/test_litellm/test_router_weighted_failover.py b/tests/test_litellm/test_router_weighted_failover.py index 6f26f329953..0115638e1fe 100644 --- a/tests/test_litellm/test_router_weighted_failover.py +++ b/tests/test_litellm/test_router_weighted_failover.py @@ -13,6 +13,7 @@ from unittest.mock import AsyncMock, patch import pytest +import litellm from litellm import Router from litellm.utils import _get_excluded_filtered_deployments @@ -222,6 +223,37 @@ async def test_acompletion_stamps_dynamic_id_for_clientside_credentials(): assert failed_deployment_id != "dep-a" +@pytest.mark.asyncio +async def test_acompletion_stamps_dynamic_id_for_clientside_credentials_on_timeout(): + """Same bug as the RuntimeError case above, but for the separate `except litellm.Timeout` + branch in `_acompletion`: it has its own call to the stamping helper, so a fix that only + covers the generic `except Exception` branch would leave a caller-supplied timeout + (`litellm.Timeout` is what `x-litellm-timeout` maps to) stamping the shared static id.""" + router = Router( + model_list=[ + { + "model_name": "test-model", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "test-key"}, + "model_info": {"id": "dep-a"}, + } + ], + ) + + timeout_exc = litellm.Timeout(message="boom", model="test-model", llm_provider="openai") + with patch("litellm.acompletion", new_callable=AsyncMock, side_effect=timeout_exc): + with pytest.raises(litellm.Timeout) as exc_info: + await router._acompletion( + model="test-model", + messages=[{"role": "user", "content": "Hello"}], + api_key="tenant-supplied-key", + metadata={"model_group": "test-model"}, + ) + + failed_deployment_id = getattr(exc_info.value, "failed_deployment_id", None) + assert failed_deployment_id is not None + assert failed_deployment_id != "dep-a" + + def test_completion_stamps_dynamic_id_for_clientside_credentials(): """Sync counterpart: _completion's exception handler must stamp the dynamic client-side-credential deployment id, not the shared static deployment's id."""