Skip to content
Open
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.
1 change: 1 addition & 0 deletions news/+memo-body-analysis.performance.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Reuse unchanged memo-body analysis during module emission to reduce repeated rendering and artifact collection.
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 @@
Evaluate generated passthrough memo bodies once and retain their render and artifacts so module emission does not repeat the work.
Original file line number Diff line number Diff line change
Expand Up @@ -324,6 +324,7 @@ def _finalize_fields(


_COMPILE_CACHE_ATTRS = (
"_memo_analysis_key",
"_cached_render_result",
"_vars_cache",
"_imports_cache",
Expand Down
140 changes: 113 additions & 27 deletions packages/reflex-base/src/reflex_base/components/memo.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,9 +29,14 @@

from reflex_base import constants
from reflex_base.components.app_wraps import collect_subtree_app_wraps
from reflex_base.components.component import Component
from reflex_base.components.component import (
BaseComponent,
Component,
_field_values_equal,
)
from reflex_base.components.memoize_helpers import (
MemoizationStrategy,
_var_data_key,
get_memoization_strategy,
)
from reflex_base.constants.compiler import (
Expand All @@ -44,7 +49,7 @@
from reflex_base.registry import RegistrationContext
from reflex_base.utils import console, format, memo_paths
from reflex_base.utils.deterministic_hash import deterministic_hash
from reflex_base.utils.imports import ImportVar
from reflex_base.utils.imports import ImportVar, ParsedImportDict
from reflex_base.utils.types import safe_issubclass, typehint_issubclass
from reflex_base.vars import VarData
from reflex_base.vars.base import LiteralVar, Var
Expand Down Expand Up @@ -1911,6 +1916,70 @@ def _create_component_wrapper(
return _MemoComponentWrapper(definition)


@dataclasses.dataclass(frozen=True, slots=True)
class _MemoBodyAnalysis:
"""Artifacts of a memo body, reusable until its compilation caches are cleared."""

component_type: type[Component]
rendered: dict
style: Any
style_data_key: tuple | None
imports: ParsedImportDict
internal_hooks: dict[str, VarData | None]
hook: str | None
added_hooks: dict[str, VarData | None]
custom_code: str | None
added_custom_code: tuple[list[str], ...]
dynamic_import: str | None
app_wraps: dict[tuple[int, str], Component]

def can_reuse(self, styled: Component) -> bool:
"""Check whether root styling and copying preserved the analyzed inputs.

Args:
styled: Its copy after applying the current app's root style.

Returns:
Whether emission can use the recorded render and artifacts.
"""
return (
type(styled) is self.component_type
and type(styled).__copy__ is BaseComponent.__copy__
and _var_data_key(styled.style._var_data) == self.style_data_key
and _field_values_equal(styled.style, self.style)
)


def _analyze_memo_body(
component: Component, rendered: dict, artifacts: tuple[Any, ...]
) -> _MemoBodyAnalysis:
"""Retain the already-collected passthrough artifacts for module emission.

Args:
component: The body whose children have been replaced by a hole.
rendered: The body's rendered JSX representation.
artifacts: The existing content-hash inputs from ``_component_artifacts``.

Returns:
Analysis shared by content hashing and module emission.
"""
_, imports, internal, hook, added, custom, *remaining = artifacts
return _MemoBodyAnalysis(
component_type=type(component),
rendered=rendered,
style=copy(component.style),
style_data_key=_var_data_key(component.style._var_data),
imports=imports,
internal_hooks=internal,
hook=hook,
added_hooks=added,
custom_code=custom,
added_custom_code=tuple(remaining[:-2]),
dynamic_import=remaining[-2],
app_wraps=remaining[-1],
)


def _component_artifacts(component: Component, *, recursive: bool) -> Iterator[Any]:
"""Yield everything besides the render that identifies a memo body.

Expand Down Expand Up @@ -1969,9 +2038,18 @@ def component_hash(component: Component, *, recursive: bool) -> str:
Returns:
The hex digest content hash.
"""
return deterministic_hash(
component.render(), *_component_artifacts(component, recursive=recursive)
)
if recursive or not component.children:
return deterministic_hash(
component.render(), *_component_artifacts(component, recursive=recursive)
)
rendered = component.render()
artifacts = tuple(_component_artifacts(component, recursive=False))
digest = deterministic_hash(rendered, *artifacts)
analyses = RegistrationContext.ensure_context()._memo_body_analyses
Comment thread
FarhanAliRaza marked this conversation as resolved.
if digest not in analyses:
analyses[digest] = _analyze_memo_body(component, rendered, artifacts)
vars(component)["_memo_analysis_key"] = digest
return digest


def memo_tag(component: Component) -> str:
Expand All @@ -1996,6 +2074,20 @@ def memo_tag(component: Component) -> str:
).capitalize()


_PASSTHROUGH_PARAMS = (
MemoParam(
name="children",
kind=MemoParamKind.CHILDREN,
annotation=Var[Component],
parameter_kind=inspect.Parameter.POSITIONAL_OR_KEYWORD,
js_prop_name="children",
placeholder_name="children",
kind_data=None,
default=inspect.Parameter.empty,
),
)


def create_passthrough_component_memo(
component: Component,
source_module: str | None = None,
Expand Down Expand Up @@ -2068,36 +2160,30 @@ def passthrough(children: Var[Component]) -> Component:
object.__setattr__(new_component, "_get_all_refs", component._get_all_refs)
return new_component

# Evaluate once to compute the tag from the rendered memo body shape.
# ``_create_component_definition`` evaluates again internally; that second
# pass appends another, identical hole to ``captured_hole_child``, and the
# ``captured_hole_child[0]`` read below picks up the first.
params = _analyze_params(passthrough, for_component=True)
preview = _normalize_component_return(_evaluate_memo_function(passthrough, params))
if preview is None:
msg = (
"`create_passthrough_component_memo` requires a component that "
"normalizes to `rx.Component`."
)
raise TypeError(msg)
# The compiler owns this fixed signature; no user annotations need resolving.
params = _PASSTHROUGH_PARAMS
rest_target_fields: set[str] = set()
preview = _evaluate_component_body(passthrough, params, rest_target_fields)
tag = memo_tag(preview)

passthrough.__name__ = format.to_snake_case(tag)
passthrough.__qualname__ = passthrough.__name__
passthrough.__module__ = __name__

definition = _create_component_definition(passthrough, Component, source_module)
# ``export_name`` is the content-hashed tag, which reads as noise in the
# React DevTools tree. Name the memo after the Python class it wraps.
replacements: dict[str, Any] = {
"auto_memo_wrapper": True,
"display_name": type(component).__qualname__,
}
if definition.export_name != tag:
replacements["export_name"] = tag
if captured_hole_child:
replacements["passthrough_hole_child"] = captured_hole_child[0]
definition = dataclasses.replace(definition, **replacements)
definition = MemoComponentDefinition(
fn=passthrough,
python_name=passthrough.__name__,
params=params,
source_module=source_module,
export_name=tag,
_component=_LazyBody.ready(preview),
_rest_target_fields=rest_target_fields,
auto_memo_wrapper=True,
display_name=type(component).__qualname__,
passthrough_hole_child=captured_hole_child[0] if captured_hole_child else None,
)

return _create_component_wrapper(definition), definition

Expand Down
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,12 +150,37 @@ 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


def _var_data_key(data: VarData | None) -> tuple | None:
"""Identify compilation metadata without invoking JavaScript equality on Vars.

Args:
data: The metadata used to compile a component or event wrapper.

Returns:
A key preserving dependency and provider identity, or None.
"""
if not data:
return None
return (
data.state,
data.field_name,
data.imports,
data.hooks,
tuple(id(dep) for dep in data.deps),
data.position,
tuple(id(component) for component in data.components),
tuple((priority, id(component)) for priority, component in data.app_wraps),
)


def fix_event_triggers_for_memo(
component: Component, page_context: PageContext
) -> Component:
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)
return chain


@dataclasses.dataclass(
Expand Down
25 changes: 23 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,15 @@
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.components.memo import _MemoBodyAnalysis
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 +75,24 @@ 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)
_memo_body_analyses: dict[str, _MemoBodyAnalysis] = dataclasses.field(
default_factory=dict, repr=False
)

def _reset_compile_caches(self) -> None:
"""Drop the memo and event caches that only need to outlive one compile."""
self._memoized_event_triggers.clear()
self._bound_event_chains.clear()
self._memo_body_analyses.clear()

@property
def app(self) -> App:
Expand Down
2 changes: 1 addition & 1 deletion pyi_hashes.json
Original file line number Diff line number Diff line change
Expand Up @@ -120,5 +120,5 @@
"packages/reflex-components-sonner/src/reflex_components_sonner/toast.pyi": "f170ac685b6ba5892370166c80684db3",
"reflex/__init__.pyi": "a3e1782fab4a9aed55f66cc98af8c217",
"reflex/components/__init__.pyi": "9facd05a776d0641432696bbf8e34388",
"reflex/experimental/memo.pyi": "27a73a66e238746e5da5accf99a8fdfd"
"reflex/experimental/memo.pyi": "3d05a929d95fd6dd3bcf608bbf716d15"
}
Loading
Loading