From 870787ab2267a5ddfb1eb0927d533a4d485e6d3b Mon Sep 17 00:00:00 2001 From: Farhan Date: Sat, 12 Sep 2026 04:01:29 +0500 Subject: [PATCH 1/2] perf(compile): read only set props and cache literal Var dispatch Component render, Var collection, and the prop-component scan walked every declared prop through the field descriptor to find the few that are set. Iterate the instance dict plus class-level defaults instead. Cache the literal Var class per exact value type, short-circuit app-wrap dedupe on identity, skip the generic tag protocol for plain tags, and hoist the memoize plugin's component imports. Docs site dry compile (511 pages): 47 s to 40 s. Claude-Session: https://claude.ai/code/session_01PmizE1eQhtYZyVs1RK2ke3 --- .../+compile-prop-hot-paths.performance.md | 1 + .../src/reflex_base/components/component.py | 58 ++++++-- .../src/reflex_base/components/tags/tag.py | 3 + .../reflex-base/src/reflex_base/vars/base.py | 52 ++++++-- reflex/compiler/plugins/memoize.py | 7 +- tests/units/components/test_component.py | 49 +++++++ tests/units/components/test_tag.py | 26 ++++ tests/units/reflex_base/vars/test_base.py | 124 +++++++++++++++++- 8 files changed, 293 insertions(+), 27 deletions(-) create mode 100644 packages/reflex-base/news/+compile-prop-hot-paths.performance.md diff --git a/packages/reflex-base/news/+compile-prop-hot-paths.performance.md b/packages/reflex-base/news/+compile-prop-hot-paths.performance.md new file mode 100644 index 00000000000..8342a4a07b8 --- /dev/null +++ b/packages/reflex-base/news/+compile-prop-hot-paths.performance.md @@ -0,0 +1 @@ +Speed up compilation by reading only the props a component sets, caching literal Var dispatch by value type, and trimming render and app-wrap bookkeeping. diff --git a/packages/reflex-base/src/reflex_base/components/component.py b/packages/reflex-base/src/reflex_base/components/component.py index adf8d36fd0e..8cf95482f9c 100644 --- a/packages/reflex-base/src/reflex_base/components/component.py +++ b/packages/reflex-base/src/reflex_base/components/component.py @@ -24,6 +24,7 @@ from reflex_base.components.dynamic import load_dynamic_serializer from reflex_base.components.field import BaseField, FieldBasedMeta from reflex_base.components.tags import Tag +from reflex_base.components.tags.tag import render_prop from reflex_base.constants import Dirs, EventTriggers, Hooks, Imports, MemoizationMode from reflex_base.constants.compiler import SpecialAttributes from reflex_base.event import ( @@ -1161,7 +1162,7 @@ def _render(self, props: dict[str, Any] | None = None) -> Tag: if props is None: # Add component props to the tag. props = { - attr.removesuffix("_"): getattr(self, attr) for attr in self.get_props() + prop.removesuffix("_"): value for prop, value in self._iter_set_props() } # Add ref to element if `ref` is None and `id` is not None. @@ -1201,6 +1202,39 @@ def get_props(cls) -> Iterable[str]: """ return cls.get_js_fields() + @classmethod + @functools.cache + def _get_defaulted_props(cls) -> frozenset[str]: + """Get the props whose field supplies a value when unset. + + Returns: + The props with a default other than ``None`` or a default factory. + """ + return frozenset( + prop + for prop, field_ in cls.get_js_fields().items() + if field_.default_factory is not None + or (field_.default is not MISSING and field_.default is not None) + ) + + def _iter_set_props(self) -> Iterator[tuple[str, Any]]: + """Walk the props that carry a value, in declaration order. + + An unset prop resolves to ``None`` through its field descriptor and + every consumer drops ``None``, so only props present on the instance + or backed by a class default are read. + + Yields: + Each prop name with its value. + """ + values = self.__dict__ + defaulted = self._get_defaulted_props() + for prop in self.get_props(): + if prop in values: + yield prop, values[prop] + elif prop in defaulted: + yield prop, getattr(self, prop) + @classmethod @functools.cache def get_initial_props(cls) -> set[str]: @@ -1215,9 +1249,8 @@ def get_initial_props(cls) -> set[str]: def _get_component_prop_property(self) -> Sequence[BaseComponent]: return [ component - for prop in self.get_props() - if (value := getattr(self, prop)) is not None - and isinstance(value, (BaseComponent, Var)) + for _, value in self._iter_set_props() + if isinstance(value, (BaseComponent, Var)) for component in _components_from(value) ] @@ -1438,11 +1471,15 @@ def render(self) -> dict: except AttributeError: pass tag = self._render() - rendered_dict = dict( - tag.set( - children=[child.render() for child in self.children], - ) - ) + children = [child.render() for child in self.children] + if type(tag) is Tag: + rendered_dict = {} + if (name := render_prop(tag.name)) is not None: + rendered_dict["name"] = name + rendered_dict["props"] = tag.format_props() + rendered_dict["children"] = children + else: + rendered_dict = dict(tag.set(children=children)) self._replace_prop_names(rendered_dict) self._cached_render_result = rendered_dict return rendered_dict @@ -1581,8 +1618,7 @@ def _get_vars( vars.extend(event_vars) # Get Vars associated with component props. - for prop in self.get_props(): - prop_var = getattr(self, prop) + for _, prop_var in self._iter_set_props(): if isinstance(prop_var, Var): vars.append(prop_var) diff --git a/packages/reflex-base/src/reflex_base/components/tags/tag.py b/packages/reflex-base/src/reflex_base/components/tags/tag.py index 6921121c4fa..cc607b90c82 100644 --- a/packages/reflex-base/src/reflex_base/components/tags/tag.py +++ b/packages/reflex-base/src/reflex_base/components/tags/tag.py @@ -20,6 +20,9 @@ def render_prop(value: Any) -> Any: Returns: The rendered value. """ + if type(value) in (str, dict): + return value + from reflex_base.components.component import BaseComponent if isinstance(value, BaseComponent): diff --git a/packages/reflex-base/src/reflex_base/vars/base.py b/packages/reflex-base/src/reflex_base/vars/base.py index 3d541ada126..e0125af1e62 100644 --- a/packages/reflex-base/src/reflex_base/vars/base.py +++ b/packages/reflex-base/src/reflex_base/vars/base.py @@ -13,7 +13,6 @@ import logging import re import string -import uuid import warnings from abc import ABCMeta from collections.abc import Callable, Coroutine, Iterable, Mapping, Sequence @@ -114,6 +113,37 @@ class VarSubclassEntry: _var_subclasses: list[VarSubclassEntry] = [] _var_literal_subclasses: list[tuple[type[LiteralVar], VarSubclassEntry]] = [] +# Exact value type -> the literal class claiming it, or None when no literal +# class does. Reset whenever a literal subclass registers. +_literal_var_by_type: dict[type, type[LiteralVar] | None] = {} + + +def _literal_var_for(value: Any) -> type[LiteralVar] | None: + """Find the literal Var class claiming ``value``'s type. + + Args: + value: The python value to wrap. + + Returns: + The matching literal class, or None if no registered class claims it. + """ + value_type = type(value) + try: + return _literal_var_by_type[value_type] + except KeyError: + pass + literal_subclass = next( + ( + literal + for literal, var_subclass in reversed(_var_literal_subclasses) + if isinstance(value, var_subclass.python_types) + ), + None, + ) + # A class object's type is its metaclass, which other classes share. + if not isinstance(value, type): + _literal_var_by_type[value_type] = literal_subclass + return literal_subclass @functools.cache @@ -235,7 +265,7 @@ def insert_app_wraps( if seen is None: seen = target.get(key) if seen is not None: - if seen != wrapper: + if seen is not wrapper and seen != wrapper: msg = ( f"Conflicting app wraps for {key!r}: two different " "components claim the same (priority, tag) slot." @@ -1650,6 +1680,7 @@ def __init_subclass__(cls, **kwargs): _var_literal_subclasses.remove(var_literal_subclass) _var_literal_subclasses.append((cls, var_subclass)) + _literal_var_by_type.clear() @classmethod def _create_literal_var( @@ -1677,9 +1708,8 @@ def _create_literal_var( return value return value._replace(merge_var_data=_var_data) - for literal_subclass, var_subclass in _var_literal_subclasses[::-1]: - if isinstance(value, var_subclass.python_types): - return literal_subclass.create(value, _var_data=_var_data) + if (literal_subclass := _literal_var_for(value)) is not None: + return literal_subclass.create(value, _var_data=_var_data) if ( (as_var_method := getattr(value, "_as_var", None)) is not None @@ -1759,9 +1789,8 @@ def _get_all_var_data_without_creating_var_dispatch( if isinstance(value, Var): return value._get_all_var_data() - for literal_subclass, var_subclass in _var_literal_subclasses[::-1]: - if isinstance(value, var_subclass.python_types): - return literal_subclass._get_all_var_data_without_creating_var(value) + if (literal_subclass := _literal_var_for(value)) is not None: + return literal_subclass._get_all_var_data_without_creating_var(value) if ( (as_var_method := getattr(value, "_as_var", None)) is not None @@ -2019,6 +2048,8 @@ def __set_name__(self, owner: Any, name: str): """ if self._attrname is None: self._attrname = name + self._cached_field_name = "_reflex_cache_" + name + cached_field_name = self._cached_field_name original_del = getattr(owner, "__del__", None) @@ -2028,7 +2059,6 @@ def delete_property(this: Any): Args: this: The object to delete the cached property from. """ - cached_field_name = "_reflex_cache_" + name try: unique_id = object.__getattribute__(this, cached_field_name) except AttributeError: @@ -2065,11 +2095,11 @@ def __get__(self, instance: Any, owner: type | None = None): if self._attrname is None: msg = "Cannot use cached_property on a class without __set_name__." raise TypeError(msg) - cached_field_name = "_reflex_cache_" + self._attrname + cached_field_name = self._cached_field_name try: unique_id = object.__getattribute__(instance, cached_field_name) except AttributeError: - unique_id = uuid.uuid4().int + unique_id = object() object.__setattr__(instance, cached_field_name, unique_id) if unique_id not in GLOBAL_CACHE: GLOBAL_CACHE[unique_id] = self._func(instance) diff --git a/reflex/compiler/plugins/memoize.py b/reflex/compiler/plugins/memoize.py index a50e3e30569..3f71eb1409b 100644 --- a/reflex/compiler/plugins/memoize.py +++ b/reflex/compiler/plugins/memoize.py @@ -35,6 +35,9 @@ from reflex_base.constants.compiler import MemoizationDisposition from reflex_base.plugins import ComponentAndChildren, PageContext from reflex_base.plugins.base import Plugin +from reflex_components_core.base.bare import Bare +from reflex_components_core.core.cond import Cond +from reflex_components_core.core.match import Match from reflex.compiler.plugins.builtin import ( collect_var_app_wraps_for_component, @@ -146,10 +149,6 @@ def _should_memoize(component: Component) -> bool: Returns: True if the component should be wrapped in a memo definition. """ - from reflex_components_core.base.bare import Bare - from reflex_components_core.core.cond import Cond - from reflex_components_core.core.match import Match - strategy = get_memoization_strategy(component) if component._memoization_mode.disposition == MemoizationDisposition.NEVER: diff --git a/tests/units/components/test_component.py b/tests/units/components/test_component.py index 687157d9725..592c5591433 100644 --- a/tests/units/components/test_component.py +++ b/tests/units/components/test_component.py @@ -5,6 +5,7 @@ import pytest from reflex_base.components.component import Component, field +from reflex_base.components.tags import Tag from reflex_base.constants import EventTriggers from reflex_base.constants.state import FIELD_MARKER from reflex_base.event import ( @@ -45,6 +46,33 @@ from reflex.utils import imports +@pytest.mark.parametrize("name", ["div", "", None]) +def test_plain_tag_render_matches_tag_protocol(name, monkeypatch): + """Direct rendering preserves names, props, children, and render caching.""" + tag = Tag(name=name).add_props(title="hello") + component = Component._create(children=[Bare.create("child")]) + monkeypatch.setattr(component, "_render", lambda: tag) + expected = dict(tag.set(children=[child.render() for child in component.children])) + assert component.render() == expected + assert component.render() is component.render() + assert not tag.children + + +def test_custom_tag_render_uses_subclass_protocol(monkeypatch): + """Custom tag iteration can depend on its supplied children.""" + + class ChildrenTag(Tag): + """A tag with custom child-dependent rendering.""" + + def __iter__(self): + """Yield a value derived from the child list.""" + yield "child_count", len(self.children) + + component = Component._create(children=[Bare.create("child")]) + monkeypatch.setattr(component, "_render", lambda: ChildrenTag()) + assert component.render() == {"child_count": 1} + + class TestState(BaseState): """A test state with various methods for event handling.""" @@ -2398,3 +2426,24 @@ def test_get_all_hooks_internal_does_not_mutate_hooks_cache(): assert dict(parent._get_hooks_internal()) == parent_own_hooks # And repeated collection yields the same result. assert parent._get_all_hooks_internal() == combined + + +def test_set_props_iteration_skips_unset_props_and_keeps_defaults(): + """Only set props and class defaults are visited, in declaration order.""" + + class DefaultedProps(Component): + first: Var[str] + second: Var[str] = LiteralVar.create("second-default") + third: Var[str] + + component = DefaultedProps._create(children=(), third="set") + assert [(prop, str(value)) for prop, value in component._iter_set_props()] == [ + ("second", '"second-default"'), + ("third", '"set"'), + ] + assert [str(var) for var in component._get_vars()] == ['"second-default"', '"set"'] + assert {prop: str(value) for prop, value in component._render().props.items()} == { + "second": '"second-default"', + "third": '"set"', + } + assert "first" not in vars(component) diff --git a/tests/units/components/test_tag.py b/tests/units/components/test_tag.py index f79065d5a02..47a117a8a35 100644 --- a/tests/units/components/test_tag.py +++ b/tests/units/components/test_tag.py @@ -1,5 +1,6 @@ import pytest from reflex_base.components.tags import CondTag, Tag, tagless +from reflex_base.components.tags.tag import render_prop from reflex_base.vars.base import LiteralVar, Var @@ -127,3 +128,28 @@ def test_tagless_string_representation(): tag = tagless.Tagless(contents="Hello world") expected_output = "Hello world" assert str(tag) == expected_output + + +def test_render_prop_preserves_plain_values_and_subclass_dispatch(): + """Already-rendered dictionaries pass through; callable subclasses do not.""" + + class CallableString(str): + """A string whose callability must still be inspected.""" + + def __call__(self): + """Return a marker value.""" + return "called" + + class CallableDict(dict): + """A mapping whose callability must still be inspected.""" + + def __call__(self): + """Return a marker value.""" + return "called" + + rendered = {"name": "div", "children": []} + assert render_prop(rendered) is rendered + assert render_prop("text") == "text" + assert render_prop(CallableString("text")) is None + assert render_prop(CallableDict(rendered)) is None + assert render_prop(("text", rendered)) == ["text", rendered] diff --git a/tests/units/reflex_base/vars/test_base.py b/tests/units/reflex_base/vars/test_base.py index df1c6046bcd..21b822635bc 100644 --- a/tests/units/reflex_base/vars/test_base.py +++ b/tests/units/reflex_base/vars/test_base.py @@ -1,12 +1,23 @@ """Tests for reflex_base.vars.base state metaclass field handling.""" +import gc +import pickle import threading import typing +import weakref from typing import Any, Literal, TypeVar import pytest from reflex_base.utils.types import get_field_type -from reflex_base.vars.base import EvenMoreBasicBaseState, Var, _linearize_bases, field +from reflex_base.vars.base import ( + GLOBAL_CACHE, + EvenMoreBasicBaseState, + LiteralVar, + Var, + _linearize_bases, + cached_property, + field, +) from reflex_base.vars.object import ObjectVar from reflex_base.vars.sequence import ArrayVar, StringVar from typing_extensions import TypeAliasType, TypeVarTuple, Unpack @@ -246,3 +257,114 @@ def __hash__(cls) -> int: _linearize_bases((b, c)), created.__mro__[1:], strict=True ) ) + + +class _CachedValue: + """A mutable input with an explicitly resettable derived value.""" + + _reflex_cache_result: object + + def __init__(self, value: str): + """Store the input. + + Args: + value: The value to cache. + """ + self.value = value + + @cached_property + def result(self) -> list[str]: + """Return the derived value. + + Returns: + A fresh list containing the input. + """ + return [self.value] + + +def test_cached_property_identity_and_reset(): + """Local keys isolate instances and survive explicit cache resets.""" + first = _CachedValue("first") + second = _CachedValue("second") + result = first.result + assert first.result is result + assert second.result == ["second"] + first.value = "changed" + assert first.result is result + GLOBAL_CACHE.clear() + assert first.result == ["changed"] + assert first.result is not result + + +def test_cached_property_pickle_does_not_reuse_another_instances_key(): + """Deserialized keys must not collide with live cache entries.""" + original = _CachedValue("original") + assert original.result == ["original"] + restored = pickle.loads(pickle.dumps(original)) + restored.value = "restored" + assert restored.result == ["restored"] + assert original.result == ["original"] + + +def test_cached_property_releases_entry_with_instance(): + """Destroying an instance removes its cached value.""" + value = _CachedValue("temporary") + assert value.result == ["temporary"] + key = value._reflex_cache_result + reference = weakref.ref(value) + del value + gc.collect() + assert reference() is None + assert key not in GLOBAL_CACHE + + +def test_literal_var_dispatch_follows_later_registrations(): + """A literal class registered after a lookup wins the next lookup for its type.""" + + class Coordinate: + """A value no literal Var claims yet.""" + + def __init__(self, x: int): + """Store the coordinate. + + Args: + x: The coordinate value. + """ + self.x = x + + from reflex_base.utils import serializers + + @serializers.serializer + def serialize_coordinate(value: Coordinate) -> str: + """Serialize a coordinate. + + Args: + value: The coordinate. + + Returns: + Its string form. + """ + return f"coordinate-{value.x}" + + assert str(LiteralVar.create(Coordinate(1))) == '"coordinate-1"' + + class CoordinateVar(Var[Coordinate], python_types=Coordinate): + """A Var holding a coordinate.""" + + class LiteralCoordinateVar(LiteralVar, CoordinateVar): + """A literal coordinate Var.""" + + @classmethod + def create(cls, value: Coordinate, _var_data=None): + """Create the literal. + + Args: + value: The coordinate. + _var_data: Unused metadata. + + Returns: + A Var with the coordinate's expression. + """ + return Var(_js_expr=f"[{value.x}]", _var_type=Coordinate) + + assert str(LiteralVar.create(Coordinate(2))) == "[2]" From c9103569b32c5a3f0aad318da1c9444c9f62bcb6 Mon Sep 17 00:00:00 2001 From: Farhan Date: Sat, 12 Sep 2026 04:02:17 +0500 Subject: [PATCH 2/2] perf(events): share one chain per handler and trigger across call sites EventChain.create rebuilt an identical chain for every component that bound the same handler to the same trigger, and the memoize pass then rendered each chain again to name its useCallback wrapper. Intern the chain on the handler keyed by args spec and trigger, and key the wrapper cache by chain identity so repeated call sites reuse the wrapper without rendering. Claude-Session: https://claude.ai/code/session_01PmizE1eQhtYZyVs1RK2ke3 --- .../+event-chain-interning.performance.md | 1 + .../reflex_base/components/memoize_helpers.py | 28 ++-- .../src/reflex_base/event/__init__.py | 16 +- .../reflex-base/src/reflex_base/registry.py | 4 + .../components/test_memoize_helpers.py | 143 ++++++++++++++++++ tests/units/test_event.py | 45 ++++++ 6 files changed, 226 insertions(+), 11 deletions(-) create mode 100644 packages/reflex-base/news/+event-chain-interning.performance.md create mode 100644 tests/units/reflex_base/components/test_memoize_helpers.py diff --git a/packages/reflex-base/news/+event-chain-interning.performance.md b/packages/reflex-base/news/+event-chain-interning.performance.md new file mode 100644 index 00000000000..417d948848b --- /dev/null +++ b/packages/reflex-base/news/+event-chain-interning.performance.md @@ -0,0 +1 @@ +Share one event chain per handler and trigger across call sites, and reuse memoized event wrappers by chain identity during compilation. diff --git a/packages/reflex-base/src/reflex_base/components/memoize_helpers.py b/packages/reflex-base/src/reflex_base/components/memoize_helpers.py index 5c8f714465a..05dacd43c0f 100644 --- a/packages/reflex-base/src/reflex_base/components/memoize_helpers.py +++ b/packages/reflex-base/src/reflex_base/components/memoize_helpers.py @@ -26,6 +26,7 @@ from reflex_base.components.component import BaseComponent, Component from reflex_base.constants import EventTriggers from reflex_base.event import EventChain, EventSpec +from reflex_base.registry import RegistrationContext from reflex_base.utils.imports import ImportVar from reflex_base.vars import VarData from reflex_base.vars.base import LiteralVar, Var @@ -100,6 +101,9 @@ def get_memoized_event_triggers( A dict mapping event trigger name to memoized_triger. """ trigger_memo: dict[str, Var] = {} + if not component.event_triggers: + return trigger_memo + cache = RegistrationContext.ensure_context()._memoized_event_triggers for event_trigger, event_args in component._get_vars_from_event_triggers( component.event_triggers ): @@ -112,8 +116,17 @@ def get_memoized_event_triggers( continue event = component.event_triggers[event_trigger] - rendered_chain = LiteralVar.create(event) + cache_key = (event_trigger, id(event)) + cached = cache.get(cache_key) + if cached is not None and cached[0] is event: + trigger_memo[event_trigger] = cached[1] + continue + rendered_chain = LiteralVar.create(event) + rendered_data = rendered_chain._get_all_var_data() + event_var_data = [ + data for arg in event_args if (data := arg._get_all_var_data()) is not None + ] chain_hash = md5( str(rendered_chain).encode("utf-8"), usedforsecurity=False ).hexdigest() @@ -122,18 +135,13 @@ def get_memoized_event_triggers( var_deps = ["addEvents", "ReflexEvent"] var_deps.extend(_get_deps_from_event_trigger(event)) - event_var_data = [] - for arg in event_args: - var_data = arg._get_all_var_data() - if var_data is None: - continue - event_var_data.append(var_data) + for var_data in event_var_data: for hook in var_data.hooks: var_deps.extend(_get_hook_deps(hook)) memo_var_data = VarData.merge( *event_var_data, - rendered_chain._get_all_var_data(), + rendered_data, VarData( hooks=[ f"const {memo_name} = useCallback({rendered_chain!s}, [{', '.join(var_deps)}])" @@ -142,9 +150,11 @@ def get_memoized_event_triggers( ), ) - trigger_memo[event_trigger] = Var( + trigger_memo[event_trigger] = memo_var = Var( _js_expr=memo_name, _var_type=EventChain, _var_data=memo_var_data ) + # Hold the chain so its id cannot be recycled while the entry lives. + cache[cache_key] = event, memo_var return trigger_memo diff --git a/packages/reflex-base/src/reflex_base/event/__init__.py b/packages/reflex-base/src/reflex_base/event/__init__.py index b191bd93349..2b837afed8b 100644 --- a/packages/reflex-base/src/reflex_base/event/__init__.py +++ b/packages/reflex-base/src/reflex_base/event/__init__.py @@ -913,6 +913,16 @@ def create( # Trust that the caller knows what they're doing passing an EventChain directly return value + # A handler bound to one trigger always produces the same chain, so + # every call site sharing the handler shares one instance. The cache + # lives on the handler, which is never pickled or copied. + bound_chains = None + if not event_chain_kwargs and isinstance(value, EventHandler): + bound_chains = value.__dict__.setdefault("_bound_chains", {}) + bound = bound_chains.get((id(args_spec), key)) + if bound is not None and bound[0] is args_spec: + return bound[1] + # If the input is a single event handler, wrap it in a list. if isinstance(value, (EventHandler, EventSpec)): value = [value] @@ -952,12 +962,14 @@ def create( for e in events ] - # Return the event chain. - return cls( + chain = cls( events=events, args_spec=args_spec, **event_chain_kwargs, ) + if bound_chains is not None: + bound_chains[id(args_spec), key] = args_spec, chain + return chain @dataclasses.dataclass( diff --git a/packages/reflex-base/src/reflex_base/registry.py b/packages/reflex-base/src/reflex_base/registry.py index 61963cb538f..3348efcfa4d 100644 --- a/packages/reflex-base/src/reflex_base/registry.py +++ b/packages/reflex-base/src/reflex_base/registry.py @@ -17,6 +17,7 @@ from reflex.state import BaseState from reflex_base.config import Config from reflex_base.event import EventHandler + from reflex_base.vars.base import Var def _default_bundled_libraries() -> list[str]: @@ -69,6 +70,9 @@ class RegistrationContext(BaseContext): repr=False, ) _app: App | None = dataclasses.field(default=None, repr=False) + _memoized_event_triggers: dict[tuple[str, int], tuple[Any, Var]] = ( + dataclasses.field(default_factory=dict, repr=False) + ) @property def app(self) -> App: diff --git a/tests/units/reflex_base/components/test_memoize_helpers.py b/tests/units/reflex_base/components/test_memoize_helpers.py new file mode 100644 index 00000000000..e207bb71438 --- /dev/null +++ b/tests/units/reflex_base/components/test_memoize_helpers.py @@ -0,0 +1,143 @@ +"""Tests for sharing prepared event wrappers within a registration context.""" + +import dataclasses + +import pytest +from reflex_base.components.component import Component +from reflex_base.components.memoize_helpers import get_memoized_event_triggers +from reflex_base.event import EventChain, EventHandler, no_args_event_spec +from reflex_base.registry import RegistrationContext +from reflex_base.utils.imports import ImportVar +from reflex_base.vars.base import LiteralVar, Var, VarData + + +def test_event_wrappers_are_reused_and_reset_with_context(): + """Identical wrappers share work only within their owning context.""" + component = Component._create( + children=(), event_triggers={"on_click": Var("handler", EventChain)} + ) + with RegistrationContext.ensure_context().fork() as context: + first = get_memoized_event_triggers(component)["on_click"] + assert get_memoized_event_triggers(component)["on_click"] is first + with context.fork() as fork: + assert not fork._memoized_event_triggers + assert get_memoized_event_triggers(component)["on_click"] is not first + context._memoized_event_triggers.clear() + assert get_memoized_event_triggers(component)["on_click"] is not first + + +@pytest.mark.parametrize( + ("first_data", "second_data"), + [ + (VarData(state="first"), VarData(state="second")), + ( + VarData(hooks=["const first = useFirst()"]), + VarData(hooks=["const second = useSecond()"]), + ), + ( + VarData(imports={"first": [ImportVar("value")]}), + VarData(imports={"second": [ImportVar("value")]}), + ), + (VarData(deps=[Var("first")]), VarData(deps=[Var("second")])), + ], +) +def test_event_wrapper_cache_preserves_dependencies( + first_data: VarData, second_data: VarData +): + """Identical expressions with different metadata must keep their dependencies.""" + with RegistrationContext.ensure_context().fork(): + first = get_memoized_event_triggers( + Component._create( + children=(), + event_triggers={"on_click": Var("handler", EventChain, first_data)}, + ) + )["on_click"] + second = get_memoized_event_triggers( + Component._create( + children=(), + event_triggers={"on_click": Var("handler", EventChain, second_data)}, + ) + )["on_click"] + assert first is not second + assert repr(first._get_all_var_data()) != repr(second._get_all_var_data()) + + +def test_event_wrapper_cache_preserves_provider_identity(): + """Providers sharing a role can still carry distinct component props.""" + first_provider = Component._create( + children=(), tag="Provider", custom_attrs={"value": "first"} + ) + second_provider = Component._create( + children=(), tag="Provider", custom_attrs={"value": "second"} + ) + with RegistrationContext.ensure_context().fork(): + for provider in (first_provider, second_provider): + event = Var("handler", EventChain, VarData(app_wraps=[(10, provider)])) + wrapper = get_memoized_event_triggers( + Component._create(children=(), event_triggers={"on_click": event}) + )["on_click"] + data = wrapper._get_all_var_data() + assert data is not None + assert data.app_wraps[0][1] is provider + + +def test_event_wrapper_cache_does_not_compare_vars_as_python_booleans(): + """Equivalent dependency expressions may belong to different Var objects.""" + with RegistrationContext.ensure_context().fork(): + for _ in range(2): + event = Var("handler", EventChain, VarData(deps=[Var("dependency")])) + wrapper = get_memoized_event_triggers( + Component._create(children=(), event_triggers={"on_click": event}) + )["on_click"] + data = wrapper._get_all_var_data() + assert data is not None + assert {str(dep) for dep in data.deps} == {"dependency"} + + +def test_event_wrapper_reflects_captured_arguments_and_actions(): + """Chains differing in nested data compile to different wrappers.""" + + def handler(value: str): + """Accept an event argument.""" + + def chain(argument: str, **actions: bool) -> EventChain: + """Build a chain for one handler call. + + Args: + argument: The captured handler argument. + **actions: Event actions applied to the nested event. + + Returns: + The chain wrapping the handler call. + """ + spec = EventHandler(fn=handler)(argument) + if actions: + spec = dataclasses.replace(spec, event_actions=actions) + return EventChain(events=[spec], args_spec=no_args_event_spec) + + component = Component._create(children=(), event_triggers={}) + with RegistrationContext.ensure_context().fork(): + rendered = [] + for event in ( + chain("first"), + dataclasses.replace(chain("first"), event_actions={"preventDefault": True}), + chain("first", stopPropagation=True), + chain("second"), + ): + component.event_triggers["on_click"] = event + rendered.append(str(get_memoized_event_triggers(component)["on_click"])) + assert len(set(rendered)) == len(rendered) + + +def test_event_wrappers_are_shared_by_chain_identity(monkeypatch): + """Components bound to one chain object share one wrapper without rendering it.""" + chain = Var("handler", EventChain) + first = Component._create(children=(), event_triggers={"on_click": chain}) + second = Component._create(children=(), event_triggers={"on_click": chain}) + other_trigger = Component._create(children=(), event_triggers={"on_blur": chain}) + with RegistrationContext.ensure_context().fork(): + wrapper = get_memoized_event_triggers(first)["on_click"] + monkeypatch.setattr(LiteralVar, "create", pytest.fail) + assert get_memoized_event_triggers(second)["on_click"] is wrapper + monkeypatch.undo() + assert get_memoized_event_triggers(other_trigger)["on_blur"] is not wrapper diff --git a/tests/units/test_event.py b/tests/units/test_event.py index bd5d9484c0a..e315cfce2a9 100644 --- a/tests/units/test_event.py +++ b/tests/units/test_event.py @@ -1367,3 +1367,48 @@ def handle_submit(form_data: dict[str, str]): log._reset() assert "expects (dict[str, typing.Any]) -> () but got (dict[str, str]) -> ()" in out assert "\\" not in out + + +def test_event_chain_create_shares_chains_bound_from_one_handler(): + """A handler bound to one trigger yields one chain for every call site.""" + + class ChainState(BaseState): + @event + def handler(self): + pass + + def args_spec(): + return () + + chain = EventChain.create(ChainState.handler, args_spec=args_spec, key="on_click") + assert isinstance(chain, EventChain) + assert ( + EventChain.create(ChainState.handler, args_spec=args_spec, key="on_click") + is chain + ) + assert ( + EventChain.create(ChainState.handler, args_spec=args_spec, key="on_blur") + is not chain + ) + assert ( + EventChain.create(ChainState.handler, args_spec=lambda: (), key="on_click") + is not chain + ) + with_actions = EventChain.create( + ChainState.handler, args_spec=args_spec, key="on_click", event_actions={"x": 1} + ) + assert with_actions is not chain + assert ( + EventChain.create(ChainState.handler, args_spec=args_spec, key="on_click") + is chain + ) + assert ( + EventChain.create( + ChainState.handler.prevent_default, args_spec=args_spec, key="on_click" + ) + is not chain + ) + assert ( + EventChain.create([ChainState.handler], args_spec=args_spec, key="on_click") + is not chain + )