Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions news/6934.bugfix.md
Original file line number Diff line number Diff line change
@@ -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.
15 changes: 9 additions & 6 deletions reflex/istate/shared.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Comment thread
greptile-apps[bot] marked this conversation as resolved.
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)
Expand Down
82 changes: 74 additions & 8 deletions reflex/utils/token_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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
Comment thread
greptile-apps[bot] marked this conversation as resolved.
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,
Expand Down
111 changes: 111 additions & 0 deletions tests/units/istate/test_shared.py
Original file line number Diff line number Diff line change
@@ -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]
49 changes: 49 additions & 0 deletions tests/units/utils/test_token_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Comment thread
cubic-dev-ai[bot] marked this conversation as resolved.
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.

Expand Down
Loading