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/+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 f5962b18150..7874b7b2ff5 100644 --- a/packages/reflex-base/src/reflex_base/event/__init__.py +++ b/packages/reflex-base/src/reflex_base/event/__init__.py @@ -920,6 +920,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] @@ -959,12 +975,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/reflex/compiler/compiler.py b/reflex/compiler/compiler.py index 54f8d90d296..f06f6bb9426 100644 --- a/reflex/compiler/compiler.py +++ b/reflex/compiler/compiler.py @@ -1278,6 +1278,11 @@ def compile_app( # ``library`` from the current module layout (handles a module flipping to # a package across hot reloads). reset_memo_component_classes() + # Page evaluation rebuilds every chain that is not interned by handler, so + # entries from an earlier compile can only retain dead chains. + context = RegistrationContext.ensure_context() + context._bound_event_chains.clear() + context._memoized_event_triggers.clear() for plugin in compiler_plugins: for dependency in plugin.get_frontend_dependencies(): _bundle_library(dependency) diff --git a/tests/units/compiler/test_compiler.py b/tests/units/compiler/test_compiler.py index 736c5ae653c..07307221829 100644 --- a/tests/units/compiler/test_compiler.py +++ b/tests/units/compiler/test_compiler.py @@ -1765,3 +1765,29 @@ def test_no_ssr_dynamic_import_names_the_client_side_wrapper(): from reflex_components_plotly.plotly import Plotly assert Plotly.create()._get_dynamic_imports().endswith(', "Plot")') + + +def test_compile_app_drops_event_caches_from_earlier_compiles( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, mocker: MockerFixture +): + """Chains and wrappers cached by an earlier compile do not outlive it. + + Args: + tmp_path: Directory for compiler output. + monkeypatch: Fixture for changing the app directory. + mocker: Fixture for configuring the test app. + """ + monkeypatch.chdir(tmp_path) + with RegistrationContext() as context: + config = rx.Config(app_name="event_cache_test", plugins=[]) + mocker.patch("reflex_base.config._get_config", return_value=config) + app = rx.App() + app.add_page(lambda: rx.el.div("hello"), route="/") + stale = object() + context._bound_event_chains[0, 0, None] = stale # pyright: ignore[reportArgumentType] + context._memoized_event_triggers["on_click", 0] = stale # pyright: ignore[reportArgumentType] + + compiler.compile_app(app, dry_run=True, use_rich=False) + + assert (0, 0, None) not in context._bound_event_chains + assert ("on_click", 0) not in context._memoized_event_triggers 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 bf3c537bb7c..c142ef5bc27 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, @@ -1399,3 +1400,90 @@ 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 + # The event_actions call above must not replace the cached 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 + )