This commit is contained in:
Charan Rathore 2026-09-30 10:30:04 -04:00 • committed by GitHub
commit 21cee60d7c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 84 additions and 18 deletions

View file

@ -14,6 +14,7 @@ import sys
import threading
import time
from collections.abc import Callable
from copy import deepcopy
from typing import TYPE_CHECKING, Any, Final
if TYPE_CHECKING:
@ -192,6 +193,8 @@ class InMemoryCache(BaseCache):
"""
# get the value
init_value: Final = self.get_cache(key=key) or set()
if not isinstance(init_value, set):
raise TypeError("Cached value is not a set")
for val in value:
init_value.add(val)
self.set_cache(key, init_value, ttl=ttl)
@ -215,8 +218,10 @@ class InMemoryCache(BaseCache):
original_cached_response: Final = self.cache_dict[key]
try:
cached_response = json.loads(original_cached_response)
except Exception:
cached_response = original_cached_response
except (TypeError, ValueError):
if isinstance(original_cached_response, (dict, list, BaseModel)):
return deepcopy(original_cached_response)
return original_cached_response
return cached_response
return None
@ -231,6 +236,8 @@ class InMemoryCache(BaseCache):
with self._increment_lock:
# keep read-modify-write atomic
init_value: Final = self.get_cache(key=key) or 0
if not isinstance(init_value, (int, float)):
raise TypeError("Cached value is not numeric")
value = init_value + value
self.set_cache(key, value, **kwargs)
return value

View file

@ -1,16 +1,9 @@
import asyncio
import json
import threading
import time
from concurrent.futures import ThreadPoolExecutor
from unittest.mock import MagicMock, patch
import httpx
import pytest
import respx
from fastapi.testclient import TestClient
from unittest.mock import AsyncMock
from pydantic import BaseModel
from litellm.caching.in_memory_cache import InMemoryCache
@ -21,6 +14,61 @@ class _SlowInt(int):
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
@ -38,6 +86,22 @@ def test_increment_cache_is_atomic_under_thread_concurrency():
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
@ -93,7 +157,7 @@ def test_in_memory_cache_ttl():
new_ttl_time = in_memory_cache.ttl_dict["new-fake-key"]
assert new_ttl_time is not None
time.sleep(1)
cached_obj = in_memory_cache.get_cache(key="new-fake-key")
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
@ -189,14 +253,9 @@ def test_in_memory_cache_eviction_order():
in_memory_cache = InMemoryCache(max_size_in_memory=2)
# Add items with different TTLs
now = time.time()
in_memory_cache.set_cache(
key="early_expire", value="value_1", ttl=100
) # expires in 100 seconds
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
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"]