diff --git a/news/6743.performance.md b/news/6743.performance.md new file mode 100644 index 00000000000..9be8f6e8bbe --- /dev/null +++ b/news/6743.performance.md @@ -0,0 +1 @@ +Dirty propagation now only walks newly-dirty vars per mutation and skips the computed var expiry scan for classes with no interval vars, roughly halving per-mutation overhead in states with computed vars. diff --git a/packages/reflex-base/news/6743.performance.md b/packages/reflex-base/news/6743.performance.md new file mode 100644 index 00000000000..9be8f6e8bbe --- /dev/null +++ b/packages/reflex-base/news/6743.performance.md @@ -0,0 +1 @@ +Dirty propagation now only walks newly-dirty vars per mutation and skips the computed var expiry scan for classes with no interval vars, roughly halving per-mutation overhead in states with computed vars. diff --git a/packages/reflex-base/src/reflex_base/utils/types.py b/packages/reflex-base/src/reflex_base/utils/types.py index f33e5c8c28b..95927a53f23 100644 --- a/packages/reflex-base/src/reflex_base/utils/types.py +++ b/packages/reflex-base/src/reflex_base/utils/types.py @@ -168,6 +168,8 @@ def __call__( "_was_touched", "_mixin", "_mutable_proxy_cache", + "_propagated_dirty_vars", + "_propagated_generation", } diff --git a/packages/reflex-base/src/reflex_base/vars/base.py b/packages/reflex-base/src/reflex_base/vars/base.py index d50a7ea6666..8cdeeb8cd64 100644 --- a/packages/reflex-base/src/reflex_base/vars/base.py +++ b/packages/reflex-base/src/reflex_base/vars/base.py @@ -2255,6 +2255,19 @@ def is_computed_var(obj: Any) -> TypeGuard[ComputedVar]: return isinstance(obj, FakeComputedVarBaseClass) +# Incremented whenever a cached computed var is recomputed. State dirty +# propagation uses this to know when an already-propagated dependency needs to +# be re-propagated (a recompute re-materializes a cache that a later mutation +# of its dependencies must invalidate again). +_computed_var_recompute_generation: int = 0 + + +def _bump_computed_var_recompute_generation() -> None: + """Record that a cached computed var was recomputed.""" + global _computed_var_recompute_generation + _computed_var_recompute_generation += 1 + + @dataclasses.dataclass( eq=False, frozen=True, @@ -2598,6 +2611,7 @@ def __get__(self, instance: BaseState | None, owner: type): instance._was_touched = True # Set the last updated timestamp on the state instance. setattr(instance, self._last_updated_attr, datetime.datetime.now()) + _bump_computed_var_recompute_generation() value = getattr(instance, self._cache_attr) # Only validate the return type when the value was just computed. self._check_deprecated_return_type(instance, value) @@ -2863,6 +2877,7 @@ async def _awaitable_result(instance: BaseState = instance) -> RETURN_TYPE: instance._was_touched = True # Set the last updated timestamp on the state instance. setattr(instance, self._last_updated_attr, datetime.datetime.now()) + _bump_computed_var_recompute_generation() value = getattr(instance, self._cache_attr) # Only validate the return type when the value was just computed. self._check_deprecated_return_type(instance, value) diff --git a/reflex/state.py b/reflex/state.py index a4e169945d0..09e56802650 100644 --- a/reflex/state.py +++ b/reflex/state.py @@ -55,6 +55,7 @@ from reflex_base.utils.serializers import serializer from reflex_base.utils.types import _isinstance, _validation_depth from reflex_base.vars import Field, VarData, field +from reflex_base.vars import base as reflex_base_vars_base from reflex_base.vars.base import ( ComputedVar, DynamicRouteVar, @@ -368,6 +369,7 @@ def _is_user_descriptor(value: Any) -> bool: "_always_dirty_computed_vars", "_always_dirty_substates", "_potentially_dirty_states", + "_interval_computed_vars", }) @@ -407,6 +409,11 @@ class BaseState(EvenMoreBasicBaseState): # Set of states which might need to be recomputed if vars in this state change. _potentially_dirty_states: ClassVar[set[str]] = set() + # Names of computed vars with an update interval, refreshed whenever + # computed_vars changes. Empty for most classes, which lets dirty + # propagation skip the per-mutation expiry scan. + _interval_computed_vars: ClassVar[tuple[str, ...]] = () + # The parent state. parent_state: BaseState | None = field(default=None, is_var=False) @@ -442,6 +449,12 @@ class BaseState(EvenMoreBasicBaseState): _mutable_proxy_cache: builtins.dict[str, MutableProxy] = field( default_factory=builtins.dict, is_var=False ) + # Dirty vars whose dependency closure was already propagated this cycle. + # Transient: cleared by _clean and never pickled. + _propagated_dirty_vars: set[str] = field(default_factory=set, is_var=False) + + # Recompute generation the propagation frontier was built in. + _propagated_generation: int = field(default=0, is_var=False) # A special event handler for setting base vars. setvar: ClassVar[EventHandler] @@ -723,9 +736,22 @@ def __init_subclass__(cls, mixin: bool = False, **kwargs): # Initialize per-class var dependency tracking. cls._var_dependencies = {} cls._init_var_dependency_dicts() + cls._refresh_interval_computed_vars() all_base_state_classes[cls.get_full_name()] = None + @classmethod + def _refresh_interval_computed_vars(cls) -> None: + """Recompute the names of computed vars that have an update interval. + + Must be called whenever cls.computed_vars is modified. + """ + cls._interval_computed_vars = tuple( + name + for name, cvar in cls.computed_vars.items() + if cvar._update_interval is not None + ) + @classmethod def _add_event_handler( cls, @@ -830,6 +856,7 @@ def computed_var_func(state: Self): setattr(cls, unique_var_name, computed_var_func_arg) cls.computed_vars[unique_var_name] = computed_var_func_arg + cls._refresh_interval_computed_vars() cls.vars[unique_var_name] = computed_var_func_arg cls._update_substate_inherited_vars({unique_var_name: computed_var_func_arg}) cls._always_dirty_computed_vars.add(unique_var_name) @@ -1403,6 +1430,7 @@ def inner_func(self: BaseState) -> list[str]: # Update tracking dicts. cls.computed_vars.update(dynamic_vars) + cls._refresh_interval_computed_vars() cls.vars.update(dynamic_vars) cls._update_substate_inherited_vars(dynamic_vars) @@ -1812,14 +1840,31 @@ async def get_var_value(self, var: Var[VAR_TYPE]) -> VAR_TYPE: def _mark_dirty_computed_vars(self) -> None: """Mark ComputedVars that need to be recalculated based on dirty_vars.""" - # Append expired computed vars to dirty_vars to trigger recalculation - self.dirty_vars.update(self._expired_computed_vars()) + if self._interval_computed_vars: + # Append expired computed vars to dirty_vars to trigger recalculation + self.dirty_vars.update(self._expired_computed_vars()) # Append always dirty computed vars to dirty_vars to trigger recalculation self.dirty_vars.update(self._always_dirty_computed_vars) - dirty_vars = self.dirty_vars - while dirty_vars: - calc_vars, dirty_vars = dirty_vars, set() + # Track which dirty vars already had their dependency closure + # propagated, so repeated mutations only process newly-dirty vars. + # Recomputing any cached var re-materializes a cache that a later + # mutation must invalidate again, so the frontier is only valid for + # the recompute generation it was built in. + # Go through __dict__ (not setattr) so this also works when self is a + # StateProxy: the proxy exposes the wrapped state's __dict__, while + # object.__setattr__ is rejected by wrapt's C ObjectProxy. + instance_dict = self.__dict__ + propagated = instance_dict["_propagated_dirty_vars"] + current_gen = reflex_base_vars_base._computed_var_recompute_generation + if instance_dict["_propagated_generation"] != current_gen: + propagated.clear() + instance_dict["_propagated_generation"] = current_gen # pyright: ignore[reportIndexIssue] + + new_dirty = self.dirty_vars - propagated + while new_dirty: + propagated |= new_dirty + calc_vars, new_dirty = new_dirty, set() for state_name, cvar in self._dirty_computed_vars(from_vars=calc_vars): if state_name == self.get_full_name(): defining_state = self @@ -1832,7 +1877,8 @@ def _mark_dirty_computed_vars(self) -> None: if actual_var is not None: actual_var.mark_dirty(instance=defining_state) if defining_state is self: - dirty_vars.add(cvar) + if cvar not in propagated: + new_dirty.add(cvar) else: # mark dirty where this var is defined defining_state._mark_dirty() @@ -1843,10 +1889,11 @@ def _expired_computed_vars(self) -> set[str]: Returns: Set of computed vars to include in the delta. """ + computed_vars = self.computed_vars return { cvar - for cvar, cvar_obj in self.computed_vars.items() - if cvar_obj.needs_update(instance=self) + for cvar in self._interval_computed_vars + if computed_vars[cvar].needs_update(instance=self) } def _dirty_computed_vars( @@ -1966,6 +2013,8 @@ def _clean(self): # Clean this state. self.dirty_vars = set() self.dirty_substates = set() + # Discard the propagation frontier along with the dirty vars it tracked. + self.__dict__["_propagated_dirty_vars"].clear() def get_value(self, key: str) -> Any: """Get the value of a field (without proxying). @@ -2085,6 +2134,9 @@ def __getstate__(self): state.pop("_was_touched", None) # Proxies wrap live state references and are rebuilt on access. state.pop("_mutable_proxy_cache", None) + # The propagation frontier is transient and rebuilt on demand. + state.pop("_propagated_dirty_vars", None) + state.pop("_propagated_generation", None) # Remove all inherited vars. for inherited_var_name in self.inherited_vars: state.pop(inherited_var_name, None) @@ -2102,6 +2154,9 @@ def __setstate__(self, state: dict[str, Any]): state["substates"] = {} # The proxy cache is never pickled; recreate it on the restored instance. state.setdefault("_mutable_proxy_cache", {}) + # The propagation frontier is never pickled; recreate it on restore. + state.setdefault("_propagated_dirty_vars", set()) + state.setdefault("_propagated_generation", 0) for key, value in state.items(): object.__setattr__(self, key, value) diff --git a/tests/units/test_state.py b/tests/units/test_state.py index 5c5c49bd13b..35fec2068f6 100644 --- a/tests/units/test_state.py +++ b/tests/units/test_state.py @@ -1453,6 +1453,64 @@ class WrongTypeState(BaseState): assert mock_error.call_count == 1 +def test_computed_var_recompute_after_mid_cycle_read(): + """A dependency mutated again after a mid-cycle read still invalidates the cache.""" + + class FrontierState(BaseState): + v: int = 0 + + @rx.var + def doubled(self) -> int: + return self.v * 2 + + s = FrontierState() + s.v = 1 + # Reading mid-cycle recomputes and re-caches the value. + assert s.doubled == 2 + # Mutating the same dependency again must invalidate the fresh cache, + # even though dirty_vars already contained it. + s.v = 2 + assert s.doubled == 4 + assert s.get_delta()[s.get_full_name()]["doubled" + FIELD_MARKER] == 4 + + +def test_computed_var_recompute_after_mid_cycle_read_across_states(): + """Cross-state dependency invalidation survives a mid-cycle recompute.""" + + class FrontierParentState(BaseState): + v: int = 0 + + class FrontierChildState(FrontierParentState): + @rx.var + def doubled(self) -> int: + return self.v * 2 + + parent = FrontierParentState() + child = parent.substates[FrontierChildState.get_name()] + assert isinstance(child, FrontierChildState) + parent.v = 1 + assert child.doubled == 2 + parent.v = 2 + assert child.doubled == 4 + + +def test_interval_computed_vars_precomputed(): + """Classes precompute which computed vars carry an update interval.""" + + class IntervalFreeState(BaseState): + @rx.var + def untimed(self) -> int: + return 2 + + class IntervalState(BaseState): + @rx.var(interval=15) + def timed(self) -> int: + return 1 + + assert IntervalFreeState._interval_computed_vars == () + assert IntervalState._interval_computed_vars == ("timed",) + + def test_computed_var_cached_depends_on_non_cached(): """Test that a cached var is recalculated if it depends on non-cached ComputedVar."""