mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 8c938fd3e4 into b781d157d7
This commit is contained in:
commit
21cee60d7c
2 changed files with 84 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:
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue