From 9b8ee01606906e7e8e63fda35e154f2f1d9fa16a Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 7 Oct 2026 10:41:33 +0000 Subject: [PATCH] feat: @cached decorator for plain functions (#248) `@cached(ttl=..., key=..., manager=None)` caches what an async or sync function returns through `CacheManager.get_or_set()`, so it gets the manager's prefix, JSON round-trip and stampede protection. The default key is `module.qualname:` plus a SHA-256 of the JSON-serialized arguments bound to the signature; `key=` takes a format template or a callable for methods and non-JSON arguments. The decorated function has `cache_key()` and `invalidate()`, and binds as a method. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01UTXhK1BtwXTBDaouEcKtpW --- README.md | 3 +- changelog.d/248.added.md | 9 + docs/APP_CACHE.md | 37 ++++ docs/api/cache-manager.md | 4 + examples/app_cache.py | 29 ++- fastapi_cachex/__init__.py | 4 + fastapi_cachex/cached.py | 261 +++++++++++++++++++++++++++ i18n/zh-TW/docs/APP_CACHE.md | 15 ++ tests/test_cached.py | 335 +++++++++++++++++++++++++++++++++++ tests/test_examples.py | 8 + 10 files changed, 703 insertions(+), 2 deletions(-) create mode 100644 changelog.d/248.added.md create mode 100644 fastapi_cachex/cached.py create mode 100644 tests/test_cached.py diff --git a/README.md b/README.md index 4a4075c..a530880 100644 --- a/README.md +++ b/README.md @@ -23,7 +23,8 @@ A high-performance caching extension for FastAPI: a server-side response cache w - **HTTP caching** — a `@cache` decorator for GET routes with `Cache-Control`, `ETag` / `If-None-Match` (304) and per-route invalidation. - **Application cache** — `CacheManager` for caching arbitrary JSON values in - your own code, with compute-on-miss `get_or_set()` and atomic store-if-absent `add()`. + your own code, with compute-on-miss `get_or_set()` and atomic store-if-absent `add()`, + and `@cached` for a plain function, keyed on its arguments. - **Backends** — in-memory, Redis and Memcached, with atomic counters, one-shot values and locks. - **Sessions and OAuth state (deprecated)**: signed session tokens and one-time diff --git a/changelog.d/248.added.md b/changelog.d/248.added.md new file mode 100644 index 0000000..00f2f39 --- /dev/null +++ b/changelog.d/248.added.md @@ -0,0 +1,9 @@ +**`@cached` caches a plain function's result, keyed on its arguments.** +`@cached(ttl=60)` on an `async` or sync function stores its result through +`CacheManager.get_or_set()`, so concurrent misses run it once; the decorated +function is always awaited. The default key is `module.qualname:` plus a +SHA-256 of the JSON-serialized arguments, bound to the signature with defaults +applied; `key="user:{user_id}"` or `key=lambda ...: ...` names it instead, for +a method or a non-JSON argument. `fn.cache_key(...)` returns the key a call +uses and `await fn.invalidate(...)` drops its value. `manager=` picks the +`CacheManager`; without it the application's (`AppCache`) is used. diff --git a/docs/APP_CACHE.md b/docs/APP_CACHE.md index 315230f..003aa30 100644 --- a/docs/APP_CACHE.md +++ b/docs/APP_CACHE.md @@ -44,6 +44,43 @@ await manager.clear_pattern("user:*") # matches "myapp:user:*" Complete runnable example: [`examples/app_cache.py`](https://github.com/allen0099/FastAPI-CacheX/blob/master/examples/app_cache.py). +## Caching a function + +`@cached` does what `get_or_set()` does for a plain function, keyed on its +arguments, so a loader or a call to another service is written once and +cached wherever it is called: + + +```python +--8<-- "examples/app_cache.py:cached" +``` + + +- The decorated function is always `await`-ed, even if it was `def` (a sync + function then runs on the event loop, as a `get_or_set()` factory does). + Its result goes through the manager's [JSON round-trip](#json-round-trip), + so it must be JSON-serializable and comes back as JSON gives it, on the + first call too. +- `ttl` takes seconds or a `timedelta`; without it the manager's + `default_ttl` applies. `manager=` names the `CacheManager` to store + through; without it the application's is used (what `AppCache` returns), + resolved on each call, so one registered with `CacheManagerProxy.set()` at + startup is picked up. `lock=False` skips the manager's + [stampede protection](#stampede-protection) for that function. +- Without `key=` the key is `module.qualname:` followed by a SHA-256 of the + arguments, bound to the signature with defaults applied, so `load(1)`, + `load(user_id=1)` and `load(1, locale="en")` share one entry. The + arguments are hashed as JSON, so they must be JSON-serializable; a call + with one that is not raises `CacheXError`. `key="user:{user_id}"` is a + `str.format` template over the arguments by name (a plain string is a fixed + key), and `key=lambda self, user_id: f"user:{user_id}"` is called with + them: use either for a method, whose `self` cannot be hashed, or for a + model. The manager's `key_prefix` goes in front, as for every key it stores. +- `fn.cache_key(*args, **kwargs)` is the key a call uses (without the + prefix) and `await fn.invalidate(*args, **kwargs)` drops its value, + returning whether one was cached. On a method, `obj.load.invalidate(1)` + binds `obj` as on a call. + ## Behavior - `get()` returns `None` (or a supplied `default=`) on a cache miss — it never diff --git a/docs/api/cache-manager.md b/docs/api/cache-manager.md index b6056b4..fb85f88 100644 --- a/docs/api/cache-manager.md +++ b/docs/api/cache-manager.md @@ -4,6 +4,10 @@ Application-level caching of arbitrary JSON-serializable values. ::: fastapi_cachex.manager.CacheManager +::: fastapi_cachex.cached.cached + +::: fastapi_cachex.cached.CachedFunction + ::: fastapi_cachex.manager_proxy.CacheManagerProxy options: inherited_members: true diff --git a/examples/app_cache.py b/examples/app_cache.py index da9d1d0..b0d7c3d 100644 --- a/examples/app_cache.py +++ b/examples/app_cache.py @@ -2,7 +2,8 @@ ``get_or_set`` computes a value once and serves it from the cache until it expires; ``add`` stores a key only if it is absent, which makes a simple -idempotency check. Both come through the ``AppCache`` dependency. +idempotency check. Both come through the ``AppCache`` dependency. ``@cached`` +does the same for a plain function, keyed on its arguments. Run it from a checkout (see ``examples/README.md``):: @@ -22,6 +23,7 @@ from fastapi_cachex import BackendProxy from fastapi_cachex import CacheManager from fastapi_cachex import CacheManagerProxy +from fastapi_cachex import cached from fastapi_cachex.backends import MemoryBackend backend = MemoryBackend() @@ -77,3 +79,28 @@ async def create_order( async def forget_rates(app_cache: AppCache) -> dict[str, bool]: """Drop the cached rates; the next read calls the upstream again.""" return {"deleted": await app_cache.delete("rates")} + + +# --8<-- [start:cached] +# Cache a plain function on its arguments. The value goes through the +# registered CacheManager, so it gets its prefix, TTL default and lock. +@cached(ttl=60, key="rate:{currency}") +async def fetch_rate(currency: str) -> float: + """Stand-in for a slow call to another service, once per currency.""" + upstream_calls["rates"] += 1 + return {"EUR": 0.92, "JPY": 151.3}.get(currency.upper(), 1.0) + + +@app.get("/rates/{currency}") +async def read_rate(currency: str) -> dict[str, float]: + """``fetch_rate`` runs once per currency until its value expires.""" + return {currency: await fetch_rate(currency)} + + +@app.delete("/rates/{currency}") +async def forget_rate(currency: str) -> dict[str, bool]: + """Drop the value cached for this currency: ``invalidate`` takes the same arguments.""" + return {"deleted": await fetch_rate.invalidate(currency)} + + +# --8<-- [end:cached] diff --git a/fastapi_cachex/__init__.py b/fastapi_cachex/__init__.py index f793606..295839e 100644 --- a/fastapi_cachex/__init__.py +++ b/fastapi_cachex/__init__.py @@ -11,6 +11,8 @@ from .cache import default_key_builder as default_key_builder from .cache import invalidate as invalidate from .cache_key import CacheKey as CacheKey +from .cached import CachedFunction as CachedFunction +from .cached import cached as cached from .dependencies import AppCache as AppCache from .dependencies import CacheBackend as CacheBackend from .dependencies import get_app_cache as get_app_cache @@ -135,6 +137,7 @@ def __getattr__(name: str) -> object: "CacheManager", "CacheManagerProxy", "CacheXError", + "CachedFunction", "LockTimeoutError", "ProxyNotSetError", "RequestNotFoundError", @@ -142,6 +145,7 @@ def __getattr__(name: str) -> object: "add_routes", "build_cache_key", "cache", + "cached", "default_key_builder", "get_app_cache", "get_cache_backend", diff --git a/fastapi_cachex/cached.py b/fastapi_cachex/cached.py new file mode 100644 index 0000000..9d816c0 --- /dev/null +++ b/fastapi_cachex/cached.py @@ -0,0 +1,261 @@ +"""``@cached``: cache the result of a plain function through ``CacheManager``.""" + +import hashlib +import inspect +import json +import logging +from collections.abc import Awaitable +from collections.abc import Callable +from datetime import timedelta +from functools import wraps +from typing import Any +from typing import Generic +from typing import ParamSpec +from typing import Protocol +from typing import TypeVar +from typing import overload + +from .backends.base import validate_ttl +from .dependencies import get_app_cache +from .exceptions import CacheXError +from .manager import CacheManager + +logger = logging.getLogger(__name__) + +P = ParamSpec("P") +R = TypeVar("R") + +# Separates the function's name from the digest of its arguments in a default +# key, as `CacheManager` keys separate their parts. +_KEY_SEPARATOR = ":" + + +class CachedFunction(Generic[P, R]): + """What ``@cached`` turns a function into. + + Calling it is always ``await``-ed, whether the function was ``async`` or + not, and returns the cached value: on a miss the function runs through + ``CacheManager.get_or_set()``, with its stampede protection, and its + result is stored. ``cache_key()`` and ``invalidate()`` take the same + arguments as the function. + """ + + # Copied from the function by `functools.wraps`. + __name__: str + __qualname__: str + + def __init__( + self, + func: Callable[P, Awaitable[R]] | Callable[P, R], + *, + ttl: int | None, + key: str | Callable[P, str] | None, + manager: CacheManager | None, + lock: bool | None, + ) -> None: + """Wrap ``func``; ``cached()`` builds these.""" + self.__wrapped__ = func + self._signature = inspect.signature(func) + self._ttl = ttl + self._key = key + self._manager = manager + self._lock = lock + self._name = f"{func.__module__}.{func.__qualname__}" + wraps(func)(self) + + @property + def manager(self) -> CacheManager: + """The ``CacheManager`` the values go through. + + The one given to ``@cached``, or else the application's (what the + ``AppCache`` dependency returns), resolved on every call so a manager + registered with ``CacheManagerProxy.set()`` at startup is used. + """ + if self._manager is not None: + return self._manager + return get_app_cache() + + @overload + def __get__( + self, instance: None, owner: type | None = None + ) -> "CachedFunction[P, R]": ... + @overload + def __get__( + self, instance: object, owner: type | None = None + ) -> "_BoundCachedFunction[R]": ... + def __get__( + self, instance: object | None, owner: type | None = None + ) -> "CachedFunction[P, R] | _BoundCachedFunction[R]": + """Bind ``instance`` as the first argument, so methods work. + + ``obj.load(1)``, ``obj.load.cache_key(1)`` and ``obj.load.invalidate(1)`` + all pass ``obj`` as ``self``; a default key then needs ``key=`` since + ``obj`` is not JSON-serializable. + """ + if instance is None: + return self + return _BoundCachedFunction(self, instance) + + def cache_key(self, *args: P.args, **kwargs: P.kwargs) -> str: + """The key a call with these arguments reads and writes. + + Without ``key`` it is ``module.qualname:``: + the arguments are bound to the function's signature, defaults + applied, so a value passed by position or by name, or left to its + default, gives the same key. They are hashed as JSON, so they must be + JSON-serializable; ``@cached(key=...)`` names the key for anything + else (a method's ``self``, a model). A ``str`` ``key`` is a + ``str.format`` template over the bound arguments; a callable is called + with them. The manager's ``key_prefix`` is not part of it. + + Raises: + CacheXError: If an argument is not JSON-serializable and no + ``key`` was given, or ``key`` does not give a ``str`` + TypeError: If the arguments do not fit the function's signature + """ + bound = self._signature.bind(*args, **kwargs) + bound.apply_defaults() + if self._key is None: + return self._default_key(bound.arguments) + if isinstance(self._key, str): + return self._key.format(**bound.arguments) + key = self._key(*args, **kwargs) + if not isinstance(key, str): + msg = f"key must return a str, got {type(key).__name__}" # type: ignore[unreachable] + raise CacheXError(msg) + return key + + def _default_key(self, arguments: dict[str, Any]) -> str: + try: + canonical = json.dumps(arguments, sort_keys=True, separators=(",", ":")) + except (TypeError, ValueError) as e: + msg = ( + f"@cached cannot build a key for {self._name}: an argument is not " + f"JSON-serializable ({e}). Pass key= to name the key yourself." + ) + raise CacheXError(msg) from e + digest = hashlib.sha256(canonical.encode("utf-8")).hexdigest() + return f"{self._name}{_KEY_SEPARATOR}{digest}" + + async def __call__(self, *args: P.args, **kwargs: P.kwargs) -> R: + """Return the cached value, running the function on a miss.""" + key = self.cache_key(*args, **kwargs) + + def factory() -> Awaitable[R] | R: + return self.__wrapped__(*args, **kwargs) + + value: R = await self.manager.get_or_set( + key, factory, ttl=self._ttl, lock=self._lock + ) + return value + + async def invalidate(self, *args: P.args, **kwargs: P.kwargs) -> bool: + """Remove the value cached for a call with these arguments. + + Returns: + Whether a value was cached for them + """ + removed = await self.manager.delete(self.cache_key(*args, **kwargs)) + logger.debug("@cached INVALIDATE; function=%s removed=%s", self._name, removed) + return removed + + +class _BoundCachedFunction(Generic[R]): + """A ``CachedFunction`` looked up on an instance (see ``__get__``).""" + + def __init__(self, function: "CachedFunction[Any, R]", instance: object) -> None: + """Bind ``instance`` to ``function``.""" + self.__wrapped__ = function + self.__self__ = instance + + def cache_key(self, *args: Any, **kwargs: Any) -> str: + """``CachedFunction.cache_key`` with the instance as ``self``.""" + return self.__wrapped__.cache_key(self.__self__, *args, **kwargs) + + async def __call__(self, *args: Any, **kwargs: Any) -> R: + """``CachedFunction.__call__`` with the instance as ``self``.""" + return await self.__wrapped__(self.__self__, *args, **kwargs) + + async def invalidate(self, *args: Any, **kwargs: Any) -> bool: + """``CachedFunction.invalidate`` with the instance as ``self``.""" + return await self.__wrapped__.invalidate(self.__self__, *args, **kwargs) + + +class _Decorator(Protocol): + @overload + def __call__(self, func: Callable[P, Awaitable[R]]) -> CachedFunction[P, R]: ... + @overload + def __call__(self, func: Callable[P, R]) -> CachedFunction[P, R]: ... + + +def cached( + ttl: int | timedelta | None = None, + *, + key: str | Callable[..., str] | None = None, + manager: CacheManager | None = None, + lock: bool | None = None, +) -> _Decorator: + """Cache what a plain function returns, keyed on its arguments. + + For a function that is not a route: a loader, a slow computation, a call + to another service. The decorated function is always ``await``-ed, even + if it was ``def``, and returns the cached value on a hit; on a miss it + runs through ``CacheManager.get_or_set()``, so concurrent misses of one + key run it once (the manager's stampede protection), and its result is + stored. A sync function runs on the event loop, as a ``get_or_set`` + factory does; wrap blocking I/O in ``run_in_threadpool`` yourself. + + The value goes through the manager's JSON round-trip, so the function's + result must be JSON-serializable and comes back as JSON gives it (a tuple + as a list, integer dict keys as strings), on the first call too. + + The decorated function has ``cache_key(*args, **kwargs)`` for the key a + call uses and ``await invalidate(*args, **kwargs)`` to drop its value. + + Args: + ttl: How long a value is cached, in seconds or as a ``timedelta``. + ``None`` uses the manager's ``default_ttl``. + key: How to name the key, without the manager's prefix. ``None`` + hashes the arguments: ``module.qualname:``, which needs + JSON-serializable arguments (a method's ``self`` is not). A + ``str`` is a ``str.format`` template over the arguments by name, + defaults applied (``"user:{user_id}"``); a plain string is a + fixed key. A callable gets the call's arguments and returns the + key (``lambda self, user_id: f"user:{user_id}"``). + manager: The ``CacheManager`` to store through. ``None`` uses the + application's, what the ``AppCache`` dependency returns, + resolved on each call. + lock: Whether concurrent misses take the manager's distributed lock. + ``None`` uses the manager's ``lock`` setting. + + Returns: + A decorator turning the function into a ``CachedFunction`` + + Raises: + TypeError: When the decorator is applied, if ``ttl`` is not an + ``int``, a ``timedelta`` or ``None``, or ``lock`` is not a + ``bool`` or ``None`` + ValueError: When the decorator is applied, if ``ttl`` is zero, + negative, larger than ``MAX_TTL`` or not a whole number of + seconds + CacheXError: When the decorator is applied, if ``key`` is not a + ``str``, a callable or ``None``; and when the function is + called, if the key cannot be built (see ``cache_key``) + """ + ttl_seconds = validate_ttl(ttl) + if lock is not None and not isinstance(lock, bool): + msg = f"lock must be a bool or None, got {type(lock).__name__}" # type: ignore[unreachable] + raise TypeError(msg) + if key is not None and not isinstance(key, str) and not callable(key): + msg = f"key must be a str, a callable or None, got {type(key).__name__}" # type: ignore[unreachable] + raise CacheXError(msg) + + def decorator(func: Callable[P, Any]) -> CachedFunction[P, Any]: + return CachedFunction( + func, ttl=ttl_seconds, key=key, manager=manager, lock=lock + ) + + return decorator + + +__all__ = ["CachedFunction", "cached"] diff --git a/i18n/zh-TW/docs/APP_CACHE.md b/i18n/zh-TW/docs/APP_CACHE.md index 2de69d1..391173c 100644 --- a/i18n/zh-TW/docs/APP_CACHE.md +++ b/i18n/zh-TW/docs/APP_CACHE.md @@ -42,6 +42,21 @@ await manager.clear_pattern("user:*") # 比對 "myapp:user:*" 完整可執行範例(英文):[`examples/app_cache.py`](https://github.com/allen0099/FastAPI-CacheX/blob/master/examples/app_cache.py)。 +## 快取一個函式 {#caching-a-function} + +`@cached` 對一般函式做的事與 `get_or_set()` 相同,以函式的引數作為鍵,因此載入資料或呼叫其他服務的函式只要寫一次,在任何地方呼叫都會被快取: + + +```python +--8<-- "examples/app_cache.py:cached" +``` + + +- 被裝飾的函式一律要 `await`,即使它原本是 `def`(同步函式會在事件迴圈上執行,和 `get_or_set()` 的 factory 一樣)。它的結果會經過 manager 的 [JSON 往返](#json-round-trip),因此必須可 JSON 序列化,拿回來的也是 JSON 解碼後的形式,第一次呼叫也一樣。 +- `ttl` 接受秒數或 `timedelta`;未設定時套用 manager 的 `default_ttl`。`manager=` 指定要透過哪個 `CacheManager` 儲存;未設定時使用應用程式的那一個(也就是 `AppCache` 回傳的),並在每次呼叫時解析,因此在啟動時以 `CacheManagerProxy.set()` 註冊的 manager 也會被採用。`lock=False` 讓這個函式略過 manager 的 [cache stampede 保護](#stampede-protection)。 +- 未設定 `key=` 時,鍵是 `module.qualname:` 加上引數的 SHA-256;引數會先依函式簽章綁定並套用預設值,因此 `load(1)`、`load(user_id=1)` 與 `load(1, locale="en")` 共用同一個項目。引數以 JSON 雜湊,因此必須可 JSON 序列化;帶有無法序列化之引數的呼叫會拋出 `CacheXError`。`key="user:{user_id}"` 是以引數名稱填入的 `str.format` 樣板(不含欄位的字串就是固定的鍵),`key=lambda self, user_id: f"user:{user_id}"` 則會以引數呼叫:方法的 `self` 無法雜湊,或引數是模型時,請用其中一種。manager 的 `key_prefix` 會加在前面,和它儲存的每個鍵一樣。 +- `fn.cache_key(*args, **kwargs)` 是該次呼叫使用的鍵(不含前綴),`await fn.invalidate(*args, **kwargs)` 則丟棄它的值,並回傳原本是否有快取。在方法上,`obj.load.invalidate(1)` 會和呼叫時一樣綁定 `obj`。 + ## 行為 {#behavior} - `get()` 在快取未命中時回傳 `None`(或你提供的 `default=`),遇到不存在或損毀的項目也絕不會拋出例外。 diff --git a/tests/test_cached.py b/tests/test_cached.py new file mode 100644 index 0000000..858edf3 --- /dev/null +++ b/tests/test_cached.py @@ -0,0 +1,335 @@ +"""``@cached`` caches a plain function's result through ``CacheManager`` (#248).""" + +import asyncio +import hashlib +import json +from collections.abc import AsyncIterator +from datetime import timedelta +from typing import Any + +import pytest +import pytest_asyncio + +from fastapi_cachex import BackendProxy +from fastapi_cachex import CachedFunction +from fastapi_cachex import CacheManager +from fastapi_cachex import CacheManagerProxy +from fastapi_cachex import cached +from fastapi_cachex.backends.memory import MemoryBackend +from fastapi_cachex.exceptions import CacheXError +from tests.conftest import Clock + + +@pytest_asyncio.fixture +async def manager(memory_backend: MemoryBackend) -> AsyncIterator[CacheManager]: + BackendProxy.set(memory_backend) + CacheManagerProxy.set(None) + yield CacheManager(backend=memory_backend, key_prefix="t:") + CacheManagerProxy.set(None) + + +def _digest(**arguments: Any) -> str: + canonical = json.dumps(arguments, sort_keys=True, separators=(",", ":")) + return hashlib.sha256(canonical.encode()).hexdigest() + + +async def test_an_async_function_runs_once_per_arguments( + manager: CacheManager, +) -> None: + calls: list[int] = [] + + @cached(ttl=60, manager=manager) + async def load(user_id: int) -> dict[str, int]: + calls.append(user_id) + return {"id": user_id} + + assert await load(1) == {"id": 1} + assert await load(1) == {"id": 1} + assert await load(2) == {"id": 2} + assert calls == [1, 2] + + +async def test_a_sync_function_is_awaited_too(manager: CacheManager) -> None: + calls: list[int] = [] + + @cached(ttl=60, manager=manager) + def square(n: int) -> int: + calls.append(n) + return n * n + + assert await square(3) == 9 + assert await square(3) == 9 + assert calls == [3] + + +async def test_the_default_key_hashes_the_bound_arguments( + manager: CacheManager, +) -> None: + @cached(ttl=60, manager=manager) + async def load(user_id: int, *, locale: str = "en") -> str: + return f"{user_id}:{locale}" + + expected = f"{__name__}.{load.__qualname__}:{_digest(user_id=1, locale='en')}" + assert load.cache_key(1) == expected + # By position or by name, defaults applied: the same key. + assert load.cache_key(user_id=1) == expected + assert load.cache_key(1, locale="en") == expected + assert load.cache_key(1, locale="fr") != expected + + await load(1) + assert await manager.has(expected) + assert await manager.get(expected) == "1:en" + + +async def test_a_non_json_argument_needs_a_key(manager: CacheManager) -> None: + @cached(ttl=60, manager=manager) + async def load(when: object) -> int: + return 1 + + with pytest.raises(CacheXError, match=r"not JSON-serializable.*Pass key="): + load.cache_key(object()) + with pytest.raises(CacheXError, match="not JSON-serializable"): + await load(object()) + + +async def test_arguments_must_fit_the_signature(manager: CacheManager) -> None: + @cached(ttl=60, manager=manager) + async def load(user_id: int) -> int: + return user_id + + with pytest.raises(TypeError): + await load(1, 2) # type: ignore[call-arg] + with pytest.raises(TypeError): + load.cache_key(other=1) # type: ignore[call-arg] + + +async def test_a_str_key_is_a_template_over_the_arguments( + manager: CacheManager, +) -> None: + @cached(ttl=60, manager=manager, key="user:{user_id}:{locale}") + async def load(user_id: int, locale: str = "en") -> str: + return f"{user_id}:{locale}" + + assert load.cache_key(42) == "user:42:en" + assert load.cache_key(42, locale="fr") == "user:42:fr" + await load(42) + assert await manager.get("user:42:en") == "42:en" + + +async def test_a_plain_str_key_is_fixed(manager: CacheManager) -> None: + @cached(ttl=60, manager=manager, key="rates") + async def rates() -> dict[str, float]: + return {"EUR": 0.9} + + assert rates.cache_key() == "rates" + assert await rates() == {"EUR": 0.9} + assert await manager.get("rates") == {"EUR": 0.9} + + +async def test_a_callable_key_gets_the_arguments(manager: CacheManager) -> None: + @cached(ttl=60, manager=manager, key=lambda item, *, verbose=False: f"i:{item}") + async def load(item: int, *, verbose: bool = False) -> int: + return item + + assert load.cache_key(7, verbose=True) == "i:7" + assert await load(7) == 7 + assert await manager.get("i:7") == 7 + + +async def test_a_callable_key_must_return_a_str(manager: CacheManager) -> None: + @cached(ttl=60, manager=manager, key=lambda item: item) + async def load(item: int) -> int: + return item + + with pytest.raises(CacheXError, match="key must return a str, got int"): + await load(7) + + +async def test_invalidate_drops_the_value_for_those_arguments( + manager: CacheManager, +) -> None: + calls: list[int] = [] + + @cached(ttl=60, manager=manager) + async def load(user_id: int) -> int: + calls.append(user_id) + return user_id + + await load(1) + await load(2) + assert await load.invalidate(1) is True + assert await load.invalidate(1) is False + await load(1) + await load(2) + assert calls == [1, 2, 1] + + +async def test_ttl_expires_the_value(manager: CacheManager, clock: Clock) -> None: + calls: list[int] = [] + + @cached(ttl=timedelta(minutes=1), manager=manager) + async def load(n: int) -> int: + calls.append(n) + return n + + await load(1) + clock.advance(59) + await load(1) + clock.advance(2) + await load(1) + assert calls == [1, 1] + + +async def test_without_ttl_the_managers_default_applies( + memory_backend: MemoryBackend, clock: Clock +) -> None: + manager = CacheManager(backend=memory_backend, default_ttl=10) + calls: list[int] = [] + + @cached(manager=manager) + async def load(n: int) -> int: + calls.append(n) + return n + + await load(1) + clock.advance(11) + await load(1) + assert calls == [1, 1] + + +async def test_the_value_comes_back_as_json_gives_it(manager: CacheManager) -> None: + @cached(ttl=60, manager=manager) + async def pair() -> tuple[int, int]: + return (1, 2) + + assert await pair() == [1, 2] # type: ignore[comparison-overlap] + assert await pair() == [1, 2] # type: ignore[comparison-overlap] + + +async def test_a_non_json_result_is_not_stored(manager: CacheManager) -> None: + @cached(ttl=60, manager=manager) + async def when() -> object: + return object() + + with pytest.raises(TypeError): + await when() + assert await manager.has(when.cache_key()) is False + + +async def test_without_a_manager_the_app_cache_is_used( + memory_backend: MemoryBackend, +) -> None: + BackendProxy.set(memory_backend) + CacheManagerProxy.set(None) + + @cached(ttl=60) + async def load(n: int) -> int: + return n + + # Registered later, at startup say, and still picked up: the manager is + # resolved on each call. + app_manager = CacheManager(backend=memory_backend, key_prefix="app:") + CacheManagerProxy.set(app_manager) + try: + assert load.manager is app_manager + await load(1) + assert await app_manager.has(load.cache_key(1)) + finally: + CacheManagerProxy.set(None) + + +async def test_concurrent_misses_run_the_function_once( + manager: CacheManager, +) -> None: + calls: list[int] = [] + started = asyncio.Event() + + @cached(ttl=60, manager=manager) + async def load(n: int) -> int: + calls.append(n) + started.set() + await asyncio.sleep(0.05) + return n + + results = await asyncio.gather(*(load(1) for _ in range(5))) + assert results == [1] * 5 + assert calls == [1] + + +async def test_lock_false_skips_the_lock(memory_backend: MemoryBackend) -> None: + manager = CacheManager(backend=memory_backend, key_prefix="t:") + + @cached(ttl=60, manager=manager, lock=False) + async def load(n: int) -> int: + assert await memory_backend.get("lock:" + "t:" + load.cache_key(n)) is None + return n + + assert await load(1) == 1 + + +async def test_methods_are_bound(manager: CacheManager) -> None: + class Repo: + def __init__(self) -> None: + self.calls: list[int] = [] + + @cached(ttl=60, manager=manager, key=lambda self, item_id: f"item:{item_id}") + async def load(self, item_id: int) -> int: + self.calls.append(item_id) + return item_id + + repo = Repo() + assert repo.load.cache_key(5) == "item:5" + assert await repo.load(5) == 5 + assert await repo.load(5) == 5 + assert repo.calls == [5] + assert await repo.load.invalidate(5) is True + assert await repo.load(5) == 5 + assert repo.calls == [5, 5] + assert isinstance(Repo.load, CachedFunction) + assert Repo.load.cache_key(repo, 5) == "item:5" + + +async def test_a_method_without_a_key_names_the_problem( + manager: CacheManager, +) -> None: + class Repo: + @cached(ttl=60, manager=manager) + async def load(self, item_id: int) -> int: + return item_id + + with pytest.raises(CacheXError, match="Pass key="): + await Repo().load(1) + + +def test_the_wrapper_keeps_the_functions_identity(manager: CacheManager) -> None: + @cached(ttl=60, manager=manager) + async def load(n: int) -> int: + """Load it.""" + return n + + assert load.__name__ == "load" + assert load.__doc__ == "Load it." + assert load.__wrapped__.__name__ == "load" + assert isinstance(load, CachedFunction) + + +@pytest.mark.parametrize( + ("kwargs", "error", "match"), + [ + pytest.param({"ttl": 0}, ValueError, "ttl must be a positive", id="ttl-0"), + pytest.param({"ttl": 1.5}, TypeError, "got float", id="ttl-float"), + pytest.param( + {"ttl": timedelta(seconds=1.5)}, + ValueError, + "whole number of seconds", + id="ttl-fraction", + ), + pytest.param({"lock": 1}, TypeError, "lock must be a bool", id="lock"), + pytest.param({"key": 3}, CacheXError, "key must be a str", id="key"), + ], +) +def test_bad_arguments_are_rejected_when_applied( + kwargs: dict[str, Any], error: type[Exception], match: str +) -> None: + with pytest.raises(error, match=match): + cached(**kwargs) diff --git a/tests/test_examples.py b/tests/test_examples.py index 7292cab..9b76b85 100644 --- a/tests/test_examples.py +++ b/tests/test_examples.py @@ -178,6 +178,14 @@ def test_app_cache() -> None: other = client.post("/orders", headers={"Idempotency-Key": "def"}) assert other.status_code == 201 + calls = example.upstream_calls["rates"] + assert client.get("/rates/eur").json() == {"eur": 0.92} + assert client.get("/rates/eur").json() == {"eur": 0.92} + assert example.upstream_calls["rates"] == calls + 1 + assert client.delete("/rates/eur").json() == {"deleted": True} + client.get("/rates/eur") + assert example.upstream_calls["rates"] == calls + 2 + def _session_cookie(client: TestClient) -> str: token = client.cookies.get("session")