Coverage for /pythoncovmergedfiles/medio/medio/src/aiohttp/aiohttp/_websocket/reader_py.py: 16%
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
Shortcuts on this page
r m x toggle line displays
j k next/prev highlighted chunk
0 (zero) top of page
1 (one) first highlighted chunk
1"""Reader for WebSocket protocol versions 13 and 8."""
3import asyncio
4import builtins
5import sys
6import weakref
7from collections import deque
9from ..base_protocol import BaseProtocol
10from ..compression_utils import TooManyMembersError, ZLibDecompressor
11from ..helpers import _EXC_SENTINEL, set_exception
12from ..log import ws_logger
13from ..streams import EofStream
14from .helpers import UNPACK_CLOSE_CODE, UNPACK_LEN3, websocket_mask
15from .models import (
16 WS_DEFLATE_TRAILING,
17 WebSocketError,
18 WSCloseCode,
19 WSMessage,
20 WSMessageBinary,
21 WSMessageClose,
22 WSMessagePing,
23 WSMessagePong,
24 WSMessageText,
25 WSMessageTextBytes,
26 WSMsgType,
27)
29# ABNORMAL_CLOSURE is used internally, should never be accepted from a client.
30# https://datatracker.ietf.org/doc/html/rfc6455#section-7.4.1
31ALLOWED_CLOSE_CODES = {int(i) for i in WSCloseCode} - {
32 int(WSCloseCode.ABNORMAL_CLOSURE)
33}
35# States for the reader, used to parse the WebSocket frame
36# integer values are used so they can be cythonized
37READ_HEADER = 1
38READ_PAYLOAD_LENGTH = 2
39READ_PAYLOAD_MASK = 3
40READ_PAYLOAD = 4
42# Largest declared payload length the reader can represent: the compiled
43# reader stores it in a Py_ssize_t, which holds 2**31-1 on the 32-bit builds
44# (the win32 and armv7l wheels) and 2**63-1 everywhere else.
45# TODO: Remove when we drop 32 bit support (and from reader_c.pxd).
46MAX_PAYLOAD_LEN = sys.maxsize
48WS_MSG_TYPE_BINARY = WSMsgType.BINARY
49WS_MSG_TYPE_TEXT = WSMsgType.TEXT
51# WSMsgType values unpacked so they can by cythonized to ints
52OP_CODE_NOT_SET = -1
53OP_CODE_CONTINUATION = WSMsgType.CONTINUATION.value
54OP_CODE_TEXT = WSMsgType.TEXT.value
55OP_CODE_BINARY = WSMsgType.BINARY.value
56OP_CODE_CLOSE = WSMsgType.CLOSE.value
57OP_CODE_PING = WSMsgType.PING.value
58OP_CODE_PONG = WSMsgType.PONG.value
60EMPTY_FRAME_ERROR = (True, b"")
61EMPTY_FRAME = (False, b"")
63COMPRESSED_NOT_SET = -1
64COMPRESSED_FALSE = 0
65COMPRESSED_TRUE = 1
67TUPLE_NEW = tuple.__new__
69# Overhead added to each message to ensure that tiny messages can't use
70# unreasonable amounts of memory.
71MSG_SIZE_OVERHEAD = 128
73STALLED_READER_COLLECTED = (
74 "WebSocketReader was garbage collected while stalled; "
75 "callers of set_parser() must hold a strong reference"
76)
78cython_int = int # Typed to int in Python, but cython with use a signed int in the pxd
81class WebSocketDataQueue:
82 """WebSocketDataQueue resumes and pauses an underlying stream.
84 It is a destination for WebSocket data.
85 """
87 def __init__(
88 self, protocol: BaseProtocol, limit: int, *, loop: asyncio.AbstractEventLoop
89 ) -> None:
90 self._size = 0
91 self._protocol = protocol
92 self._limit = limit * 2
93 self._loop = loop
94 self._eof = False
95 self._waiter: asyncio.Future[None] | None = None
96 self._exception: type[BaseException] | BaseException | None = None
97 self._buffer: deque[WSMessage] = deque()
98 self._get_buffer = self._buffer.popleft
99 self._put_buffer = self._buffer.append
100 self._stalled_reader: "weakref.ref[WebSocketReader] | None" = None
102 def is_eof(self) -> bool:
103 return self._eof
105 def exception(self) -> type[BaseException] | BaseException | None:
106 return self._exception
108 def set_exception(
109 self,
110 exc: type[BaseException] | BaseException,
111 exc_cause: builtins.BaseException = _EXC_SENTINEL,
112 ) -> None:
113 self._eof = True
114 self._exception = exc
115 if (waiter := self._waiter) is not None:
116 self._waiter = None
117 set_exception(waiter, exc, exc_cause)
119 def _release_waiter(self) -> None:
120 if (waiter := self._waiter) is None:
121 return
122 self._waiter = None
123 if not waiter.done():
124 waiter.set_result(None)
126 def feed_eof(self) -> None:
127 self._eof = True
128 self._release_waiter()
129 self._exception = None # Break cyclic references
131 def feed_data(self, data: "WSMessage") -> None:
132 # Unbox into the typed local before adding, so Cython keeps the sum in
133 # C instead of boxing MSG_SIZE_OVERHEAD for a Python-level add.
134 size = data.size
135 self._size += size + MSG_SIZE_OVERHEAD
136 self._put_buffer(data)
137 self._release_waiter()
138 if self._size > self._limit and not self._protocol._reading_paused:
139 self._protocol.pause_reading()
141 async def read(self) -> WSMessage:
142 if not self._buffer and not self._eof:
143 assert not self._waiter
144 self._waiter = self._loop.create_future()
145 try:
146 await self._waiter
147 except (asyncio.CancelledError, asyncio.TimeoutError):
148 self._waiter = None
149 raise
150 return self._read_from_buffer()
152 def _read_from_buffer(self) -> WSMessage:
153 if self._buffer:
154 data = self._get_buffer()
155 size = data.size
156 self._size -= size + MSG_SIZE_OVERHEAD
157 if self._stalled_reader is not None and self._size <= self._limit // 2:
158 # Resume parsing once the queue drains to the low-water mark.
159 # Each resume re-slices the parser's unparsed tail, so waiting
160 # for headroom makes a drain cost one copy per batch of
161 # messages instead of one per message.
162 if (reader := self._stalled_reader()) is not None:
163 reader.feed_data(b"")
164 else:
165 # The stash died with the reader. Deliver what was already
166 # queued, then surface the contract violation on the next
167 # read instead of hanging. Log as well, since a caller that
168 # stops reading early never sees the deferred exception. A
169 # real failure that was already recorded stays the reported
170 # cause.
171 self._stalled_reader = None
172 ws_logger.warning(STALLED_READER_COLLECTED)
173 if self._exception is None:
174 self.set_exception(RuntimeError(STALLED_READER_COLLECTED))
175 # Resuming the transport while a stash remains would admit
176 # another socket read into the tail for every couple of messages
177 # drained, moving the memory bound from the queue into the tail.
178 if (
179 self._stalled_reader is None
180 and self._size < self._limit
181 and self._protocol._reading_paused
182 ):
183 self._protocol.resume_reading()
184 return data
185 if self._exception is not None:
186 raise self._exception
187 raise EofStream
190class WebSocketReader:
191 def __init__(
192 self,
193 queue: WebSocketDataQueue,
194 max_msg_size: int,
195 compress: bool,
196 decode_text: bool,
197 ) -> None:
198 self.queue = queue
199 self._max_msg_size = max_msg_size
200 self._decode_text = decode_text
201 # Parked on the queue while parsing is stalled; created once so
202 # stalling does not allocate.
203 self._weak_self = weakref.ref(self)
205 self._exc: Exception | None = None
206 self._partial = bytearray()
207 self._state = READ_HEADER
209 self._opcode: int = OP_CODE_NOT_SET
210 self._frame_fin = False
211 self._frame_opcode: int = OP_CODE_NOT_SET
212 # Reads of an in-flight frame, joined once when it completes.
213 self._payload_fragments: list[bytes] = []
214 # Fold reads into _payload_buffer past this count to bound the object
215 # count (bytes are bounded by max_msg_size).
216 self._max_fragments = max(1024, max_msg_size // 256) if max_msg_size else 0
217 self._payload_buffer = bytearray()
218 self._frame_payload_len = 0
220 self._tail: bytes = b""
221 self._has_mask = False
222 self._frame_mask: bytes | None = None
223 self._payload_bytes_to_read = 0
224 self._payload_len_flag = 0
225 self._compressed: int = COMPRESSED_NOT_SET
226 self._decompressobj: ZLibDecompressor | None = None
227 self._compress = compress
229 def feed_eof(self) -> None:
230 self.queue.feed_eof()
232 # data can be bytearray on Windows because proactor event loop uses bytearray
233 # and asyncio types this to Union[bytes, bytearray, memoryview] so we need
234 # coerce data to bytes if it is not
235 def feed_data(self, data: bytes | bytearray | memoryview) -> tuple[bool, bytes]:
236 if type(data) is not bytes:
237 data = bytes(data)
239 if self._exc is not None:
240 return True, data
242 try:
243 self._feed_data(data)
244 except Exception as exc:
245 self._exc = exc
246 set_exception(self.queue, exc)
247 return EMPTY_FRAME_ERROR
249 return EMPTY_FRAME
251 def _handle_frame(
252 self,
253 fin: bool,
254 opcode: int | cython_int, # Union intended: Cython pxd uses C int
255 payload: bytes | bytearray,
256 compressed: int | cython_int, # Union intended: Cython pxd uses C int
257 ) -> None:
258 msg: WSMessage
259 if opcode in {OP_CODE_TEXT, OP_CODE_BINARY, OP_CODE_CONTINUATION}:
260 # Validate continuation frames before processing
261 if opcode == OP_CODE_CONTINUATION and self._opcode == OP_CODE_NOT_SET:
262 raise WebSocketError(
263 WSCloseCode.PROTOCOL_ERROR,
264 "Continuation frame for non started message",
265 )
267 # load text/binary
268 if not fin:
269 # got partial frame payload
270 if opcode != OP_CODE_CONTINUATION:
271 self._opcode = opcode
272 self._partial += payload
273 return
275 has_partial = bool(self._partial)
276 if opcode == OP_CODE_CONTINUATION:
277 opcode = self._opcode
278 self._opcode = OP_CODE_NOT_SET
279 # previous frame was non finished
280 # we should get continuation opcode
281 elif has_partial:
282 raise WebSocketError(
283 WSCloseCode.PROTOCOL_ERROR,
284 "The opcode in non-fin frame is expected "
285 f"to be zero, got {opcode!r}",
286 )
288 assembled_payload: bytes | bytearray
289 if has_partial:
290 assembled_payload = self._partial + payload
291 self._partial.clear()
292 else:
293 assembled_payload = payload
295 # Decompress process must to be done after all packets
296 # received.
297 if compressed:
298 if not self._decompressobj:
299 self._decompressobj = ZLibDecompressor(suppress_deflate_header=True)
300 # XXX: It's possible that the zlib backend (isal is known to
301 # do this, maybe others too?) will return max_length bytes,
302 # but internally buffer more data such that the payload is
303 # >max_length, so we return one extra byte and if we're able
304 # to do that, then the message is too big.
305 try:
306 payload_merged = self._decompressobj.decompress_sync(
307 assembled_payload + WS_DEFLATE_TRAILING,
308 (
309 self._max_msg_size + 1
310 if self._max_msg_size
311 else self._max_msg_size
312 ),
313 )
314 except TooManyMembersError as exc:
315 raise WebSocketError(
316 WSCloseCode.MESSAGE_TOO_BIG,
317 "Compressed message has too many deflate members",
318 ) from exc
319 if self._max_msg_size and len(payload_merged) > self._max_msg_size:
320 raise WebSocketError(
321 WSCloseCode.MESSAGE_TOO_BIG,
322 f"Decompressed message exceeds size limit {self._max_msg_size}",
323 )
324 elif type(assembled_payload) is bytes:
325 payload_merged = assembled_payload
326 else:
327 payload_merged = bytes(assembled_payload)
329 size = len(payload_merged)
330 if opcode == OP_CODE_TEXT:
331 if self._decode_text:
332 try:
333 text = payload_merged.decode("utf-8")
334 except UnicodeDecodeError as exc:
335 raise WebSocketError(
336 WSCloseCode.INVALID_TEXT, "Invalid UTF-8 text message"
337 ) from exc
339 # XXX: The Text and Binary messages here can be a performance
340 # bottleneck, so we use tuple.__new__ to improve performance.
341 # This is not type safe, but many tests should fail in
342 # test_client_ws_functional.py if this is wrong.
343 msg = TUPLE_NEW(WSMessageText, (text, size, "", WS_MSG_TYPE_TEXT))
344 else:
345 # Return raw bytes for TEXT messages when decode_text=False
346 msg = TUPLE_NEW(
347 WSMessageTextBytes, (payload_merged, size, "", WS_MSG_TYPE_TEXT)
348 )
349 else:
350 msg = TUPLE_NEW(
351 WSMessageBinary, (payload_merged, size, "", WS_MSG_TYPE_BINARY)
352 )
354 self.queue.feed_data(msg)
355 elif opcode == OP_CODE_CLOSE:
356 payload_len = len(payload)
357 if payload_len >= 2:
358 close_code = UNPACK_CLOSE_CODE(payload[:2])[0]
359 # https://datatracker.ietf.org/doc/html/rfc6455#section-7.4.2
360 if close_code > 4999 or (
361 close_code < 3000 and close_code not in ALLOWED_CLOSE_CODES
362 ):
363 raise WebSocketError(
364 WSCloseCode.PROTOCOL_ERROR,
365 f"Invalid close code: {close_code}",
366 )
367 try:
368 close_message = payload[2:].decode("utf-8")
369 except UnicodeDecodeError as exc:
370 raise WebSocketError(
371 WSCloseCode.INVALID_TEXT, "Invalid UTF-8 text message"
372 ) from exc
373 msg = WSMessageClose(
374 data=close_code, size=payload_len, extra=close_message
375 )
376 elif payload:
377 raise WebSocketError(
378 WSCloseCode.PROTOCOL_ERROR,
379 f"Invalid close frame: {fin} {opcode} {payload!r}",
380 )
381 else:
382 msg = WSMessageClose(data=0, size=payload_len, extra="")
384 self.queue.feed_data(msg)
385 elif opcode == OP_CODE_PING:
386 self.queue.feed_data(
387 WSMessagePing(data=bytes(payload), size=len(payload), extra="")
388 )
389 elif opcode == OP_CODE_PONG:
390 self.queue.feed_data(
391 WSMessagePong(data=bytes(payload), size=len(payload), extra="")
392 )
393 else:
394 raise WebSocketError(
395 WSCloseCode.PROTOCOL_ERROR, f"Unexpected opcode={opcode!r}"
396 )
398 def _feed_data(self, data: bytes) -> None:
399 """Return the next frame from the socket."""
400 self.queue._stalled_reader = None
401 if self._tail:
402 data, self._tail = self._tail + data, b""
404 start_pos = 0
405 data_len = len(data)
406 data_cstr = data
408 while True:
409 if start_pos < data_len and self.queue._size > self.queue._limit:
410 # Over the high-water mark with unparsed bytes left: stash the
411 # remainder and stall. Gating on unparsed bytes keeps a read
412 # that ended on a frame boundary from arming an empty stall,
413 # which would hold the transport paused with nothing to drain.
414 self.queue._stalled_reader = self._weak_self
415 break
417 # read header
418 if self._state == READ_HEADER:
419 if data_len - start_pos < 2:
420 break
421 first_byte = data_cstr[start_pos]
422 second_byte = data_cstr[start_pos + 1]
423 start_pos += 2
425 fin = (first_byte >> 7) & 1
426 rsv1 = (first_byte >> 6) & 1
427 rsv2 = (first_byte >> 5) & 1
428 rsv3 = (first_byte >> 4) & 1
429 opcode = first_byte & 0xF
431 # frame-fin = %x0 ; more frames of this message follow
432 # / %x1 ; final frame of this message
433 # frame-rsv1 = %x0 ;
434 # 1 bit, MUST be 0 unless negotiated otherwise
435 # frame-rsv2 = %x0 ;
436 # 1 bit, MUST be 0 unless negotiated otherwise
437 # frame-rsv3 = %x0 ;
438 # 1 bit, MUST be 0 unless negotiated otherwise
439 #
440 # Remove rsv1 from this test for deflate development
441 if rsv2 or rsv3 or (rsv1 and not self._compress):
442 raise WebSocketError(
443 WSCloseCode.PROTOCOL_ERROR,
444 "Received frame with non-zero reserved bits",
445 )
447 if opcode not in {
448 OP_CODE_CONTINUATION,
449 OP_CODE_TEXT,
450 OP_CODE_BINARY,
451 OP_CODE_CLOSE,
452 OP_CODE_PING,
453 OP_CODE_PONG,
454 }:
455 raise WebSocketError(
456 WSCloseCode.PROTOCOL_ERROR,
457 f"Unexpected opcode={opcode!r}",
458 )
460 if opcode > 0x7 and fin == 0:
461 raise WebSocketError(
462 WSCloseCode.PROTOCOL_ERROR,
463 "Received fragmented control frame",
464 )
466 has_mask = (second_byte >> 7) & 1
467 length = second_byte & 0x7F
469 # Control frames MUST have a payload
470 # length of 125 bytes or less
471 if opcode > 0x7 and length > 125:
472 raise WebSocketError(
473 WSCloseCode.PROTOCOL_ERROR,
474 "Control frame payload cannot be larger than 125 bytes",
475 )
477 # Control frames (opcode > 0x7) may be interleaved between the
478 # fragments of a data message and never carry the per-message
479 # compressed bit, so they must not touch the compression state.
480 # https://datatracker.ietf.org/doc/html/rfc6455#section-5.4
481 # https://datatracker.ietf.org/doc/html/rfc7692#section-6.1
482 if opcode > 0x7:
483 if rsv1:
484 raise WebSocketError(
485 WSCloseCode.PROTOCOL_ERROR,
486 "Received frame with non-zero reserved bits",
487 )
488 else:
489 # Set compress status if last package is FIN
490 # OR set compress status if this is first fragment
491 # Raise error if not first fragment with rsv1 = 0x1
492 if self._frame_fin or self._compressed == COMPRESSED_NOT_SET:
493 self._compressed = COMPRESSED_TRUE if rsv1 else COMPRESSED_FALSE
494 elif rsv1:
495 raise WebSocketError(
496 WSCloseCode.PROTOCOL_ERROR,
497 "Received frame with non-zero reserved bits",
498 )
499 self._frame_fin = bool(fin)
501 self._frame_opcode = opcode
502 self._has_mask = bool(has_mask)
503 self._payload_len_flag = length
504 self._state = READ_PAYLOAD_LENGTH
506 # read payload length
507 if self._state == READ_PAYLOAD_LENGTH:
508 len_flag = self._payload_len_flag
509 if len_flag == 126:
510 if data_len - start_pos < 2:
511 break
512 first_byte = data_cstr[start_pos]
513 second_byte = data_cstr[start_pos + 1]
514 start_pos += 2
515 self._payload_bytes_to_read = first_byte << 8 | second_byte
516 elif len_flag > 126:
517 if data_len - start_pos < 8:
518 break
519 # The declared length is an unsigned 64-bit integer that
520 # does not necessarily fit _payload_bytes_to_read.
521 frame_len = UNPACK_LEN3(data, start_pos)[0]
522 if frame_len > MAX_PAYLOAD_LEN:
523 raise WebSocketError(
524 WSCloseCode.MESSAGE_TOO_BIG,
525 f"Message size {int(frame_len) + len(self._partial)} "
526 f"exceeds limit {self._max_msg_size or MAX_PAYLOAD_LEN}",
527 )
528 self._payload_bytes_to_read = frame_len
529 start_pos += 8
530 else:
531 self._payload_bytes_to_read = len_flag
533 # Reject oversized data frames before buffering any payload
534 # bytes. Control frames are capped at 125 bytes (checked in
535 # READ_HEADER) so only text/binary/continuation need this.
536 if self._max_msg_size and self._frame_opcode in {
537 OP_CODE_TEXT,
538 OP_CODE_BINARY,
539 OP_CODE_CONTINUATION,
540 }:
541 # partial_len declared in reader_c.pxd to keep it in C.
542 partial_len = len(self._partial)
543 # payload_bytes_to_read is a signed Py_ssize_t C value,
544 # use subtraction here to avoid an integer overflow.
545 if self._payload_bytes_to_read >= self._max_msg_size - partial_len:
546 raise WebSocketError(
547 WSCloseCode.MESSAGE_TOO_BIG,
548 f"Message size {int(self._payload_bytes_to_read) + partial_len} "
549 f"exceeds limit {self._max_msg_size}",
550 )
552 self._state = READ_PAYLOAD_MASK if self._has_mask else READ_PAYLOAD
554 # read payload mask
555 if self._state == READ_PAYLOAD_MASK:
556 if data_len - start_pos < 4:
557 break
558 self._frame_mask = data_cstr[start_pos : start_pos + 4]
559 start_pos += 4
560 self._state = READ_PAYLOAD
562 if self._state == READ_PAYLOAD:
563 chunk_len = data_len - start_pos
564 if self._payload_bytes_to_read >= chunk_len:
565 f_end_pos = data_len
566 self._payload_bytes_to_read -= chunk_len
567 else:
568 f_end_pos = start_pos + self._payload_bytes_to_read
569 self._payload_bytes_to_read = 0
571 had_fragments = self._frame_payload_len
572 self._frame_payload_len += f_end_pos - start_pos
573 f_start_pos = start_pos
574 start_pos = f_end_pos
576 if self._payload_bytes_to_read != 0:
577 if f_start_pos < f_end_pos: # skip a header-only read
578 self._payload_fragments.append(data_cstr[f_start_pos:f_end_pos])
579 if (
580 self._max_fragments
581 and len(self._payload_fragments) > self._max_fragments
582 ):
583 # Fold to bound the object count. Not a pause: nothing
584 # resumes reading until the frame is queued.
585 self._payload_buffer += b"".join(self._payload_fragments)
586 self._payload_fragments.clear()
587 break
589 payload: bytes | bytearray
590 if had_fragments:
591 self._payload_fragments.append(data_cstr[f_start_pos:f_end_pos])
592 if self._payload_buffer: # folded prefix
593 self._payload_buffer += b"".join(self._payload_fragments)
594 if self._has_mask:
595 assert self._frame_mask is not None
596 websocket_mask(self._frame_mask, self._payload_buffer)
597 payload = self._payload_buffer
598 self._payload_buffer = bytearray() # detach; payload aliases it
599 elif self._has_mask:
600 assert self._frame_mask is not None
601 payload_bytearray = bytearray(b"".join(self._payload_fragments))
602 websocket_mask(self._frame_mask, payload_bytearray)
603 payload = payload_bytearray
604 else:
605 payload = b"".join(self._payload_fragments)
606 self._payload_fragments.clear()
607 elif self._has_mask:
608 assert self._frame_mask is not None
609 payload_bytearray = data_cstr[f_start_pos:f_end_pos] # type: ignore[assignment]
610 if type(payload_bytearray) is not bytearray:
611 # Cython will do the conversion for us
612 # but we need to do it for Python and we
613 # will always get here in Python
614 payload_bytearray = bytearray(payload_bytearray)
615 websocket_mask(self._frame_mask, payload_bytearray)
616 payload = payload_bytearray
617 else:
618 payload = data_cstr[f_start_pos:f_end_pos]
620 self._handle_frame(
621 self._frame_fin, self._frame_opcode, payload, self._compressed
622 )
623 self._frame_payload_len = 0
624 self._state = READ_HEADER
626 # XXX: Cython needs slices to be bounded, so we can't omit the slice end here.
627 self._tail = data_cstr[start_pos:data_len] if start_pos < data_len else b""