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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions docs/library/forms/form.md
Original file line number Diff line number Diff line change
Expand Up @@ -298,7 +298,7 @@ class DynamicFormState(rx.State):
]

@rx.event
def add_field(self, form_data: dict):
def add_form_field(self, form_data: dict):
new_field = form_data.get("new_field")
if not new_field:
return
Expand Down Expand Up @@ -331,7 +331,7 @@ def dynamic_form():
rx.input(placeholder="New Field", name="new_field"),
rx.button("+", type="submit"),
),
on_submit=DynamicFormState.add_field,
on_submit=DynamicFormState.add_form_field,
reset_on_submit=True,
),
rx.divider(),
Expand Down
6 changes: 6 additions & 0 deletions docs/state/overview.md
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,12 @@ A state class is made up of two parts: vars and event handlers.

**Event handlers** are functions that modify these vars in response to events.

State declarations cannot reuse framework method or bookkeeping names, such as
`get_state`, `_get_was_touched`, or `dirty_vars`. Reflex checks these names when
creating a state class and when adding vars, event handlers, or route arguments
dynamically. Rename a conflicting declaration and update its references. Ordinary
backend names such as `_count` remain supported.

These are the main concepts to understand how state works in Reflex:

```python eval
Expand Down
1 change: 1 addition & 0 deletions news/+reserved-state-names.breaking.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
State vars, event handlers, and dynamic route arguments now reject names reserved by framework methods and bookkeeping before registration. Rename conflicting members.
122 changes: 122 additions & 0 deletions reflex/istate/validation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,122 @@
"""Validate the framework namespace before constructing or extending a state."""

from functools import cache
from types import FunctionType
from typing import Any

from reflex_base.utils.compat import annotations_from_namespace
from reflex_base.utils.exceptions import (
EventHandlerShadowsBuiltInStateMethodError,
StateValueError,
)
from reflex_base.vars.base import (
BaseStateMeta,
EvenMoreBasicBaseState,
_linearize_bases,
)

_FIELD_MAP_NAMES = frozenset({"__fields__", "__own_fields__", "__inherited_fields__"})


@cache
def _reserved_state_members() -> dict[str, Any]:
"""Return framework members, excluding state vars and Python protocols.

Returns:
Reserved names and their original descriptors, without invoking them.
"""
# BaseState must exist before its namespace can be inspected.
from reflex.state import BaseState

members = {}
for base in reversed(BaseState.__mro__[:-1]):
namespace = vars(base)
members.update(
(name, namespace.get(name))
for name in namespace.keys() | annotations_from_namespace(namespace).keys()
if not name.startswith("__") or name in _FIELD_MAP_NAMES
)
for name, field in BaseState.__fields__.items():
if field.is_var:
members.pop(name, None)
Comment thread
masenf marked this conversation as resolved.
return members


def _validate_state_name(name: str, value: Any = None) -> None:
"""Reject declarations that replace framework methods or bookkeeping.

Args:
name: The declared or dynamically registered name.
value: The raw class declaration, when available.

Raises:
StateValueError: If a declaration uses a reserved name.
EventHandlerShadowsBuiltInStateMethodError: If a method overrides a builtin.
"""
members = _reserved_state_members()
if name not in members:
return
method = value.__func__ if isinstance(value, (classmethod, staticmethod)) else value
if isinstance(method, FunctionType):
if value is members[name] or getattr(method, "__override_base_method__", False):
return
msg = f"The event handler name `{name}` shadows a builtin State method; use a different name instead"
raise EventHandlerShadowsBuiltInStateMethodError(msg)
msg = f"State name `{name}` is reserved by BaseState; use a different name instead."
raise StateValueError(msg)

Comment thread
FarhanAliRaza marked this conversation as resolved.

def _validate_inherited_members(base: type, seen: set[str]) -> None:
"""Check the members a Python mixin or model base adds to a state.

Args:
base: A base class that is not itself a validated state.
seen: Names an earlier base already provides in the MRO.
"""
is_model = isinstance(base, BaseStateMeta)
if is_model:
# Model fields are inherited even when an earlier base masks their
# class attributes in the MRO.
for member in base.__own_fields__:
_validate_state_name(member)
seen.update(base.__own_fields__)
for member, value in vars(base).items():
if member not in seen and not (
is_model and (member in _FIELD_MAP_NAMES or member == "_mixin")
):
_validate_state_name(member, value)


class _StateMeta(BaseStateMeta):
"""Check state declarations before field collection and subclass initialization."""

def __new__(
cls,
name: str,
bases: tuple[type, ...],
namespace: dict[str, Any],
mixin: bool = False,
) -> type:
"""Construct a state after checking its declarations and Python mixins.

Args:
name: The class name.
bases: The parent classes.
namespace: The unmodified class namespace.
mixin: Whether the class is a state mixin.

Returns:
The validated state class.
"""
if any(isinstance(base, _StateMeta) for base in bases):
seen = namespace.keys() | annotations_from_namespace(namespace).keys()
for member in seen:
_validate_state_name(member, namespace.get(member))
for base in _linearize_bases(bases):
if not isinstance(base, _StateMeta) and base not in (
EvenMoreBasicBaseState,
object,
):
_validate_inherited_members(base, seen)
seen.update(vars(base))
return super().__new__(cls, name, bases, namespace, mixin=mixin)
98 changes: 24 additions & 74 deletions reflex/state.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,7 +45,6 @@
ComputedVarShadowsStateVarError,
DynamicComponentInvalidSignatureError,
DynamicRouteArgShadowsStateVarError,
EventHandlerShadowsBuiltInStateMethodError,
ReflexRuntimeError,
SetUndefinedStateVarError,
StateMismatchError,
Expand Down Expand Up @@ -78,6 +77,7 @@
from reflex.istate.proxy import ImmutableMutableProxy as ImmutableMutableProxy
from reflex.istate.proxy import MutableProxy, is_mutable_type
from reflex.istate.storage import ClientStorageBase
from reflex.istate.validation import _StateMeta, _validate_state_name
from reflex.utils import console, format, types
from reflex.utils.exec import is_testing_env

Expand Down Expand Up @@ -382,8 +382,8 @@ def _is_user_descriptor(value: Any) -> bool:
# Instance bookkeeping fields and framework methods read on every event. They
# bypass the var-resolution logic below, so nothing stored in `_backend_vars`
# (e.g. `_reflex_internal_links`) or delegated to the parent (`router_data`)
# may appear here. A subclass that defines one of these names itself (as a var
# or an event handler) drops it from its own `_fast_attr_names`.
# may appear here. A subclass that overrides one of these methods drops the
# name from its own `_fast_attr_names`.
_FRAMEWORK_ATTR_NAMES = frozenset({
"dirty_vars",
"dirty_substates",
Expand Down Expand Up @@ -427,7 +427,7 @@ def _is_user_descriptor(value: Any) -> bool:
})


class BaseState(EvenMoreBasicBaseState):
class BaseState(EvenMoreBasicBaseState, metaclass=_StateMeta):
"""The state of the app."""

# A map from the var name to the var.
Expand Down Expand Up @@ -617,9 +617,6 @@ def __init_subclass__(cls, mixin: bool = False, **kwargs):
# Validate the module name.
cls._validate_module_name()

# Event handlers should not shadow builtin state methods.
cls._check_overridden_methods()

# Computed vars should not shadow builtin state props.
cls._check_overridden_basevars()

Expand Down Expand Up @@ -784,38 +781,15 @@ def __init_subclass__(cls, mixin: bool = False, **kwargs):
cls._var_dependencies = {}
cls._init_var_dependency_dicts()

cls._prune_fast_attr_names()

all_base_state_classes[cls.get_full_name()] = None

@classmethod
def _prune_fast_attr_names(cls) -> None:
"""Recompute which framework attribute names this state tree may fast-path.

A name the state defines (as a var, a backend var, an event handler or
a marked method override) must keep going through the full lookup in
``_get_attribute``. The set is rebuilt from the parent's current set
minus this class's own names, then recomputed for every substate, so a
var or handler registered after class creation (dynamic route args,
``add_var``, ...) drops the name for the whole subtree that inherits it.
"""
# A marked override of a framework method must keep the full lookup.
parent_state = cls.get_parent_state()
inherited = (
cls._fast_attr_names = (
parent_state._fast_attr_names
if parent_state is not None
else _FRAMEWORK_ATTR_NAMES
)
cls._fast_attr_names = inherited - (
_FRAMEWORK_ATTR_NAMES
& (
set(cls.__dict__)
| set(cls.vars)
| set(cls.backend_vars)
| set(cls.event_handlers)
)
)
for substate_class in cls.get_substates():
substate_class._prune_fast_attr_names()
) - cls.__dict__.keys()

all_base_state_classes[cls.get_full_name()] = None

@classmethod
def _add_event_handler(
Expand All @@ -829,10 +803,10 @@ def _add_event_handler(
name: The name of the event handler.
fn: The function to call when the event is triggered.
"""
_validate_state_name(name)
handler = cls._create_event_handler(fn)
cls.event_handlers[name] = handler
setattr(cls, name, handler)
cls._prune_fast_attr_names()

@staticmethod
def _copy_fn(fn: Callable) -> Callable:
Expand Down Expand Up @@ -1063,29 +1037,6 @@ def _iter_functions(cls) -> Iterator[tuple[str, FunctionType]]:
if isinstance(value, FunctionType):
yield name, value

@classmethod
def _check_overridden_methods(cls):
"""Check for shadow methods and raise error if any.

Raises:
EventHandlerShadowsBuiltInStateMethodError: When an event handler shadows an inbuilt state method.
"""
overridden_methods = set()
state_base_functions = cls._get_base_functions()
for name, method in cls._iter_functions():
# Check if the method is overridden and not a dunder method
if (
not name.startswith("__")
and method.__name__ in state_base_functions
and state_base_functions[method.__name__] != method
and not getattr(method, "__override_base_method__", False)
):
overridden_methods.add(method.__name__)

for method_name in overridden_methods:
msg = f"The event handler name `{method_name}` shadows a builtin State method; use a different name instead"
raise EventHandlerShadowsBuiltInStateMethodError(msg)

@classmethod
def _check_overridden_basevars(cls):
"""Check for shadow base vars and raise error if any.
Expand Down Expand Up @@ -1332,6 +1283,18 @@ def _init_var(cls, name: str, prop: Var):
cls._create_setter(name, prop)
cls._set_default_value(name, prop)

@classmethod
def add_field(cls, name: str, var: Var, default_value: Any):
"""Validate a dynamically added field before updating the field map.

Args:
name: The name of the field to add.
var: The variable to add a field for.
default_value: The default value of the field.
"""
_validate_state_name(name)
super().add_field(name, var, default_value)

@classmethod
def add_var(cls, name: str, type_: Any, default_value: Any = None):
"""Add dynamically a variable to the State.
Expand Down Expand Up @@ -1373,7 +1336,6 @@ def add_var(cls, name: str, type_: Any, default_value: Any = None):
# let substates know about the new variable
for substate_class in cls.get_substates():
substate_class.vars.setdefault(name, var)
cls._prune_fast_attr_names()

# Reinitialize dependency tracking dicts.
cls._init_var_dependency_dicts()
Expand Down Expand Up @@ -1484,19 +1446,6 @@ def _get_var_default(cls, name: str, annotation_value: Any) -> Any:
except TypeError:
return None

@staticmethod
def _get_base_functions() -> builtins.dict[str, FunctionType]:
"""Get all functions of the state class excluding dunder methods.

Returns:
The functions of rx.State class as a dict.
"""
return {
func[0]: func[1]
for func in inspect.getmembers(BaseState, predicate=inspect.isfunction)
if not func[0].startswith("__")
}

@classmethod
def _update_substate_inherited_vars(cls, vars_to_add: builtins.dict[str, Var]):
"""Update the inherited vars of substates recursively when new vars are added.
Expand All @@ -1517,7 +1466,6 @@ def _update_substate_inherited_vars(cls, vars_to_add: builtins.dict[str, Var]):
substate_class._update_substate_inherited_vars(vars_to_add)
# Reinitialize dependency tracking dicts.
cls._init_var_dependency_dicts()
cls._prune_fast_attr_names()

@classmethod
def _dynamic_route_arg_types(cls) -> builtins.dict[str, str]:
Expand Down Expand Up @@ -1549,6 +1497,8 @@ def setup_dynamic_args(cls, args: builtins.dict[str, str]):
if not args:
return

for name in args:
_validate_state_name(name)
cls._check_overwritten_dynamic_args(list(args.keys()))

def argsingle_factory(param: str):
Expand Down
23 changes: 0 additions & 23 deletions tests/units/istate/manager/test_redis.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,29 +115,6 @@ async def test_basic_get_set(
)


async def test_set_state_with_shadowed_touched_method(clean_registration_context):
"""Persist a backend var that shadows the touched-state method.

Args:
clean_registration_context: A fresh, empty registration context.
"""

class ShadowState(BaseState):
"""State with an intentional framework-method collision."""

_get_was_touched: int = 7

state = ShadowState()
state._get_was_touched = 8
manager = StateManagerRedis(redis=mock_redis())
token = BaseStateToken(ident="shadowed", cls=ShadowState)

await manager.set_state(token, state)

restored = BaseState._deserialize(data=await manager.redis.get(str(token)))
assert restored._get_was_touched == 8


async def test_modify(
state_manager_redis: StateManagerRedis,
root_state: type[RedisTestState],
Expand Down
Loading
Loading