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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions news/+event-chain-interning.performance.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Share one event chain per handler and trigger across call sites, and reuse memoized event wrappers by chain identity during compilation.
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Share one event chain per handler and trigger across call sites, and reuse memoized event wrappers by chain identity during compilation.
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
):
Expand All @@ -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()
Expand All @@ -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)}])"
Expand All @@ -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


Expand Down
24 changes: 22 additions & 2 deletions packages/reflex-base/src/reflex_base/event/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down Expand Up @@ -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)
Comment thread
cubic-dev-ai[bot] marked this conversation as resolved.
return chain


@dataclasses.dataclass(
Expand Down
15 changes: 13 additions & 2 deletions packages/reflex-base/src/reflex_base/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down Expand Up @@ -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]] = (
Comment thread
FarhanAliRaza marked this conversation as resolved.
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:
Expand Down
5 changes: 5 additions & 0 deletions reflex/compiler/compiler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
26 changes: 26 additions & 0 deletions tests/units/compiler/test_compiler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
143 changes: 143 additions & 0 deletions tests/units/reflex_base/components/test_memoize_helpers.py
Original file line number Diff line number Diff line change
@@ -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)}
Comment thread
masenf marked this conversation as resolved.
)
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
Loading
Loading