diff --git a/news/+forked-worker-token-identity.bugfix.md b/news/+forked-worker-token-identity.bugfix.md new file mode 100644 index 00000000000..37ad6502ac1 --- /dev/null +++ b/news/+forked-worker-token-identity.bugfix.md @@ -0,0 +1 @@ +Give forked backend workers distinct socket-owner identities so Redis can deliver backend-initiated state updates to clients connected to another worker. diff --git a/reflex/app.py b/reflex/app.py index 866f1a01e9f..f6a69692fd1 100644 --- a/reflex/app.py +++ b/reflex/app.py @@ -718,13 +718,22 @@ async def context_middleware(scope: Scope, receive: Receive, send: Send): @contextlib.asynccontextmanager async def _setup_event_processor(self) -> AsyncIterator[None]: + """Configure event processing with a fresh worker socket identity. + + Yields: + None while the event processor is active. + """ + # The app may have been imported before the server forked its workers. + event_namespace = self.event_namespace + if event_namespace is not None: + event_namespace._token_manager._reset_instance_id() # Create the event processor. self._event_processor = BaseStateEventProcessor( middleware=self, backend_exception_handler=self.backend_exception_handler ) async with self._event_processor.configure( state_manager=self.state_manager, - event_namespace=self.event_namespace, + event_namespace=event_namespace, ): yield diff --git a/reflex/utils/token_manager.py b/reflex/utils/token_manager.py index cdafbb0b12d..afb1685a5ad 100644 --- a/reflex/utils/token_manager.py +++ b/reflex/utils/token_manager.py @@ -55,12 +55,16 @@ class TokenManager(ABC): def __init__(self): """Initialize the token manager with local dictionaries.""" # Each process has an instance_id to identify its own sockets. - self.instance_id: str = _get_new_token() + self._reset_instance_id() # Keep a mapping between client token and socket ID. self.token_to_socket: dict[str, SocketRecord] = {} # Keep a mapping between socket ID and client token. self.sid_to_token: dict[str, str] = {} + def _reset_instance_id(self) -> None: + """Assign a fresh socket-owner identity when a server worker starts.""" + self.instance_id: str = _get_new_token() + @property def token_to_sid(self) -> MappingProxyType[str, str]: """Read-only compatibility property for token_to_socket mapping. diff --git a/tests/units/test_app.py b/tests/units/test_app.py index c78fec088e1..a341b76def7 100644 --- a/tests/units/test_app.py +++ b/tests/units/test_app.py @@ -7,6 +7,8 @@ import io import json import logging +import multiprocessing +import pickle import re import unittest.mock import uuid @@ -72,9 +74,16 @@ from reflex.istate.manager.token import BaseStateToken from reflex.istate.storage import Cookie, LocalStorage, SessionStorage from reflex.model import Model -from reflex.state import BaseState, OnLoadInternalState, State, reload_state_module +from reflex.state import ( + BaseState, + OnLoadInternalState, + State, + StateUpdate, + reload_state_module, +) from reflex.utils import build from reflex.utils import exec as exec_utils +from reflex.utils.token_manager import RedisTokenManager, SocketRecord from .conftest import active_tracer, chdir, metric_points from .states import GenState @@ -86,6 +95,8 @@ ) if TYPE_CHECKING: + from multiprocessing.connection import Connection + from sqlalchemy.engine.base import Engine @@ -3180,6 +3191,85 @@ def test_call_app(): assert isinstance(api, Starlette) +def _probe_worker_token_identity(app: App, redis: AsyncMock, sender: Connection): + """Run the server startup path in a forked worker and report delta routing. + + Args: + app: The application created before the fork. + redis: The mock Redis connection carrying another worker's socket record. + sender: The pipe used to report the worker's result. + """ + + async def probe(): + """Start event processing and publish a delta to the socket owner.""" + async with app._setup_event_processor(): + assert app.event_namespace is not None + manager = app.event_namespace._token_manager + assert isinstance(manager, RedisTokenManager) + published = await manager.emit_lost_and_found( + "client", StateUpdate(delta={"state": {"count": 1}}) + ) + sender.send(( + manager.instance_id, + published, + redis.publish.call_args.args if published else None, + )) + + try: + asyncio.run(probe()) + finally: + sender.close() + + +@pytest.mark.skipif( + "fork" not in multiprocessing.get_all_start_methods(), + reason="Requires a server that forks after importing the app", +) +def test_forked_workers_publish_deltas_to_the_socket_owner(): + """Give forked workers distinct identities before routing backend deltas.""" + app = App() + app._state_manager = StateManagerMemory() + redis = AsyncMock() + manager = RedisTokenManager(redis) + assert app.event_namespace is not None + app.event_namespace._token_manager = manager + owner_id = manager.instance_id + redis.get.return_value = pickle.dumps( + SocketRecord(instance_id=owner_id, sid="remote") + ) + workers = multiprocessing.get_context("fork") + worker_ids = set() + + for _ in range(2): + receiver, sender = workers.Pipe(duplex=False) + process = workers.Process( + target=_probe_worker_token_identity, args=(app, redis, sender) + ) + process.start() + sender.close() + try: + assert receiver.poll(10), "The forked worker did not report a result" + worker_id, published, publish_args = receiver.recv() + process.join(timeout=10) + assert process.exitcode == 0 + assert published, "A live socket owned by another worker was discarded" + assert worker_id != owner_id + assert worker_id not in worker_ids + worker_ids.add(worker_id) + assert publish_args[0] == f"channel:token_manager_lost_and_found_{owner_id}" + record = pickle.loads(publish_args[1]) + assert record.token == "client" + assert record.update.delta == {"state": {"count": 1}} + finally: + receiver.close() + if process.is_alive(): + process.terminate() + process.join(timeout=10) + process.close() + + assert manager.instance_id == owner_id + + @pytest.fixture def upload_enabled(monkeypatch): """Fixture that enables Upload and cleans up afterward."""