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/6841.bugfix.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Avoid leaking temporary router dirtiness into linked/shared state event deltas.
251 changes: 245 additions & 6 deletions reflex/istate/shared.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,13 +5,13 @@
import logging
from collections.abc import AsyncIterator
from dataclasses import MISSING
from typing import TYPE_CHECKING, TypeVar
from typing import TYPE_CHECKING, Any, TypeVar

from reflex_base.constants import ROUTER_DATA, ROUTER_VARS
from reflex_base.event import Event, get_hydrate_event
from reflex_base.registry import RegistrationContext
from reflex_base.utils.exceptions import ReflexRuntimeError
from reflex_base.vars.base import _owner_state
from reflex_base.vars.base import _delta_value_key, _owner_state
from typing_extensions import Self

from reflex.istate.delta import _suppress_delta_recording
Expand Down Expand Up @@ -115,14 +115,253 @@ async def _patch_state(
linked_state.dirty_vars.update(linked_state.computed_vars)
linked_state._mark_dirty()
# Apply the updates into the existing state tree for rehydrate.
# For regular linked-state events this router dirtiness is temporary:
# it forces router-dependent computed vars to resolve for the patched
# tree, but should not leak into the event's final delta.
root_state = original_state._get_root_state()
dirty_state_snapshots: list[
tuple[
BaseState,
set[str],
set[str],
set[str],
bool,
]
] = []
states_by_name: dict[str, BaseState] = {}
computed_var_cache_snapshots: dict[
str, dict[str, tuple[bool, Any, bool, Any]]
] = {}
if not full_delta:
states_to_snapshot = [root_state]
while states_to_snapshot:
state = states_to_snapshot.pop()
# Interval and always-dirty computed vars are valid output from
# this refresh and must not be mistaken for temporary router dirtiness.
computed_vars_to_preserve = state._expired_computed_vars().union(
state._always_dirty_computed_vars
)
states_by_name[state.get_full_name()] = state
dirty_state_snapshots.append((
state,
set(state.dirty_vars),
set(state.dirty_substates),
computed_vars_to_preserve,
state._was_touched,
))
states_to_snapshot.extend(state.substates.values())
computed_vars_to_snapshot = {
state.get_full_name(): (
dirty_vars.intersection(state.computed_vars)
| computed_vars_to_preserve
| state._interval_computed_var_names
)
for (
state,
dirty_vars,
_,
computed_vars_to_preserve,
_,
) in dirty_state_snapshots
}
# Track only caches that dirty propagation can invalidate, before
# marking router vars deletes those cached values.
pending_dependencies = [
(state, name)
for (
state,
dirty_vars,
_,
computed_vars_to_preserve,
_,
) in dirty_state_snapshots
for name in dirty_vars | computed_vars_to_preserve
]
pending_dependencies.extend((root_state, name) for name in ROUTER_VARS)
pending_dependencies.append((root_state, ROUTER_DATA))
seen_dependencies: set[tuple[str, str]] = set()
while pending_dependencies:
state, name = pending_dependencies.pop()
dependency = (state.get_full_name(), name)
if dependency in seen_dependencies:
continue
seen_dependencies.add(dependency)
for target_name, computed_var_name in state._var_dependencies.get(
name, ()
):
target_state = states_by_name.get(target_name)
if target_state is None:
continue
target_vars = computed_vars_to_snapshot[target_name]
target_vars.add(computed_var_name)
pending_dependencies.append((target_state, computed_var_name))
for (
state,
_,
_,
_,
_,
) in dirty_state_snapshots:
computed_var_cache_snapshots[state.get_full_name()] = {
name: (
hasattr(state, state.computed_vars[name]._cache_attr),
getattr(state, state.computed_vars[name]._cache_attr, None),
hasattr(state, state.computed_vars[name]._last_updated_attr),
getattr(
state, state.computed_vars[name]._last_updated_attr, None
),
)
for name in computed_vars_to_snapshot[state.get_full_name()]
}
root_state.dirty_vars.update(ROUTER_VARS)
root_state.dirty_vars.add(ROUTER_DATA)
root_state._mark_dirty()
# The delta is discarded: it is only resolved to refresh computed vars,
# so its values must not count as sent to the client.
with _suppress_delta_recording():
await root_state._get_resolved_delta()
router_dirty_snapshots: list[
tuple[
BaseState,
set[str],
set[str],
set[str],
set[str],
set[str],
dict[str, tuple[bool, Any]],
dict[str, tuple[bool, Any, bool, Any]],
]
] = []
if not full_delta:
for (
state,
dirty_vars,
dirty_substates,
computed_vars_to_preserve,
_,
) in dirty_state_snapshots:
computed_var_snapshots = computed_var_cache_snapshots[
state.get_full_name()
]
router_dirty_vars = (
state.dirty_vars - dirty_vars - computed_vars_to_preserve
)
router_computed_snapshots = {
name: (
computed_var_snapshots[name][0],
_delta_value_key(computed_var_snapshots[name][1])
if computed_var_snapshots[name][0]
else None,
)
for name in router_dirty_vars.intersection(computed_var_snapshots)
}
router_dirty_snapshots.append((
state,
dirty_vars,
dirty_substates,
router_dirty_vars,
state.dirty_substates - dirty_substates,
computed_vars_to_preserve,
router_computed_snapshots,
computed_var_snapshots,
))
try:
# The delta is discarded: it is only resolved to refresh computed vars,
# so its values must not count as sent to the client.
with _suppress_delta_recording():
await root_state._get_resolved_delta()
except BaseException:
if not full_delta:
for (
state,
dirty_vars,
dirty_substates,
_,
was_touched,
) in dirty_state_snapshots:
computed_var_snapshots = computed_var_cache_snapshots[
state.get_full_name()
]
state.dirty_vars = dirty_vars
state.dirty_substates = dirty_substates
Comment thread
greptile-apps[bot] marked this conversation as resolved.
state._was_touched = was_touched
for name, (
had_cache,
cached_value,
had_last_updated,
last_updated,
) in computed_var_snapshots.items():
computed_var = state.computed_vars[name]
if had_cache:
setattr(state, computed_var._cache_attr, cached_value)
else:
with contextlib.suppress(AttributeError):
delattr(state, computed_var._cache_attr)
if had_last_updated:
setattr(
state,
computed_var._last_updated_attr,
last_updated,
)
else:
with contextlib.suppress(AttributeError):
delattr(state, computed_var._last_updated_attr)
raise
else:
if not full_delta:
computed_refresh_states: list[BaseState] = []
for (
state,
dirty_vars,
dirty_substates,
router_dirty_vars,
_,
computed_vars_to_preserve,
router_computed_snapshots,
computed_var_snapshots,
) in router_dirty_snapshots:
computed_vars_refreshed = set(computed_vars_to_preserve) | (
state.dirty_vars - router_dirty_vars
).intersection(state.computed_vars)
computed_vars_refreshed.update(
name
for name in state._interval_computed_var_names
if (
computed_var_snapshots[name][2]
!= hasattr(
state,
state.computed_vars[name]._last_updated_attr,
)
or computed_var_snapshots[name][3]
!= getattr(
state,
state.computed_vars[name]._last_updated_attr,
None,
)
)
)
computed_vars_refreshed.update(
name
for name, (
had_cache,
cached_value_key,
) in router_computed_snapshots.items()
if hasattr(state, state.computed_vars[name]._cache_attr)
and (
not had_cache
or cached_value_key
!= _delta_value_key(
getattr(state, state.computed_vars[name]._cache_attr)
)
)
)
if computed_vars_refreshed:
computed_refresh_states.append(state)
state.dirty_vars = (
dirty_vars | computed_vars_to_preserve | computed_vars_refreshed
)
state.dirty_substates = dirty_substates
for refreshed_state in computed_refresh_states:
state = refreshed_state
while state.parent_state is not None:
state.parent_state.dirty_substates.add(state.get_name())
state = state.parent_state
yield
finally:
original_parent_state.substates[state_name] = original_state
Expand Down
Loading
Loading