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..44a18592698 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. @@ -240,22 +251,29 @@ async def enumerate_tokens(self) -> AsyncIterator[str]: if not cursor: break - async def _handle_socket_record_del( - self, token: str, expired: bool = False - ) -> None: + async def _handle_socket_record_del(self, token: str) -> None: """Handle deletion of a socket record from Redis. + Disconnects pop the local mapping before touching redis, so a + self-owned record still present locally was created by a newer link + and refers to a live socket; the deletion is an expiration or an + outdated notification, and the record is re-stored to keep it alive. + A cached foreign record is dropped and refetched on demand. + Args: token: The client token whose record was deleted. - expired: Whether the deletion was due to expiration. """ - if ( - socket_record := self.token_to_socket.pop(token, None) - ) is not None and socket_record.instance_id == self.instance_id: - self.sid_to_token.pop(socket_record.sid, None) - if expired: - # Keep the record alive as long as this process is alive and not deleted. - await self.link_token_to_sid(token, socket_record.sid) + if (socket_record := self.token_to_socket.get(token)) is None: + return + if socket_record.instance_id == self.instance_id: + # Restore only if the key is still absent: a newer record (a + # relink here or a claim by another instance) must win. + if not await self._store_socket_record(token, socket_record, nx=True): + # Adopt the newer record without waiting for its set + # notification, which may be lost across pubsub reconnects. + await self._get_token_owner(token, refresh=True) + else: + self.token_to_socket.pop(token, None) async def _subscribe_socket_record_updates(self) -> None: """Subscribe to Redis keyspace notifications for socket record updates.""" @@ -277,10 +295,7 @@ async def _subscribe_socket_record_updates(self) -> None: event = message["data"].decode() if event in ("del", "expired", "evicted"): - await self._handle_socket_record_del( - token, - expired=(event == "expired"), - ) + await self._handle_socket_record_del(token) elif event == "set": await self._get_token_owner(token, refresh=True) @@ -325,7 +340,6 @@ async def link_token_to_sid(self, token: str, sid: str) -> str | None: if token_exists_in_redis: # Duplicate exists somewhere - generate new token token = new_token = _get_new_token() - redis_key = self._get_redis_key(new_token) # Store in local dicts socket_record = self.token_to_socket[token] = SocketRecord( @@ -333,17 +347,35 @@ async def link_token_to_sid(self, token: str, sid: str) -> str | None: ) self.sid_to_token[sid] = token - # Store in Redis if possible + await self._store_socket_record(token, socket_record) + # Return the new token if one was generated + return new_token + + async def _store_socket_record( + self, token: str, socket_record: SocketRecord, nx: bool = False + ) -> bool: + """Store a socket record in Redis, logging errors instead of raising. + + Args: + token: The client token. + socket_record: The record to store. + nx: Only store the record if the key does not already exist. + + Returns: + True if the record was stored, False if rejected (nx) or on error. + """ try: - await self.redis.set( - redis_key, - pickle.dumps(socket_record), - ex=self.token_expiration, + return bool( + await self.redis.set( + self._get_redis_key(token), + pickle.dumps(socket_record), + ex=self.token_expiration, + nx=nx, + ) ) except Exception as e: logger.error(f"Redis error storing token: {e}") - # Return the new token if one was generated - return new_token + return False async def disconnect_token(self, token: str, sid: str) -> None: """Clean up token mapping when client disconnects. @@ -358,16 +390,17 @@ async def disconnect_token(self, token: str, sid: str) -> None: and socket_record.sid == sid and socket_record.instance_id == self.instance_id ): - # Clean up Redis + # Drop the local mapping before the redis round-trip so the + # locally-owned fast paths stop treating the token as connected + # while the delete is in flight. + await super().disconnect_token(token, sid) + redis_key = self._get_redis_key(token) try: await self.redis.delete(redis_key) except Exception as e: logger.error(f"Redis error deleting token: {e}") - # Clean up local dicts (always do this) - await super().disconnect_token(token, sid) - @staticmethod def _get_lost_and_found_key(instance_id: str) -> str: """Get the Redis key for lost and found deltas for an instance. @@ -431,17 +464,79 @@ 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. + + Redis errors propagate to the caller so it can distinguish a lookup + failure from an absent record. This instance is authoritative for its + own sockets: a record claiming this instance without a live local + link is a stale leftover of a disconnected socket (its delete is + still in flight or failed) and is treated as absent. + + Args: + token: The client token. + + Returns: + The refreshed socket record, or None if the token has no live socket. + """ + record_pkl = await self.redis.get(self._get_redis_key(token)) + if not record_pkl: + return None + socket_record = pickle.loads(record_pkl) + # Stale leftover of one of this instance's own disconnected sockets. + if ( + socket_record.instance_id == self.instance_id + and self.sid_to_token.get(socket_record.sid) != token + ): + return None + # 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, so it is refreshed from redis + (dropped if the client is gone) and trusted as-is if the refresh fails. + + 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..f8690b0d6b5 100644 --- a/tests/units/utils/test_token_manager.py +++ b/tests/units/utils/test_token_manager.py @@ -303,6 +303,7 @@ async def test_link_token_to_sid_normal_case(self, manager, mock_redis): f"token_manager_socket_record_{token}", pickle.dumps(SocketRecord(instance_id=manager.instance_id, sid=sid)), ex=3600, + nx=False, ) assert manager.token_to_socket[token].sid == sid assert manager.sid_to_token[sid] == token @@ -350,6 +351,7 @@ async def test_link_token_to_sid_duplicate_detected(self, manager, mock_redis): f"token_manager_socket_record_{result}", pickle.dumps(SocketRecord(instance_id=manager.instance_id, sid=sid)), ex=3600, + nx=False, ) assert manager.token_to_sid[result] == sid assert manager.sid_to_token[sid] == result @@ -427,6 +429,121 @@ async def test_disconnect_token_not_owned_locally(self, manager, mock_redis): mock_redis.delete.assert_not_called() + async def test_disconnect_token_clears_local_before_redis_delete( + self, manager, mock_redis + ): + """Local mappings are dropped before the redis delete round-trip starts.""" + token, sid = "token1", "sid1" + manager.token_to_socket[token] = SocketRecord( + instance_id=manager.instance_id, sid=sid + ) + manager.sid_to_token[sid] = token + + cached_at_delete = {} + + def delete(key): + cached_at_delete["token"] = token in manager.token_to_socket + cached_at_delete["sid"] = sid in manager.sid_to_token + + mock_redis.delete.side_effect = delete + + await manager.disconnect_token(token, sid) + + mock_redis.delete.assert_called_once_with( + f"token_manager_socket_record_{token}" + ) + assert cached_at_delete == {"token": False, "sid": False} + + async def test_is_token_connected_ignores_stale_own_record( + self, manager, mock_redis + ): + """A redis record claiming this instance without a live local link is stale.""" + stale = SocketRecord(instance_id=manager.instance_id, sid="sid1") + mock_redis.get = AsyncMock(return_value=pickle.dumps(stale)) + + 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_fetch_socket_record_own_record_with_live_link( + self, manager, mock_redis + ): + """A redis record claiming this instance with a live local link is cached.""" + record = SocketRecord(instance_id=manager.instance_id, sid="sid1") + manager.sid_to_token["sid1"] = "token1" + mock_redis.get = AsyncMock(return_value=pickle.dumps(record)) + + assert await manager._fetch_socket_record("token1") == record + assert manager.token_to_socket["token1"] == record + + async def test_socket_record_del_keeps_live_local_mapping( + self, manager, mock_redis + ): + """A late redis del notification must not drop a live local mapping.""" + token, sid = "token1", "sid2" + record = SocketRecord(instance_id=manager.instance_id, sid=sid) + manager.token_to_socket[token] = record + manager.sid_to_token[sid] = token + + await manager._handle_socket_record_del(token) + + assert manager.token_to_socket[token] == record + assert manager.sid_to_token[sid] == token + mock_redis.set.assert_called_once_with( + f"token_manager_socket_record_{token}", + pickle.dumps(record), + ex=3600, + nx=True, + ) + + async def test_socket_record_del_adopts_newer_record_when_restore_loses( + self, manager, mock_redis + ): + """A rejected NX restore adopts the newer redis record immediately.""" + token, sid = "token1", "sid1" + stale = SocketRecord(instance_id=manager.instance_id, sid=sid) + manager.token_to_socket[token] = stale + manager.sid_to_token[sid] = token + newer = SocketRecord(instance_id="other-instance", sid="sid2") + mock_redis.set.return_value = None + mock_redis.get = AsyncMock(return_value=pickle.dumps(newer)) + + await manager._handle_socket_record_del(token) + + assert manager.token_to_socket[token] == newer + assert manager.sid_to_token["sid2"] == token + assert sid not in manager.sid_to_token + + async def test_socket_record_del_drops_foreign_cache(self, manager, mock_redis): + """A del notification drops a cached foreign record.""" + manager.token_to_socket["token1"] = SocketRecord( + instance_id="other-instance", sid="sid9" + ) + + await manager._handle_socket_record_del("token1") + + assert "token1" not in manager.token_to_socket + mock_redis.set.assert_not_called() + + async def test_is_token_connected_after_late_disconnect_notification( + self, manager, mock_redis + ): + """A late del notification for a relinked token must not break delivery.""" + token = "token1" + mock_redis.exists.return_value = False + await manager.link_token_to_sid(token, "sid1") + # Old socket disconnects; the client relinks before the notification arrives. + await manager.disconnect_token(token, "sid1") + await manager.link_token_to_sid(token, "sid2") + # The keyspace notification for the old delete arrives late. + await manager._handle_socket_record_del(token) + + record = manager.token_to_socket.get(token) + mock_redis.get = AsyncMock( + return_value=pickle.dumps(record) if record else None + ) + assert await manager.is_token_connected(token) + async def test_disconnect_token_redis_error(self, manager, mock_redis): """Test disconnect continues with local cleanup even if Redis fails. @@ -477,6 +594,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. @@ -647,6 +813,41 @@ async def _wait_for_call_count_positive(mock: Mock, timeout: float = 5.0): await asyncio.sleep(0.1) +@pytest.mark.usefixtures("redis_url") +@pytest.mark.asyncio +async def test_redis_token_manager_restore_does_not_clobber_new_owner( + event_namespace_factory: Callable[[], EventNamespace], +): + """A late keep-alive restore must not overwrite another instance's record. + + Args: + event_namespace_factory: Factory fixture for EventNamespace instances. + """ + event_namespace1 = event_namespace_factory() + event_namespace2 = event_namespace_factory() + + manager1 = event_namespace1._token_manager + manager2 = event_namespace2._token_manager + assert isinstance(manager1, RedisTokenManager) + assert isinstance(manager2, RedisTokenManager) + + await event_namespace1.on_connect(sid="sid1", environ=query_string_for("token1")) + # Stop the live subscriber so the deletion notification is handled manually. + assert manager1._socket_record_task is not None + manager1._socket_record_task.cancel() + # The record expires and another instance claims the token. + await manager1.redis.delete(manager1._get_redis_key("token1")) + await event_namespace2.on_connect(sid="sid2", environ=query_string_for("token1")) + # The first instance processes the expiration notification late. + await manager1._handle_socket_record_del("token1") + + assert await manager2._get_token_owner("token1", refresh=True) == ( + manager2.instance_id + ) + # The losing restore adopted the newer record locally. + assert manager1.token_to_socket["token1"].instance_id == manager2.instance_id + + @pytest.mark.usefixtures("redis_url") @pytest.mark.asyncio async def test_redis_token_manager_lost_and_found(