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")