From 866117320f4752e0f3ab9a7ebd23f3d4947742b7 Mon Sep 17 00:00:00 2001 From: allen0099 Date: Sun, 27 Sep 2026 13:39:04 +0000 Subject: [PATCH] feat(cache): add vary to key cached responses on request headers @cache(vary=["Accept-Language"]) appends a name=value component per listed header to the key, after whatever the key builder returns, and adds the names to the Vary header of every GET response, including 304s and responses that bypass or are not stored in the backend. invalidate() takes the same list. The Vary helper moves from the session middleware to fastapi_cachex/headers.py so both use it. --- changelog.d/268.added.md | 12 ++ docs/CACHE_FLOW.md | 5 +- docs/HTTP_CACHING.md | 71 ++++++- fastapi_cachex/cache.py | 127 +++++++++++-- fastapi_cachex/headers.py | 27 +++ fastapi_cachex/session/middleware.py | 30 +-- i18n/zh-TW/docs/CACHE_FLOW.md | 2 +- i18n/zh-TW/docs/HTTP_CACHING.md | 43 ++++- tests/test_cache_vary.py | 273 +++++++++++++++++++++++++++ 9 files changed, 545 insertions(+), 45 deletions(-) create mode 100644 changelog.d/268.added.md create mode 100644 fastapi_cachex/headers.py create mode 100644 tests/test_cache_vary.py diff --git a/changelog.d/268.added.md b/changelog.d/268.added.md new file mode 100644 index 0000000..f92a597 --- /dev/null +++ b/changelog.d/268.added.md @@ -0,0 +1,12 @@ +**`@cache(vary=[...])` caches one entry per value of the listed request +headers.** A route whose response depends on `Accept-Language` or `Accept` no +longer serves the first cached variant to everyone: each listed header adds a +`name=value` component (name lower-cased, value trimmed, missing as empty) to +the key, after whatever the `key_builder` returns, and the names are added to +the `Vary` header of every GET response, 200 or 304, stored or not, without +repeating names already there and leaving `Vary: *` alone. `vary` is checked +when the decorator is applied; a bare string such as `vary="Accept"` is +rejected. `invalidate()` takes the same `vary` list, and `clear_path()` clears +every variant of a path. Routes without `vary` keep their keys. The header +values are client-controlled, so HTTP_CACHING.md shows how to normalise them +with a `key_builder` and `build_cache_key` instead. diff --git a/docs/CACHE_FLOW.md b/docs/CACHE_FLOW.md index 09e6b1e..dc23c9c 100644 --- a/docs/CACHE_FLOW.md +++ b/docs/CACHE_FLOW.md @@ -86,7 +86,10 @@ The monitoring routes decode them again for display. A custom `key_builder` can add components after the query string with `build_cache_key(request, *components)`; they are encoded the same way, and `clear_path()` still matches the path (see "Adding components to the key" in -[HTTP caching](HTTP_CACHING.md#adding-components-to-the-key)). +[HTTP caching](HTTP_CACHING.md#adding-components-to-the-key)). `@cache(vary=[...])` +appends one `name=value` component per listed request header after whatever +the key builder returns, and adds the names to the response's `Vary` header +(see [Varying on request headers](HTTP_CACHING.md#varying-on-request-headers)). Query parameters are joined in the order the request sent them (`str(request.query_params)`) and are **not sorted**, so `?page=1&limit=10` and diff --git a/docs/HTTP_CACHING.md b/docs/HTTP_CACHING.md index 19256c3..0d9884e 100644 --- a/docs/HTTP_CACHING.md +++ b/docs/HTTP_CACHING.md @@ -228,6 +228,71 @@ default) in front of every key, so other applications can share the server; namespace instead of this `|||`-separated format, since its keys aren't tied to HTTP requests. +### Varying on request headers + +The key includes no request header, so a route whose response depends on, say, +`Accept-Language` would serve the first cached language to everyone. List such +headers in `vary`: + +```python +@app.get("/greeting") +@cache(ttl=300, vary=["Accept-Language"]) +async def greeting(request: Request): + return {"text": translate("hello", request.headers.get("accept-language"))} +``` + +Each listed header adds a `name=value` component to the key: the name +lower-cased, the value trimmed (repeated header lines joined with `,`), and a +missing header treated as an empty one. The components are escaped like the +rest of the key and come after whatever the `key_builder` returns, so `vary` +and a custom key builder compose: +`GET|||example.com|||/greeting|||||||tenant-1|||accept-language=de` for +`key_builder` returning `build_cache_key(request, "tenant-1")`. Routes without +`vary` keep their keys. + +The names are also added to the `Vary` header of every response to a GET +request on the route, on a 200 or a 304, served from the backend or not +(`private`, `no_store`, a bypassed `Authorization` request or a response that +is not stored), so a shared cache in front of the app keys on them too. A name +the response already lists (in any case) is not repeated, and a response with +`Vary: *` is left alone. + +`vary` must be a list (or tuple) of header field names. A bare string such as +`vary="Accept"` is rejected when the decorator is applied, as are empty names, +`*` and anything that is not a valid field name. + +> [!WARNING] +> **Every distinct header value is its own entry, and the values come from the +> client.** `Accept-Language: de`, `de-DE`, `de-DE,de;q=0.9` and every other +> spelling are separate keys, and a client can send a new one on every request +> to fill the backend. Each listed header multiplies the number of entries. +> When only a few values matter, normalise in a `key_builder` instead, and add +> the header to `Vary` yourself: +> +> ```python +> SUPPORTED = ("en", "de", "fr") +> +> +> def locale_key(request: Request) -> str: +> wanted = request.headers.get("accept-language", "") +> locale = next( +> (tag for tag in SUPPORTED if wanted.lower().startswith(tag)), "en" +> ) +> return build_cache_key(request, locale) +> +> +> @app.get("/greeting") +> @cache(ttl=300, key_builder=locale_key) +> async def greeting(request: Request, response: Response): +> response.headers["Vary"] = "Accept-Language" +> return {"text": translate("hello", locale_of(request))} +> ``` + +`clear_path()` clears every variant of a path (see +[Adding components to the key](#adding-components-to-the-key)). +`invalidate()` takes the same `vary` list and deletes only the variant the +request it is given selects. + ### Authenticated endpoints > [!WARNING] @@ -375,12 +440,14 @@ async def update_item(item_id: int, request: Request): return {"invalidated": await invalidate(StarletteRequest(scope))} ``` -`invalidate(request, key_builder=None)` returns `True` when an entry existed and +`invalidate(request, key_builder=None, vary=None)` returns `True` when an entry existed and was removed, `False` otherwise, including when no backend is configured. An error from the backend itself is raised to the caller (see [When the backend fails](#when-the-backend-fails)). The request you hand it must produce the cached route's key: same method, host, path and query string. If the cached route uses a custom -`key_builder`, pass the same one here, or the key will not match. +`key_builder` or `vary`, pass the same here, or the key will not match; with +`vary` only the variant selected by the request's own header values is +deleted, and `clear_path()` removes all of them. ## Monitoring routes diff --git a/fastapi_cachex/cache.py b/fastapi_cachex/cache.py index 167ff0a..dae6c90 100644 --- a/fastapi_cachex/cache.py +++ b/fastapi_cachex/cache.py @@ -6,6 +6,7 @@ from collections.abc import Awaitable from collections.abc import Callable from collections.abc import Mapping +from collections.abc import Sequence from functools import update_wrapper from functools import wraps from inspect import Parameter @@ -35,6 +36,7 @@ from .exceptions import BackendNotFoundError from .exceptions import CacheXError from .exceptions import RequestNotFoundError +from .headers import add_vary from .proxy import BackendProxy from .proxy import get_backend_or_fallback from .types import CACHE_KEY_SEPARATOR @@ -92,12 +94,24 @@ def per_user_key(request: Request) -> str: rejected too), e.g. ``None`` from a missing user ID, which would otherwise put every such caller under one ``"None"`` key. """ - parts = [ - request.method, - escape_key_component(request.headers.get("host", "unknown")), - escape_key_component(request.url.path), - str(request.query_params), - ] + key = _append_key_components( + CACHE_KEY_SEPARATOR.join( + [ + request.method, + escape_key_component(request.headers.get("host", "unknown")), + escape_key_component(request.url.path), + str(request.query_params), + ] + ), + components, + ) + logger.debug("Built cache key: %s", key) + return key + + +def _append_key_components(key: str, components: Sequence[str | int]) -> str: + """Append each component to ``key``, escaped, after another separator.""" + parts = [key] for component in components: if isinstance(component, bool) or not isinstance(component, (str, int)): msg = ( @@ -106,9 +120,54 @@ def per_user_key(request: Request) -> str: ) raise TypeError(msg) parts.append(escape_key_component(str(component))) - key = CACHE_KEY_SEPARATOR.join(parts) - logger.debug("Built cache key: %s", key) - return key + return CACHE_KEY_SEPARATOR.join(parts) + + +# RFC 9110 §5.1: a field name is a token. +_FIELD_NAME_CHARS = frozenset( + "!#$%&'*+-.^_`|~0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ" +) + + +def _validate_vary(vary: Sequence[str] | None) -> list[str]: + """Check ``@cache(vary=...)`` and return the names, first spelling of each. + + Raises: + CacheXError: If ``vary`` is a single string instead of a sequence of + names, or a name is not a non-empty header field name, or is ``*``. + """ + if vary is None: + return [] + if isinstance(vary, (str, bytes)) or not isinstance(vary, Sequence): + msg = ( + "vary must be a list of header names, e.g. vary=['Accept-Language'], " + f"got {type(vary).__name__}" + ) + raise CacheXError(msg) + names: dict[str, str] = {} + for name in vary: + if not isinstance(name, str) or not name or not set(name) <= _FIELD_NAME_CHARS: + msg = f"vary entries must be header field names, got {name!r}" + raise CacheXError(msg) + if name == "*": + msg = "vary cannot contain '*': the key can only vary on named headers" + raise CacheXError(msg) + names.setdefault(name.lower(), name) + return list(names.values()) + + +def _vary_components(request: Request, names: Sequence[str]) -> list[str]: + """The key components for ``@cache(vary=names)``: ``name=value`` each. + + The name is lower-cased and the value trimmed; repeated header lines are + joined with ``,`` as RFC 9110 §5.3 allows, and a missing header gives an + empty value, the same as an empty one. + """ + return [ + f"{name.lower()}=" + + ",".join(value.strip() for value in request.headers.getlist(name)) + for name in names + ] def default_key_builder(request: Request) -> str: @@ -131,6 +190,7 @@ def default_key_builder(request: Request) -> str: async def invalidate( request: Request, key_builder: CacheKeyBuilder | None = None, + vary: Sequence[str] | None = None, ) -> bool: """Invalidate the cache entry a ``@cache``-decorated route would use. @@ -145,12 +205,21 @@ async def invalidate( it via ``request.app.url_path_for(...)`` for a GET route). key_builder: Custom key builder used by the target route's ``@cache`` decorator, if any. If None, uses ``default_key_builder``. + vary: The target route's ``vary`` names, if any. Only the variant + selected by ``request``'s own values for those headers is + deleted; ``clear_path()`` clears every variant of a path. Returns: True if a cache entry existed and was deleted, False otherwise. + + Raises: + CacheXError: If ``vary`` is not a list of header names. """ builder = key_builder or default_key_builder - cache_key = builder(request) + vary_names = _validate_vary(vary) + cache_key = _append_key_components( + builder(request), _vary_components(request, vary_names) + ) try: cache_backend = BackendProxy.get() @@ -591,6 +660,7 @@ def cache( key_builder: CacheKeyBuilder | None = None, fail_open: bool = True, cache_authorized: bool = False, + vary: Sequence[str] | None = None, ) -> Callable[[HandlerCallable], AsyncResponseCallable]: """Cache decorator for FastAPI route handlers. @@ -650,6 +720,15 @@ def cache( ``key_builder`` puts the verified caller's identity into the key; with the default key builder one user's response would be served to the next. + vary: Request header names the response depends on, e.g. + ``["Accept-Language"]``. Each header's value (trimmed; empty when + missing) is appended to the key, after whatever ``key_builder`` + returns, as a ``name=value`` component, so every distinct value + gets its own entry. The names are also added to the ``Vary`` + header of every response to a GET request, unless it already + lists them or ``*``. Values are client-controlled: each listed + header multiplies the number of entries, so normalise them in a + ``key_builder`` when only a few values matter. Returns: Decorator function that wraps route handlers with caching logic @@ -657,8 +736,9 @@ def cache( Raises: CacheXError: When the decorator is applied, if ``stale`` and ``stale_ttl`` are not given together, if ``public`` and - ``private`` are both set, or if ``ttl`` is not an ``int``, is - negative or is larger than ``MAX_TTL``. + ``private`` are both set, if ``ttl`` is not an ``int``, is + negative or is larger than ``MAX_TTL``, or if ``vary`` is not a + list of header field names (a single string is rejected). """ def decorator(func: HandlerCallable) -> AsyncResponseCallable: @@ -683,6 +763,7 @@ def decorator(func: HandlerCallable) -> AsyncResponseCallable: if ttl is not None and ttl > MAX_TTL: msg = f"ttl must be at most {MAX_TTL} seconds" raise CacheXError(msg) + vary_names = _validate_vary(vary) # Analyze the original function's signature sig: Signature = inspect.signature(func) @@ -761,7 +842,7 @@ def decorator(func: HandlerCallable) -> AsyncResponseCallable: bypass_backend = private or not ttl @wraps(func) - async def wrapper(*args: Any, **kwargs: Any) -> Response: + async def serve(*args: Any, **kwargs: Any) -> Response: # Resolve backend on every request to support lifespan-configured backends cache_backend = get_backend_or_fallback() @@ -850,7 +931,9 @@ async def wrapper(*args: Any, **kwargs: Any) -> Response: # Built only here: the branches above never touch the backend, so a # custom key builder would run for nothing. - cache_key = builder(req) + cache_key = _append_key_components( + builder(req), _vary_components(req, vary_names) + ) try: cached_data = await cache_backend.get(cache_key) @@ -989,6 +1072,22 @@ async def wrapper(*args: Any, **kwargs: Any) -> Response: current_response, cache_control, private_cache_control ) + wrapper: AsyncResponseCallable = serve + if vary_names: + + @wraps(func) + async def with_vary(*args: Any, **kwargs: Any) -> Response: + # Read before `serve` pops an injected request parameter. + req: Request | None = kwargs.get(request_name) + response = await serve(*args, **kwargs) + if req is not None and req.method == "GET": + # Every GET answer, served from the backend or not: a + # shared cache downstream keys on these headers too. + add_vary(response.headers, vary_names) + return response + + wrapper = with_vary + # Update the wrapper with the new signature update_wrapper(wrapper, func) wrapper.__signature__ = sig # type: ignore[attr-defined] diff --git a/fastapi_cachex/headers.py b/fastapi_cachex/headers.py new file mode 100644 index 0000000..33fe63a --- /dev/null +++ b/fastapi_cachex/headers.py @@ -0,0 +1,27 @@ +"""Response header helpers shared by ``@cache`` and the session middleware.""" + +from starlette.datastructures import MutableHeaders + + +def add_vary(headers: MutableHeaders, names: list[str]) -> None: + """Add header names to ``Vary``, keeping the values already there. + + Starlette's ``add_vary_header`` appends unconditionally, so names the + response already varies on (compared case-insensitively) are skipped, as + is everything when it already varies on ``*``. + + Args: + headers: Mutable response headers to write to + names: Request header names the response depends on + """ + present = { + value.strip().lower() + for line in headers.getlist("vary") + for value in line.split(",") + } + if "*" in present: + return + for name in names: + if name.lower() not in present: + headers.add_vary_header(name) + present.add(name.lower()) diff --git a/fastapi_cachex/session/middleware.py b/fastapi_cachex/session/middleware.py index 9942e72..f521ea4 100644 --- a/fastapi_cachex/session/middleware.py +++ b/fastapi_cachex/session/middleware.py @@ -18,6 +18,8 @@ from starlette.types import Scope from starlette.types import Send +from fastapi_cachex.headers import add_vary + from .config import SessionConfig from .exceptions import SessionError from .manager import SessionManager @@ -146,30 +148,6 @@ def _read_header_token( return None, consulted -def _add_vary(headers: MutableHeaders, names: list[str]) -> None: - """Add header names to ``Vary``, keeping the values already there. - - Starlette's ``add_vary_header`` appends unconditionally, so names the - response already varies on (compared case-insensitively) are skipped, as - is everything when it already varies on ``*``. - - Args: - headers: Mutable response headers to write to - names: Request header names the response depends on - """ - present = { - value.strip().lower() - for line in headers.getlist("vary") - for value in line.split(",") - } - if "*" in present: - return - for name in names: - if name.lower() not in present: - headers.add_vary_header(name) - present.add(name.lower()) - - def _forbid_storing(headers: MutableHeaders) -> None: """Keep a response that carries a session token out of every cache. @@ -298,7 +276,7 @@ async def dispatch( if response_token is not None: response.headers[self.config.header_name] = response_token - _add_vary(response.headers, _read_header_token(request, self.config)[1]) + add_vary(response.headers, _read_header_token(request, self.config)[1]) _forbid_storing(response.headers) return response @@ -459,7 +437,7 @@ async def send_wrapper(message: Message) -> None: ) if session.accessed or sent_token: - _add_vary(headers, vary_on) + add_vary(headers, vary_on) if sent_token: _forbid_storing(headers) diff --git a/i18n/zh-TW/docs/CACHE_FLOW.md b/i18n/zh-TW/docs/CACHE_FLOW.md index 6c3ee23..107356e 100644 --- a/i18n/zh-TW/docs/CACHE_FLOW.md +++ b/i18n/zh-TW/docs/CACHE_FLOW.md @@ -74,7 +74,7 @@ cache_key = CACHE_KEY_SEPARATOR.join( host 與路徑會先經過百分比編碼:`|` 變成 `%7C`,`%` 變成 `%25`(`fastapi_cachex/types.py` 中的 `escape_key_component`)。兩者都來自用戶端,其中若出現未編碼的 `|||`,各段就會錯位,使某個請求的快取鍵可能與另一個請求相同。查詢字串本來就經過 URL 編碼。監控路由顯示時會再解碼。 -自訂的 `key_builder` 可以用 `build_cache_key(request, *components)` 在查詢字串之後加入其他段;這些段以同樣方式編碼,`clear_path()` 也仍會比對路徑(見 [HTTP 快取](HTTP_CACHING.md#adding-components-to-the-key)中的「在鍵中加入其他段」)。 +自訂的 `key_builder` 可以用 `build_cache_key(request, *components)` 在查詢字串之後加入其他段;這些段以同樣方式編碼,`clear_path()` 也仍會比對路徑(見 [HTTP 快取](HTTP_CACHING.md#adding-components-to-the-key)中的「在鍵中加入其他段」)。`@cache(vary=[...])` 會在 key builder 回傳的鍵之後,為每個列出的請求標頭附加一個 `name=value` 段,並把這些名稱加入回應的 `Vary` 標頭(見 [HTTP 快取](HTTP_CACHING.md#varying-on-request-headers)中的「依請求標頭區分」)。 查詢參數依請求送出的順序串接(`str(request.query_params)`),**不會排序**,因此 `?page=1&limit=10` 與 `?limit=10&page=1` 是兩個不同的快取項目。若希望兩者視為同一個,請傳入自訂的 `key_builder` 將查詢字串正規化。 diff --git a/i18n/zh-TW/docs/HTTP_CACHING.md b/i18n/zh-TW/docs/HTTP_CACHING.md index 8a656b6..aa32cac 100644 --- a/i18n/zh-TW/docs/HTTP_CACHING.md +++ b/i18n/zh-TW/docs/HTTP_CACHING.md @@ -150,6 +150,47 @@ def per_tenant_key(request: Request) -> str: Redis 與 Memcached 後端還會在每個鍵前面加上自己的前綴(預設為 `fastapi_cachex:`),讓其他應用程式可以共用同一台伺服器;`MemoryBackend` 沒有前綴。`CacheManager`(見[應用層快取](APP_CACHE.md))則使用另一個較簡單、以 `cache:` 為前綴的鍵命名空間,而不是這種以 `|||` 分隔的格式,因為它的鍵與 HTTP 請求無關。 +### 依請求標頭區分 {#varying-on-request-headers} + +快取鍵不包含任何請求標頭,因此回應內容取決於 `Accept-Language` 等標頭的路由,會把第一個快取下來的語言提供給所有人。請把這類標頭列在 `vary` 中: + +```python +@app.get("/greeting") +@cache(ttl=300, vary=["Accept-Language"]) +async def greeting(request: Request): + return {"text": translate("hello", request.headers.get("accept-language"))} +``` + +每個列出的標頭都會在鍵中加入一個 `name=value` 段:名稱轉為小寫,值去除前後空白(重複的標頭行以 `,` 串接),缺少的標頭視同空值。這些段與鍵的其他部分一樣經過編碼,並接在 `key_builder` 回傳的鍵之後,因此 `vary` 可以與自訂的 key builder 一起使用:`key_builder` 回傳 `build_cache_key(request, "tenant-1")` 時,鍵為 `GET|||example.com|||/greeting|||||||tenant-1|||accept-language=de`。沒有設定 `vary` 的路由,鍵維持不變。 + +這些名稱也會加入該路由對 GET 請求的每個回應的 `Vary` 標頭,不論是 200 或 304,也不論是否由後端提供(`private`、`no_store`、繞過後端的 `Authorization` 請求,或未儲存的回應),讓應用程式前方的共用快取也依它們區分。回應已列出的名稱(不分大小寫)不會重複加入,帶有 `Vary: *` 的回應則維持原樣。 + +`vary` 必須是由標頭欄位名稱組成的 list(或 tuple)。套用裝飾器時,會拒絕 `vary="Accept"` 這類單一字串,以及空名稱、`*` 與任何不是有效欄位名稱的值。 + +> [!WARNING] +> **每個不同的標頭值都是一筆獨立的項目,而這些值來自用戶端。** `Accept-Language: de`、`de-DE`、`de-DE,de;q=0.9` 以及其他寫法都是不同的鍵,用戶端可以在每個請求送出新的值來塞滿後端。每多列一個標頭,項目數量就會成倍增加。若只有少數幾個值有意義,請改在 `key_builder` 中正規化,並自行把該標頭加入 `Vary`: +> +> ```python +> SUPPORTED = ("en", "de", "fr") +> +> +> def locale_key(request: Request) -> str: +> wanted = request.headers.get("accept-language", "") +> locale = next( +> (tag for tag in SUPPORTED if wanted.lower().startswith(tag)), "en" +> ) +> return build_cache_key(request, locale) +> +> +> @app.get("/greeting") +> @cache(ttl=300, key_builder=locale_key) +> async def greeting(request: Request, response: Response): +> response.headers["Vary"] = "Accept-Language" +> return {"text": translate("hello", locale_of(request))} +> ``` + +`clear_path()` 會清除某個路徑的所有變體(見[在鍵中加入其他段](#adding-components-to-the-key))。`invalidate()` 接受同樣的 `vary` 清單,只會刪除傳入的請求所選中的那個變體。 + ### 需驗證身分的端點 {#authenticated-endpoints} > [!WARNING] @@ -266,7 +307,7 @@ async def update_item(item_id: int, request: Request): return {"invalidated": await invalidate(StarletteRequest(scope))} ``` -`invalidate(request, key_builder=None)` 在項目存在且已移除時回傳 `True`,否則回傳 `False`,包括尚未設定後端的情況。後端本身的錯誤則會拋給呼叫端(見[後端發生錯誤時](#when-the-backend-fails))。傳入的請求必須能產生快取路由的鍵:相同的方法、主機、路徑與查詢字串。如果快取路由使用自訂的 `key_builder`,這裡也要傳入同一個,否則鍵不會相符。 +`invalidate(request, key_builder=None, vary=None)` 在項目存在且已移除時回傳 `True`,否則回傳 `False`,包括尚未設定後端的情況。後端本身的錯誤則會拋給呼叫端(見[後端發生錯誤時](#when-the-backend-fails))。傳入的請求必須能產生快取路由的鍵:相同的方法、主機、路徑與查詢字串。如果快取路由使用自訂的 `key_builder` 或 `vary`,這裡也要傳入相同的值,否則鍵不會相符;使用 `vary` 時,只會刪除請求本身的標頭值所選中的變體,`clear_path()` 則會移除所有變體。 ## 監控路由 {#monitoring-routes} diff --git a/tests/test_cache_vary.py b/tests/test_cache_vary.py new file mode 100644 index 0000000..848a011 --- /dev/null +++ b/tests/test_cache_vary.py @@ -0,0 +1,273 @@ +"""``@cache(vary=[...])`` keys on request headers and sends ``Vary`` (#268).""" + +from typing import Any + +import pytest +from fastapi import FastAPI +from fastapi import Request +from fastapi import Response +from fastapi.testclient import TestClient + +from fastapi_cachex import build_cache_key +from fastapi_cachex import invalidate +from fastapi_cachex.cache import cache +from fastapi_cachex.exceptions import CacheXError +from fastapi_cachex.proxy import BackendProxy + +BASE_KEY = "GET|||testserver|||/greet|||" + + +def _app(**cache_kwargs: Any) -> tuple[TestClient, dict[str, int]]: + app = FastAPI() + calls = {"n": 0} + + @app.api_route("/greet", methods=["GET", "POST"]) + @cache(ttl=60, **cache_kwargs) + async def greet(request: Request) -> dict[str, Any]: + calls["n"] += 1 + return {"lang": request.headers.get("accept-language"), "n": calls["n"]} + + return TestClient(app), calls + + +async def test_each_header_value_gets_its_own_entry() -> None: + client, calls = _app(vary=["Accept-Language"]) + + de = client.get("/greet", headers={"Accept-Language": "de"}) + en = client.get("/greet", headers={"Accept-Language": "en"}) + de_again = client.get("/greet", headers={"Accept-Language": "de"}) + + assert de.json() == {"lang": "de", "n": 1} + assert en.json() == {"lang": "en", "n": 2} + assert de_again.json() == {"lang": "de", "n": 1} + assert calls["n"] == 2 + assert sorted(await BackendProxy.get().get_all_keys()) == [ + f"{BASE_KEY}|||accept-language=de", + f"{BASE_KEY}|||accept-language=en", + ] + + +async def test_routes_without_vary_keep_their_key() -> None: + client, _ = _app() + + response = client.get("/greet", headers={"Accept-Language": "de"}) + + assert await BackendProxy.get().get_all_keys() == [BASE_KEY] + assert "vary" not in response.headers + + +async def test_values_are_trimmed_joined_and_escaped() -> None: + client, _ = _app(vary=["accept-language", "X-Variant"]) + + client.get("/greet", headers={"Accept-Language": " de "}) + client.get( + "/greet", + headers=[ + ("Accept-Language", "fr"), + ("Accept-Language", " en "), + ("X-Variant", "a|||b"), + ], + ) + + assert sorted(await BackendProxy.get().get_all_keys()) == [ + f"{BASE_KEY}|||accept-language=de|||x-variant=", + f"{BASE_KEY}|||accept-language=fr,en|||x-variant=a%7C%7C%7Cb", + ] + + +async def test_missing_and_empty_headers_share_an_entry() -> None: + client, calls = _app(vary=["Accept-Language"]) + + client.get("/greet") + client.get("/greet", headers={"Accept-Language": ""}) + + assert calls["n"] == 1 + assert await BackendProxy.get().get_all_keys() == [f"{BASE_KEY}|||accept-language="] + + +async def test_vary_components_follow_a_custom_key_builder() -> None: + def per_tenant(request: Request) -> str: + return build_cache_key(request, "tenant-1") + + client, _ = _app(vary=["Accept-Language"], key_builder=per_tenant) + client.get("/greet", headers={"Accept-Language": "de"}) + + assert await BackendProxy.get().get_all_keys() == [ + f"{BASE_KEY}|||tenant-1|||accept-language=de" + ] + + +async def test_clear_path_clears_every_variant() -> None: + client, _ = _app(vary=["Accept-Language"]) + for lang in ("de", "en", "fr"): + client.get("/greet", headers={"Accept-Language": lang}) + client.get("/greet?page=2", headers={"Accept-Language": "de"}) + + backend = BackendProxy.get() + assert await backend.clear_path("/greet") == 3 + assert await backend.clear_path("/greet", include_params=True) == 1 + + +async def test_invalidate_deletes_the_requested_variant() -> None: + app = FastAPI() + + @app.get("/greet") + @cache(ttl=60, vary=["Accept-Language"]) + async def greet(request: Request) -> dict[str, str | None]: + return {"lang": request.headers.get("accept-language")} + + @app.post("/greet") + async def reset(request: Request) -> dict[str, bool]: + request.scope["method"] = "GET" + return { + "without_vary": await invalidate(request), + "with_vary": await invalidate(request, vary=["Accept-Language"]), + } + + client = TestClient(app) + client.get("/greet", headers={"Accept-Language": "de"}) + client.get("/greet", headers={"Accept-Language": "en"}) + + result = client.post("/greet", headers={"Accept-Language": "de"}).json() + + assert result == {"without_vary": False, "with_vary": True} + assert await BackendProxy.get().get_all_keys() == [ + f"{BASE_KEY}|||accept-language=en" + ] + + +def test_vary_header_on_miss_hit_and_304() -> None: + client, _ = _app(vary=["Accept-Language"]) + headers = {"Accept-Language": "de"} + + miss = client.get("/greet", headers=headers) + hit = client.get("/greet", headers=headers) + not_modified = client.get( + "/greet", headers={**headers, "If-None-Match": miss.headers["etag"]} + ) + + assert not_modified.status_code == 304 + for response in (miss, hit, not_modified): + assert response.headers["vary"] == "Accept-Language" + + +@pytest.mark.parametrize( + ("cache_kwargs", "request_headers"), + [ + pytest.param({"private": True}, {}, id="private"), + pytest.param({"no_store": True}, {}, id="no_store"), + pytest.param({"no_cache": True}, {}, id="no_cache"), + pytest.param({}, {"Authorization": "Bearer a"}, id="authorization-bypass"), + ], +) +def test_vary_header_on_responses_that_skip_the_backend( + cache_kwargs: dict[str, Any], request_headers: dict[str, str] +) -> None: + client, _ = _app(vary=["Accept-Language"], **cache_kwargs) + + first = client.get("/greet", headers=request_headers) + revalidated = ( + client.get( + "/greet", + headers={**request_headers, "If-None-Match": first.headers["etag"]}, + ) + if "etag" in first.headers + else first + ) + + for response in (first, revalidated): + assert response.headers["vary"] == "Accept-Language" + + +def test_authorization_bypass_keeps_private_cache_control() -> None: + client, _ = _app(vary=["Accept-Language"], public=False) + + response = client.get("/greet", headers={"Authorization": "Bearer a"}) + + assert response.headers["cache-control"] == "private, max-age=60" + assert response.headers["vary"] == "Accept-Language" + + +async def test_unstored_cookie_response_gets_vary_and_private() -> None: + app = FastAPI() + + @app.get("/greet") + @cache(ttl=60, public=True, vary=["Accept-Language"]) + async def greet(response: Response) -> dict[str, str]: + response.set_cookie("seen", "1") + return {"ok": "yes"} + + response = TestClient(app).get("/greet") + + assert response.headers["vary"] == "Accept-Language" + assert response.headers["cache-control"] == "private, max-age=60" + assert await BackendProxy.get().get_all_keys() == [] + + +def test_names_already_in_vary_are_not_repeated() -> None: + app = FastAPI() + + @app.get("/greet") + @cache(ttl=60, vary=["accept-language", "Accept"]) + async def greet(response: Response) -> dict[str, str]: + response.headers["Vary"] = "Accept-Language, Origin" + return {"ok": "yes"} + + client = TestClient(app) + miss = client.get("/greet") + hit = client.get("/greet") + + for response in (miss, hit): + assert response.headers.get_list("vary") == ["Accept-Language, Origin, Accept"] + + +def test_vary_star_is_left_alone() -> None: + app = FastAPI() + + @app.get("/greet") + @cache(ttl=60, vary=["Accept-Language"]) + async def greet(response: Response) -> dict[str, str]: + response.headers["Vary"] = "*" + return {"ok": "yes"} + + client = TestClient(app) + + assert client.get("/greet").headers.get_list("vary") == ["*"] + assert client.get("/greet").headers.get_list("vary") == ["*"] + + +def test_non_get_requests_get_no_vary() -> None: + client, _ = _app(vary=["Accept-Language"]) + + assert "vary" not in client.post("/greet").headers + + +def test_duplicate_names_are_listed_once() -> None: + client, _ = _app(vary=["Accept-Language", "accept-language"]) + + assert client.get("/greet").headers["vary"] == "Accept-Language" + + +@pytest.mark.parametrize( + "vary", + [ + pytest.param("Accept", id="bare-str"), + pytest.param(b"Accept", id="bytes"), + pytest.param({"Accept"}, id="set"), + pytest.param([""], id="empty-name"), + pytest.param(["Accept Language"], id="space"), + pytest.param(["Accept,Origin"], id="comma"), + pytest.param([None], id="none"), + pytest.param(["*"], id="star"), + ], +) +def test_invalid_vary_is_rejected_at_decoration(vary: object) -> None: + with pytest.raises(CacheXError, match="vary"): + cache(ttl=60, vary=vary)(lambda: None) # type: ignore[arg-type] + + +async def test_invalid_vary_is_rejected_by_invalidate() -> None: + request = Request({"type": "http", "method": "GET", "path": "/", "headers": []}) + + with pytest.raises(CacheXError, match="vary"): + await invalidate(request, vary="Accept")