litellm/tests/unit/caching/test_in_memory_cache.py
2026-09-28 10:16:33 +05:30

335 lines
11 KiB
Python

import threading
import time
from concurrent.futures import ThreadPoolExecutor
import pytest
from pydantic import BaseModel
from litellm.caching.in_memory_cache import InMemoryCache
class _SlowInt(int):
def __add__(self, value: int) -> "_SlowInt":
time.sleep(0.05)
return _SlowInt(int(self) + value)
@pytest.mark.parametrize(
"value",
[
{"budget": {"spend": 1.0, "events": [1]}},
[{"spend": 1.0, "events": [1]}],
],
)
def test_get_cache_isolates_nested_mutable_values(value):
cache = InMemoryCache()
cache.set_cache("auth", value)
first = cache.get_cache("auth")
nested = first["budget"] if isinstance(first, dict) else first[0]
nested["spend"] = 50.0
nested["events"].append(2)
expected = {"budget": {"spend": 1.0, "events": [1]}} if isinstance(value, dict) else [{"spend": 1.0, "events": [1]}]
assert cache.get_cache("auth") == expected
assert cache.cache_dict["auth"] == expected
assert first is not cache.cache_dict["auth"]
class _CachedBudget(BaseModel):
spend: float
events: list[int]
def test_get_cache_isolates_nested_model_state():
cache = InMemoryCache()
cache.set_cache("budget", _CachedBudget(spend=1.0, events=[1]))
first = cache.get_cache("budget")
first.spend = 50.0
first.events.append(2)
assert cache.get_cache("budget") == _CachedBudget(spend=1.0, events=[1])
def test_get_cache_keeps_unrelated_client_identity():
cache = InMemoryCache()
client = object()
cache.set_cache("client", client)
assert cache.get_cache("client") is client
def test_get_cache_preserves_json_string_and_scalar_behavior():
cache = InMemoryCache()
cache.set_cache("json", '{"budget": {"spend": 1}}')
cache.set_cache("plain", "plain text")
cache.set_cache("count", 7)
first = cache.get_cache("json")
first["budget"]["spend"] = 50
assert cache.get_cache("json") == {"budget": {"spend": 1}}
assert cache.get_cache("plain") == "plain text"
assert cache.get_cache("count") == 7
def test_increment_cache_is_atomic_under_thread_concurrency():
cache = InMemoryCache()
seed = 1000
cache.set_cache("counter", _SlowInt(seed))
thread_count = 8
barrier = threading.Barrier(thread_count)
def increment(_: int) -> float:
barrier.wait()
return cache.increment_cache("counter", 1)
with ThreadPoolExecutor(max_workers=thread_count) as executor:
tuple(executor.map(increment, range(thread_count)))
assert cache.get_cache("counter") == seed + thread_count
@pytest.mark.parametrize("bad_value", [{"count": 1}, [1], _CachedBudget(spend=1.0, events=[1])])
def test_increment_cache_rejects_non_numeric_cached_values(bad_value):
cache = InMemoryCache()
cache.set_cache("counter", bad_value)
with pytest.raises(TypeError, match="Cached value is not numeric"):
cache.increment_cache("counter", 1)
@pytest.mark.asyncio
async def test_async_set_cache_sadd_rejects_non_set_cached_value():
cache = InMemoryCache()
cache.set_cache("members", ["one"])
with pytest.raises(TypeError, match="Cached value is not a set"):
await cache.async_set_cache_sadd("members", ["two"], ttl=None)
async def test_async_increment_delegates_to_locked_sync_path():
cache = InMemoryCache()
assert await cache.async_increment("counter", 2) == 2
assert await cache.async_increment("counter", 3) == 5
assert cache.get_cache("counter") == 5
def test_in_memory_openai_obj_cache():
from openai import OpenAI
openai_obj = OpenAI(api_key="my-fake-key")
in_memory_cache = InMemoryCache()
in_memory_cache.set_cache(key="my-fake-key", value=openai_obj)
cached_obj = in_memory_cache.get_cache(key="my-fake-key")
assert cached_obj is not None
assert cached_obj == openai_obj
def test_in_memory_cache_max_size_per_item():
"""
Test that the cache will not store items larger than the max size per item
"""
in_memory_cache = InMemoryCache(max_size_per_item=100)
result = in_memory_cache.check_value_size("a" * 100000000)
assert result is False
def test_in_memory_cache_ttl():
"""
Check that
- if ttl is not set, it will be set to default ttl
- if object expires, the ttl is also removed
"""
in_memory_cache = InMemoryCache()
in_memory_cache.set_cache(key="my-fake-key", value="my-fake-value", ttl=10)
initial_ttl_time = in_memory_cache.ttl_dict["my-fake-key"]
assert initial_ttl_time is not None
in_memory_cache.set_cache(key="my-fake-key", value="my-fake-value-2", ttl=10)
new_ttl_time = in_memory_cache.ttl_dict["my-fake-key"]
assert new_ttl_time == initial_ttl_time # ttl should not be updated
## On object expiration, the ttl should be removed
in_memory_cache.set_cache(key="new-fake-key", value="new-fake-value", ttl=1)
new_ttl_time = in_memory_cache.ttl_dict["new-fake-key"]
assert new_ttl_time is not None
time.sleep(1)
in_memory_cache.get_cache(key="new-fake-key")
new_ttl_time = in_memory_cache.ttl_dict.get("new-fake-key")
assert new_ttl_time is None
def test_in_memory_cache_ttl_allow_override():
"""
Check that
- if ttl is not set, it will be set to default ttl
- if object expires, the ttl is also removed
"""
in_memory_cache = InMemoryCache()
## On object expiration, but no get_cache, the override should be allowed
in_memory_cache.set_cache(key="new-fake-key", value="new-fake-value", ttl=1)
initial_ttl_time = in_memory_cache.ttl_dict["new-fake-key"]
assert initial_ttl_time is not None
time.sleep(1)
in_memory_cache.set_cache(key="new-fake-key", value="new-fake-value-2", ttl=1)
new_ttl_time = in_memory_cache.ttl_dict["new-fake-key"]
assert new_ttl_time is not None
assert new_ttl_time != initial_ttl_time
def test_in_memory_cache_max_size_with_ttl():
"""
Test that max_size_in_memory is respected even when all items have long TTLs.
This tests the fix for the unbounded growth issue.
"""
in_memory_cache = InMemoryCache(max_size_in_memory=3)
long_ttl = 86400 # 1 day
# Fill the cache to max capacity
for i in range(3):
in_memory_cache.set_cache(key=f"key_{i}", value=f"value_{i}", ttl=long_ttl)
time.sleep(0.01) # Small delay to ensure different timestamps
assert len(in_memory_cache.cache_dict) == 3
assert len(in_memory_cache.ttl_dict) == 3
# Add another item - should evict the earliest item
in_memory_cache.set_cache(key="key_3", value="value_3", ttl=long_ttl)
# Cache should still be at max size, not larger
assert len(in_memory_cache.cache_dict) == 3
assert len(in_memory_cache.ttl_dict) == 3
# key_0 should have been evicted (it was added first)
assert "key_0" not in in_memory_cache.cache_dict
assert "key_0" not in in_memory_cache.ttl_dict
# Other keys should still be present
assert "key_1" in in_memory_cache.cache_dict
assert "key_2" in in_memory_cache.cache_dict
assert "key_3" in in_memory_cache.cache_dict
def test_in_memory_cache_expired_items_evicted_first():
"""
Test that expired items are evicted before non-expired items when cache is full.
"""
in_memory_cache = InMemoryCache(max_size_in_memory=3)
# Add items with short TTL that will expire
in_memory_cache.set_cache(key="expired_1", value="value_1", ttl=1)
in_memory_cache.set_cache(key="expired_2", value="value_2", ttl=1)
# Add item with long TTL
in_memory_cache.set_cache(key="long_lived", value="value_long", ttl=86400)
assert len(in_memory_cache.cache_dict) == 3
# Wait for short TTL items to expire
time.sleep(2)
# Add new item - should evict expired items first, not the long-lived one
in_memory_cache.set_cache(key="new_item", value="new_value", ttl=86400)
# Long-lived item should still be present
assert "long_lived" in in_memory_cache.cache_dict
assert "new_item" in in_memory_cache.cache_dict
# Expired items should be gone
assert "expired_1" not in in_memory_cache.cache_dict
assert "expired_2" not in in_memory_cache.cache_dict
assert "expired_1" not in in_memory_cache.ttl_dict
assert "expired_2" not in in_memory_cache.ttl_dict
def test_in_memory_cache_eviction_order():
"""
Test that when non-expired items need to be evicted, those with earliest expiration times are evicted first.
"""
in_memory_cache = InMemoryCache(max_size_in_memory=2)
# Add items with different TTLs
in_memory_cache.set_cache(key="early_expire", value="value_1", ttl=100) # expires in 100 seconds
time.sleep(0.01)
in_memory_cache.set_cache(key="late_expire", value="value_2", ttl=200) # expires in 200 seconds
# Verify TTL order
early_ttl = in_memory_cache.ttl_dict["early_expire"]
late_ttl = in_memory_cache.ttl_dict["late_expire"]
assert early_ttl < late_ttl, "early_expire should have earlier expiration time"
assert len(in_memory_cache.cache_dict) == 2
# Add third item - should evict the one with earliest expiration time
in_memory_cache.set_cache(key="new_item", value="value_3", ttl=300)
assert len(in_memory_cache.cache_dict) == 2
# Item with earliest expiration should be evicted
assert "early_expire" not in in_memory_cache.cache_dict
assert "early_expire" not in in_memory_cache.ttl_dict
# 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
def test_in_memory_cache_prunes_expired_heap_entries_below_capacity():
"""
Re-inserting expired keys below capacity should not grow expiration_heap
without bound.
"""
in_memory_cache = InMemoryCache(max_size_in_memory=200, default_ttl=1)
for cycle in range(3):
for i in range(5):
in_memory_cache.set_cache(key=f"key_{i}", value=f"value_{cycle}_{i}", ttl=1)
time.sleep(1.1)
for i in range(5):
in_memory_cache.set_cache(key=f"key_{i}", value=f"value_final_{i}", ttl=1)
assert len(in_memory_cache.cache_dict) == 5
assert len(in_memory_cache.ttl_dict) == 5
assert len(in_memory_cache.expiration_heap) == 5
def test_in_memory_cache_injected_clock_controls_expiry_and_eviction() -> None:
class Clock:
now = 0.0
def __call__(self) -> float:
return self.now
clock = Clock()
cache = InMemoryCache(max_size_in_memory=2, default_ttl=60, clock=clock)
cache.set_cache("first", "original", ttl=10)
clock.now = 9.0
cache.set_cache("second", "survivor")
assert cache.get_cache("first") == "original"
clock.now = 10.001
assert cache.get_cache("first") is None
cache.set_cache("third", "replacement")
assert cache.get_cache("second") == "survivor"
clock.now = 69.001
cache.set_cache("fourth", "new")
assert cache.get_cache("second") is None
assert cache.get_cache("third") == "replacement"
assert cache.get_cache("fourth") == "new"