From cec2c387f39020189547255c3e6f7f3c5ecdb393 Mon Sep 17 00:00:00 2001 From: Anas Khan <83116240+anxkhn@users.noreply.github.com> Date: Mon, 13 Jul 2026 21:59:02 +0530 Subject: [PATCH 1/2] fix(tokens): claim redis socket tokens atomically Signed-off-by: Anas Khan <83116240+anxkhn@users.noreply.github.com> --- news/6740.bugfix.md | 1 + reflex/utils/token_manager.py | 37 ++++++-------- tests/units/utils/test_token_manager.py | 65 ++++++++++++++++++++----- 3 files changed, 67 insertions(+), 36 deletions(-) create mode 100644 news/6740.bugfix.md diff --git a/news/6740.bugfix.md b/news/6740.bugfix.md new file mode 100644 index 00000000000..037d7627122 --- /dev/null +++ b/news/6740.bugfix.md @@ -0,0 +1 @@ +Redis-backed socket token claims are now atomic across workers, preventing simultaneous connections from sharing client state. diff --git a/reflex/utils/token_manager.py b/reflex/utils/token_manager.py index e8f92fa96d0..256438ffc0f 100644 --- a/reflex/utils/token_manager.py +++ b/reflex/utils/token_manager.py @@ -309,36 +309,27 @@ async def link_token_to_sid(self, token: str, sid: str) -> str | None: # Make sure the update subscriber is running self._ensure_socket_record_task() - # Check Redis for cross-worker duplicates - redis_key = self._get_redis_key(token) + socket_record = SocketRecord(instance_id=self.instance_id, sid=sid) + socket_record_data = pickle.dumps(socket_record) + new_token = None try: - token_exists_in_redis = await self.redis.exists(redis_key) + while not await self.redis.set( + self._get_redis_key(token), + socket_record_data, + ex=self.token_expiration, + nx=True, + ): + token = new_token = _get_new_token() except Exception as e: - console.error(f"Redis error checking token existence: {e}") - return await super().link_token_to_sid(token, sid) - - new_token = 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) + console.error(f"Redis error claiming token: {e}") + fallback_token = await super().link_token_to_sid(token, sid) + return fallback_token or new_token # Store in local dicts - socket_record = self.token_to_socket[token] = SocketRecord( - instance_id=self.instance_id, sid=sid - ) + self.token_to_socket[token] = socket_record self.sid_to_token[sid] = token - # Store in Redis if possible - try: - await self.redis.set( - redis_key, - pickle.dumps(socket_record), - ex=self.token_expiration, - ) - except Exception as e: - console.error(f"Redis error storing token: {e}") # Return the new token if one was generated return new_token diff --git a/tests/units/utils/test_token_manager.py b/tests/units/utils/test_token_manager.py index 208387cd86b..956f042cb65 100644 --- a/tests/units/utils/test_token_manager.py +++ b/tests/units/utils/test_token_manager.py @@ -291,18 +291,16 @@ async def test_link_token_to_sid_normal_case(self, manager, mock_redis): mock_redis: Mock Redis client fixture. """ token, sid = "token1", "sid1" - mock_redis.exists.return_value = False + mock_redis.set.return_value = True result = await manager.link_token_to_sid(token, sid) assert result is None - mock_redis.exists.assert_called_once_with( - f"token_manager_socket_record_{token}" - ) mock_redis.set.assert_called_once_with( f"token_manager_socket_record_{token}", pickle.dumps(SocketRecord(instance_id=manager.instance_id, sid=sid)), ex=3600, + nx=True, ) assert manager.token_to_socket[token].sid == sid assert manager.sid_to_token[sid] == token @@ -324,7 +322,6 @@ async def test_link_token_to_sid_reconnection_skips_redis( result = await manager.link_token_to_sid(token, sid) assert result is None - mock_redis.exists.assert_not_called() mock_redis.set.assert_not_called() async def test_link_token_to_sid_duplicate_detected(self, manager, mock_redis): @@ -335,7 +332,7 @@ async def test_link_token_to_sid_duplicate_detected(self, manager, mock_redis): mock_redis: Mock Redis client fixture. """ token, sid = "token1", "sid1" - mock_redis.exists.return_value = True + mock_redis.set.side_effect = [False, True] result = await manager.link_token_to_sid(token, sid) @@ -343,13 +340,12 @@ async def test_link_token_to_sid_duplicate_detected(self, manager, mock_redis): assert result != token assert len(result) == 36 # UUID4 length - mock_redis.exists.assert_called_once_with( - f"token_manager_socket_record_{token}" - ) - mock_redis.set.assert_called_once_with( + assert mock_redis.set.await_count == 2 + mock_redis.set.assert_awaited_with( f"token_manager_socket_record_{result}", pickle.dumps(SocketRecord(instance_id=manager.instance_id, sid=sid)), ex=3600, + nx=True, ) assert manager.token_to_sid[result] == sid assert manager.sid_to_token[sid] == result @@ -362,7 +358,7 @@ async def test_link_token_to_sid_redis_error_fallback(self, manager, mock_redis) mock_redis: Mock Redis client fixture. """ token, sid = "token1", "sid1" - mock_redis.exists.side_effect = Exception("Redis connection error") + mock_redis.set.side_effect = Exception("Redis connection error") with patch.object( LocalTokenManager, "link_token_to_sid", new_callable=AsyncMock @@ -384,7 +380,6 @@ async def test_link_token_to_sid_redis_set_error_continues( mock_redis: Mock Redis client fixture. """ token, sid = "token1", "sid1" - mock_redis.exists.return_value = False mock_redis.set.side_effect = Exception("Redis set error") result = await manager.link_token_to_sid(token, sid) @@ -393,6 +388,50 @@ async def test_link_token_to_sid_redis_set_error_continues( assert manager.token_to_sid[token] == sid assert manager.sid_to_token[sid] == token + async def test_link_token_to_sid_claims_token_atomically(self, mock_redis): + """Test concurrent managers cannot both claim the same token. + + Args: + mock_redis: Mock Redis client fixture. + """ + redis_values = {} + exists_calls = 0 + both_checked = asyncio.Event() + + async def exists(key): + nonlocal exists_calls + result = key in redis_values + exists_calls += 1 + if exists_calls == 2: + both_checked.set() + await both_checked.wait() + return result + + def set_value(key, value, *, ex, nx=False): + if nx and key in redis_values: + return False + redis_values[key] = value + return True + + mock_redis.exists.side_effect = exists + mock_redis.set.side_effect = set_value + with patch("reflex_base.config.get_config") as mock_get_config: + mock_get_config.return_value.redis_token_expiration = 3600 + managers = [RedisTokenManager(mock_redis), RedisTokenManager(mock_redis)] + + with patch.object(RedisTokenManager, "_ensure_socket_record_task"): + results = await asyncio.gather( + managers[0].link_token_to_sid("token1", "sid1"), + managers[1].link_token_to_sid("token1", "sid2"), + ) + + assert sum(result is None for result in results) == 1 + claimed_tokens = { + manager.sid_to_token[sid] + for manager, sid in zip(managers, ("sid1", "sid2"), strict=True) + } + assert len(claimed_tokens) == 2 + async def test_disconnect_token_owned_locally(self, manager, mock_redis): """Test disconnect cleans up both Redis and local mappings when owned locally. @@ -465,7 +504,7 @@ async def test_various_redis_errors_handled_gracefully( redis_error: Exception to test error handling. """ token, sid = "token1", "sid1" - mock_redis.exists.side_effect = redis_error + mock_redis.set.side_effect = redis_error with patch.object( LocalTokenManager, "link_token_to_sid", new_callable=AsyncMock From 3e0cf0f354c2e93a481a7bf4359d1e623655650d Mon Sep 17 00:00:00 2001 From: Anas Khan <83116240+anxkhn@users.noreply.github.com> Date: Fri, 17 Jul 2026 21:00:15 +0530 Subject: [PATCH 2/2] fix(tokens): address atomic claim review Signed-off-by: Anas Khan <83116240+anxkhn@users.noreply.github.com> --- news/{6740.bugfix.md => 6771.bugfix.md} | 0 reflex/utils/token_manager.py | 18 +++++++++++------ tests/units/utils/test_token_manager.py | 27 ++++++++++++++----------- 3 files changed, 27 insertions(+), 18 deletions(-) rename news/{6740.bugfix.md => 6771.bugfix.md} (100%) diff --git a/news/6740.bugfix.md b/news/6771.bugfix.md similarity index 100% rename from news/6740.bugfix.md rename to news/6771.bugfix.md diff --git a/reflex/utils/token_manager.py b/reflex/utils/token_manager.py index 256438ffc0f..1cff6cf83fd 100644 --- a/reflex/utils/token_manager.py +++ b/reflex/utils/token_manager.py @@ -314,13 +314,19 @@ async def link_token_to_sid(self, token: str, sid: str) -> str | None: new_token = None try: - while not await self.redis.set( - self._get_redis_key(token), - socket_record_data, - ex=self.token_expiration, - nx=True, - ): + for _ in range(3): + if await self.redis.set( + self._get_redis_key(token), + socket_record_data, + ex=self.token_expiration, + nx=True, + ): + break token = new_token = _get_new_token() + else: + console.error("Redis failed to claim a unique token after 3 attempts") + fallback_token = await super().link_token_to_sid(token, sid) + return fallback_token or new_token except Exception as e: console.error(f"Redis error claiming token: {e}") fallback_token = await super().link_token_to_sid(token, sid) diff --git a/tests/units/utils/test_token_manager.py b/tests/units/utils/test_token_manager.py index 956f042cb65..e12f5a70a99 100644 --- a/tests/units/utils/test_token_manager.py +++ b/tests/units/utils/test_token_manager.py @@ -395,17 +395,6 @@ async def test_link_token_to_sid_claims_token_atomically(self, mock_redis): mock_redis: Mock Redis client fixture. """ redis_values = {} - exists_calls = 0 - both_checked = asyncio.Event() - - async def exists(key): - nonlocal exists_calls - result = key in redis_values - exists_calls += 1 - if exists_calls == 2: - both_checked.set() - await both_checked.wait() - return result def set_value(key, value, *, ex, nx=False): if nx and key in redis_values: @@ -413,7 +402,6 @@ def set_value(key, value, *, ex, nx=False): redis_values[key] = value return True - mock_redis.exists.side_effect = exists mock_redis.set.side_effect = set_value with patch("reflex_base.config.get_config") as mock_get_config: mock_get_config.return_value.redis_token_expiration = 3600 @@ -432,6 +420,21 @@ def set_value(key, value, *, ex, nx=False): } assert len(claimed_tokens) == 2 + async def test_link_token_to_sid_limits_claim_attempts(self, manager, mock_redis): + """Test persistent failed claims fall back to local token storage. + + Args: + manager: RedisTokenManager fixture instance. + mock_redis: Mock Redis client fixture. + """ + mock_redis.set.return_value = False + + result = await manager.link_token_to_sid("token1", "sid1") + + assert mock_redis.set.await_count == 3 + assert result is not None + assert manager.sid_to_token["sid1"] == result + async def test_disconnect_token_owned_locally(self, manager, mock_redis): """Test disconnect cleans up both Redis and local mappings when owned locally.