From f7cbca6a4d751ee04a2e95f79c7eff2f0f989662 Mon Sep 17 00:00:00 2001 From: Benedikt Bartscher Date: Mon, 24 Aug 2026 12:18:13 +0200 Subject: [PATCH 1/4] 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/4] 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/4] 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/4] 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.