diff --git a/news/7068.breaking.md b/news/7068.breaking.md new file mode 100644 index 00000000000..d54b2d8e8c0 --- /dev/null +++ b/news/7068.breaking.md @@ -0,0 +1 @@ +The root state gained five base vars holding the router data: `rx_router_session`, `rx_router_headers`, `rx_router_page`, `rx_router_url` and `rx_router_route_id`. A substate that declares one of these names now raises `BaseVarShadowsInheritedVarError`, the same error any other shadowed inherited var raises, and must rename its field. `State.router` itself is unchanged. diff --git a/news/7068.deprecation.md b/news/7068.deprecation.md new file mode 100644 index 00000000000..cac4544b070 --- /dev/null +++ b/news/7068.deprecation.md @@ -0,0 +1 @@ +Declaring a computed var dependency on the `router` var (`deps=["router"]`) is deprecated; depend on the router Var instead, e.g. `deps=[State.router.url]` for a single field or `deps=[State.router]` to keep tracking all of them. diff --git a/news/7068.performance.md b/news/7068.performance.md new file mode 100644 index 00000000000..9908f7676a3 --- /dev/null +++ b/news/7068.performance.md @@ -0,0 +1,3 @@ +Store router data in separate base vars (session, headers, page, url, route_id) so a navigation delta only re-sends the fields that changed instead of the whole router, and gather the connection-scoped router data (headers, client IP, session id) once at connect time rather than on every event. `State.router` is unchanged for app code. + +The page URL is also persisted as the URL itself rather than as its parsed pieces: `ReflexURL` and `URLData` re-split on the way out of the state store instead of writing scheme, netloc, origin, path, query, query parameters and fragment alongside the href on every state write. diff --git a/packages/reflex-base/news/7068.feature.md b/packages/reflex-base/news/7068.feature.md new file mode 100644 index 00000000000..1240ddc60e4 --- /dev/null +++ b/packages/reflex-base/news/7068.feature.md @@ -0,0 +1 @@ +`VarData` now tracks every state field a var is built from in `field_dependencies`, a mapping of state name to that state's field names, unioned and deduped as vars merge. A computed var depending on a composite var (`deps=[SomeState.composite]`) is now invalidated when any of its underlying fields changes — including fields belonging to a different state, which previously went untracked. `state` and `field_name` still report the first state and its first field, so existing readers are unaffected; read `field_dependencies` when you need every state a var reaches. diff --git a/packages/reflex-base/news/7068.performance.md b/packages/reflex-base/news/7068.performance.md new file mode 100644 index 00000000000..41445cf7826 --- /dev/null +++ b/packages/reflex-base/news/7068.performance.md @@ -0,0 +1 @@ +The event processor now refreshes only the router vars whose backing `router_data` keys actually changed, so a navigation no longer rebuilds and re-sends the connection-scoped session and header data. `ROUTER_VARS` names the per-field router vars that replaced the single `router` var on the root state. diff --git a/packages/reflex-base/src/reflex_base/constants/__init__.py b/packages/reflex-base/src/reflex_base/constants/__init__.py index f83bcfd1e25..3f46a794461 100644 --- a/packages/reflex-base/src/reflex_base/constants/__init__.py +++ b/packages/reflex-base/src/reflex_base/constants/__init__.py @@ -60,6 +60,12 @@ ROUTER, ROUTER_DATA, ROUTER_DATA_INCLUDE, + ROUTER_HEADERS, + ROUTER_PAGE, + ROUTER_ROUTE_ID, + ROUTER_SESSION, + ROUTER_URL, + ROUTER_VARS, DefaultPage, Page404, RouteArgType, @@ -86,6 +92,12 @@ "ROUTER", "ROUTER_DATA", "ROUTER_DATA_INCLUDE", + "ROUTER_HEADERS", + "ROUTER_PAGE", + "ROUTER_ROUTE_ID", + "ROUTER_SESSION", + "ROUTER_URL", + "ROUTER_VARS", "ROUTE_NOT_FOUND", "SESSION_STORAGE", "SETTER_PREFIX", diff --git a/packages/reflex-base/src/reflex_base/constants/route.py b/packages/reflex-base/src/reflex_base/constants/route.py index 30e7b32170e..4c8606561eb 100644 --- a/packages/reflex-base/src/reflex_base/constants/route.py +++ b/packages/reflex-base/src/reflex_base/constants/route.py @@ -11,10 +11,30 @@ class RouteArgType(SimpleNamespace): LIST = "arg_list" -# the name of the backend var containing path and client information +# the name of the state attribute exposing path and client information ROUTER = "router" ROUTER_DATA = "router_data" +# The names of the per-field base vars holding router data on the root state. +# Session and headers are constant for the lifetime of a websocket connection, +# while page, url, and route_id change on every navigation; keeping them in +# separate vars means a navigation delta only re-sends the navigation fields. +# The `rx_` prefix keeps them from colliding with a field an app already +# defines; `router` itself stays unprefixed as the public switchboard. +ROUTER_SESSION = "rx_router_session" +ROUTER_HEADERS = "rx_router_headers" +ROUTER_PAGE = "rx_router_page" +ROUTER_URL = "rx_router_url" +ROUTER_ROUTE_ID = "rx_router_route_id" + +ROUTER_VARS = ( + ROUTER_SESSION, + ROUTER_HEADERS, + ROUTER_PAGE, + ROUTER_URL, + ROUTER_ROUTE_ID, +) + class RouteVar(SimpleNamespace): """Names of variables used in the router_data dict stored in State.""" diff --git a/packages/reflex-base/src/reflex_base/event/processor/base_state_processor.py b/packages/reflex-base/src/reflex_base/event/processor/base_state_processor.py index b9f1a982e72..0e60983b572 100644 --- a/packages/reflex-base/src/reflex_base/event/processor/base_state_processor.py +++ b/packages/reflex-base/src/reflex_base/event/processor/base_state_processor.py @@ -13,7 +13,6 @@ from time import perf_counter from typing import TYPE_CHECKING, Any -from reflex.istate.data import RouterData from reflex.istate.manager.token import BaseStateToken from reflex.istate.proxy import StateProxy from reflex.utils import types @@ -431,12 +430,20 @@ async def _execute_event( ) # re-assign only when the value is set and different - if router_data and state.router_data != router_data: - # assignment will recurse into substates and force recalculation of - # dependent ComputedVar (dynamic route variables) - state.router_data = router_data - if state.router != (router := RouterData.from_router_data(router_data)): - state.router = router + if ( + router_data + and (previous_router_data := state.router_data) != router_data + ): + # only the router vars whose backing keys changed are rebuilt + # and re-sent; session/headers stay put across navigations. + merged_router_data = state._update_router_vars( + router_data, previous_router_data + ) + # Store what it merged, not the payload, so a partial payload + # does not drop the keys it omits. Only on a real change: the + # assignment dirties router_data and marks the state touched. + if merged_router_data != previous_router_data: + state.router_data = merged_router_data # Preprocess the event. if ( diff --git a/packages/reflex-base/src/reflex_base/vars/base.py b/packages/reflex-base/src/reflex_base/vars/base.py index e425919a7b5..026ba4bb7db 100644 --- a/packages/reflex-base/src/reflex_base/vars/base.py +++ b/packages/reflex-base/src/reflex_base/vars/base.py @@ -278,6 +278,34 @@ def insert_app_wraps( target[key] = wrapper +def _normalize_field_dependencies( + field_dependencies: Mapping[str, Iterable[str]] | None, + state: str, + field_name: str, +) -> Mapping[str, tuple[str, ...]]: + """Build the canonical state -> fields mapping from the accepted shorthands. + + Args: + field_dependencies: The canonical mapping, if the caller gave one. + state: The single enclosing state, for the shorthand form. + field_name: A single field of `state`. + + Returns: + A mapping of state name to its deduped field names. + """ + if field_dependencies is not None: + return { + state_name: tuple(dict.fromkeys(names)) + for state_name, names in field_dependencies.items() + } + names = (field_name,) if field_name else () + # A state with no named field still has to be recorded: plenty of vars + # carry only the state (for imports and hooks) and nothing reads a field. + if not state and not names: + return {} + return {state: names} + + @dataclasses.dataclass( eq=True, frozen=True, @@ -285,11 +313,17 @@ def insert_app_wraps( class VarData: """Metadata associated with a x.""" - # The name of the enclosing state. - state: str = dataclasses.field(default="") - - # The name of the field in the state. - field_name: str = dataclasses.field(default="") + # Every state field this var is built from, grouped by the state that owns + # it. A var normally stands for a single field of a single state, but one + # composed of several -- possibly spanning several states -- names all of + # them, so a dependency on it tracks each field it actually reads. + # Built fresh for every VarData and never mutated afterwards, so it is + # effectively frozen like the tuples beside it. A plain dict rather than a + # MappingProxyType because VarData is pickled along with the states holding + # it, and mappingproxy cannot be pickled. + field_dependencies: Mapping[str, tuple[str, ...]] = dataclasses.field( + default_factory=dict + ) # Imports needed to render this var imports: ParsedImportTuple = dataclasses.field(default_factory=tuple) @@ -322,18 +356,26 @@ def __init__( position: Hooks.HookPosition | None = None, components: Iterable[BaseComponent] | None = None, app_wraps: Iterable[tuple[int, BaseComponent]] | None = None, + field_dependencies: Mapping[str, Iterable[str]] | None = None, ): """Initialize the var data. Args: - state: The name of the enclosing state. - field_name: The name of the field in the state. + state: The name of the enclosing state. Shorthand for a + single-state ``field_dependencies``; ignored when that is given. + field_name: The name of the field in ``state``. Ignored when + ``field_dependencies`` is given. imports: Imports needed to render this var. hooks: Hooks that need to be present in the component to render this var. deps: Dependencies of the var for useCallback. position: Position of the hook in the component. components: Components that are part of this var. app_wraps: App-level wrapper components this var requires when used. + field_dependencies: Every state field this var is built from, + grouped by owning state. The canonical form; ``state``, + and ``field_name`` are the shorthand for a single state with a + single field. Keyword-only in practice: it trails the older + parameters so positional callers of those are unaffected. """ if isinstance(hooks, str): hooks = [hooks] @@ -342,8 +384,11 @@ def __init__( immutable_imports: ParsedImportTuple = tuple( (k, tuple(v)) for k, v in parse_imports(imports or {}).items() ) - object.__setattr__(self, "state", state) - object.__setattr__(self, "field_name", field_name) + object.__setattr__( + self, + "field_dependencies", + _normalize_field_dependencies(field_dependencies, state, field_name), + ) object.__setattr__(self, "imports", immutable_imports) object.__setattr__(self, "hooks", tuple(hooks or {})) object.__setattr__(self, "deps", tuple(deps or [])) @@ -355,8 +400,11 @@ def __init__( # Merge our dependencies first, so they can be referenced. merged_var_data = VarData.merge(*hooks.values(), self) if merged_var_data is not None: - object.__setattr__(self, "state", merged_var_data.state) - object.__setattr__(self, "field_name", merged_var_data.field_name) + object.__setattr__( + self, + "field_dependencies", + merged_var_data.field_dependencies, + ) object.__setattr__(self, "imports", merged_var_data.imports) object.__setattr__(self, "hooks", merged_var_data.hooks) object.__setattr__(self, "deps", merged_var_data.deps) @@ -364,6 +412,31 @@ def __init__( object.__setattr__(self, "components", merged_var_data.components) object.__setattr__(self, "app_wraps", merged_var_data.app_wraps) + @property + def state(self) -> str: + """The name of the enclosing state. + + A var may be built from fields of more than one state; this reports + only the first. Read ``field_dependencies`` to see every state. + + Returns: + The first state name, or an empty string if there is none. + """ + return next(iter(self.field_dependencies), "") + + @property + def field_name(self) -> str: + """The name of the field in the state. + + A var built from several fields reports only the first, of the first + state. Read ``field_dependencies`` to see all of them. + + Returns: + The first field name, or an empty string if there is none. + """ + field_names = self.field_dependencies.get(self.state, ()) + return field_names[0] if field_names else "" + def old_school_imports(self) -> ImportDict: """Return the imports as a mutable dict. @@ -394,16 +467,20 @@ def merge(*all: VarData | None) -> VarData | None: if len(all_var_datas) == 1: return all_var_datas[0] - # Get the first non-empty field name or default to empty string. - field_name = next( - (var_data.field_name for var_data in all_var_datas if var_data.field_name), - "", - ) - - # Get the first non-empty state or default to empty string. - state = next( - (var_data.state for var_data in all_var_datas if var_data.state), "" - ) + # Union every state's fields, in order and deduped, so a var composed + # of several fields -- across as many states as it reaches -- carries + # all of them and a dependency on it tracks each one. Accumulated as + # ordered sets and materialized once: this runs for every var + # operation, so rebuilding a tuple per contributing var costs. + seen_fields: dict[str, dict[str, None]] = {} + for var_data in all_var_datas: + for state_name, names in var_data.field_dependencies.items(): + seen = seen_fields.get(state_name) + if seen is None: + seen_fields[state_name] = dict.fromkeys(names) + else: + for name in names: + seen[name] = None hooks: dict[str, VarData | None] = { hook: None for var_data in all_var_datas for hook in var_data.hooks @@ -445,8 +522,7 @@ def merge(*all: VarData | None) -> VarData | None: insert_app_wraps(app_wraps, var_data.app_wraps) return VarData( - state=state, - field_name=field_name, + field_dependencies=seen_fields, imports=imports_, hooks=hooks, deps=deps, @@ -464,10 +540,9 @@ def __bool__(self) -> bool: True if any field is set to a non-default value. """ return bool( - self.state + self.field_dependencies or self.imports or self.hooks - or self.field_name or self.deps or self.position or self.components @@ -493,8 +568,7 @@ def _identity_key(self) -> tuple: A hashable tuple uniquely identifying this VarData. """ return ( - self.state, - self.field_name, + tuple(self.field_dependencies.items()), self.imports, self.hooks, tuple(dep._hash_key() for dep in self.deps), @@ -787,6 +861,23 @@ def _get_all_var_data(self) -> VarData | None: """ return self._var_data + def _dependency_fields(self) -> Mapping[str, tuple[str, ...]]: + """The state fields a ComputedVar depending on this Var must track. + + A Var normally stands for a single field of a single state, but one + composed of several must name all of them, or a ``deps=[that_var]`` + dependency would track only some of the fields it reads and leave the + computed var stale when any of the others change. A composite var may + also span several states, so the fields stay grouped by their owner. + ``VarData.merge`` unions them as vars combine, so the merged VarData + already knows every one. + + Returns: + The fields to register the dependency against, by state name. + """ + all_var_data = self._get_all_var_data() + return all_var_data.field_dependencies if all_var_data is not None else {} + def __deepcopy__(self, memo: dict[int, Any]) -> Self: """Deepcopy the var. @@ -2507,16 +2598,17 @@ def _add_static_dep( if deps is None: deps = self._static_deps if isinstance(dep, Var): - state_name = ( - all_var_data.state - if (all_var_data := dep._get_all_var_data()) and all_var_data.state - else None - ) - if all_var_data is not None: - var_name = all_var_data.field_name + if (all_var_data := dep._get_all_var_data()) is not None: + # A composite Var names every state field it is built from, in + # each state that owns them. + field_dependencies = all_var_data.field_dependencies + if field_dependencies: + for state_name, field_names in field_dependencies.items(): + deps.setdefault(state_name or None, set()).update(field_names) + else: + deps.setdefault(None, set()) else: - var_name = dep._js_expr - deps.setdefault(state_name, set()).add(var_name) + deps.setdefault(None, set()).add(dep._js_expr) elif isinstance(dep, str) and dep != "": deps.setdefault(None, set()).add(dep) else: @@ -2837,24 +2929,30 @@ def add_dependency(self, objclass: type[BaseState], dep: Var): state and field name """ if all_var_data := dep._get_all_var_data(): - state_name = all_var_data.state - if state_name: - var_name = all_var_data.field_name - if var_name: - self._static_deps.setdefault(state_name, set()).add(var_name) - target_state_class = objclass.get_root_state().get_class_substate( - state_name - ) + # A composite Var names every state field it is built from, and may + # span several states; register against each of them. + registered = False + for state_name, field_names in all_var_data.field_dependencies.items(): + var_names = tuple(filter(None, field_names)) + if not state_name or not var_names: + continue + self._static_deps.setdefault(state_name, set()).update(var_names) + target_state_class = objclass.get_root_state().get_class_substate( + state_name + ) + for var_name in var_names: target_state_class._var_dependencies.setdefault( var_name, set() ).add(( objclass.get_full_name(), self._name, )) - target_state_class._potentially_dirty_states.add( - objclass.get_full_name() - ) - return + target_state_class._potentially_dirty_states.add( + objclass.get_full_name() + ) + registered = True + if registered: + return msg = ( "ComputedVar dependencies must be Var instances with a state and " f"field name, got {dep!r}." diff --git a/packages/reflex-base/src/reflex_base/vars/dep_tracking.py b/packages/reflex-base/src/reflex_base/vars/dep_tracking.py index 4d4ad5e8333..837badc9d8e 100644 --- a/packages/reflex-base/src/reflex_base/vars/dep_tracking.py +++ b/packages/reflex-base/src/reflex_base/vars/dep_tracking.py @@ -388,9 +388,8 @@ def handle_getting_var(self, instruction: dis.Instruction) -> None: if the_var_data is None: msg = f"Cannot determine the source code for the var in {self.func!r}." raise VarValueError(msg) - self.dependencies.setdefault(the_var_data.state, set()).add( - the_var_data.field_name - ) + for state_name, field_names in the_var._dependency_fields().items(): + self.dependencies.setdefault(state_name, set()).update(field_names) self.scan_status = ScanStatus.SCANNING def _populate_dependencies(self) -> None: diff --git a/reflex/app.py b/reflex/app.py index cf08b167e89..4e4b86c1dfe 100644 --- a/reflex/app.py +++ b/reflex/app.py @@ -76,7 +76,7 @@ from reflex.app_mixins import AppMixin, LifespanMixin, MiddlewareMixin from reflex.compiler import compiler from reflex.compiler.compiler import readable_name_from_component -from reflex.istate.data import RouterData +from reflex.istate.data import SessionData from reflex.istate.manager import StateManager, StateModificationContext from reflex.istate.manager.token import BaseStateToken from reflex.route import ( @@ -2093,6 +2093,10 @@ def __init__(self, namespace: str, app: App): # Number of client_error reports logged per SID, for rate limiting. self._client_error_counts: dict[str, int] = {} + # Connection-scoped router_data entries per SID, computed once at + # connect time instead of for every event on the connection. + self._static_router_data: dict[str, dict[str, Any]] = {} + # Start time and count of the current process-wide client_error window. self._client_error_window_start = 0.0 self._client_error_window_count = 0 @@ -2144,6 +2148,51 @@ async def on_connect(self, sid: str, environ: dict): if otel.enabled: otel.record_connection(1) + # Headers, client IP, and session id cannot change for the lifetime of + # the connection; compute them once instead of on every event. + self._static_router_data[sid] = self._build_static_router_data(sid, environ) + + def _build_static_router_data(self, sid: str, environ: dict) -> dict[str, Any]: + """Build the connection-scoped router_data entries for a socket. + + Args: + sid: The Socket.IO session id. + environ: The request information, including HTTP headers. + + Returns: + The router_data entries that are constant for the connection. + """ + asgi_scope = environ.get("asgi.scope", {}) + + # Get the client headers. + headers = { + k.decode("utf-8"): v.decode("utf-8") + for (k, v) in asgi_scope.get("headers", []) + } + + # Get the client IP + try: + client_ip = asgi_scope["client"][0] + headers["asgi-scope-client"] = client_ip + except (KeyError, IndexError): + client_ip = environ.get("REMOTE_ADDR", "0.0.0.0") + + # Unroll reverse proxy forwarded headers. + client_ip = ( + headers + .get( + "x-forwarded-for", + client_ip, + ) + .partition(",")[0] + .strip() + ) + return { + constants.RouteVar.SESSION_ID: sid, + constants.RouteVar.HEADERS: headers, + constants.RouteVar.CLIENT_IP: client_ip, + } + def on_disconnect(self, sid: str) -> asyncio.Task | None: """Event for when the websocket disconnects. @@ -2156,6 +2205,7 @@ def on_disconnect(self, sid: str) -> asyncio.Task | None: if otel.enabled: otel.record_connection(-1) self._client_error_counts.pop(sid, None) + self._static_router_data.pop(sid, None) # Get token before cleaning up disconnect_token = self.sid_to_token.get(sid) if disconnect_token: @@ -2248,45 +2298,33 @@ async def on_event(self, sid: str, data: Any): msg = f"Failed to deserialize event data: {fields}." raise exceptions.EventDeserializationError(msg) from ex - # Get the event environment. - if self.app.sio is None: - msg = "Socket.IO is not initialized." - raise RuntimeError(msg) - environ = self.app.sio.get_environ(sid, self.namespace) - if environ is None: - msg = "Socket.IO environ is not initialized." - raise RuntimeError(msg) - - # Get the client headers. - headers = { - k.decode("utf-8"): v.decode("utf-8") - for (k, v) in environ["asgi.scope"]["headers"] - } - - # Get the client IP - try: - client_ip = environ["asgi.scope"]["client"][0] - headers["asgi-scope-client"] = client_ip - except (KeyError, IndexError): - client_ip = environ.get("REMOTE_ADDR", "0.0.0.0") - - # Unroll reverse proxy forwarded headers. - client_ip = ( - headers - .get( - "x-forwarded-for", - client_ip, + static_router_data = self._static_router_data.get(sid) + if static_router_data is None: + # The connection was not seen by on_connect (e.g. namespace created + # after the socket connected); fall back to the connection environ. + if self.app.sio is None: + msg = "Socket.IO is not initialized." + raise RuntimeError(msg) + environ = self.app.sio.get_environ(sid, self.namespace) + if environ is None: + msg = "Socket.IO environ is not initialized." + raise RuntimeError(msg) + static_router_data = self._static_router_data[sid] = ( + self._build_static_router_data(sid, environ) ) - .partition(",")[0] - .strip() - ) router_data = event.router_data + router_data.update(static_router_data) + # The cached headers reach the event, and from there `state.router_data`, + # which is a plain mutable dict: sharing the mapping would let a handler + # mutating `self.router_data["headers"]` corrupt the connection cache for + # every later event on this socket. The shallow copy is ~17x cheaper than + # the per-event header decode it replaced, so the cache still pays off. + router_data[constants.RouteVar.HEADERS] = static_router_data[ + constants.RouteVar.HEADERS + ].copy() router_data.update({ constants.RouteVar.QUERY: format.format_query_params(event.router_data), constants.RouteVar.CLIENT_TOKEN: token, - constants.RouteVar.SESSION_ID: sid, - constants.RouteVar.HEADERS: headers, - constants.RouteVar.CLIENT_IP: client_ip, }) router_data[constants.RouteVar.PATH] = "/" + ( self.app.router(path) or "404" @@ -2408,4 +2446,11 @@ async def link_token_to_sid(self, sid: str, token: str): BaseStateToken(ident=new_token or token, cls=self.app._state) ) as state: state.router_data[constants.RouteVar.SESSION_ID] = sid - state.router = RouterData.from_router_data(state.router_data) + # Record the identity the state was loaded under; duplicate-token + # handling can hand back a fresh one here. + state.router_data[constants.RouteVar.CLIENT_TOKEN] = new_token or token + # Rebuild from router_data to keep the session var in step with it. + if ( + session := SessionData.from_router_data(state.router_data) + ) != state.rx_router_session: + state.rx_router_session = session diff --git a/reflex/istate/data.py b/reflex/istate/data.py index 74e13c60899..d70962e8b2c 100644 --- a/reflex/istate/data.py +++ b/reflex/istate/data.py @@ -1,9 +1,9 @@ """This module contains the dataclasses representing the router object.""" import dataclasses -from collections.abc import Mapping +from collections.abc import Callable, Mapping from types import MappingProxyType -from typing import TYPE_CHECKING, Any, ClassVar +from typing import TYPE_CHECKING, Any, ClassVar, Final, NoReturn from urllib.parse import _NetlocResultMixinStr, parse_qsl, urlsplit from reflex_base import constants @@ -159,6 +159,50 @@ def __new__(cls, url: str): object.__setattr__(obj, "fragment", fragment) return obj + def __setattr__(self, name: str, value: Any) -> NoReturn: + """Reject attribute assignment. + + A `ReflexURL` is a parsed view of an immutable `str`, and the empty + one is the class-level default of `URLData.href`, so it is shared by + every state that has not navigated yet. Letting a component assign to + a parsed component would rewrite that shared object for every state. + `__new__` fills the components with `object.__setattr__`. + + Args: + name: The attribute being assigned. + value: The value it would take. + + Raises: + AttributeError: Always. + """ + msg = f"cannot assign to {name!r}: ReflexURL is immutable" + raise AttributeError(msg) + + def __delattr__(self, name: str) -> NoReturn: + """Reject attribute deletion. + + Args: + name: The attribute being deleted. + + Raises: + AttributeError: Always. + """ + msg = f"cannot delete {name!r}: ReflexURL is immutable" + raise AttributeError(msg) + + def __reduce__(self) -> tuple[type["ReflexURL"], tuple[str]]: + """Persist only the URL itself, re-splitting it on the way back in. + + Every parsed component is derived from the string by ``__new__``, so + pickling them as well writes the URL into the state store several times + over. Reconstructing costs one ``urlsplit`` and is cheaper than reading + the components back. + + Returns: + The callable and argument that rebuild this URL. + """ + return (type(self), (str.__str__(self),)) + @serializer(to=dict) def _serialize_reflex_url(obj: ReflexURL) -> dict: @@ -382,6 +426,115 @@ def _serialize_page_data(obj: PageData) -> dict: return {key.name: getattr(obj, key.name) for key in dataclasses.fields(obj)} +def _url_from_router_data(router_data: dict) -> ReflexURL: + """Build the browser URL for the page described by a router_data dict. + + Args: + router_data: the router_data dict. + + Returns: + The parsed browser URL (origin header + prefixed path). + """ + return ReflexURL( + router_data.get(constants.RouteVar.HEADERS, {}).get("origin", "") + + get_config().prepend_frontend_path( + router_data.get(constants.RouteVar.ORIGIN, "") + ) + ) + + +# The parsed empty URL, shared as every URLData default: it is immutable, so +# one instance can back every state that has not navigated yet. +_EMPTY_URL: Final = ReflexURL("") + + +@dataclasses.dataclass(frozen=True) +class URLData: + """The parsed components of the current page URL. + + Storage form of ``RouterData.url`` in the state: unlike ``ReflexURL`` (a + ``str`` subclass, which ``json.dumps`` would serialize as a bare string), + a dataclass goes through the registered serializer, so the frontend + receives the parsed component dict. + """ + + # Read off the empty URL so `URLData()` is exactly + # `URLData.from_url(ReflexURL(""))` -- note `ReflexURL("").origin` is "://". + scheme: str = _EMPTY_URL.scheme + netloc: str = _EMPTY_URL.netloc + origin: str = _EMPTY_URL.origin + path: str = _EMPTY_URL.path + query: str = _EMPTY_URL.query + query_parameters: Mapping[str, str] = _EMPTY_URL.query_parameters + fragment: str = _EMPTY_URL.fragment + # Annotated str so the frontend var for this field renders the raw href + # string, but always holds a ReflexURL at runtime so the backend keeps + # parsed-component access without re-splitting the URL. + href: str = _EMPTY_URL + + @classmethod + def from_url(cls, url: ReflexURL) -> "URLData": + """Create a URLData object from an already-parsed ReflexURL. + + Args: + url: the parsed URL. + + Returns: + A URLData object mirroring the URL's components. + """ + return cls( + scheme=url.scheme, + netloc=url.netloc, + origin=url.origin, + path=url.path, + query=url.query, + query_parameters=url.query_parameters, + fragment=url.fragment, + href=url, + ) + + @classmethod + def from_router_data(cls, router_data: dict) -> "URLData": + """Create a URLData object from the given router_data. + + Args: + router_data: the router_data dict. + + Returns: + A URLData object for the page described by the router_data. + """ + return cls.from_url(_url_from_router_data(router_data)) + + def __reduce__(self) -> tuple[Callable[[str], "URLData"], tuple[str]]: + """Persist only the href, deriving the components again on the way back. + + Every other field is a parsed piece of ``href``, so storing them too + writes the URL into the state store eight times over. This is the + storage form of a router var, so it is pickled on every state write. + + Returns: + The callable and argument that rebuild this URLData. + """ + return (_url_data_from_href, (str.__str__(self.href),)) + + +def _url_data_from_href(href: str) -> URLData: + """Rebuild a URLData from the raw href alone. + + Args: + href: the full URL string. + + Returns: + A URLData with every component re-derived from the URL. + """ + return URLData.from_url(ReflexURL(href)) + + +@serializer(to=dict) +def _serialize_url_data(obj: URLData) -> dict: + return {key.name: getattr(obj, key.name) for key in dataclasses.fields(obj)} + + @dataclasses.dataclass(frozen=True) class SessionData: """An object containing session data.""" @@ -451,16 +604,21 @@ def from_router_data(cls, router_data: dict) -> "RouterData": session=SessionData.from_router_data(router_data), headers=HeaderData.from_router_data(router_data), _page=PageData.from_router_data(router_data), - url=ReflexURL( - router_data.get(constants.RouteVar.HEADERS, {}).get("origin", "") - + get_config().prepend_frontend_path( - router_data.get(constants.RouteVar.ORIGIN, "") - ) - ), + url=_url_from_router_data(router_data), route_id=router_data.get(constants.RouteVar.PATH, ""), ) +# Keys of the serialized RouterData: the object shape the frontend receives. +# `serialize_router_data` emits it, and `RouterDataVar` composes the same shape +# when the whole router is rendered, so the two must not drift apart. +SESSION_KEY: Final = "session" +HEADERS_KEY: Final = "headers" +PAGE_KEY: Final = "page" +URL_KEY: Final = "url" +ROUTE_ID_KEY: Final = "route_id" + + @serializer(to=dict) def serialize_router_data(obj: RouterData) -> dict: """Serialize a RouterData object to a dict. @@ -472,13 +630,156 @@ def serialize_router_data(obj: RouterData) -> dict: A dict representation of the RouterData object. """ return { - "session": obj.session, - "headers": obj.headers, - "page": obj._page, + SESSION_KEY: obj.session, + HEADERS_KEY: obj.headers, + PAGE_KEY: obj._page, # ReflexURL is a str subclass, so json.dumps handles it natively and # never invokes the `default=serialize` hook. Call the URL serializer # eagerly here so the frontend receives the parsed component dict # instead of just the raw URL string. - "url": _serialize_reflex_url(obj.url), - "route_id": obj.route_id, + URL_KEY: _serialize_reflex_url(obj.url), + ROUTE_ID_KEY: obj.route_id, } + + +def _null_var() -> Var: + """Placeholder default for RouterDataVar component fields. + + Returns: + A null Var. + """ + return Var(_js_expr="null", _var_type=None) + + +@dataclasses.dataclass( + eq=False, + frozen=True, + slots=True, +) +class RouterDataVar(CachedVarOperation, ObjectVar[RouterData]): + """Switchboard Var for ``State.router``. + + Router data is stored in separate per-field base vars on the root state + (session, headers, page, url, route_id) so that unchanged + connection-scoped data is not re-sent in the delta on every navigation. + This var stitches them back together: each attribute resolves directly to + the underlying per-field base var, and rendering the var itself produces + an object literal matching the pre-split serialized router shape. + """ + + _url_var: Var = dataclasses.field(default_factory=_null_var) + _page_var: Var = dataclasses.field(default_factory=_null_var) + _session_var: Var = dataclasses.field(default_factory=_null_var) + _headers_var: Var = dataclasses.field(default_factory=_null_var) + _route_id_var: Var = dataclasses.field(default_factory=_null_var) + _default_var_type: ClassVar[Any] = RouterData + + @cached_property_no_lock + def _cached_var_name(self) -> str: + """Render the router as an object literal over the per-field vars. + + Returns: + The JS expression for the assembled router object. + """ + return ( + "({ " + + ", ".join(f'"{key}": {var!s}' for key, var in self._wire_fields().items()) + + " })" + ) + + def _wire_fields(self) -> dict[str, Var]: + """Map each serialized RouterData key to the var backing it. + + Returns: + The keys of the serialized router shape, in order, to their vars. + """ + return { + SESSION_KEY: self._session_var, + HEADERS_KEY: self._headers_var, + PAGE_KEY: self._page_var, + URL_KEY: self._url_var, + ROUTE_ID_KEY: self._route_id_var, + } + + @property + def session(self) -> ObjectVar[SessionData]: + """The per-connection session data. + + Returns: + ObjectVar for the ``rx_router_session`` base var. + """ + return self._session_var.to(ObjectVar, SessionData) + + @property + def headers(self) -> ObjectVar[HeaderData]: + """The headers of the websocket connection request. + + Returns: + ObjectVar for the ``rx_router_headers`` base var. + """ + return self._headers_var.to(ObjectVar, HeaderData) + + @property + def page(self) -> ObjectVar[PageData]: + """The page data for the current page (deprecated, use ``url``). + + Returns: + ObjectVar for the ``rx_router_page`` base var. + """ + return self._page_var.to(ObjectVar, PageData) + + # RouterData exposes the page data under both `page` and `_page`. + _page = page + + @property + def url(self) -> ReflexURLCastedVar: + """The parsed URL of the current page. + + Returns: + ReflexURLCastedVar over the ``rx_router_url`` base var. + """ + return ReflexURLCastedVar.create(self._url_var) + + @property + def route_id(self) -> StringVar: + """The route pattern that matched the current page. + + Returns: + StringVar for the ``rx_router_route_id`` base var. + """ + return self._route_id_var.to(str) + + @classmethod + def create( + cls, + *, + session: Var, + headers: Var, + page: Var, + url: Var, + route_id: Var, + _var_data: VarData | None = None, + ) -> "RouterDataVar": + """Create a RouterDataVar over the per-field router base vars. + + Args: + session: The ``rx_router_session`` base var. + headers: The ``rx_router_headers`` base var. + page: The ``rx_router_page`` base var. + url: The ``rx_router_url`` base var. + route_id: The ``rx_router_route_id`` base var. + _var_data: Additional VarData to merge in. + + Returns: + The new RouterDataVar. + """ + return cls( + _js_expr="", + _var_type=RouterData, + _var_data=_var_data, + _url_var=url, + _page_var=page, + _session_var=session, + _headers_var=headers, + _route_id_var=route_id, + ) diff --git a/reflex/istate/shared.py b/reflex/istate/shared.py index 4d5f359c066..d87a2f33569 100644 --- a/reflex/istate/shared.py +++ b/reflex/istate/shared.py @@ -6,7 +6,7 @@ from collections.abc import AsyncIterator from typing import TypeVar -from reflex_base.constants import ROUTER_DATA +from reflex_base.constants import ROUTER_DATA, ROUTER_VARS from reflex_base.event import Event, get_hydrate_event from reflex_base.registry import RegistrationContext from reflex_base.utils.exceptions import ReflexRuntimeError @@ -114,7 +114,7 @@ async def _patch_state( linked_state._mark_dirty() # Apply the updates into the existing state tree for rehydrate. root_state = original_state._get_root_state() - root_state.dirty_vars.add("router") + root_state.dirty_vars.update(ROUTER_VARS) root_state.dirty_vars.add(ROUTER_DATA) root_state._mark_dirty() # The delta is discarded: it is only resolved to refresh computed vars, @@ -248,7 +248,7 @@ async def _link_to(self, token: str) -> Self: return self # already linked to this token if self._linked_to and self._linked_to != token: # Disassociate from previous linked token since unlink will not be called. - self._linked_from.discard(self.router.session.client_token) + self._linked_from.discard(self.rx_router_session.client_token) # TODO: Change StateManager to accept token + class instead of combining them in a string. if "_" in token: msg = f"Invalid token {token} for linking state {self.get_full_name()}, cannot use underscore (_) in the token name." @@ -283,12 +283,12 @@ async def _unlink(self): # Break the linkage for future events. self._reflex_internal_links.pop(state_name) - self._linked_from.discard(self.router.session.client_token) + self._linked_from.discard(self.rx_router_session.client_token) # Patch in the original state, apply updates, then rehydrate. private_root_state = await get_state_manager().get_state( BaseStateToken( - ident=self.router.session.client_token, + ident=self.rx_router_session.client_token, cls=type(self), ) ) @@ -337,14 +337,13 @@ async def _internal_patch_linked_state( # Set client_token on the linked root so that subsequent get_state # calls when directly modifying a linked token will load the # associated instance. - if linked_root_state.router.session.client_token != token: + if ( + session := linked_root_state.rx_router_session + ).client_token != token: import dataclasses as dc - linked_root_state.router = dc.replace( - linked_root_state.router, - session=dc.replace( - linked_root_state.router.session, client_token=token - ), + linked_root_state.rx_router_session = dc.replace( + session, client_token=token ) if linked_root_state is None: linked_root_state = await get_state_manager().get_state( @@ -357,8 +356,8 @@ async def _internal_patch_linked_state( # Avoid unnecessary dirtiness of shared state when there are no changes. if type(self) not in self._held_locks[token]: self._held_locks[token][type(self)] = linked_state - if self.router.session.client_token not in linked_state._linked_from: - linked_state._linked_from.add(self.router.session.client_token) + if self.rx_router_session.client_token not in linked_state._linked_from: + linked_state._linked_from.add(self.rx_router_session.client_token) if linked_state._linked_to != token: linked_state._linked_to = token await self._exit_stack.enter_async_context( @@ -449,7 +448,7 @@ async def _modify_linked_states( affected_tokens.update( token for token in linked_state._linked_from - if token != self.router.session.client_token + if token != self.rx_router_session.client_token ) # When modifying a shared token directly (empty _reflex_internal_links), # the held locks will be empty. Check SharedState substates for linked diff --git a/reflex/state.py b/reflex/state.py index 3601ae366f1..d332aef5a14 100644 --- a/reflex/state.py +++ b/reflex/state.py @@ -26,7 +26,9 @@ Final, ParamSpec, TypeVar, + cast, get_type_hints, + overload, ) from reflex_base import constants @@ -74,7 +76,15 @@ import reflex.istate.dynamic from reflex import event from reflex.istate import HANDLED_PICKLE_ERRORS, debug_failed_pickles -from reflex.istate.data import RouterData +from reflex.istate.data import ( + HeaderData, + PageData, + ReflexURL, + RouterData, + RouterDataVar, + SessionData, + URLData, +) from reflex.istate.proxy import ImmutableMutableProxy as ImmutableMutableProxy from reflex.istate.proxy import MutableProxy, is_mutable_type from reflex.istate.storage import ClientStorageBase @@ -82,6 +92,17 @@ from reflex.utils import console, format, types from reflex.utils.exec import is_testing_env +# The key a pre-split pickle stored the whole RouterData under. Not +# `constants.ROUTER`: this name is frozen into payloads already on disk. +_LEGACY_ROUTER_PICKLE_KEY = "router" + +# Shared empty router defaults. Each is a frozen dataclass whose members are +# themselves immutable, so one instance can back every state's field instead +# of being rebuilt per state. +_DEFAULT_SESSION_DATA = SessionData() +_DEFAULT_HEADER_DATA = HeaderData() +_DEFAULT_URL_DATA = URLData() + logger = logging.getLogger(__name__) if TYPE_CHECKING: @@ -430,6 +451,131 @@ def _is_user_descriptor(value: Any) -> bool: return not is_computed_var(value) +def _router_fget(self: BaseState) -> RouterData: + """Assemble the RouterData view over the per-field router vars. + + Args: + self: The state instance. + + Returns: + The RouterData for the current connection and page. + """ + return RouterData( + session=self.rx_router_session, + headers=self.rx_router_headers, + _page=self.rx_router_page, + # URLData.href always holds a ReflexURL at runtime (see URLData). + url=cast("ReflexURL", self.rx_router_url.href), + route_id=self.rx_router_route_id, + ) + + +def _router_fset(self: BaseState, value: RouterData) -> None: + """Decompose a RouterData assignment into the per-field router vars. + + Args: + self: The state instance. + value: The RouterData to store. + """ + self.rx_router_session = value.session + self.rx_router_headers = value.headers + self.rx_router_page = value._page + self.rx_router_url = URLData.from_url(value.url) + self.rx_router_route_id = value.route_id + + +def _get_router_var(cls: type[BaseState]) -> RouterDataVar: + """Get (or build and cache) the router switchboard var for a state class. + + Args: + cls: The state class the ``router`` attribute was accessed on. + + Returns: + The RouterDataVar over the root state's per-field router vars. + """ + root_cls = cls.get_root_state() + router_var = root_cls.__dict__.get("_reflex_router_var") + if router_var is None: + base_vars = root_cls.base_vars + if constants.ROUTER_SESSION not in base_vars: + # BaseState itself and mixins never initialize base vars; give + # introspection-style access an unbound switchboard. + return RouterDataVar(_js_expr="", _var_type=RouterData) + router_var = RouterDataVar.create( + session=base_vars[constants.ROUTER_SESSION], + headers=base_vars[constants.ROUTER_HEADERS], + page=base_vars[constants.ROUTER_PAGE], + url=base_vars[constants.ROUTER_URL], + route_id=base_vars[constants.ROUTER_ROUTE_ID], + # Name the `router` attribute the switchboard stands for, so + # `get_var_value(State.router)` resolves it through the property + # and hands back the composed RouterData, as it does on a state + # with a single `router` base var. + _var_data=VarData( + state=root_cls.get_full_name(), field_name=constants.ROUTER + ), + ) + setattr(root_cls, "_reflex_router_var", router_var) # noqa: B010 + return router_var + + +class _RouterDescriptor(property): + """Property exposing the per-field router vars as a single ``router`` attribute. + + Instance access composes a ``RouterData`` view from the per-field router + vars and assignment decomposes one into them, so existing reads and writes + of ``state.router`` keep working unchanged. Class-level access returns the + ``RouterDataVar`` switchboard, resolving ``State.router.`` to the + underlying per-field base var. Subclassing ``property`` keeps the state + field machinery from treating this as a base var and lets ComputedVar + dependency tracking recurse into the getter, so any computed var reading + ``self.router`` depends on the per-field vars. + """ + + if TYPE_CHECKING: + + @overload + def __get__(self, instance: None, owner: type, /) -> RouterDataVar: ... + + @overload + def __get__(self, instance: BaseState, owner: type, /) -> RouterData: ... + + def __get__(self, instance: Any, owner: type | None = None, /) -> Any: + """Get the switchboard var (class) or RouterData view (instance). + + Args: + instance: The state instance, or None for class access. + owner: The class through which the attribute was accessed. + + Returns: + The RouterDataVar for class access, or the RouterData view. + """ + + def __set__(self, instance: Any, value: RouterData) -> None: + """Set the router data on the instance. + + Args: + instance: The state instance. + value: The RouterData to store. + """ + + else: + + def __get__(self, instance: Any, owner: type | None = None, /): + """Get the switchboard var (class) or RouterData view (instance). + + Args: + instance: The state instance, or None for class access. + owner: The class through which the attribute was accessed. + + Returns: + The RouterDataVar for class access, or the RouterData view. + """ + if instance is None: + return _get_router_var(owner) + return super().__get__(instance, owner) + + all_base_state_classes: dict[str, None] = {} # Instance bookkeeping fields and framework methods read on every event. They @@ -548,8 +694,32 @@ class BaseState(EvenMoreBasicBaseState, metaclass=_StateMeta): default_factory=builtins.dict, is_var=False ) - # The router data for the current page - router: Field[RouterData] = field(default_factory=RouterData) + # The per-connection session data (constant for the socket lifetime). + # These three defaults are frozen dataclasses holding only immutable + # members, so every state can share one instance instead of building a + # fresh one per field per state. `field()` cannot be used for that: it + # only shares a `default` whose type is in `IMMUTABLE_TYPES`, and + # otherwise deep-copies it per instance. + rx_router_session: Field[SessionData] = Field(default=_DEFAULT_SESSION_DATA) + + # The headers of the connection request (constant for the socket lifetime). + rx_router_headers: Field[HeaderData] = Field(default=_DEFAULT_HEADER_DATA) + + # The page data for the current page (deprecated; params feeds dynamic route vars). + rx_router_page: Field[PageData] = field(default_factory=PageData) + + # The parsed URL of the current page. + rx_router_url: Field[URLData] = Field(default=_DEFAULT_URL_DATA) + + # The route pattern that matched the current page. + rx_router_route_id: Field[str] = field(default="") + + # Switchboard for the router vars above: instance reads compose a + # RouterData view, writes decompose into the per-field vars, and class + # access returns the RouterDataVar. Deliberately not a Field: storing each + # kind of router data in its own base var means a navigation delta only + # re-sends the navigation-scoped vars, not session/headers. + router = _RouterDescriptor(_router_fget, _router_fset) # Whether the state has ever been touched since instantiation. _was_touched: bool = field(default=False, is_var=False) @@ -654,7 +824,7 @@ def __init_subclass__(cls, mixin: bool = False, **kwargs): **kwargs: The kwargs to pass to the init_subclass method. Raises: - StateValueError: If a substate class shadows another. + StateValueError: If a substate shadows another. """ from reflex_base.utils.exceptions import StateValueError @@ -784,6 +954,11 @@ def __init_subclass__(cls, mixin: bool = False, **kwargs): **cls.inherited_vars, **cls.base_vars, **cls.computed_vars, + # `router` is a switchboard over the per-field router vars rather + # than a field of its own, but it is usable as a Var everywhere one + # is accepted, so it is listed here (and thus inherited by + # substates). It has no backing field, so it never reaches a delta. + constants.ROUTER: _get_router_var(cls), } cls.event_handlers = {} @@ -1026,6 +1201,23 @@ def _init_var_dependency_dicts(cls): # Do not perform dep calculation when cache=False (these are always dirty). continue for state_name, dvar_set in cvar._deps(objclass=cls).items(): + if constants.ROUTER in dvar_set: + # `router` names the switchboard, which has no field of its + # own: depend on the per-field router vars instead. The Var + # form already carries them, so only the legacy string form + # arrives here without them, and only it is deprecated. + if dvar_set.isdisjoint(constants.ROUTER_VARS): + console.deprecate( + feature_name=f'ComputedVar deps=["router"] on {cls.__name__}.{cvar_name}', + reason="the router var was split; depend on the router" + " Var instead (e.g. deps=[State.router.url] for one" + " field, or deps=[State.router] for all of them).", + deprecation_version="0.9.12", + removal_version="1.0", + ) + dvar_set = (dvar_set - {constants.ROUTER}) | set( + constants.ROUTER_VARS + ) state_cls = cls.get_root_state().get_class_substate(state_name) for dvar in dvar_set: defining_state_cls = state_cls @@ -1099,7 +1291,9 @@ def _check_overridden_basevars(cls): """ hints = cls._get_type_hints() for name, computed_var_ in cls._get_computed_vars(): - if name in hints: + # `router` is not a field, but shadowing the descriptor would + # silently break router access for the whole state tree. + if name in hints or name == constants.ROUTER: msg = f"The computed var name `{computed_var_._js_expr}` shadows a base var in {cls.__module__}.{cls.__name__}; use a different name instead" raise ComputedVarShadowsBaseVarsError(msg) @@ -1171,6 +1365,11 @@ def get_skip_vars(cls) -> set[str]: "dirty_vars", "dirty_substates", "router_data", + # Listed in `vars` but backed by no field of its own, so a + # `router` annotation must never become a base var that would + # half-shadow the descriptor. Substates are already covered by + # `inherited_vars` above; this catches a root state class. + constants.ROUTER, } | types.RESERVED_BACKEND_VAR_NAMES ) @@ -1556,7 +1755,7 @@ def setup_dynamic_args(cls, args: builtins.dict[str, str]): def argsingle_factory(param: str): def inner_func(self: BaseState) -> str: - return self.router._page.params.get(param, "") + return self.rx_router_page.params.get(param, "") inner_func.__name__ = param @@ -1564,7 +1763,7 @@ def inner_func(self: BaseState) -> str: def arglist_factory(param: str): def inner_func(self: BaseState) -> list[str]: - return self.router._page.params.get(param, []) + return self.rx_router_page.params.get(param, []) inner_func.__name__ = param @@ -1581,7 +1780,7 @@ def inner_func(self: BaseState) -> list[str]: dynamic_vars[param] = DynamicRouteVar( fget=func, auto_deps=False, - deps=["router"], + deps=[constants.ROUTER_PAGE], _var_data=VarData.from_state(cls, param), ) setattr(cls, param, dynamic_vars[param]) @@ -1759,7 +1958,7 @@ def reset(self): # Reset the base vars. fields = self.get_fields() for prop_name in self.base_vars: - if prop_name == constants.ROUTER: + if prop_name in constants.ROUTER_VARS: continue # never reset the router data field = fields[prop_name] if default_factory := field.default_factory: @@ -1777,6 +1976,90 @@ def reset(self): for substate in self.substates.values(): substate.reset() + def _update_router_vars( + self, + router_data: builtins.dict[str, Any], + previous_router_data: builtins.dict[str, Any], + ) -> builtins.dict[str, Any]: + """Update the per-field router vars from a new router_data dict. + + Each var is rebuilt only when the router_data keys it derives from + changed, so connection-scoped data (session, headers) is not recomputed + on every navigation, and is then assigned only when the rebuilt value + actually differs -- different keys can still yield an equal value (an + absent key and an empty one both produce the default), and assigning + regardless would dirty the var, mark the state touched, and persist it. + + A key missing from ``router_data`` carries no information about the + value it feeds, so the previous one is carried forward rather than + letting the constructors default it away: a payload holding only the + navigation keys must not empty the connection-scoped vars, nor rebuild + the page and URL without the origin header that gives them their host. + + Args: + router_data: The new router_data dict. + previous_router_data: The router_data dict this state last saw. + + Returns: + The router_data to store on the state: the new values over the + previous ones, so a partial payload does not drop keys for the + next comparison either. + """ + # Merging also makes an absent key compare equal to what it replaced, + # so it is not read as a change without a special case for it. + merged = ( + {**previous_router_data, **router_data} + if previous_router_data + else router_data + ) + get = merged.get + prev_get = previous_router_data.get + + headers_changed = prev_get(constants.RouteVar.HEADERS) != get( + constants.RouteVar.HEADERS + ) + # Only the origin header feeds the URL/page host, so the navigation + # vars must not be rebuilt for a change to any other header. + origin_changed = headers_changed and ( + prev_get(constants.RouteVar.HEADERS, {}).get("origin", "") + != get(constants.RouteVar.HEADERS, {}).get("origin", "") + ) + + if ( + any( + prev_get(key) != get(key) + for key in ( + constants.RouteVar.CLIENT_TOKEN, + constants.RouteVar.SESSION_ID, + constants.RouteVar.CLIENT_IP, + ) + ) + and (session := SessionData.from_router_data(merged)) + != self.rx_router_session + ): + self.rx_router_session = session + if ( + headers_changed + and (headers := HeaderData.from_router_data(merged)) + != self.rx_router_headers + ): + self.rx_router_headers = headers + if ( + origin_changed + or prev_get(constants.RouteVar.PATH) != get(constants.RouteVar.PATH) + or prev_get(constants.RouteVar.ORIGIN) != get(constants.RouteVar.ORIGIN) + or prev_get(constants.RouteVar.QUERY) != get(constants.RouteVar.QUERY) + ): + if (page := PageData.from_router_data(merged)) != self.rx_router_page: + self.rx_router_page = page + if (url := URLData.from_router_data(merged)) != self.rx_router_url: + self.rx_router_url = url + if ( + route_id := get(constants.RouteVar.PATH, "") + ) != self.rx_router_route_id: + self.rx_router_route_id = route_id + return merged + @classmethod @functools.lru_cache def _is_client_storage(cls, prop_name_or_field: str | Field) -> bool: @@ -1889,7 +2172,9 @@ async def _get_state_from_redis(self, state_cls: type[T_STATE]) -> T_STATE: ) raise RuntimeError(msg) state_in_redis = await state_manager.get_state( - token=BaseStateToken(ident=self.router.session.client_token, cls=state_cls), + token=BaseStateToken( + ident=self.rx_router_session.client_token, cls=state_cls + ), top_level=False, for_state_instance=self, ) @@ -2285,7 +2570,8 @@ def __getstate__(self): state = state.copy() if state.get("parent_state") is not None: # Do not serialize router data in substates (only the root state). - state.pop("router", None) + for router_var in constants.ROUTER_VARS: + state.pop(router_var, None) state.pop("router_data", None) # Never serialize parent_state or substates. state.pop("parent_state", None) @@ -2306,6 +2592,9 @@ def __setstate__(self, state: builtins.dict[str, Any]): """ state["parent_state"] = None state["substates"] = {} + # Pre-split pickles stored a RouterData under this key, now a + # descriptor; drop it so unpickling does not route through the setter. + state.pop(_LEGACY_ROUTER_PICKLE_KEY, None) for key, value in state.items(): object.__setattr__(self, key, value) @@ -2682,7 +2971,7 @@ def on_load_internal(self) -> list[Event | EventSpec | event.EventCallback] | No The list of events to queue for on load handling. """ load_events = RegistrationContext.get().app.get_load_events( - self.router.url.path + self.rx_router_url.path ) if not load_events: self.is_hydrated = True diff --git a/tests/benchmarks/test_event_processing.py b/tests/benchmarks/test_event_processing.py index c4fb3e13d95..6a714dfabae 100644 --- a/tests/benchmarks/test_event_processing.py +++ b/tests/benchmarks/test_event_processing.py @@ -207,3 +207,85 @@ async def test_table_event_deltas(): assert delta["total_amount" + FIELD_MARKER] == sum( row["amount"] for row in expected ) + + +@pytest.fixture +def on_event_harness(): + """Set up an EventNamespace with a connected socket for benchmarking on_event. + + The event processor's enqueue is mocked out so the benchmark isolates the + per-event router_data preparation (which reuses the connection-scoped + data gathered once in on_connect). + + Yields: + An async callable that feeds the given number of events through + ``EventNamespace.on_event``, and the event loop to drive it with. + """ + from reflex.app import App, EventNamespace + + app = App() + app._event_processor = mock.Mock(enqueue=mock.AsyncMock()) + namespace = EventNamespace("/event", app) + + sid = "benchmark-sid" + environ = { + "QUERY_STRING": "token=benchmark-token", + "asgi.scope": { + "headers": [ + (b"host", b"localhost:3000"), + (b"origin", b"http://localhost:3000"), + (b"user-agent", b"Mozilla/5.0 (X11; Linux x86_64) benchmark"), + (b"accept-encoding", b"gzip, deflate, br"), + (b"accept-language", b"en-US,en;q=0.9"), + (b"cookie", b"session=abc123; theme=dark"), + (b"upgrade", b"websocket"), + (b"connection", b"Upgrade"), + (b"sec-websocket-version", b"13"), + (b"sec-websocket-key", b"dGhlIHNhbXBsZSBub25jZQ=="), + (b"x-forwarded-for", b"203.0.113.7, 10.0.0.1"), + ], + "client": ("127.0.0.1", 54321), + }, + } + + async def run_events(num_events: int) -> None: + """Feed events through on_event. + + Args: + num_events: Number of events to process. + """ + for _ in range(num_events): + await namespace.on_event( + sid, + { + "name": "state.hydrate", + "router_data": {"pathname": "/", "query": {}, "asPath": "/"}, + "payload": {}, + }, + ) + + loop = asyncio.new_event_loop() + loop.run_until_complete(namespace.on_connect(sid, environ)) + yield run_events, loop + loop.close() + + +def test_on_event_router_data( + on_event_harness, + benchmark: BenchmarkFixture, +): + """Benchmark the per-event router_data preparation in on_event. + + Headers and client IP are gathered once at connect time, so the + per-event path is reduced to merging the cached connection-scoped dict + with the event's navigation data. + + Args: + on_event_harness: The run_events async callable and its event loop. + benchmark: The codspeed benchmark fixture. + """ + run_events, loop = on_event_harness + + @benchmark + def _(): + loop.run_until_complete(run_events(num_events=10)) diff --git a/tests/units/istate/test_data.py b/tests/units/istate/test_data.py index 6ff0b6e805a..d7af58128dc 100644 --- a/tests/units/istate/test_data.py +++ b/tests/units/istate/test_data.py @@ -1,8 +1,10 @@ """Tests for ReflexURL parsing, serialization, and Var attribute access.""" from collections.abc import Mapping +from typing import cast from urllib.parse import parse_qsl +import pytest from reflex_base.vars.object import ObjectVar from reflex_base.vars.sequence import StringVar @@ -146,3 +148,204 @@ def test_router_url_var_renders_as_href_at_top_level(): """ url_var = rx.State.router.url assert str(url_var) == f'{url_var._original!s}?.["href"]' + + +def test_url_data_serializes_like_reflex_url(): + """URLData (the per-field storage form of the router URL) must serialize + to the same component dict shape as the eager ReflexURL serialization, so + the frontend var access patterns are unchanged by the router var split. + """ + import json + + from reflex_base.utils.format import json_dumps + + from reflex.istate.data import URLData, _serialize_reflex_url + + url = ReflexURL(SAMPLE_URL) + payload = json.loads(json_dumps(URLData.from_url(url))) + assert payload == json.loads(json_dumps(_serialize_reflex_url(url))) + # The runtime value of href keeps parsed-component access on the backend. + assert isinstance(URLData.from_url(url).href, ReflexURL) + + +def test_router_var_resolves_to_per_field_base_vars(): + """State.router is a switchboard: each attribute must resolve directly to + the per-field base var, so a navigation delta that only carries the + navigation-scoped vars still updates every rendered router expression. + """ + prefix = "reflex___state____state" + assert ( + str(rx.State.router.session.client_token) + == f'{prefix}.rx_router_session_rx_state_?.["client_token"]' + ) + assert ( + str(rx.State.router.headers.user_agent) + == f'{prefix}.rx_router_headers_rx_state_?.["user_agent"]' + ) + assert ( + str(rx.State.router.page.raw_path) + == f'{prefix}.rx_router_page_rx_state_?.["raw_path"]' + ) + assert str(rx.State.router.url) == f'{prefix}.rx_router_url_rx_state_?.["href"]' + assert ( + str(rx.State.router.url.path) == f'{prefix}.rx_router_url_rx_state_?.["path"]' + ) + assert str(rx.State.router.route_id) == f"{prefix}.rx_router_route_id_rx_state_" + + +def test_router_var_renders_composed_object(): + """Rendering State.router itself produces an object literal over the + per-field vars, matching the pre-split serialized router shape. + """ + prefix = "reflex___state____state" + assert str(rx.State.router) == ( + "({ " + f'"session": {prefix}.rx_router_session_rx_state_, ' + f'"headers": {prefix}.rx_router_headers_rx_state_, ' + f'"page": {prefix}.rx_router_page_rx_state_, ' + f'"url": {prefix}.rx_router_url_rx_state_, ' + f'"route_id": {prefix}.rx_router_route_id_rx_state_' + " })" + ) + + +def test_router_var_shape_matches_the_serializer(): + """The composed router literal and the serializer must emit the same keys. + + Rendering `State.router` as a whole has to produce the object shape the + backend serializes a `RouterData` into, or a component reading the whole + router would see different keys from the ones the delta carries. The two + are built in different places, so pin them to each other. + """ + import json + + from reflex_base.utils.format import json_dumps + + from reflex.istate.data import RouterData, serialize_router_data + + rendered_keys = list(rx.State.router._wire_fields()) + assert rendered_keys == list(serialize_router_data(RouterData())) + # And that is what actually reaches the client for a whole-router value. + assert rendered_keys == list(json.loads(json_dumps(RouterData()))) + + +def test_router_var_carries_state_var_data(): + """The switchboard var must merge the per-field vars' VarData so hooks + and context wiring for the root state are set up when it renders. + """ + var_data = rx.State.router._get_all_var_data() + assert var_data is not None + assert var_data.state == rx.State.get_full_name() + + +@pytest.mark.parametrize("attr", ["path", "scheme", "netloc", "query", "fragment"]) +def test_reflex_url_rejects_attribute_assignment(attr: str): + """A parsed component must not be assignable. + + `URLData.href` defaults to a class-level `ReflexURL("")`, so the empty URL + object is shared by every state that has not navigated yet. If a component + could be assigned, writing through one state's `router.url` would rewrite + that shared object for all of them. + """ + url = ReflexURL(SAMPLE_URL) + before = getattr(url, attr) + + with pytest.raises(AttributeError, match="immutable"): + setattr(url, attr, "/mutated") + with pytest.raises(AttributeError, match="immutable"): + delattr(url, attr) + + assert getattr(url, attr) == before + + +def test_shared_empty_url_default_cannot_be_mutated_through_a_state(): + """Writing through one state's router.url must not leak into another.""" + from reflex.istate.data import URLData + from reflex.state import BaseState + + # A root state, so the router fields live on the instance under test + # rather than being delegated to a parent that is not in a tree here. + class _URLIsolationState(BaseState): + pass + + one = _URLIsolationState(_reflex_internal_init=True) # pyright: ignore [reportCallIssue] + two = _URLIsolationState(_reflex_internal_init=True) # pyright: ignore [reportCallIssue] + + with pytest.raises(AttributeError, match="immutable"): + one.router.url.path = "/mutated" + + one.rx_router_url = URLData.from_url(ReflexURL("https://example.com/real")) + + assert one.router.url.path == "/real" + assert two.router.url.path == "" + assert cast("ReflexURL", URLData().href).path == "" + + +def test_url_data_default_matches_the_parsed_empty_url(): + """`URLData()` must equal `URLData.from_url(ReflexURL(""))`. + + The two are different construction paths to the same "not navigated yet" + value: the dataclass defaults back every fresh state, while `from_url` is + what `__reduce__` rebuilds a persisted one through. Written out by hand the + defaults drifted -- `ReflexURL("").origin` is "://", not "" -- so a state + reported one origin before a save and another after. + """ + from reflex.istate.data import URLData + + assert URLData() == URLData.from_url(ReflexURL("")) + + +@pytest.mark.parametrize( + "raw", + [ + "", + SAMPLE_URL, + "http://x/", + "https://a.b/c?d=1&d=2#f", + ], +) +def test_reflex_url_and_url_data_survive_pickling(raw: str): + """Both persist through a pickle round-trip with every component intact. + + `ReflexURL` and `URLData` persist only the URL itself and re-split it on + the way back, so this pins that the derived components come back equal + rather than being silently dropped or recomputed differently. + """ + import pickle + + from reflex.istate.data import URLData + + components = ("scheme", "netloc", "origin", "path", "query", "fragment") + + url = ReflexURL(raw) + restored_url = pickle.loads(pickle.dumps(url)) + assert type(restored_url) is ReflexURL + assert str(restored_url) == raw + for component in components: + assert getattr(restored_url, component) == getattr(url, component) + assert dict(restored_url.query_parameters) == dict(url.query_parameters) + + data = URLData.from_url(url) + restored_data = pickle.loads(pickle.dumps(data)) + assert restored_data == data + # The runtime href must still be a ReflexURL, or backend component access + # through `self.router.url` breaks after a state is loaded from the store. + assert isinstance(restored_data.href, ReflexURL) + + +def test_pickling_a_url_does_not_store_its_derived_components(): + """The persisted form must carry the URL once, not every parsed piece. + + `URLData` is the storage form of a router var, so it is pickled on every + state write. Storing the seven derived components alongside `href` wrote + the URL into the state store eight times over. + """ + import pickle + + from reflex.istate.data import URLData + + blob = pickle.dumps(URLData.from_url(ReflexURL(SAMPLE_URL))) + # Every component is derivable from the href, so the href is the only + # occurrence of the URL text in the payload. + assert blob.count(b"example.com") == 1 + assert b"query_parameters" not in blob diff --git a/tests/units/reflex_base/event/processor/test_base_state_processor.py b/tests/units/reflex_base/event/processor/test_base_state_processor.py index 6e0db179927..8d6ecefdb23 100644 --- a/tests/units/reflex_base/event/processor/test_base_state_processor.py +++ b/tests/units/reflex_base/event/processor/test_base_state_processor.py @@ -12,7 +12,7 @@ import pytest import pytest_asyncio from opentelemetry.trace import SpanKind, StatusCode -from reflex_base import otel +from reflex_base import constants, otel from reflex_base.constants import CompileVars, RouteVar from reflex_base.constants.state import FIELD_MARKER from reflex_base.environment import environment @@ -1338,3 +1338,163 @@ def noop(self): for p in metric_points(otel_metrics, otel.METRIC_STATE_ACQUIRE_DURATION) } assert Event.from_event_type(AcquireState.noop())[0].name in names + + +async def test_no_op_partial_router_data_leaves_the_state_untouched( + wired_app: App, + real_base_state_processor: BaseStateEventProcessor, + emitted_deltas: list, + token: str, +): + """A payload that merges to what is already there must not touch the state. + + A partial router_data (only the navigation keys, as `fix_events` produces) + is never equal to the full dict the state holds, so it reaches the merge. + If it merges to the same thing, nothing moved: assigning it anyway would + dirty router_data, mark the state touched, and persist it for an event + that changed nothing. + + Args: + wired_app: The App wired to the processor's state manager. + real_base_state_processor: The unmocked BaseStateEventProcessor. + emitted_deltas: List of deltas captured from the processor. + token: The client token. + """ + + class NoOpRouterState(State): + n: int = 0 + + @event + def bump(self): + self.n += 1 + + full_view = { + "pathname": "/a", + "asPath": "/a", + "query": {}, + "token": token, + "sid": "sid1", + "ip": "127.0.0.1", + "headers": {"origin": "http://localhost:3000"}, + } + # Same navigation, but carrying only the keys a chained event keeps. + navigation_only = {"pathname": "/a", "asPath": "/a", "query": {}} + + def client_event(router_data: dict[str, Any]) -> Event: + return dataclasses.replace( + Event.from_event_type(NoOpRouterState.bump())[0], router_data=router_data + ) + + async with real_base_state_processor as processor: + await processor.enqueue(token, client_event(full_view)) + await processor.join(10) + + root_ctx = real_base_state_processor._root_context + assert root_ctx is not None + state = await root_ctx.state_manager.get_state( + BaseStateToken(ident=token, cls=State) + ) + state._was_touched = False + emitted_deltas.clear() + + async with real_base_state_processor as processor: + await processor.enqueue(token, client_event(navigation_only)) + await processor.join(10) + + # The connection-scoped data survived the partial payload... + assert state.router_data["headers"] == full_view["headers"] + assert state.rx_router_session.client_token == token + # ...and nothing about the router was re-sent or marked dirty. + assert not any( + key.removesuffix(FIELD_MARKER) in constants.ROUTER_VARS + for _token, delta in emitted_deltas + for key in delta.get(State.get_full_name(), {}) + ) + assert not state._get_was_touched() + + +async def test_navigation_delta_elides_connection_scoped_router_vars( + wired_app: App, + real_base_state_processor: BaseStateEventProcessor, + emitted_deltas: list, + token: str, +): + """A navigation only re-sends the navigation-scoped router vars. + + Session and headers cannot change without going through a reconnect, so + re-shipping them in the delta of every client event is pure overhead. + The router is stored in per-field base vars precisely so that a + navigation marks only page/url/route_id dirty; a reconnect (new sid) + marks only the session dirty. + + Args: + wired_app: The App wired to the processor's state manager. + real_base_state_processor: The unmocked BaseStateEventProcessor. + emitted_deltas: List of deltas captured from the processor. + token: The client token. + """ + + class NavState(State): + n: int = 0 + + @event + def bump(self): + self.n += 1 + + headers = {"origin": "http://localhost:3000", "user-agent": "test-agent"} + + def view(path: str, sid: str = "sid1") -> dict[str, Any]: + return { + "pathname": path, + "asPath": path, + "query": {}, + "token": token, + "sid": sid, + "ip": "127.0.0.1", + "headers": headers, + } + + def client_event(router_data: dict[str, Any]) -> Event: + return dataclasses.replace( + Event.from_event_type(NavState.bump())[0], router_data=router_data + ) + + def router_vars_in_deltas() -> set[str]: + return { + key.removesuffix(FIELD_MARKER) + for _token, delta in emitted_deltas + for key in delta.get(State.get_full_name(), {}) + if key.removesuffix(FIELD_MARKER) in constants.ROUTER_VARS + } + + async def run_event(router_data: dict[str, Any]) -> None: + emitted_deltas.clear() + async with real_base_state_processor as processor: + await processor.enqueue(token, client_event(router_data)) + await processor.join(10) + + # First event on the connection populates every router var. + await run_event(view("/a")) + assert router_vars_in_deltas() == { + "rx_router_session", + "rx_router_headers", + "rx_router_page", + "rx_router_url", + "rx_router_route_id", + } + + # A navigation only re-sends the navigation-scoped vars. + await run_event(view("/b")) + assert router_vars_in_deltas() == { + "rx_router_page", + "rx_router_url", + "rx_router_route_id", + } + + # An event without a route change re-sends no router vars at all. + await run_event(view("/b")) + assert router_vars_in_deltas() == set() + + # A reconnect (new sid, same headers) re-sends only the session. + await run_event(view("/b", sid="sid2")) + assert router_vars_in_deltas() == {"rx_router_session"} diff --git a/tests/units/reflex_base/vars/test_base.py b/tests/units/reflex_base/vars/test_base.py index 3314f541384..f8eb80e8205 100644 --- a/tests/units/reflex_base/vars/test_base.py +++ b/tests/units/reflex_base/vars/test_base.py @@ -271,6 +271,64 @@ def __hash__(cls) -> int: ) +def test_var_data_merge_collects_field_names(): + """Merging vars of one state keeps every field name, deduped and in order.""" + merged = VarData.merge( + VarData(state="s", field_name="a"), + VarData(state="s", field_name="b"), + VarData(state="s", field_name="a"), + ) + + assert merged is not None + assert dict(merged.field_dependencies) == {"s": ("a", "b")} + # `field_name` stays the first, so existing single-field readers are intact. + assert merged.field_name == "a" + + +def test_var_data_merge_keeps_field_names_of_every_state(): + """A var spanning several states keeps each state's own fields. + + Fields stay grouped by the state that owns them, so a dependency on a + composite var tracks every field it reads rather than only those of + whichever state happened to merge first. + """ + merged = VarData.merge( + VarData(state="s", field_name="a"), + VarData(state="other", field_name="b"), + VarData(state="s", field_name="c"), + ) + + assert merged is not None + assert dict(merged.field_dependencies) == {"s": ("a", "c"), "other": ("b",)} + # The fallback accessors report the first state and its first field only. + assert merged.state == "s" + assert merged.field_name == "a" + + +def test_var_data_field_dependencies_round_trip(): + """`state`/`field_name` are the shorthand for a single-field mapping.""" + assert dict(VarData(state="s", field_name="a").field_dependencies) == {"s": ("a",)} + # A state with no named field is still recorded: many vars carry only the + # state, for its imports and hooks, and read no field. + assert dict(VarData(state="s").field_dependencies) == {"s": ()} + assert dict(VarData().field_dependencies) == {} + # The canonical form wins over the shorthand. + assert dict( + VarData( + state="ignored", + field_name="ignored", + field_dependencies={"s": ("a",), "other": ("b",)}, + ).field_dependencies + ) == {"s": ("a",), "other": ("b",)} + + +def test_var_data_field_name_reports_the_first_field(): + """`field_name` reports the first field of the first state.""" + assert VarData(field_name="a").field_name == "a" + assert VarData(field_dependencies={"s": ("a", "b")}).field_name == "a" + assert VarData().field_name == "" + + def test_serializer_attribute_error_is_not_masked() -> None: """An AttributeError raised inside a serializer surfaces chained, with its own frame.""" diff --git a/tests/units/test_app.py b/tests/units/test_app.py index 3eadd66418b..7d247286abb 100644 --- a/tests/units/test_app.py +++ b/tests/units/test_app.py @@ -16,15 +16,16 @@ import threading import unittest.mock import uuid -from collections.abc import Generator +from collections.abc import AsyncGenerator, Generator from concurrent.futures import ThreadPoolExecutor from contextlib import nullcontext as does_not_raise from importlib.util import find_spec from pathlib import Path -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, cast from unittest.mock import AsyncMock, Mock import pytest +import pytest_asyncio import reflex_base from opentelemetry import trace from opentelemetry.sdk.trace import TracerProvider @@ -77,7 +78,7 @@ ) from reflex.compiler.plugins import default_page_plugins from reflex.environment import environment -from reflex.istate.data import RouterData +from reflex.istate.data import RouterData, URLData from reflex.istate.manager.disk import StateManagerDisk from reflex.istate.manager.memory import StateManagerMemory from reflex.istate.manager.redis import StateManagerRedis @@ -373,9 +374,9 @@ def test_add_page_set_route_dynamic(index_page: ComponentCallable): assert app._pages.keys() == {"test/[dynamic]"} assert "dynamic" in app._state.computed_vars assert app._state.computed_vars["dynamic"]._deps(objclass=EmptyState) == { - EmptyState.get_full_name(): {constants.ROUTER}, + EmptyState.get_full_name(): {"rx_router_page"}, } - assert constants.ROUTER in app._state()._var_dependencies + assert "rx_router_page" in app._state()._var_dependencies def test_add_page_set_route_nested(app: App, index_page: ComponentCallable): @@ -2034,9 +2035,9 @@ async def test_dynamic_route_var_route_change_completed_on_load( assert arg_name in app._state.vars assert arg_name in app._state.computed_vars assert app._state.computed_vars[arg_name]._deps(objclass=DynamicState) == { - DynamicState.get_full_name(): {constants.ROUTER}, + DynamicState.get_full_name(): {"rx_router_page"}, } - assert constants.ROUTER in app._state()._var_dependencies + assert "rx_router_page" in app._state()._var_dependencies substate_token = BaseStateToken(ident=token, cls=DynamicState) exp_vals = ["foo", "foobar", "baz"] @@ -2070,6 +2071,16 @@ def _dynamic_state_event(name, val, **kwargs): val=exp_val, ) exp_router = RouterData.from_router_data(on_load_internal.router_data) + # Only the navigation-scoped router vars change (no session/headers in + # the router_data), so only those land in the delta. + exp_router_delta = { + "rx_router_page" + FIELD_MARKER: exp_router._page, + "rx_router_url" + FIELD_MARKER: URLData.from_url(exp_router.url), + } + if exp_index == 0: + # Every navigation here matches the same route, so the route_id + # only changes on the first one. + exp_router_delta["rx_router_route_id" + FIELD_MARKER] = exp_router.route_id async with mock_base_state_event_processor as processor: await processor.enqueue( token, @@ -2083,7 +2094,7 @@ def _dynamic_state_event(name, val, **kwargs): State.get_full_name(): { arg_name + FIELD_MARKER: exp_val, constants.CompileVars.IS_HYDRATED + FIELD_MARKER: False, - "router" + FIELD_MARKER: exp_router, + **exp_router_delta, }, DynamicState.get_full_name(): { f"comp_{arg_name}" + FIELD_MARKER: exp_val, @@ -4724,6 +4735,210 @@ def test_client_error_constants_match_frontend(): ) +@pytest_asyncio.fixture +async def event_namespace_with_processor_mock() -> AsyncGenerator[EventNamespace, None]: + """An EventNamespace whose app has a mocked event processor. + + Yields: + The EventNamespace instance. + """ + app = App() + app._event_processor = Mock(enqueue=AsyncMock()) + event_namespace = EventNamespace("/event", app) + yield event_namespace + # The token manager is backed by redis when one is configured; drop the + # tokens these tests link so they do not show up in another test's + # enumeration of the shared instance. Awaited rather than run in a fresh + # loop via asyncio.run: the redis client is bound to the test's loop. + await event_namespace._token_manager.disconnect_all() + + +def _connect_environ(token: str) -> dict[str, Any]: + return { + "QUERY_STRING": f"token={token}", + "asgi.scope": { + "headers": [ + (b"origin", b"http://localhost:3000"), + (b"user-agent", b"test-agent"), + ], + "client": ("127.0.0.1", 1234), + }, + } + + +def _client_event_payload() -> dict[str, Any]: + return { + "name": "state.hydrate", + "router_data": {"pathname": "/", "query": {}, "asPath": "/"}, + "payload": {}, + } + + +@pytest.mark.asyncio +async def test_on_event_uses_connect_time_router_data( + token: str, + event_namespace_with_processor_mock: EventNamespace, +): + """on_event merges the connection-scoped router_data gathered at connect. + + Headers, client IP, and session id are computed once in on_connect; the + per-event path must not re-read the connection environ at all. + + Args: + token: A token. + event_namespace_with_processor_mock: The event namespace fixture. + """ + event_namespace = event_namespace_with_processor_mock + await event_namespace.on_connect("sid1", _connect_environ(token)) + assert "sid1" in event_namespace._static_router_data + + # The per-event path must not re-read the connection environ. + event_namespace.app.sio = Mock( + get_environ=Mock(side_effect=AssertionError("environ must not be consulted")) + ) + await event_namespace.on_event("sid1", _client_event_payload()) + + enqueue_mock = cast(AsyncMock, event_namespace.app.event_processor.enqueue) + enqueue_mock.assert_called_once() + enqueued_token, event = enqueue_mock.call_args[0] + assert enqueued_token == token + assert event.router_data[constants.RouteVar.CLIENT_TOKEN] == token + assert event.router_data[constants.RouteVar.SESSION_ID] == "sid1" + assert event.router_data[constants.RouteVar.CLIENT_IP] == "127.0.0.1" + assert event.router_data[constants.RouteVar.HEADERS] == { + "origin": "http://localhost:3000", + "user-agent": "test-agent", + "asgi-scope-client": "127.0.0.1", + } + assert event.router_data[constants.RouteVar.PATH] == "/404" + assert event.router_data[constants.RouteVar.QUERY] == {} + + # Disconnect drops the cached connection data. + event_namespace.on_disconnect("sid1") + assert "sid1" not in event_namespace._static_router_data + + +@pytest.mark.asyncio +async def test_link_token_to_sid_records_the_connecting_identity( + token: str, + event_namespace_with_processor_mock: EventNamespace, + mocker: MockerFixture, +): + """The session var carries the token the state was loaded under. + + Duplicate-token handling hands back a fresh token, and the state is loaded + under it. Leaving `rx_router_session.client_token` empty until the first event + would let anything reading it in between -- a background task, a + shared-state link -- address the wrong state tree. + + Args: + token: A token. + event_namespace_with_processor_mock: The event namespace fixture. + mocker: pytest-mock fixture. + """ + event_namespace = event_namespace_with_processor_mock + state = Mock() + state.router_data = {} + mocker.patch.object( + event_namespace.app.state_manager, + "modify_state", + Mock(return_value=AsyncMock(__aenter__=AsyncMock(return_value=state))), + ) + + # No duplicate: the connecting token is recorded. + await event_namespace.link_token_to_sid("sid1", token) + assert state.router_data[constants.RouteVar.CLIENT_TOKEN] == token + assert state.rx_router_session.client_token == token + assert state.rx_router_session.session_id == "sid1" + + # Duplicate: the *new* token is recorded, not the one the client sent. + # The duplicate branch emits the replacement token to the client, which + # needs a server the bare namespace does not have. + event_namespace.emit = AsyncMock() # pyright: ignore[reportAttributeAccessIssue] + new_token = "a-fresh-token" + mocker.patch.object( + event_namespace._token_manager, + "link_token_to_sid", + AsyncMock(return_value=new_token), + ) + await event_namespace.link_token_to_sid("sid2", token) + assert state.router_data[constants.RouteVar.CLIENT_TOKEN] == new_token + assert state.rx_router_session.client_token == new_token + assert state.rx_router_session.session_id == "sid2" + + +@pytest.mark.asyncio +async def test_on_event_does_not_share_the_cached_headers( + token: str, + event_namespace_with_processor_mock: EventNamespace, +): + """Each event gets its own headers mapping, not the cached one. + + The headers reach `state.router_data`, a plain mutable dict, so sharing + the cached mapping would let a handler mutating it corrupt the connection + cache for every later event on the socket. + + Args: + token: A token. + event_namespace_with_processor_mock: The event namespace fixture. + """ + event_namespace = event_namespace_with_processor_mock + await event_namespace.on_connect("sid1", _connect_environ(token)) + cached_headers = event_namespace._static_router_data["sid1"][ + constants.RouteVar.HEADERS + ] + + await event_namespace.on_event("sid1", _client_event_payload()) + enqueue_mock = cast(AsyncMock, event_namespace.app.event_processor.enqueue) + _, event = enqueue_mock.call_args[0] + event_headers = event.router_data[constants.RouteVar.HEADERS] + + assert event_headers == cached_headers + assert event_headers is not cached_headers + # Mutating what the handler sees must not reach the connection cache. + event_headers["user-agent"] = "mutated" + assert cached_headers["user-agent"] == "test-agent" + + enqueue_mock.reset_mock() + await event_namespace.on_event("sid1", _client_event_payload()) + _, next_event = enqueue_mock.call_args[0] + assert ( + next_event.router_data[constants.RouteVar.HEADERS]["user-agent"] == "test-agent" + ) + + +@pytest.mark.asyncio +async def test_on_event_falls_back_to_environ_without_connect( + token: str, + event_namespace_with_processor_mock: EventNamespace, +): + """on_event computes and caches the static router_data if connect was missed. + + Args: + token: A token. + event_namespace_with_processor_mock: The event namespace fixture. + """ + event_namespace = event_namespace_with_processor_mock + await event_namespace._token_manager.link_token_to_sid(token, "sid1") + event_namespace.app.sio = Mock( + get_environ=Mock(return_value=_connect_environ(token)) + ) + + await event_namespace.on_event("sid1", _client_event_payload()) + await event_namespace.on_event("sid1", _client_event_payload()) + + # The environ is only consulted once; the result is cached for the sid. + event_namespace.app.sio.get_environ.assert_called_once() + enqueue_mock = cast(AsyncMock, event_namespace.app.event_processor.enqueue) + assert enqueue_mock.call_count == 2 + for call in enqueue_mock.call_args_list: + _, event = call[0] + assert event.router_data[constants.RouteVar.SESSION_ID] == "sid1" + assert ( + event.router_data[constants.RouteVar.HEADERS]["user-agent"] == "test-agent" + ) + + @pytest.mark.parametrize("compile_raises", [False, True]) def test_compile_releases_memo_naming_caches( mocker: MockerFixture, compile_raises: bool diff --git a/tests/units/test_state.py b/tests/units/test_state.py index da3320c1cf4..f4f6c134fa8 100644 --- a/tests/units/test_state.py +++ b/tests/units/test_state.py @@ -14,7 +14,7 @@ import threading from collections.abc import AsyncGenerator, Callable, Mapping from textwrap import dedent -from typing import Any, ClassVar, Literal, TypeVar +from typing import Any, ClassVar, Literal, TypeVar, cast from unittest.mock import AsyncMock, Mock import pytest @@ -45,7 +45,14 @@ import reflex as rx from reflex.app import App from reflex.environment import environment -from reflex.istate.data import HeaderData, RouterData, SessionData, _FrozenDictStrStr +from reflex.istate.data import ( + HeaderData, + RouterData, + RouterDataVar, + SessionData, + URLData, + _FrozenDictStrStr, +) from reflex.istate.manager import StateManager from reflex.istate.manager.disk import StateManagerDisk from reflex.istate.manager.memory import StateManagerMemory @@ -81,9 +88,9 @@ LOCK_EXPIRE_SLEEP = 2.5 if CI else 0.4 -formatted_router = { - "route_id": "", - "url": { +formatted_router_vars = { + "rx_router_route_id" + FIELD_MARKER: "", + "rx_router_url" + FIELD_MARKER: { "scheme": "", "netloc": "", "origin": "://", @@ -93,8 +100,12 @@ "fragment": "", "href": "", }, - "session": {"client_token": "", "client_ip": "", "session_id": ""}, - "headers": { + "rx_router_session" + FIELD_MARKER: { + "client_token": "", + "client_ip": "", + "session_id": "", + }, + "rx_router_headers" + FIELD_MARKER: { "host": "", "origin": "", "upgrade": "", @@ -110,7 +121,7 @@ "accept_language": "", "raw_headers": {}, }, - "page": { + "rx_router_page" + FIELD_MARKER: { "host": "", "path": "", "raw_path": "", @@ -389,7 +400,8 @@ def test_class_vars(test_state): """ cls = type(test_state) assert cls.vars.keys() == { - "router", + constants.ROUTER, + *constants.ROUTER_VARS, "num1", "num2", "key", @@ -470,8 +482,10 @@ def test_dict(test_state: TestState): } test_state_dict = test_state.dict() assert set(test_state_dict) == substates + # Only vars with a backing field are serialized; `router` is a switchboard + # over the per-field router vars and has no field of its own. assert set(test_state_dict[test_state.get_name()]) == { - var + FIELD_MARKER for var in test_state.vars + var + FIELD_MARKER for var in (*test_state.base_vars, *test_state.computed_vars) } assert set(test_state.dict(include_computed=False)[test_state.get_name()]) == { var + FIELD_MARKER for var in test_state.base_vars @@ -1225,7 +1239,8 @@ def test_interdependent_state_initial_dict() -> None: s = InterdependentState() state_name = s.get_name() d = s.dict(initial=True)[state_name] - d.pop("router" + FIELD_MARKER) + for router_var in constants.ROUTER_VARS: + d.pop(router_var + FIELD_MARKER) assert d == { "x" + FIELD_MARKER: 0, "v1" + FIELD_MARKER: 0, @@ -1798,19 +1813,19 @@ def dep_v(self) -> int: dict1 = json.loads(json_dumps(ps.dict())) assert dict1[ps.get_full_name()] == { "no_cache_v" + FIELD_MARKER: 1, - "router" + FIELD_MARKER: formatted_router, + **formatted_router_vars, } assert dict1[cs.get_full_name()] == {"dep_v" + FIELD_MARKER: 2} dict2 = json.loads(json_dumps(ps.dict())) assert dict2[ps.get_full_name()] == { "no_cache_v" + FIELD_MARKER: 3, - "router" + FIELD_MARKER: formatted_router, + **formatted_router_vars, } assert dict2[cs.get_full_name()] == {"dep_v" + FIELD_MARKER: 4} dict3 = json.loads(json_dumps(ps.dict())) assert dict3[ps.get_full_name()] == { "no_cache_v" + FIELD_MARKER: 5, - "router" + FIELD_MARKER: formatted_router, + **formatted_router_vars, } assert dict3[cs.get_full_name()] == {"dep_v" + FIELD_MARKER: 6} assert counter == 6 @@ -2695,7 +2710,13 @@ async def test_state_proxy( ( token, { - TestState.get_full_name(): {"router" + FIELD_MARKER: router_data}, + TestState.get_full_name(): { + "rx_router_session" + FIELD_MARKER: router_data.session, + "rx_router_headers" + FIELD_MARKER: router_data.headers, + "rx_router_page" + FIELD_MARKER: router_data._page, + "rx_router_url" + FIELD_MARKER: URLData.from_url(router_data.url), + "rx_router_route_id" + FIELD_MARKER: router_data.route_id, + }, grandchild_state.get_full_name(): { "value2" + FIELD_MARKER: "42", }, @@ -3392,7 +3413,7 @@ class MutableContainsBase(BaseState): assert json.loads(val) == { MutableContainsBase.get_full_name(): { f"items{FIELD_MARKER}": [{"tags": ["123", "456"]}], - f"router{FIELD_MARKER}": formatted_router, + **formatted_router_vars, } } @@ -3691,7 +3712,10 @@ def index(): assert len(emitted_deltas) == 1 + len(expected) first_token, first_delta = emitted_deltas[0] assert first_token == token - assert first_delta[State.get_full_name()].pop("router" + FIELD_MARKER) is not None + first_state_delta = first_delta[State.get_full_name()] + assert first_state_delta.pop("rx_router_url" + FIELD_MARKER) is not None + for router_var in constants.ROUTER_VARS: + first_state_delta.pop(router_var + FIELD_MARKER, None) assert first_delta == exp_is_hydrated(State, False) # Find the deltas containing the test handler's state change @@ -3751,7 +3775,10 @@ def index(): # First delta: router + is_hydrated=False assert len(emitted_deltas) >= 2 first_delta = emitted_deltas[0][1] - assert first_delta[State.get_full_name()].pop("router" + FIELD_MARKER) is not None + first_state_delta = first_delta[State.get_full_name()] + assert first_state_delta.pop("rx_router_url" + FIELD_MARKER) is not None + for router_var in constants.ROUTER_VARS: + first_state_delta.pop(router_var + FIELD_MARKER, None) assert first_delta == exp_is_hydrated(State, False) # Find deltas containing the test handler's state change (num incremented twice) @@ -4028,12 +4055,15 @@ def foo(self) -> str: foo = RouterVarDepState.computed_vars["foo"] State._init_var_dependency_dicts() + # Reading self.router recurses into the router property getter, so the + # dependency lands on each of the per-field router vars. assert foo._deps(objclass=RouterVarDepState) == { - RouterVarDepState.get_full_name(): {"router"} + RouterVarDepState.get_full_name(): set(constants.ROUTER_VARS) } - assert (RouterVarDepState.get_full_name(), "foo") in State._var_dependencies[ - "router" - ] + for router_var in constants.ROUTER_VARS: + assert (RouterVarDepState.get_full_name(), "foo") in State._var_dependencies[ + router_var + ] # Get state from state manager. rx_state = await state_manager.get_state(BaseStateToken(ident=token, cls=State)) @@ -4046,10 +4076,369 @@ def foo(self) -> str: # Reassign router var state.router = state.router - assert rx_state.dirty_vars == {"router"} + assert rx_state.dirty_vars == set(constants.ROUTER_VARS) assert state.dirty_vars == {"foo"} assert parent_state.dirty_substates == {RouterVarDepState.get_name()} + # The locally-defined states above registered themselves in the class-level + # dependency maps on State, which outlive this test. Left behind, a later + # test that dirties a router var on a fresh State tree resolves the stale + # entry and raises on the missing substate. Drop them. + for dep_set in State._var_dependencies.values(): + dep_set.difference_update({ + (RouterVarDepState.get_full_name(), "foo"), + }) + State._potentially_dirty_states.discard(RouterVarDepState.get_full_name()) + + +@pytest.mark.parametrize("name", constants.ROUTER_VARS) +def test_router_field_names_are_reserved(name): + """A substate cannot replace framework-owned router storage. + + The guard is the framework's general inherited-var shadow detection, not + anything router-specific: the router fields live on `BaseState`, so a + substate redeclaring one shadows an inherited var like any other. Note + this covers substates only -- a direct `BaseState` subclass starts its own + root and has no inherited var to shadow. + """ + with pytest.raises(BaseVarShadowsInheritedVarError): + type( + "InvalidRouterState", + (State,), + {"__module__": __name__, "__annotations__": {name: int}, name: 1}, + ) + + +def test_router_var_dep_legacy_string() -> None: + """An explicit deps=["router"] still fires when any router var changes. + + The `router` base var was split into per-field vars; a legacy string dep + on "router" is expanded to all of them (with a deprecation warning). + """ + + class LegacyRouterDepState(State): + """A state with a legacy string dependency on the router var.""" + + @rx.var(deps=["router"], auto_deps=False) + def foo(self) -> str: + return self.router.url.path + + for router_var in constants.ROUTER_VARS: + assert ( + LegacyRouterDepState.get_full_name(), + "foo", + ) in State._var_dependencies[router_var] + assert "router" not in State._var_dependencies + + # Drop the class-level registrations this locally-defined state made; see + # the note in test_router_var_dep. + for dep_set in State._var_dependencies.values(): + dep_set.discard((LegacyRouterDepState.get_full_name(), "foo")) + State._potentially_dirty_states.discard(LegacyRouterDepState.get_full_name()) + + +def test_router_var_dep_legacy_string_still_compiles() -> None: + """An app declaring deps=["router"] must still pass dependency validation. + + `_validate_var_dependencies` checks the raw `_deps()` names against + `state_cls.vars` rather than the expanded registrations, so the deprecated + string only keeps working while `router` is itself listed as a var. + """ + + class LegacyRouterCompileState(State): + """A state with a legacy string dependency on the router var.""" + + @rx.var(deps=["router"], auto_deps=False) + def foo(self) -> str: + return self.router.url.path + + assert constants.ROUTER in State.vars + # Raises VarDependencyError if the dependency does not resolve to a var. + App()._validate_var_dependencies() + + for dep_set in State._var_dependencies.values(): + dep_set.discard((LegacyRouterCompileState.get_full_name(), "foo")) + State._potentially_dirty_states.discard(LegacyRouterCompileState.get_full_name()) + + +@pytest.mark.asyncio +async def test_get_var_value_of_the_whole_router() -> None: + """`get_var_value(State.router)` must hand back the composed RouterData. + + The switchboard renders as an object literal over the five per-field vars, + so it has no field of its own to read. Without naming the `router` + attribute it stands for, this raised UnretrievableVarValueError, while a + state with a single `router` base var resolved it. + """ + state = State(_reflex_internal_init=True) # pyright: ignore [reportCallIssue] + + router = await state.get_var_value(State.router) + + assert isinstance(router, RouterData) + # The per-field vars resolve too, which the pre-split single var could not do. + assert await state.get_var_value(State.router.route_id) == router.route_id + assert ( + await state.get_var_value(State.router.session) + ).client_token == router.session.client_token + + +def test_router_var_dep_does_not_warn_for_the_var_form( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Only the legacy string form is deprecated, and it must name the var. + + `State.router` carries the per-field names as well as `router` itself, so + the expansion has nothing to warn about; `deps=["router"]` arrives with + only `router` and does. The warning has to identify the computed var, + because the lazy dep scan means the reported caller frame is unrelated to + the declaration. + """ + # `console.deprecate` logs and dedupes rather than printing, so record the + # calls instead of scraping output. + from reflex import state as state_module + + deprecations: list[str] = [] + monkeypatch.setattr( + state_module.console, + "deprecate", + lambda *, feature_name, **kwargs: deprecations.append(feature_name), + ) + + class VarFormRouterDepState(State): + """A state depending on the router through the Var.""" + + @rx.var(deps=[State.router], auto_deps=False) + def from_var(self) -> str: + return "" + + assert deprecations == [] + + class StringFormRouterDepState(State): + """A state depending on the router through the legacy string.""" + + @rx.var(deps=["router"], auto_deps=False) + def from_string(self) -> str: + return "" + + assert len(deprecations) == 1 + assert "StringFormRouterDepState.from_string" in deprecations[0] + + for dep_set in State._var_dependencies.values(): + dep_set.discard((VarFormRouterDepState.get_full_name(), "from_var")) + dep_set.discard((StringFormRouterDepState.get_full_name(), "from_string")) + State._potentially_dirty_states.discard(VarFormRouterDepState.get_full_name()) + State._potentially_dirty_states.discard(StringFormRouterDepState.get_full_name()) + + +def test_router_var_dep_whole_router() -> None: + """deps=[State.router] must track every per-field router var. + + The switchboard is composed of the five per-field vars, so its VarData + must carry all five field names; if it reported only one, a cached var + declaring the whole router would go stale when any other router field + changed -- a reconnect updates the session without touching the URL, for + instance. + """ + + class WholeRouterDepState(State): + """A state depending on the whole router var.""" + + @rx.var(deps=[State.router], auto_deps=False) + def summary(self) -> str: + return "" + + # The declared set also names `router` itself, the switchboard the five + # fields were read through; it is expanded away before registration. + assert WholeRouterDepState.computed_vars["summary"]._static_deps == { + State.get_full_name(): {constants.ROUTER, *constants.ROUTER_VARS} + } + for router_var in constants.ROUTER_VARS: + assert ( + WholeRouterDepState.get_full_name(), + "summary", + ) in State._var_dependencies[router_var] + # `router` has no backing field, so nothing may be registered against it -- + # it would never be dirtied and the dependent var would go stale. + assert ( + WholeRouterDepState.get_full_name(), + "summary", + ) not in State._var_dependencies.get(constants.ROUTER, set()) + + # Drop the class-level registrations; see the note in test_router_var_dep. + for dep_set in State._var_dependencies.values(): + dep_set.discard((WholeRouterDepState.get_full_name(), "summary")) + State._potentially_dirty_states.discard(WholeRouterDepState.get_full_name()) + + +def test_router_is_listed_as_a_var_and_inherited_by_substates() -> None: + """`router` is usable as a Var, so it is listed in vars and inherited. + + It has no backing field of its own, so it must stay out of anything that + serializes vars: the switchboard resolves to the root state's per-field + base vars instead. + """ + + class RouterVarListingState(State): + """A substate that only inherits the router.""" + + assert constants.ROUTER in State.vars + assert constants.ROUTER in RouterVarListingState.inherited_vars + assert constants.ROUTER not in State.base_vars + assert constants.ROUTER not in State.computed_vars + + # The substate's entry is the root's switchboard, resolving to the root's + # per-field base vars rather than to anything on the substate. + router_var = RouterVarListingState.vars[constants.ROUTER] + assert isinstance(router_var, RouterDataVar) + assert router_var.equals(State.router) + assert str(router_var.route_id) == str(State.rx_router_route_id) + + +def test_update_router_vars_ignores_omitted_static_keys( + test_state: TestState, +) -> None: + """A navigation-only payload must not reset the connection-scoped vars. + + A router_data carrying only the navigation keys says nothing about the + session or headers; treating the omission as a change would wipe them to + their defaults and ship a destructive delta. + + Args: + test_state: A state. + """ + full_router_data = { + RouteVar.PATH: "/a", + RouteVar.ORIGIN: "/a", + RouteVar.QUERY: {}, + RouteVar.CLIENT_TOKEN: "tok", + RouteVar.SESSION_ID: "sid1", + RouteVar.CLIENT_IP: "127.0.0.1", + RouteVar.HEADERS: {"origin": "http://localhost:3000", "cookie": "a=b"}, + } + test_state._update_router_vars(full_router_data, {}) + test_state._clean() + + navigation_only = { + RouteVar.PATH: "/b", + RouteVar.ORIGIN: "/b", + RouteVar.QUERY: {}, + } + merged = test_state._update_router_vars(navigation_only, full_router_data) + assert test_state.dirty_vars & set(constants.ROUTER_VARS) == { + "rx_router_page", + "rx_router_url", + "rx_router_route_id", + } + assert test_state.router.session.client_token == "tok" + assert test_state.router.session.session_id == "sid1" + assert test_state.router.headers.cookie == "a=b" + # The rebuilt navigation vars keep the host from the headers the payload + # omitted, rather than being reconstructed from the partial dict alone. + assert test_state.router.url.origin == "http://localhost:3000" + assert test_state.router.url.path == "/b" + assert test_state.router.page.host == "http://localhost:3000" + # The merged data is what the caller stores, so the omitted keys are still + # there to compare against next time. + assert merged[RouteVar.CLIENT_TOKEN] == "tok" + assert merged[RouteVar.HEADERS] == full_router_data[RouteVar.HEADERS] + + # A second consecutive partial payload still has the full picture. + test_state._clean() + merged2 = test_state._update_router_vars( + {RouteVar.PATH: "/c", RouteVar.ORIGIN: "/c", RouteVar.QUERY: {}}, merged + ) + assert test_state.router.url.origin == "http://localhost:3000" + assert test_state.router.session.client_token == "tok" + assert merged2[RouteVar.HEADERS] == full_router_data[RouteVar.HEADERS] + + +def test_update_router_vars_non_origin_header_leaves_navigation_clean( + test_state: TestState, +) -> None: + """Only the origin header feeds the page/URL, so other headers leave them alone. + + Args: + test_state: A state. + """ + router_data = { + RouteVar.PATH: "/a", + RouteVar.ORIGIN: "/a", + RouteVar.QUERY: {}, + RouteVar.HEADERS: {"origin": "http://localhost:3000", "cookie": "a=b"}, + } + test_state._update_router_vars(router_data, {}) + test_state._clean() + + new_cookie = { + **router_data, + RouteVar.HEADERS: {"origin": "http://localhost:3000", "cookie": "c=d"}, + } + test_state._update_router_vars(new_cookie, router_data) + assert test_state.dirty_vars & set(constants.ROUTER_VARS) == {"rx_router_headers"} + + +def test_update_router_vars_granular_delta(test_state: TestState) -> None: + """_update_router_vars only dirties the vars whose source keys changed. + + Args: + test_state: A state. + """ + full_router_data = { + RouteVar.PATH: "/a", + RouteVar.ORIGIN: "/a", + RouteVar.QUERY: {}, + RouteVar.CLIENT_TOKEN: "tok", + RouteVar.SESSION_ID: "sid1", + RouteVar.CLIENT_IP: "127.0.0.1", + RouteVar.HEADERS: {"origin": "http://localhost:3000"}, + } + test_state._update_router_vars(full_router_data, {}) + assert set(constants.ROUTER_VARS) <= test_state.dirty_vars + test_state._clean() + + # Navigation: only the navigation-scoped vars are rebuilt. + nav_router_data = {**full_router_data, RouteVar.PATH: "/b", RouteVar.ORIGIN: "/b"} + test_state._update_router_vars(nav_router_data, full_router_data) + assert test_state.dirty_vars & set(constants.ROUTER_VARS) == { + "rx_router_page", + "rx_router_url", + "rx_router_route_id", + } + assert test_state.router.url.path == "/b" + assert test_state.router.session.session_id == "sid1" + test_state._clean() + + # Reconnect: only the session var is rebuilt. + reconnect_router_data = {**nav_router_data, RouteVar.SESSION_ID: "sid2"} + test_state._update_router_vars(reconnect_router_data, nav_router_data) + assert test_state.dirty_vars & set(constants.ROUTER_VARS) == {"rx_router_session"} + assert test_state.router.session.session_id == "sid2" + test_state._clean() + + # Header change: headers, and the page/URL whose host derives from them. + # route_id derives from the path alone, so it is left clean. + new_headers_router_data = { + **reconnect_router_data, + RouteVar.HEADERS: {"origin": "http://example.com"}, + } + test_state._update_router_vars(new_headers_router_data, reconnect_router_data) + assert test_state.dirty_vars & set(constants.ROUTER_VARS) == { + "rx_router_headers", + "rx_router_page", + "rx_router_url", + } + assert test_state.router.url.origin == "http://example.com" + test_state._clean() + + # Keys that differ but derive the same values leave every var clean: an + # absent key and an empty one both produce the default, and dirtying on + # that alone would mark the state touched and persist it. + equivalent_router_data = { + k: v for k, v in new_headers_router_data.items() if k != RouteVar.QUERY + } + test_state._update_router_vars(equivalent_router_data, new_headers_router_data) + assert test_state.dirty_vars & set(constants.ROUTER_VARS) == set() + @pytest.mark.asyncio async def test_setvar( @@ -5769,3 +6158,82 @@ class ReannotatingChild(ReannotatedParent): reannotated_value: int # pyright: ignore[reportGeneralTypeIssues] assert isinstance(ReannotatingChild.reannotated_value, Var) + + +def test_composite_var_dep_tracks_fields_in_every_state(): + """A dependency on a var spanning two states must track both states' fields. + + `VarData` groups field names by the state that owns them, so merging a var + built from `StateA.a_field` with one built from `StateB.b_field` keeps + both. Before that grouping the merge kept only the first state's fields and + a computed var depending on the composite went stale whenever the other + state changed. + """ + from reflex_base.vars.base import Var, VarData + + class _CompositeDepStateA(rx.State): + a_field: str = "a" + + class _CompositeDepStateB(rx.State): + b_field: str = "b" + + composite = Var( + "combo", + _var_data=VarData.merge( + cast("Var", _CompositeDepStateA.a_field)._get_all_var_data(), + cast("Var", _CompositeDepStateB.b_field)._get_all_var_data(), + ), + ) + + a_name = _CompositeDepStateA.get_full_name() + b_name = _CompositeDepStateB.get_full_name() + assert dict(composite._dependency_fields()) == { + a_name: ("a_field",), + b_name: ("b_field",), + } + + class _CompositeDepConsumer(rx.State): + @rx.var(deps=[composite], cache=True) + def combined(self) -> str: + return "x" + + static_deps = _CompositeDepConsumer.__dict__["combined"]._static_deps + assert "a_field" in static_deps.get(a_name, set()) + assert "b_field" in static_deps.get(b_name, set()) + + # The consumer registered itself in both source states' class-level + # dependency maps, which outlive this test. Left behind, a later test that + # dirties a_field or b_field resolves the stale entry and raises on the + # missing substate. Drop them. + consumer_name = _CompositeDepConsumer.get_full_name() + for state_cls in (_CompositeDepStateA, _CompositeDepStateB): + for dep_set in state_cls._var_dependencies.values(): + dep_set.difference_update({(consumer_name, "combined")}) + state_cls._potentially_dirty_states.discard(consumer_name) + + +def test_setstate_drops_the_legacy_router_entry(): + """Unpickling a pre-split state must not route `router` through the setter. + + Older pickles stored the whole `RouterData` under `router`, which is now a + descriptor. Restoring it with `object.__setattr__` would shadow that + descriptor on the instance; assigning it would decompose into the per-field + vars and resurrect stale connection data. The schema check in + `_deserialize` discards such states anyway, so the entry is simply dropped. + """ + state = BaseState(_reflex_internal_init=True) # pyright: ignore [reportCallIssue] + legacy = { + "parent_state": None, + "substates": {}, + "router": RouterData.from_router_data({ + constants.RouteVar.CLIENT_TOKEN: "stale-token", + }), + "dirty_vars": set(), + } + + state.__setstate__(legacy) + + # The entry is gone rather than shadowing the descriptor... + assert "router" not in state.__dict__ + # ...and `router` still resolves through the switchboard to live fields. + assert state.router.session.client_token == "" diff --git a/tests/units/utils/test_format.py b/tests/units/utils/test_format.py index 83e718411fd..3b89da9c455 100644 --- a/tests/units/utils/test_format.py +++ b/tests/units/utils/test_format.py @@ -657,9 +657,9 @@ def test_format_query_params(input, output): assert format.format_query_params(input) == output -formatted_router = { - "route_id": "", - "url": { +formatted_router_vars = { + "rx_router_route_id" + FIELD_MARKER: "", + "rx_router_url" + FIELD_MARKER: { "scheme": "", "netloc": "", "origin": "://", @@ -669,8 +669,12 @@ def test_format_query_params(input, output): "fragment": "", "href": "", }, - "session": {"client_token": "", "client_ip": "", "session_id": ""}, - "headers": { + "rx_router_session" + FIELD_MARKER: { + "client_token": "", + "client_ip": "", + "session_id": "", + }, + "rx_router_headers" + FIELD_MARKER: { "host": "", "origin": "", "upgrade": "", @@ -686,7 +690,7 @@ def test_format_query_params(input, output): "accept_language": "", "raw_headers": {}, }, - "page": { + "rx_router_page" + FIELD_MARKER: { "host": "", "path": "", "raw_path": "", @@ -720,7 +724,7 @@ def test_format_query_params(input, output): "obj" + FIELD_MARKER: {"prop1": 42, "prop2": "hello"}, "sum" + FIELD_MARKER: 3.15, "upper" + FIELD_MARKER: "", - "router" + FIELD_MARKER: formatted_router, + **formatted_router_vars, "asynctest" + FIELD_MARKER: 0, }, ChildState.get_full_name(): { @@ -742,7 +746,7 @@ def test_format_query_params(input, output): "dt" + FIELD_MARKER: "1989-11-09 18:53:00+01:00", "t" + FIELD_MARKER: "18:53:00+01:00", "td" + FIELD_MARKER: "11 days, 0:11:00", - "router" + FIELD_MARKER: formatted_router, + **formatted_router_vars, }, }, ), diff --git a/tests/units/vars/test_dep_tracking.py b/tests/units/vars/test_dep_tracking.py index 6c3c1782316..7f5fd27c482 100644 --- a/tests/units/vars/test_dep_tracking.py +++ b/tests/units/vars/test_dep_tracking.py @@ -300,6 +300,27 @@ async def func_with_get_var_value(self: DependencyTestState): assert tracker.dependencies == expected_deps +@pytest.mark.skipif( + sys.version_info < (3, 11), reason="Requires Python 3.11+ for positions" +) +def test_get_var_value_tracks_all_composed_fields(): + """Composed get_var_value arguments register every state field they read.""" + composed_var = DependencyTestState.count + DependencyTestState.items.length() + + async def composed(self: DependencyTestState): + """Read a composite expression. + + Returns: + The combined field value. + """ + return await self.get_var_value(composed_var) + + tracker = DependencyTracker(composed, DependencyTestState) + assert tracker.dependencies == { + DependencyTestState.get_full_name(): {"count", "items"} + } + + @pytest.mark.skipif( sys.version_info < (3, 11), reason="Requires Python 3.11+ for positions" )