codekingpro/portable-devtools
115k
1"""Reader for WebSocket protocol versions 13 and 8."""
2
3import asyncio
4import builtins
5from collections import deque
6from typing import Deque, Final, Optional, Set, Tuple, Union
7
8from ..base_protocol import BaseProtocol
9from ..compression_utils import ZLibDecompressor
10from ..helpers import _EXC_SENTINEL, set_exception
11from ..streams import EofStream
12from .helpers import UNPACK_CLOSE_CODE, UNPACK_LEN3, websocket_mask
13from .models import (
14 WS_DEFLATE_TRAILING,
15 WebSocketError,
16 WSCloseCode,
17 WSMessage,
18 WSMsgType,
19)
20
21ALLOWED_CLOSE_CODES: Final[Set[int]] = {int(i) for i in WSCloseCode}
22
23# States for the reader, used to parse the WebSocket frame
24# integer values are used so they can be cythonized
25READ_HEADER = 1
26READ_PAYLOAD_LENGTH = 2
27READ_PAYLOAD_MASK = 3
28READ_PAYLOAD = 4
29
30WS_MSG_TYPE_BINARY = WSMsgType.BINARY
31WS_MSG_TYPE_TEXT = WSMsgType.TEXT
32
33# WSMsgType values unpacked so they can by cythonized to ints
34OP_CODE_NOT_SET = -1
35OP_CODE_CONTINUATION = WSMsgType.CONTINUATION.value
36OP_CODE_TEXT = WSMsgType.TEXT.value
37OP_CODE_BINARY = WSMsgType.BINARY.value
38OP_CODE_CLOSE = WSMsgType.CLOSE.value
39OP_CODE_PING = WSMsgType.PING.value
40OP_CODE_PONG = WSMsgType.PONG.value
41
42EMPTY_FRAME_ERROR = (True, b"")
43EMPTY_FRAME = (False, b"")
44
45COMPRESSED_NOT_SET = -1
46COMPRESSED_FALSE = 0
47COMPRESSED_TRUE = 1
48
49TUPLE_NEW = tuple.__new__
50
51cython_int = int # Typed to int in Python, but cython with use a signed int in the pxd
52
53
54class WebSocketDataQueue:
55 """WebSocketDataQueue resumes and pauses an underlying stream.
56
57 It is a destination for WebSocket data.
58 """
59
60 def __init__(
61 self, protocol: BaseProtocol, limit: int, *, loop: asyncio.AbstractEventLoop
62 ) -> None:
63 self._size = 0
64 self._protocol = protocol
65 self._limit = limit * 2
66 self._loop = loop
67 self._eof = False
68 self._waiter: Optional[asyncio.Future[None]] = None
69 self._exception: Union[BaseException, None] = None
70 self._buffer: Deque[Tuple[WSMessage, int]] = deque()
71 self._get_buffer = self._buffer.popleft
72 self._put_buffer = self._buffer.append
73
74 def is_eof(self) -> bool:
75 return self._eof
76
77 def exception(self) -> Optional[BaseException]:
78 return self._exception
79
80 def set_exception(
81 self,
82 exc: BaseException,
83 exc_cause: builtins.BaseException = _EXC_SENTINEL,
84 ) -> None:
85 self._eof = True
86 self._exception = exc
87 if (waiter := self._waiter) is not None:
88 self._waiter = None
89 set_exception(waiter, exc, exc_cause)
90
91 def _release_waiter(self) -> None:
92 if (waiter := self._waiter) is None:
93 return
94 self._waiter = None
95 if not waiter.done():
96 waiter.set_result(None)
97
98 def feed_eof(self) -> None:
99 self._eof = True
100 self._release_waiter()
101 self._exception = None # Break cyclic references
102
103 def feed_data(self, data: "WSMessage", size: "cython_int") -> None:
104 self._size += size
105 self._put_buffer((data, size))
106 self._release_waiter()
107 if self._size > self._limit and not self._protocol._reading_paused:
108 self._protocol.pause_reading()
109
110 async def read(self) -> WSMessage:
111 if not self._buffer and not self._eof:
112 assert not self._waiter
113 self._waiter = self._loop.create_future()
114 try:
115 await self._waiter
116 except (asyncio.CancelledError, asyncio.TimeoutError):
117 self._waiter = None
118 raise
119 return self._read_from_buffer()
120
121 def _read_from_buffer(self) -> WSMessage:
122 if self._buffer:
123 data, size = self._get_buffer()
124 self._size -= size
125 if self._size < self._limit and self._protocol._reading_paused:
126 self._protocol.resume_reading()
127 return data
128 if self._exception is not None:
129 raise self._exception
130 raise EofStream
131
132
133class WebSocketReader:
134 def __init__(
135 self, queue: WebSocketDataQueue, max_msg_size: int, compress: bool = True
136 ) -> None:
137 self.queue = queue
138 self._max_msg_size = max_msg_size
139
140 self._exc: Optional[Exception] = None
141 self._partial = bytearray()
142 self._state = READ_HEADER
143
144 self._opcode: int = OP_CODE_NOT_SET
145 self._frame_fin = False
146 self._frame_opcode: int = OP_CODE_NOT_SET
147 self._payload_fragments: list[bytes] = []
148 self._frame_payload_len = 0
149
150 self._tail: bytes = b""
151 self._has_mask = False
152 self._frame_mask: Optional[bytes] = None
153 self._payload_bytes_to_read = 0
154 self._payload_len_flag = 0
155 self._compressed: int = COMPRESSED_NOT_SET
156 self._decompressobj: Optional[ZLibDecompressor] = None
157 self._compress = compress
158
159 def feed_eof(self) -> None:
160 self.queue.feed_eof()
161
162 # data can be bytearray on Windows because proactor event loop uses bytearray
163 # and asyncio types this to Union[bytes, bytearray, memoryview] so we need
164 # coerce data to bytes if it is not
165 def feed_data(
166 self, data: Union[bytes, bytearray, memoryview]
167 ) -> Tuple[bool, bytes]:
168 if type(data) is not bytes:
169 data = bytes(data)
170
171 if self._exc is not None:
172 return True, data
173
174 try:
175 self._feed_data(data)
176 except Exception as exc:
177 self._exc = exc
178 set_exception(self.queue, exc)
179 return EMPTY_FRAME_ERROR
180
181 return EMPTY_FRAME
182
183 def _handle_frame(
184 self,
185 fin: bool,
186 opcode: Union[int, cython_int], # Union intended: Cython pxd uses C int
187 payload: Union[bytes, bytearray],
188 compressed: Union[int, cython_int], # Union intended: Cython pxd uses C int
189 ) -> None:
190 msg: WSMessage
191 if opcode in {OP_CODE_TEXT, OP_CODE_BINARY, OP_CODE_CONTINUATION}:
192 # Validate continuation frames before processing
193 if opcode == OP_CODE_CONTINUATION and self._opcode == OP_CODE_NOT_SET:
194 raise WebSocketError(
195 WSCloseCode.PROTOCOL_ERROR,
196 "Continuation frame for non started message",
197 )
198
199 # load text/binary
200 if not fin:
201 # got partial frame payload
202 if opcode != OP_CODE_CONTINUATION:
203 self._opcode = opcode
204 self._partial += payload
205 if self._max_msg_size and len(self._partial) >= self._max_msg_size:
206 raise WebSocketError(
207 WSCloseCode.MESSAGE_TOO_BIG,
208 f"Message size {len(self._partial)} "
209 f"exceeds limit {self._max_msg_size}",
210 )
211 return
212
213 has_partial = bool(self._partial)
214 if opcode == OP_CODE_CONTINUATION:
215 opcode = self._opcode
216 self._opcode = OP_CODE_NOT_SET
217 # previous frame was non finished
218 # we should get continuation opcode
219 elif has_partial:
220 raise WebSocketError(
221 WSCloseCode.PROTOCOL_ERROR,
222 "The opcode in non-fin frame is expected "
223 f"to be zero, got {opcode!r}",
224 )
225
226 assembled_payload: Union[bytes, bytearray]
227 if has_partial:
228 assembled_payload = self._partial + payload
229 self._partial.clear()
230 else:
231 assembled_payload = payload
232
233 if self._max_msg_size and len(assembled_payload) >= self._max_msg_size:
234 raise WebSocketError(
235 WSCloseCode.MESSAGE_TOO_BIG,
236 f"Message size {len(assembled_payload)} "
237 f"exceeds limit {self._max_msg_size}",
238 )
239
240 # Decompress process must to be done after all packets
241 # received.
242 if compressed:
243 if not self._decompressobj:
244 self._decompressobj = ZLibDecompressor(suppress_deflate_header=True)
245 # XXX: It's possible that the zlib backend (isal is known to
246 # do this, maybe others too?) will return max_length bytes,
247 # but internally buffer more data such that the payload is
248 # >max_length, so we return one extra byte and if we're able
249 # to do that, then the message is too big.
250 payload_merged = self._decompressobj.decompress_sync(
251 assembled_payload + WS_DEFLATE_TRAILING,
252 (
253 self._max_msg_size + 1
254 if self._max_msg_size
255 else self._max_msg_size
256 ),
257 )
258 if self._max_msg_size and len(payload_merged) > self._max_msg_size:
259 raise WebSocketError(
260 WSCloseCode.MESSAGE_TOO_BIG,
261 f"Decompressed message exceeds size limit {self._max_msg_size}",
262 )
263 elif type(assembled_payload) is bytes:
264 payload_merged = assembled_payload
265 else:
266 payload_merged = bytes(assembled_payload)
267
268 if opcode == OP_CODE_TEXT:
269 try:
270 text = payload_merged.decode("utf-8")
271 except UnicodeDecodeError as exc:
272 raise WebSocketError(
273 WSCloseCode.INVALID_TEXT, "Invalid UTF-8 text message"
274 ) from exc
275
276 # XXX: The Text and Binary messages here can be a performance
277 # bottleneck, so we use tuple.__new__ to improve performance.
278 # This is not type safe, but many tests should fail in
279 # test_client_ws_functional.py if this is wrong.
280 self.queue.feed_data(
281 TUPLE_NEW(WSMessage, (WS_MSG_TYPE_TEXT, text, "")),
282 len(payload_merged),
283 )
284 else:
285 self.queue.feed_data(
286 TUPLE_NEW(WSMessage, (WS_MSG_TYPE_BINARY, payload_merged, "")),
287 len(payload_merged),
288 )
289 elif opcode == OP_CODE_CLOSE:
290 if len(payload) >= 2:
291 close_code = UNPACK_CLOSE_CODE(payload[:2])[0]
292 if close_code < 3000 and close_code not in ALLOWED_CLOSE_CODES:
293 raise WebSocketError(
294 WSCloseCode.PROTOCOL_ERROR,
295 f"Invalid close code: {close_code}",
296 )
297 try:
298 close_message = payload[2:].decode("utf-8")
299 except UnicodeDecodeError as exc:
300 raise WebSocketError(
301 WSCloseCode.INVALID_TEXT, "Invalid UTF-8 text message"
302 ) from exc
303 msg = TUPLE_NEW(WSMessage, (WSMsgType.CLOSE, close_code, close_message))
304 elif payload:
305 raise WebSocketError(
306 WSCloseCode.PROTOCOL_ERROR,
307 f"Invalid close frame: {fin} {opcode} {payload!r}",
308 )
309 else:
310 msg = TUPLE_NEW(WSMessage, (WSMsgType.CLOSE, 0, ""))
311
312 self.queue.feed_data(msg, 0)
313 elif opcode == OP_CODE_PING:
314 msg = TUPLE_NEW(WSMessage, (WSMsgType.PING, payload, ""))
315 self.queue.feed_data(msg, len(payload))
316 elif opcode == OP_CODE_PONG:
317 msg = TUPLE_NEW(WSMessage, (WSMsgType.PONG, payload, ""))
318 self.queue.feed_data(msg, len(payload))
319 else:
320 raise WebSocketError(
321 WSCloseCode.PROTOCOL_ERROR, f"Unexpected opcode={opcode!r}"
322 )
323
324 def _feed_data(self, data: bytes) -> None:
325 """Return the next frame from the socket."""
326 if self._tail:
327 data, self._tail = self._tail + data, b""
328
329 start_pos: int = 0
330 data_len = len(data)
331 data_cstr = data
332
333 while True:
334 # read header
335 if self._state == READ_HEADER:
336 if data_len - start_pos < 2:
337 break
338 first_byte = data_cstr[start_pos]
339 second_byte = data_cstr[start_pos + 1]
340 start_pos += 2
341
342 fin = (first_byte >> 7) & 1
343 rsv1 = (first_byte >> 6) & 1
344 rsv2 = (first_byte >> 5) & 1
345 rsv3 = (first_byte >> 4) & 1
346 opcode = first_byte & 0xF
347
348 # frame-fin = %x0 ; more frames of this message follow
349 # / %x1 ; final frame of this message
350 # frame-rsv1 = %x0 ;
351 # 1 bit, MUST be 0 unless negotiated otherwise
352 # frame-rsv2 = %x0 ;
353 # 1 bit, MUST be 0 unless negotiated otherwise
354 # frame-rsv3 = %x0 ;
355 # 1 bit, MUST be 0 unless negotiated otherwise
356 #
357 # Remove rsv1 from this test for deflate development
358 if rsv2 or rsv3 or (rsv1 and not self._compress):
359 raise WebSocketError(
360 WSCloseCode.PROTOCOL_ERROR,
361 "Received frame with non-zero reserved bits",
362 )
363
364 if opcode > 0x7 and fin == 0:
365 raise WebSocketError(
366 WSCloseCode.PROTOCOL_ERROR,
367 "Received fragmented control frame",
368 )
369
370 has_mask = (second_byte >> 7) & 1
371 length = second_byte & 0x7F
372
373 # Control frames MUST have a payload
374 # length of 125 bytes or less
375 if opcode > 0x7 and length > 125:
376 raise WebSocketError(
377 WSCloseCode.PROTOCOL_ERROR,
378 "Control frame payload cannot be larger than 125 bytes",
379 )
380
381 # Set compress status if last package is FIN
382 # OR set compress status if this is first fragment
383 # Raise error if not first fragment with rsv1 = 0x1
384 if self._frame_fin or self._compressed == COMPRESSED_NOT_SET:
385 self._compressed = COMPRESSED_TRUE if rsv1 else COMPRESSED_FALSE
386 elif rsv1:
387 raise WebSocketError(
388 WSCloseCode.PROTOCOL_ERROR,
389 "Received frame with non-zero reserved bits",
390 )
391
392 self._frame_fin = bool(fin)
393 self._frame_opcode = opcode
394 self._has_mask = bool(has_mask)
395 self._payload_len_flag = length
396 self._state = READ_PAYLOAD_LENGTH
397
398 # read payload length
399 if self._state == READ_PAYLOAD_LENGTH:
400 len_flag = self._payload_len_flag
401 if len_flag == 126:
402 if data_len - start_pos < 2:
403 break
404 first_byte = data_cstr[start_pos]
405 second_byte = data_cstr[start_pos + 1]
406 start_pos += 2
407 self._payload_bytes_to_read = first_byte << 8 | second_byte
408 elif len_flag > 126:
409 if data_len - start_pos < 8:
410 break
411 self._payload_bytes_to_read = UNPACK_LEN3(data, start_pos)[0]
412 start_pos += 8
413 else:
414 self._payload_bytes_to_read = len_flag
415
416 self._state = READ_PAYLOAD_MASK if self._has_mask else READ_PAYLOAD
417
418 # read payload mask
419 if self._state == READ_PAYLOAD_MASK:
420 if data_len - start_pos < 4:
421 break
422 self._frame_mask = data_cstr[start_pos : start_pos + 4]
423 start_pos += 4
424 self._state = READ_PAYLOAD
425
426 if self._state == READ_PAYLOAD:
427 chunk_len = data_len - start_pos
428 if self._payload_bytes_to_read >= chunk_len:
429 f_end_pos = data_len
430 self._payload_bytes_to_read -= chunk_len
431 else:
432 f_end_pos = start_pos + self._payload_bytes_to_read
433 self._payload_bytes_to_read = 0
434
435 had_fragments = self._frame_payload_len
436 self._frame_payload_len += f_end_pos - start_pos
437 f_start_pos = start_pos
438 start_pos = f_end_pos
439
440 if self._payload_bytes_to_read != 0:
441 # If we don't have a complete frame, we need to save the
442 # data for the next call to feed_data.
443 self._payload_fragments.append(data_cstr[f_start_pos:f_end_pos])
444 break
445
446 payload: Union[bytes, bytearray]
447 if had_fragments:
448 # We have to join the payload fragments get the payload
449 self._payload_fragments.append(data_cstr[f_start_pos:f_end_pos])
450 if self._has_mask:
451 assert self._frame_mask is not None
452 payload_bytearray = bytearray(b"".join(self._payload_fragments))
453 websocket_mask(self._frame_mask, payload_bytearray)
454 payload = payload_bytearray
455 else:
456 payload = b"".join(self._payload_fragments)
457 self._payload_fragments.clear()
458 elif self._has_mask:
459 assert self._frame_mask is not None
460 payload_bytearray = data_cstr[f_start_pos:f_end_pos] # type: ignore[assignment]
461 if type(payload_bytearray) is not bytearray: # pragma: no branch
462 # Cython will do the conversion for us
463 # but we need to do it for Python and we
464 # will always get here in Python
465 payload_bytearray = bytearray(payload_bytearray)
466 websocket_mask(self._frame_mask, payload_bytearray)
467 payload = payload_bytearray
468 else:
469 payload = data_cstr[f_start_pos:f_end_pos]
470
471 self._handle_frame(
472 self._frame_fin, self._frame_opcode, payload, self._compressed
473 )
474 self._frame_payload_len = 0
475 self._state = READ_HEADER
476
477 # XXX: Cython needs slices to be bounded, so we can't omit the slice end here.
478 self._tail = data_cstr[start_pos:data_len] if start_pos < data_len else b""
479 