diff --git a/CHANGELOG.md b/CHANGELOG.md index afeadfe..5d320da 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -17,6 +17,11 @@ Note that 0.3.3 was never released; 0.3.4 follows 0.3.2. ## [Unreleased] +### Added + +- **`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)) + ### Changed - **GitHub release notes list one line per change.** Each changelog entry now 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/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/docs/LOCK.md b/docs/LOCK.md new file mode 100644 index 0000000..0a74702 --- /dev/null +++ b/docs/LOCK.md @@ -0,0 +1,40 @@ +# 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. 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:`. + +> [!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/__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..f5c0f88 --- /dev/null +++ b/fastapi_cachex/lock.py @@ -0,0 +1,200 @@ +"""Distributed lock helper for FastAPI-CacheX.""" + +import asyncio +import logging +import secrets +import time +from types import TracebackType + +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 + +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) + 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) + self._is_held = False + + @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). 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) + + 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 BaseException: + self._is_held = False + raise + + 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 + """ + self._is_held = False + 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) -> "CacheLock": # noqa: PYI034 + """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/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..4008223 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 @@ -667,3 +668,58 @@ 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}' + + +@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_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..db67bcd 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 @@ -886,6 +887,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.""" @@ -975,3 +1018,18 @@ 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 new file mode 100644 index 0000000..dba7299 --- /dev/null +++ b/tests/test_lock.py @@ -0,0 +1,247 @@ +"""Tests for the CacheLock distributed lock helper.""" + +import asyncio +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 +from fastapi_cachex.types import CacheEntry + + +@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: + 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) + + +@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 + + +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 = YieldingMemoryBackend() + 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) + + +@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 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" },