From 64fcee285f39ccd97e896a822c70a6790f777e55 Mon Sep 17 00:00:00 2001 From: Shivansh Shukla Date: Fri, 25 Sep 2026 21:14:43 +0530 Subject: [PATCH 1/7] feat: add expire_if_equals backend primitive and CacheLock async context manager (#62) - Implement expire_if_equals(key, expected, ttl) atomic primitive on BaseCacheBackend, MemoryBackend, AsyncRedisCacheBackend (Lua script), and MemcachedBackend (CAS). - Document expire_if_equals in docs/BACKENDS.md under 'Atomic backend primitives'. - Implement CacheLock distributed lock helper and async context manager in fastapi_cachex/lock.py. - Add LockTimeoutError exception raised on context manager timeout. - Export CacheLock and LockTimeoutError from package root. - Add comprehensive unit test suite in tests/test_lock.py and tests/backends/. - Add CHANGELOG entry under [Unreleased]. --- CHANGELOG.md | 5 + docs/BACKENDS.md | 10 +- fastapi_cachex/__init__.py | 4 + fastapi_cachex/backends/base.py | 26 ++++ fastapi_cachex/backends/memcached.py | 23 ++++ fastapi_cachex/backends/memory.py | 15 +++ fastapi_cachex/backends/redis.py | 36 ++++++ fastapi_cachex/exceptions.py | 4 + fastapi_cachex/lock.py | 165 ++++++++++++++++++++++++ fastapi_cachex/types.py | 9 ++ pyproject.toml | 1 + tests/backends/test_base.py | 15 +++ tests/backends/test_memcached.py | 41 ++++++ tests/backends/test_memory.py | 25 ++++ tests/backends/test_redis.py | 42 ++++++ tests/test_lock.py | 183 +++++++++++++++++++++++++++ 16 files changed, 602 insertions(+), 2 deletions(-) create mode 100644 fastapi_cachex/lock.py create mode 100644 tests/test_lock.py diff --git a/CHANGELOG.md b/CHANGELOG.md index a960843..899bd6c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -15,6 +15,11 @@ Note that 0.3.3 was never released; 0.3.4 follows 0.3.2. ## [Unreleased] +### Added + +- `CacheLock` distributed lock async context manager helper built on backend primitives, with `acquire`, `release`, `extend`, `locked`, and `LockTimeoutError`. +- `expire_if_equals(key, expected, ttl) -> bool` atomic backend primitive added to `BaseCacheBackend`, `MemoryBackend`, `AsyncRedisCacheBackend`, and `MemcachedBackend`. + ## [0.3.7] - 2026-09-25 ### Added diff --git a/docs/BACKENDS.md b/docs/BACKENDS.md index 047519f..e7a12ed 100644 --- a/docs/BACKENDS.md +++ b/docs/BACKENDS.md @@ -161,8 +161,14 @@ if await backend.set_if_absent(f"stream:{user_id}", owner, ttl=300): through a Lua script that re-checks the value it compared, and Memcached uses `GETS` + a `CAS` write that expires the entry immediately (the classic protocol's `DELETE` takes no CAS token). - -All four have a non-atomic fallback on `BaseCacheBackend`, so a third-party backend +- `expire_if_equals(key, expected, ttl) -> bool` — updates the TTL on `key` to `ttl` + seconds only while it still holds `expected`, so a long-running lock holder can + renew its lease without risking overwriting someone else's lock if it expired. + Memory updates under its lock, Redis compares in Python then executes a Lua + script (`GET` compare + `EXPIRE`), and Memcached uses `GETS` + `CAS` writing the + same bytes with the new exptime (`TOUCH` takes no CAS token). + +All five have a non-atomic fallback on `BaseCacheBackend`, so a third-party backend that only implements the abstract methods keeps working; override them to get real atomicity. diff --git a/fastapi_cachex/__init__.py b/fastapi_cachex/__init__.py index 19e0b89..78f8b20 100644 --- a/fastapi_cachex/__init__.py +++ b/fastapi_cachex/__init__.py @@ -11,6 +11,8 @@ from .dependencies import CacheBackend as CacheBackend from .dependencies import get_app_cache as get_app_cache from .dependencies import get_cache_backend as get_cache_backend +from .exceptions import LockTimeoutError as LockTimeoutError +from .lock import CacheLock as CacheLock from .manager import CacheManager as CacheManager from .manager_proxy import CacheManagerProxy as CacheManagerProxy from .proxy import BackendProxy as BackendProxy @@ -69,10 +71,12 @@ def _read_version() -> str: "BackendProxy", "CacheBackend", "CacheKeyBuilder", + "CacheLock", "CacheManager", "CacheManagerProxy", "FastAPICacheXSessionMiddleware", "InvalidStateError", + "LockTimeoutError", "Session", "SessionConfig", "SessionError", diff --git a/fastapi_cachex/backends/base.py b/fastapi_cachex/backends/base.py index d9e7d27..9a7f392 100644 --- a/fastapi_cachex/backends/base.py +++ b/fastapi_cachex/backends/base.py @@ -159,6 +159,32 @@ async def delete_if_equals(self, key: str, expected: CacheEntry) -> bool: await self.delete(key) return True + async def expire_if_equals(self, key: str, expected: CacheEntry, ttl: int) -> bool: + """Update expiry on ``key`` to ``ttl`` seconds only while it still holds ``expected``. + + Re-setting a lock's TTL with a plain ``set`` is unsafe: if the holder's + entry expired and someone else claimed the key in the meantime, a plain + ``set`` overwrites the new holder's entry. Comparing against the value + the caller stored makes the renewal a no-op in that case. + + The base implementation is a best-effort, NON-atomic get-compare-set + fallback for third-party backends; the built-in backends override it + with an atomic implementation. + + Args: + key: Cache key to update expiry for + expected: The entry the caller stored (compared with ``==``) + ttl: Time to live in seconds + + Returns: + Whether the expiry was updated + """ + validate_ttl(ttl) + if await self.get(key) != expected: + return False + await self.set(key, expected, ttl=ttl) + return True + async def increment(self, key: str, delta: int = 1, ttl: int | None = None) -> int: """Atomically add ``delta`` to the integer counter stored at ``key``. diff --git a/fastapi_cachex/backends/memcached.py b/fastapi_cachex/backends/memcached.py index f5c0fef..1355e01 100644 --- a/fastapi_cachex/backends/memcached.py +++ b/fastapi_cachex/backends/memcached.py @@ -215,6 +215,29 @@ async def delete_if_equals(self, key: str, expected: CacheEntry) -> bool: ) return bool(deleted) + async def expire_if_equals(self, key: str, expected: CacheEntry, ttl: int) -> bool: + """Atomically update expiry on ``key`` while it holds ``expected`` (see base class). + + TOUCH in the Memcached protocol takes no CAS token, so the renewal is a + CAS write with the same bytes read by GETS and the new exptime. + """ + validate_ttl(ttl) + prefixed_key = self._make_key(key) + raw, cas_token = await asyncio.to_thread(self.client.gets, prefixed_key) + if raw is None or decode_entry(raw) != expected: + logger.debug("Memcached EXPIRE_IF_EQUALS MISMATCH; key=%s", key) + return False + updated = await asyncio.to_thread( + self.client.cas, prefixed_key, raw, cas_token, _expiry(ttl), noreply=False + ) + logger.debug( + "Memcached EXPIRE_IF_EQUALS %s; key=%s ttl=%s", + "HIT" if updated else "LOST RACE", + key, + ttl, + ) + return bool(updated) + def _add_delta(self, prefixed_key: str, delta: int) -> int | None: """Apply ``delta`` with INCR/DECR; ``None`` when the key does not exist.""" if delta < 0: diff --git a/fastapi_cachex/backends/memory.py b/fastapi_cachex/backends/memory.py index 4aa4e53..2b31c1c 100644 --- a/fastapi_cachex/backends/memory.py +++ b/fastapi_cachex/backends/memory.py @@ -191,6 +191,21 @@ async def delete_if_equals(self, key: str, expected: CacheEntry) -> bool: logger.debug("Memory cache DELETE_IF_EQUALS HIT; key=%s", key) return True + async def expire_if_equals(self, key: str, expected: CacheEntry, ttl: int) -> bool: + """Atomically update expiry on ``key`` while it holds ``expected`` (see base class).""" + validate_ttl(ttl) + async with self.lock: + item = self.cache.get(key) + if item is None or not _is_live(item, time.time()): + logger.debug("Memory cache EXPIRE_IF_EQUALS MISS; key=%s", key) + return False + if item.value != expected: + logger.debug("Memory cache EXPIRE_IF_EQUALS MISMATCH; key=%s", key) + return False + item.expiry = time.time() + ttl + logger.debug("Memory cache EXPIRE_IF_EQUALS HIT; key=%s ttl=%s", key, ttl) + return True + async def increment(self, key: str, delta: int = 1, ttl: int | None = None) -> int: """Atomically add ``delta`` to the counter at ``key`` (see base class). diff --git a/fastapi_cachex/backends/redis.py b/fastapi_cachex/backends/redis.py index 7596e95..5f79c16 100644 --- a/fastapi_cachex/backends/redis.py +++ b/fastapi_cachex/backends/redis.py @@ -60,6 +60,15 @@ def _escape_glob(text: str) -> str: return 0 """ +# EXPIRE that only fires while the key still holds the exact bytes the caller read, +# so a value replaced in the meantime survives. KEYS[1] = key, ARGV[1] = bytes, ARGV[2] = ttl. +_EXPIRE_IF_EQUALS_SCRIPT = """ +if redis.call('GET', KEYS[1]) == ARGV[1] then + return redis.call('EXPIRE', KEYS[1], ARGV[2]) +end +return 0 +""" + class AsyncRedisCacheBackend(BaseCacheBackend): """Async Redis cache backend implementation. @@ -135,6 +144,9 @@ def __init__( self._delete_if_equals_script = self.client.register_script( _DELETE_IF_EQUALS_SCRIPT ) + self._expire_if_equals_script = self.client.register_script( + _EXPIRE_IF_EQUALS_SCRIPT + ) @staticmethod def load_from_config(config: RedisConfig) -> "AsyncRedisCacheBackend": @@ -259,6 +271,30 @@ async def delete_if_equals(self, key: str, expected: CacheEntry) -> bool: ) return bool(deleted) + async def expire_if_equals(self, key: str, expected: CacheEntry, ttl: int) -> bool: + """Atomically update expiry on ``key`` while it holds ``expected`` (see base class). + + The stored value is decoded and compared here, then a Lua script + updates the expiry on the key only if it still holds the bytes that + were compared, so a value written in between is never overwritten. + """ + validate_ttl(ttl) + prefixed_key = self._make_key(key) + raw = await self.client.get(prefixed_key) + if raw is None or decode_entry(raw) != expected: + logger.debug("Redis EXPIRE_IF_EQUALS MISMATCH; key=%s", key) + return False + updated = await self._expire_if_equals_script( + keys=[prefixed_key], args=[raw, ttl] + ) + logger.debug( + "Redis EXPIRE_IF_EQUALS %s; key=%s ttl=%s", + "HIT" if updated else "LOST RACE", + key, + ttl, + ) + return bool(updated) + async def increment(self, key: str, delta: int = 1, ttl: int | None = None) -> int: """Atomically add ``delta`` to the counter at ``key`` (see base class). diff --git a/fastapi_cachex/exceptions.py b/fastapi_cachex/exceptions.py index dd2e667..552a68d 100644 --- a/fastapi_cachex/exceptions.py +++ b/fastapi_cachex/exceptions.py @@ -15,3 +15,7 @@ class BackendNotFoundError(CacheXError): class RequestNotFoundError(CacheXError): """Exception raised when a request is not found.""" + + +class LockTimeoutError(CacheXError): + """Exception raised when acquiring a lock times out.""" diff --git a/fastapi_cachex/lock.py b/fastapi_cachex/lock.py new file mode 100644 index 0000000..af7bc16 --- /dev/null +++ b/fastapi_cachex/lock.py @@ -0,0 +1,165 @@ +"""Distributed lock helper for FastAPI-CacheX.""" + +import asyncio +import logging +import secrets +import time +from types import TracebackType + +from typing_extensions import Self + +from fastapi_cachex.backends.base import BaseCacheBackend +from fastapi_cachex.exceptions import LockTimeoutError +from fastapi_cachex.proxy import BackendProxy +from fastapi_cachex.types import CacheEntry +from fastapi_cachex.types import lock_entry + +logger = logging.getLogger(__name__) + + +class CacheLock: + """Distributed lock helper built on backend primitives. + + Can be used as an async context manager or with direct acquire/release calls. + + Args: + name: Unique lock identifier + ttl: Time to live in seconds (default: 60) + blocking: Whether acquire() waits for the lock if unavailable (default: True) + timeout: Maximum seconds to wait in blocking mode (None = wait indefinitely) + poll_interval: Seconds between retry attempts in blocking mode (default: 0.1) + backend: Explicit backend instance to use (default: BackendProxy.get()) + key_prefix: Prefix for lock keys (default: "lock:") + """ + + def __init__( + self, + name: str, + ttl: int = 60, + blocking: bool = True, + timeout: float | None = None, + poll_interval: float = 0.1, + backend: BaseCacheBackend | None = None, + key_prefix: str = "lock:", + ) -> None: + """Initialize a CacheLock instance.""" + self.name = name + self.ttl = ttl + self.blocking = blocking + self.timeout = timeout + self.poll_interval = poll_interval + self._backend = backend + self.key_prefix = key_prefix + self._token = secrets.token_hex(16) + self._entry: CacheEntry = lock_entry(self._token) + + @property + def key(self) -> str: + """The fully-prefixed backend key for this lock.""" + return f"{self.key_prefix}{self.name}" + + def _get_backend(self) -> BaseCacheBackend: + """Return configured backend or resolve via BackendProxy.get().""" + if self._backend is not None: + return self._backend + return BackendProxy.get() + + async def acquire( + self, + blocking: bool | None = None, + timeout: float | None = None, + poll_interval: float | None = None, + ttl: int | None = None, + ) -> bool: + """Attempt to acquire the lock. + + Args: + blocking: Override default blocking mode + timeout: Override default timeout (seconds) + poll_interval: Override default poll interval (seconds) + ttl: Override default TTL (seconds) + + Returns: + True if the lock was acquired, False on timeout or failure + """ + is_blocking = self.blocking if blocking is None else blocking + effective_timeout = self.timeout if timeout is None else timeout + interval = self.poll_interval if poll_interval is None else poll_interval + effective_ttl = self.ttl if ttl is None else ttl + + backend = self._get_backend() + + if not is_blocking: + return await backend.set_if_absent(self.key, self._entry, ttl=effective_ttl) + + start = time.monotonic() + while True: + if await backend.set_if_absent(self.key, self._entry, ttl=effective_ttl): + return True + if effective_timeout is not None: + elapsed = time.monotonic() - start + if elapsed >= effective_timeout: + return False + remaining = effective_timeout - elapsed + sleep_time = min(interval, remaining) + else: + sleep_time = interval + await asyncio.sleep(sleep_time) + + async def release(self) -> bool: + """Release the lock if still held by this instance. + + Returns: + True if the lock was released, False if expired or owned by another caller + """ + backend = self._get_backend() + released = await backend.delete_if_equals(self.key, self._entry) + if not released: + logger.debug( + "CacheLock release failed for key=%s (lock lost or expired)", + self.key, + ) + return released + + async def extend(self, ttl: int | None = None) -> bool: + """Renew the TTL on the lock if still held by this instance. + + Args: + ttl: New TTL in seconds (None = use default instance TTL) + + Returns: + True if the TTL was updated, False if expired or owned by another caller + """ + effective_ttl = self.ttl if ttl is None else ttl + backend = self._get_backend() + return await backend.expire_if_equals(self.key, self._entry, ttl=effective_ttl) + + async def locked(self) -> bool: + """Check whether the lock is currently held by any holder. + + Returns: + True if the lock key currently exists in the backend, False otherwise + """ + backend = self._get_backend() + return (await backend.get(self.key)) is not None + + async def __aenter__(self) -> Self: + """Acquire lock as async context manager. + + Raises: + LockTimeoutError: If the lock cannot be acquired within the timeout + """ + acquired = await self.acquire() + if not acquired: + msg = f"Failed to acquire lock '{self.name}' within timeout" + raise LockTimeoutError(msg) + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_val: BaseException | None, + exc_tb: TracebackType | None, + ) -> None: + """Release lock when exiting async context manager.""" + await self.release() diff --git a/fastapi_cachex/types.py b/fastapi_cachex/types.py index 5fb6d74..9b507c5 100644 --- a/fastapi_cachex/types.py +++ b/fastapi_cachex/types.py @@ -71,3 +71,12 @@ def counter_value(entry: CacheEntry) -> int: except ValueError as e: msg = "Cache key holds a value that is not a counter" raise CacheXError(msg) from e + + +# Fingerprint every backend reports for a key that holds a lock entry +LOCK_FINGERPRINT = "lock" + + +def lock_entry(token: str) -> CacheEntry: + """Wrap a lock token string in the entry model shared by every backend.""" + return CacheEntry(fingerprint=LOCK_FINGERPRINT, content=token.encode("utf-8")) diff --git a/pyproject.toml b/pyproject.toml index 2f9b256..6ff777a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -127,6 +127,7 @@ keep-runtime-typing = true ] "fastapi_cachex/backends/memcached.py" = ["PLC0415"] # Optional dependency "fastapi_cachex/backends/redis.py" = ["PLR0913", "PLC0415"] # Optional dependency, Redis config +"fastapi_cachex/lock.py" = ["PLR0913"] # Configurable lock parameters "fastapi_cachex/proxy.py" = ["PLW0603"] # Global backend management by design "fastapi_cachex/session/manager.py" = ["PLR0915"] # Session/security validation branches needed diff --git a/tests/backends/test_base.py b/tests/backends/test_base.py index d40450c..45c0585 100644 --- a/tests/backends/test_base.py +++ b/tests/backends/test_base.py @@ -113,3 +113,18 @@ async def test_delete_if_equals_fallback_removes_only_a_matching_entry( assert await backend.delete_if_equals("slot", theirs) is True assert "slot" not in backend.store assert await backend.delete_if_equals("slot", theirs) is False + + +@pytest.mark.asyncio +async def test_expire_if_equals_fallback_updates_ttl_only_when_matching( + backend: DictBackend, +) -> None: + mine = CacheEntry(fingerprint="lock", content=b"owner-a") + theirs = CacheEntry(fingerprint="lock", content=b"owner-b") + await backend.set("slot", theirs, ttl=30) + + assert await backend.expire_if_equals("slot", mine, ttl=60) is False + assert backend.store["slot"] == (theirs, 30) + assert await backend.expire_if_equals("slot", theirs, ttl=60) is True + assert backend.store["slot"] == (theirs, 60) + assert await backend.expire_if_equals("missing", theirs, ttl=60) is False diff --git a/tests/backends/test_memcached.py b/tests/backends/test_memcached.py index a8483cc..e40db88 100644 --- a/tests/backends/test_memcached.py +++ b/tests/backends/test_memcached.py @@ -667,3 +667,44 @@ def gets_then_overwrite(key): raw = await asyncio.to_thread(client.get, memcached_backend._make_key("slot")) assert raw == b'{"overwritten": true}' + + +@requires_memcached +@pytest.mark.asyncio +async def test_memcached_expire_if_equals_updates_ttl_only_when_matching( + memcached_backend: MemcachedBackend, +): + mine = CacheEntry(fingerprint="lock", content=b"owner-a") + theirs = CacheEntry(fingerprint="lock", content=b"owner-b") + await memcached_backend.set("slot", theirs, 30) + + assert await memcached_backend.expire_if_equals("slot", mine, 60) is False + assert await memcached_backend.expire_if_equals("slot", theirs, 60) is True + assert await memcached_backend.get("slot") == theirs + assert await memcached_backend.expire_if_equals("missing", theirs, 60) is False + + +@requires_memcached +@pytest.mark.asyncio +async def test_memcached_expire_if_equals_keeps_a_value_written_after_the_compare( + memcached_backend: MemcachedBackend, +): + mine = CacheEntry(fingerprint="lock", content=b"owner-a") + await memcached_backend.set("slot", mine, 60) + + client = memcached_backend.client + original_gets = client.gets + + def gets_then_overwrite(key): + result = original_gets(key) + client.set(key, b'{"overwritten": true}', 60) + return result + + client.gets = gets_then_overwrite + try: + assert await memcached_backend.expire_if_equals("slot", mine, 120) is False + finally: + client.gets = original_gets + + raw = await asyncio.to_thread(client.get, memcached_backend._make_key("slot")) + assert raw == b'{"overwritten": true}' diff --git a/tests/backends/test_memory.py b/tests/backends/test_memory.py index c9b48b7..40009fa 100644 --- a/tests/backends/test_memory.py +++ b/tests/backends/test_memory.py @@ -799,3 +799,28 @@ async def test_memory_lock_release_after_expiry_keeps_the_new_holder( assert await memory_backend.delete_if_equals("slot", owner_a) is False assert await memory_backend.get("slot") == owner_b + + +@pytest.mark.asyncio +async def test_memory_expire_if_equals_updates_ttl_only_when_matching( + memory_backend: MemoryBackend, +): + mine = CacheEntry(fingerprint="lock", content=b"owner-a") + theirs = CacheEntry(fingerprint="lock", content=b"owner-b") + await memory_backend.set("slot", theirs, 30) + + assert await memory_backend.expire_if_equals("slot", mine, 60) is False + assert await memory_backend.expire_if_equals("slot", theirs, 60) is True + assert memory_backend.cache["slot"].expiry is not None + assert await memory_backend.expire_if_equals("missing", theirs, 60) is False + + +@pytest.mark.asyncio +async def test_memory_expire_if_equals_ignores_an_expired_entry( + memory_backend: MemoryBackend, +): + entry = CacheEntry(fingerprint="lock", content=b"owner-a") + await memory_backend.set("slot", entry, 60) + memory_backend.cache["slot"].expiry = time.time() - 1 + + assert await memory_backend.expire_if_equals("slot", entry, 60) is False diff --git a/tests/backends/test_redis.py b/tests/backends/test_redis.py index 8e457b5..93acacd 100644 --- a/tests/backends/test_redis.py +++ b/tests/backends/test_redis.py @@ -886,6 +886,48 @@ async def test_redis_lock_release_after_expiry_keeps_the_new_holder( assert await async_redis_backend.get("slot") == owner_b +@requires_redis +@pytest.mark.asyncio +async def test_redis_expire_if_equals_updates_ttl_only_when_matching( + async_redis_backend: AsyncRedisCacheBackend, +) -> None: + mine = CacheEntry(fingerprint="lock", content=b"owner-a") + theirs = CacheEntry(fingerprint="lock", content=b"owner-b") + await async_redis_backend.set("slot", theirs, 30) + + assert await async_redis_backend.expire_if_equals("slot", mine, 60) is False + assert await async_redis_backend.expire_if_equals("slot", theirs, 60) is True + pttl = await async_redis_backend.client.pttl(async_redis_backend._make_key("slot")) + assert 55000 <= pttl <= 60000 + assert await async_redis_backend.expire_if_equals("missing", theirs, 60) is False + + +@requires_redis +@pytest.mark.asyncio +async def test_redis_expire_if_equals_keeps_a_value_written_after_the_compare( + async_redis_backend: AsyncRedisCacheBackend, +) -> None: + mine = CacheEntry(fingerprint="lock", content=b"owner-a") + theirs = CacheEntry(fingerprint="lock", content=b"owner-b") + await async_redis_backend.set("slot", mine, 60) + + client = async_redis_backend.client + original_get = client.get + + async def get_then_overwrite(name): + raw = await original_get(name) + await async_redis_backend.set("slot", theirs, 60) + return raw + + client.get = get_then_overwrite # type: ignore[method-assign] + try: + assert await async_redis_backend.expire_if_equals("slot", mine, 120) is False + finally: + client.get = original_get # type: ignore[method-assign] + + assert await async_redis_backend.get("slot") == theirs + + @requires_redis def test_cached_route_with_ttl_zero_is_served() -> None: """End to end: `@cache(ttl=0)` used to send `SET ... EX 0` and answer 500.""" diff --git a/tests/test_lock.py b/tests/test_lock.py new file mode 100644 index 0000000..e58f2f9 --- /dev/null +++ b/tests/test_lock.py @@ -0,0 +1,183 @@ +"""Tests for the CacheLock distributed lock helper.""" + +import time + +import pytest + +from fastapi_cachex.backends.memory import MemoryBackend +from fastapi_cachex.exceptions import BackendNotFoundError +from fastapi_cachex.exceptions import LockTimeoutError +from fastapi_cachex.lock import CacheLock +from fastapi_cachex.proxy import BackendProxy + + +@pytest.mark.asyncio +async def test_lock_basic_context_manager() -> None: + backend = MemoryBackend() + BackendProxy.set(backend) + + async with CacheLock("test_job", ttl=30) as lock: + assert lock.name == "test_job" + assert lock.key == "lock:test_job" + assert await lock.locked() is True + assert await backend.get("lock:test_job") is not None + + assert await backend.get("lock:test_job") is None + + +@pytest.mark.asyncio +async def test_lock_acquire_non_blocking() -> None: + backend = MemoryBackend() + BackendProxy.set(backend) + + lock1 = CacheLock("job", ttl=30) + lock2 = CacheLock("job", ttl=30) + + assert await lock1.acquire(blocking=False) is True + assert await lock2.acquire(blocking=False) is False + assert await lock1.locked() is True + + assert await lock1.release() is True + assert await lock2.release() is False + assert await lock1.locked() is False + + +@pytest.mark.asyncio +async def test_lock_acquire_blocking_indefinite_retries_until_available() -> None: + import asyncio + + backend = MemoryBackend() + BackendProxy.set(backend) + + lock1 = CacheLock("job", ttl=30) + lock2 = CacheLock("job", ttl=30, poll_interval=0.01) + + assert await lock1.acquire(blocking=False) is True + + async def release_later() -> None: + await asyncio.sleep(0.03) + await lock1.release() + + task = asyncio.create_task(release_later()) + assert await lock2.acquire(blocking=True, poll_interval=0.01) is True + await task + assert await lock2.release() is True + + +@pytest.mark.asyncio +async def test_lock_acquire_blocking_timeout_returns_false() -> None: + backend = MemoryBackend() + BackendProxy.set(backend) + + lock1 = CacheLock("job", ttl=30) + lock2 = CacheLock("job", ttl=30, timeout=0.05, poll_interval=0.01) + + assert await lock1.acquire(blocking=False) is True + assert await lock2.acquire(blocking=True) is False + assert await lock1.release() is True + + +@pytest.mark.asyncio +async def test_lock_context_manager_timeout_raises_lock_timeout_error() -> None: + backend = MemoryBackend() + BackendProxy.set(backend) + + lock1 = CacheLock("job", ttl=30) + assert await lock1.acquire(blocking=False) is True + + with pytest.raises(LockTimeoutError, match="Failed to acquire lock 'job'"): + async with CacheLock("job", ttl=30, timeout=0.05, poll_interval=0.01): + pass + + assert await lock1.release() is True + + +@pytest.mark.asyncio +async def test_lock_extend() -> None: + backend = MemoryBackend() + BackendProxy.set(backend) + + lock1 = CacheLock("job", ttl=30) + lock2 = CacheLock("job", ttl=30) + + assert await lock1.acquire(blocking=False) is True + assert await lock1.extend(60) is True + assert await lock2.extend(60) is False + + assert await lock1.release() is True + assert await lock1.extend(60) is False + + +@pytest.mark.asyncio +async def test_lock_locked_status() -> None: + backend = MemoryBackend() + BackendProxy.set(backend) + + lock = CacheLock("job") + assert await lock.locked() is False + + await lock.acquire(blocking=False) + assert await lock.locked() is True + + await lock.release() + assert await lock.locked() is False + + +@pytest.mark.asyncio +async def test_lock_custom_key_prefix() -> None: + backend = MemoryBackend() + BackendProxy.set(backend) + + lock = CacheLock("custom", key_prefix="custom_lock:") + assert lock.key == "custom_lock:custom" + + async with lock: + assert await backend.get("custom_lock:custom") is not None + + assert await backend.get("custom_lock:custom") is None + + +@pytest.mark.asyncio +async def test_lock_custom_backend() -> None: + default_backend = MemoryBackend() + custom_backend = MemoryBackend() + BackendProxy.set(default_backend) + + lock = CacheLock("job", backend=custom_backend) + async with lock: + assert await custom_backend.get("lock:job") is not None + assert await default_backend.get("lock:job") is None + + assert await custom_backend.get("lock:job") is None + + +@pytest.mark.asyncio +async def test_lock_release_on_expired_lock_returns_false() -> None: + backend = MemoryBackend() + BackendProxy.set(backend) + + lock = CacheLock("job", ttl=30) + assert await lock.acquire(blocking=False) is True + + # Simulate TTL expiration + backend.cache["lock:job"].expiry = time.time() - 1 + + assert await lock.release() is False + + +@pytest.mark.asyncio +async def test_lock_token_uniqueness() -> None: + lock1 = CacheLock("job") + lock2 = CacheLock("job") + + assert lock1._token != lock2._token + assert lock1._entry != lock2._entry + + +@pytest.mark.asyncio +async def test_lock_raises_backend_not_found_if_no_backend_set() -> None: + BackendProxy.set(None) + lock = CacheLock("job") + + with pytest.raises(BackendNotFoundError): + await lock.acquire(blocking=False) From b6fb03a2584a3b010b007d0028a4c2c53fd491a4 Mon Sep 17 00:00:00 2001 From: Shivansh Shukla Date: Fri, 25 Sep 2026 22:25:05 +0530 Subject: [PATCH 2/7] refactor: address maintainer review feedback for CacheLock (#64) - Move _LOCK_FINGERPRINT and _lock_entry into lock.py as private helpers. - Add _is_held state guard to CacheLock to prevent re-entry or instance sharing across concurrent tasks. - Remove Self import from typing_extensions; use string return annotation for __aenter__. - Update acquire() docstring timeout parameter description. - Add RuntimeError re-entry and shared instance unit tests in tests/test_lock.py. - Add live Redis and Memcached integration tests for CacheLock. - Add docs/LOCK.md guide, docs/api/lock.md API reference, and update navigation in zensical.toml, zensical.zh-TW.toml, and README.md. - Update CHANGELOG.md entry with issue link #64. --- CHANGELOG.md | 2 +- README.md | 1 + docs/LOCK.md | 39 ++++++++++++++++++++++ docs/api/lock.md | 5 +++ fastapi_cachex/lock.py | 42 ++++++++++++++++++++---- fastapi_cachex/types.py | 9 ----- pyproject.toml | 2 +- tests/test_lock.py | 73 +++++++++++++++++++++++++++++++++++++++-- zensical.toml | 2 ++ zensical.zh-TW.toml | 1 + 10 files changed, 156 insertions(+), 20 deletions(-) create mode 100644 docs/LOCK.md create mode 100644 docs/api/lock.md diff --git a/CHANGELOG.md b/CHANGELOG.md index 899bd6c..d10bd97 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -17,7 +17,7 @@ Note that 0.3.3 was never released; 0.3.4 follows 0.3.2. ### Added -- `CacheLock` distributed lock async context manager helper built on backend primitives, with `acquire`, `release`, `extend`, `locked`, and `LockTimeoutError`. +- `CacheLock` distributed lock async context manager helper built on backend primitives, with `acquire`, `release`, `extend`, `locked`, and `LockTimeoutError`. ([#64](https://github.com/allen0099/FastAPI-CacheX/issues/64)) - `expire_if_equals(key, expected, ttl) -> bool` atomic backend primitive added to `BaseCacheBackend`, `MemoryBackend`, `AsyncRedisCacheBackend`, and `MemcachedBackend`. ## [0.3.7] - 2026-09-25 diff --git a/README.md b/README.md index 059b5f9..616e2b4 100644 --- a/README.md +++ b/README.md @@ -88,6 +88,7 @@ async def report(cache: AppCache): - [Backends](https://fastapi-cachex.readthedocs.io/en/latest/BACKENDS/) — choosing and configuring a backend, atomic primitives - [Session management](https://fastapi-cachex.readthedocs.io/en/latest/SESSION/) and [JWT claims](https://fastapi-cachex.readthedocs.io/en/latest/JWT_CLAIMS/) - [OAuth state](https://fastapi-cachex.readthedocs.io/en/latest/STATE/) — one-shot OAuth/CSRF state tokens +- [Distributed lock](https://fastapi-cachex.readthedocs.io/en/latest/LOCK/) — `CacheLock` for multi-process mutual exclusion - [API reference](https://fastapi-cachex.readthedocs.io/en/latest/api/http-caching/) - [Development guide](https://fastapi-cachex.readthedocs.io/en/latest/DEVELOPMENT/) and [contributing](https://fastapi-cachex.readthedocs.io/en/latest/CONTRIBUTING/) - [Changelog](https://github.com/allen0099/FastAPI-CacheX/blob/master/CHANGELOG.md) · [Known limitations and planned work](https://github.com/allen0099/FastAPI-CacheX/issues) diff --git a/docs/LOCK.md b/docs/LOCK.md new file mode 100644 index 0000000..1f60bc8 --- /dev/null +++ b/docs/LOCK.md @@ -0,0 +1,39 @@ +# Distributed Lock + +`CacheLock` provides a distributed lock helper built on backend atomic primitives (`set_if_absent`, `delete_if_equals`, `expire_if_equals`). It guarantees mutual exclusion across multiple processes or containers sharing the same cache backend. + +```python +from fastapi import HTTPException +from fastapi_cachex import CacheLock, LockTimeoutError + +# Usage as an async context manager: +async with CacheLock(f"report:{report_id}", ttl=30): + ... # only one process holds the lock at any time + +# Explicit acquire and release calls: +lock = CacheLock(f"stream:{user_id}", ttl=60) +if not await lock.acquire(blocking=False): + raise HTTPException(409, detail="Lock already held") +try: + ... + await lock.extend(60) # long-running tasks renew before expiry +finally: + await lock.release() +``` + +## Behavior + +- **Safety & Token Ownership**: Each `CacheLock` instance generates a unique token (`secrets.token_hex(16)`) stored inside a `CacheEntry`. Releases (`release()`) and extensions (`extend()`) use owner-checked backend primitives (`delete_if_equals` and `expire_if_equals`), so a holder whose lock expired cannot release or renew a lock claimed by someone else. +- **Blocking & Non-blocking Modes**: + - Non-blocking (`acquire(blocking=False)`): Performs a single atomic `set_if_absent` and immediately returns `True` if acquired or `False` if held. + - Blocking (`acquire(blocking=True, timeout=None, poll_interval=0.1)`): Retries at `poll_interval` seconds until acquired or until `timeout` seconds elapse. If `timeout` is reached, `acquire()` returns `False`. +- **Context Manager Timeouts**: Entering a context manager (`async with CacheLock(...)`) invokes `acquire()`. If acquisition fails or times out, it raises `LockTimeoutError`. +- **TTL Renewal (`extend`)**: `extend(ttl)` updates the key's TTL only while the lock is still owned by this holder instance, preventing race conditions on expired locks. +- **One Instance per Acquisition Rule**: A single `CacheLock` instance tracks its active ownership state. Re-entering or sharing a single `CacheLock` instance across concurrent tasks raises a `RuntimeError`. Instantiate a new `CacheLock` instance for each acquisition. +- **Namespace**: Lock keys live under their own `lock:` prefix by default (e.g. `lock:report:123`), separate from `cache:` and `oauth_state:`. + +> [!NOTE] +> `CacheLock` works on all built-in backends (`MemoryBackend`, `AsyncRedisCacheBackend`, `MemcachedBackend`). +> Redis executes Lua scripts for atomic operations, Memcached uses `ADD` and `CAS`, and the memory backend operates under its internal lock. + +The full class signature and options are documented in the [API reference](api/lock.md). diff --git a/docs/api/lock.md b/docs/api/lock.md new file mode 100644 index 0000000..221bbc0 --- /dev/null +++ b/docs/api/lock.md @@ -0,0 +1,5 @@ +# Distributed lock + +Distributed locking using backend atomic primitives. + +::: fastapi_cachex.lock.CacheLock diff --git a/fastapi_cachex/lock.py b/fastapi_cachex/lock.py index af7bc16..72367fa 100644 --- a/fastapi_cachex/lock.py +++ b/fastapi_cachex/lock.py @@ -6,22 +6,31 @@ import time from types import TracebackType -from typing_extensions import Self - from fastapi_cachex.backends.base import BaseCacheBackend from fastapi_cachex.exceptions import LockTimeoutError from fastapi_cachex.proxy import BackendProxy from fastapi_cachex.types import CacheEntry -from fastapi_cachex.types import lock_entry logger = logging.getLogger(__name__) +# Fingerprint every backend reports for a key that holds a lock entry +_LOCK_FINGERPRINT = "lock" + + +def _lock_entry(token: str) -> CacheEntry: + """Wrap a lock token string in the entry model shared by every backend.""" + return CacheEntry(fingerprint=_LOCK_FINGERPRINT, content=token.encode("utf-8")) + class CacheLock: """Distributed lock helper built on backend primitives. Can be used as an async context manager or with direct acquire/release calls. + Note: + A single CacheLock instance cannot be shared across concurrent tasks or + re-entered. Instantiate a new CacheLock for each acquisition. + Args: name: Unique lock identifier ttl: Time to live in seconds (default: 60) @@ -51,7 +60,8 @@ def __init__( self._backend = backend self.key_prefix = key_prefix self._token = secrets.token_hex(16) - self._entry: CacheEntry = lock_entry(self._token) + self._entry: CacheEntry = _lock_entry(self._token) + self._is_held = False @property def key(self) -> str: @@ -75,13 +85,24 @@ async def acquire( Args: blocking: Override default blocking mode - timeout: Override default timeout (seconds) + timeout: Override default timeout (seconds). If timeout=None, the lock will + use the instance's default timeout. poll_interval: Override default poll interval (seconds) ttl: Override default TTL (seconds) Returns: True if the lock was acquired, False on timeout or failure + + Raises: + RuntimeError: If this CacheLock instance is already held. """ + if self._is_held: + msg = ( + "This CacheLock instance is already held. Use a new CacheLock instance " + "for each acquisition." + ) + raise RuntimeError(msg) + is_blocking = self.blocking if blocking is None else blocking effective_timeout = self.timeout if timeout is None else timeout interval = self.poll_interval if poll_interval is None else poll_interval @@ -90,11 +111,17 @@ async def acquire( backend = self._get_backend() if not is_blocking: - return await backend.set_if_absent(self.key, self._entry, ttl=effective_ttl) + acquired = await backend.set_if_absent( + self.key, self._entry, ttl=effective_ttl + ) + if acquired: + self._is_held = True + return acquired start = time.monotonic() while True: if await backend.set_if_absent(self.key, self._entry, ttl=effective_ttl): + self._is_held = True return True if effective_timeout is not None: elapsed = time.monotonic() - start @@ -112,6 +139,7 @@ async def release(self) -> bool: Returns: True if the lock was released, False if expired or owned by another caller """ + self._is_held = False backend = self._get_backend() released = await backend.delete_if_equals(self.key, self._entry) if not released: @@ -143,7 +171,7 @@ async def locked(self) -> bool: backend = self._get_backend() return (await backend.get(self.key)) is not None - async def __aenter__(self) -> Self: + async def __aenter__(self) -> "CacheLock": """Acquire lock as async context manager. Raises: diff --git a/fastapi_cachex/types.py b/fastapi_cachex/types.py index 9b507c5..5fb6d74 100644 --- a/fastapi_cachex/types.py +++ b/fastapi_cachex/types.py @@ -71,12 +71,3 @@ def counter_value(entry: CacheEntry) -> int: except ValueError as e: msg = "Cache key holds a value that is not a counter" raise CacheXError(msg) from e - - -# Fingerprint every backend reports for a key that holds a lock entry -LOCK_FINGERPRINT = "lock" - - -def lock_entry(token: str) -> CacheEntry: - """Wrap a lock token string in the entry model shared by every backend.""" - return CacheEntry(fingerprint=LOCK_FINGERPRINT, content=token.encode("utf-8")) diff --git a/pyproject.toml b/pyproject.toml index 6ff777a..c2b4830 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -127,7 +127,7 @@ keep-runtime-typing = true ] "fastapi_cachex/backends/memcached.py" = ["PLC0415"] # Optional dependency "fastapi_cachex/backends/redis.py" = ["PLR0913", "PLC0415"] # Optional dependency, Redis config -"fastapi_cachex/lock.py" = ["PLR0913"] # Configurable lock parameters +"fastapi_cachex/lock.py" = ["PLR0913", "PYI034"] # Configurable lock parameters, string return annotation for __aenter__ "fastapi_cachex/proxy.py" = ["PLW0603"] # Global backend management by design "fastapi_cachex/session/manager.py" = ["PLR0915"] # Session/security validation branches needed diff --git a/tests/test_lock.py b/tests/test_lock.py index e58f2f9..1b2c3ec 100644 --- a/tests/test_lock.py +++ b/tests/test_lock.py @@ -1,6 +1,8 @@ """Tests for the CacheLock distributed lock helper.""" +import asyncio import time +from typing import TYPE_CHECKING import pytest @@ -9,6 +11,13 @@ from fastapi_cachex.exceptions import LockTimeoutError from fastapi_cachex.lock import CacheLock from fastapi_cachex.proxy import BackendProxy +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 import AsyncRedisCacheBackend + from fastapi_cachex.backends import MemcachedBackend @pytest.mark.asyncio @@ -44,8 +53,6 @@ async def test_lock_acquire_non_blocking() -> None: @pytest.mark.asyncio async def test_lock_acquire_blocking_indefinite_retries_until_available() -> None: - import asyncio - backend = MemoryBackend() BackendProxy.set(backend) @@ -181,3 +188,65 @@ async def test_lock_raises_backend_not_found_if_no_backend_set() -> None: with pytest.raises(BackendNotFoundError): await lock.acquire(blocking=False) + + +@pytest.mark.asyncio +async def test_lock_reentry_raises_runtime_error() -> None: + backend = MemoryBackend() + BackendProxy.set(backend) + + lock = CacheLock("job", ttl=30) + assert await lock.acquire(blocking=False) is True + + with pytest.raises( + RuntimeError, + match="This CacheLock instance is already held", + ): + await lock.acquire(blocking=False) + + assert await lock.release() is True + + +@pytest.mark.asyncio +async def test_lock_shared_instance_raises_runtime_error() -> None: + backend = MemoryBackend() + BackendProxy.set(backend) + + lock = CacheLock("shared_job", ttl=30) + + async def task_worker() -> None: + async with lock: + await asyncio.sleep(0.05) + + results = await asyncio.gather(task_worker(), task_worker(), return_exceptions=True) + assert any(isinstance(res, RuntimeError) for res in results) + assert any(res is None for res in results) + + +@requires_redis +@requires_redis_package +@pytest.mark.asyncio +async def test_lock_lifecycle_with_redis( + async_redis_backend: "AsyncRedisCacheBackend", +) -> None: + lock1 = CacheLock("redis_job", ttl=30, backend=async_redis_backend) + lock2 = CacheLock("redis_job", ttl=30, backend=async_redis_backend) + + assert await lock1.acquire(blocking=False) is True + assert await lock2.acquire(blocking=False) is False + assert await lock1.extend(60) is True + assert await lock1.release() is True + + +@requires_memcached +@pytest.mark.asyncio +async def test_lock_lifecycle_with_memcached( + memcached_backend: "MemcachedBackend", +) -> None: + lock1 = CacheLock("memcached_job", ttl=30, backend=memcached_backend) + lock2 = CacheLock("memcached_job", ttl=30, backend=memcached_backend) + + assert await lock1.acquire(blocking=False) is True + assert await lock2.acquire(blocking=False) is False + assert await lock1.extend(60) is True + assert await lock1.release() is True diff --git a/zensical.toml b/zensical.toml index 0295cf4..95c0b9c 100644 --- a/zensical.toml +++ b/zensical.toml @@ -26,11 +26,13 @@ nav = [ { "Session management" = "SESSION.md" }, { "JWT claims" = "JWT_CLAIMS.md" }, { "State management" = "STATE.md" }, + { "Distributed lock" = "LOCK.md" }, ] }, { "API reference" = [ { "HTTP caching" = "api/http-caching.md" }, { "Backends" = "api/backends.md" }, { "CacheManager" = "api/cache-manager.md" }, + { "CacheLock" = "api/lock.md" }, { "Session" = "api/session.md" }, { "State" = "api/state.md" }, { "Types and exceptions" = "api/types.md" }, diff --git a/zensical.zh-TW.toml b/zensical.zh-TW.toml index 897a0ec..c32c3f6 100644 --- a/zensical.zh-TW.toml +++ b/zensical.zh-TW.toml @@ -30,6 +30,7 @@ nav = [ { "Session 管理" = "SESSION.md" }, { "JWT claims" = "JWT_CLAIMS.md" }, { "OAuth state" = "STATE.md" }, + { "Distributed lock" = "LOCK.md" }, ] }, { "API 參考(英文)" = "https://fastapi-cachex.readthedocs.io/en/latest/api/http-caching/" }, { "開發(英文)" = [ From 4c4024ff5b98210e3a8255f191f7c282853b514a Mon Sep 17 00:00:00 2001 From: Shivansh Shukla Date: Fri, 25 Sep 2026 23:23:08 +0530 Subject: [PATCH 3/7] fix: address maintainer review #2 feedback for CacheLock (#64) - Set _is_held = True synchronously at entry of acquire() before the first await to prevent shared-instance race conditions. - Format CHANGELOG.md entries with bold summaries and link #64 on both. - Revert zensical.zh-TW.toml edit. - Update docs/LOCK.md to detail TTL expiration behavior and default timeout=None indefinite blocking. - Move PYI034 ignore inline on __aenter__ line in lock.py. - Update test_lock_shared_instance_raises_runtime_error in tests/test_lock.py to use YieldingMemoryBackend. --- CHANGELOG.md | 4 +-- docs/LOCK.md | 3 +- fastapi_cachex/lock.py | 67 +++++++++++++++++++++++------------------- pyproject.toml | 2 +- tests/test_lock.py | 11 ++++++- zensical.zh-TW.toml | 1 - 6 files changed, 52 insertions(+), 36 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index d10bd97..cd9d94a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -17,8 +17,8 @@ Note that 0.3.3 was never released; 0.3.4 follows 0.3.2. ### Added -- `CacheLock` distributed lock async context manager helper built on backend primitives, with `acquire`, `release`, `extend`, `locked`, and `LockTimeoutError`. ([#64](https://github.com/allen0099/FastAPI-CacheX/issues/64)) -- `expire_if_equals(key, expected, ttl) -> bool` atomic backend primitive added to `BaseCacheBackend`, `MemoryBackend`, `AsyncRedisCacheBackend`, and `MemcachedBackend`. +- **`CacheLock`, a distributed lock built on backend primitives.** Usable as an async context manager or with direct `acquire`/`release`/`extend`/`locked` calls, raising `LockTimeoutError` on timeout. ([#64](https://github.com/allen0099/FastAPI-CacheX/issues/64)) +- **`expire_if_equals()` backend primitive for owner-checked TTL renewal.** Added to `BaseCacheBackend`, `MemoryBackend`, `AsyncRedisCacheBackend`, and `MemcachedBackend`. ([#64](https://github.com/allen0099/FastAPI-CacheX/issues/64)) ## [0.3.7] - 2026-09-25 diff --git a/docs/LOCK.md b/docs/LOCK.md index 1f60bc8..0a74702 100644 --- a/docs/LOCK.md +++ b/docs/LOCK.md @@ -26,8 +26,9 @@ finally: - **Safety & Token Ownership**: Each `CacheLock` instance generates a unique token (`secrets.token_hex(16)`) stored inside a `CacheEntry`. Releases (`release()`) and extensions (`extend()`) use owner-checked backend primitives (`delete_if_equals` and `expire_if_equals`), so a holder whose lock expired cannot release or renew a lock claimed by someone else. - **Blocking & Non-blocking Modes**: - Non-blocking (`acquire(blocking=False)`): Performs a single atomic `set_if_absent` and immediately returns `True` if acquired or `False` if held. - - Blocking (`acquire(blocking=True, timeout=None, poll_interval=0.1)`): Retries at `poll_interval` seconds until acquired or until `timeout` seconds elapse. If `timeout` is reached, `acquire()` returns `False`. + - Blocking (`acquire(blocking=True, timeout=None, poll_interval=0.1)`): Retries at `poll_interval` seconds until acquired or until `timeout` seconds elapse. The default `timeout=None` makes a blocking `acquire()` (and `async with CacheLock(...)`) wait indefinitely until the lock becomes free. If a finite `timeout` is reached, `acquire()` returns `False`. - **Context Manager Timeouts**: Entering a context manager (`async with CacheLock(...)`) invokes `acquire()`. If acquisition fails or times out, it raises `LockTimeoutError`. +- **TTL Expiration**: If a task takes longer than its `ttl` and fails to renew, the lock entry expires in the backend and becomes free. Another process or container can then acquire the lock while the original code is still running. Subsequent calls to `extend()` or `release()` by the original holder will safely return `False` without throwing an error. Always choose a `ttl` longer than the expected work, or call `extend()` periodically during long-running operations. - **TTL Renewal (`extend`)**: `extend(ttl)` updates the key's TTL only while the lock is still owned by this holder instance, preventing race conditions on expired locks. - **One Instance per Acquisition Rule**: A single `CacheLock` instance tracks its active ownership state. Re-entering or sharing a single `CacheLock` instance across concurrent tasks raises a `RuntimeError`. Instantiate a new `CacheLock` instance for each acquisition. - **Namespace**: Lock keys live under their own `lock:` prefix by default (e.g. `lock:report:123`), separate from `cache:` and `oauth_state:`. diff --git a/fastapi_cachex/lock.py b/fastapi_cachex/lock.py index 72367fa..a1f356b 100644 --- a/fastapi_cachex/lock.py +++ b/fastapi_cachex/lock.py @@ -103,35 +103,42 @@ async def acquire( ) raise RuntimeError(msg) - is_blocking = self.blocking if blocking is None else blocking - effective_timeout = self.timeout if timeout is None else timeout - interval = self.poll_interval if poll_interval is None else poll_interval - effective_ttl = self.ttl if ttl is None else ttl - - backend = self._get_backend() - - if not is_blocking: - acquired = await backend.set_if_absent( - self.key, self._entry, ttl=effective_ttl - ) - if acquired: - self._is_held = True - return acquired - - start = time.monotonic() - while True: - if await backend.set_if_absent(self.key, self._entry, ttl=effective_ttl): - self._is_held = True - return True - if effective_timeout is not None: - elapsed = time.monotonic() - start - if elapsed >= effective_timeout: - return False - remaining = effective_timeout - elapsed - sleep_time = min(interval, remaining) - else: - sleep_time = interval - await asyncio.sleep(sleep_time) + self._is_held = True + try: + is_blocking = self.blocking if blocking is None else blocking + effective_timeout = self.timeout if timeout is None else timeout + interval = self.poll_interval if poll_interval is None else poll_interval + effective_ttl = self.ttl if ttl is None else ttl + + backend = self._get_backend() + + if not is_blocking: + acquired = await backend.set_if_absent( + self.key, self._entry, ttl=effective_ttl + ) + if not acquired: + self._is_held = False + return acquired + + start = time.monotonic() + while True: + if await backend.set_if_absent( + self.key, self._entry, ttl=effective_ttl + ): + return True + if effective_timeout is not None: + elapsed = time.monotonic() - start + if elapsed >= effective_timeout: + self._is_held = False + return False + remaining = effective_timeout - elapsed + sleep_time = min(interval, remaining) + else: + sleep_time = interval + await asyncio.sleep(sleep_time) + except Exception: + self._is_held = False + raise async def release(self) -> bool: """Release the lock if still held by this instance. @@ -171,7 +178,7 @@ async def locked(self) -> bool: backend = self._get_backend() return (await backend.get(self.key)) is not None - async def __aenter__(self) -> "CacheLock": + async def __aenter__(self) -> "CacheLock": # noqa: PYI034 """Acquire lock as async context manager. Raises: diff --git a/pyproject.toml b/pyproject.toml index c2b4830..6ff777a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -127,7 +127,7 @@ keep-runtime-typing = true ] "fastapi_cachex/backends/memcached.py" = ["PLC0415"] # Optional dependency "fastapi_cachex/backends/redis.py" = ["PLR0913", "PLC0415"] # Optional dependency, Redis config -"fastapi_cachex/lock.py" = ["PLR0913", "PYI034"] # Configurable lock parameters, string return annotation for __aenter__ +"fastapi_cachex/lock.py" = ["PLR0913"] # Configurable lock parameters "fastapi_cachex/proxy.py" = ["PLW0603"] # Global backend management by design "fastapi_cachex/session/manager.py" = ["PLR0915"] # Session/security validation branches needed diff --git a/tests/test_lock.py b/tests/test_lock.py index 1b2c3ec..4a70c37 100644 --- a/tests/test_lock.py +++ b/tests/test_lock.py @@ -11,6 +11,7 @@ from fastapi_cachex.exceptions import LockTimeoutError from fastapi_cachex.lock import CacheLock from fastapi_cachex.proxy import BackendProxy +from fastapi_cachex.types import CacheEntry from tests.live_servers import requires_memcached from tests.live_servers import requires_redis from tests.live_servers import requires_redis_package @@ -207,9 +208,17 @@ async def test_lock_reentry_raises_runtime_error() -> None: assert await lock.release() is True +class YieldingMemoryBackend(MemoryBackend): + async def set_if_absent( + self, key: str, value: CacheEntry, ttl: int | None = None + ) -> bool: + await asyncio.sleep(0) + return await super().set_if_absent(key, value, ttl=ttl) + + @pytest.mark.asyncio async def test_lock_shared_instance_raises_runtime_error() -> None: - backend = MemoryBackend() + backend = YieldingMemoryBackend() BackendProxy.set(backend) lock = CacheLock("shared_job", ttl=30) diff --git a/zensical.zh-TW.toml b/zensical.zh-TW.toml index c32c3f6..897a0ec 100644 --- a/zensical.zh-TW.toml +++ b/zensical.zh-TW.toml @@ -30,7 +30,6 @@ nav = [ { "Session 管理" = "SESSION.md" }, { "JWT claims" = "JWT_CLAIMS.md" }, { "OAuth state" = "STATE.md" }, - { "Distributed lock" = "LOCK.md" }, ] }, { "API 參考(英文)" = "https://fastapi-cachex.readthedocs.io/en/latest/api/http-caching/" }, { "開發(英文)" = [ From 3c5efea0ef3d076795b8cd1d61be3c2b8154ee66 Mon Sep 17 00:00:00 2001 From: Shivansh Shukla Date: Sat, 26 Sep 2026 13:46:57 +0530 Subject: [PATCH 4/7] test: move CacheLock live backend tests to backend test files (#64) --- tests/backends/test_memcached.py | 16 ++++++++++++++ tests/backends/test_redis.py | 17 +++++++++++++++ tests/test_lock.py | 37 -------------------------------- 3 files changed, 33 insertions(+), 37 deletions(-) diff --git a/tests/backends/test_memcached.py b/tests/backends/test_memcached.py index e40db88..643258f 100644 --- a/tests/backends/test_memcached.py +++ b/tests/backends/test_memcached.py @@ -7,6 +7,7 @@ from fastapi_cachex.backends import MemcachedBackend from fastapi_cachex.exceptions import CacheXError +from fastapi_cachex.lock import CacheLock from fastapi_cachex.types import CacheEntry from fastapi_cachex.types import counter_entry from tests.live_servers import MEMCACHED_SERVER @@ -708,3 +709,18 @@ def gets_then_overwrite(key): raw = await asyncio.to_thread(client.get, memcached_backend._make_key("slot")) assert raw == b'{"overwritten": true}' + + +@requires_memcached +@pytest.mark.asyncio +async def test_lock_lifecycle_with_memcached( + memcached_backend: MemcachedBackend, +) -> None: + lock1 = CacheLock("memcached_job", ttl=30, backend=memcached_backend) + lock2 = CacheLock("memcached_job", ttl=30, backend=memcached_backend) + + assert await lock1.acquire(blocking=False) is True + assert await lock2.acquire(blocking=False) is False + assert await lock1.extend(60) is True + assert await lock1.release() is True + diff --git a/tests/backends/test_redis.py b/tests/backends/test_redis.py index 93acacd..24b4c8e 100644 --- a/tests/backends/test_redis.py +++ b/tests/backends/test_redis.py @@ -10,6 +10,7 @@ from fastapi_cachex.backends import AsyncRedisCacheBackend from fastapi_cachex.backends.redis import _BATCH_SIZE from fastapi_cachex.exceptions import CacheXError +from fastapi_cachex.lock import CacheLock from fastapi_cachex.types import CacheEntry from fastapi_cachex.types import counter_entry from tests.live_servers import REDIS_HOST @@ -1017,3 +1018,19 @@ def make(prefix: str) -> AsyncRedisCacheBackend: finally: await globbed.clear() await other.clear() + + +@requires_redis +@requires_redis_package +@pytest.mark.asyncio +async def test_lock_lifecycle_with_redis( + async_redis_backend: AsyncRedisCacheBackend, +) -> None: + lock1 = CacheLock("redis_job", ttl=30, backend=async_redis_backend) + lock2 = CacheLock("redis_job", ttl=30, backend=async_redis_backend) + + assert await lock1.acquire(blocking=False) is True + assert await lock2.acquire(blocking=False) is False + assert await lock1.extend(60) is True + assert await lock1.release() is True + diff --git a/tests/test_lock.py b/tests/test_lock.py index 4a70c37..88515be 100644 --- a/tests/test_lock.py +++ b/tests/test_lock.py @@ -2,7 +2,6 @@ import asyncio import time -from typing import TYPE_CHECKING import pytest @@ -12,13 +11,6 @@ from fastapi_cachex.lock import CacheLock from fastapi_cachex.proxy import BackendProxy from fastapi_cachex.types import CacheEntry -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 import AsyncRedisCacheBackend - from fastapi_cachex.backends import MemcachedBackend @pytest.mark.asyncio @@ -230,32 +222,3 @@ async def task_worker() -> None: results = await asyncio.gather(task_worker(), task_worker(), return_exceptions=True) assert any(isinstance(res, RuntimeError) for res in results) assert any(res is None for res in results) - - -@requires_redis -@requires_redis_package -@pytest.mark.asyncio -async def test_lock_lifecycle_with_redis( - async_redis_backend: "AsyncRedisCacheBackend", -) -> None: - lock1 = CacheLock("redis_job", ttl=30, backend=async_redis_backend) - lock2 = CacheLock("redis_job", ttl=30, backend=async_redis_backend) - - assert await lock1.acquire(blocking=False) is True - assert await lock2.acquire(blocking=False) is False - assert await lock1.extend(60) is True - assert await lock1.release() is True - - -@requires_memcached -@pytest.mark.asyncio -async def test_lock_lifecycle_with_memcached( - memcached_backend: "MemcachedBackend", -) -> None: - lock1 = CacheLock("memcached_job", ttl=30, backend=memcached_backend) - lock2 = CacheLock("memcached_job", ttl=30, backend=memcached_backend) - - assert await lock1.acquire(blocking=False) is True - assert await lock2.acquire(blocking=False) is False - assert await lock1.extend(60) is True - assert await lock1.release() is True From 025b552ac36670e6b6f5b4c17d5cc600d7c4b177 Mon Sep 17 00:00:00 2001 From: allen0099 Date: Sat, 26 Sep 2026 08:33:45 +0000 Subject: [PATCH 5/7] style: drop the trailing blank lines ruff format flags in the lock tests Lint failed on ruff format --check for the two CacheLock lifecycle tests moved into tests/backends/. --- tests/backends/test_memcached.py | 1 - tests/backends/test_redis.py | 1 - 2 files changed, 2 deletions(-) diff --git a/tests/backends/test_memcached.py b/tests/backends/test_memcached.py index 643258f..4008223 100644 --- a/tests/backends/test_memcached.py +++ b/tests/backends/test_memcached.py @@ -723,4 +723,3 @@ async def test_lock_lifecycle_with_memcached( assert await lock2.acquire(blocking=False) is False assert await lock1.extend(60) is True assert await lock1.release() is True - diff --git a/tests/backends/test_redis.py b/tests/backends/test_redis.py index 24b4c8e..db67bcd 100644 --- a/tests/backends/test_redis.py +++ b/tests/backends/test_redis.py @@ -1033,4 +1033,3 @@ async def test_lock_lifecycle_with_redis( assert await lock2.acquire(blocking=False) is False assert await lock1.extend(60) is True assert await lock1.release() is True - From 09b1fdb589eb61186463dbc146b8db09b99f8eca Mon Sep 17 00:00:00 2001 From: Shivansh Shukla Date: Sat, 26 Sep 2026 14:45:52 +0530 Subject: [PATCH 6/7] fix: reset CacheLock._is_held on BaseException to handle task cancellation (#64) --- fastapi_cachex/lock.py | 2 +- tests/test_lock.py | 24 ++++++++++++++++++++++++ 2 files changed, 25 insertions(+), 1 deletion(-) diff --git a/fastapi_cachex/lock.py b/fastapi_cachex/lock.py index a1f356b..f5c0f88 100644 --- a/fastapi_cachex/lock.py +++ b/fastapi_cachex/lock.py @@ -136,7 +136,7 @@ async def acquire( else: sleep_time = interval await asyncio.sleep(sleep_time) - except Exception: + except BaseException: self._is_held = False raise diff --git a/tests/test_lock.py b/tests/test_lock.py index 88515be..08edcda 100644 --- a/tests/test_lock.py +++ b/tests/test_lock.py @@ -222,3 +222,27 @@ async def task_worker() -> None: results = await asyncio.gather(task_worker(), task_worker(), return_exceptions=True) assert any(isinstance(res, RuntimeError) for res in results) assert any(res is None for res in results) + + +@pytest.mark.asyncio +async def test_lock_acquire_cancelled_resets_is_held() -> None: + backend = MemoryBackend() + BackendProxy.set(backend) + + lock1 = CacheLock("job", ttl=30) + lock2 = CacheLock("job", ttl=30, poll_interval=0.01) + + assert await lock1.acquire(blocking=False) is True + + task = asyncio.create_task(lock2.acquire(blocking=True, poll_interval=0.01)) + await asyncio.sleep(0.02) + + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + await lock1.release() + + assert await lock2.acquire(blocking=False) is True + assert await lock2.release() is True + From 7cfe73f6d0e2a3985683562c81ac8e804af8eab4 Mon Sep 17 00:00:00 2001 From: allen0099 Date: Sat, 26 Sep 2026 10:21:29 +0000 Subject: [PATCH 7/7] style: drop the trailing blank line ruff format flags in tests/test_lock.py --- tests/test_lock.py | 1 - 1 file changed, 1 deletion(-) diff --git a/tests/test_lock.py b/tests/test_lock.py index 08edcda..dba7299 100644 --- a/tests/test_lock.py +++ b/tests/test_lock.py @@ -245,4 +245,3 @@ async def test_lock_acquire_cancelled_resets_is_held() -> None: assert await lock2.acquire(blocking=False) is True assert await lock2.release() is True -