From d3310db53ab88c7a0edda69a484f86e40ed07fd4 Mon Sep 17 00:00:00 2001 From: Shivansh Shukla Date: Mon, 28 Sep 2026 10:45:44 +0530 Subject: [PATCH] feat(cache-manager): add stampede protection to get_or_set Introduce opt-in lock-based stampede protection in CacheManager.get_or_set() with double-check, polling backoff, re-entrancy prevention, crashed winner recovery, and comprehensive multi-backend tests. --- changelog.d/66.added.md | 6 + docs/APP_CACHE.md | 39 +- fastapi_cachex/exceptions.py | 4 +- fastapi_cachex/manager.py | 269 +++++++++++-- i18n/zh-TW/docs/APP_CACHE.md | 31 +- tests/conftest.py | 10 + tests/test_cache_manager.py | 736 ++++++++++++++++++++++++++++++++++- 7 files changed, 1062 insertions(+), 33 deletions(-) create mode 100644 changelog.d/66.added.md diff --git a/changelog.d/66.added.md b/changelog.d/66.added.md new file mode 100644 index 0000000..41f3dfc --- /dev/null +++ b/changelog.d/66.added.md @@ -0,0 +1,6 @@ +**`CacheManager.get_or_set()` supports lock-based stampede protection.** Pass +`lock=True` (or configure `lock=True` on `CacheManager`) to coordinate +concurrent misses for the same key through `CacheLock` so only one caller +executes `factory` while others wait for the cached value. `LockTimeoutError` +now also subclasses the standard library `TimeoutError`, so `except TimeoutError` +catches it. diff --git a/docs/APP_CACHE.md b/docs/APP_CACHE.md index 5d94860..4f0a1e4 100644 --- a/docs/APP_CACHE.md +++ b/docs/APP_CACHE.md @@ -45,8 +45,11 @@ Complete runnable example: [`examples/app_cache.py`](https://github.com/allen009 - `get()` returns `None` (or a supplied `default=`) on a cache miss — it never raises for missing or corrupted entries. - `set()` lets `TypeError` propagate for values that are not JSON-serializable. -- `get_or_set()` provides no stampede protection: concurrent misses for the same - key each run `factory`. +- `get_or_set()` supports opt-in stampede protection via `lock=True` (or + manager-level default `CacheManager(lock=True)`), preventing concurrent misses + from running `factory` simultaneously, with graceful fallback to computing + directly if waiting times out. The distributed lock key is + `lock:` (`lock:cache:user:42` by default). - `get_or_set()` returns the JSON-decoded value on a miss as well as on a hit (see [JSON round-trip](#json-round-trip)), so both paths give the same result. - `add()` stores a value only when the key is free and returns whether it did. @@ -86,6 +89,38 @@ Complete runnable example: [`examples/app_cache.py`](https://github.com/allen009 > `get()`/`set()`/`add()`/`delete()`/`has()` work normally. Use Redis or the in-memory > backend if you need bulk clearing. +## Stampede protection + +When the factory is expensive (a slow database query, a rate-limited upstream +API) and the key is hot, cache expiry turns into simultaneous recomputations. +You can enable distributed stampede protection built on `CacheLock` either +per-call or manager-wide: + +```python +# Per-call protection: +profile = await manager.get_or_set( + "user:42", + lambda: load_user(42), + ttl=300, + lock=True, + lock_ttl=30, # lease upper bound for factory run (default: 60) + wait_timeout=10, # caller latency budget in seconds (default: None, wait while held) + raise_on_timeout=False, # True raises LockTimeoutError, False falls back to factory (default: False) +) + +# Or manager-wide default: +manager = CacheManager(lock=True, lock_ttl=60) +``` + +1. **Miss**: On a miss, callers attempt non-blocking lock acquisition using `CacheLock` under the key `lock:` (`lock:cache:user:42` by default). +2. **Winner**: The winner re-checks the cache, invokes `factory`, stores the value in the backend, and releases the lock. +3. **Waiters**: Other callers poll the cache with exponential backoff (50ms base, 1.5x factor, 500ms cap) and randomized jitter (±10%) until the value appears. When `wait_timeout` is omitted, callers wait bounded by the holder's `lock_ttl` without an arbitrary deadline. +4. **Takeover**: If the winner fails or its lock expires, a waiting caller takes over the lock, re-checks the cache, and computes if necessary. +5. **Re-entrancy**: Recursive calls to `get_or_set()` for the same key in the same task automatically skip locking to avoid self-deadlock. +6. **Timeouts**: When an explicit `wait_timeout` elapses, `raise_on_timeout=False` logs a warning and falls back to direct factory computation (graceful degradation), while `raise_on_timeout=True` raises `LockTimeoutError`. + +Ensure `lock_ttl` exceeds the expected execution time of `factory`. If `factory` outlives `lock_ttl`, the lock expires mid-run and a waiting caller may start a second computation. + ## JSON round-trip Values are stored as JSON (`json.dumps` with its defaults) and read back with diff --git a/fastapi_cachex/exceptions.py b/fastapi_cachex/exceptions.py index c79f859..ead9a89 100644 --- a/fastapi_cachex/exceptions.py +++ b/fastapi_cachex/exceptions.py @@ -40,8 +40,8 @@ class RequestNotFoundError(CacheXError): """Exception raised when a request is not found.""" -class LockTimeoutError(CacheXError): - """Exception raised when acquiring a lock times out.""" +class LockTimeoutError(CacheXError, TimeoutError): + """Exception raised when acquiring a lock or waiting for stampede protection times out.""" def __getattr__(name: str) -> type[CacheXError]: diff --git a/fastapi_cachex/manager.py b/fastapi_cachex/manager.py index 1d06db6..9e22047 100644 --- a/fastapi_cachex/manager.py +++ b/fastapi_cachex/manager.py @@ -1,10 +1,14 @@ """Generic application-level cache manager for FastAPI-CacheX.""" +import asyncio +import contextvars import fnmatch import hashlib import inspect import json import logging +import secrets +import time import warnings from collections.abc import Awaitable from collections.abc import Callable @@ -12,6 +16,8 @@ from .backends.base import BaseCacheBackend from .backends.base import validate_ttl +from .exceptions import LockTimeoutError +from .lock import CacheLock from .proxy import BackendProxy from .types import CacheEntry from .types import log_ref @@ -25,6 +31,62 @@ # holding one still cannot be passed through to Redis unescaped). _GLOB_SPECIAL = frozenset("*?[]\\") +# Keys currently held under a stampede-protection lock by the running task. +# Stored as a set of unique backend-and-key identifiers to skip locking on +# re-entrancy within the same task. +_HELD_LOCKS: contextvars.ContextVar[frozenset[str]] = contextvars.ContextVar( + "_HELD_LOCKS", default=frozenset() +) + +# Polling configuration for stampede protection waiters +_INITIAL_POLL_INTERVAL: float = 0.05 +_MAX_POLL_INTERVAL: float = 0.5 +_BACKOFF_FACTOR: float = 1.5 +_JITTER_RATIO: float = 0.1 + +_system_random = secrets.SystemRandom() +_SENTINEL = object() + +# Hooks for tests to inject artificial clocks and sleeping +_sleep: Callable[[float], Awaitable[None]] = asyncio.sleep +_monotonic: Callable[[], float] = time.monotonic + + +def _validate_lock(lock: Any, *, allow_none: bool = False) -> None: + if lock is None: + if allow_none: + return + msg = "lock must be a bool, got None" + raise TypeError(msg) + if not isinstance(lock, bool): + expected = "a bool or None" if allow_none else "a bool" + msg = f"lock must be {expected}, got {type(lock).__name__}" + raise TypeError(msg) + + +def _validate_get_or_set_args( + lock: bool | None, + ttl: int | None, + lock_ttl: int | None, + wait_timeout: float | None, +) -> float | None: + _validate_lock(lock, allow_none=True) + validate_ttl(ttl) + if lock_ttl is not None: + validate_ttl(lock_ttl) + if wait_timeout is not None: + if isinstance(wait_timeout, bool) or not isinstance(wait_timeout, (int, float)): + msg = ( + f"wait_timeout must be a number of seconds or None, " + f"got {type(wait_timeout).__name__}" + ) + raise TypeError(msg) + if wait_timeout <= 0: + msg = f"wait_timeout must be greater than zero, got {wait_timeout!r}" + raise ValueError(msg) + return float(wait_timeout) + return None + class CacheManager: """Provides convenient get/set/delete access to the configured cache backend. @@ -40,6 +102,9 @@ def __init__( backend: BaseCacheBackend | None = None, key_prefix: str = "cache:", default_ttl: int | None = None, + *, + lock: bool = False, + lock_ttl: int = 60, ) -> None: r"""Initialize CacheManager. @@ -48,11 +113,15 @@ def __init__( key_prefix: Prefix prepended to all logical keys in the cache backend. default_ttl: Default TTL (seconds) applied when set() is called without an explicit ttl. None means no expiry by default. + lock: Whether get_or_set() uses distributed locking by default to + prevent cache stampedes (default: False). + lock_ttl: Default TTL in seconds for stampede protection locks (default: 60). Raises: BackendNotFoundError: If ``backend`` is None and no backend has been set with ``BackendProxy.set()``. - ValueError: If ``default_ttl`` is zero or negative. + TypeError: If ``lock`` is not a bool or ``lock_ttl`` is not an int. + ValueError: If ``default_ttl`` or ``lock_ttl`` is zero or negative. Warns: UserWarning: If ``key_prefix`` contains a glob metacharacter @@ -63,6 +132,16 @@ def __init__( self.backend = backend if backend is not None else BackendProxy.get() self.key_prefix = key_prefix self.default_ttl = validate_ttl(default_ttl) + + _validate_lock(lock, allow_none=False) + + effective_lock_ttl = validate_ttl(lock_ttl) + if effective_lock_ttl is None: + msg = "lock_ttl must be a positive int, got None" + raise ValueError(msg) + self.lock = lock + self.lock_ttl: int = effective_lock_ttl + if self._prefix_has_glob: warnings.warn( f"CacheManager key_prefix {key_prefix!r} contains a glob " @@ -188,18 +267,133 @@ async def has(self, key: str) -> bool: """ return await self.backend.get(self._cache_key(key)) is not None - async def get_or_set( + async def _compute_and_store( + self, + key: str, + factory: Callable[[], Any] | Callable[[], Awaitable[Any]], + ttl: int | None, + ) -> Any: + value = factory() + if inspect.isawaitable(value): + value = await value + + effective_ttl = ttl if ttl is not None else self.default_ttl + entry = self._encode(value) + await self.backend.set(self._cache_key(key), entry, ttl=effective_ttl) + logger.debug("Cache SET; key=%s ttl=%s", key, effective_ttl) + return json.loads(entry.content) + + async def _execute_as_winner( + self, + key: str, + factory: Callable[[], Any] | Callable[[], Awaitable[Any]], + ttl: int | None, + lock_instance: CacheLock, + ) -> Any: + lock_id = f"{id(self.backend)}:{self._cache_key(key)}" + token = _HELD_LOCKS.set(_HELD_LOCKS.get() | {lock_id}) + try: + cached = await self.get(key, default=_SENTINEL) + if cached is not _SENTINEL: + return cached + return await self._compute_and_store(key, factory, ttl) + finally: + try: + try: + await lock_instance.release() + except Exception: + logger.warning( + "Failed to release stampede protection lock for key_ref=%s", + log_ref(key), + exc_info=True, + ) + finally: + _HELD_LOCKS.reset(token) + + async def _poll_for_value( # noqa: PLR0913 + self, + key: str, + factory: Callable[[], Any] | Callable[[], Awaitable[Any]], + ttl: int | None, + lock_ttl: int, + wait_timeout: float | None, + *, + raise_on_timeout: bool, + ) -> Any: + start = _monotonic() + interval = _INITIAL_POLL_INTERVAL + while True: + if wait_timeout is not None: + elapsed = _monotonic() - start + if elapsed >= wait_timeout: + cached = await self.get(key, default=_SENTINEL) + if cached is not _SENTINEL: + return cached + if raise_on_timeout: + msg = f"Waiting for cache key {key!r} timed out after {wait_timeout}s" + raise LockTimeoutError(msg) + logger.warning( + "Cache stampede wait timeout exceeded for key_ref=%s; computing directly", + log_ref(key), + ) + return await self._compute_and_store(key, factory, ttl) + + jitter = _system_random.uniform( + -interval * _JITTER_RATIO, interval * _JITTER_RATIO + ) + sleep_time = max(0.0, interval + jitter) + if wait_timeout is not None: + remaining = wait_timeout - elapsed + sleep_time = min(sleep_time, remaining) + await _sleep(sleep_time) + + interval = min(interval * _BACKOFF_FACTOR, _MAX_POLL_INTERVAL) + + # 1. Check cache, return on hit + cached = await self.get(key, default=_SENTINEL) + if cached is not _SENTINEL: + return cached + + # 2. Otherwise try to take lock + lock_instance = CacheLock( + name=self._cache_key(key), + ttl=lock_ttl, + backend=self.backend, + ) + if await lock_instance.acquire(blocking=False): + return await self._execute_as_winner(key, factory, ttl, lock_instance) + + async def get_or_set( # noqa: PLR0913 self, key: str, factory: Callable[[], Any] | Callable[[], Awaitable[Any]], ttl: int | None = None, + *, + lock: bool | None = None, + lock_ttl: int | None = None, + wait_timeout: float | None = None, + raise_on_timeout: bool = False, ) -> Any: """Get a cached value, computing and storing it via ``factory`` on a miss. ``factory`` is only invoked when ``key`` is missing, expired, or its stored content cannot be decoded; on a hit the cached value is - returned directly. This method does not provide stampede protection: - concurrent misses for the same key may each invoke ``factory``. + returned directly. + + When stampede protection is enabled (via ``lock=True`` or the + manager's ``lock`` default), concurrent misses for the same key acquire + a distributed lock built on :class:`~fastapi_cachex.lock.CacheLock`. + The winner re-checks the cache, invokes ``factory``, stores the result, + and releases the lock. Waiting callers poll the cache with exponential + backoff and jitter until the value appears, taking over the lock if the + winner fails or the lock lease expires. + + Re-entrant calls to ``get_or_set()`` for the same key within the + current task skip locking to prevent self-deadlock. + + Ensure ``lock_ttl`` exceeds the expected execution time of ``factory``. + If ``factory`` outlives ``lock_ttl``, the lock expires mid-run and a + waiting caller may start a second computation. A miss returns the value as it will be read back from the cache, not the object ``factory`` returned: it goes through the same JSON @@ -216,36 +410,63 @@ async def get_or_set( the result is awaited. ttl: Time-to-live in seconds for a newly created value. If None, uses ``self.default_ttl``. + lock: Whether to use distributed locking for stampede protection. + If None, inherits the manager's ``lock`` setting. + lock_ttl: Upper bound in seconds for the lock lease. If None, + inherits the manager's ``lock_ttl``. + wait_timeout: Maximum seconds waiting callers poll the cache before + timing out. If None, callers wait without a fixed deadline, + bounded by the holder's lock lease, and attempt to take over the + lock if it expires. + raise_on_timeout: If True, raise :exc:`~fastapi_cachex.exceptions.LockTimeoutError` + when ``wait_timeout`` elapses. If False (default), log a warning and + fall back to invoking ``factory`` directly. Returns: The cached value (existing or newly created), JSON-decoded in both cases. Raises: - TypeError: If the value produced by ``factory`` is not JSON-serializable. - ValueError: If ``ttl`` is zero or negative. + TypeError: If the value produced by ``factory`` is not JSON-serializable, + or if ``lock``, ``ttl``, ``lock_ttl``, or ``wait_timeout`` + have invalid types. + ValueError: If ``ttl``, ``lock_ttl``, or ``wait_timeout`` is zero or negative. + LockTimeoutError: If ``raise_on_timeout=True`` and waiting exceeds ``wait_timeout``. """ - # Reject a bad ttl before the factory does any (possibly costly) work. - validate_ttl(ttl) - sentinel = object() - cached = await self.get(key, default=sentinel) - if cached is not sentinel: + validated_wait_timeout = _validate_get_or_set_args( + lock, ttl, lock_ttl, wait_timeout + ) + + cached = await self.get(key, default=_SENTINEL) + if cached is not _SENTINEL: return cached - # Await whatever comes back awaitable, not just from coroutine - # functions: `lambda: load(42)` and `functools.partial` return one too. - value = factory() - if inspect.isawaitable(value): - value = await value + use_lock = self.lock if lock is None else lock + if not use_lock: + return await self._compute_and_store(key, factory, ttl) - # Encode once, store those bytes and return them decoded, so a miss - # returns exactly what a later hit will: a tuple comes back as a - # list, int dict keys as strings. - effective_ttl = ttl if ttl is not None else self.default_ttl - entry = self._encode(value) - await self.backend.set(self._cache_key(key), entry, ttl=effective_ttl) - logger.debug("Cache SET; key=%s ttl=%s", key, effective_ttl) - return json.loads(entry.content) + lock_id = f"{id(self.backend)}:{self._cache_key(key)}" + if lock_id in _HELD_LOCKS.get(): + logger.debug("Re-entrant get_or_set call for key=%s; skipping lock", key) + return await self._compute_and_store(key, factory, ttl) + + effective_lock_ttl = self.lock_ttl if lock_ttl is None else lock_ttl + lock_instance = CacheLock( + name=self._cache_key(key), + ttl=effective_lock_ttl, + backend=self.backend, + ) + if await lock_instance.acquire(blocking=False): + return await self._execute_as_winner(key, factory, ttl, lock_instance) + + return await self._poll_for_value( + key, + factory, + ttl, + effective_lock_ttl, + validated_wait_timeout, + raise_on_timeout=raise_on_timeout, + ) async def clear_pattern(self, pattern: str) -> int: r"""Clear all keys under this manager's namespace matching a glob pattern. diff --git a/i18n/zh-TW/docs/APP_CACHE.md b/i18n/zh-TW/docs/APP_CACHE.md index 32603e1..f1f972a 100644 --- a/i18n/zh-TW/docs/APP_CACHE.md +++ b/i18n/zh-TW/docs/APP_CACHE.md @@ -42,7 +42,7 @@ await manager.clear_pattern("user:*") # 比對 "myapp:user:*" - `get()` 在快取未命中時回傳 `None`(或你提供的 `default=`),遇到不存在或損毀的項目也絕不會拋出例外。 - `set()` 遇到無法 JSON 序列化的值時,會讓 `TypeError` 直接往外拋出。 -- `get_or_set()` 不提供 cache stampede 保護:同一個鍵同時發生多次未命中時,每一次都會執行 `factory`。 +- `get_or_set()` 支援以 `lock=True`(或 manager 全域設定 `CacheManager(lock=True)`)選用 cache stampede 保護,避免多次並行未命中時同時執行 `factory`,若等待超時則具備直接計算的優雅降級回退。分散式鎖的鍵名格式為 `lock:`(預設為 `lock:cache:user:42`)。 - `get_or_set()` 在未命中與命中時都回傳經 JSON 解碼後的值(見 [JSON 往返](#json-round-trip)),因此兩條路徑的結果相同。 - `add()` 只在鍵尚未被占用時寫入值,並回傳是否有寫入。檢查與寫入是同一個後端原子操作(`set_if_absent`),因此適合「每個鍵只做一次」的工作,例如 webhook 或電子郵件的去重。已過期的鍵視為未被占用;存放無法解碼之值的鍵則不算,即使 `get()` 會把它當成未命中。 - 鍵預設位於獨立、以 `cache:` 為前綴的命名空間,與 HTTP 路由快取及 OAuth state 分開,因此 `clear()`/`clear_prefix()` 絕不會動到無關的快取項目。 @@ -53,6 +53,35 @@ await manager.clear_pattern("user:*") # 比對 "myapp:user:*" > [!NOTE] > `clear()`/`clear_prefix()` 是以後端的 `get_all_keys()` 與 `delete_many()` 實作(在 Redis 上是一次批次 `DEL`)。由於 Memcached 不支援列舉鍵(見[後端](BACKENDS.md#memcached)),這些方法以及 `clear_pattern()` 在 Memcached 後端上不會有任何作用;`get()`/`set()`/`add()`/`delete()`/`has()` 則照常運作。若需要大量清除,請使用 Redis 或記憶體後端。 +## Cache stampede 保護 {#stampede-protection} + +當 `factory` 的運算成本很高(例如慢速資料庫查詢、受速率限制的外部 API)且該鍵又是熱門鍵時,快取過期會導致多個請求同時重新計算。你可以透過 `CacheLock` 在單次呼叫或 manager 全域啟用分散式 cache stampede 保護: + +```python +# 單次呼叫保護: +profile = await manager.get_or_set( + "user:42", + lambda: load_user(42), + ttl=300, + lock=True, + lock_ttl=30, # factory 執行的租約上限(秒,預設:60) + wait_timeout=10, # 呼叫者等待的延遲預算(秒,預設:None,持鎖期間持續等待) + raise_on_timeout=False, # True 拋出 LockTimeoutError,False 回退至執行 factory(預設:False) +) + +# 或 manager 全域預設: +manager = CacheManager(lock=True, lock_ttl=60) +``` + +1. **未命中**:未命中時,呼叫者嘗試使用 `CacheLock` 進行非阻塞式取鎖,鎖鍵名稱為 `lock:`(預設為 `lock:cache:user:42`)。 +2. **勝出者**:取鎖成功的勝出者會再次檢查快取、執行 `factory`、將值存入後端並釋放鎖。 +3. **等待者**:其他呼叫者以指數退避(基準 50ms、倍率 1.5 倍、上限 500ms)與隨機抖動(±10%)輪詢快取,直到值出現為止。未指定 `wait_timeout` 時,呼叫者受限於持鎖者的 `lock_ttl` 期間持續等待,不會有固定的期限。 +4. **接手**:若勝出者失敗或其鎖已過期,等待中的呼叫者會接手取得鎖、重新檢查快取並在需要時執行計算。 +5. **可重新進入(Re-entrancy)**:同一工作(task)內對同一個鍵遞迴呼叫 `get_or_set()` 會自動跳過取鎖,避免自我死鎖。 +6. **逾時**:當明確指定的 `wait_timeout` 到期時,`raise_on_timeout=False` 會記錄警告並回退為直接執行 factory 計算(優雅降級),而 `raise_on_timeout=True` 則拋出 `LockTimeoutError`。 + +請確保 `lock_ttl` 超過 `factory` 的預期執行時間。若 `factory` 執行時間超過 `lock_ttl`,鎖會在執行途中過期,導致等待中的呼叫者發起第二次計算。 + ## JSON 往返 {#json-round-trip} 值以 JSON 儲存(`json.dumps` 的預設設定),讀取時以 `json.loads` 解碼,因此取回的是 JSON 解碼後的形式,而不是你存入的物件: diff --git a/tests/conftest.py b/tests/conftest.py index 76dc170..cdfdf3b 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -9,6 +9,7 @@ import pytest import pytest_asyncio +from fastapi_cachex import manager from fastapi_cachex.backends import MemcachedBackend from fastapi_cachex.backends import memory from fastapi_cachex.backends.base import BaseCacheBackend @@ -106,9 +107,16 @@ def __init__(self) -> None: def time(self) -> float: return self.now + def monotonic(self) -> float: + return self.now + def advance(self, seconds: float) -> None: self.now += seconds + async def sleep(self, seconds: float) -> None: + self.advance(seconds) + await asyncio.sleep(0) + async def wait(self, seconds: float, backend: BaseCacheBackend) -> None: """Let `seconds` pass as far as `backend` can tell. @@ -135,6 +143,8 @@ def now(cls, tz: tzinfo | None = None) -> datetime: # type: ignore[override] return datetime.fromtimestamp(clock.now, tz) monkeypatch.setattr(memory, "time", SimpleNamespace(time=clock.time)) + monkeypatch.setattr(manager, "_monotonic", clock.monotonic) + monkeypatch.setattr(manager, "_sleep", clock.sleep) for module in (session_manager, session_models, state_manager, state_models): monkeypatch.setattr(module, "datetime", ClockDatetime) return clock diff --git a/tests/test_cache_manager.py b/tests/test_cache_manager.py index 4dac446..014b890 100644 --- a/tests/test_cache_manager.py +++ b/tests/test_cache_manager.py @@ -4,30 +4,33 @@ import logging import time from collections.abc import AsyncGenerator +from contextlib import asynccontextmanager from functools import partial -from typing import TYPE_CHECKING from typing import Any import pytest import pytest_asyncio +from fastapi_cachex.backends.base import BaseCacheBackend from fastapi_cachex.backends.memory import MemoryBackend from fastapi_cachex.dependencies import get_app_cache from fastapi_cachex.exceptions import BackendNotFoundError +from fastapi_cachex.exceptions import LockTimeoutError +from fastapi_cachex.lock import CacheLock from fastapi_cachex.manager import CacheManager from fastapi_cachex.manager_proxy import CacheManagerProxy from fastapi_cachex.proxy import BackendProxy from fastapi_cachex.types import CacheEntry from fastapi_cachex.types import log_ref from tests.conftest import Clock +from tests.live_servers import MEMCACHED_SERVER from tests.live_servers import REDIS_HOST from tests.live_servers import REDIS_PORT +from tests.live_servers import flush_memcached +from tests.live_servers import requires_memcached from tests.live_servers import requires_redis from tests.live_servers import requires_redis_package -if TYPE_CHECKING: - from fastapi_cachex.backends.base import BaseCacheBackend - @pytest_asyncio.fixture( params=[ @@ -713,3 +716,728 @@ def test_get_app_cache_reuses_existing_proxy_instance( assert get_app_cache() is existing finally: CacheManagerProxy.set(None) + + +# --- Stampede protection (get_or_set) -------------------------------------------- + + +async def test_stampede_protection_default_off( + memory_backend: MemoryBackend, +) -> None: + """By default (lock=False), concurrent misses for the same key each invoke factory.""" + manager = CacheManager(backend=memory_backend) + calls = 0 + + async def factory() -> str: + nonlocal calls + calls += 1 + await asyncio.sleep(0.02) + return "data" + + results = await asyncio.gather( + manager.get_or_set("key", factory), + manager.get_or_set("key", factory), + ) + assert list(results) == ["data", "data"] + assert calls == 2 + + +async def test_stampede_protection_explicit_lock_false_overrides_manager( + memory_backend: MemoryBackend, +) -> None: + """An explicit lock=False on get_or_set overrides manager.lock=True.""" + manager = CacheManager(backend=memory_backend, lock=True) + calls = 0 + + async def factory() -> str: + nonlocal calls + calls += 1 + await asyncio.sleep(0.02) + return "data" + + results = await asyncio.gather( + manager.get_or_set("key", factory, lock=False), + manager.get_or_set("key", factory, lock=False), + ) + assert list(results) == ["data", "data"] + assert calls == 2 + + +async def test_stampede_protection_manager_default_lock( + memory_backend: MemoryBackend, +) -> None: + """When manager has lock=True, get_or_set without lock parameter uses locking.""" + manager = CacheManager(backend=memory_backend, lock=True) + calls = 0 + + async def factory() -> str: + nonlocal calls + calls += 1 + await asyncio.sleep(0.02) + return "computed" + + results = await asyncio.gather( + manager.get_or_set("key", factory), + manager.get_or_set("key", factory), + manager.get_or_set("key", factory), + ) + assert list(results) == ["computed", "computed", "computed"] + assert calls == 1 + + +async def test_stampede_protection_single_winner_concurrent_misses( + memory_backend: MemoryBackend, +) -> None: + """Concurrent calls with lock=True execute factory once, and all callers get the value.""" + manager = CacheManager(backend=memory_backend) + calls = 0 + + async def slow_factory() -> dict[str, int]: + nonlocal calls + calls += 1 + await asyncio.sleep(0.03) + return {"count": 42} + + results = await asyncio.gather( + *[manager.get_or_set("hot_key", slow_factory, lock=True) for _ in range(8)] + ) + assert all(r == {"count": 42} for r in results) + assert calls == 1 + assert await manager.get("hot_key") == {"count": 42} + + +async def test_stampede_protection_winner_rechecks_cache( + memory_backend: MemoryBackend, +) -> None: + """Winner re-checks the cache after acquiring lock and does not run factory if populated.""" + manager = CacheManager(backend=memory_backend) + calls = 0 + + def factory() -> str: + nonlocal calls + calls += 1 + return "from_factory" + + # Pre-populate cache directly before get_or_set + await manager.set("recheck_key", "pre_existing") + + result = await manager.get_or_set("recheck_key", factory, lock=True) + assert result == "pre_existing" + assert calls == 0 + + +async def test_stampede_protection_winner_exception_releases_lock( + memory_backend: MemoryBackend, +) -> None: + """If the winner's factory raises, the lock is released in finally, allowing a waiter to take over.""" + manager = CacheManager(backend=memory_backend) + calls = 0 + + async def failing_factory() -> str: + nonlocal calls + calls += 1 + msg = "factory failure" + raise ValueError(msg) + + async def succeeding_factory() -> str: + nonlocal calls + calls += 1 + await asyncio.sleep(0.01) + return "recovered" + + task1 = asyncio.create_task( + manager.get_or_set("fail_key", failing_factory, lock=True) + ) + # Ensure task1 starts first and acquires lock + await asyncio.sleep(0.005) + task2 = asyncio.create_task( + manager.get_or_set("fail_key", succeeding_factory, lock=True) + ) + + with pytest.raises(ValueError, match="factory failure"): + await task1 + + result2 = await task2 + assert result2 == "recovered" + assert calls == 2 + assert await manager.get("fail_key") == "recovered" + + +async def test_stampede_protection_crashed_winner_timeout_fallback( + memory_backend: MemoryBackend, + clock: Clock, + caplog: pytest.LogCaptureFixture, +) -> None: + """If winner holds lock indefinitely, waiter times out, logs warning, and computes value.""" + manager = CacheManager(backend=memory_backend) + + # Acquire lock directly, simulating a crashed winner that never released + lock = CacheLock( + name=manager._cache_key("stuck_key"), + ttl=10, + backend=memory_backend, + ) + assert await lock.acquire(blocking=False) is True + + calls = 0 + + def factory() -> str: + nonlocal calls + calls += 1 + return "fallback_value" + + with caplog.at_level(logging.WARNING, logger="fastapi_cachex.manager"): + result = await manager.get_or_set( + "stuck_key", + factory, + lock=True, + wait_timeout=1.0, + raise_on_timeout=False, + ) + + assert result == "fallback_value" + assert calls == 1 + assert any( + "Cache stampede wait timeout exceeded" in r.getMessage() for r in caplog.records + ) + await lock.release() + + +async def test_stampede_protection_crashed_winner_timeout_raises( + memory_backend: MemoryBackend, + clock: Clock, +) -> None: + """If raise_on_timeout=True and wait_timeout elapses, LockTimeoutError is raised.""" + manager = CacheManager(backend=memory_backend) + + # Simulate crashed lock holder + lock = CacheLock( + name=manager._cache_key("timeout_key"), + ttl=10, + backend=memory_backend, + ) + assert await lock.acquire(blocking=False) is True + + with pytest.raises(LockTimeoutError, match="timed out") as exc_info: + await manager.get_or_set( + "timeout_key", + lambda: "value", + lock=True, + wait_timeout=1.0, + raise_on_timeout=True, + ) + + assert isinstance(exc_info.value, TimeoutError) + assert isinstance(exc_info.value, LockTimeoutError) + await lock.release() + + +async def test_stampede_protection_crashed_winner_single_takeover( + memory_backend: MemoryBackend, +) -> None: + """When a winner crashes with default wait_timeout=None, waiters wait and exactly one takes over.""" + manager = CacheManager(backend=memory_backend, lock=True, lock_ttl=1) + calls = 0 + + async def factory() -> str: + nonlocal calls + calls += 1 + return "recovered" + + # Simulate crashed lock holder + crashed_lock = CacheLock( + name=manager._cache_key("crashed_key"), ttl=1, backend=memory_backend + ) + assert await crashed_lock.acquire(blocking=False) is True + + # 10 concurrent misses: with wait_timeout=None, they do not time out at t=1. + # When crashed_lock expires at 1s, one waiter acquires the lock, executes factory once, + # and populates the cache. The other 9 waiters receive the cached value. + results = await asyncio.wait_for( + asyncio.gather( + *(manager.get_or_set("crashed_key", factory) for _ in range(10)) + ), + timeout=10, + ) + assert list(results) == ["recovered"] * 10 + assert calls == 1 + assert await manager.get("crashed_key") == "recovered" + + +async def test_stampede_protection_reentrancy_same_key( + memory_backend: MemoryBackend, +) -> None: + """A factory recursively calling get_or_set on the same key skips locking and avoids deadlock.""" + manager = CacheManager(backend=memory_backend, lock=True) + calls = 0 + + async def outer_factory() -> dict[str, Any]: + nonlocal calls + calls += 1 + inner = await manager.get_or_set("recursive", lambda: {"inner": 1}) + return {"outer": inner} + + result = await manager.get_or_set("recursive", outer_factory) + assert result == {"outer": {"inner": 1}} + assert calls == 1 + + +async def test_stampede_protection_reentrancy_different_keys( + memory_backend: MemoryBackend, +) -> None: + """A factory calling get_or_set on a different key acquires a lock for that key normally.""" + manager = CacheManager(backend=memory_backend, lock=True) + + async def outer_factory() -> dict[str, str]: + inner = await manager.get_or_set("child", lambda: "child_val") + return {"parent": inner} + + result = await manager.get_or_set("parent", outer_factory) + assert result == {"parent": "child_val"} + assert await manager.get("child") == "child_val" + assert await manager.get("parent") == {"parent": "child_val"} + + +async def test_stampede_protection_ordering_store_before_release( + memory_backend: MemoryBackend, +) -> None: + """The winner stores the value in backend before releasing the lock.""" + import unittest.mock + + manager = CacheManager(backend=memory_backend) + release_observed_cache: list[bool] = [] + + original_release = CacheLock.release + + async def tracking_release(self_lock: CacheLock) -> bool: + has_val = await memory_backend.get(manager._cache_key("order_key")) is not None + release_observed_cache.append(has_val) + return await original_release(self_lock) + + with unittest.mock.patch.object(CacheLock, "release", tracking_release): + await manager.get_or_set("order_key", lambda: "stored_val", lock=True) + + assert release_observed_cache == [True] + + +async def test_stampede_protection_non_json_serializable_releases_lock( + memory_backend: MemoryBackend, +) -> None: + """Non-JSON-serializable value raises TypeError and releases lock without corrupting cache.""" + manager = CacheManager(backend=memory_backend) + + with pytest.raises(TypeError): + await manager.get_or_set("bad_val", object, lock=True) + + lock = CacheLock(name=manager._cache_key("bad_val"), backend=memory_backend) + assert await lock.locked() is False + assert await manager.get("bad_val") is None + + assert ( + await manager.get_or_set("bad_val", lambda: "recovered", lock=True) + == "recovered" + ) + + +async def test_stampede_protection_validation(memory_backend: MemoryBackend) -> None: + """Validation rejects invalid lock_ttl, wait_timeout, and lock settings.""" + with pytest.raises(TypeError, match="lock must be a bool"): + CacheManager(backend=memory_backend, lock="no") # type: ignore[arg-type] + + with pytest.raises(TypeError, match="lock must be a bool"): + CacheManager(backend=memory_backend, lock=1) # type: ignore[arg-type] + + with pytest.raises(TypeError, match="lock must be a bool"): + CacheManager(backend=memory_backend, lock=None) # type: ignore[arg-type] + + with pytest.raises(ValueError, match="ttl must be a positive number"): + CacheManager(backend=memory_backend, lock_ttl=0) + + with pytest.raises(ValueError, match="ttl must be a positive number"): + CacheManager(backend=memory_backend, lock_ttl=-10) + + with pytest.raises(ValueError, match="lock_ttl must be a positive int, got None"): + CacheManager(backend=memory_backend, lock_ttl=None) # type: ignore[arg-type] + + with pytest.raises(TypeError, match="ttl must be an int"): + CacheManager(backend=memory_backend, lock_ttl=1.5) # type: ignore[arg-type] + + manager = CacheManager(backend=memory_backend) + + with pytest.raises(TypeError, match="lock must be a bool or None"): + await manager.get_or_set("key", lambda: 1, lock="yes") # type: ignore[arg-type] + + with pytest.raises(TypeError, match="lock must be a bool or None"): + await manager.get_or_set("key", lambda: 1, lock=1) # type: ignore[arg-type] + + with pytest.raises(ValueError, match="ttl must be a positive number"): + await manager.get_or_set("key", lambda: 1, lock_ttl=0) + + with pytest.raises(TypeError, match="ttl must be an int"): + await manager.get_or_set("key", lambda: 1, lock_ttl=1.5) # type: ignore[arg-type] + + with pytest.raises(ValueError, match="wait_timeout must be greater than zero"): + await manager.get_or_set("key", lambda: 1, wait_timeout=0) + + with pytest.raises(ValueError, match="wait_timeout must be greater than zero"): + await manager.get_or_set("key", lambda: 1, wait_timeout=-1.0) + + with pytest.raises(TypeError, match="wait_timeout"): + await manager.get_or_set("key", lambda: 1, wait_timeout=True) + + with pytest.raises(TypeError, match="wait_timeout"): + await manager.get_or_set("key", lambda: 1, wait_timeout="5") # type: ignore[arg-type] + + +async def test_stampede_protection_module_clock_and_sleep_hooks( + memory_backend: MemoryBackend, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Module-level _sleep and _monotonic hooks advance simulated time without real-world delay.""" + virtual_time = 0.0 + sleep_calls: list[float] = [] + + def fake_monotonic() -> float: + return virtual_time + + async def fake_sleep(duration: float) -> None: + nonlocal virtual_time + sleep_calls.append(duration) + virtual_time += duration + await asyncio.sleep(0) + + from fastapi_cachex import manager as manager_module + + monkeypatch.setattr(manager_module, "_sleep", fake_sleep) + monkeypatch.setattr(manager_module, "_monotonic", fake_monotonic) + + manager = CacheManager(backend=memory_backend) + + lock = CacheLock( + name=manager._cache_key("tick_key"), ttl=10, backend=memory_backend + ) + assert await lock.acquire(blocking=False) is True + + start_real = time.monotonic() + result = await manager.get_or_set( + "tick_key", + lambda: "after_timeout", + lock=True, + wait_timeout=1.0, + raise_on_timeout=False, + ) + elapsed_real = time.monotonic() - start_real + + assert elapsed_real < 0.2 + assert result == "after_timeout" + assert len(sleep_calls) >= 1 + assert virtual_time >= 1.0 + await lock.release() + + +async def test_stampede_protection_release_failure_logged_not_propagated( + memory_backend: MemoryBackend, + caplog: pytest.LogCaptureFixture, +) -> None: + """A failed lock release logs a warning and does not overwrite the computed result.""" + import unittest.mock + + manager = CacheManager(backend=memory_backend) + + async def failing_release(_self: CacheLock) -> bool: + msg = "backend release failure" + raise RuntimeError(msg) + + with ( + unittest.mock.patch.object(CacheLock, "release", failing_release), + caplog.at_level(logging.WARNING, logger="fastapi_cachex.manager"), + ): + result = await manager.get_or_set("release_fail", lambda: "success", lock=True) + + assert result == "success" + assert any( + "Failed to release stampede protection lock" in r.getMessage() + for r in caplog.records + ) + + +async def test_stampede_protection_release_failure_preserves_factory_exception( + memory_backend: MemoryBackend, + caplog: pytest.LogCaptureFixture, +) -> None: + """If factory raises and lock release also fails, the factory exception is preserved.""" + import unittest.mock + + manager = CacheManager(backend=memory_backend) + + async def failing_release(_self: CacheLock) -> bool: + msg = "backend release failure" + raise RuntimeError(msg) + + def failing_factory() -> None: + msg = "factory exception" + raise ValueError(msg) + + with ( + unittest.mock.patch.object(CacheLock, "release", failing_release), + caplog.at_level(logging.WARNING, logger="fastapi_cachex.manager"), + pytest.raises(ValueError, match="factory exception"), + ): + await manager.get_or_set("release_fail_factory", failing_factory, lock=True) + + assert any( + "Failed to release stampede protection lock" in r.getMessage() + for r in caplog.records + ) + + +@asynccontextmanager +async def _stampede_backend_context( + backend_name: str, +) -> AsyncGenerator[BaseCacheBackend, None]: + backend: BaseCacheBackend + if backend_name == "memory": + backend = MemoryBackend() + backend.start_cleanup() + elif backend_name == "redis": + from fastapi_cachex.backends import AsyncRedisCacheBackend + + backend = AsyncRedisCacheBackend( + host=REDIS_HOST, + port=REDIS_PORT, + socket_timeout=1.0, + socket_connect_timeout=1.0, + key_prefix="test_stampede:", + ) + else: + from fastapi_cachex.backends import MemcachedBackend + + backend = MemcachedBackend( + servers=[MEMCACHED_SERVER], + key_prefix="test_stampede:", + ) + + BackendProxy.set(backend) + try: + yield backend + finally: + if backend_name == "memcached": + from fastapi_cachex.backends import MemcachedBackend + + if isinstance(backend, MemcachedBackend): + await flush_memcached(backend) + else: + await backend.clear() + if backend_name == "memory" and isinstance(backend, MemoryBackend): + backend.stop_cleanup() + await backend.aclose() + + +_STAMPEDE_BACKEND_PARAMS = [ + pytest.param("memory", id="MemoryBackend"), + pytest.param( + "redis", + id="RedisBackend", + marks=[requires_redis, requires_redis_package], + ), + pytest.param( + "memcached", + id="MemcachedBackend", + marks=[requires_memcached], + ), +] + + +@pytest.mark.parametrize("backend_name", _STAMPEDE_BACKEND_PARAMS) +async def test_stampede_protection_across_backends(backend_name: str) -> None: + """Stampede protection works across Memory, Redis, and Memcached backends.""" + async with _stampede_backend_context(backend_name) as backend: + manager = CacheManager(backend=backend) + calls = 0 + + async def expensive_factory() -> dict[str, str]: + nonlocal calls + calls += 1 + await asyncio.sleep(0.05) + return {"status": "ok"} + + results = await asyncio.gather( + manager.get_or_set("report", expensive_factory, lock=True), + manager.get_or_set("report", expensive_factory, lock=True), + manager.get_or_set("report", expensive_factory, lock=True), + ) + + assert list(results) == [{"status": "ok"}, {"status": "ok"}, {"status": "ok"}] + assert calls == 1 + assert await manager.get("report") == {"status": "ok"} + + +@pytest.mark.parametrize("backend_name", _STAMPEDE_BACKEND_PARAMS) +async def test_stampede_protection_crashed_winner_across_backends( + backend_name: str, +) -> None: + """When a winner crashes, waiters wait for lock expiry and take over across backends.""" + async with _stampede_backend_context(backend_name) as backend: + manager = CacheManager(backend=backend, lock=True, lock_ttl=1) + calls = 0 + + async def factory() -> str: + nonlocal calls + calls += 1 + return "recovered" + + # Simulate crashed winner holding the lock for 1 second + crashed_lock = CacheLock( + name=manager._cache_key("crashed_key"), ttl=1, backend=backend + ) + assert await crashed_lock.acquire(blocking=False) is True + + results = await asyncio.wait_for( + asyncio.gather( + *(manager.get_or_set("crashed_key", factory) for _ in range(5)) + ), + timeout=10, + ) + assert list(results) == ["recovered"] * 5 + assert calls == 1 + assert await manager.get("crashed_key") == "recovered" + + +@pytest.mark.parametrize("backend_name", _STAMPEDE_BACKEND_PARAMS) +async def test_stampede_protection_raising_factory_across_backends( + backend_name: str, +) -> None: + """If factory raises, lock is released cleanly and subsequent callers can compute across backends.""" + async with _stampede_backend_context(backend_name) as backend: + manager = CacheManager(backend=backend, lock=True) + + def failing_factory() -> None: + msg = "database error" + raise ValueError(msg) + + with pytest.raises(ValueError, match="database error"): + await manager.get_or_set("fail_key", failing_factory) + + # Lock is released and key is not cached + lock = CacheLock(name=manager._cache_key("fail_key"), backend=backend) + assert await lock.locked() is False + assert await manager.get("fail_key") is None + + # Subsequent caller can compute and store normally + result = await manager.get_or_set("fail_key", lambda: "recovered") + assert result == "recovered" + assert await manager.get("fail_key") == "recovered" + + +@pytest.mark.parametrize("backend_name", _STAMPEDE_BACKEND_PARAMS) +async def test_stampede_protection_wait_timeout_across_backends( + backend_name: str, +) -> None: + """Wait timeout raises LockTimeoutError or falls back to factory across backends.""" + async with _stampede_backend_context(backend_name) as backend: + manager = CacheManager(backend=backend, lock=True) + + # Hold lock with longer TTL + lock = CacheLock( + name=manager._cache_key("timeout_key"), ttl=10, backend=backend + ) + assert await lock.acquire(blocking=False) is True + + try: + # 1. raise_on_timeout=True raises LockTimeoutError + with pytest.raises(LockTimeoutError) as exc_info: + await manager.get_or_set( + "timeout_key", + lambda: "value", + wait_timeout=0.05, + raise_on_timeout=True, + ) + assert isinstance(exc_info.value, TimeoutError) + + # 2. raise_on_timeout=False falls back to factory + result = await manager.get_or_set( + "timeout_key", + lambda: "fallback_value", + wait_timeout=0.05, + raise_on_timeout=False, + ) + assert result == "fallback_value" + finally: + await lock.release() + + +async def test_stampede_protection_waiter_takes_over_lock( + memory_backend: MemoryBackend, +) -> None: + """A waiter acquires the lock on a later tick when the previous lock is released.""" + manager = CacheManager(backend=memory_backend) + calls = 0 + + lock = CacheLock( + name=manager._cache_key("takeover"), ttl=10, backend=memory_backend + ) + assert await lock.acquire(blocking=False) is True + + async def release_soon() -> None: + await asyncio.sleep(0.06) + await lock.release() + + async def factory() -> str: + nonlocal calls + calls += 1 + return "taken_over" + + task = asyncio.create_task(release_soon()) + result = await manager.get_or_set("takeover", factory, lock=True, wait_timeout=1.0) + await task + assert result == "taken_over" + assert calls == 1 + + +async def test_stampede_protection_value_found_on_timeout_check( + memory_backend: MemoryBackend, +) -> None: + """If the value appears just as wait_timeout elapses, it returns the cached value.""" + manager = CacheManager(backend=memory_backend) + + lock = CacheLock( + name=manager._cache_key("late_val"), ttl=10, backend=memory_backend + ) + assert await lock.acquire(blocking=False) is True + + async def populate_late() -> None: + await asyncio.sleep(0.06) + await manager.set("late_val", "late_result") + await lock.release() + + task = asyncio.create_task(populate_late()) + result = await manager.get_or_set( + "late_val", + lambda: "should_not_run", + lock=True, + wait_timeout=0.08, + raise_on_timeout=True, + ) + await task + assert result == "late_result" + + +async def test_stampede_protection_winner_rechecks_inside_execution( + memory_backend: MemoryBackend, +) -> None: + """If value appears after lock acquisition check, _execute_as_winner returns cached value.""" + import unittest.mock + + manager = CacheManager(backend=memory_backend) + first_get = True + + async def sneaky_get(_key: str, default: Any = None) -> Any: + nonlocal first_get + if first_get: + first_get = False + return default + return "sneaky_cached" + + with unittest.mock.patch.object(manager, "get", side_effect=sneaky_get): + result = await manager.get_or_set("sneaky", lambda: "from_factory", lock=True) + assert result == "sneaky_cached"