mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge d098536b92 into 7204942756
This commit is contained in:
commit
06b1d0c004
5 changed files with 95 additions and 7 deletions
|
|
@ -8,13 +8,14 @@ Has 4 methods:
|
|||
- async_get_cache
|
||||
"""
|
||||
|
||||
import copy
|
||||
import heapq
|
||||
import json
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final, TypeVar, cast
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
|
|
@ -26,6 +27,7 @@ from litellm.constants import MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB
|
|||
from .base_cache import BaseCache
|
||||
|
||||
DEFAULT_MAX_SIZE_IN_MEMORY: Final = 200
|
||||
_T = TypeVar("_T")
|
||||
|
||||
|
||||
class InMemoryCache(BaseCache):
|
||||
|
|
@ -208,7 +210,25 @@ class InMemoryCache(BaseCache):
|
|||
return True
|
||||
return False
|
||||
|
||||
def get_cache(self, key, **kwargs):
|
||||
@staticmethod
|
||||
def _copy_cached_value(value: _T) -> _T:
|
||||
"""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 cast(_T, 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):
|
||||
if key in self.cache_dict:
|
||||
if self.evict_element_if_expired(key):
|
||||
return None
|
||||
|
|
@ -217,9 +237,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)
|
||||
|
|
|
|||
|
|
@ -138,6 +138,13 @@ class AlertingHangingRequestCheck:
|
|||
# flag so the entry is skipped on later ticks; one alert per hang,
|
||||
# with the existing TTL still handling cleanup
|
||||
hanging_request_data.alerted = True
|
||||
# InMemoryCache returns read-isolated mutable values. Persist the
|
||||
# state transition explicitly so the next checker tick observes
|
||||
# that this request has already been alerted.
|
||||
await self.hanging_request_cache.async_set_cache(
|
||||
key=request_id,
|
||||
value=hanging_request_data,
|
||||
)
|
||||
|
||||
return
|
||||
|
||||
|
|
|
|||
|
|
@ -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,54 @@ 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_cache_preserves_non_copyable_mutable_values():
|
||||
cache = InMemoryCache()
|
||||
cache.set_cache("payload", ["unchanged"])
|
||||
|
||||
with patch(
|
||||
"litellm.caching.in_memory_cache.copy.deepcopy",
|
||||
side_effect=RuntimeError("not copyable"),
|
||||
):
|
||||
cached = cache.get_cache("payload")
|
||||
|
||||
assert cached is cache.cache_dict["payload"]
|
||||
|
||||
|
||||
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