diff --git a/litellm/router_utils/cooldown_cache.py b/litellm/router_utils/cooldown_cache.py index ef29f7d8fd3..4fd433d9391 100644 --- a/litellm/router_utils/cooldown_cache.py +++ b/litellm/router_utils/cooldown_cache.py @@ -118,6 +118,7 @@ class CooldownCache: ) # Set the cache with a TTL equal to the cooldown time + self._drop_in_memory_entry_if_extending(cooldown_key, _cooldown_time) self.cooldown_store.set_cache( value=cooldown_data, key=cooldown_key, @@ -127,6 +128,14 @@ class CooldownCache: verbose_logger.error("CooldownCache::add_deployment_to_cooldown - Exception occurred - %s", e) raise e + def _drop_in_memory_entry_if_extending(self, cooldown_key: str, new_cooldown_time: float) -> None: + """InMemoryCache keeps a live key's expiry, so a longer cooldown lands only if the entry is deleted first.""" + current_expiry: Final = self.in_memory_cache.ttl_dict.get(cooldown_key) + if current_expiry is None: + return + if float(current_expiry) < time.time() + float(new_cooldown_time): + self.in_memory_cache.delete_cache(cooldown_key) + @staticmethod @functools.lru_cache(maxsize=1024) def get_cooldown_cache_key(model_id: str) -> str: diff --git a/tests/unit/router_utils/test_cooldown_cache.py b/tests/unit/router_utils/test_cooldown_cache.py index 6f90fa8465f..fc59a2054f5 100644 --- a/tests/unit/router_utils/test_cooldown_cache.py +++ b/tests/unit/router_utils/test_cooldown_cache.py @@ -582,3 +582,55 @@ class TestCooldownSurvivesUnrelatedCacheTraffic: assert [model_id] == [entry[0] for entry in active], ( "unrelated router cache traffic must not evict a cooldown that is still running" ) + + +class TestCooldownExtension: + """A second failure asking for a longer cooldown must move the deadline out.""" + + @staticmethod + def _cache() -> CooldownCache: + return CooldownCache(cache=DualCache(), default_cooldown_time=60.0) + + def test_longer_cooldown_extends_the_deadline(self): + cc = self._cache() + key = CooldownCache.get_cooldown_cache_key("dep-a") + + cc.add_deployment_to_cooldown( + model_id="dep-a", original_exception=Exception("429"), exception_status=429, cooldown_time=2.0 + ) + first_expiry = cc.in_memory_cache.ttl_dict[key] + + started = time.time() + cc.add_deployment_to_cooldown( + model_id="dep-a", original_exception=Exception("429"), exception_status=429, cooldown_time=60.0 + ) + second_expiry = cc.in_memory_cache.ttl_dict[key] + + assert second_expiry > first_expiry, "the longer cooldown was dropped" + assert second_expiry - started == pytest.approx(60.0, abs=2.0) + + def test_shorter_cooldown_does_not_cut_an_active_one_short(self): + cc = self._cache() + key = CooldownCache.get_cooldown_cache_key("dep-b") + + cc.add_deployment_to_cooldown( + model_id="dep-b", original_exception=Exception("429"), exception_status=429, cooldown_time=60.0 + ) + long_expiry = cc.in_memory_cache.ttl_dict[key] + + cc.add_deployment_to_cooldown( + model_id="dep-b", original_exception=Exception("500"), exception_status=500, cooldown_time=1.0 + ) + + assert cc.in_memory_cache.ttl_dict[key] == long_expiry, "a brief cooldown shortened a long one" + + def test_first_cooldown_is_unaffected(self): + cc = self._cache() + key = CooldownCache.get_cooldown_cache_key("dep-c") + + started = time.time() + cc.add_deployment_to_cooldown( + model_id="dep-c", original_exception=Exception("429"), exception_status=429, cooldown_time=30.0 + ) + + assert cc.in_memory_cache.ttl_dict[key] - started == pytest.approx(30.0, abs=2.0)