Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
websocket.py865 linesDownload Raw Back to eventlet
1import base642import codecs3import collections4import errno5from random import Random6from socket import error as SocketError7import string8import struct9import sys10import time11 12import zlib13 14try:15    from hashlib import md5, sha116except ImportError:  # pragma NO COVER17    from md5 import md518    from sha import sha as sha119 20from eventlet import semaphore21from eventlet import wsgi22from eventlet.green import socket23from eventlet.support import get_errno24import six25 26# Python 2's utf8 decoding is more lenient than we'd like27# In order to pass autobahn's testsuite we need stricter validation28# if available...29for _mod in ('wsaccel.utf8validator', 'autobahn.utf8validator'):30    # autobahn has it's own python-based validator. in newest versions31    # this prefers to use wsaccel, a cython based implementation, if available.32    # wsaccel may also be installed w/out autobahn, or with a earlier version.33    try:34        utf8validator = __import__(_mod, {}, {}, [''])35    except ImportError:36        utf8validator = None37    else:38        break39 40ACCEPTABLE_CLIENT_ERRORS = set((errno.ECONNRESET, errno.EPIPE))41DEFAULT_MAX_FRAME_LENGTH = 8 << 2042 43__all__ = ["WebSocketWSGI", "WebSocket"]44PROTOCOL_GUID = b'258EAFA5-E914-47DA-95CA-C5AB0DC85B11'45VALID_CLOSE_STATUS = set(46    list(range(1000, 1004)) +47    list(range(1007, 1012)) +48    # 3000-3999: reserved for use by libraries, frameworks,49    # and applications50    list(range(3000, 4000)) +51    # 4000-4999: reserved for private use and thus can't52    # be registered53    list(range(4000, 5000))54)55 56 57class BadRequest(Exception):58    def __init__(self, status='400 Bad Request', body=None, headers=None):59        super(Exception, self).__init__()60        self.status = status61        self.body = body62        self.headers = headers63 64 65class WebSocketWSGI(object):66    """Wraps a websocket handler function in a WSGI application.67 68    Use it like this::69 70      @websocket.WebSocketWSGI71      def my_handler(ws):72          from_browser = ws.wait()73          ws.send("from server")74 75    The single argument to the function will be an instance of76    :class:`WebSocket`.  To close the socket, simply return from the77    function.  Note that the server will log the websocket request at78    the time of closure.79 80    An optional argument max_frame_length can be given, which will set the81    maximum incoming *uncompressed* payload length of a frame. By default, this82    is set to 8MiB. Note that excessive values here might create a DOS attack83    vector.84    """85 86    def __init__(self, handler, max_frame_length=DEFAULT_MAX_FRAME_LENGTH):87        self.handler = handler88        self.protocol_version = None89        self.support_legacy_versions = True90        self.supported_protocols = []91        self.origin_checker = None92        self.max_frame_length = max_frame_length93 94    @classmethod95    def configured(cls,96                   handler=None,97                   supported_protocols=None,98                   origin_checker=None,99                   support_legacy_versions=False):100        def decorator(handler):101            inst = cls(handler)102            inst.support_legacy_versions = support_legacy_versions103            inst.origin_checker = origin_checker104            if supported_protocols:105                inst.supported_protocols = supported_protocols106            return inst107        if handler is None:108            return decorator109        return decorator(handler)110 111    def __call__(self, environ, start_response):112        http_connection_parts = [113            part.strip()114            for part in environ.get('HTTP_CONNECTION', '').lower().split(',')]115        if not ('upgrade' in http_connection_parts and116                environ.get('HTTP_UPGRADE', '').lower() == 'websocket'):117            # need to check a few more things here for true compliance118            start_response('400 Bad Request', [('Connection', 'close')])119            return []120 121        try:122            if 'HTTP_SEC_WEBSOCKET_VERSION' in environ:123                ws = self._handle_hybi_request(environ)124            elif self.support_legacy_versions:125                ws = self._handle_legacy_request(environ)126            else:127                raise BadRequest()128        except BadRequest as e:129            status = e.status130            body = e.body or b''131            headers = e.headers or []132            start_response(status,133                           [('Connection', 'close'), ] + headers)134            return [body]135 136        try:137            self.handler(ws)138        except socket.error as e:139            if get_errno(e) not in ACCEPTABLE_CLIENT_ERRORS:140                raise141        # Make sure we send the closing frame142        ws._send_closing_frame(True)143        # use this undocumented feature of eventlet.wsgi to ensure that it144        # doesn't barf on the fact that we didn't call start_response145        wsgi.WSGI_LOCAL.already_handled = True146        return []147 148    def _handle_legacy_request(self, environ):149        if 'eventlet.input' in environ:150            sock = environ['eventlet.input'].get_socket()151        elif 'gunicorn.socket' in environ:152            sock = environ['gunicorn.socket']153        else:154            raise Exception('No eventlet.input or gunicorn.socket present in environ.')155 156        if 'HTTP_SEC_WEBSOCKET_KEY1' in environ:157            self.protocol_version = 76158            if 'HTTP_SEC_WEBSOCKET_KEY2' not in environ:159                raise BadRequest()160        else:161            self.protocol_version = 75162 163        if self.protocol_version == 76:164            key1 = self._extract_number(environ['HTTP_SEC_WEBSOCKET_KEY1'])165            key2 = self._extract_number(environ['HTTP_SEC_WEBSOCKET_KEY2'])166            # There's no content-length header in the request, but it has 8167            # bytes of data.168            environ['wsgi.input'].content_length = 8169            key3 = environ['wsgi.input'].read(8)170            key = struct.pack(">II", key1, key2) + key3171            response = md5(key).digest()172 173        # Start building the response174        scheme = 'ws'175        if environ.get('wsgi.url_scheme') == 'https':176            scheme = 'wss'177        location = '%s://%s%s%s' % (178            scheme,179            environ.get('HTTP_HOST'),180            environ.get('SCRIPT_NAME'),181            environ.get('PATH_INFO')182        )183        qs = environ.get('QUERY_STRING')184        if qs is not None:185            location += '?' + qs186        if self.protocol_version == 75:187            handshake_reply = (188                b"HTTP/1.1 101 Web Socket Protocol Handshake\r\n"189                b"Upgrade: WebSocket\r\n"190                b"Connection: Upgrade\r\n"191                b"WebSocket-Origin: " + six.b(environ.get('HTTP_ORIGIN')) + b"\r\n"192                b"WebSocket-Location: " + six.b(location) + b"\r\n\r\n"193            )194        elif self.protocol_version == 76:195            handshake_reply = (196                b"HTTP/1.1 101 WebSocket Protocol Handshake\r\n"197                b"Upgrade: WebSocket\r\n"198                b"Connection: Upgrade\r\n"199                b"Sec-WebSocket-Origin: " + six.b(environ.get('HTTP_ORIGIN')) + b"\r\n"200                b"Sec-WebSocket-Protocol: " +201                six.b(environ.get('HTTP_SEC_WEBSOCKET_PROTOCOL', 'default')) + b"\r\n"202                b"Sec-WebSocket-Location: " + six.b(location) + b"\r\n"203                b"\r\n" + response204            )205        else:  # pragma NO COVER206            raise ValueError("Unknown WebSocket protocol version.")207        sock.sendall(handshake_reply)208        return WebSocket(sock, environ, self.protocol_version)209 210    def _parse_extension_header(self, header):211        if header is None:212            return None213        res = {}214        for ext in header.split(","):215            parts = ext.split(";")216            config = {}217            for part in parts[1:]:218                key_val = part.split("=")219                if len(key_val) == 1:220                    config[key_val[0].strip().lower()] = True221                else:222                    config[key_val[0].strip().lower()] = key_val[1].strip().strip('"').lower()223            res.setdefault(parts[0].strip().lower(), []).append(config)224        return res225 226    def _negotiate_permessage_deflate(self, extensions):227        if not extensions:228            return None229        deflate = extensions.get("permessage-deflate")230        if deflate is None:231            return None232        for config in deflate:233            # We'll evaluate each config in the client's preferred order and pick234            # the first that we can support.235            want_config = {236                # These are bool options, we can support both237                "server_no_context_takeover": config.get("server_no_context_takeover", False),238                "client_no_context_takeover": config.get("client_no_context_takeover", False)239            }240            # These are either bool OR int options. True means the client can accept a value241            # for the option, a number means the client wants that specific value.242            max_wbits = min(zlib.MAX_WBITS, 15)243            mwb = config.get("server_max_window_bits")244            if mwb is not None:245                if mwb is True:246                    want_config["server_max_window_bits"] = max_wbits247                else:248                    want_config["server_max_window_bits"] = \249                        int(config.get("server_max_window_bits", max_wbits))250                    if not (8 <= want_config["server_max_window_bits"] <= 15):251                        continue252            mwb = config.get("client_max_window_bits")253            if mwb is not None:254                if mwb is True:255                    want_config["client_max_window_bits"] = max_wbits256                else:257                    want_config["client_max_window_bits"] = \258                        int(config.get("client_max_window_bits", max_wbits))259                    if not (8 <= want_config["client_max_window_bits"] <= 15):260                        continue261            return want_config262        return None263 264    def _format_extension_header(self, parsed_extensions):265        if not parsed_extensions:266            return None267        parts = []268        for name, config in parsed_extensions.items():269            ext_parts = [six.b(name)]270            for key, value in config.items():271                if value is False:272                    pass273                elif value is True:274                    ext_parts.append(six.b(key))275                else:276                    ext_parts.append(six.b("%s=%s" % (key, str(value))))277            parts.append(b"; ".join(ext_parts))278        return b", ".join(parts)279 280    def _handle_hybi_request(self, environ):281        if 'eventlet.input' in environ:282            sock = environ['eventlet.input'].get_socket()283        elif 'gunicorn.socket' in environ:284            sock = environ['gunicorn.socket']285        else:286            raise Exception('No eventlet.input or gunicorn.socket present in environ.')287 288        hybi_version = environ['HTTP_SEC_WEBSOCKET_VERSION']289        if hybi_version not in ('8', '13', ):290            raise BadRequest(status='426 Upgrade Required',291                             headers=[('Sec-WebSocket-Version', '8, 13')])292        self.protocol_version = int(hybi_version)293        if 'HTTP_SEC_WEBSOCKET_KEY' not in environ:294            # That's bad.295            raise BadRequest()296        origin = environ.get(297            'HTTP_ORIGIN',298            (environ.get('HTTP_SEC_WEBSOCKET_ORIGIN', '')299             if self.protocol_version <= 8 else ''))300        if self.origin_checker is not None:301            if not self.origin_checker(environ.get('HTTP_HOST'), origin):302                raise BadRequest(status='403 Forbidden')303        protocols = environ.get('HTTP_SEC_WEBSOCKET_PROTOCOL', None)304        negotiated_protocol = None305        if protocols:306            for p in (i.strip() for i in protocols.split(',')):307                if p in self.supported_protocols:308                    negotiated_protocol = p309                    break310 311        key = environ['HTTP_SEC_WEBSOCKET_KEY']312        response = base64.b64encode(sha1(six.b(key) + PROTOCOL_GUID).digest())313        handshake_reply = [b"HTTP/1.1 101 Switching Protocols",314                           b"Upgrade: websocket",315                           b"Connection: Upgrade",316                           b"Sec-WebSocket-Accept: " + response]317        if negotiated_protocol:318            handshake_reply.append(b"Sec-WebSocket-Protocol: " + six.b(negotiated_protocol))319 320        parsed_extensions = {}321        extensions = self._parse_extension_header(environ.get("HTTP_SEC_WEBSOCKET_EXTENSIONS"))322 323        deflate = self._negotiate_permessage_deflate(extensions)324        if deflate is not None:325            parsed_extensions["permessage-deflate"] = deflate326 327        formatted_ext = self._format_extension_header(parsed_extensions)328        if formatted_ext is not None:329            handshake_reply.append(b"Sec-WebSocket-Extensions: " + formatted_ext)330 331        sock.sendall(b'\r\n'.join(handshake_reply) + b'\r\n\r\n')332        return RFC6455WebSocket(sock, environ, self.protocol_version,333                                protocol=negotiated_protocol,334                                extensions=parsed_extensions,335                                max_frame_length=self.max_frame_length)336 337    def _extract_number(self, value):338        """339        Utility function which, given a string like 'g98sd  5[]221@1', will340        return 9852211. Used to parse the Sec-WebSocket-Key headers.341        """342        out = ""343        spaces = 0344        for char in value:345            if char in string.digits:346                out += char347            elif char == " ":348                spaces += 1349        return int(out) // spaces350 351 352class WebSocket(object):353    """A websocket object that handles the details of354    serialization/deserialization to the socket.355 356    The primary way to interact with a :class:`WebSocket` object is to357    call :meth:`send` and :meth:`wait` in order to pass messages back358    and forth with the browser.  Also available are the following359    properties:360 361    path362        The path value of the request.  This is the same as the WSGI PATH_INFO variable,363        but more convenient.364    protocol365        The value of the Websocket-Protocol header.366    origin367        The value of the 'Origin' header.368    environ369        The full WSGI environment for this request.370 371    """372 373    def __init__(self, sock, environ, version=76):374        """375        :param socket: The eventlet socket376        :type socket: :class:`eventlet.greenio.GreenSocket`377        :param environ: The wsgi environment378        :param version: The WebSocket spec version to follow (default is 76)379        """380        self.log = environ.get('wsgi.errors', sys.stderr)381        self.log_context = 'server={shost}/{spath} client={caddr}:{cport}'.format(382            shost=environ.get('HTTP_HOST'),383            spath=environ.get('SCRIPT_NAME', '') + environ.get('PATH_INFO', ''),384            caddr=environ.get('REMOTE_ADDR'), cport=environ.get('REMOTE_PORT'),385        )386        self.socket = sock387        self.origin = environ.get('HTTP_ORIGIN')388        self.protocol = environ.get('HTTP_WEBSOCKET_PROTOCOL')389        self.path = environ.get('PATH_INFO')390        self.environ = environ391        self.version = version392        self.websocket_closed = False393        self._buf = b""394        self._msgs = collections.deque()395        self._sendlock = semaphore.Semaphore()396 397    def _pack_message(self, message):398        """Pack the message inside ``00`` and ``FF``399 400        As per the dataframing section (5.3) for the websocket spec401        """402        if isinstance(message, six.text_type):403            message = message.encode('utf-8')404        elif not isinstance(message, six.binary_type):405            message = six.b(str(message))406        packed = b"\x00" + message + b"\xFF"407        return packed408 409    def _parse_messages(self):410        """ Parses for messages in the buffer *buf*.  It is assumed that411        the buffer contains the start character for a message, but that it412        may contain only part of the rest of the message.413 414        Returns an array of messages, and the buffer remainder that415        didn't contain any full messages."""416        msgs = []417        end_idx = 0418        buf = self._buf419        while buf:420            frame_type = six.indexbytes(buf, 0)421            if frame_type == 0:422                # Normal message.423                end_idx = buf.find(b"\xFF")424                if end_idx == -1:  # pragma NO COVER425                    break426                msgs.append(buf[1:end_idx].decode('utf-8', 'replace'))427                buf = buf[end_idx + 1:]428            elif frame_type == 255:429                # Closing handshake.430                assert six.indexbytes(buf, 1) == 0, "Unexpected closing handshake: %r" % buf431                self.websocket_closed = True432                break433            else:434                raise ValueError("Don't understand how to parse this type of message: %r" % buf)435        self._buf = buf436        return msgs437 438    def send(self, message):439        """Send a message to the browser.440 441        *message* should be convertable to a string; unicode objects should be442        encodable as utf-8.  Raises socket.error with errno of 32443        (broken pipe) if the socket has already been closed by the client."""444        packed = self._pack_message(message)445        # if two greenthreads are trying to send at the same time446        # on the same socket, sendlock prevents interleaving and corruption447        self._sendlock.acquire()448        try:449            self.socket.sendall(packed)450        finally:451            self._sendlock.release()452 453    def wait(self):454        """Waits for and deserializes messages.455 456        Returns a single message; the oldest not yet processed. If the client457        has already closed the connection, returns None.  This is different458        from normal socket behavior because the empty string is a valid459        websocket message."""460        while not self._msgs:461            # Websocket might be closed already.462            if self.websocket_closed:463                return None464            # no parsed messages, must mean buf needs more data465            delta = self.socket.recv(8096)466            if delta == b'':467                return None468            self._buf += delta469            msgs = self._parse_messages()470            self._msgs.extend(msgs)471        return self._msgs.popleft()472 473    def _send_closing_frame(self, ignore_send_errors=False):474        """Sends the closing frame to the client, if required."""475        if self.version == 76 and not self.websocket_closed:476            try:477                self.socket.sendall(b"\xff\x00")478            except SocketError:479                # Sometimes, like when the remote side cuts off the connection,480                # we don't care about this.481                if not ignore_send_errors:  # pragma NO COVER482                    raise483            self.websocket_closed = True484 485    def close(self):486        """Forcibly close the websocket; generally it is preferable to487        return from the handler method."""488        try:489            self._send_closing_frame(True)490            self.socket.shutdown(True)491        except SocketError as e:492            if e.errno != errno.ENOTCONN:493                self.log.write('{ctx} socket shutdown error: {e}'.format(ctx=self.log_context, e=e))494        finally:495            self.socket.close()496 497 498class ConnectionClosedError(Exception):499    pass500 501 502class FailedConnectionError(Exception):503    def __init__(self, status, message):504        super(FailedConnectionError, self).__init__(status, message)505        self.message = message506        self.status = status507 508 509class ProtocolError(ValueError):510    pass511 512 513class RFC6455WebSocket(WebSocket):514    def __init__(self, sock, environ, version=13, protocol=None, client=False, extensions=None,515                 max_frame_length=DEFAULT_MAX_FRAME_LENGTH):516        super(RFC6455WebSocket, self).__init__(sock, environ, version)517        self.iterator = self._iter_frames()518        self.client = client519        self.protocol = protocol520        self.extensions = extensions or {}521 522        self._deflate_enc = None523        self._deflate_dec = None524        self.max_frame_length = max_frame_length525        self._remote_close_data = None526 527    class UTF8Decoder(object):528        def __init__(self):529            if utf8validator:530                self.validator = utf8validator.Utf8Validator()531            else:532                self.validator = None533            decoderclass = codecs.getincrementaldecoder('utf8')534            self.decoder = decoderclass()535 536        def reset(self):537            if self.validator:538                self.validator.reset()539            self.decoder.reset()540 541        def decode(self, data, final=False):542            if self.validator:543                valid, eocp, c_i, t_i = self.validator.validate(data)544                if not valid:545                    raise ValueError('Data is not valid unicode')546            return self.decoder.decode(data, final)547 548    def _get_permessage_deflate_enc(self):549        options = self.extensions.get("permessage-deflate")550        if options is None:551            return None552 553        def _make():554            return zlib.compressobj(zlib.Z_DEFAULT_COMPRESSION, zlib.DEFLATED,555                                    -options.get("client_max_window_bits" if self.client556                                                 else "server_max_window_bits",557                                                 zlib.MAX_WBITS))558 559        if options.get("client_no_context_takeover" if self.client560                       else "server_no_context_takeover"):561            # This option means we have to make a new one every time562            return _make()563        else:564            if self._deflate_enc is None:565                self._deflate_enc = _make()566            return self._deflate_enc567 568    def _get_permessage_deflate_dec(self, rsv1):569        options = self.extensions.get("permessage-deflate")570        if options is None or not rsv1:571            return None572 573        def _make():574            return zlib.decompressobj(-options.get("server_max_window_bits" if self.client575                                                   else "client_max_window_bits",576                                                   zlib.MAX_WBITS))577 578        if options.get("server_no_context_takeover" if self.client579                       else "client_no_context_takeover"):580            # This option means we have to make a new one every time581            return _make()582        else:583            if self._deflate_dec is None:584                self._deflate_dec = _make()585            return self._deflate_dec586 587    def _get_bytes(self, numbytes):588        data = b''589        while len(data) < numbytes:590            d = self.socket.recv(numbytes - len(data))591            if not d:592                raise ConnectionClosedError()593            data = data + d594        return data595 596    class Message(object):597        def __init__(self, opcode, max_frame_length, decoder=None, decompressor=None):598            self.decoder = decoder599            self.data = []600            self.finished = False601            self.opcode = opcode602            self.decompressor = decompressor603            self.max_frame_length = max_frame_length604 605        def push(self, data, final=False):606            self.finished = final607            self.data.append(data)608 609        def getvalue(self):610            data = b"".join(self.data)611            if not self.opcode & 8 and self.decompressor:612                data = self.decompressor.decompress(data + b"\x00\x00\xff\xff", self.max_frame_length)613                if self.decompressor.unconsumed_tail:614                    raise FailedConnectionError(615                        1009,616                        "Incoming compressed frame exceeds length limit of {} bytes.".format(self.max_frame_length))617 618            if self.decoder:619                data = self.decoder.decode(data, self.finished)620            return data621 622    @staticmethod623    def _apply_mask(data, mask, length=None, offset=0):624        if length is None:625            length = len(data)626        cnt = range(length)627        return b''.join(six.int2byte(six.indexbytes(data, i) ^ mask[(offset + i) % 4]) for i in cnt)628 629    def _handle_control_frame(self, opcode, data):630        if opcode == 8:  # connection close631            self._remote_close_data = data632            if not data:633                status = 1000634            elif len(data) > 1:635                status = struct.unpack_from('!H', data)[0]636                if not status or status not in VALID_CLOSE_STATUS:637                    raise FailedConnectionError(638                        1002,639                        "Unexpected close status code.")640                try:641                    data = self.UTF8Decoder().decode(data[2:], True)642                except (UnicodeDecodeError, ValueError):643                    raise FailedConnectionError(644                        1002,645                        "Close message data should be valid UTF-8.")646            else:647                status = 1002648            self.close(close_data=(status, ''))649            raise ConnectionClosedError()650        elif opcode == 9:  # ping651            self.send(data, control_code=0xA)652        elif opcode == 0xA:  # pong653            pass654        else:655            raise FailedConnectionError(656                1002, "Unknown control frame received.")657 658    def _iter_frames(self):659        fragmented_message = None660        try:661            while True:662                message = self._recv_frame(message=fragmented_message)663                if message.opcode & 8:664                    self._handle_control_frame(665                        message.opcode, message.getvalue())666                    continue667                if fragmented_message and message is not fragmented_message:668                    raise RuntimeError('Unexpected message change.')669                fragmented_message = message670                if message.finished:671                    data = fragmented_message.getvalue()672                    fragmented_message = None673                    yield data674        except FailedConnectionError:675            exc_typ, exc_val, exc_tb = sys.exc_info()676            self.close(close_data=(exc_val.status, exc_val.message))677        except ConnectionClosedError:678            return679        except Exception:680            self.close(close_data=(1011, 'Internal Server Error'))681            raise682 683    def _recv_frame(self, message=None):684        recv = self._get_bytes685 686        # Unpacking the frame described in Section 5.2 of RFC6455687        # (https://tools.ietf.org/html/rfc6455#section-5.2)688        header = recv(2)689        a, b = struct.unpack('!BB', header)690        finished = a >> 7 == 1691        rsv123 = a >> 4 & 7692        rsv1 = rsv123 & 4693        if rsv123:694            if rsv1 and "permessage-deflate" not in self.extensions:695                # must be zero - unless it's compressed then rsv1 is true696                raise FailedConnectionError(697                    1002,698                    "RSV1, RSV2, RSV3: MUST be 0 unless an extension is"699                    " negotiated that defines meanings for non-zero values.")700        opcode = a & 15701        if opcode not in (0, 1, 2, 8, 9, 0xA):702            raise FailedConnectionError(1002, "Unknown opcode received.")703        masked = b & 128 == 128704        if not masked and not self.client:705            raise FailedConnectionError(1002, "A client MUST mask all frames"706                                        " that it sends to the server")707        length = b & 127708        if opcode & 8:709            if not finished:710                raise FailedConnectionError(1002, "Control frames must not"711                                            " be fragmented.")712            if length > 125:713                raise FailedConnectionError(714                    1002,715                    "All control frames MUST have a payload length of 125"716                    " bytes or less")717        elif opcode and message:718            raise FailedConnectionError(719                1002,720                "Received a non-continuation opcode within"721                " fragmented message.")722        elif not opcode and not message:723            raise FailedConnectionError(724                1002,725                "Received continuation opcode with no previous"726                " fragments received.")727        if length == 126:728            length = struct.unpack('!H', recv(2))[0]729        elif length == 127:730            length = struct.unpack('!Q', recv(8))[0]731 732        if length > self.max_frame_length:733            raise FailedConnectionError(1009, "Incoming frame of {} bytes is above length limit of {} bytes.".format(734                length, self.max_frame_length))735        if masked:736            mask = struct.unpack('!BBBB', recv(4))737        received = 0738        if not message or opcode & 8:739            decoder = self.UTF8Decoder() if opcode == 1 else None740            decompressor = self._get_permessage_deflate_dec(rsv1)741            message = self.Message(opcode, self.max_frame_length, decoder=decoder, decompressor=decompressor)742        if not length:743            message.push(b'', final=finished)744        else:745            while received < length:746                d = self.socket.recv(length - received)747                if not d:748                    raise ConnectionClosedError()749                dlen = len(d)750                if masked:751                    d = self._apply_mask(d, mask, length=dlen, offset=received)752                received = received + dlen753                try:754                    message.push(d, final=finished)755                except (UnicodeDecodeError, ValueError):756                    raise FailedConnectionError(757                        1007, "Text data must be valid utf-8")758        return message759 760    def _pack_message(self, message, masked=False,761                      continuation=False, final=True, control_code=None):762        is_text = False763        if isinstance(message, six.text_type):764            message = message.encode('utf-8')765            is_text = True766 767        compress_bit = 0768        compressor = self._get_permessage_deflate_enc()769        # Control frames are identified by opcodes where the most significant770        # bit of the opcode is 1.  Currently defined opcodes for control frames771        # include 0x8 (Close), 0x9 (Ping), and 0xA (Pong).  Opcodes 0xB-0xF are772        # reserved for further control frames yet to be defined.773        # https://datatracker.ietf.org/doc/html/rfc6455#section-5.5774        is_control_frame = (control_code or 0) & 8775        # An endpoint MUST NOT set the "Per-Message Compressed" bit of control776        # frames and non-first fragments of a data message.  An endpoint777        # receiving such a frame MUST _Fail the WebSocket Connection_.778        # https://datatracker.ietf.org/doc/html/rfc7692#section-6.1779        if message and compressor and not is_control_frame:780            message = compressor.compress(message)781            message += compressor.flush(zlib.Z_SYNC_FLUSH)782            assert message[-4:] == b"\x00\x00\xff\xff"783            message = message[:-4]784            compress_bit = 1 << 6785 786        length = len(message)787        if not length:788            # no point masking empty data789            masked = False790        if control_code:791            if control_code not in (8, 9, 0xA):792                raise ProtocolError('Unknown control opcode.')793            if continuation or not final:794                raise ProtocolError('Control frame cannot be a fragment.')795            if length > 125:796                raise ProtocolError('Control frame data too large (>125).')797            header = struct.pack('!B', control_code | 1 << 7)798        else:799            opcode = 0 if continuation else ((1 if is_text else 2) | compress_bit)800            header = struct.pack('!B', opcode | (1 << 7 if final else 0))801        lengthdata = 1 << 7 if masked else 0802        if length > 65535:803            lengthdata = struct.pack('!BQ', lengthdata | 127, length)804        elif length > 125:805            lengthdata = struct.pack('!BH', lengthdata | 126, length)806        else:807            lengthdata = struct.pack('!B', lengthdata | length)808        if masked:809            # NOTE: RFC6455 states:810            # A server MUST NOT mask any frames that it sends to the client811            rand = Random(time.time())812            mask = [rand.getrandbits(8) for _ in six.moves.xrange(4)]813            message = RFC6455WebSocket._apply_mask(message, mask, length)814            maskdata = struct.pack('!BBBB', *mask)815        else:816            maskdata = b''817 818        return b''.join((header, lengthdata, maskdata, message))819 820    def wait(self):821        for i in self.iterator:822            return i823 824    def _send(self, frame):825        self._sendlock.acquire()826        try:827            self.socket.sendall(frame)828        finally:829            self._sendlock.release()830 831    def send(self, message, **kw):832        kw['masked'] = self.client833        payload = self._pack_message(message, **kw)834        self._send(payload)835 836    def _send_closing_frame(self, ignore_send_errors=False, close_data=None):837        if self.version in (8, 13) and not self.websocket_closed:838            if close_data is not None:839                status, msg = close_data840                if isinstance(msg, six.text_type):841                    msg = msg.encode('utf-8')842                data = struct.pack('!H', status) + msg843            else:844                data = ''845            try:846                self.send(data, control_code=8)847            except SocketError:848                # Sometimes, like when the remote side cuts off the connection,849                # we don't care about this.850                if not ignore_send_errors:  # pragma NO COVER851                    raise852            self.websocket_closed = True853 854    def close(self, close_data=None):855        """Forcibly close the websocket; generally it is preferable to856        return from the handler method."""857        try:858            self._send_closing_frame(close_data=close_data, ignore_send_errors=True)859            self.socket.shutdown(socket.SHUT_WR)860        except SocketError as e:861            if e.errno != errno.ENOTCONN:862                self.log.write('{ctx} socket shutdown error: {e}'.format(ctx=self.log_context, e=e))863        finally:864            self.socket.close()865 
codekingpro/portable-devtools · Team Ai