diff --git a/packages/reflex-base/src/reflex_base/.templates/web/isr-render.mjs b/packages/reflex-base/src/reflex_base/.templates/web/isr-render.mjs new file mode 100644 index 00000000000..c4edace1a51 --- /dev/null +++ b/packages/reflex-base/src/reflex_base/.templates/web/isr-render.mjs @@ -0,0 +1,109 @@ +/** + * ISR render tier for Reflex. + * + * A small HTTP service that renders the react-router **server bundle** to HTML + * for a given route, so the Python ISR page-server (reflex.isr.HttpRenderer) + * can cache and serve it. This is the one component that must run JavaScript; + * it is deployed as its own service (off the request hot path, scale-to-zero + * friendly) rather than in every frontend pod. + * + * Contract (matches reflex.isr.HttpRenderer): + * POST / { "path": "/blog/hello" } + * 200 -> { "html": "…", "revalidate"?: number, "tags"?: string[] } + * non-200 -> ISR treats it as a render failure (serves stale/shell) + * + * Requirements: + * - An SSR build of the app (react-router build with `ssr: true`) producing + * `build/server/index.js`. The app root `loader()` fetches page state from + * the Python `/_ssr_data` endpoint, so this service needs network access to + * the backend (set SSR_DATA / the backend URL in the app env). + * - Node deps: `react-router`, `express`. + * + * Run: node isr-render.mjs (listens on $PORT, default 8600) + * + * Optional per-page cache hints: the app may set response headers + * `X-Reflex-Revalidate: ` and `X-Reflex-Tags: a,b,c` + * which are forwarded to the ISR cache; otherwise Python applies its defaults. + */ +import { createRequestHandler } from "react-router"; +import express from "express"; +import { dirname, join } from "node:path"; +import { fileURLToPath } from "node:url"; + +const __dirname = + import.meta.dirname ?? dirname(fileURLToPath(import.meta.url)); + +const serverBuild = join(__dirname, "build", "server", "index.js"); +const build = await import(serverBuild); +const handler = createRequestHandler(build, "production"); + +const app = express(); +app.use(express.json({ limit: "1mb" })); + +// Health endpoint for k8s probes. +app.get("/_health", (_req, res) => res.status(200).send("OK")); + +app.post("/", async (req, res) => { + const path = + typeof req.body?.path === "string" && req.body.path ? req.body.path : "/"; + + try { + // Render the target route through react-router as a normal GET document + // request. The app's loader runs here and pulls state from /_ssr_data. + const url = new URL(path, "http://reflex-isr-render.local"); + const request = new Request(url, { + method: "GET", + headers: { + "User-Agent": "ReflexISRRenderer/1.0", + Accept: "text/html", + // Forward the caller's cookies so authenticated/personalized loaders + // resolve the same way they would for the original visitor, if the + // page-server chooses to pass them along. + ...(req.body?.cookie ? { Cookie: String(req.body.cookie) } : {}), + }, + }); + + const response = await handler(request); + const html = await response.text(); + + if (!response.ok || !html) { + return res + .status(502) + .json({ error: "render_failed", status: response.status }); + } + + const revalidate = headerNumber(response, "x-reflex-revalidate"); + const tags = headerList(response, "x-reflex-tags"); + + return res.json({ + html, + ...(revalidate != null ? { revalidate } : {}), + ...(tags.length ? { tags } : {}), + }); + } catch (err) { + console.error("[isr-render] render error for", path, err); + return res.status(500).json({ error: String(err) }); + } +}); + +function headerNumber(response, name) { + const raw = response.headers.get(name); + if (raw == null) return null; + const n = Number(raw); + return Number.isFinite(n) ? n : null; +} + +function headerList(response, name) { + const raw = response.headers.get(name); + return raw + ? raw + .split(",") + .map((s) => s.trim()) + .filter(Boolean) + : []; +} + +const port = parseInt(process.env.PORT || "8600", 10); +app.listen(port, () => { + console.log(`[isr-render] listening on http://localhost:${port}`); +}); diff --git a/packages/reflex-base/src/reflex_base/compiler/templates.py b/packages/reflex-base/src/reflex_base/compiler/templates.py index ca94aa2c482..e358e8da467 100644 --- a/packages/reflex-base/src/reflex_base/compiler/templates.py +++ b/packages/reflex-base/src/reflex_base/compiler/templates.py @@ -208,6 +208,50 @@ def app_root_template( f' "{lib_path}": {lib_alias},' for lib_alias, lib_path in window_libraries ]) + from reflex_base import ssr + + if ssr.is_enabled(): + ssr_imports = ( + 'import { useLoaderData } from "react-router";\n' + 'import { getBackendURL } from "$/utils/state";\n' + 'import { SSRStateContext } from "$/utils/context";\n' + 'import env from "$/env.json";' + ) + ssr_loader = """ +export async function loader({ request }) { + // Fetch on_load-hydrated state from the Python /_ssr_data endpoint so the + // server-rendered HTML already contains real data. + try { + const res = await fetch(getBackendURL(env.SSR_DATA), { + method: "POST", + headers: { + "Content-Type": "application/json", + ...(request.headers.get("cookie") ? { Cookie: request.headers.get("cookie") } : {}), + }, + body: JSON.stringify({ + path: new URL(request.url).pathname, + headers: Object.fromEntries(request.headers), + }), + }); + if (res.ok) return res.json(); + } catch (e) { + console.error("SSR data fetch failed:", e); + } + return { state: null }; +} +""" + layout_body = ( + "const loaderData = useLoaderData();\n" + " const ssrState = loaderData?.state ?? null;\n" + " return jsx(AppLayout, {}, " + "jsx(SSRStateContext.Provider, {value: ssrState}, " + "jsx(ReflexProviders, {}, children)));" + ) + else: + ssr_imports = "" + ssr_loader = "" + layout_body = "return jsx(AppLayout, {}, jsx(ReflexProviders, {}, children));" + return f""" {imports_str} {dynamic_imports_str} @@ -215,10 +259,11 @@ def app_root_template( import {{ ThemeProvider }} from '$/utils/react-theme'; import {{ Layout as AppLayout }} from './_document'; import {{ Outlet }} from 'react-router'; +{ssr_imports} {import_window_libraries} {custom_code_str} - +{ssr_loader} function ReflexProviders({{children}}) {{ useEffect(() => {{ // Make contexts and state objects available globally for dynamic eval'd components @@ -240,7 +285,7 @@ def app_root_template( }} export function Layout({{children}}) {{ - return jsx(AppLayout, {{}}, jsx(ReflexProviders, {{}}, children)); + {layout_body} }} // Used by entry.client.js when mount_target is configured: skips the document @@ -289,6 +334,9 @@ def context_template( Returns: Rendered context file content as string. """ + from reflex_base import ssr + + ssr_enabled = ssr.is_enabled() initial_state = initial_state or {} state_contexts_str = "".join([ f"{format_state_name(state_name)}: createContext(null)," @@ -342,10 +390,16 @@ def context_template( """ ) - state_reducer_str = "\n".join( - rf'const [{format_state_name(state_name)}, dispatch_{format_state_name(state_name)}] = useReducer(applyDelta, initialState["{state_name}"])' - for state_name in initial_state - ) + if ssr_enabled: + state_reducer_str = "\n".join( + rf'const [{format_state_name(state_name)}, dispatch_{format_state_name(state_name)}] = useReducer(applyDelta, (ssrState && ssrState["{state_name}"] != null) ? ssrState["{state_name}"] : initialState["{state_name}"])' + for state_name in initial_state + ) + else: + state_reducer_str = "\n".join( + rf'const [{format_state_name(state_name)}, dispatch_{format_state_name(state_name)}] = useReducer(applyDelta, initialState["{state_name}"])' + for state_name in initial_state + ) create_state_contexts_str = "\n".join( rf"createElement(StateContexts.{format_state_name(state_name)},{{value: {format_state_name(state_name)}}}," @@ -374,6 +428,7 @@ def context_template( export const DispatchContext = createContext(null); export const StateContexts = {{{state_contexts_str}}}; export const EventLoopContext = createContext(null); +{"export const SSRStateContext = createContext(null);" if ssr_enabled else ""} export const clientStorage = {"{}" if client_storage is None else json.dumps(client_storage)} {state_str} @@ -446,6 +501,7 @@ def context_template( }} export function StateProvider({{ children }}) {{ + {"const ssrState = useContext(SSRStateContext);" if ssr_enabled else ""} {state_reducer_str} const dispatchers = useMemo(() => {{ return {{ diff --git a/packages/reflex-base/src/reflex_base/config.py b/packages/reflex-base/src/reflex_base/config.py index be02f00a937..57a70f99d20 100644 --- a/packages/reflex-base/src/reflex_base/config.py +++ b/packages/reflex-base/src/reflex_base/config.py @@ -186,6 +186,10 @@ class BaseConfig: plugins: List of plugins to use in the app. disable_plugins: List of plugin types to disable in the app. transport: The transport method for client-server communication. + ssr_mode: Server-side rendering mode ("off", "bot_only", or "always"). + isr_render_url: ISR render-tier URL; when set, enables ISR page serving. + isr_revalidate: Default seconds until an ISR-cached page goes stale. + isr_build_id: Identifier that rotates the ISR cache on each deploy. """ app_name: str @@ -271,6 +275,20 @@ class BaseConfig: transport: Literal["websocket", "polling"] = "websocket" + # Server-side rendering mode: "off" (default), "bot_only", or "always". + ssr_mode: constants.SsrMode = constants.SsrMode.OFF + + # ISR render-tier URL. When set, the ISR page-server regenerates pages by + # POSTing paths here; when None, ISR is disabled. + isr_render_url: str | None = None + + # Default seconds until an ISR-cached page is considered stale. + isr_revalidate: int = 60 + + # Identifier that rotates the ISR cache on each deploy. Defaults at runtime + # to the REFLEX_BUILD_ID env var (see reflex.isr.get_build_id). + isr_build_id: str | None = None + # Whether to skip plugin checks. _skip_plugins_checks: bool = dataclasses.field(default=False, repr=False) diff --git a/packages/reflex-base/src/reflex_base/constants/__init__.py b/packages/reflex-base/src/reflex_base/constants/__init__.py index 714cf0faa84..347b7345361 100644 --- a/packages/reflex-base/src/reflex_base/constants/__init__.py +++ b/packages/reflex-base/src/reflex_base/constants/__init__.py @@ -25,6 +25,7 @@ Reflex, ReflexHostingCLI, RunningMode, + SsrMode, Templates, ) from .compiler import ( @@ -129,6 +130,7 @@ "RouteVar", "RunningMode", "SocketEvent", + "SsrMode", "StateManagerMode", "Templates", "UvLock", diff --git a/packages/reflex-base/src/reflex_base/constants/base.py b/packages/reflex-base/src/reflex_base/constants/base.py index c1b91dd5014..fabda3589c8 100644 --- a/packages/reflex-base/src/reflex_base/constants/base.py +++ b/packages/reflex-base/src/reflex_base/constants/base.py @@ -165,8 +165,8 @@ class ReactRouter(Javascript): DEV_FRONTEND_LISTENING_REGEX = r"Local:[\s]+" # Regex to pattern the route path in the config file - # INFO Accepting connections at http://localhost:3000 - PROD_FRONTEND_LISTENING_REGEX = r"Accepting connections at[\s]+" + # Matches output from sirv ("Accepting connections at") or ssr-serve ("[ssr-serve]"). + PROD_FRONTEND_LISTENING_REGEX = r"(?:Accepting connections at|\[ssr-serve\])[\s]+" FRONTEND_LISTENING_REGEX = ( rf"(?:{DEV_FRONTEND_LISTENING_REGEX}|{PROD_FRONTEND_LISTENING_REGEX})(.*)" @@ -203,6 +203,19 @@ class Env(str, Enum): PROD = "prod" +class SsrMode(str, Enum): + """Server-side rendering modes. + + OFF: No SSR — static SPA (default, same as before). + BOT_ONLY: SSR for crawlers/bots only — regular users get the SPA shell. + ALWAYS: SSR for all users — better FCP/LCP, higher server load. + """ + + OFF = "off" + BOT_ONLY = "bot_only" + ALWAYS = "always" + + class RunningMode(str, Enum): """The running modes.""" diff --git a/packages/reflex-base/src/reflex_base/constants/event.py b/packages/reflex-base/src/reflex_base/constants/event.py index 515b34574e9..ff37f2be00b 100644 --- a/packages/reflex-base/src/reflex_base/constants/event.py +++ b/packages/reflex-base/src/reflex_base/constants/event.py @@ -13,6 +13,7 @@ class Endpoint(Enum): AUTH_CODESPACE = "auth-codespace" HEALTH = "_health" ALL_ROUTES = "_all_routes" + SSR_DATA = "_ssr_data" def __str__(self) -> str: """Get the string representation of the endpoint. diff --git a/packages/reflex-base/src/reflex_base/ssr.py b/packages/reflex-base/src/reflex_base/ssr.py new file mode 100644 index 00000000000..6875c3b8c0b --- /dev/null +++ b/packages/reflex-base/src/reflex_base/ssr.py @@ -0,0 +1,46 @@ +"""Server-side rendering (SSR) config gate for the base package. + +The request-time backend (the ``/_ssr_data`` endpoint) lives in ``reflex.ssr`` +because it depends on the state/app machinery that is not part of +``reflex_base``. Only the config gate lives here so ``reflex_base`` code can +check whether SSR is enabled without importing the top-level ``reflex`` package. + +SSR is a no-op unless ``config.ssr_mode`` is not ``OFF``. +""" + +from __future__ import annotations + +import os +from typing import TYPE_CHECKING + +from reflex_base import constants +from reflex_base.config import get_config + +if TYPE_CHECKING: + from reflex_base.config import Config + + +def is_enabled(config: Config | None = None) -> bool: + """Whether SSR is enabled. + + Args: + config: The config to read from, or the global config when omitted. + + Returns: + True when ``config.ssr_mode`` is not ``OFF``. + """ + config = config if config is not None else get_config() + return config.ssr_mode != constants.SsrMode.OFF + + +def ssr_build_enabled() -> bool: + """Whether to emit an ``ssr: true`` react-router build. + + Controlled by the ``REFLEX_SSR_BUILD`` env var so the same app can produce + both the default static SPA build (served by nginx) and a server-render + build (consumed by the ISR render tier) from separate build invocations. + + Returns: + True when ``REFLEX_SSR_BUILD`` is set to ``1``. + """ + return os.environ.get("REFLEX_SSR_BUILD") == "1" diff --git a/pyi_hashes.json b/pyi_hashes.json index 481c3e8ef8b..d86c5f02478 100644 --- a/pyi_hashes.json +++ b/pyi_hashes.json @@ -118,7 +118,7 @@ "packages/reflex-components-recharts/src/reflex_components_recharts/polar.pyi": "99ebcfc07868061bdc3c2010d85a153f", "packages/reflex-components-recharts/src/reflex_components_recharts/recharts.pyi": "4f6c26f8c76543cc41e2b9dc400ece8a", "packages/reflex-components-sonner/src/reflex_components_sonner/toast.pyi": "58521fcd1b514804f534d97624e82c9a", - "reflex/__init__.pyi": "56385a4f0d9431eb0056dbc5553a58f9", + "reflex/__init__.pyi": "5f82b639e054023bc5e23a867c0f98ff", "reflex/components/__init__.pyi": "9facd05a776d0641432696bbf8e34388", "reflex/experimental/memo.pyi": "36e5d5f97eb64e94c0974e909e7e2952" } diff --git a/reflex/__init__.py b/reflex/__init__.py index 6b23b495d3e..f395b841d4c 100644 --- a/reflex/__init__.py +++ b/reflex/__init__.py @@ -182,7 +182,7 @@ "app": ["App", "UploadFile"], "assets": ["asset"], "config": ["Config", "DBConfig"], - "constants": ["Env"], + "constants": ["Env", "SsrMode"], "constants.colors": ["Color"], "_upload": [ "UploadChunk", diff --git a/reflex/app.py b/reflex/app.py index 9364c547436..71eec12668f 100644 --- a/reflex/app.py +++ b/reflex/app.py @@ -66,6 +66,7 @@ from starlette.staticfiles import StaticFiles from typing_extensions import Unpack +from reflex import ssr from reflex._upload import UploadedFilesHeadersMiddleware, upload from reflex._upload import UploadFile as UploadFile from reflex.admin import AdminDash @@ -787,6 +788,12 @@ def _add_default_endpoints(self): health, methods=["GET"], ) + if ssr.is_enabled(): + self._api.add_route( + config.prepend_backend_path(str(constants.Endpoint.SSR_DATA)), + ssr.ssr_data(self), + methods=["POST"], + ) def _add_optional_endpoints(self): """Add optional api endpoints (_upload).""" diff --git a/reflex/constants/__init__.py b/reflex/constants/__init__.py index b37780a3acb..a99bbdb59ba 100644 --- a/reflex/constants/__init__.py +++ b/reflex/constants/__init__.py @@ -20,6 +20,7 @@ ReactRouter, Reflex, ReflexHostingCLI, + SsrMode, Templates, ) from .compiler import ( @@ -114,6 +115,7 @@ "RouteRegex", "RouteVar", "SocketEvent", + "SsrMode", "StateManagerMode", "Templates", "UvLock", diff --git a/reflex/isr.py b/reflex/isr.py new file mode 100644 index 00000000000..28db1a9003c --- /dev/null +++ b/reflex/isr.py @@ -0,0 +1,836 @@ +"""Incremental Static Regeneration (ISR) for Reflex. + +Serves prerendered HTML from a shared cache with stale-while-revalidate (SWR) +semantics and cross-worker single-flight regeneration, so multiple workers +(e.g. Kubernetes replicas) share one cache and never double-render a page. + +Responsibilities split: + +* This module owns **caching, coordination, and revalidation** — pure Python, + no JS runtime. It works across workers when backed by Redis. +* Producing the HTML is delegated to a pluggable :data:`Renderer` (e.g. a Node + render tier that runs the react-router server bundle and fetches state from + the ``/_ssr_data`` endpoint). ISR itself never renders React. + +Cache entries are keyed by ``{build_id}:{path}`` so a new deploy (new build id) +transparently rotates the cache and never serves HTML that references +content-hashed assets from a previous build. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import dataclasses +import json +import os +import time +from collections.abc import Awaitable, Callable +from typing import TYPE_CHECKING, Protocol, runtime_checkable + +from reflex.utils import console + +if TYPE_CHECKING: + import httpx + from redis.asyncio import Redis + from starlette.applications import Starlette + + from reflex.config import Config + + +@dataclasses.dataclass(frozen=True) +class RenderResult: + """The result of rendering a page. + + Attributes: + html: The rendered HTML document. + revalidate: Seconds until the entry goes stale (``None`` = use default). + tags: Revalidation tags for on-demand invalidation of related pages. + """ + + html: str + revalidate: float | None = None + tags: tuple[str, ...] = () + + +# A Renderer turns a concrete request path into a RenderResult, or None on +# failure (in which case ISR serves stale/fallback rather than caching an error). +Renderer = Callable[[str], Awaitable["RenderResult | None"]] + + +@dataclasses.dataclass(frozen=True) +class CachedPage: + """A cached rendered page. + + Attributes: + html: The rendered HTML document. + generated_at: Unix timestamp when the page was rendered. + revalidate: Seconds until the page is considered stale (``<= 0`` never). + tags: Revalidation tags associated with the page. + """ + + html: str + generated_at: float + revalidate: float + tags: tuple[str, ...] = () + + def is_stale(self, now: float | None = None) -> bool: + """Whether the page is past its revalidation window. + + Args: + now: The current time (defaults to ``time.time()``). + + Returns: + True if the page should be revalidated. + """ + if self.revalidate <= 0: + return False + now = time.time() if now is None else now + return (now - self.generated_at) >= self.revalidate + + def to_json(self) -> str: + """Serialize the page for storage. + + Returns: + A JSON string. + """ + return json.dumps(dataclasses.asdict(self)) + + @classmethod + def from_json(cls, raw: str | bytes) -> CachedPage: + """Deserialize a page from storage. + + Args: + raw: The JSON string/bytes produced by :meth:`to_json`. + + Returns: + The reconstructed page. + """ + data = json.loads(raw) + return cls( + html=data["html"], + generated_at=data["generated_at"], + revalidate=data["revalidate"], + tags=tuple(data.get("tags", ())), + ) + + +@runtime_checkable +class ISRCache(Protocol): + """Shared cache + coordination backend for ISR. + + A Redis-backed implementation makes the cache, single-flight lock, and tag + index shared across all workers; the in-memory implementation is per-process + (dev/test only). + """ + + async def get(self, key: str) -> CachedPage | None: + """Get a cached page. + + Args: + key: The cache key. + + Returns: + The cached page, or None if absent. + """ + ... + + async def set(self, key: str, page: CachedPage) -> None: + """Store a page and index its tags. + + Args: + key: The cache key. + page: The page to store. + """ + ... + + async def delete(self, key: str) -> None: + """Remove a cached page. + + Args: + key: The cache key. + """ + ... + + async def acquire_lock(self, key: str, ttl: float) -> bool: + """Try to acquire the single-flight render lock for ``key``. + + Args: + key: The cache key. + ttl: Lock expiry in seconds. + + Returns: + True if this caller now holds the lock. + """ + ... + + async def release_lock(self, key: str) -> None: + """Release the single-flight render lock for ``key``. + + Args: + key: The cache key. + """ + ... + + async def invalidate_tag(self, tag: str) -> int: + """Delete every cached page carrying ``tag``. + + Args: + tag: The tag to invalidate. + + Returns: + The number of pages removed. + """ + ... + + +class MemoryISRCache: + """In-process ISR cache (single worker; dev/test only). + + Not shared across workers — use :class:`RedisISRCache` in production. + """ + + def __init__(self) -> None: + """Initialize empty stores.""" + self._pages: dict[str, CachedPage] = {} + self._tags: dict[str, set[str]] = {} + self._locks: set[str] = set() + + async def get(self, key: str) -> CachedPage | None: + """Get a cached page. + + Args: + key: The cache key. + + Returns: + The cached page, or None. + """ + return self._pages.get(key) + + async def set(self, key: str, page: CachedPage) -> None: + """Store a page and index its tags. + + Args: + key: The cache key. + page: The page to store. + """ + self._pages[key] = page + for tag in page.tags: + self._tags.setdefault(tag, set()).add(key) + + async def delete(self, key: str) -> None: + """Remove a cached page. + + Args: + key: The cache key. + """ + self._pages.pop(key, None) + + async def acquire_lock(self, key: str, ttl: float) -> bool: + """Acquire the render lock (no TTL needed in-process). + + Args: + key: The cache key. + ttl: Unused for the in-memory backend. + + Returns: + True if the lock was acquired. + """ + if key in self._locks: + return False + self._locks.add(key) + return True + + async def release_lock(self, key: str) -> None: + """Release the render lock. + + Args: + key: The cache key. + """ + self._locks.discard(key) + + async def invalidate_tag(self, tag: str) -> int: + """Delete all pages carrying ``tag``. + + Args: + tag: The tag to invalidate. + + Returns: + The number of pages removed. + """ + keys = self._tags.pop(tag, set()) + for key in keys: + self._pages.pop(key, None) + return len(keys) + + +class RedisISRCache: + """Redis-backed ISR cache shared across all workers.""" + + _PAGE_PREFIX = "reflex_isr:page:" + _LOCK_PREFIX = "reflex_isr:lock:" + _TAG_PREFIX = "reflex_isr:tag:" + + def __init__(self, redis: Redis) -> None: + """Store the async Redis client. + + Args: + redis: An async Redis client. + """ + self._redis = redis + + async def get(self, key: str) -> CachedPage | None: + """Get a cached page. + + Args: + key: The cache key. + + Returns: + The cached page, or None. + """ + raw = await self._redis.get(self._PAGE_PREFIX + key) + return CachedPage.from_json(raw) if raw is not None else None + + async def set(self, key: str, page: CachedPage) -> None: + """Store a page and index its tags. + + The stored value gets a generous TTL (well past ``revalidate``) so stale + entries remain available for stale-while-revalidate but do not leak + forever. + + Args: + key: The cache key. + page: The page to store. + """ + # Keep stale entries around for SWR: TTL = revalidate window + 24h grace. + ttl = int(page.revalidate) + 86400 if page.revalidate > 0 else None + pipe = self._redis.pipeline() + pipe.set(self._PAGE_PREFIX + key, page.to_json(), ex=ttl) + for tag in page.tags: + pipe.sadd(self._TAG_PREFIX + tag, key) + await pipe.execute() + + async def delete(self, key: str) -> None: + """Remove a cached page. + + Args: + key: The cache key. + """ + await self._redis.delete(self._PAGE_PREFIX + key) + + async def acquire_lock(self, key: str, ttl: float) -> bool: + """Acquire the single-flight lock via ``SET NX EX``. + + Args: + key: The cache key. + ttl: Lock expiry in seconds (guards against a crashed holder). + + Returns: + True if the lock was acquired. + """ + acquired = await self._redis.set( + self._LOCK_PREFIX + key, b"1", nx=True, ex=max(1, int(ttl)) + ) + return bool(acquired) + + async def release_lock(self, key: str) -> None: + """Release the single-flight lock. + + Args: + key: The cache key. + """ + await self._redis.delete(self._LOCK_PREFIX + key) + + async def invalidate_tag(self, tag: str) -> int: + """Delete all pages carrying ``tag``. + + Args: + tag: The tag to invalidate. + + Returns: + The number of pages removed. + """ + tag_key = self._TAG_PREFIX + tag + # redis.asyncio stubs type set-commands with their sync (non-awaitable) + # return type; smembers is awaitable at runtime. + keys = await self._redis.smembers(tag_key) # pyright: ignore[reportGeneralTypeIssues] + if not keys: + return 0 + page_keys = [ + self._PAGE_PREFIX + (k.decode() if isinstance(k, bytes) else k) + for k in keys + ] + await self._redis.delete(*page_keys, tag_key) + return len(page_keys) + + +def cache_key(build_id: str, path: str) -> str: + """Build the cache key for a page. + + Args: + build_id: The current build identifier (rotates the cache per deploy). + path: The request path. + + Returns: + The cache key. + """ + return f"{build_id}:{path}" + + +class ISRManager: + """Coordinates ISR serving: cache lookup, SWR, and single-flight rendering. + + Args: + cache: The shared cache/coordination backend. + renderer: Produces HTML for a path (e.g. calls the Node render tier). + build_id: Identifier that rotates the cache on each deploy. + default_revalidate: Default seconds-until-stale when a render does not + specify its own. + lock_ttl: Max seconds a single-flight render may hold the lock. + wait_timeout: Max seconds a request waits for another worker's in-flight + render on a cache miss before falling back. + """ + + def __init__( + self, + cache: ISRCache, + renderer: Renderer, + *, + build_id: str = "dev", + default_revalidate: float = 60.0, + lock_ttl: float = 30.0, + wait_timeout: float = 5.0, + ) -> None: + """Store configuration.""" + self.cache = cache + self.renderer = renderer + self.build_id = build_id + self.default_revalidate = default_revalidate + self.lock_ttl = lock_ttl + self.wait_timeout = wait_timeout + self._background: set[asyncio.Task] = set() + + async def get_html(self, path: str) -> str | None: + """Return HTML for ``path`` using ISR semantics. + + * Fresh cache hit -> serve immediately. + * Stale hit -> serve stale immediately, revalidate in the background + (single-flight). + * Miss -> single-flight render; losing waiters briefly poll for the + winner's result, then fall back to ``None`` (caller serves the shell). + + Args: + path: The concrete request path. + + Returns: + The HTML string, or None if it could not be produced. + """ + key = cache_key(self.build_id, path) + page = await self.cache.get(key) + + if page is not None and not page.is_stale(): + return page.html + + if page is not None: + # Stale-while-revalidate: serve stale now, refresh in the background. + self._spawn_revalidate(key, path) + return page.html + + # Cache miss: block on a single-flight render. + return await self._render_single_flight(key, path) + + async def _render_single_flight(self, key: str, path: str) -> str | None: + """Render ``path`` once across workers, or wait for the winner. + + Args: + key: The cache key. + path: The request path. + + Returns: + The rendered HTML, or None on failure/timeout. + """ + if await self.cache.acquire_lock(key, self.lock_ttl): + try: + return await self._render_and_store(key, path) + finally: + await self.cache.release_lock(key) + + # Another worker is rendering; poll briefly for its result. + deadline = time.monotonic() + self.wait_timeout + while time.monotonic() < deadline: + await asyncio.sleep(0.05) + page = await self.cache.get(key) + if page is not None: + return page.html + return None + + async def _render_and_store(self, key: str, path: str) -> str | None: + """Render ``path`` and cache the result. + + Args: + key: The cache key. + path: The request path. + + Returns: + The rendered HTML, or None on failure. + """ + try: + result = await self.renderer(path) + except Exception: + import traceback + + console.warn(f"ISR render failed for {path}: {traceback.format_exc()}") + return None + if result is None: + return None + revalidate = ( + result.revalidate + if result.revalidate is not None + else self.default_revalidate + ) + await self.cache.set( + key, + CachedPage( + html=result.html, + generated_at=time.time(), + revalidate=revalidate, + tags=result.tags, + ), + ) + return result.html + + def _spawn_revalidate(self, key: str, path: str) -> None: + """Schedule a background, single-flight revalidation of ``key``. + + Args: + key: The cache key. + path: The request path. + """ + task = asyncio.create_task(self._revalidate(key, path)) + self._background.add(task) + task.add_done_callback(self._background.discard) + + async def _revalidate(self, key: str, path: str) -> None: + """Re-render and refresh a stale entry if no other worker is doing so. + + Args: + key: The cache key. + path: The request path. + """ + if not await self.cache.acquire_lock(key, self.lock_ttl): + return + try: + await self._render_and_store(key, path) + finally: + await self.cache.release_lock(key) + + async def revalidate_path(self, path: str) -> None: + """Invalidate the cached page for a specific path. + + Args: + path: The path to invalidate. + """ + await self.cache.delete(cache_key(self.build_id, path)) + + async def revalidate_tag(self, tag: str) -> int: + """Invalidate every cached page carrying ``tag``. + + Args: + tag: The tag to invalidate. + + Returns: + The number of pages invalidated. + """ + return await self.cache.invalidate_tag(tag) + + +# --------------------------------------------------------------------------- +# HTTP render tier: the concrete Renderer that calls a separate render service. +# --------------------------------------------------------------------------- + + +class HttpRenderer: + """Renderer that delegates to a render-tier HTTP service. + + The service receives ``{"path": path}`` and returns + ``{"html": ..., "revalidate"?: number, "tags"?: string[]}``. Any error or + non-200 response yields ``None`` so ISR serves stale/shell instead of + caching a failure. + + Args: + render_url: The render-tier endpoint URL. + timeout: Per-request timeout in seconds. + """ + + def __init__(self, render_url: str, *, timeout: float = 10.0) -> None: + """Store the render endpoint and timeout.""" + self._url = render_url + self._timeout = timeout + # A single pooled client is created lazily on first use (inside the + # event loop) and reused, so page-server -> render-tier connections stay + # keep-alive instead of a fresh TCP/TLS handshake per render. + self._client: httpx.AsyncClient | None = None + + def _get_client(self) -> httpx.AsyncClient: + """Return the shared client, creating it on first use. + + Returns: + The pooled async HTTP client. + """ + import httpx + + if self._client is None: + self._client = httpx.AsyncClient(timeout=self._timeout) + return self._client + + async def aclose(self) -> None: + """Close the pooled client (call on page-server shutdown).""" + if self._client is not None: + await self._client.aclose() + self._client = None + + async def __call__(self, path: str) -> RenderResult | None: + """Render ``path`` via the render tier. + + Args: + path: The concrete request path. + + Returns: + The render result, or None on failure. + """ + try: + resp = await self._get_client().post(self._url, json={"path": path}) + if resp.status_code != 200: + console.warn(f"ISR render tier returned {resp.status_code} for {path}") + return None + data = resp.json() + except Exception as exc: + console.warn(f"ISR render tier request failed for {path}: {exc}") + return None + + html = data.get("html") + if not html: + return None + return RenderResult( + html=html, + revalidate=data.get("revalidate"), + tags=tuple(data.get("tags", ())), + ) + + +# --------------------------------------------------------------------------- +# Factories + the dedicated ISR page-server (sits behind nginx proxy_cache). +# --------------------------------------------------------------------------- + + +def is_enabled(config: Config | None = None) -> bool: + """Whether ISR is configured (a render tier URL is set). + + Args: + config: The config to read (defaults to the global config). + + Returns: + True when ``config.isr_render_url`` is set. + """ + from reflex.config import get_config + + config = config if config is not None else get_config() + return config.isr_render_url is not None + + +def get_build_id(config: Config | None = None) -> str: + """Resolve the cache-rotating build id. + + Precedence: ``config.isr_build_id`` -> ``REFLEX_BUILD_ID`` env -> ``"dev"``. + + Args: + config: The config to read (defaults to the global config). + + Returns: + The build id. + """ + from reflex.config import get_config + + config = config if config is not None else get_config() + return config.isr_build_id or os.environ.get("REFLEX_BUILD_ID") or "dev" + + +def get_cache() -> ISRCache: + """Build the ISR cache backend: Redis when available, else in-memory. + + Returns: + A shared Redis-backed cache, or an in-process fallback. + """ + from reflex.utils import prerequisites + + redis = prerequisites.get_redis() + if redis is not None: + return RedisISRCache(redis) + console.warn( + "ISR: no redis_url configured; using an in-memory cache that is NOT " + "shared across workers. Set redis_url for multi-worker deployments." + ) + return MemoryISRCache() + + +async def revalidate_path(path: str, *, config: Config | None = None) -> None: + """Invalidate the ISR cache entry for ``path`` from application code. + + Call this from an event handler after content changes so the next request + re-renders the page. It operates on the shared Redis cache, so a single + call propagates to every page-server worker. No-op when no ``redis_url`` + is configured (nothing is shared to invalidate). + + Args: + path: The path to invalidate (e.g. ``/blog/hello-world``). + config: The config to read the build id from (defaults to global). + """ + from reflex.utils import prerequisites + + redis = prerequisites.get_redis() + if redis is not None: + await RedisISRCache(redis).delete(cache_key(get_build_id(config), path)) + + +async def revalidate_tag(tag: str) -> int: + """Invalidate every ISR page carrying ``tag`` from application code. + + Propagates to all page-server workers via the shared Redis cache. + + Args: + tag: The revalidation tag (e.g. ``post-123``). + + Returns: + The number of pages invalidated (0 when no ``redis_url`` is configured). + """ + from reflex.utils import prerequisites + + redis = prerequisites.get_redis() + if redis is None: + return 0 + return await RedisISRCache(redis).invalidate_tag(tag) + + +def create_manager( + config: Config | None = None, + *, + cache: ISRCache | None = None, + renderer: Renderer | None = None, +) -> ISRManager: + """Assemble an :class:`ISRManager` from config (and optional overrides). + + Args: + config: The config to read (defaults to the global config). + cache: Cache backend override (defaults to :func:`get_cache`). + renderer: Renderer override (defaults to an :class:`HttpRenderer`). + + Returns: + A configured ISR manager. + + Raises: + ValueError: If no renderer is given and ``isr_render_url`` is unset. + """ + from reflex.config import get_config + + config = config if config is not None else get_config() + if renderer is None: + if not config.isr_render_url: + msg = "ISR requires config.isr_render_url (or an explicit renderer)." + raise ValueError(msg) + renderer = HttpRenderer(config.isr_render_url) + return ISRManager( + cache if cache is not None else get_cache(), + renderer, + build_id=get_build_id(config), + default_revalidate=float(config.isr_revalidate), + ) + + +def _load_shell() -> str | None: + """Read the SPA shell HTML from the compiled static build, if present. + + Returns: + The shell HTML, or None if it cannot be found. + """ + from reflex import constants + from reflex.utils import prerequisites + + static_dir = prerequisites.get_web_dir() / constants.Dirs.STATIC + for name in (constants.ReactRouter.SPA_FALLBACK, "index.html"): + candidate = static_dir / name + if candidate.exists(): + return candidate.read_text() + return None + + +def create_isr_app( + manager: ISRManager | None = None, + *, + shell: str | None = None, +) -> Starlette: + """Create the dedicated ISR page-server ASGI app. + + Serves documents via ISR (nginx ``proxy_cache`` is expected in front for + edge caching), falling back to the SPA shell on a miss/render failure, and + exposes ``POST /_isr/revalidate`` for on-demand invalidation. The + revalidation endpoint requires a matching ``X-Reflex-ISR-Token`` header when + the ``REFLEX_ISR_REVALIDATE_TOKEN`` env var is set. + + Run in production as an app factory, e.g. + ``granian --factory reflex.isr:create_isr_app``. + + Args: + manager: The ISR manager (defaults to one built from config). + shell: The SPA shell HTML (defaults to reading the compiled build). + + Returns: + A Starlette app. + """ + from starlette.applications import Starlette + from starlette.requests import Request + from starlette.responses import HTMLResponse, JSONResponse, Response + from starlette.routing import Route + + mgr = manager if manager is not None else create_manager() + shell_html = shell if shell is not None else _load_shell() + revalidate_token = os.environ.get("REFLEX_ISR_REVALIDATE_TOKEN") + + async def serve(request: Request) -> Response: + html = await mgr.get_html(request.url.path) + if html is not None: + return HTMLResponse(html) + if shell_html is not None: + # Miss/failure: serve the SPA shell so the client hydrates. + return HTMLResponse(shell_html) + return Response("Not found", status_code=404) + + async def revalidate(request: Request) -> Response: + if revalidate_token and ( + request.headers.get("x-reflex-isr-token") != revalidate_token + ): + return JSONResponse({"error": "unauthorized"}, status_code=401) + body = await request.json() + if path := body.get("path"): + await mgr.revalidate_path(path) + count = 0 + if tag := body.get("tag"): + count = await mgr.revalidate_tag(tag) + return JSONResponse({"revalidated": True, "tag_pages": count}) + + @contextlib.asynccontextmanager + async def _lifespan(_app: Starlette): + """Close the renderer's pooled HTTP client on shutdown, if any.""" + try: + yield + finally: + aclose = getattr(mgr.renderer, "aclose", None) + if aclose is not None: + await aclose() + + return Starlette( + routes=[ + Route("/_isr/revalidate", revalidate, methods=["POST"]), + Route("/{path:path}", serve, methods=["GET"]), + ], + lifespan=_lifespan, + ) diff --git a/reflex/route.py b/reflex/route.py index 4b450fb6e25..567a902be71 100644 --- a/reflex/route.py +++ b/reflex/route.py @@ -96,6 +96,39 @@ def _add_route_arg(arg_name: str, type_: str): return args +def extract_route_params(path: str, route: str) -> dict[str, str]: + """Extract dynamic route parameter values from a concrete path. + + Given a concrete path (e.g. "/blog/hello-world") and a route pattern + (e.g. "blog/[slug]"), extract the parameter values by positional matching. + + Args: + path: The concrete URL path (e.g. "/blog/hello-world"). + route: The route pattern with brackets (e.g. "blog/[slug]"). + + Returns: + Dict mapping parameter names to their values, + e.g. {"slug": "hello-world"}. + """ + route_args = get_route_args(route) + if not route_args: + return {} + + params: dict[str, str] = {} + route_parts = route.strip("/").split("/") + path_parts = path.strip("/").split("/") + for i, route_part in enumerate(route_parts): + if i < len(path_parts): + arg_match = re.match( + r"^\[{1,2}(?:\.\.\.)?([a-zA-Z_]\w*)\]{1,2}$", route_part + ) + if arg_match: + param_name = arg_match.group(1) + if param_name in route_args: + params[param_name] = path_parts[i] + return params + + def replace_brackets_with_keywords(input_string: str) -> str: """Replace brackets and everything inside it in a string with a keyword. diff --git a/reflex/ssr.py b/reflex/ssr.py new file mode 100644 index 00000000000..149e961352b --- /dev/null +++ b/reflex/ssr.py @@ -0,0 +1,181 @@ +"""Server-side rendering (SSR) backend for Reflex apps. + +This is the request-time half of SSR: the ``/_ssr_data`` endpoint that computes +the hydrated initial state for a route by reusing the regular event machinery. +It lives in the top-level ``reflex`` package because it depends on the state and +app machinery that is not part of ``reflex_base``. + +The config gate (:func:`is_enabled`) and the compiler JS snippets live in +:mod:`reflex_base.ssr` so the compiler templates can use them without importing +``reflex``. ``is_enabled`` is re-exported here for convenience. + +Everything here is a no-op unless ``config.ssr_mode`` is not ``OFF``. +""" + +from __future__ import annotations + +import asyncio +import inspect +import traceback +from typing import TYPE_CHECKING, Any + +from reflex_base import constants +from reflex_base.ssr import is_enabled as is_enabled + +from reflex.utils import console + +if TYPE_CHECKING: + from starlette.requests import Request + from starlette.responses import Response + + from reflex.app import App + from reflex.event import Event + from reflex.state import BaseState + +# Sentinel token/session id used for the stateless SSR render. +SSR_TOKEN = "__ssr__" + + +def _build_router_data( + app: App, path: str, headers: dict[str, str], client_ip: str +) -> dict[str, Any]: + """Build the ``router_data`` dict for a stateless SSR render. + + Args: + app: The app, used to resolve the concrete path to a route pattern. + path: The concrete request path (e.g. ``/blog/hello-world``). + headers: The forwarded request headers. + client_ip: The client IP address. + + Returns: + A ``router_data`` dict with the same shape ``process()`` produces. + """ + from reflex.route import extract_route_params + + resolved_route = app.router(path) or "404" + params = extract_route_params(path, resolved_route) + return { + constants.RouteVar.PATH: "/" + resolved_route.removeprefix("/"), + constants.RouteVar.ORIGIN: path, + constants.RouteVar.QUERY: dict(params), + constants.RouteVar.CLIENT_TOKEN: SSR_TOKEN, + constants.RouteVar.SESSION_ID: SSR_TOKEN, + constants.RouteVar.HEADERS: { + "origin": headers.get("origin", headers.get("host", "http://localhost")), + **headers, + }, + constants.RouteVar.CLIENT_IP: client_ip, + } + + +async def _run_on_load_event(state: BaseState, event: Event, path: str) -> None: + """Run a single resolved on_load event on the ephemeral SSR state. + + Resolves the handler's target substate from ``event.name`` and invokes it + directly (rather than going through the session-managed event queue, which + needs a real client token and state manager). Handlers mutate the state in + place; their return value is consumed but discarded since the whole tree is + serialized later. + + Args: + state: The ephemeral root state instance. + event: The resolved on_load event (``event.name`` is the dotted handler). + path: The URL path (for error logging). + """ + try: + # e.g. "reflex___state____state.blog_state.on_load" -> substate + method. + *substate_path, method = event.name.split(".") + substate = state.get_substate(substate_path) + handler = substate.event_handlers[method] + + result = handler.fn(substate) + if asyncio.iscoroutine(result): + result = await result + if inspect.isgenerator(result): + for _ in result: + pass + elif inspect.isasyncgen(result): + async for _ in result: + pass + except Exception: + console.warn(f"SSR on_load handler failed for {path}: {traceback.format_exc()}") + + +async def _run_on_load_events(app: App, state: BaseState, path: str) -> None: + """Run the route's on_load handlers on the ephemeral SSR state. + + Args: + app: The app to get load events from. + state: The ephemeral root state instance. + path: The URL path (for error logging). + """ + from reflex.event import Event + + load_events = app.get_load_events(path) + if not load_events: + return + for event in Event.from_event_type(load_events, router_data=state.router_data): + await _run_on_load_event(state, event, path) + + +def ssr_data(app: App): + """Build the ``/_ssr_data`` endpoint handler. + + The handler creates an ephemeral state, applies route data, runs on_load + handlers, and returns the serialized state tree for server-side rendering. + + Args: + app: The app to get SSR data for. + + Returns: + The SSR data request handler. + """ + from starlette.responses import Response + + from reflex.state import RouterData, State + from reflex.utils import format + + async def ssr_data_handler(request: Request) -> Response: + """Handle an SSR data request. + + Args: + request: The Starlette request object. + + Returns: + Response with the serialized state as JSON. + """ + body = await request.json() + path = body.get("path", "/") + headers = body.get("headers", {}) + + if not app._state: + return Response( + content='{"state": null}', + media_type="application/json", + ) + + # Ephemeral root state — no persistent session is created. Use State + # (root) rather than app._state which may be a subclass whose inherited + # vars can't be set without a parent. + state = State(_reflex_internal_init=True) # pyright: ignore[reportCallIssue] + + router_data = _build_router_data( + app, + path, + headers, + request.client.host if request.client else "0.0.0.0", + ) + # Assigning router_data recomputes dependent DynamicRouteVars. + state.router_data = router_data + state.router = RouterData.from_router_data(router_data) + + await _run_on_load_events(app, state, path) + + json_str = format.json_dumps({"state": state.dict()}) + return Response( + content=json_str, + media_type="application/json", + headers={"Cache-Control": "no-cache"}, + ) + + return ssr_data_handler diff --git a/reflex/utils/build.py b/reflex/utils/build.py index dec54c075bb..3ed0bfb42f8 100644 --- a/reflex/utils/build.py +++ b/reflex/utils/build.py @@ -147,10 +147,11 @@ def zip_app( } if frontend: + web_dir = prerequisites.get_web_dir() _zip( component_name=constants.ComponentName.FRONTEND, target=zip_dest_dir / constants.ComponentName.FRONTEND.zip(), - root_directory=prerequisites.get_web_dir() / constants.Dirs.STATIC, + root_directory=web_dir / constants.Dirs.STATIC, files_to_exclude=files_to_exclude, exclude_venv_directories=False, ) diff --git a/reflex/utils/frontend_skeleton.py b/reflex/utils/frontend_skeleton.py index de9f29e7bbd..8c71d16759f 100644 --- a/reflex/utils/frontend_skeleton.py +++ b/reflex/utils/frontend_skeleton.py @@ -472,12 +472,14 @@ def update_react_router_config(prerender_routes: bool = False): def _update_react_router_config(config: Config, prerender_routes: bool = False): + from reflex_base import ssr + react_router_config = { "basename": config.prepend_frontend_path("/"), "future": { "unstable_optimizeDeps": True, }, - "ssr": False, + "ssr": ssr.ssr_build_enabled(), } if prerender_routes: diff --git a/tests/units/test_isr.py b/tests/units/test_isr.py new file mode 100644 index 00000000000..ab698264ebf --- /dev/null +++ b/tests/units/test_isr.py @@ -0,0 +1,194 @@ +"""Unit tests for the ISR (Incremental Static Regeneration) core.""" + +from __future__ import annotations + +import asyncio +import time + +import pytest + +from reflex.isr import CachedPage, ISRManager, MemoryISRCache, RenderResult, cache_key + + +class FakeRenderer: + """A controllable renderer that counts calls and can block on demand.""" + + def __init__( + self, + html: str = "rendered", + *, + revalidate: float | None = None, + tags: tuple[str, ...] = (), + gate: asyncio.Event | None = None, + fail: bool = False, + ) -> None: + """Configure the fake renderer.""" + self.html = html + self.revalidate = revalidate + self.tags = tags + self.gate = gate + self.fail = fail + self.calls = 0 + + async def __call__(self, path: str) -> RenderResult | None: + """Render the path, optionally blocking on the gate or failing. + + Args: + path: The path being rendered. + + Returns: + A render result, or None when configured to fail. + """ + self.calls += 1 + if self.gate is not None: + await self.gate.wait() + if self.fail: + return None + return RenderResult( + html=f"{self.html}:{path}:{self.calls}", + revalidate=self.revalidate, + tags=self.tags, + ) + + +async def _drain_background(manager: ISRManager) -> None: + """Wait for any background revalidation tasks to finish.""" + if manager._background: + await asyncio.gather(*list(manager._background)) + + +def _manager(renderer, **kwargs) -> ISRManager: + kwargs.setdefault("build_id", "b1") + kwargs.setdefault("default_revalidate", 60.0) + return ISRManager(MemoryISRCache(), renderer, **kwargs) + + +def test_cached_page_staleness(): + """A page is stale once past its revalidate window; 0 means never.""" + now = time.time() + fresh = CachedPage(html="x", generated_at=now, revalidate=60) + assert not fresh.is_stale(now + 30) + assert fresh.is_stale(now + 61) + never = CachedPage(html="x", generated_at=now, revalidate=0) + assert not never.is_stale(now + 10_000) + + +def test_cached_page_roundtrip(): + """CachedPage survives JSON serialization.""" + page = CachedPage(html="

