diff --git a/litellm/caching/in_memory_cache.py b/litellm/caching/in_memory_cache.py index 082cac791f..5239fa1f4b 100644 --- a/litellm/caching/in_memory_cache.py +++ b/litellm/caching/in_memory_cache.py @@ -11,6 +11,7 @@ Has 4 methods: import json import sys import time +import heapq from typing import TYPE_CHECKING, Any, List, Optional if TYPE_CHECKING: @@ -46,6 +47,7 @@ class InMemoryCache(BaseCache): # in-memory cache self.cache_dict: dict = {} self.ttl_dict: dict = {} + self.expiration_heap: list[tuple[float, str]] = [] def check_value_size(self, value: Any): """ @@ -114,19 +116,27 @@ class InMemoryCache(BaseCache): """ current_time = time.time() - - # Step 1: Remove expired items - expired_keys = [key for key, ttl in self.ttl_dict.items() if current_time > ttl] - for key in expired_keys: - self._remove_key(key) - # Step 2: If cache is still full, evict items with earliest expiration times - if len(self.cache_dict) >= self.max_size_in_memory: - # Sort by expiration time (earliest first) and evict until we're under the limit - items_by_expiration = sorted(self.ttl_dict.items(), key=lambda x: x[1]) - keys_to_evict = items_by_expiration[:len(self.cache_dict) - self.max_size_in_memory + 1] - - for key, _ in keys_to_evict: + # Step 1: Remove expired or outdated items + while self.expiration_heap: + expiration_time, key = self.expiration_heap[0] + + # Case 1: Heap entry is outdated + if expiration_time != self.ttl_dict.get(key): + heapq.heappop(self.expiration_heap) + # Case 2: Entry is valid but expired + elif expiration_time <= current_time: + heapq.heappop(self.expiration_heap) + self._remove_key(key) + else: + # Case 3: Entry is valid and not expired + break + + # Step 2: Evict if cache is still full + while len(self.cache_dict) >= self.max_size_in_memory: + expiration_time, key = heapq.heappop(self.expiration_heap) + # Skip if key was removed or updated + if self.ttl_dict.get(key) == expiration_time: self._remove_key(key) # de-reference the removed item @@ -150,7 +160,7 @@ class InMemoryCache(BaseCache): # Handle the edge case where max_size_in_memory is 0 if self.max_size_in_memory == 0: return # Don't cache anything if max size is 0 - + if len(self.cache_dict) >= self.max_size_in_memory: # only evict when cache is full self.evict_cache() @@ -161,8 +171,10 @@ class InMemoryCache(BaseCache): if self.allow_ttl_override(key): # if ttl is not set, set it to default ttl if "ttl" in kwargs and kwargs["ttl"] is not None: self.ttl_dict[key] = time.time() + float(kwargs["ttl"]) + heapq.heappush(self.expiration_heap, (self.ttl_dict[key], key)) else: self.ttl_dict[key] = time.time() + self.default_ttl + heapq.heappush(self.expiration_heap, (self.ttl_dict[key], key)) async def async_set_cache(self, key, value, **kwargs): self.set_cache(key=key, value=value, **kwargs) @@ -253,6 +265,7 @@ class InMemoryCache(BaseCache): def flush_cache(self): self.cache_dict.clear() self.ttl_dict.clear() + self.expiration_heap.clear() async def disconnect(self): pass diff --git a/tests/test_litellm/caching/test_in_memory_cache.py b/tests/test_litellm/caching/test_in_memory_cache.py index 616c60c74a..e7cc7f80ab 100644 --- a/tests/test_litellm/caching/test_in_memory_cache.py +++ b/tests/test_litellm/caching/test_in_memory_cache.py @@ -186,3 +186,16 @@ def test_in_memory_cache_eviction_order(): # Items with later expiration should remain assert "late_expire" in in_memory_cache.cache_dict assert "new_item" in in_memory_cache.cache_dict + + +def test_in_memory_cache_heap_size_staus_bounded(): + """ + Test that the expiration_heap does not grow unbounded when the same key is updated repeaatedly. + """ + in_memory_cache = InMemoryCache(max_size_in_memory=10) + + for i in range(1_000): + in_memory_cache.set_cache(key="hot_key", value=f"value_{i}", ttl=60) + + # Expiration heap should only have 1 entry + assert len(in_memory_cache.expiration_heap) == 1 diff --git a/tests/test_litellm/litellm_core_utils/specialty_caches/test_dynamic_logging_cache.py b/tests/test_litellm/litellm_core_utils/specialty_caches/test_dynamic_logging_cache.py index 85fcc2700b..f21cd56750 100644 --- a/tests/test_litellm/litellm_core_utils/specialty_caches/test_dynamic_logging_cache.py +++ b/tests/test_litellm/litellm_core_utils/specialty_caches/test_dynamic_logging_cache.py @@ -41,8 +41,10 @@ class TestLangfuseInMemoryCache: "litellm.integrations.langfuse.langfuse.LangFuseLogger", MockLangFuseLogger ): # Add the mock logger to cache with expired TTL + expired_time = time.time() - 1 # Already expired self.cache.cache_dict["test_key"] = mock_logger - self.cache.ttl_dict["test_key"] = time.time() - 1 # Already expired + self.cache.ttl_dict["test_key"] = expired_time + self.cache.expiration_heap = [(expired_time, "test_key")] initial_count = litellm.initialized_langfuse_clients