Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
handshake.py516 linesDownload Raw Back to wsproto
1"""2wsproto/handshake3~~~~~~~~~~~~~~~~~~4 5An implementation of WebSocket handshakes.6"""7from __future__ import annotations8 9from collections import deque10from typing import (11    TYPE_CHECKING,12    cast,13)14 15import h1116 17from .connection import Connection, ConnectionState, ConnectionType18from .events import AcceptConnection, Event, RejectConnection, RejectData, Request19from .extensions import Extension20from .utilities import (21    LocalProtocolError,22    RemoteProtocolError,23    generate_accept_token,24    generate_nonce,25    normed_header_dict,26    split_comma_header,27)28 29if TYPE_CHECKING:30    from collections.abc import Generator, Iterable, Sequence31 32    from .typing import Headers33 34# RFC6455, Section 4.2.1/6 - Reading the Client's Opening Handshake35WEBSOCKET_VERSION = b"13"36 37# RFC6455, Section 4.2.1/3 - Value of the Upgrade header38WEBSOCKET_UPGRADE = b"websocket"39 40 41class H11Handshake:42    """A Handshake implementation for HTTP/1.1 connections."""43 44    def __init__(self, connection_type: ConnectionType) -> None:45        self.client = connection_type is ConnectionType.CLIENT46        self._state = ConnectionState.CONNECTING47 48        if self.client:49            self._h11_connection = h11.Connection(h11.CLIENT)50        else:51            self._h11_connection = h11.Connection(h11.SERVER)52 53        self._connection: Connection | None = None54        self._events: deque[Event] = deque()55        self._initiating_request: Request | None = None56        self._nonce: bytes | None = None57 58    @property59    def state(self) -> ConnectionState:60        return self._state61 62    @property63    def connection(self) -> Connection | None:64        """65        Return the established connection.66 67        This will either return the connection or raise a68        LocalProtocolError if the connection has not yet been69        established.70 71        :rtype: h11.Connection72        """73        return self._connection74 75    def initiate_upgrade_connection(76        self, headers: Headers, path: bytes | str,77    ) -> None:78        """79        Initiate an upgrade connection.80 81        This should be used if the request has already be received and82        parsed.83 84        :param list headers: HTTP headers represented as a list of 2-tuples.85        :param str path: A URL path.86        """87        if self.client:88            msg = "Cannot initiate an upgrade connection when acting as the client"89            raise LocalProtocolError(90                msg,91            )92        upgrade_request = h11.Request(method=b"GET", target=path, headers=headers)93        h11_client = h11.Connection(h11.CLIENT)94        self.receive_data(h11_client.send(upgrade_request))95 96    def send(self, event: Event) -> bytes:97        """98        Send an event to the remote.99 100        This will return the bytes to send based on the event or raise101        a LocalProtocolError if the event is not valid given the102        state.103 104        :returns: Data to send to the WebSocket peer.105        :rtype: bytes106        """107        data = b""108        if isinstance(event, Request):109            data += self._initiate_connection(event)110        elif isinstance(event, AcceptConnection):111            data += self._accept(event)112        elif isinstance(event, RejectConnection):113            data += self._reject(event)114        elif isinstance(event, RejectData):115            data += self._send_reject_data(event)116        else:117            msg = f"Event {event} cannot be sent during the handshake"118            raise LocalProtocolError(119                msg,120            )121        return data122 123    def receive_data(self, data: bytes | None) -> None:124        """125        Receive data from the remote.126 127        A list of events that the remote peer triggered by sending128        this data can be retrieved with :meth:`events`.129 130        :param bytes data: Data received from the WebSocket peer.131        """132        self._h11_connection.receive_data(data or b"")133        while True:134            try:135                event = self._h11_connection.next_event()136            except h11.RemoteProtocolError:137                msg = "Bad HTTP message"138                raise RemoteProtocolError(139                    msg, event_hint=RejectConnection(),140                )141            if (142                isinstance(event, h11.ConnectionClosed)143                or event is h11.NEED_DATA144                or event is h11.PAUSED145            ):146                break147 148            if self.client:149                if isinstance(event, h11.InformationalResponse):150                    if event.status_code == 101:151                        self._events.append(self._establish_client_connection(event))152                    else:153                        self._events.append(154                            RejectConnection(155                                headers=list(event.headers),156                                status_code=event.status_code,157                                has_body=False,158                            ),159                        )160                        self._state = ConnectionState.CLOSED161                elif isinstance(event, h11.Response):162                    self._state = ConnectionState.REJECTING163                    self._events.append(164                        RejectConnection(165                            headers=list(event.headers),166                            status_code=event.status_code,167                            has_body=True,168                        ),169                    )170                elif isinstance(event, h11.Data):171                    self._events.append(172                        RejectData(data=event.data, body_finished=False),173                    )174                elif isinstance(event, h11.EndOfMessage):175                    self._events.append(RejectData(data=b"", body_finished=True))176                    self._state = ConnectionState.CLOSED177            elif isinstance(event, h11.Request):178                self._events.append(self._process_connection_request(event))179 180    def events(self) -> Generator[Event, None, None]:181        """182        Return a generator that provides any events that have been generated183        by protocol activity.184 185        :returns: a generator that yields H11 events.186        """187        while self._events:188            yield self._events.popleft()189 190    # Server mode methods191 192    def _process_connection_request(193        self, event: h11.Request,194    ) -> Request:195        if event.method != b"GET":196            msg = "Request method must be GET"197            raise RemoteProtocolError(198                msg, event_hint=RejectConnection(),199            )200        connection_tokens = None201        extensions: list[str] = []202        host = None203        key = None204        subprotocols: list[str] = []205        upgrade = b""206        version = None207        headers: Headers = []208        for name, value in event.headers:209            name = name.lower()210            if name == b"connection":211                connection_tokens = split_comma_header(value)212            elif name == b"host":213                host = value.decode("idna")214                continue  # Skip appending to headers215            elif name == b"sec-websocket-extensions":216                extensions.extend(split_comma_header(value))217                continue  # Skip appending to headers218            elif name == b"sec-websocket-key":219                key = value220            elif name == b"sec-websocket-protocol":221                subprotocols.extend(split_comma_header(value))222                continue  # Skip appending to headers223            elif name == b"sec-websocket-version":224                version = value225            elif name == b"upgrade":226                upgrade = value227            headers.append((name, value))228        if connection_tokens is None or not any(229            token.lower() == "upgrade" for token in connection_tokens230        ):231            msg = "Missing header, 'Connection: Upgrade'"232            raise RemoteProtocolError(233                msg, event_hint=RejectConnection(),234            )235        if version != WEBSOCKET_VERSION:236            msg = "Missing header, 'Sec-WebSocket-Version'"237            raise RemoteProtocolError(238                msg,239                event_hint=RejectConnection(240                    headers=[(b"Sec-WebSocket-Version", WEBSOCKET_VERSION)],241                    status_code=426 if version else 400,242                ),243            )244        if key is None:245            msg = "Missing header, 'Sec-WebSocket-Key'"246            raise RemoteProtocolError(247                msg, event_hint=RejectConnection(),248            )249        if upgrade.lower() != WEBSOCKET_UPGRADE:250            msg = f"Missing header, 'Upgrade: {WEBSOCKET_UPGRADE.decode()}'"251            raise RemoteProtocolError(252                msg,253                event_hint=RejectConnection(),254            )255        if host is None:256            msg = "Missing header, 'Host'"257            raise RemoteProtocolError(258                msg, event_hint=RejectConnection(),259            )260 261        self._initiating_request = Request(262            extensions=extensions,263            extra_headers=headers,264            host=host,265            subprotocols=subprotocols,266            target=event.target.decode("ascii"),267        )268        return self._initiating_request269 270    def _accept(self, event: AcceptConnection) -> bytes:271        # _accept is always called after _process_connection_request.272        assert self._initiating_request is not None273        request_headers = normed_header_dict(self._initiating_request.extra_headers)274 275        nonce = request_headers[b"sec-websocket-key"]276        accept_token = generate_accept_token(nonce)277 278        headers = [279            (b"Upgrade", WEBSOCKET_UPGRADE),280            (b"Connection", b"Upgrade"),281            (b"Sec-WebSocket-Accept", accept_token),282        ]283 284        if event.subprotocol is not None:285            if event.subprotocol not in self._initiating_request.subprotocols:286                msg = f"unexpected subprotocol {event.subprotocol}"287                raise LocalProtocolError(msg)288            headers.append(289                (b"Sec-WebSocket-Protocol", event.subprotocol.encode("ascii")),290            )291 292        if event.extensions:293            accepts = server_extensions_handshake(294                cast("Sequence[str]", self._initiating_request.extensions),295                event.extensions,296            )297            if accepts:298                headers.append((b"Sec-WebSocket-Extensions", accepts))299 300        response = h11.InformationalResponse(301            status_code=101,302            headers=headers + event.extra_headers,303            reason=b"Switching Protocols",304        )305        self._connection = Connection(306            ConnectionType.CLIENT if self.client else ConnectionType.SERVER,307            event.extensions,308        )309        self._state = ConnectionState.OPEN310        return self._h11_connection.send(response) or b""311 312    def _reject(self, event: RejectConnection) -> bytes:313        if self.state != ConnectionState.CONNECTING:314            msg = f"Connection cannot be rejected in state {self.state}"315            raise LocalProtocolError(316                msg,317            )318 319        headers = list(event.headers)320        if not event.has_body:321            headers.append((b"content-length", b"0"))322        response = h11.Response(status_code=event.status_code, headers=headers)323        data = self._h11_connection.send(response) or b""324        self._state = ConnectionState.REJECTING325        if not event.has_body:326            data += self._h11_connection.send(h11.EndOfMessage()) or b""327            self._state = ConnectionState.CLOSED328        return data329 330    def _send_reject_data(self, event: RejectData) -> bytes:331        if self.state != ConnectionState.REJECTING:332            msg = f"Cannot send rejection data in state {self.state}"333            raise LocalProtocolError(334                msg,335            )336 337        data = self._h11_connection.send(h11.Data(data=event.data)) or b""338        if event.body_finished:339            data += self._h11_connection.send(h11.EndOfMessage()) or b""340            self._state = ConnectionState.CLOSED341        return data342 343    # Client mode methods344 345    def _initiate_connection(self, request: Request) -> bytes:346        self._initiating_request = request347        self._nonce = generate_nonce()348 349        headers = [350            (b"Host", request.host.encode("idna")),351            (b"Upgrade", WEBSOCKET_UPGRADE),352            (b"Connection", b"Upgrade"),353            (b"Sec-WebSocket-Key", self._nonce),354            (b"Sec-WebSocket-Version", WEBSOCKET_VERSION),355        ]356 357        if request.subprotocols:358            headers.append(359                (360                    b"Sec-WebSocket-Protocol",361                    (", ".join(request.subprotocols)).encode("ascii"),362                ),363            )364 365        if request.extensions:366            offers: dict[str, str | bool] = {}367            for e in request.extensions:368                assert isinstance(e, Extension)369                offers[e.name] = e.offer()370            extensions = []371            for name, params in offers.items():372                bname = name.encode("ascii")373                if isinstance(params, bool):374                    if params:375                        extensions.append(bname)376                else:377                    extensions.append(b"%s; %s" % (bname, params.encode("ascii")))378            if extensions:379                headers.append((b"Sec-WebSocket-Extensions", b", ".join(extensions)))380 381        upgrade = h11.Request(382            method=b"GET",383            target=request.target.encode("ascii"),384            headers=headers + request.extra_headers,385        )386        return self._h11_connection.send(upgrade) or b""387 388    def _establish_client_connection(389        self, event: h11.InformationalResponse,390    ) -> AcceptConnection:391        # _establish_client_connection is always called after _initiate_connection.392        assert self._initiating_request is not None393        assert self._nonce is not None394 395        accept = None396        connection_tokens = None397        accepts: list[str] = []398        subprotocol = None399        upgrade = b""400        headers: Headers = []401        for name, value in event.headers:402            name = name.lower()403            if name == b"connection":404                connection_tokens = split_comma_header(value)405                continue  # Skip appending to headers406            if name == b"sec-websocket-extensions":407                accepts = split_comma_header(value)408                continue  # Skip appending to headers409            if name == b"sec-websocket-accept":410                accept = value411                continue  # Skip appending to headers412            if name == b"sec-websocket-protocol":413                subprotocol = value.decode("ascii")414                continue  # Skip appending to headers415            if name == b"upgrade":416                upgrade = value417                continue  # Skip appending to headers418            headers.append((name, value))419 420        if connection_tokens is None or not any(421            token.lower() == "upgrade" for token in connection_tokens422        ):423            msg = "Missing header, 'Connection: Upgrade'"424            raise RemoteProtocolError(425                msg, event_hint=RejectConnection(),426            )427        if upgrade.lower() != WEBSOCKET_UPGRADE:428            msg = f"Missing header, 'Upgrade: {WEBSOCKET_UPGRADE.decode()}'"429            raise RemoteProtocolError(430                msg,431                event_hint=RejectConnection(),432            )433        accept_token = generate_accept_token(self._nonce)434        if accept != accept_token:435            msg = "Bad accept token"436            raise RemoteProtocolError(msg, event_hint=RejectConnection())437        if subprotocol is not None and subprotocol not in self._initiating_request.subprotocols:438            msg = f"unrecognized subprotocol {subprotocol}"439            raise RemoteProtocolError(440                msg,441                event_hint=RejectConnection(),442            )443        extensions = client_extensions_handshake(444            accepts, cast("Sequence[Extension]", self._initiating_request.extensions),445        )446 447        self._connection = Connection(448            ConnectionType.CLIENT if self.client else ConnectionType.SERVER,449            extensions,450            self._h11_connection.trailing_data[0],451        )452        self._state = ConnectionState.OPEN453        return AcceptConnection(454            extensions=extensions, extra_headers=headers, subprotocol=subprotocol,455        )456 457    def __repr__(self) -> str:458        return f"{self.__class__.__name__}(client={self.client}, state={self.state})"459 460 461def server_extensions_handshake(462    requested: Iterable[str], supported: list[Extension],463) -> bytes | None:464    """465    Agree on the extensions to use returning an appropriate header value.466 467    This returns None if there are no agreed extensions468    """469    accepts: dict[str, bool | bytes] = {}470    for offer in requested:471        name = offer.split(";", 1)[0].strip()472        for extension in supported:473            if extension.name == name:474                accept = extension.accept(offer)475                if isinstance(accept, bool):476                    if accept:477                        accepts[extension.name] = True478                elif accept is not None:479                    accepts[extension.name] = accept.encode("ascii")480 481    if accepts:482        extensions: list[bytes] = []483        for name, params in accepts.items():484            name_bytes = name.encode("ascii")485            if isinstance(params, bool):486                assert params487                extensions.append(name_bytes)488            elif params == b"":489                extensions.append(b"%s" % (name_bytes))490            else:491                extensions.append(b"%s; %s" % (name_bytes, params))492        return b", ".join(extensions)493 494    return None495 496 497def client_extensions_handshake(498    accepted: Iterable[str], supported: Sequence[Extension],499) -> list[Extension]:500    # This raises RemoteProtocolError is the accepted extension is not501    # supported.502    extensions = []503    for accept in accepted:504        name = accept.split(";", 1)[0].strip()505        for extension in supported:506            if extension.name == name:507                extension.finalize(accept)508                extensions.append(extension)509                break510        else:511            msg = f"unrecognized extension {name}"512            raise RemoteProtocolError(513                msg, event_hint=RejectConnection(),514            )515    return extensions516 
codekingpro/portable-devtools · Team Ai