mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(cache): make in-memory and disk cache increments atomic (#34013)
* fix(cache): make in-memory and disk increments atomic * refactor(cache): narrow in-memory increment lock scope * fix(cache): address follow-up review on increment tests/types * fix(cache): refresh atomic increment coverage * test(cache): widen increment race window with non-zero _SlowInt seed The zero seed was falsy, so InMemoryCache.increment_cache's `get_cache(...) or 0` and DiskCache.get_cache's truthiness guard both discarded the _SlowInt before __add__ could run, leaving the sleep-based window-widening inert. Seed a non-zero value and return _SlowInt from __add__ so the sleep fires on every read-modify-write in both backends, making the concurrency regression deterministic. * test(cache): cover InMemoryCache.async_increment delegation Add a focused async test asserting async_increment accumulates through the locked sync path, exercising the previously uncovered delegation line. --------- Co-authored-by: Emerson Gomes <emerson.gomes@thalesgroup.com>
This commit is contained in:
parent
583ddaf199
commit
432954a2ab
4 changed files with 75 additions and 23 deletions
|
|
@ -58,12 +58,12 @@ class DiskCache(BaseCache):
|
|||
return return_val
|
||||
|
||||
def increment_cache(self, key, value: int, **kwargs) -> int:
|
||||
# get the value
|
||||
cached_value = self.get_cache(key=key)
|
||||
init_value = cached_value if isinstance(cached_value, int) else 0
|
||||
value = init_value + value
|
||||
self.set_cache(key, value, **kwargs)
|
||||
return value
|
||||
with self.disk_cache.transact():
|
||||
cached_value = self.get_cache(key=key)
|
||||
init_value = cached_value if isinstance(cached_value, int) else 0
|
||||
new_value = init_value + value
|
||||
self.set_cache(key, new_value, **kwargs)
|
||||
return new_value
|
||||
|
||||
async def async_get_cache(self, key, **kwargs):
|
||||
return self.get_cache(key=key, **kwargs)
|
||||
|
|
@ -76,12 +76,7 @@ class DiskCache(BaseCache):
|
|||
return return_val
|
||||
|
||||
async def async_increment(self, key, value: int, **kwargs) -> int:
|
||||
# get the value
|
||||
cached_value = await self.async_get_cache(key=key)
|
||||
init_value = cached_value if isinstance(cached_value, int) else 0
|
||||
value = init_value + value
|
||||
await self.async_set_cache(key, value, **kwargs)
|
||||
return value
|
||||
return self.increment_cache(key=key, value=value, **kwargs)
|
||||
|
||||
def flush_cache(self):
|
||||
self.disk_cache.clear()
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import json
|
|||
import sys
|
||||
import time
|
||||
import heapq
|
||||
import threading
|
||||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -46,6 +47,7 @@ class InMemoryCache(BaseCache):
|
|||
self.cache_dict: dict = {}
|
||||
self.ttl_dict: dict = {}
|
||||
self.expiration_heap: list[tuple[float, str]] = []
|
||||
self._increment_lock = threading.Lock()
|
||||
|
||||
def check_value_size(self, value: Any):
|
||||
"""
|
||||
|
|
@ -223,12 +225,13 @@ class InMemoryCache(BaseCache):
|
|||
return_val.append(val)
|
||||
return return_val
|
||||
|
||||
def increment_cache(self, key, value: int, **kwargs) -> int:
|
||||
# get the value
|
||||
init_value = self.get_cache(key=key) or 0
|
||||
value = init_value + value
|
||||
self.set_cache(key, value, **kwargs)
|
||||
return value
|
||||
def increment_cache(self, key, value: float, **kwargs) -> float:
|
||||
with self._increment_lock:
|
||||
# keep read-modify-write atomic
|
||||
init_value = self.get_cache(key=key) or 0
|
||||
value = init_value + value
|
||||
self.set_cache(key, value, **kwargs)
|
||||
return value
|
||||
|
||||
async def async_get_cache(self, key, **kwargs):
|
||||
return self.get_cache(key=key, **kwargs)
|
||||
|
|
@ -241,11 +244,7 @@ class InMemoryCache(BaseCache):
|
|||
return return_val
|
||||
|
||||
async def async_increment(self, key, value: float, **kwargs) -> float:
|
||||
# get the value
|
||||
init_value = await self.async_get_cache(key=key) or 0
|
||||
value = init_value + value
|
||||
await self.async_set_cache(key, value, **kwargs)
|
||||
return value
|
||||
return self.increment_cache(key=key, value=value, **kwargs)
|
||||
|
||||
async def async_increment_pipeline(
|
||||
self, increment_list: List["RedisPipelineIncrementOperation"], **kwargs
|
||||
|
|
|
|||
|
|
@ -1,3 +1,7 @@
|
|||
import threading
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("diskcache")
|
||||
|
|
@ -5,6 +9,12 @@ pytest.importorskip("diskcache")
|
|||
from litellm.caching.disk_cache import DiskCache
|
||||
|
||||
|
||||
class _SlowInt(int):
|
||||
def __add__(self, value: int) -> "_SlowInt":
|
||||
time.sleep(0.05)
|
||||
return _SlowInt(int(self) + value)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def cache(tmp_path):
|
||||
return DiskCache(disk_cache_dir=str(tmp_path))
|
||||
|
|
@ -27,6 +37,22 @@ def test_increment_cache_treats_non_int_cached_value_as_zero(cache):
|
|||
assert cache.get_cache("counter") == 4
|
||||
|
||||
|
||||
def test_increment_cache_is_atomic_under_thread_concurrency(cache):
|
||||
seed = 1000
|
||||
cache.set_cache("counter", _SlowInt(seed))
|
||||
thread_count = 8
|
||||
barrier = threading.Barrier(thread_count)
|
||||
|
||||
def increment(_: int) -> int:
|
||||
barrier.wait()
|
||||
return cache.increment_cache("counter", 1)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=thread_count) as executor:
|
||||
tuple(executor.map(increment, range(thread_count)))
|
||||
|
||||
assert cache.get_cache("counter") == seed + thread_count
|
||||
|
||||
|
||||
async def test_async_increment_starts_from_zero_when_key_missing(cache):
|
||||
assert await cache.async_increment("counter", 2) == 2
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,9 @@ import asyncio
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -18,6 +20,36 @@ from unittest.mock import AsyncMock
|
|||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
|
||||
|
||||
class _SlowInt(int):
|
||||
def __add__(self, value: int) -> "_SlowInt":
|
||||
time.sleep(0.05)
|
||||
return _SlowInt(int(self) + value)
|
||||
|
||||
|
||||
def test_increment_cache_is_atomic_under_thread_concurrency():
|
||||
cache = InMemoryCache()
|
||||
seed = 1000
|
||||
cache.set_cache("counter", _SlowInt(seed))
|
||||
thread_count = 8
|
||||
barrier = threading.Barrier(thread_count)
|
||||
|
||||
def increment(_: int) -> float:
|
||||
barrier.wait()
|
||||
return cache.increment_cache("counter", 1)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=thread_count) as executor:
|
||||
tuple(executor.map(increment, range(thread_count)))
|
||||
|
||||
assert cache.get_cache("counter") == seed + thread_count
|
||||
|
||||
|
||||
async def test_async_increment_delegates_to_locked_sync_path():
|
||||
cache = InMemoryCache()
|
||||
assert await cache.async_increment("counter", 2) == 2
|
||||
assert await cache.async_increment("counter", 3) == 5
|
||||
assert cache.get_cache("counter") == 5
|
||||
|
||||
|
||||
def test_in_memory_openai_obj_cache():
|
||||
from openai import OpenAI
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue