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/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/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/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() 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 }) => {};