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. diff --git a/reflex/istate/shared.py b/reflex/istate/shared.py index e38517fef66..432cdadbb52 100644 --- a/reflex/istate/shared.py +++ b/reflex/istate/shared.py @@ -54,20 +54,23 @@ def _do_update_other_tokens( """ 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; 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), 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/reflex/utils/token_manager.py b/reflex/utils/token_manager.py index 93d5f88393c..de0d00a11cf 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. @@ -431,17 +442,72 @@ 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. + + 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. If the + refresh fails, the cached record is preserved and trusted. + + 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 + 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 + ): + 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, diff --git a/tests/units/istate/test_shared.py b/tests/units/istate/test_shared.py new file mode 100644 index 00000000000..c8b16c874dd --- /dev/null +++ b/tests/units/istate/test_shared.py @@ -0,0 +1,111 @@ +"""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 + # 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..295f40604f8 100644 --- a/tests/units/utils/test_token_manager.py +++ b/tests/units/utils/test_token_manager.py @@ -477,6 +477,55 @@ 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 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.