mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(cache): isolate mutable in-memory reads
This commit is contained in:
parent
f4308bc124
commit
da8769284b
4 changed files with 73 additions and 6 deletions
|
|
@ -8,6 +8,7 @@ Has 4 methods:
|
|||
- async_get_cache
|
||||
"""
|
||||
|
||||
import copy
|
||||
import heapq
|
||||
import json
|
||||
import sys
|
||||
|
|
@ -208,7 +209,25 @@ class InMemoryCache(BaseCache):
|
|||
return True
|
||||
return False
|
||||
|
||||
def get_cache(self, key, **kwargs):
|
||||
@staticmethod
|
||||
def _copy_cached_value(value: Any) -> Any:
|
||||
"""Return a read-isolated value without making the cache brittle.
|
||||
|
||||
In-memory cache entries can contain Pydantic models and mutable
|
||||
containers. Live resources such as SDK clients are deliberately left
|
||||
alone; attempting to deepcopy them can create partially initialized
|
||||
transports before failing. A failed deepcopy should not turn a cache
|
||||
hit into a request failure, so retain the existing value as a last
|
||||
resort.
|
||||
"""
|
||||
if not isinstance(value, (dict, list, set, tuple, frozenset, bytearray, BaseModel)):
|
||||
return value
|
||||
try:
|
||||
return copy.deepcopy(value)
|
||||
except Exception: # noqa: BLE001 - cache reads must tolerate non-copyable values
|
||||
return value
|
||||
|
||||
def _get_cache_value(self, key, *, copy_value: bool) -> Any:
|
||||
if key in self.cache_dict:
|
||||
if self.evict_element_if_expired(key):
|
||||
return None
|
||||
|
|
@ -217,9 +236,12 @@ class InMemoryCache(BaseCache):
|
|||
cached_response = json.loads(original_cached_response)
|
||||
except Exception:
|
||||
cached_response = original_cached_response
|
||||
return cached_response
|
||||
return self._copy_cached_value(cached_response) if copy_value else cached_response
|
||||
return None
|
||||
|
||||
def get_cache(self, key, **kwargs):
|
||||
return self._get_cache_value(key=key, copy_value=True)
|
||||
|
||||
def batch_get_cache(self, keys: list, **kwargs):
|
||||
return_val: Final = []
|
||||
for k in keys:
|
||||
|
|
|
|||
|
|
@ -72,7 +72,9 @@ class LLMClientCache(InMemoryCache):
|
|||
key = self.update_cache_key_with_event_loop(key)
|
||||
self.evicted_client_closer.reap()
|
||||
|
||||
return super().get_cache(key, **kwargs)
|
||||
# SDK/httpx clients are live resources. Returning a deepcopy would
|
||||
# detach the caller from the client tracked by the eviction closer.
|
||||
return super()._get_cache_value(key=key, copy_value=False)
|
||||
|
||||
async def async_get_cache(self, key, **kwargs):
|
||||
key = self.update_cache_key_with_event_loop(key)
|
||||
|
|
|
|||
|
|
@ -3,14 +3,13 @@ import json
|
|||
import threading
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import AsyncMock, 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
|
||||
|
||||
|
|
@ -45,6 +44,41 @@ async def test_async_increment_delegates_to_locked_sync_path():
|
|||
assert cache.get_cache("counter") == 5
|
||||
|
||||
|
||||
def test_in_memory_cache_returns_read_isolated_mutable_values():
|
||||
cache = InMemoryCache()
|
||||
cache.set_cache("user", {"spend": 1.0, "metadata": {"region": "us"}})
|
||||
|
||||
cached = cache.get_cache("user")
|
||||
assert cached == {"spend": 1.0, "metadata": {"region": "us"}}
|
||||
assert cached is not cache.cache_dict["user"]
|
||||
|
||||
cached["spend"] = 50
|
||||
cached["metadata"]["region"] = "eu"
|
||||
|
||||
assert cache.get_cache("user") == {"spend": 1.0, "metadata": {"region": "us"}}
|
||||
|
||||
|
||||
def test_in_memory_cache_returns_read_isolated_pydantic_models():
|
||||
class CachedBudget(BaseModel):
|
||||
spend: float
|
||||
metadata: dict[str, str]
|
||||
|
||||
cache = InMemoryCache()
|
||||
cache.set_cache("budget", CachedBudget(spend=1.0, metadata={"region": "us"}))
|
||||
|
||||
cached = cache.get_cache("budget")
|
||||
assert isinstance(cached, CachedBudget)
|
||||
assert cached is not cache.cache_dict["budget"]
|
||||
|
||||
cached.spend = 50
|
||||
cached.metadata["region"] = "eu"
|
||||
|
||||
stored = cache.get_cache("budget")
|
||||
assert isinstance(stored, CachedBudget)
|
||||
assert stored.spend == 1.0
|
||||
assert stored.metadata == {"region": "us"}
|
||||
|
||||
|
||||
def test_in_memory_openai_obj_cache():
|
||||
from openai import OpenAI
|
||||
|
||||
|
|
|
|||
|
|
@ -38,6 +38,15 @@ class MockSyncClient:
|
|||
self.closed = True
|
||||
|
||||
|
||||
def test_client_cache_returns_live_client_reference():
|
||||
cache = LLMClientCache()
|
||||
client = MockSyncClient()
|
||||
|
||||
cache.set_cache("client", client)
|
||||
|
||||
assert cache.get_cache("client") is client
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remove_key_does_not_close_async_client():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue