From ef8865150607540bfc1d565d9c823d9ff0723e63 Mon Sep 17 00:00:00 2001 From: Benjamin Hindman Date: Wed, 30 Sep 2026 00:21:20 +0000 Subject: [PATCH 1/2] React: add a websocket for the mutations of all states A browser sends the mutations of a state over a websocket that has the state in its path, which is what routes it to the server that is authoritative for the state. That takes a websocket for every state that is being mutated, and browsers limit how many websockets can be open, to around 200 (255 in Chrome). This adds a websocket without a state in its path, that can be used for the mutations of any state: every `MutateRequest` says what state it is a mutation of, and so does every `MutateResponse`. It is routed to any server, which is fine because performing a mutation that was sent over a websocket has always meant calling the server that is authoritative for the state. The mutations of a state are performed, and responded to, in the order that they were sent, just like when they have a websocket of their own. The mutations of different states are performed concurrently, so that a mutation only ever waits for mutations of the same state. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01WceTY5nXYmn6txE4haF8tw --- rbt/v1alpha1/react.proto | 14 + reboot/aio/react.py | 163 +++++++- tests/reboot/aio/BUILD.bazel | 19 + .../reboot/aio/react_websocket_mutate_test.py | 391 ++++++++++++++++++ 4 files changed, 581 insertions(+), 6 deletions(-) create mode 100644 tests/reboot/aio/react_websocket_mutate_test.py diff --git a/rbt/v1alpha1/react.proto b/rbt/v1alpha1/react.proto index 5eb6b1a09..26532d1da 100644 --- a/rbt/v1alpha1/react.proto +++ b/rbt/v1alpha1/react.proto @@ -63,9 +63,23 @@ message MutateRequest { // Authorization bearer token. optional string bearer_token = 4; + + // The state to mutate, as found in the path of a request, e.g., + // `my.package.MyState:my-id`. + // + // Only expected when sent over the websocket for the mutations of + // all states, since a websocket for the mutations of a single state + // has the state in its path. + optional string state_ref = 5; } message MutateResponse { + // The `state_ref` of the request that this is the response to, if + // it had one. Mutations of the same state are performed, and + // responded to, in the order that they were sent, but mutations of + // different states are performed concurrently. + optional string state_ref = 3; + oneof response_or_status { // Serialized response from the method specified in the request. bytes response = 1; diff --git a/reboot/aio/react.py b/reboot/aio/react.py index 6264ebf30..762fb3969 100644 --- a/reboot/aio/react.py +++ b/reboot/aio/react.py @@ -3,8 +3,10 @@ import logging import reboot.aio.placement import traceback +import urllib.parse import uuid import websockets +from collections import deque from google.protobuf.json_format import MessageToJson from google.rpc import code_pb2, status_pb2 from grpc_health.v1 import health_pb2 @@ -30,6 +32,10 @@ logger = get_logger(__name__) +# Path of the websocket for the mutations of all states, rather than +# of the single state that is in the path. +MUTATE_WEBSOCKET_PATH = '/__/reboot/websocket/mutate' + class _SuppressInvalidHandshakeFilter(logging.Filter): """Drop spurious `opening handshake failed` logs from non-WebSocket @@ -155,15 +161,21 @@ async def _serve(self, websocket): with use_application_id(self._application_id): try: - application_id, state_ref = ( - websocket.request.headers[APPLICATION_ID_HEADER], - StateRef.from_maybe_readable( - websocket.request.headers[STATE_REF_HEADER] - ), - ) + application_id = websocket.request.headers[ + APPLICATION_ID_HEADER] assert self._application_id == application_id + if websocket.request.path == MUTATE_WEBSOCKET_PATH: + return await self._websocket_mutate_states( + websocket, + application_id=application_id, + ) + + state_ref = StateRef.from_maybe_readable( + websocket.request.headers[STATE_REF_HEADER] + ) + state_type_name = self._state_type_name_for_state_ref( state_ref ) @@ -303,6 +315,145 @@ async def _websocket_mutate( ).SerializeToString() ) + async def _mutate_state( + self, + request: react_pb2.MutateRequest, + *, + application_id: ApplicationId, + ) -> react_pb2.MutateResponse: + """Performs the mutation in `request`, of the state in `request`, + and returns its response, which is a status if it failed.""" + try: + # The state is what would otherwise be in the path, see + # `mangled_http_path.lua`. + state_ref = StateRef.from_maybe_readable( + urllib.parse.unquote(request.state_ref) + ) + + state_type_name = self._state_type_name_for_state_ref(state_ref) + + if state_type_name is None: + log_at_most_once_per( + seconds=60, + log_method=logger.error, + message=_unknown_query_or_mutation_error_message( + is_query=False, + state_type=state_ref.state_type, + ), + ) + raise SystemAborted(UnknownService()) + + middleware = self._middleware_by_state_type[state_type_name] + + # NOTE: we might not be the server that is authoritative + # for this state, which is fine because `react_mutate()` + # calls the server that is. + response = await middleware.react_mutate( + Headers( + application_id=application_id, + state_ref=state_ref, + idempotency_key=uuid.UUID(request.idempotency_key), + bearer_token=request.bearer_token, + # This request came in over a websocket, so this is + # a frontend client calling, not another Reboot + # application. Therefore there is no caller ID. + caller_id=None, + ), + request.method, + request.request, + ) + + return react_pb2.MutateResponse( + response=response.SerializeToString(), + ) + except asyncio.CancelledError: + raise + except Aborted as aborted: + return react_pb2.MutateResponse( + status=MessageToJson(aborted.to_status()), + ) + except BaseException as exception: + # See the comment in `_serve()` for why we log the stack + # trace but don't send it. + error_message = ( + 'Failed to execute mutation via websocket; ' + f'{type(exception).__name__}: {exception}' + ) + + if should_print_stacktrace(): + error_message += f'\n{traceback.format_exc()}' + + logger.error(error_message) + + return react_pb2.MutateResponse( + status=MessageToJson( + status_pb2.Status( + code=code_pb2.Code.UNKNOWN, + message=f'{type(exception).__name__}: {exception}', + ) + ), + ) + + async def _websocket_mutate_states( + self, + websocket, + *, + application_id: ApplicationId, + ): + """Performs the mutations of any state that are sent over + `websocket`: those of the same state in the order that they + were sent, and those of different states concurrently, so that + a mutation only ever waits for mutations of the same state.""" + # The mutations of each state that we have yet to respond to, + # by `state_ref`. A state is in here for as long as it has a + # task performing its mutations. + requests_by_state_ref: dict[str, deque[react_pb2.MutateRequest]] = {} + + tasks: set[asyncio.Task] = set() + + async def mutate_state(state_ref: str): + requests = requests_by_state_ref[state_ref] + while len(requests) > 0: + response = await self._mutate_state( + requests[0], + application_id=application_id, + ) + response.state_ref = state_ref + await websocket.send(response.SerializeToString()) + requests.popleft() + del requests_by_state_ref[state_ref] + + def done(task: asyncio.Task): + tasks.discard(task) + if not task.cancelled(): + # Retrieve the exception, if any, which is what we + # expect when the websocket gets closed. + task.exception() + + try: + async for request_bytes in websocket: + request = react_pb2.MutateRequest() + request.ParseFromString(request_bytes) + + requests = requests_by_state_ref.get(request.state_ref) + + if requests is not None: + requests.append(request) + continue + + requests_by_state_ref[request.state_ref] = deque([request]) + + task = asyncio.create_task( + mutate_state(request.state_ref), + name=f'mutate_state({request.state_ref}) in {__name__}', + ) + tasks.add(task) + task.add_done_callback(done) + finally: + # Just like a websocket for the mutations of a single + # state, once the websocket is closed we are done. + await wait_for_tasks(list(tasks), cancel=True) + async def _websocket_health_check(self, websocket): """ Handle health check requests received via websocket. diff --git a/tests/reboot/aio/BUILD.bazel b/tests/reboot/aio/BUILD.bazel index 069fd4dec..70b356125 100644 --- a/tests/reboot/aio/BUILD.bazel +++ b/tests/reboot/aio/BUILD.bazel @@ -1,3 +1,4 @@ +load("@rbt_pypi//:requirements.bzl", "requirement") load("@rules_python//python:defs.bzl", "py_test") py_test( @@ -98,3 +99,21 @@ py_test( "//reboot/aio:types_py", ], ) + +py_test( + name = "react_websocket_mutate_test_py", + srcs = [":react_websocket_mutate_test.py"], + main = "react_websocket_mutate_test.py", + deps = [ + requirement("protobuf"), + requirement("websockets"), + "//rbt/v1alpha1:react_py_proto", + "//reboot/aio:applications_py", + "//reboot/aio:external_py", + "//reboot/aio:headers_py", + "//reboot/aio:react_py", + "//reboot/aio:tests_py", + "//reboot/aio:types_py", + "//tests/reboot:greeter_servicers_py", + ], +) diff --git a/tests/reboot/aio/react_websocket_mutate_test.py b/tests/reboot/aio/react_websocket_mutate_test.py new file mode 100644 index 000000000..33d2d3d2b --- /dev/null +++ b/tests/reboot/aio/react_websocket_mutate_test.py @@ -0,0 +1,391 @@ +import asyncio +import unittest +import uuid +import websockets +from google.protobuf.message import Message +from rbt.v1alpha1 import react_pb2 +from reboot.aio.applications import Application +from reboot.aio.external import ExternalContext +from reboot.aio.headers import Headers +from reboot.aio.react import MUTATE_WEBSOCKET_PATH, ReactServicer +from reboot.aio.tests import Reboot +from reboot.aio.types import ApplicationId, StateTypeName +from tests.reboot.greeter_rbt import ( + Greeter, + SetAdjectiveRequest, + SetAdjectiveResponse, +) +from tests.reboot.greeter_servicers import MyGreeterServicer +from typing import Optional + +STATE_TYPE_NAME = StateTypeName('tests.reboot.Greeter') + +# More states than the servers that we have by default, so that some +# of them must be on a different server than the websocket is. +STATE_IDS = [f'greeter-{i}' for i in range(8)] + +# How long we wait before we believe that we are not getting a +# response, rather than just being slow. +WAITING_SECONDS = 2 + + +def state_ref(state_id: str) -> str: + """Returns what a browser uses for the state, see `stateIdToRef()`.""" + return f'{STATE_TYPE_NAME}:{state_id}' + + +def mutate_request( + *, + state_ref: str, + adjective: str, + method: str = 'SetAdjective', +) -> bytes: + return react_pb2.MutateRequest( + method=method, + request=SetAdjectiveRequest(adjective=adjective).SerializeToString(), + idempotency_key=str(uuid.uuid4()), + state_ref=state_ref, + ).SerializeToString() + + +def mutate_response(response_bytes: bytes) -> react_pb2.MutateResponse: + response = react_pb2.MutateResponse() + response.ParseFromString(response_bytes) + return response + + +class WebSocketMutateTestCase(unittest.IsolatedAsyncioTestCase): + """Tests the websocket for the mutations of all states the way + that a browser uses it, i.e., through Envoy.""" + + async def asyncSetUp(self) -> None: + self.rbt = Reboot() + await self.rbt.start() + + await self.rbt.up( + Application(servicers=[MyGreeterServicer]), + local_envoy=True, + ) + + self.context: ExternalContext = self.rbt.create_external_context( + name=self.id() + ) + + for state_id in STATE_IDS + ['google-oauth2|123']: + await Greeter.Create( + self.context, + state_id, + title='Dr', + name='Jonathan', + adjective='initial', + ) + + self.websocket = await websockets.connect( + f'ws://localhost:{self.rbt.envoy_port()}{MUTATE_WEBSOCKET_PATH}' + ) + + async def asyncTearDown(self) -> None: + await self.websocket.close() + await self.rbt.stop() + + async def receive(self) -> react_pb2.MutateResponse: + return mutate_response( + await asyncio.wait_for( + self.websocket.recv(), + timeout=30, + ) + ) + + async def adjective(self, state_id: str) -> str: + state = await Greeter.ref(state_id).GetWholeState(self.context) + return state.adjective + + async def test_states(self) -> None: + """Tests mutating more than one state, wherever they are.""" + for state_id in STATE_IDS: + await self.websocket.send( + mutate_request( + state_ref=state_ref(state_id), + adjective=f'adjective of {state_id}', + ) + ) + + responses = [await self.receive() for _ in STATE_IDS] + + for response in responses: + self.assertEqual( + 'response', + response.WhichOneof('response_or_status'), + response, + ) + + self.assertCountEqual( + [state_ref(state_id) for state_id in STATE_IDS], + [response.state_ref for response in responses], + ) + + for state_id in STATE_IDS: + self.assertEqual( + f'adjective of {state_id}', + await self.adjective(state_id), + ) + + async def test_order(self) -> None: + """Tests that the mutations of a state are performed in the + order that they were sent.""" + adjectives = [f'adjective {i}' for i in range(10)] + + for adjective in adjectives: + for state_id in STATE_IDS: + await self.websocket.send( + mutate_request( + state_ref=state_ref(state_id), + adjective=adjective, + ) + ) + + for _ in range(len(adjectives) * len(STATE_IDS)): + response = await self.receive() + self.assertEqual( + 'response', + response.WhichOneof('response_or_status'), + response, + ) + + for state_id in STATE_IDS: + self.assertEqual(adjectives[-1], await self.adjective(state_id)) + + async def test_state_id_that_needs_encoding(self) -> None: + """Tests a state whose ID is percent-encoded by a browser.""" + await self.websocket.send( + mutate_request( + state_ref=state_ref('google-oauth2%7C123'), + adjective='encoded', + ) + ) + + response = await self.receive() + + self.assertEqual( + 'response', + response.WhichOneof('response_or_status'), + response, + ) + + # The response has what the request had, since that is what a + # browser is looking for. + self.assertEqual(state_ref('google-oauth2%7C123'), response.state_ref) + + self.assertEqual('encoded', await self.adjective('google-oauth2|123')) + + async def test_status(self) -> None: + """Tests that a mutation that fails has a status for a response, + and that the websocket can still be used.""" + await self.websocket.send( + mutate_request( + state_ref='tests.reboot.Unknown:unknown', + adjective='unknown', + ) + ) + + response = await self.receive() + + self.assertEqual( + 'status', + response.WhichOneof('response_or_status'), + response, + ) + self.assertEqual('tests.reboot.Unknown:unknown', response.state_ref) + + await self.websocket.send( + mutate_request( + state_ref=state_ref(STATE_IDS[0]), + adjective='unknown', + method='TestLongRunningWriter', + ) + ) + + response = await self.receive() + + self.assertEqual( + 'status', + response.WhichOneof('response_or_status'), + response, + ) + self.assertEqual(state_ref(STATE_IDS[0]), response.state_ref) + + await self.websocket.send( + mutate_request( + state_ref=state_ref(STATE_IDS[0]), + adjective='friendly', + ) + ) + + response = await self.receive() + + self.assertEqual( + 'response', + response.WhichOneof('response_or_status'), + response, + ) + + self.assertEqual('friendly', await self.adjective(STATE_IDS[0])) + + +class FakeMiddleware: + """A `Middleware` whose mutations are performed once the test says + so.""" + + def __init__(self) -> None: + # The adjectives of the mutations that have been performed, + # in the order that they were. + self.performed: list[str] = [] + + self._events: dict[str, asyncio.Event] = {} + + def event(self, adjective: str) -> asyncio.Event: + return self._events.setdefault(adjective, asyncio.Event()) + + async def react_mutate( + self, + headers: Headers, + method: str, + request_bytes: bytes, + ) -> Message: + request = SetAdjectiveRequest() + request.ParseFromString(request_bytes) + await self.event(request.adjective).wait() + self.performed.append(request.adjective) + return SetAdjectiveResponse() + + +class FakeWebSocket: + """A websocket that receives what the test says it does.""" + + def __init__(self) -> None: + self.received: asyncio.Queue[Optional[bytes]] = asyncio.Queue() + self.sent: asyncio.Queue[bytes] = asyncio.Queue() + + def __aiter__(self): + return self + + async def __anext__(self) -> bytes: + request_bytes = await self.received.get() + if request_bytes is None: + raise StopAsyncIteration + return request_bytes + + async def send(self, response_bytes: bytes) -> None: + self.sent.put_nowait(response_bytes) + + +class ConcurrentlyTestCase(unittest.IsolatedAsyncioTestCase): + """Tests which mutations wait for which, which needs mutations + that take as long as the test wants them to.""" + + async def asyncSetUp(self) -> None: + self.middleware = FakeMiddleware() + + self.websocket = FakeWebSocket() + + self.task = asyncio.create_task( + ReactServicer( + ApplicationId('application'), + { + STATE_TYPE_NAME: + self.middleware, # type: ignore[dict-item] + }, + )._websocket_mutate_states( + self.websocket, + application_id=ApplicationId('application'), + ) + ) + + async def asyncTearDown(self) -> None: + self.task.cancel() + await asyncio.wait([self.task]) + + def receive(self, *, state_id: str, adjective: str) -> None: + self.websocket.received.put_nowait( + mutate_request( + state_ref=state_ref(state_id), + adjective=adjective, + ) + ) + + async def sent(self) -> str: + """Returns the state of the next response.""" + response = mutate_response( + await asyncio.wait_for( + self.websocket.sent.get(), + timeout=WAITING_SECONDS, + ) + ) + return response.state_ref + + async def test_different_states(self) -> None: + """Tests that a mutation does not wait for a mutation of a + different state.""" + self.receive(state_id='slow', adjective='slow') + self.receive(state_id='fast', adjective='fast') + + self.middleware.event('fast').set() + + self.assertEqual(state_ref('fast'), await self.sent()) + + self.assertEqual(['fast'], self.middleware.performed) + + self.middleware.event('slow').set() + + self.assertEqual(state_ref('slow'), await self.sent()) + + self.assertEqual(['fast', 'slow'], self.middleware.performed) + + async def test_same_state(self) -> None: + """Tests that a mutation waits for the mutations of the same + state that were sent before it.""" + self.receive(state_id='state', adjective='first') + self.receive(state_id='state', adjective='second') + self.receive(state_id='state', adjective='third') + + # Even though the mutations after it could be performed. + self.middleware.event('third').set() + self.middleware.event('second').set() + + with self.assertRaises(asyncio.TimeoutError): + await self.sent() + + self.assertEqual([], self.middleware.performed) + + self.middleware.event('first').set() + + for _ in range(3): + self.assertEqual(state_ref('state'), await self.sent()) + + self.assertEqual( + ['first', 'second', 'third'], + self.middleware.performed, + ) + + async def test_closed(self) -> None: + """Tests that mutations are cancelled once the websocket is + closed, just like for a websocket for a single state.""" + self.receive(state_id='state', adjective='never') + + # Wait for the mutation to be waiting. + while 'never' not in self.middleware._events: + await asyncio.sleep(0.01) + + self.websocket.received.put_nowait(None) + + await asyncio.wait_for(self.task, timeout=WAITING_SECONDS) + + self.middleware.event('never').set() + + await asyncio.sleep(0.1) + + self.assertEqual([], self.middleware.performed) + + +if __name__ == '__main__': + unittest.main() From a256bca3a47084aed8ef3eeff6dbfd024837ef22 Mon Sep 17 00:00:00 2001 From: Benjamin Hindman Date: Wed, 30 Sep 2026 00:21:23 +0000 Subject: [PATCH 2/2] React: send the mutations of all states over one websocket Rather than a websocket for every state whose mutators are bound there is now one for every backend, that is opened once the first mutator is bound and closed once no state is used anymore. A state can not tell the difference: what it sends its mutations over looks like a websocket of its own, and because the backend responds to the mutations of a state in order it can still tell what mutation a response is for by counting. A backend from before this responds without saying what state its response is for, at which point we go back to a websocket for every state. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01WceTY5nXYmn6txE4haF8tw --- documentation/docs/rbt_cli.md | 2 +- reboot/templates/reboot_react.ts.j2 | 12 +- reboot/web/index.ts | 242 +++++++++++- tests/reboot/greeter_rbt_react.golden.js | 7 +- .../mutator_websocket.test.tsx | 360 +++++++++++++++++- 5 files changed, 600 insertions(+), 23 deletions(-) diff --git a/documentation/docs/rbt_cli.md b/documentation/docs/rbt_cli.md index af65e9e8f..5ff47491e 100644 --- a/documentation/docs/rbt_cli.md +++ b/documentation/docs/rbt_cli.md @@ -47,7 +47,7 @@ warning. Other calls share a handful of HTTP/1.1 connections per host, so many outstanding calls queue behind each other, which the client also warns about. Over HTTPS reactive readers and other calls are all streams of one HTTP/2 connection. A WebSocket is then only -used for the mutators of a state, and only once one of them is used. +used for mutations, a single one for the mutations of all states. To enable HTTPS you must provide your own TLS certificate when running `rbt dev run`. diff --git a/reboot/templates/reboot_react.ts.j2 b/reboot/templates/reboot_react.ts.j2 index c528f968a..46e19a0cb 100644 --- a/reboot/templates/reboot_react.ts.j2 +++ b/reboot/templates/reboot_react.ts.j2 @@ -638,7 +638,7 @@ class {{ client.proto.state_name | to_camel }}Instance { private runningMutates: reboot_react.Mutate[] = []; private queuedMutates: reboot_react.Mutate[] = []; private flushMutates?: reboot_api.Event = undefined; - private websocket?: WebSocket = undefined; + private websocket?: reboot_web.StateWebSocket = undefined; private wantsWebSocket = false; private backoff: reboot_api.Backoff = new reboot_api.Backoff(); @@ -683,12 +683,7 @@ class {{ client.proto.state_name | to_camel }}Instance { // caller while no default ID has resolved (e.g. signed out): it // opens no socket so there's nothing to connect to. if (this.websocket === undefined && this.refs > 0 && this.id !== "") { - const url = new URL(`${this.url}/__/reboot/rpc/${this.stateRef}`); - url.protocol = url.protocol === "https:" ? "wss:" : "ws:"; - - this.websocket = reboot_web.websockets.create(url); - - this.websocket.binaryType = "arraybuffer"; + this.websocket = reboot_web.websockets.mutate(this.url, this.stateRef); this.websocket.onopen = () => { if (this.websocket?.readyState === WebSocket.OPEN) { @@ -769,6 +764,9 @@ class {{ client.proto.state_name | to_camel }}Instance { ? partialRequest : new reboot_api.react_pb.MutateRequest(partialRequest); + // The websocket might be for the mutations of all states. + request.stateRef = this.stateRef; + return new Promise((resolve, _) => { if (this.loadingReaders === 0) { this.runningMutates = this.runningMutates.concat({ request, resolve, update }); diff --git a/reboot/web/index.ts b/reboot/web/index.ts index 50b0fbd54..08875b1c2 100644 --- a/reboot/web/index.ts +++ b/reboot/web/index.ts @@ -548,9 +548,11 @@ export class WebSockets { "calls (specifically reactive readers or mutations) will never " + "make it to your Reboot application (even though we keep retrying). " + (url.protocol === "wss:" - ? "When you use TLS a websocket is only used for the mutators " + - "of a state, so you are using the mutators of too many " + - "states at the same time." + ? "When you use TLS websockets are only used for mutations, " + + "and there is only a websocket for the mutators of every " + + "state if your backend is running a version of Reboot " + + "from before there was one for the mutations of all " + + "states. You can solve this by upgrading your backend." : "You can solve this for reactive readers by using HTTP/2 " + "which allows an unlimited number of concurrent streams. " + "Reboot uses HTTP/2 by default when you use TLS. You should " + @@ -568,6 +570,240 @@ export class WebSockets { return websocket; } + + // The `MutateWebSocket` for each URL that has one. + private mutateWebSockets = new Map(); + + // URLs that we have learned don't have a websocket for the + // mutations of all states, because they are running a version of + // Reboot from before it existed. + private urlsWithoutMutateWebSocket = new Set(); + + // Returns what to send the mutations of the state over, which the + // caller should `close()` once it has no more mutations to send. + mutate(url: string, stateRef: string): StateWebSocket { + if (this.urlsWithoutMutateWebSocket.has(url)) { + return this.mutateState(url, stateRef); + } + + let mutateWebSocket = this.mutateWebSockets.get(url); + + if (mutateWebSocket === undefined) { + mutateWebSocket = new MutateWebSocket(url, () => { + this.urlsWithoutMutateWebSocket.add(url); + }); + this.mutateWebSockets.set(url, mutateWebSocket); + } + + return mutateWebSocket.open(stateRef); + } + + // Returns a websocket for the mutations of only this state. + private mutateState(url: string, stateRef: string): StateWebSocket { + const websocketUrl = new URL(`${url}/__/reboot/rpc/${stateRef}`); + websocketUrl.protocol = websocketUrl.protocol === "https:" ? "wss:" : "ws:"; + + const websocket = this.create(websocketUrl); + + websocket.binaryType = "arraybuffer"; + + const stateWebSocket = new StateWebSocket({ + send: (data) => websocket.send(data), + close: () => websocket.close(), + }); + + websocket.onopen = () => stateWebSocket.opened(); + websocket.onerror = () => stateWebSocket.onerror?.(); + websocket.onclose = () => stateWebSocket.closed(); + websocket.onmessage = (event) => stateWebSocket.onmessage?.(event); + + return stateWebSocket; + } +} + +// What the mutations of a state are sent over, which as far as the +// state can tell is a websocket of its own. +// +// Browsers limit how many websockets can be open, so what it really +// is, if the backend has one, is a part of the one websocket that the +// mutations of all states are sent over. +export class StateWebSocket { + // NOTE: we can't use, e.g., `WebSocket.CONNECTING`, because there + // is no `WebSocket` when we are on the server, e.g., in Next.js. + static readonly CONNECTING = 0; + static readonly OPEN = 1; + static readonly CLOSED = 3; + + readyState: number = StateWebSocket.CONNECTING; + + onopen?: () => void; + onerror?: () => void; + onclose?: () => void; + onmessage?: (event: { data: ArrayBuffer }) => void; + + constructor( + private readonly websocket: { + send: (data: Uint8Array) => void; + close: () => void; + } + ) {} + + send(data: Uint8Array) { + this.websocket.send(data); + } + + close() { + this.websocket.close(); + } + + opened() { + if (this.readyState === StateWebSocket.CONNECTING) { + this.readyState = StateWebSocket.OPEN; + this.onopen?.(); + } + } + + closed() { + if (this.readyState !== StateWebSocket.CLOSED) { + this.readyState = StateWebSocket.CLOSED; + this.onclose?.(); + } + } +} + +// The websocket that the mutations of all states of a backend are +// sent over, which is open for as long as any state has mutations to +// send. +// +// Every mutation says what state it is a mutation of, and so does +// every response. The backend responds to the mutations of a state in +// the order that they were sent, so a state can tell what mutation a +// response is for just like it can when it has a websocket of its +// own. +class MutateWebSocket { + private websocket?: WebSocket; + + private stateWebSockets = new Map(); + + constructor( + private readonly url: string, + // Invoked if the backend turns out not to have a websocket for + // the mutations of all states. + private readonly unsupported: () => void + ) {} + + open(stateRef: string): StateWebSocket { + const stateWebSocket: StateWebSocket = new StateWebSocket({ + send: (data) => { + if (this.websocket?.readyState === StateWebSocket.OPEN) { + this.websocket.send(data); + } + }, + close: () => { + if (this.stateWebSockets.get(stateRef) === stateWebSocket) { + this.stateWebSockets.delete(stateRef); + } + + // Nobody has any mutations to send anymore. + if (this.stateWebSockets.size === 0 && this.websocket !== undefined) { + const websocket = this.websocket; + this.websocket = undefined; + websocket.close(); + } + + // A websocket closes asynchronously. + queueMicrotask(() => stateWebSocket.closed()); + }, + }); + + this.stateWebSockets.set(stateRef, stateWebSocket); + + if (this.websocket === undefined) { + this.connect(); + } else if (this.websocket.readyState === StateWebSocket.OPEN) { + // A websocket opens asynchronously, which is what gives the + // caller the chance to set `onopen`. + queueMicrotask(() => { + if ( + this.stateWebSockets.get(stateRef) === stateWebSocket && + this.websocket?.readyState === StateWebSocket.OPEN + ) { + stateWebSocket.opened(); + } + }); + } + + return stateWebSocket; + } + + private connect() { + const url = new URL(`${this.url}/__/reboot/websocket/mutate`); + url.protocol = url.protocol === "https:" ? "wss:" : "ws:"; + + const websocket = websockets.create(url); + + this.websocket = websocket; + + websocket.binaryType = "arraybuffer"; + + // NOTE: every one of these needs to check that `websocket` is + // still what we are using because we might have closed it, and + // even opened another one, by the time that they are invoked. + + websocket.onopen = () => { + if (this.websocket === websocket) { + for (const stateWebSocket of this.stateWebSockets.values()) { + stateWebSocket.opened(); + } + } + }; + + websocket.onmessage = (event) => { + if (this.websocket !== websocket) { + return; + } + + const { stateRef } = react_pb.MutateResponse.fromBinary( + new Uint8Array(event.data) + ); + + if (stateRef === undefined) { + // This is from a backend that does not know about the + // websocket for the mutations of all states, and thus assumed + // that this was one for the mutations of a single state. + // + // Every state will try to open a websocket again once this + // one is closed, which will then be one of its own. + this.unsupported(); + websocket.close(); + return; + } + + this.stateWebSockets.get(stateRef)?.onmessage?.(event); + }; + + websocket.onerror = () => { + if (this.websocket === websocket) { + for (const stateWebSocket of this.stateWebSockets.values()) { + stateWebSocket.onerror?.(); + } + } + }; + + websocket.onclose = () => { + if (this.websocket === websocket) { + this.websocket = undefined; + + const stateWebSockets = [...this.stateWebSockets.values()]; + + this.stateWebSockets.clear(); + + for (const stateWebSocket of stateWebSockets) { + stateWebSocket.closed(); + } + } + }; + } } export const websockets = new WebSockets(); diff --git a/tests/reboot/greeter_rbt_react.golden.js b/tests/reboot/greeter_rbt_react.golden.js index b77e3d50b..e8b8e22fc 100755 --- a/tests/reboot/greeter_rbt_react.golden.js +++ b/tests/reboot/greeter_rbt_react.golden.js @@ -1811,10 +1811,7 @@ class GreeterInstance { // caller while no default ID has resolved (e.g. signed out): it // opens no socket so there's nothing to connect to. if (this.websocket === undefined && this.refs > 0 && this.id !== "") { - const url = new URL(`${this.url}/__/reboot/rpc/${this.stateRef}`); - url.protocol = url.protocol === "https:" ? "wss:" : "ws:"; - this.websocket = reboot_web.websockets.create(url); - this.websocket.binaryType = "arraybuffer"; + this.websocket = reboot_web.websockets.mutate(this.url, this.stateRef); this.websocket.onopen = () => { var _a; if (((_a = this.websocket) === null || _a === void 0 ? void 0 : _a.readyState) === WebSocket.OPEN) { @@ -1872,6 +1869,8 @@ class GreeterInstance { const request = partialRequest instanceof reboot_api.react_pb.MutateRequest ? partialRequest : new reboot_api.react_pb.MutateRequest(partialRequest); + // The websocket might be for the mutations of all states. + request.stateRef = this.stateRef; return new Promise((resolve, _) => { var _a; if (this.loadingReaders === 0) { diff --git a/tests/reboot/react/test_mutator_websocket/mutator_websocket.test.tsx b/tests/reboot/react/test_mutator_websocket/mutator_websocket.test.tsx index 1468b7b66..691f507b6 100644 --- a/tests/reboot/react/test_mutator_websocket/mutator_websocket.test.tsx +++ b/tests/reboot/react/test_mutator_websocket/mutator_websocket.test.tsx @@ -15,6 +15,12 @@ import { UseGreeterApi, useGreeter } from "../../greeter_rbt_react.js"; // the only websockets are the ones for mutations. const URL = "https://reboot.test"; +// A backend from before there was a websocket for the mutations of +// all states. +const OLD_URL = "https://old.reboot.test"; + +const stateRef = (id: string) => `tests.reboot.Greeter:${id}`; + class FakeWebSocket { static readonly CONNECTING = 0; static readonly OPEN = 1; @@ -34,6 +40,9 @@ class FakeWebSocket { // Everything that was sent to the backend. sent: react_pb.MutateRequest[] = []; + // How many of those we have responded to. + private responded = 0; + readonly url: string; constructor(url: string | URL) { @@ -52,13 +61,44 @@ class FakeWebSocket { this.onopen?.(); } - respond() { - const bytes = new react_pb.MutateResponse({ - responseOrStatus: { - case: "response", - value: new SetAdjectiveResponse().toBinary(), - }, - }).toBinary(); + // Responds to the first mutation that we have not responded to + // yet, or to `request` if there is one. + respond(request?: react_pb.MutateRequest) { + if (request === undefined) { + request = this.sent[this.responded]; + this.responded += 1; + } + this.receive( + new react_pb.MutateResponse({ + stateRef: request.stateRef, + responseOrStatus: { + case: "response", + value: new SetAdjectiveResponse().toBinary(), + }, + }) + ); + } + + // Responds like a backend from before there was a websocket for the + // mutations of all states does, because it is missing the state + // that it expects to find in the path. + respondWithoutState() { + this.receive( + new react_pb.MutateResponse({ + responseOrStatus: { case: "status", value: "{}" }, + }) + ); + this.close(); + } + + fail() { + this.readyState = FakeWebSocket.CLOSED; + this.onerror?.(); + this.onclose?.(); + } + + private receive(response: react_pb.MutateResponse) { + const bytes = response.toBinary(); this.onmessage?.({ data: bytes.buffer.slice( bytes.byteOffset, @@ -92,19 +132,40 @@ const fakeFetch = (url: string) => { let use: (greeter: UseGreeterApi) => void = () => {}; // The `greeter` from the last render, so that the test can call -// mutators without having to go through the DOM. +// mutators without having to go through the DOM, and the one from the +// last render for each ID. let greeter: UseGreeterApi; +let greeters: { [id: string]: UseGreeterApi } = {}; const Greeter: React.FC<{ id: string }> = ({ id }) => { greeter = useGreeter({ id }); + greeters[id] = greeter; use(greeter); return
; }; +// What has been resolved, in the order that it was. +let resolved: string[] = []; + +const setAdjective = (id: string, adjective: string) => { + act(() => { + greeters[id].setAdjective({ adjective }).then(() => { + resolved = [...resolved, adjective]; + }); + }); +}; + +const adjectives = (websocket: FakeWebSocket) => + websocket.sent.map( + ({ request }) => SetAdjectiveRequest.fromBinary(request).adjective + ); + describe("The websocket for mutations", () => { beforeEach(() => { fetched = []; use = () => {}; + greeters = {}; + resolved = []; FakeWebSocket.instances = []; vi.stubGlobal("WebSocket", FakeWebSocket); vi.stubGlobal("fetch", fakeFetch); @@ -266,6 +327,289 @@ describe("The websocket for mutations", () => { expect(websocket.readyState).toBe(FakeWebSocket.CLOSED); }); + it("is for the mutations of all states", async () => { + use = ({ setAdjective }) => {}; + + render( + + + + + ); + + await waitFor(() => { + expect(FakeWebSocket.instances.length).toBe(1); + }); + + const [websocket] = FakeWebSocket.instances; + + expect(websocket.url).toBe("wss://reboot.test/__/reboot/websocket/mutate"); + + act(() => { + websocket.open(); + }); + + setAdjective("first", "first"); + setAdjective("second", "second"); + setAdjective("first", "first again"); + + await waitFor(() => { + expect(websocket.sent.length).toBe(3); + }); + + expect(adjectives(websocket)).toEqual(["first", "second", "first again"]); + + expect(websocket.sent.map(({ stateRef }) => stateRef)).toEqual([ + stateRef("first"), + stateRef("second"), + stateRef("first"), + ]); + + // A response is for the first mutation of its state that does + // not have one yet, no matter what other states have been up to. + act(() => { + websocket.respond(websocket.sent[1]); + }); + + await waitFor(() => { + expect(resolved).toEqual(["second"]); + }); + + act(() => { + websocket.respond(websocket.sent[0]); + }); + + await waitFor(() => { + expect(resolved).toEqual(["second", "first"]); + }); + + act(() => { + websocket.respond(websocket.sent[2]); + }); + + await waitFor(() => { + expect(resolved).toEqual(["second", "first", "first again"]); + }); + + expect(FakeWebSocket.instances.length).toBe(1); + }); + + it("is also used by a state that is used later", async () => { + use = ({ setAdjective }) => {}; + + const { rerender } = render( + + + + ); + + await waitFor(() => { + expect(FakeWebSocket.instances.length).toBe(1); + }); + + const [websocket] = FakeWebSocket.instances; + + act(() => { + websocket.open(); + }); + + rerender( + + + + + ); + + await waitFor(() => { + expect(greeters["later"]).toBeDefined(); + }); + + setAdjective("later", "later"); + + await waitFor(() => { + expect(websocket.sent.length).toBe(1); + }); + + act(() => { + websocket.respond(); + }); + + await waitFor(() => { + expect(resolved).toEqual(["later"]); + }); + + expect(FakeWebSocket.instances.length).toBe(1); + }); + + it("stays open until no state is used anymore", async () => { + use = ({ setAdjective }) => {}; + + const { rerender, unmount } = render( + + + + + ); + + await waitFor(() => { + expect(FakeWebSocket.instances.length).toBe(1); + }); + + const [websocket] = FakeWebSocket.instances; + + act(() => { + websocket.open(); + }); + + rerender( + + + + ); + + expect(websocket.readyState).toBe(FakeWebSocket.OPEN); + + setAdjective("second", "second"); + + await waitFor(() => { + expect(websocket.sent.length).toBe(1); + }); + + unmount(); + + expect(websocket.readyState).toBe(FakeWebSocket.CLOSED); + }); + + it("is opened again if it gets closed", async () => { + use = ({ setAdjective }) => {}; + + render( + + + + + ); + + await waitFor(() => { + expect(FakeWebSocket.instances.length).toBe(1); + }); + + act(() => { + FakeWebSocket.instances[0].open(); + }); + + setAdjective("first", "first"); + setAdjective("second", "second"); + + await waitFor(() => { + expect(FakeWebSocket.instances[0].sent.length).toBe(2); + }); + + act(() => { + FakeWebSocket.instances[0].fail(); + }); + + await waitFor( + () => { + expect(FakeWebSocket.instances.length).toBe(2); + }, + { timeout: 10000 } + ); + + const websocket = FakeWebSocket.instances[1]; + + expect(websocket.url).toBe("wss://reboot.test/__/reboot/websocket/mutate"); + + // Every state waits for a while of its own before it tries again. + await new Promise((resolve) => setTimeout(resolve, 4000)); + + expect(FakeWebSocket.instances.length).toBe(2); + + act(() => { + websocket.open(); + }); + + // The mutations that did not get a response are sent again. + expect([...adjectives(websocket)].sort()).toEqual(["first", "second"]); + + act(() => { + websocket.respond(); + websocket.respond(); + }); + + await waitFor(() => { + expect([...resolved].sort()).toEqual(["first", "second"]); + }); + }, 30000); + + it("is a websocket for every state if the backend requires it", async () => { + use = ({ setAdjective }) => {}; + + render( + + + + + ); + + await waitFor(() => { + expect(FakeWebSocket.instances.length).toBe(1); + }); + + expect(FakeWebSocket.instances[0].url).toBe( + "wss://old.reboot.test/__/reboot/websocket/mutate" + ); + + act(() => { + FakeWebSocket.instances[0].open(); + }); + + setAdjective("first", "first"); + setAdjective("second", "second"); + + await waitFor(() => { + expect(FakeWebSocket.instances[0].sent.length).toBe(2); + }); + + act(() => { + FakeWebSocket.instances[0].respondWithoutState(); + }); + + // Neither mutation has been resolved with that response. + expect(resolved).toEqual([]); + + await waitFor( + () => { + expect(FakeWebSocket.instances.length).toBe(3); + }, + { timeout: 10000 } + ); + + const websockets = FakeWebSocket.instances.slice(1); + + expect(websockets.map(({ url }) => url).sort()).toEqual([ + `wss://old.reboot.test/__/reboot/rpc/${stateRef("first")}`, + `wss://old.reboot.test/__/reboot/rpc/${stateRef("second")}`, + ]); + + for (const websocket of websockets) { + act(() => { + websocket.open(); + }); + + // The mutation that did not get a response is sent again. + expect(websocket.sent.length).toBe(1); + + act(() => { + websocket.respond(); + }); + } + + await waitFor(() => { + expect([...resolved].sort()).toEqual(["first", "second"]); + }); + }, 30000); + it("does not need a request of its own", async () => { use = ({ setAdjective }) => {};