From f7cbca6a4d751ee04a2e95f79c7eff2f0f989662 Mon Sep 17 00:00:00 2001 From: Benedikt Bartscher Date: Mon, 24 Aug 2026 12:18:13 +0200 Subject: [PATCH 1/5] fix: deliver linked state updates to clients connected to other instances --- reflex/istate/shared.py | 23 +++++-- tests/units/istate/test_shared.py | 108 ++++++++++++++++++++++++++++++ 2 files changed, 125 insertions(+), 6 deletions(-) create mode 100644 tests/units/istate/test_shared.py diff --git a/reflex/istate/shared.py b/reflex/istate/shared.py index e38517fef66..e178e6a9e05 100644 --- a/reflex/istate/shared.py +++ b/reflex/istate/shared.py @@ -52,22 +52,33 @@ def _do_update_other_tokens( Returns: The list of asyncio tasks created to perform the updates. """ + from reflex.utils.token_manager import RedisTokenManager + app = RegistrationContext.get().app + tasks = [] + if (event_namespace := app.event_namespace) is None: + return tasks + token_manager = event_namespace._token_manager + async def _update_client(token: str): + # Don't send updates for disconnected clients. The local + # token_to_socket map only tracks sockets owned by this instance, so + # with redis the socket record is resolved (and cached) from redis + # instead; emit_update then relays the delta to the owning instance + # via the lost-and-found channel. + if isinstance(token_manager, RedisTokenManager): + if await token_manager._get_token_owner(token) is None: + return + elif token not in token_manager.token_to_socket: + return async with app.modify_state( BaseStateToken(ident=token, cls=state_type), previous_dirty_vars=previous_dirty_vars, ): pass - tasks = [] - if (event_namespace := app.event_namespace) is None: - return tasks for affected_token in affected_tokens: - # Don't send updates for disconnected clients. - if affected_token not in event_namespace._token_manager.token_to_socket: - continue # TODO: remove disconnected clients after some time. t = asyncio.create_task(_update_client(affected_token)) UPDATE_OTHER_CLIENT_TASKS.add(t) diff --git a/tests/units/istate/test_shared.py b/tests/units/istate/test_shared.py new file mode 100644 index 00000000000..e8b99d6a8e1 --- /dev/null +++ b/tests/units/istate/test_shared.py @@ -0,0 +1,108 @@ +"""Unit tests for shared state fan-out to other linked clients.""" + +import asyncio +import pickle +from contextlib import asynccontextmanager +from unittest.mock import AsyncMock, Mock, patch + +import pytest + +from reflex.istate.shared import _do_update_other_tokens +from reflex.state import State +from reflex.utils.token_manager import ( + LocalTokenManager, + RedisTokenManager, + SocketRecord, +) + + +@pytest.fixture +def mock_redis(): + """Create a mock Redis client. + + Returns: + The mock Redis client. + """ + redis = AsyncMock() + redis.get = AsyncMock(return_value=None) + redis.get_connection_kwargs = Mock(return_value={"db": 0}) + return redis + + +@pytest.fixture +def redis_manager(mock_redis): + """Create a RedisTokenManager instance with mocked config. + + Returns: + The RedisTokenManager instance. + """ + with patch("reflex_base.config.get_config") as mock_get_config: + mock_config = Mock() + mock_config.redis_token_expiration = 3600 + mock_get_config.return_value = mock_config + + return RedisTokenManager(mock_redis) + + +def _mock_app(token_manager) -> tuple[Mock, list[str]]: + """Create a mock app recording the tokens passed to modify_state. + + Returns: + The mock app and the list collecting modified token idents. + """ + modified_tokens: list[str] = [] + + @asynccontextmanager + async def modify_state(token, previous_dirty_vars=None): + modified_tokens.append(token.ident) + yield Mock() + + app = Mock() + app.modify_state = modify_state + app.event_namespace = Mock() + app.event_namespace._token_manager = token_manager + return app, modified_tokens + + +async def _run_update_other_tokens(app, affected_tokens: set[str]) -> None: + """Run _do_update_other_tokens against a mock app and await its tasks.""" + with patch("reflex_base.registry.RegistrationContext.get") as mock_get: + mock_get.return_value = Mock(app=app) + tasks = _do_update_other_tokens( + affected_tokens=affected_tokens, + previous_dirty_vars={}, + state_type=State, + ) + await asyncio.gather(*tasks) + + +async def test_update_other_tokens_local_manager(): + """With a LocalTokenManager, only locally connected tokens are updated.""" + manager = LocalTokenManager() + manager.token_to_socket["connected"] = SocketRecord( + instance_id=manager.instance_id, sid="sid1" + ) + app, modified_tokens = _mock_app(manager) + + await _run_update_other_tokens(app, {"connected", "disconnected"}) + + assert modified_tokens == ["connected"] + + +async def test_update_other_tokens_redis_cross_instance(redis_manager, mock_redis): + """Tokens connected to another instance are resolved via redis and updated.""" + redis_manager.token_to_socket["local"] = SocketRecord( + instance_id=redis_manager.instance_id, sid="sid1" + ) + foreign_record = SocketRecord(instance_id="other-instance", sid="sid2") + foreign_key = redis_manager._get_redis_key("foreign") + mock_redis.get.side_effect = lambda key: ( + pickle.dumps(foreign_record) if key == foreign_key else None + ) + app, modified_tokens = _mock_app(redis_manager) + + await _run_update_other_tokens(app, {"local", "foreign", "disconnected"}) + + assert sorted(modified_tokens) == ["foreign", "local"] + # The foreign socket record is cached locally for later emit_update routing. + assert redis_manager.token_to_socket["foreign"] == foreign_record From 352aa4d9472d135a7f7b967cc3bf7ba30dd01bbd Mon Sep 17 00:00:00 2001 From: Benedikt Bartscher Date: Mon, 24 Aug 2026 12:29:49 +0200 Subject: [PATCH 2/5] address review comments --- reflex/istate/shared.py | 14 ++------- reflex/utils/token_manager.py | 39 +++++++++++++++++++++++++ tests/units/istate/test_shared.py | 3 ++ tests/units/utils/test_token_manager.py | 32 ++++++++++++++++++++ 4 files changed, 77 insertions(+), 11 deletions(-) diff --git a/reflex/istate/shared.py b/reflex/istate/shared.py index e178e6a9e05..432cdadbb52 100644 --- a/reflex/istate/shared.py +++ b/reflex/istate/shared.py @@ -52,8 +52,6 @@ def _do_update_other_tokens( Returns: The list of asyncio tasks created to perform the updates. """ - from reflex.utils.token_manager import RedisTokenManager - app = RegistrationContext.get().app tasks = [] @@ -62,15 +60,9 @@ def _do_update_other_tokens( token_manager = event_namespace._token_manager async def _update_client(token: str): - # Don't send updates for disconnected clients. The local - # token_to_socket map only tracks sockets owned by this instance, so - # with redis the socket record is resolved (and cached) from redis - # instead; emit_update then relays the delta to the owning instance - # via the lost-and-found channel. - if isinstance(token_manager, RedisTokenManager): - if await token_manager._get_token_owner(token) is None: - return - elif token not in token_manager.token_to_socket: + # Don't send updates for disconnected clients; emit_update relays the + # delta to the owning instance if the socket lives elsewhere. + if not await token_manager.is_token_connected(token): return async with app.modify_state( BaseStateToken(ident=token, cls=state_type), diff --git a/reflex/utils/token_manager.py b/reflex/utils/token_manager.py index 93d5f88393c..73c3da4a8fd 100644 --- a/reflex/utils/token_manager.py +++ b/reflex/utils/token_manager.py @@ -81,6 +81,17 @@ async def enumerate_tokens(self) -> AsyncIterator[str]: for token in self.token_to_socket: yield token + async def is_token_connected(self, token: str) -> bool: + """Whether the token has a connected client socket on any instance. + + Args: + token: The client token. + + Returns: + True if the token has a connected socket. + """ + return token in self.token_to_socket + @abstractmethod async def link_token_to_sid(self, token: str, sid: str) -> str | None: """Link a token to a session ID. @@ -443,6 +454,34 @@ async def _get_token_owner(self, token: str, refresh: bool = False) -> str | Non logger.error(f"Redis error getting token owner: {e}") return None + async def is_token_connected(self, token: str) -> bool: + """Whether the token has a connected client socket on any instance. + + A record owned by this instance is authoritative. A cached record + from another instance may be stale (the client may have reconnected + elsewhere), so the socket record is refreshed from redis instead, + and dropped from the local cache if the client is gone. + + Args: + token: The client token. + + Returns: + True if the token has a connected socket on any instance. + """ + if ( + socket_record := self.token_to_socket.get(token) + ) is not None and socket_record.instance_id == self.instance_id: + return True + if await self._get_token_owner(token, refresh=True) is not None: + return True + if ( + socket_record is not None + and self.token_to_socket.get(token) is socket_record + ): + self.token_to_socket.pop(token, None) + self.sid_to_token.pop(socket_record.sid, None) + return False + async def emit_lost_and_found( self, token: str, diff --git a/tests/units/istate/test_shared.py b/tests/units/istate/test_shared.py index e8b99d6a8e1..c8b16c874dd 100644 --- a/tests/units/istate/test_shared.py +++ b/tests/units/istate/test_shared.py @@ -106,3 +106,6 @@ async def test_update_other_tokens_redis_cross_instance(redis_manager, mock_redi assert sorted(modified_tokens) == ["foreign", "local"] # The foreign socket record is cached locally for later emit_update routing. assert redis_manager.token_to_socket["foreign"] == foreign_record + # Locally owned sockets are authoritative and never require a redis lookup. + local_key = redis_manager._get_redis_key("local") + assert local_key not in [call.args[0] for call in mock_redis.get.call_args_list] diff --git a/tests/units/utils/test_token_manager.py b/tests/units/utils/test_token_manager.py index 208387cd86b..2703b8c81d7 100644 --- a/tests/units/utils/test_token_manager.py +++ b/tests/units/utils/test_token_manager.py @@ -477,6 +477,38 @@ async def test_various_redis_errors_handled_gracefully( assert result is None mock_super.assert_called_once() + async def test_is_token_connected_locally_owned(self, manager, mock_redis): + """A locally owned socket record is authoritative, without a redis lookup.""" + manager.token_to_socket["token1"] = SocketRecord( + instance_id=manager.instance_id, sid="sid1" + ) + + assert await manager.is_token_connected("token1") + mock_redis.get.assert_not_called() + + async def test_is_token_connected_stale_foreign_record(self, manager, mock_redis): + """A cached foreign record is refreshed from redis and dropped when gone.""" + manager.token_to_socket["token1"] = SocketRecord( + instance_id="other-instance", sid="sid1" + ) + manager.sid_to_token["sid1"] = "token1" + mock_redis.get = AsyncMock(return_value=None) + + assert not await manager.is_token_connected("token1") + assert "token1" not in manager.token_to_socket + assert "sid1" not in manager.sid_to_token + + async def test_is_token_connected_foreign_record_moved(self, manager, mock_redis): + """A cached foreign record is replaced when the client moved instances.""" + manager.token_to_socket["token1"] = SocketRecord( + instance_id="old-instance", sid="sid1" + ) + new_record = SocketRecord(instance_id="new-instance", sid="sid2") + mock_redis.get = AsyncMock(return_value=pickle.dumps(new_record)) + + assert await manager.is_token_connected("token1") + assert manager.token_to_socket["token1"] == new_record + def test_inheritance_from_local_manager(self, manager): """Test RedisTokenManager inherits from LocalTokenManager. From 5833800f418cc6b0731cdbcaeadd51ed6c6ee81a Mon Sep 17 00:00:00 2001 From: Benedikt Bartscher Date: Mon, 24 Aug 2026 12:33:00 +0200 Subject: [PATCH 3/5] add news fragment --- news/6934.bugfix.md | 1 + 1 file changed, 1 insertion(+) create mode 100644 news/6934.bugfix.md diff --git a/news/6934.bugfix.md b/news/6934.bugfix.md new file mode 100644 index 00000000000..4ac25a0107e --- /dev/null +++ b/news/6934.bugfix.md @@ -0,0 +1 @@ +Shared state updates now reach linked clients connected to other backend instances — the fan-out previously skipped any client whose websocket was not connected to the instance processing the event, so with redis and multiple workers only same-instance clients received live updates. From b643ed80019cab68df6353600f8b3ae966ec7577 Mon Sep 17 00:00:00 2001 From: Benedikt Bartscher Date: Mon, 24 Aug 2026 13:27:38 +0200 Subject: [PATCH 4/5] cubic --- reflex/utils/token_manager.py | 49 +++++++++++++++++++------ tests/units/utils/test_token_manager.py | 19 +++++++++- 2 files changed, 56 insertions(+), 12 deletions(-) diff --git a/reflex/utils/token_manager.py b/reflex/utils/token_manager.py index 73c3da4a8fd..de0d00a11cf 100644 --- a/reflex/utils/token_manager.py +++ b/reflex/utils/token_manager.py @@ -442,17 +442,39 @@ async def _get_token_owner(self, token: str, refresh: bool = False) -> str | Non ): return socket_record.instance_id - redis_key = self._get_redis_key(token) try: - record_pkl = await self.redis.get(redis_key) - if record_pkl: - socket_record = pickle.loads(record_pkl) - self.token_to_socket[token] = socket_record - self.sid_to_token[socket_record.sid] = token - return socket_record.instance_id + socket_record = await self._fetch_socket_record(token) except Exception as e: logger.error(f"Redis error getting token owner: {e}") - return None + return None + return socket_record.instance_id if socket_record is not None else None + + async def _fetch_socket_record(self, token: str) -> SocketRecord | None: + """Fetch the socket record for a token from redis and cache it. + + Unlike _get_token_owner, redis errors propagate to the caller so it + can distinguish a lookup failure from an absent record. + + Args: + token: The client token. + + Returns: + The refreshed socket record, or None if the token has none. + """ + record_pkl = await self.redis.get(self._get_redis_key(token)) + if not record_pkl: + return None + socket_record = pickle.loads(record_pkl) + # Drop the reverse mapping of a superseded record (client moved sids). + if ( + (previous := self.token_to_socket.get(token)) is not None + and previous.sid != socket_record.sid + and self.sid_to_token.get(previous.sid) == token + ): + self.sid_to_token.pop(previous.sid, None) + self.token_to_socket[token] = socket_record + self.sid_to_token[socket_record.sid] = token + return socket_record async def is_token_connected(self, token: str) -> bool: """Whether the token has a connected client socket on any instance. @@ -460,7 +482,8 @@ async def is_token_connected(self, token: str) -> bool: A record owned by this instance is authoritative. A cached record from another instance may be stale (the client may have reconnected elsewhere), so the socket record is refreshed from redis instead, - and dropped from the local cache if the client is gone. + and dropped from the local cache if the client is gone. If the + refresh fails, the cached record is preserved and trusted. Args: token: The client token. @@ -472,8 +495,12 @@ async def is_token_connected(self, token: str) -> bool: socket_record := self.token_to_socket.get(token) ) is not None and socket_record.instance_id == self.instance_id: return True - if await self._get_token_owner(token, refresh=True) is not None: - return True + try: + if await self._fetch_socket_record(token) is not None: + return True + except Exception as e: + logger.warning(f"Redis error checking token connection: {e}") + return socket_record is not None if ( socket_record is not None and self.token_to_socket.get(token) is socket_record diff --git a/tests/units/utils/test_token_manager.py b/tests/units/utils/test_token_manager.py index 2703b8c81d7..295f40604f8 100644 --- a/tests/units/utils/test_token_manager.py +++ b/tests/units/utils/test_token_manager.py @@ -499,15 +499,32 @@ async def test_is_token_connected_stale_foreign_record(self, manager, mock_redis assert "sid1" not in manager.sid_to_token async def test_is_token_connected_foreign_record_moved(self, manager, mock_redis): - """A cached foreign record is replaced when the client moved instances.""" + """A cached foreign record and its sid mapping are replaced on a move.""" manager.token_to_socket["token1"] = SocketRecord( instance_id="old-instance", sid="sid1" ) + manager.sid_to_token["sid1"] = "token1" new_record = SocketRecord(instance_id="new-instance", sid="sid2") mock_redis.get = AsyncMock(return_value=pickle.dumps(new_record)) assert await manager.is_token_connected("token1") assert manager.token_to_socket["token1"] == new_record + assert "sid1" not in manager.sid_to_token + assert manager.sid_to_token["sid2"] == "token1" + + async def test_is_token_connected_redis_error_trusts_cache( + self, manager, mock_redis + ): + """A redis failure preserves and trusts the cached foreign record.""" + record = SocketRecord(instance_id="other-instance", sid="sid1") + manager.token_to_socket["token1"] = record + manager.sid_to_token["sid1"] = "token1" + mock_redis.get = AsyncMock(side_effect=Exception("Redis down")) + + assert await manager.is_token_connected("token1") + assert not await manager.is_token_connected("unknown-token") + assert manager.token_to_socket["token1"] == record + assert manager.sid_to_token["sid1"] == "token1" def test_inheritance_from_local_manager(self, manager): """Test RedisTokenManager inherits from LocalTokenManager. From 09d67ad55e2147919b6ef9c969eab1628141e024 Mon Sep 17 00:00:00 2001 From: Benedikt Bartscher Date: Tue, 25 Aug 2026 09:40:40 +0200 Subject: [PATCH 5/5] feat: unsubscribe disconnected clients from shared states after a grace period Resolves the TODO in _do_update_other_tokens: tokens of clients that disconnected and did not reconnect within REFLEX_SHARED_STATE_DISCONNECT_GRACE (default 30s, 0 disables) are removed from the _linked_from subscriber sets of all shared states the client was linked to, so shared mutations stop fanning out modify_state churn for them. The new SharedState._on_subscriber_disconnected hook lets shared states clean up per-client data (e.g. presence bookkeeping). Safety: a reconnect within the grace (any instance, checked via is_token_connected) is a no-op; a client reaped too eagerly re-subscribes automatically on its next event through _internal_patch_linked_state. A new disconnect for the same token restarts the grace. Only the shared subscriber sets are touched, never the client's own _reflex_internal_links. --- news/+shared-state-disconnect-reap.feature.md | 1 + .../src/reflex_base/environment.py | 4 + reflex/app.py | 5 + reflex/istate/shared.py | 104 ++++++++++++++++- tests/units/istate/test_shared.py | 107 +++++++++++++++++- 5 files changed, 218 insertions(+), 3 deletions(-) create mode 100644 news/+shared-state-disconnect-reap.feature.md diff --git a/news/+shared-state-disconnect-reap.feature.md b/news/+shared-state-disconnect-reap.feature.md new file mode 100644 index 00000000000..8db8e6b581c --- /dev/null +++ b/news/+shared-state-disconnect-reap.feature.md @@ -0,0 +1 @@ +Disconnected clients are unsubscribed from their linked shared states after a reconnect grace period (`REFLEX_SHARED_STATE_DISCONNECT_GRACE`, default 30s, 0 disables). The new `SharedState._on_subscriber_disconnected` hook lets shared states clean up per-client data (e.g. presence bookkeeping); a client reconnecting later re-subscribes automatically with its next event. diff --git a/packages/reflex-base/src/reflex_base/environment.py b/packages/reflex-base/src/reflex_base/environment.py index 74993868c32..7274695af99 100644 --- a/packages/reflex-base/src/reflex_base/environment.py +++ b/packages/reflex-base/src/reflex_base/environment.py @@ -688,6 +688,10 @@ class EnvironmentVariables: # The address to bind the HTTP client to. You can set this to "::" to enable IPv6. REFLEX_HTTP_CLIENT_BIND_ADDRESS: EnvVar[str | None] = env_var(None) + # Seconds a disconnected client may reconnect before it is unsubscribed + # from the shared states it was linked to. 0 disables the unsubscription. + REFLEX_SHARED_STATE_DISCONNECT_GRACE: EnvVar[int] = env_var(30) + # Maximum size of the message in the websocket server in bytes. REFLEX_SOCKET_MAX_HTTP_BUFFER_SIZE: EnvVar[int] = env_var( constants.POLLING_MAX_HTTP_BUFFER_SIZE diff --git a/reflex/app.py b/reflex/app.py index 63ec53a4c75..5d6c0063e16 100644 --- a/reflex/app.py +++ b/reflex/app.py @@ -2010,10 +2010,15 @@ def on_disconnect(self, sid: str) -> asyncio.Task | None: Returns: An asyncio Task for cleaning up the token, or None. """ + from reflex.istate.shared import schedule_disconnect_reap + self._client_error_counts.pop(sid, None) # Get token before cleaning up disconnect_token = self.sid_to_token.get(sid) if disconnect_token: + # Unsubscribe the client from its linked shared states unless it + # reconnects within the grace period. + schedule_disconnect_reap(self.app, disconnect_token) # Use async cleanup through token manager task = asyncio.create_task( self._token_manager.disconnect_token(disconnect_token, sid), diff --git a/reflex/istate/shared.py b/reflex/istate/shared.py index 432cdadbb52..9726443d58d 100644 --- a/reflex/istate/shared.py +++ b/reflex/istate/shared.py @@ -3,10 +3,12 @@ import asyncio import contextlib import logging +import time from collections.abc import AsyncIterator -from typing import TypeVar +from typing import TYPE_CHECKING, TypeVar from reflex_base.constants import ROUTER_DATA +from reflex_base.environment import environment from reflex_base.event import Event, get_hydrate_event from reflex_base.registry import RegistrationContext from reflex_base.utils.exceptions import ReflexRuntimeError @@ -15,9 +17,13 @@ from reflex.istate.manager.token import BaseStateToken from reflex.state import BaseState, State, _override_base_method +if TYPE_CHECKING: + from reflex.app import App + logger = logging.getLogger(__name__) UPDATE_OTHER_CLIENT_TASKS: set[asyncio.Task] = set() +DISCONNECT_REAP_TASKS: dict[str, asyncio.Task] = {} LINKED_STATE = TypeVar("LINKED_STATE", bound="SharedStateBaseInternal") @@ -71,7 +77,9 @@ async def _update_client(token: str): pass for affected_token in affected_tokens: - # TODO: remove disconnected clients after some time. + # Disconnected clients are removed by schedule_disconnect_reap after + # the reconnect grace; until then the connectivity check above skips + # them. t = asyncio.create_task(_update_client(affected_token)) UPDATE_OTHER_CLIENT_TASKS.add(t) t.add_done_callback(_log_update_client_errors) @@ -79,6 +87,86 @@ async def _update_client(token: str): return tasks +def schedule_disconnect_reap(app: "App", token: str) -> asyncio.Task | None: + """Schedule unsubscribing a disconnected client from its linked shared states. + + Called on client disconnect. The reap runs after a grace period + (REFLEX_SHARED_STATE_DISCONNECT_GRACE) so brief reconnects (page reloads, + network blips) are no-ops; a client reaped too eagerly re-subscribes on its + next event through _internal_patch_linked_state. A new disconnect for the + same token restarts the grace. + + Args: + app: The application object. + token: The client token that disconnected. + + Returns: + The scheduled reap task, or None when reaping is disabled. + """ + grace = environment.REFLEX_SHARED_STATE_DISCONNECT_GRACE.get() + if grace <= 0 or app._state is None: + return None + if (previous := DISCONNECT_REAP_TASKS.pop(token, None)) is not None: + previous.cancel() + + task = asyncio.create_task( + _reap_disconnected_client(app, token, grace), + name=f"reflex_shared_state_reap|{token}|{time.time()}", + ) + DISCONNECT_REAP_TASKS[token] = task + + def _on_done(task: asyncio.Task) -> None: + if DISCONNECT_REAP_TASKS.get(token) is task: + DISCONNECT_REAP_TASKS.pop(token, None) + if not task.cancelled() and (exc := task.exception()) is not None: + logger.warning(f"Error reaping disconnected shared state client: {exc}") + + task.add_done_callback(_on_done) + return task + + +async def _reap_disconnected_client(app: "App", token: str, grace: float) -> None: + """Unsubscribe a client from its linked shared states unless it reconnected. + + Args: + app: The application object. + token: The client token that disconnected. + grace: Seconds to wait for a reconnect before unsubscribing. + """ + await asyncio.sleep(grace) + if (event_namespace := app.event_namespace) is None: + return + if await event_namespace._token_manager.is_token_connected(token): + # The client reconnected, possibly to another instance. + return + if app._state is None: + return + # Read-only peek at the client's links; racing a concurrent event is fine, + # since any later event through the link re-adds the client anyway. + root_state = await app.state_manager.get_state( + BaseStateToken(ident=token, cls=app._state) + ) + if not isinstance(root_state, State): + return + links = dict(root_state._reflex_internal_links or {}) + for state_name, linked_token in links.items(): + try: + state_cls = app._state.get_class_substate(state_name) + async with app.modify_state( + BaseStateToken(ident=linked_token, cls=state_cls) + ) as shared_root: + shared = await shared_root.get_state(state_cls) + if ( + not isinstance(shared, SharedState) + or token not in shared._linked_from + ): + continue + shared._linked_from.discard(token) + await shared._on_subscriber_disconnected(token) + except Exception as e: + logger.warning(f"Error unsubscribing disconnected shared state client: {e}") + + @contextlib.asynccontextmanager async def _patch_state( original_state: BaseState, linked_state: BaseState, full_delta: bool = False @@ -506,6 +594,18 @@ class SharedState(SharedStateBaseInternal, mixin=True): _linked_to: str = "" _previous_dirty_vars: set[str] = set() + async def _on_subscriber_disconnected(self, client_token: str) -> None: + """Hook called when a subscribed client is unsubscribed after disconnecting. + + Override to clean up per-client data kept on the shared state (e.g. + presence bookkeeping). Runs with the shared token's state locked, after + the client token was removed from the subscriber set; mutations + propagate to the remaining linked clients. + + Args: + client_token: The client token that was unsubscribed. + """ + @classmethod def __init_subclass__(cls, **kwargs): """Initialize subclass and set up shared state fields. diff --git a/tests/units/istate/test_shared.py b/tests/units/istate/test_shared.py index c8b16c874dd..f1621acdce4 100644 --- a/tests/units/istate/test_shared.py +++ b/tests/units/istate/test_shared.py @@ -7,7 +7,13 @@ import pytest -from reflex.istate.shared import _do_update_other_tokens +from reflex.istate.shared import ( + DISCONNECT_REAP_TASKS, + SharedState, + _do_update_other_tokens, + _reap_disconnected_client, + schedule_disconnect_reap, +) from reflex.state import State from reflex.utils.token_manager import ( LocalTokenManager, @@ -109,3 +115,102 @@ async def test_update_other_tokens_redis_cross_instance(redis_manager, mock_redi # Locally owned sockets are authoritative and never require a redis lookup. local_key = redis_manager._get_redis_key("local") assert local_key not in [call.args[0] for call in mock_redis.get.call_args_list] + + +def _mock_reap_app(token_manager) -> tuple[Mock, Mock, list[str]]: + """Create a mock app with one client linked to one shared state. + + Returns: + The mock app, the mock shared state and the list collecting the + token idents passed to modify_state. + """ + modified_tokens: list[str] = [] + + shared = Mock(spec=SharedState) + shared._linked_from = {"gone", "other"} + shared._on_subscriber_disconnected = AsyncMock() + + @asynccontextmanager + async def modify_state(token, previous_dirty_vars=None): + modified_tokens.append(token.ident) + shared_root = Mock() + shared_root.get_state = AsyncMock(return_value=shared) + yield shared_root + + root_state = Mock(spec=State) + root_state._reflex_internal_links = {"state.shared_sub": "room-token"} + + app = Mock() + app.modify_state = modify_state + app.event_namespace = Mock() + app.event_namespace._token_manager = token_manager + app.state_manager.get_state = AsyncMock(return_value=root_state) + app._state = Mock() + app._state.get_class_substate = Mock(return_value=State) + return app, shared, modified_tokens + + +async def test_reap_disconnected_client_unsubscribes(): + """A client that stayed disconnected is removed from its linked states.""" + app, shared, modified_tokens = _mock_reap_app(LocalTokenManager()) + + await _reap_disconnected_client(app, "gone", grace=0) + + assert modified_tokens == ["room-token"] + assert shared._linked_from == {"other"} + shared._on_subscriber_disconnected.assert_awaited_once_with("gone") + + +async def test_reap_skips_reconnected_client(): + """A client that reconnected within the grace is left subscribed.""" + manager = LocalTokenManager() + manager.token_to_socket["gone"] = SocketRecord( + instance_id=manager.instance_id, sid="sid1" + ) + app, shared, modified_tokens = _mock_reap_app(manager) + + await _reap_disconnected_client(app, "gone", grace=0) + + assert modified_tokens == [] + assert shared._linked_from == {"gone", "other"} + shared._on_subscriber_disconnected.assert_not_awaited() + + +async def test_reap_skips_unsubscribed_client(): + """A token no longer in the subscriber set does not trigger the hook.""" + app, shared, modified_tokens = _mock_reap_app(LocalTokenManager()) + + await _reap_disconnected_client(app, "unknown", grace=0) + + assert modified_tokens == ["room-token"] + assert shared._linked_from == {"gone", "other"} + shared._on_subscriber_disconnected.assert_not_awaited() + + +def test_schedule_disconnect_reap_disabled(monkeypatch): + """Grace 0 disables scheduling entirely.""" + monkeypatch.setenv("REFLEX_SHARED_STATE_DISCONNECT_GRACE", "0") + app, _, _ = _mock_reap_app(LocalTokenManager()) + + assert schedule_disconnect_reap(app, "gone") is None + assert not DISCONNECT_REAP_TASKS + + +async def test_schedule_disconnect_reap_restarts_grace(monkeypatch): + """A new disconnect for the same token cancels the pending reap.""" + monkeypatch.setenv("REFLEX_SHARED_STATE_DISCONNECT_GRACE", "60") + app, _, modified_tokens = _mock_reap_app(LocalTokenManager()) + + first = schedule_disconnect_reap(app, "gone") + second = schedule_disconnect_reap(app, "gone") + assert first is not None + assert second is not None + with pytest.raises(asyncio.CancelledError): + await first + assert {"gone": second} == DISCONNECT_REAP_TASKS + + second.cancel() + with pytest.raises(asyncio.CancelledError): + await second + assert not DISCONNECT_REAP_TASKS + assert modified_tokens == []