", generated_at=1.0, revalidate=30, tags=("a", "b")) + assert CachedPage.from_json(page.to_json()) == page + + +@pytest.mark.asyncio +async def test_miss_renders_once_and_caches(): + """A cache miss renders and caches; a second request serves from cache.""" + renderer = FakeRenderer() + mgr = _manager(renderer) + + html1 = await mgr.get_html("/page") + html2 = await mgr.get_html("/page") + + assert html1 == html2 + assert renderer.calls == 1 # second request hit the cache + + +@pytest.mark.asyncio +async def test_stale_serves_stale_then_revalidates(): + """A stale entry is served immediately and refreshed in the background.""" + renderer = FakeRenderer() + mgr = _manager(renderer, default_revalidate=0.05) + + first = await mgr.get_html("/p") + assert renderer.calls == 1 + + # Let it go stale. + await asyncio.sleep(0.06) + + stale = await mgr.get_html("/p") + assert stale == first # served the stale copy immediately (no new render yet) + + await _drain_background(mgr) + assert renderer.calls == 2 # background revalidation happened exactly once + + # The refreshed value is now cached. + fresh = await mgr.get_html("/p") + assert fresh != first + assert renderer.calls == 2 + + +@pytest.mark.asyncio +async def test_single_flight_on_concurrent_miss(): + """Concurrent misses for one path trigger exactly one render.""" + gate = asyncio.Event() + renderer = FakeRenderer(gate=gate) + mgr = _manager(renderer, wait_timeout=5.0) + + tasks = [asyncio.create_task(mgr.get_html("/hot")) for _ in range(10)] + await asyncio.sleep(0.1) # let them all reach the lock / poll loop + gate.set() # release the single render + results = await asyncio.gather(*tasks) + + assert renderer.calls == 1 # only one worker rendered + assert len(set(results)) == 1 # everyone got the same HTML + assert results[0] is not None + + +@pytest.mark.asyncio +async def test_revalidate_path_forces_rerender(): + """Invalidating a path drops the cache so the next request re-renders.""" + renderer = FakeRenderer() + mgr = _manager(renderer) + + await mgr.get_html("/p") + assert renderer.calls == 1 + + await mgr.revalidate_path("/p") + await mgr.get_html("/p") + assert renderer.calls == 2 + + +@pytest.mark.asyncio +async def test_revalidate_tag_invalidates_all_matching(): + """Invalidating a tag drops every page carrying it.""" + renderer = FakeRenderer(tags=("blog",)) + mgr = _manager(renderer) + + await mgr.get_html("/blog/a") + await mgr.get_html("/blog/b") + assert renderer.calls == 2 + + count = await mgr.revalidate_tag("blog") + assert count == 2 + + await mgr.get_html("/blog/a") + await mgr.get_html("/blog/b") + assert renderer.calls == 4 # both re-rendered after invalidation + + +@pytest.mark.asyncio +async def test_build_id_scopes_cache(): + """A different build id yields a different key, forcing a re-render.""" + renderer = FakeRenderer() + cache = MemoryISRCache() + + mgr_old = ISRManager(cache, renderer, build_id="old") + mgr_new = ISRManager(cache, renderer, build_id="new") + + await mgr_old.get_html("/p") + await mgr_new.get_html("/p") + + assert renderer.calls == 2 + assert cache_key("old", "/p") != cache_key("new", "/p") + + +@pytest.mark.asyncio +async def test_render_failure_is_not_cached(): + """When the renderer returns None, nothing is cached and None is returned.""" + renderer = FakeRenderer(fail=True) + mgr = _manager(renderer) + + result = await mgr.get_html("/p") + assert result is None + # A subsequent request retries (nothing was cached). + await mgr.get_html("/p") + assert renderer.calls == 2 diff --git a/tests/units/test_isr_app.py b/tests/units/test_isr_app.py new file mode 100644 index 00000000000..fe8dc637af7 --- /dev/null +++ b/tests/units/test_isr_app.py @@ -0,0 +1,138 @@ +"""Tests for the ISR page-server app and factory helpers.""" + +from __future__ import annotations + +import reflex as rx +from reflex.isr import ( + ISRManager, + MemoryISRCache, + RenderResult, + create_isr_app, + get_build_id, + is_enabled, +) + + +class _Renderer: + """Minimal renderer that echoes the path and counts calls.""" + + def __init__(self) -> None: + """Initialize the call counter.""" + self.calls = 0 + + async def __call__(self, path: str) -> RenderResult: + """Render the path. + + Args: + path: The path being rendered. + + Returns: + A render result echoing the path. + """ + self.calls += 1 + return RenderResult(html=f"{path}#{self.calls}", tags=("t",)) + + +def _app(renderer: _Renderer, *, shell: str | None = "shell"): + manager = ISRManager( + MemoryISRCache(), renderer, build_id="b", default_revalidate=60 + ) + return create_isr_app(manager, shell=shell), manager + + +def _client(app): + from starlette.testclient import TestClient + + return TestClient(app) + + +def test_serves_and_caches_rendered_html(): + """The page-server renders once and then serves from cache.""" + renderer = _Renderer() + app, _ = _app(renderer) + client = _client(app) + + r1 = client.get("/blog/hello") + r2 = client.get("/blog/hello") + + assert r1.status_code == 200 + assert "/blog/hello#1" in r1.text + assert r1.text == r2.text + assert renderer.calls == 1 # second request hit the cache + + +def test_falls_back_to_shell_on_render_failure(): + """When the renderer yields no HTML, the SPA shell is served.""" + + class _NullRenderer: + async def __call__(self, path: str) -> None: + return None + + manager = ISRManager(MemoryISRCache(), _NullRenderer(), build_id="b") + app = create_isr_app(manager, shell="SHELL") + client = _client(app) + + r = client.get("/anything") + assert r.status_code == 200 + assert "SHELL" in r.text + + +def test_revalidate_path_endpoint(): + """POST /_isr/revalidate with a path drops that page from cache.""" + renderer = _Renderer() + app, _ = _app(renderer) + client = _client(app) + + client.get("/p") + assert renderer.calls == 1 + + resp = client.post("/_isr/revalidate", json={"path": "/p"}) + assert resp.status_code == 200 + assert resp.json()["revalidated"] is True + + client.get("/p") + assert renderer.calls == 2 # re-rendered after invalidation + + +def test_revalidate_tag_endpoint(): + """POST /_isr/revalidate with a tag invalidates all tagged pages.""" + renderer = _Renderer() + app, _ = _app(renderer) + client = _client(app) + + client.get("/a") + client.get("/b") + assert renderer.calls == 2 + + resp = client.post("/_isr/revalidate", json={"tag": "t"}) + assert resp.json()["tag_pages"] == 2 + + client.get("/a") + client.get("/b") + assert renderer.calls == 4 + + +def test_revalidate_requires_token_when_configured(monkeypatch): + """When REFLEX_ISR_REVALIDATE_TOKEN is set, the endpoint enforces it.""" + monkeypatch.setenv("REFLEX_ISR_REVALIDATE_TOKEN", "secret") + renderer = _Renderer() + app, _ = _app(renderer) + client = _client(app) + + assert client.post("/_isr/revalidate", json={"path": "/p"}).status_code == 401 + ok = client.post( + "/_isr/revalidate", + json={"path": "/p"}, + headers={"X-Reflex-ISR-Token": "secret"}, + ) + assert ok.status_code == 200 + + +def test_config_helpers_reflect_isr_settings(): + """is_enabled / get_build_id read from the config.""" + off = rx.Config(app_name="t") + assert is_enabled(off) is False + + on = rx.Config(app_name="t", isr_render_url="http://render:8000", isr_build_id="v9") + assert is_enabled(on) is True + assert get_build_id(on) == "v9" diff --git a/tests/units/test_ssr_compile.py b/tests/units/test_ssr_compile.py new file mode 100644 index 00000000000..d44e631a6f4 --- /dev/null +++ b/tests/units/test_ssr_compile.py @@ -0,0 +1,49 @@ +"""Unit tests for SSR route-param extraction.""" + +from __future__ import annotations + +from reflex.route import extract_route_params + + +class TestExtractRouteParams: + """Tests for the extract_route_params utility function.""" + + def test_simple_dynamic_route(self): + """Single dynamic segment is extracted correctly.""" + result = extract_route_params("/blog/hello-world", "blog/[slug]") + assert result == {"slug": "hello-world"} + + def test_multiple_dynamic_segments(self): + """Multiple dynamic segments are extracted correctly.""" + result = extract_route_params("/users/42/posts/99", "users/[id]/posts/[pid]") + assert result == {"id": "42", "pid": "99"} + + def test_no_dynamic_segments(self): + """Static route returns empty dict.""" + result = extract_route_params("/about", "about") + assert result == {} + + def test_root_path(self): + """Root path with no segments returns empty dict.""" + result = extract_route_params("/", "/") + assert result == {} + + def test_leading_slash_handling(self): + """Leading slashes on both path and route are handled.""" + result = extract_route_params("/blog/my-post", "/blog/[slug]") + assert result == {"slug": "my-post"} + + def test_optional_segment(self): + """Optional dynamic segment ([[param]]) is extracted.""" + result = extract_route_params("/docs/intro", "docs/[[section]]") + assert result == {"section": "intro"} + + def test_no_match_shorter_path(self): + """When path has fewer segments than route, missing params are skipped.""" + result = extract_route_params("/blog", "blog/[slug]") + assert result == {} + + def test_preserves_special_characters_in_value(self): + """Values with hyphens and other URL-safe chars are preserved.""" + result = extract_route_params("/blog/my-great-post-2024", "blog/[slug]") + assert result == {"slug": "my-great-post-2024"} diff --git a/tests/units/test_ssr_data.py b/tests/units/test_ssr_data.py new file mode 100644 index 00000000000..c5d87fd076f --- /dev/null +++ b/tests/units/test_ssr_data.py @@ -0,0 +1,317 @@ +"""Unit tests for the /_ssr_data endpoint handler.""" + +from __future__ import annotations + +import json +from typing import Any +from unittest.mock import AsyncMock, Mock + +import pytest +from starlette.responses import Response + +import reflex as rx +from reflex.app import App +from reflex.ssr import ssr_data + + +@pytest.fixture(autouse=True) +def _isolate_ssr_state(clean_registration_context): + """Serialize the /_ssr_data tree against only each test's own states. + + The endpoint calls ``State.dict()``, which walks every registered substate; + with a polluted global registry it would serialize (and crash on) unrelated + states from other modules (e.g. test_app's DynamicState). ``clean_registration_context`` + gives an empty substate registry, but the root ``State``'s per-class dependency + sets are global — reset them to match the clean context (and restore after), + so lookups stay consistent and nothing leaks in or out. + + Args: + clean_registration_context: A fresh, empty registration context. + + Yields: + None. + """ + from reflex.state import State + + names = ( + "_potentially_dirty_states", + "_always_dirty_substates", + "_var_dependencies", + ) + original = {name: getattr(State, name) for name in names} + State._potentially_dirty_states = set() + State._always_dirty_substates = set() + State._var_dependencies = {} + State.get_class_substate.cache_clear() + try: + yield + finally: + for name, value in original.items(): + setattr(State, name, value) + State.get_class_substate.cache_clear() + + +def _make_request(path: str = "/", headers: dict | None = None) -> Mock: + """Create a mock Starlette Request with the given path and headers. + + Args: + path: The URL path to include in the request body. + headers: Optional headers dict to include in the request body. + + Returns: + A mock Request object. + """ + body = {"path": path, "headers": headers or {}} + request = Mock() + request.json = AsyncMock(return_value=body) + request.client = Mock() + request.client.host = "127.0.0.1" + return request + + +def _parse_response(response: Response) -> dict[str, Any]: + """Parse a Starlette Response body as JSON. + + Args: + response: The response to parse. + + Returns: + Parsed JSON as a dict. + """ + assert isinstance(response.body, bytes) + return json.loads(response.body) + + +@pytest.mark.asyncio +async def test_ssr_data_no_state(): + """When the app has no state, the endpoint returns null state.""" + app = App(enable_state=False) + app.add_page(lambda: rx.text("hello"), route="/") + handler = ssr_data(app) + + response = await handler(_make_request("/")) + + assert response.status_code == 200 + data = _parse_response(response) + assert data["state"] is None + + +@pytest.mark.asyncio +async def test_ssr_data_basic_state(): + """The endpoint returns serialized state for a basic stateful app.""" + + class BasicState(rx.State): + title: str = "default" + + app = App() + app._state = BasicState + app.add_page(lambda: rx.text(BasicState.title), route="/") + + handler = ssr_data(app) + response = await handler(_make_request("/")) + + assert response.status_code == 200 + data = _parse_response(response) + assert data["state"] is not None + + # The root State should be present. + root_name = rx.State.get_full_name() + assert root_name in data["state"] + + # The user's substate should also be present with the default value. + substate_name = BasicState.get_full_name() + assert substate_name in data["state"] + assert data["state"][substate_name]["title_rx_state_"] == "default" + + +@pytest.mark.asyncio +async def test_ssr_data_dynamic_route_params(): + """Route params are extracted from the URL path and set on the state.""" + + class PostState(rx.State): + pass + + app = App() + app._state = PostState + app.add_page(lambda: rx.text("post"), route="/blog/[slug]") + + handler = ssr_data(app) + response = await handler(_make_request("/blog/hello-world")) + + assert response.status_code == 200 + data = _parse_response(response) + root_name = rx.State.get_full_name() + router = data["state"][root_name]["router_rx_state_"] + assert router["page"]["params"] == {"slug": "hello-world"} + assert router["page"]["raw_path"] == "/blog/hello-world" + + +@pytest.mark.asyncio +async def test_ssr_data_on_load_runs(): + """The on_load handler runs and mutates state before serialization.""" + + class LoadState(rx.State): + title: str = "" + + @rx.event + def on_load_post(self): + self.title = "loaded" + + app = App() + app._state = LoadState + app.add_page( + lambda: rx.text(LoadState.title), + route="/page", + on_load=LoadState.on_load_post, + ) + + handler = ssr_data(app) + response = await handler(_make_request("/page")) + + assert response.status_code == 200 + data = _parse_response(response) + substate_name = LoadState.get_full_name() + assert data["state"][substate_name]["title_rx_state_"] == "loaded" + + +@pytest.mark.asyncio +async def test_ssr_data_async_on_load(): + """An async on_load handler is properly awaited.""" + + class AsyncLoadState(rx.State): + message: str = "" + + @rx.event + async def load_data(self): + self.message = "async-loaded" + + app = App() + app._state = AsyncLoadState + app.add_page( + lambda: rx.text(AsyncLoadState.message), + route="/async", + on_load=AsyncLoadState.load_data, + ) + + handler = ssr_data(app) + response = await handler(_make_request("/async")) + + assert response.status_code == 200 + data = _parse_response(response) + substate_name = AsyncLoadState.get_full_name() + assert data["state"][substate_name]["message_rx_state_"] == "async-loaded" + + +@pytest.mark.asyncio +async def test_ssr_data_on_load_error_graceful(): + """If on_load raises, the endpoint returns state with defaults (no crash).""" + + class ErrorState(rx.State): + value: str = "untouched" + + @rx.event + def bad_handler(self): + msg = "boom" + raise RuntimeError(msg) + + app = App() + app._state = ErrorState + app.add_page( + lambda: rx.text(ErrorState.value), + route="/error", + on_load=ErrorState.bad_handler, + ) + + handler = ssr_data(app) + response = await handler(_make_request("/error")) + + assert response.status_code == 200 + data = _parse_response(response) + substate_name = ErrorState.get_full_name() + # State should still be returned with original defaults. + assert data["state"][substate_name]["value_rx_state_"] == "untouched" + + +@pytest.mark.asyncio +async def test_ssr_data_headers_forwarded(): + """Request headers are set on the state's router headers.""" + + class HeaderState(rx.State): + pass + + app = App() + app._state = HeaderState + app.add_page(lambda: rx.text("h"), route="/") + + handler = ssr_data(app) + response = await handler( + _make_request( + "/", headers={"user-agent": "Googlebot", "origin": "https://example.com"} + ) + ) + + assert response.status_code == 200 + data = _parse_response(response) + root_name = rx.State.get_full_name() + router = data["state"][root_name]["router_rx_state_"] + assert router["headers"]["user_agent"] == "Googlebot" + + +@pytest.mark.asyncio +async def test_ssr_data_unknown_route(): + """An unknown path resolves to the 404 route.""" + + class NotFoundState(rx.State): + pass + + app = App() + app._state = NotFoundState + app.add_page(lambda: rx.text("home"), route="/") + + handler = ssr_data(app) + response = await handler(_make_request("/this/does/not/exist")) + + assert response.status_code == 200 + data = _parse_response(response) + # Should still return valid state (the 404 handler path). + assert data["state"] is not None + + +@pytest.mark.asyncio +async def test_ssr_data_cache_control_header(): + """The response includes Cache-Control: no-cache.""" + + class CacheState(rx.State): + pass + + app = App() + app._state = CacheState + app.add_page(lambda: rx.text("c"), route="/") + + handler = ssr_data(app) + response = await handler(_make_request("/")) + + assert response.headers["cache-control"] == "no-cache" + + +@pytest.mark.asyncio +async def test_ssr_data_client_ip(): + """The client IP from the request is set in the state.""" + + class IpState(rx.State): + pass + + app = App() + app._state = IpState + app.add_page(lambda: rx.text("ip"), route="/") + + handler = ssr_data(app) + request = _make_request("/") + request.client.host = "10.0.0.42" + response = await handler(request) + + assert response.status_code == 200 + data = _parse_response(response) + root_name = rx.State.get_full_name() + router = data["state"][root_name]["router_rx_state_"] + assert router["session"]["client_ip"] == "10.0.0.42"