v0.166.0
  1"""Driving a view's `websocket()` in-process from a test.
  2
  3    with client.websocket("/live/") as ws:
  4        ws.send("hello")
  5        assert ws.receive() == "hello"
  6
  7The handshake runs through the same pipeline as any test-client request.
  8On acceptance the socket is served by the server's own code
  9(`HandledRequest.serve_websocket()`), over one end of a socketpair, on an
 10event loop this object owns and steps from the test's own thread whenever
 11it is asked to send, receive, or close. This object is the other end: the
 12client's side of the protocol. No background thread: the view runs on the
 13test thread, inside a copy of the test's context, so the test database
 14transaction is visible to it.
 15"""
 16
 17import asyncio
 18import base64
 19import os
 20import socket
 21from types import TracebackType
 22from typing import TYPE_CHECKING, Any, Self
 23
 24from plain.http import WebSocketClosed, WebSocketResponse
 25from plain.http.websocket_frames import (
 26    OP_BINARY,
 27    OP_CLOSE,
 28    OP_CONTINUATION,
 29    OP_PING,
 30    OP_PONG,
 31    OP_TEXT,
 32    IncompleteFrame,
 33    encode_close,
 34    encode_frame,
 35    parse_close_payload,
 36    read_frame,
 37)
 38
 39if TYPE_CHECKING:
 40    from plain.server.inprocess import HandledRequest
 41
 42    from .client import ClientResponse
 43
 44# How long a call on the connection waits, unless the test says otherwise.
 45DEFAULT_TIMEOUT = 5.0
 46
 47# Server-to-client frames are read without a size cap: the test wrote the
 48# view and knows what it sends.
 49_NO_CAP = 1 << 40
 50
 51
 52def handshake_headers(
 53    *, key: str | None = None, subprotocols: tuple[str, ...] = ()
 54) -> dict[str, str]:
 55    """The request headers of a well-formed RFC 6455 opening handshake.
 56
 57    A fresh random key unless one is given (tests pin the RFC's example
 58    key to assert on the accept value).
 59    """
 60    headers = {
 61        "Upgrade": "websocket",
 62        "Connection": "Upgrade",
 63        "Sec-WebSocket-Version": "13",
 64        "Sec-WebSocket-Key": key or base64.b64encode(os.urandom(16)).decode(),
 65    }
 66    if subprotocols:
 67        headers["Sec-WebSocket-Protocol"] = ", ".join(subprotocols)
 68    return headers
 69
 70
 71def _client_frame(opcode: int, payload: bytes = b"") -> bytes:
 72    """One client-to-server frame, masked with a fresh key as RFC 6455 requires."""
 73    return encode_frame(opcode, payload, mask=os.urandom(4))
 74
 75
 76class WebSocketRejected(Exception):
 77    """The upgrade did not produce a socket; `.response` says why (a 403, a redirect).
 78
 79    `.response` is the same kind of response any other client request returns.
 80    """
 81
 82    def __init__(self, response: ClientResponse) -> None:
 83        self.response = response
 84        super().__init__(f"WebSocket upgrade rejected with {response.status_code}")
 85
 86
 87class WebSocketTestConnection:
 88    """The client end of an accepted websocket, driven synchronously."""
 89
 90    def __init__(self, handled: HandledRequest, *, timeout: float) -> None:
 91        response = handled.response
 92        # For the type checker: `Client.websocket()` is the only caller, and
 93        # it has already turned anything else into a `WebSocketRejected`.
 94        assert isinstance(response, WebSocketResponse)
 95        # The handshake: the request that asked for the socket, and the 101
 96        # that granted it.
 97        self.request = handled.request
 98        self.response = response
 99        self.subprotocol = response.subprotocol
100        self._timeout = timeout
101        self._closed = False
102
103        self._loop = asyncio.new_event_loop()
104        server_sock, client_sock = socket.socketpair()
105        try:
106            self._reader, self._writer = self._loop.run_until_complete(
107                asyncio.open_connection(sock=client_sock)
108            )
109        except BaseException:
110            server_sock.close()
111            client_sock.close()
112            self._loop.close()
113            raise
114        # The view runs in the context the handshake ran in, a copy of the
115        # test's: it sees what the test set up (the test's database
116        # transaction included), with the handshake's request span current.
117        self._server_task = self._loop.create_task(handled.serve_websocket(server_sock))
118
119    async def _settle_transports(self) -> None:
120        """Let asyncio finish closing both ends before the loop goes away."""
121        try:
122            await asyncio.wait_for(self._writer.wait_closed(), 1)
123        except TimeoutError, OSError:
124            pass
125        await asyncio.sleep(0)
126
127    # ------------------------------------------------------------------
128
129    def __enter__(self) -> Self:
130        return self
131
132    def __exit__(
133        self,
134        exc_type: type[BaseException] | None,
135        exc: BaseException | None,
136        tb: TracebackType | None,
137    ) -> None:
138        try:
139            if not self._closed:
140                self.close()
141        finally:
142            self._writer.close()
143            self._loop.run_until_complete(self._settle_transports())
144            self._loop.close()
145        if exc is None:
146            self._raise_view_error()
147
148    def send(self, message: str | bytes, *, timeout: float | None = None) -> None:
149        """Send one message to the view."""
150        if isinstance(message, str):
151            frame = _client_frame(OP_TEXT, message.encode())
152        else:
153            frame = _client_frame(OP_BINARY, bytes(message))
154        self._writer.write(frame)
155        self._step(self._writer.drain(), timeout)
156
157    def receive(self, *, timeout: float | None = None) -> str | bytes:
158        """The next message the view sent.
159
160        Raises the view's own exception if it failed, `WebSocketClosed` if
161        it closed the socket, and `TimeoutError` if nothing arrives.
162        """
163        return self._step(self._read_message(), timeout)
164
165    def close(
166        self, code: int = 1000, reason: str = "", *, timeout: float | None = None
167    ) -> None:
168        """Close from the client side and wait for the view to finish."""
169        if self._closed:
170            return
171        self._closed = True
172        self._writer.write(encode_close(code, reason, mask=os.urandom(4)))
173        try:
174            self._step(self._writer.drain(), timeout)
175            self._step(asyncio.shield(self._server_task), timeout)
176        except TimeoutError, OSError:
177            # A view that never reads the socket does not see the CLOSE;
178            # stop it the way a worker shutdown would.
179            self._server_task.cancel()
180            self._loop.run_until_complete(
181                asyncio.gather(self._server_task, return_exceptions=True)
182            )
183
184    # ------------------------------------------------------------------
185
186    def _step(self, awaitable: Any, timeout: float | None) -> Any:
187        return self._loop.run_until_complete(
188            asyncio.wait_for(awaitable, self._timeout if timeout is None else timeout)
189        )
190
191    def _raise_view_error(self) -> None:
192        if not self._server_task.done() or self._server_task.cancelled():
193            return
194        # Two different failures: the serving code itself blew up, or the
195        # view failed and `run_websocket` logged it and returned it.
196        if (exc := self._server_task.exception()) is not None:
197            raise exc
198        if (view_error := self._server_task.result()) is not None:
199            raise view_error
200
201    async def _read_message(self) -> str | bytes:
202        fragments: list[bytes] = []
203        text = False
204        while True:
205            if self._server_task.done():
206                self._raise_view_error()
207            try:
208                frame = await read_frame(self._reader.read, _NO_CAP, require_mask=False)
209            except IncompleteFrame:
210                # The server closed the transport — after a CLOSE we already
211                # surfaced, or because the view failed.
212                await asyncio.gather(self._server_task, return_exceptions=True)
213                self._raise_view_error()
214                raise WebSocketClosed() from None
215
216            if frame.opcode == OP_PING:
217                self._writer.write(_client_frame(OP_PONG, frame.payload))
218                continue
219            if frame.opcode == OP_PONG:
220                continue
221            if frame.opcode == OP_CLOSE:
222                reason = parse_close_payload(frame.payload)
223                self._writer.write(_client_frame(OP_CLOSE, frame.payload))
224                self._closed = True
225                await asyncio.gather(self._server_task, return_exceptions=True)
226                self._raise_view_error()
227                raise WebSocketClosed(reason.code, reason.reason)
228
229            if frame.opcode == OP_TEXT:
230                text = True
231            fragments.append(frame.payload)
232            if frame.fin:
233                payload = b"".join(fragments)
234                return payload.decode() if text else payload
235            if frame.opcode != OP_CONTINUATION and len(fragments) > 1:
236                raise AssertionError("server interleaved two messages")