From 432954a2ab1d9c2b1ff4f38a677df6cbbc8d2a5a Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Mon, 20 Jul 2026 15:51:01 -0700 Subject: [PATCH] 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 --- litellm/caching/disk_cache.py | 19 ++++------- litellm/caching/in_memory_cache.py | 21 ++++++------ tests/test_litellm/caching/test_disk_cache.py | 26 +++++++++++++++ .../caching/test_in_memory_cache.py | 32 +++++++++++++++++++ 4 files changed, 75 insertions(+), 23 deletions(-) diff --git a/litellm/caching/disk_cache.py b/litellm/caching/disk_cache.py index d9f65ce949e..af8eb92849f 100644 --- a/litellm/caching/disk_cache.py +++ b/litellm/caching/disk_cache.py @@ -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() diff --git a/litellm/caching/in_memory_cache.py b/litellm/caching/in_memory_cache.py index 2ad3f3f11b7..36b477f7a8b 100644 --- a/litellm/caching/in_memory_cache.py +++ b/litellm/caching/in_memory_cache.py @@ -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 diff --git a/tests/test_litellm/caching/test_disk_cache.py b/tests/test_litellm/caching/test_disk_cache.py index b8d3b7b8d36..084370726b1 100644 --- a/tests/test_litellm/caching/test_disk_cache.py +++ b/tests/test_litellm/caching/test_disk_cache.py @@ -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 diff --git a/tests/test_litellm/caching/test_in_memory_cache.py b/tests/test_litellm/caching/test_in_memory_cache.py index 8828ebf207e..7be03d23fbe 100644 --- a/tests/test_litellm/caching/test_in_memory_cache.py +++ b/tests/test_litellm/caching/test_in_memory_cache.py @@ -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