Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 5 additions & 2 deletions docs/HTTP_CACHING.md
Original file line number Diff line number Diff line change
Expand Up @@ -335,10 +335,13 @@ add_routes(
- `GET {prefix}/cached-hits` — every cached entry split into method, host, path
and query, with its ETag and expiry, plus counts of valid and expired entries
and the distinct cached paths. It does not count hits.
- `GET {prefix}/cached-records` — every cached record with its size, expiry and
a preview of the first 100 bytes of the cached content. With
- `GET {prefix}/cached-records` — every cached record with its size, expiry,
`media_type` (the stored response's media type, `null` if it had none) and a
preview of the first 100 bytes of the cached content. With
`include_content_preview=False`, `content_preview` is `null` and no response
body leaves the server; keys, sizes and expiry are still reported.
`content_type` is always `"bytes"` and is kept for compatibility; read
`media_type` instead.

> [!WARNING]
> **These routes have no authentication of their own.** `include_in_schema=False`
Expand Down
112 changes: 71 additions & 41 deletions fastapi_cachex/cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -207,6 +207,57 @@ def _media_type_of(response: Response) -> str | None:
return response.headers.get("content-type")


def _build_cache_control(
*,
ttl: int | None,
stale: Literal["error", "revalidate"] | None,
stale_ttl: int | None,
no_cache: bool,
public: bool,
private: bool,
immutable: bool,
must_revalidate: bool,
) -> str:
"""The ``Cache-Control`` value for a ``@cache`` route's arguments.

``no_cache`` sends only ``no-cache`` (plus ``must-revalidate``); otherwise
the directives follow in a fixed order: scope, ``max-age``,
``must-revalidate``, the stale directive, ``immutable``.
"""
cache_control = CacheControl()
if no_cache:
cache_control.add(DirectiveType.NO_CACHE)
if must_revalidate:
cache_control.add(DirectiveType.MUST_REVALIDATE)
return str(cache_control)

# 1. Access scope (public/private)
if public:
cache_control.add(DirectiveType.PUBLIC)
elif private:
cache_control.add(DirectiveType.PRIVATE)

# 2. Cache time settings
if ttl is not None:
cache_control.add(DirectiveType.MAX_AGE, ttl)

# 3. Validation related
if must_revalidate:
cache_control.add(DirectiveType.MUST_REVALIDATE)

# 4. Stale response handling (stale_ttl is validated at decoration time)
if stale == "revalidate":
cache_control.add(DirectiveType.STALE_WHILE_REVALIDATE, stale_ttl)
elif stale == "error":
cache_control.add(DirectiveType.STALE_IF_ERROR, stale_ttl)

# 5. Special flags
if immutable:
cache_control.add(DirectiveType.IMMUTABLE)

return str(cache_control)


def _is_request_annotation(annotation: Any) -> bool:
"""Whether an annotation asks for a ``Request`` (or a subclass of one)."""
if get_origin(annotation) is Annotated:
Expand Down Expand Up @@ -555,42 +606,17 @@ def decorator(func: HandlerCallable) -> AsyncResponseCallable:
else:
request_name = found_request.name

def build_cache_control() -> str:
cache_control = CacheControl()
if no_cache:
cache_control.add(DirectiveType.NO_CACHE)
if must_revalidate:
cache_control.add(DirectiveType.MUST_REVALIDATE)
return str(cache_control)

# 1. Access scope (public/private)
if public:
cache_control.add(DirectiveType.PUBLIC)
elif private:
cache_control.add(DirectiveType.PRIVATE)

# 2. Cache time settings
if ttl is not None:
cache_control.add(DirectiveType.MAX_AGE, ttl)

# 3. Validation related
if must_revalidate:
cache_control.add(DirectiveType.MUST_REVALIDATE)

# 4. Stale response handling (stale_ttl is validated at decoration time)
if stale == "revalidate":
cache_control.add(DirectiveType.STALE_WHILE_REVALIDATE, stale_ttl)
elif stale == "error":
cache_control.add(DirectiveType.STALE_IF_ERROR, stale_ttl)

# 5. Special flags
if immutable:
cache_control.add(DirectiveType.IMMUTABLE)

return str(cache_control)

# The header only depends on the decorator arguments, so build it once.
cache_control = build_cache_control()
cache_control = _build_cache_control(
ttl=ttl,
stale=stale,
stale_ttl=stale_ttl,
no_cache=no_cache,
public=public,
private=private,
immutable=immutable,
must_revalidate=must_revalidate,
)
builder = key_builder or default_key_builder
# Without a positive ttl nothing may be served from storage, and a 304
# answered from a stored ETag would be exactly that: it would keep
Expand All @@ -609,7 +635,7 @@ async def wrapper(*args: Any, **kwargs: Any) -> Response:
else:
req = kwargs.pop(request_name, None)

if not req:
if req is None:
# Reached when the wrapper is called outside the router, which
# is the only caller that supplies the request parameter.
raise RequestNotFoundError
Expand All @@ -621,12 +647,12 @@ async def wrapper(*args: Any, **kwargs: Any) -> Response:
)
return await get_response(func, req, *args, **kwargs)

cache_key = builder(req)

# Handle special case: no-store (highest priority)
if no_store:
response = await get_response(func, req, *args, **kwargs)
logger.debug("no-store active; bypassed cache for key=%s", cache_key)
logger.debug(
"no-store active; bypassed cache for path=%s", req.url.path
)
return _with_cache_control(response, _NO_STORE)

client_etag = req.headers.get("if-none-match")
Expand All @@ -645,12 +671,16 @@ async def wrapper(*args: Any, **kwargs: Any) -> Response:
# StreamingResponse/FileResponse — cannot compute ETag
return _with_cache_control(response, cache_control)
if _etag_matches(client_etag, etag):
logger.debug("304 Not Modified (uncached); key=%s", cache_key)
logger.debug("304 Not Modified (uncached); path=%s", req.url.path)
return _not_modified(etag, cache_control, response.headers)
response.headers["ETag"] = etag
logger.debug("Bypassed the backend; key=%s", cache_key)
logger.debug("Bypassed the backend; path=%s", req.url.path)
return _with_cache_control(response, cache_control)

# Built only here: the branches above never touch the backend, so a
# custom key builder would run for nothing.
cache_key = builder(req)

try:
cached_data = await cache_backend.get(cache_key)
except Exception as e:
Expand Down
17 changes: 14 additions & 3 deletions fastapi_cachex/routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,11 @@

# Constants
CACHE_KEY_MIN_PARTS = 3
CACHE_KEY_MAX_PARTS = 3
# A ``maxsplit`` count for ``str.split``, not a number of parts: splitting at
# most three times keeps a query string containing the separator in one piece.
CACHE_KEY_MAX_SPLIT = 3
# Former name, kept so existing imports keep working.
CACHE_KEY_MAX_PARTS = CACHE_KEY_MAX_SPLIT
_PREVIEW_BYTES = 100


Expand Down Expand Up @@ -61,7 +65,12 @@ class CacheHitsResponse:

@dataclass
class CachedRecord:
"""Record for a single cached item."""
"""Record for a single cached item.

``media_type`` is the stored response's media type (``None`` when it had
none). ``content_type`` is always ``"bytes"``, the type of the stored body,
and is kept for compatibility.
"""

cache_key: str
method: str
Expand All @@ -74,6 +83,7 @@ class CachedRecord:
is_expired: bool
ttl_remaining: float | None
content_preview: str | None
media_type: str | None


@dataclass
Expand Down Expand Up @@ -106,7 +116,7 @@ def _parse_cache_key(cache_key: str) -> tuple[str, str, str, str]:
Returns:
Tuple of (method, host, path, query_params)
"""
key_parts = cache_key.split(CACHE_KEY_SEPARATOR, CACHE_KEY_MAX_PARTS)
key_parts = cache_key.split(CACHE_KEY_SEPARATOR, CACHE_KEY_MAX_SPLIT)
if len(key_parts) >= CACHE_KEY_MIN_PARTS:
method = key_parts[0]
host = unescape_key_component(key_parts[1])
Expand Down Expand Up @@ -218,6 +228,7 @@ def _cached_records(
if include_content_preview
else None
),
media_type=e.entry.media_type,
)
for e in entries
]
Expand Down
2 changes: 1 addition & 1 deletion i18n/zh-TW/docs/HTTP_CACHING.md
Original file line number Diff line number Diff line change
Expand Up @@ -259,7 +259,7 @@ add_routes(
```

- `GET {prefix}/cached-hits`:列出每筆快取項目,拆分為方法、主機、路徑與查詢,附上 ETag 與到期時間,另外統計有效與已過期的項目數,以及不重複的快取路徑。它不會計算命中次數。
- `GET {prefix}/cached-records`:列出每筆快取紀錄的大小、到期時間,以及快取內容前 100 個位元組的預覽。設定 `include_content_preview=False` 時,`content_preview` 為 `null`,不會有任何回應本文離開伺服器;鍵、大小與到期時間仍會回報。
- `GET {prefix}/cached-records`:列出每筆快取紀錄的大小、到期時間、`media_type`(儲存的回應的媒體類型,沒有時為 `null`),以及快取內容前 100 個位元組的預覽。設定 `include_content_preview=False` 時,`content_preview` 為 `null`,不會有任何回應本文離開伺服器;鍵、大小與到期時間仍會回報。`content_type` 一律是 `"bytes"`,只為相容而保留;請改讀 `media_type`。

> [!WARNING]
> **這些路由本身沒有任何身分驗證。** `include_in_schema=False` 只是讓它們不出現在 OpenAPI 文件中;任何猜到路徑的人都能讀取。`/cached-records` 含有快取內容的預覽(除非設定 `include_content_preview=False`),並會暴露整個路由結構。正式環境中請務必傳入 `dependencies=[Depends(your_auth)]`,或將它們掛載在僅供內部使用的應用程式上。
Expand Down
118 changes: 118 additions & 0 deletions tests/test_cache_internals.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,118 @@
"""The ``@cache`` wrapper builds a key only when it touches the backend (#182)."""

import pytest
from fastapi import FastAPI
from fastapi import Request
from fastapi.testclient import TestClient

from fastapi_cachex.cache import _build_cache_control
from fastapi_cachex.cache import cache
from fastapi_cachex.cache import default_key_builder


def _app_counting_keys(**cache_kwargs: object) -> tuple[TestClient, list[str]]:
calls: list[str] = []

def counting_builder(request: Request) -> str:
calls.append(request.url.path)
return default_key_builder(request)

app = FastAPI()

@app.api_route("/item", methods=["GET", "POST"])
@cache(key_builder=counting_builder, **cache_kwargs) # type: ignore[arg-type]
async def item() -> dict[str, str]:
return {"ok": "yes"}

return TestClient(app), calls


@pytest.mark.parametrize(
"cache_kwargs",
[
pytest.param({"ttl": 60, "no_store": True}, id="no_store"),
pytest.param({"ttl": 60, "private": True}, id="private"),
pytest.param({}, id="no-ttl"),
pytest.param({"ttl": 0}, id="ttl-0"),
],
)
def test_key_builder_not_called_when_the_backend_is_skipped(
cache_kwargs: dict[str, object],
) -> None:
"""Routes that never read or write the backend do not build a key."""
client, calls = _app_counting_keys(**cache_kwargs)

first = client.get("/item")
client.get("/item", headers={"If-None-Match": first.headers.get("etag", "x")})

assert first.status_code == 200
assert calls == []


def test_key_builder_not_called_for_non_get() -> None:
"""Only GET is cached, so other methods do not build a key."""
client, calls = _app_counting_keys(ttl=60)

assert client.post("/item").status_code == 200
assert calls == []


def test_key_builder_called_once_per_cached_request() -> None:
"""A cached route still builds its key on every GET, miss and hit alike."""
client, calls = _app_counting_keys(ttl=60)

miss = client.get("/item")
hit = client.get("/item")
not_modified = client.get("/item", headers={"If-None-Match": miss.headers["etag"]})

assert (miss.status_code, hit.status_code, not_modified.status_code) == (
200,
200,
304,
)
assert calls == ["/item", "/item", "/item"]


_NO_FLAGS = {
"ttl": None,
"stale": None,
"stale_ttl": None,
"no_cache": False,
"public": False,
"private": False,
"immutable": False,
"must_revalidate": False,
}


@pytest.mark.parametrize(
("overrides", "expected"),
[
pytest.param({}, "", id="nothing"),
pytest.param(
{
"ttl": 60,
"public": True,
"must_revalidate": True,
"stale": "revalidate",
"stale_ttl": 30,
"immutable": True,
},
"public, max-age=60, must-revalidate, stale-while-revalidate=30, immutable",
id="directive-order",
),
pytest.param(
{"ttl": 60, "private": True, "stale": "error", "stale_ttl": 5},
"private, max-age=60, stale-if-error=5",
id="private-stale-if-error",
),
pytest.param(
{"no_cache": True, "ttl": 60, "public": True, "must_revalidate": True},
"no-cache, must-revalidate",
id="no-cache-drops-the-rest",
),
],
)
def test_build_cache_control(overrides: dict[str, object], expected: str) -> None:
"""The header is a pure function of the decorator arguments."""
assert _build_cache_control(**{**_NO_FLAGS, **overrides}) == expected # type: ignore[arg-type]
Loading
Loading