fix(cache): make in-memory and disk cache increments atomic (#34013)
Some checks failed
CodSpeed Benchmarks / benchmarks (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled

* 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:
yucheng-berri 2026-07-20 15:51:01 -07:00 • committed by GitHub
parent 583ddaf199
commit 432954a2ab
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 75 additions and 23 deletions

View file

@ -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()

View file

@ -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

View file

@ -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

View file

@ -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