Skip to content
Merged
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/+forked-worker-token-identity.bugfix.md
Original file line number Diff line number Diff line change
@@ -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.
11 changes: 10 additions & 1 deletion reflex/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
6 changes: 5 additions & 1 deletion reflex/utils/token_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
92 changes: 91 additions & 1 deletion tests/units/test_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@
import io
import json
import logging
import multiprocessing
import pickle
import re
import unittest.mock
import uuid
Expand Down Expand Up @@ -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
Expand All @@ -86,6 +95,8 @@
)

if TYPE_CHECKING:
from multiprocessing.connection import Connection

from sqlalchemy.engine.base import Engine


Expand Down Expand Up @@ -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."""
Expand Down
Loading