diff --git a/litellm/caching/in_memory_cache.py b/litellm/caching/in_memory_cache.py index 56c9147e066..88708efcf14 100644 --- a/litellm/caching/in_memory_cache.py +++ b/litellm/caching/in_memory_cache.py @@ -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 diff --git a/tests/unit/caching/test_in_memory_cache.py b/tests/unit/caching/test_in_memory_cache.py index 40ad4f0c6f0..9ace9992ba9 100644 --- a/tests/unit/caching/test_in_memory_cache.py +++ b/tests/unit/caching/test_in_memory_cache.py @@ -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"]