mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix: isolate mutable in-memory cache reads
This commit is contained in:
parent
74cad08997
commit
7fb2c11eb5
2 changed files with 66 additions and 18 deletions
|
|
@ -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:
|
||||
|
|
@ -215,8 +216,12 @@ 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):
|
||||
# JSON strings already produce a new value. Isolate mutable data
|
||||
# stored as Python objects without copying cached clients or handles.
|
||||
if isinstance(original_cached_response, (dict, list, BaseModel)):
|
||||
return deepcopy(original_cached_response)
|
||||
return original_cached_response
|
||||
return cached_response
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -93,7 +141,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 +237,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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue