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")