Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
reader_c.py479 linesDownload Raw Back to _websocket
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 
codekingpro/portable-devtools · Team Ai