diff --git a/news/+compile-prop-hot-paths.performance.md b/news/+compile-prop-hot-paths.performance.md new file mode 100644 index 00000000000..8342a4a07b8 --- /dev/null +++ b/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/news/+event-chain-interning.performance.md b/news/+event-chain-interning.performance.md new file mode 100644 index 00000000000..417d948848b --- /dev/null +++ b/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/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/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/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/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/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/event/__init__.py b/packages/reflex-base/src/reflex_base/event/__init__.py index b191bd93349..310f8a7a747 100644 --- a/packages/reflex-base/src/reflex_base/event/__init__.py +++ b/packages/reflex-base/src/reflex_base/event/__init__.py @@ -913,6 +913,22 @@ 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 per + # registration context. Handlers carrying event actions are fresh + # copies at every call site, so caching them would only retain them. + bound_handler = None + if ( + not event_chain_kwargs + and isinstance(value, EventHandler) + and not value.event_actions + ): + bound_handler = value + bound_chains = RegistrationContext.ensure_context()._bound_event_chains + bound = bound_chains.get((id(value), id(args_spec), key)) + if bound is not None and bound[0] is value and bound[1] is args_spec: + return bound[2] + # If the input is a single event handler, wrap it in a list. if isinstance(value, (EventHandler, EventSpec)): value = [value] @@ -952,12 +968,16 @@ 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_handler is not None: + RegistrationContext.ensure_context()._bound_event_chains[ + id(bound_handler), id(args_spec), key + ] = (bound_handler, 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 fe337669286..27393504505 100644 --- a/packages/reflex-base/src/reflex_base/registry.py +++ b/packages/reflex-base/src/reflex_base/registry.py @@ -11,12 +11,14 @@ from reflex_base.utils.exceptions import ReflexRuntimeError, StateValueError if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Callable, Sequence from reflex.app import App from reflex.state import BaseState from reflex_base.config import Config - from reflex_base.event import EventHandler + from reflex_base.event import EventChain, EventHandler + from reflex_base.utils.types import ArgsSpec + from reflex_base.vars.base import Var def _default_bundled_libraries() -> list[str]: @@ -72,6 +74,15 @@ class RegistrationContext(BaseContext): default_factory=dict, 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) + ) + # (handler id, args_spec id, trigger key) -> the handler, spec and their + # bound chain. The referents keep the ids valid for the map's lifetime. + _bound_event_chains: dict[ + tuple[int, int, str | None], + tuple[EventHandler, ArgsSpec | Sequence[ArgsSpec], EventChain], + ] = dataclasses.field(default_factory=dict, repr=False) @property def app(self) -> App: diff --git a/packages/reflex-base/src/reflex_base/vars/base.py b/packages/reflex-base/src/reflex_base/vars/base.py index 79ebe8bd283..ef3de2d9877 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 @@ -115,6 +114,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 @@ -236,7 +266,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." @@ -1651,6 +1681,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( @@ -1678,9 +1709,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 @@ -1760,9 +1790,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 @@ -2020,6 +2049,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) @@ -2029,7 +2060,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: @@ -2067,11 +2097,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: try: 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/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/reflex_base/vars/test_base.py b/tests/units/reflex_base/vars/test_base.py index 4a2a72a4347..f2186cd5f9e 100644 --- a/tests/units/reflex_base/vars/test_base.py +++ b/tests/units/reflex_base/vars/test_base.py @@ -1,9 +1,12 @@ """Tests for reflex_base.vars.base state metaclass field handling.""" import dataclasses +import gc +import pickle import threading import traceback import typing +import weakref from typing import Any, Literal, TypeVar import pytest @@ -11,11 +14,13 @@ from reflex_base.utils.exceptions import ReflexRuntimeError from reflex_base.utils.types import get_field_type from reflex_base.vars.base import ( + GLOBAL_CACHE, CachedVarOperation, EvenMoreBasicBaseState, LiteralVar, Var, _linearize_bases, + cached_property, cached_property_no_lock, field, ) @@ -302,3 +307,128 @@ def _cached_get_all_var_data(self): BrokenVar(_js_expr="")._get_all_var_data() assert isinstance(exc_info.value.__cause__, AttributeError) assert str(exc_info.value.__cause__) == "the real error message" + + +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 + from reflex_base.vars import base + + var_subclasses = len(base._var_subclasses) + literal_subclasses = len(base._var_literal_subclasses) + try: + + @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]" + finally: + serializers.SERIALIZERS.pop(Coordinate) + serializers.SERIALIZER_TYPES.pop(Coordinate) + serializers.get_serializer.cache_clear() + serializers.get_serializer_type.cache_clear() + del base._var_subclasses[var_subclasses:] + del base._var_literal_subclasses[literal_subclasses:] + base._clear_var_subclass_lookup_caches() + base._literal_var_by_type.clear() diff --git a/tests/units/test_event.py b/tests/units/test_event.py index ad964ab821e..2e2263c7f9b 100644 --- a/tests/units/test_event.py +++ b/tests/units/test_event.py @@ -19,6 +19,7 @@ on_submit_event, on_submit_string_event, ) +from reflex_base.registry import RegistrationContext from reflex_base.utils import format, log from reflex_base.utils.exceptions import ( EventHandlerArgTypeMismatchError, @@ -1368,3 +1369,89 @@ 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_cache_lives_on_the_registration_context( + forked_registration_context: RegistrationContext, +): + """Bound chains are shared per context and leave the handler stateless.""" + + class ChainState(BaseState): + @event + def handler(self): + pass + + def args_spec(): + return () + + chain = EventChain.create(ChainState.handler, args_spec=args_spec, key="on_click") + with forked_registration_context.fork(): + forked = EventChain.create( + ChainState.handler, args_spec=args_spec, key="on_click" + ) + assert forked is not chain + assert ( + EventChain.create(ChainState.handler, args_spec=args_spec, key="on_click") + is forked + ) + assert ( + EventChain.create(ChainState.handler, args_spec=args_spec, key="on_click") + is chain + ) + + def retains(value: Any) -> bool: + if isinstance(value, dict): + value = tuple(value.values()) + if isinstance(value, (tuple, list)): + return any(retains(item) for item in value) + return value is chain + + assert not any(retains(value) for value in vars(ChainState.handler).values()) + + +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 + ) + bound_chains = RegistrationContext.ensure_context()._bound_event_chains + cached = len(bound_chains) + assert ( + EventChain.create( + ChainState.handler.prevent_default, args_spec=args_spec, key="on_click" + ) + is not chain + ) + assert len(bound_chains) == cached + assert ( + EventChain.create([ChainState.handler], args_spec=args_spec, key="on_click") + is not chain + )