litellm/tests/test_litellm/router_utils/test_cooldown_cache.py
Deepanshu Lulla 05943b47a3
fix(router): cool down failed fallback deployments and correct cooldown TTL after Redis backfill (#35104)
* fix(router): cool down failed fallback deployments and correct cooldown TTL after Redis backfill

A deployment that failed partway through a fallback chain (any attempt after
the first) was silently exempt from cooldown, because the has_logged_async_failure
dedup flag blocks the normal failure callback for every attempt past the first.
_trigger_cooldown_for_failed_deployment now explicitly evaluates cooldown for
that deployment when the dedup flag is set, using the same deployment-config >
response-header > router-default precedence as the primary failure path, and
skips advisor-orchestration failures. Deployment-ID resolution prefers the
exception's stamped failed_deployment_id, now also set from the generic-API-call
fallback path (rerank, embeddings, /v1/messages, etc.), falling back to metadata
inspection for call paths that don't stamp it yet.

CooldownCache also recomputes the remaining TTL when DualCache promotes a Redis
entry into the in-memory layer: before this, a cooldown entry restored from Redis
kept the in-memory layer's default 600s TTL regardless of the deployment's real
cooldown_time, so a deployment could stay excluded from routing for up to 10
minutes after a much shorter cooldown had already expired.

* fix(router): address Greptile review on the fallback-cooldown trigger

Two P1 findings on PR #35104:

- _trigger_cooldown_for_failed_deployment never incremented the deployment's
  per-minute failure counter before evaluating cooldown, so a fallback
  deployment's repeated retryable failures never accumulated toward the
  default percent-fail-rate threshold that _should_cooldown_deployment checks.

- The metadata-bucket fallback (checking "metadata" before "litellm_metadata"
  for a deployment_model_name marker) could be fooled by a caller with
  permission to set metadata, since neither bucket's authorship can be
  determined without knowing the call's function_name. Removed it entirely;
  cooldown now requires the server-stamped failed_deployment_id, matching
  what the primary chat-completions path and the generic-API-call path
  (rerank, embeddings, /v1/messages, etc.) already set unconditionally.

* fix(router): freeze the litellm_params fallback mapping to satisfy the type-discipline gate

* fix(router): defer f-string interpolation in fallback-cooldown debug logs

* fix(router): annotate cooldown-path locals with Final to satisfy the LIT010 budget

* fix(router): don't cool down deployments for request-scoped 404s on generic API fallbacks

* fix(router): stamp the dynamic client-side-credential deployment id, not the shared static one

* fix(router): don't cool down deployments for a caller-supplied x-litellm-timeout

* fix(router): stamp dynamic client-side-credential id in completion fallback paths too

The generic-API-call helper already stamped the effective (dynamic-if-client-side-credential)
deployment id on exceptions, but the regular _completion/_acompletion exception handlers still
stamped the static shared deployment's id. A tenant using invalid forwarded credentials could
generate repeated failures attributed to, and eventually cooling down, the shared deployment
other tenants rely on. Extracted the stamping logic into one shared helper used by all three
call sites (generic API, sync completion, async completion) so the fix and future changes to it
stay in one place.

* test(router): add direct-reference unit tests for the new stamping helper

router_code_coverage.py's coverage gate flags _stamp_failed_deployment_id_with_effective_model_info
as untested because it only sees the function invoked indirectly through _completion/_acompletion's
exception handlers. Added two tests that call it directly, covering both the dynamic-id-present and
static-fallback branches.

* test(router): cover the timeout stamping branch and async active-cooldown append

_acompletion's litellm.Timeout handler and async_get_active_cooldowns' happy
path both lacked direct coverage despite their sibling branches (the generic
Exception handler, the sync get_active_cooldowns) being tested.

* test(router): remove duplicate cooldown-trigger and fallback-helper tests

#34416 landed its own TestTriggerCooldownForFailedDeployment/
TestRunAsyncFallbackTriggersCooldown classes and
test_ageneric_api_call_with_fallbacks_helper_stamps_failed_deployment_id
covering the exact same scenarios as this branch's earlier flat-function
tests, once its version of fallback_event_handlers.py was taken as-is
during the last merge. Dropping the redundant copies.

---------

Co-authored-by: Deepanshu <deepanshu.lulla@alpha-sense.com>
2026-08-10 16:51:55 -07:00

437 lines
18 KiB
Python

"""
Unit tests for CooldownCache exception masking functionality
"""
import os
import sys
import time
from unittest.mock import MagicMock
import pytest
# Add the parent directory to the system path
sys.path.insert(0, os.path.abspath("../../.."))
from litellm.caching.dual_cache import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.router_utils.cooldown_cache import CooldownCache, CooldownCacheValue
class TestCooldownCacheExceptionMasking:
"""Test suite for CooldownCache exception masking functionality"""
@pytest.fixture
def cooldown_cache(self):
"""Create a CooldownCache instance for testing"""
mock_dual_cache = MagicMock(spec=DualCache)
return CooldownCache(cache=mock_dual_cache, default_cooldown_time=60.0)
def test_exception_masker_initialization(self, cooldown_cache):
"""Test that the exception masker is properly initialized"""
assert isinstance(cooldown_cache.exception_masker, SensitiveDataMasker)
assert cooldown_cache.exception_masker.visible_prefix == 50
assert cooldown_cache.exception_masker.visible_suffix == 0
assert cooldown_cache.exception_masker.mask_char == "*"
def test_short_exception_string_not_masked(self, cooldown_cache):
"""Test that short exception strings are not masked"""
short_exception = "Short error"
model_id = "test-model"
exception_status = 500
cooldown_time = 30.0
cooldown_key, cooldown_data = cooldown_cache._common_add_cooldown_logic(
model_id=model_id,
original_exception=Exception(short_exception),
exception_status=exception_status,
cooldown_time=cooldown_time,
)
# Short exception should not be masked
assert cooldown_data["exception_received"] == short_exception
assert cooldown_key == f"deployment:{model_id}:cooldown"
def test_long_exception_string_masked(self, cooldown_cache):
"""Test that long exception strings are properly masked"""
# Create a long exception string that simulates prompt leakage
long_exception = (
"litellm.proxy.proxy_server._handle_llm_api_exception(): Exception occurred - "
"No deployments available for selected model, Try again in 5 seconds. "
"Passed model=anthropic_claude_sonnet_4_v1_0. pre-call-checks=False, "
"cooldown_list=[('deepseek_r1-eastus', {'exception_received': "
"'litellm.RateLimitError: RateLimitError: Azure_aiException - "
'{"error":{"code":"Invalid input","status":422,"message":"invalid input error",'
'"details":[{"type":"model_attributes_type","loc":["body"],'
'"msg":"Tell me a story about a dragon and a princess in a magical kingdom '
"where the dragon is actually protecting the princess from an evil wizard "
'who wants to steal her magical powers and use them to conquer the world"}]}'
)
model_id = "test-model"
exception_status = 429
cooldown_time = 60.0
cooldown_key, cooldown_data = cooldown_cache._common_add_cooldown_logic(
model_id=model_id,
original_exception=Exception(long_exception),
exception_status=exception_status,
cooldown_time=cooldown_time,
)
masked_exception = cooldown_data["exception_received"]
# Should start with first 50 characters
assert masked_exception.startswith(long_exception[:50])
# Should contain masking characters
assert "*" in masked_exception
# Should be same length (prefix + asterisks)
assert len(masked_exception) == len(long_exception)
# Should not contain the sensitive prompt content
assert "Tell me a story about a dragon" not in masked_exception
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")
def test_exception_with_api_keys_masked(self, cooldown_cache):
"""Test that API keys in exceptions are properly masked"""
exception_with_key = (
"Authentication failed with api_key=sk-1234567890abcdefghijklmnopqrstuvwxyz "
"and token=bearer_token_123456789 for model gpt-4"
)
model_id = "test-model"
exception_status = 401
cooldown_time = 30.0
cooldown_key, cooldown_data = cooldown_cache._common_add_cooldown_logic(
model_id=model_id,
original_exception=Exception(exception_with_key),
exception_status=exception_status,
cooldown_time=cooldown_time,
)
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 "*" in masked_exception
assert len(masked_exception) == len(exception_with_key)
def test_cooldown_data_structure(self, cooldown_cache):
"""Test that the cooldown data structure is correctly formed"""
exception_msg = "Test exception for structure validation"
model_id = "test-model"
exception_status = 500
cooldown_time = 45.0
cooldown_key, cooldown_data = cooldown_cache._common_add_cooldown_logic(
model_id=model_id,
original_exception=Exception(exception_msg),
exception_status=exception_status,
cooldown_time=cooldown_time,
)
# Verify cooldown data structure
assert isinstance(cooldown_data, dict)
assert "exception_received" in cooldown_data
assert "status_code" in cooldown_data
assert "timestamp" in cooldown_data
assert "cooldown_time" in cooldown_data
# Verify data types
assert isinstance(cooldown_data["exception_received"], str)
assert isinstance(cooldown_data["status_code"], str)
assert isinstance(cooldown_data["timestamp"], float)
assert isinstance(cooldown_data["cooldown_time"], float)
# Verify values
assert cooldown_data["status_code"] == str(exception_status)
assert cooldown_data["cooldown_time"] == cooldown_time
assert cooldown_data["exception_received"] == exception_msg
def test_exception_object_conversion(self, cooldown_cache):
"""Test that different exception types are properly converted to strings"""
# Test with different exception types
exceptions = [
ValueError("Invalid value provided"),
KeyError("Missing required key"),
RuntimeError("Runtime error occurred"),
Exception("Generic exception"),
]
for exc in exceptions:
model_id = f"test-model-{exc.__class__.__name__}"
cooldown_key, cooldown_data = cooldown_cache._common_add_cooldown_logic(
model_id=model_id,
original_exception=exc,
exception_status=500,
cooldown_time=30.0,
)
# Should successfully convert exception to string
assert isinstance(cooldown_data["exception_received"], str)
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"""
debugging_exception = (
"RateLimitError: Rate limit exceeded for model gpt-4. "
"Current usage: 1000 tokens/minute. Limit: 500 tokens/minute. "
"Request details: model=gpt-4, user_id=user123, "
"prompt='Write a comprehensive analysis of the economic implications "
"of artificial intelligence adoption in the healthcare sector, including "
"potential cost savings, job displacement, and regulatory challenges'"
)
model_id = "gpt-4-deployment"
exception_status = 429
cooldown_time = 120.0
cooldown_key, cooldown_data = cooldown_cache._common_add_cooldown_logic(
model_id=model_id,
original_exception=Exception(debugging_exception),
exception_status=exception_status,
cooldown_time=cooldown_time,
)
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-")
# Should mask the prompt content
assert "Write a comprehensive analysis" not in masked_exception
assert "healthcare sector" not in masked_exception
# Should contain masking indicator
assert "*" in masked_exception
def test_error_handling_in_common_add_cooldown_logic(self, cooldown_cache):
"""Test error handling in the _common_add_cooldown_logic method"""
# This test ensures that edge cases are properly handled
model_id = "test-model"
# Test with None exception (edge case) - should be handled gracefully
cooldown_key, cooldown_data = cooldown_cache._common_add_cooldown_logic(
model_id=model_id,
original_exception=None,
exception_status=500,
cooldown_time=30.0,
)
# Should handle None by converting to string
assert cooldown_data["exception_received"] == "None"
assert cooldown_key == f"deployment:{model_id}:cooldown"
def test_custom_masker_settings(self):
"""Test that custom masker settings work correctly"""
mock_dual_cache = MagicMock(spec=DualCache)
# Create cooldown cache and verify default settings
cache = CooldownCache(cache=mock_dual_cache, default_cooldown_time=60.0)
# Test that we can access and verify the masker configuration
assert cache.exception_masker.visible_prefix == 50
assert cache.exception_masker.visible_suffix == 0
assert cache.exception_masker.mask_char == "*"
# Test masking behavior with these settings
long_string = "A" * 100 # 100 character string
masked = cache.exception_masker._mask_value(long_string)
# Should show first 50 characters, then all asterisks
expected = "A" * 50 + "*" * 50
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()
dual_cache = DualCache(in_memory_cache=in_memory)
return CooldownCache(cache=dual_cache, default_cooldown_time=60.0)
def _entry(self, timestamp: float, cooldown_time: float) -> CooldownCacheValue:
return CooldownCacheValue(
exception_received="Rate limit",
status_code="429",
timestamp=timestamp,
cooldown_time=cooldown_time,
)
def test_expired_entry_returns_none_and_evicts(self):
cc = self._make_cooldown_cache()
key = "deployment:expired-dep:cooldown"
entry = self._entry(timestamp=time.time() - 120.0, cooldown_time=60.0)
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=600)
result = cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
assert result is None
assert cc.cache.in_memory_cache.get_cache(key) is None
def test_active_entry_within_window_returns_value(self):
cc = self._make_cooldown_cache()
key = "deployment:active-dep:cooldown"
entry = self._entry(timestamp=time.time(), cooldown_time=60.0)
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=60)
result = cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
assert result is not None
assert result["status_code"] == "429"
def test_inflated_ttl_is_corrected(self):
cc = self._make_cooldown_cache()
key = "deployment:backfilled-dep:cooldown"
remaining = 30.0
entry = self._entry(timestamp=time.time() - (60.0 - remaining), cooldown_time=60.0)
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=600)
result = cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
assert result is not None
corrected_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
assert corrected_expiry is not None
assert corrected_expiry - time.time() <= 60.0
def test_normal_ttl_not_modified(self):
cc = self._make_cooldown_cache()
key = "deployment:normal-dep:cooldown"
entry = self._entry(timestamp=time.time(), cooldown_time=60.0)
cc.cache.in_memory_cache.set_cache(key, dict(entry), ttl=60)
original_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
cc._corrected_active_cooldown(key, dict(entry), current_time=time.time())
after_expiry = cc.cache.in_memory_cache.ttl_dict.get(key)
assert after_expiry == original_expiry