Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
server.py590 linesDownload Raw Back to websockets
1from __future__ import annotations
2
3import base64
4import binascii
5import email.utils
6import http
7import re
8import warnings
9from collections.abc import Generator, Sequence
10from typing import Any, Callable, cast
11
12from .datastructures import Headers, MultipleValuesError
13from .exceptions import (
14    InvalidHandshake,
15    InvalidHeader,
16    InvalidHeaderValue,
17    InvalidMessage,
18    InvalidOrigin,
19    InvalidUpgrade,
20    NegotiationError,
21)
22from .extensions import Extension, ServerExtensionFactory
23from .headers import (
24    build_extension,
25    parse_connection,
26    parse_extension,
27    parse_subprotocol,
28    parse_upgrade,
29)
30from .http11 import Request, Response
31from .imports import lazy_import
32from .protocol import CONNECTING, OPEN, SERVER, Protocol, State
33from .typing import (
34    ConnectionOption,
35    ExtensionHeader,
36    LoggerLike,
37    Origin,
38    StatusLike,
39    Subprotocol,
40    UpgradeProtocol,
41)
42from .utils import accept_key
43
44
45__all__ = ["ServerProtocol"]
46
47
48class ServerProtocol(Protocol):
49    """
50    Sans-I/O implementation of a WebSocket server connection.
51
52    Args:
53        origins: Acceptable values of the ``Origin`` header. Values can be
54            :class:`str` to test for an exact match or regular expressions
55            compiled by :func:`re.compile` to test against a pattern. Include
56            :obj:`None` in the list if the lack of an origin is acceptable.
57            This is useful for defending against Cross-Site WebSocket
58            Hijacking attacks.
59        extensions: List of supported extensions, in order in which they
60            should be tried.
61        subprotocols: List of supported subprotocols, in order of decreasing
62            preference.
63        select_subprotocol: Callback for selecting a subprotocol among
64            those supported by the client and the server. It has the same
65            signature as the :meth:`select_subprotocol` method, including a
66            :class:`ServerProtocol` instance as first argument.
67        state: Initial state of the WebSocket connection.
68        max_size: Maximum size of incoming messages in bytes.
69            :obj:`None` disables the limit. You may pass a ``(max_message_size,
70            max_fragment_size)`` tuple to set different limits for messages and
71            fragments when you expect long messages sent in short fragments.
72        logger: Logger for this connection;
73            defaults to ``logging.getLogger("websockets.server")``;
74            see the :doc:`logging guide <../../topics/logging>` for details.
75
76    """
77
78    def __init__(
79        self,
80        *,
81        origins: Sequence[Origin | re.Pattern[str] | None] | None = None,
82        extensions: Sequence[ServerExtensionFactory] | None = None,
83        subprotocols: Sequence[Subprotocol] | None = None,
84        select_subprotocol: (
85            Callable[
86                [ServerProtocol, Sequence[Subprotocol]],
87                Subprotocol | None,
88            ]
89            | None
90        ) = None,
91        state: State = CONNECTING,
92        max_size: int | None | tuple[int | None, int | None] = 2**20,
93        logger: LoggerLike | None = None,
94    ) -> None:
95        super().__init__(
96            side=SERVER,
97            state=state,
98            max_size=max_size,
99            logger=logger,
100        )
101        self.origins = origins
102        self.available_extensions = extensions
103        self.available_subprotocols = subprotocols
104        if select_subprotocol is not None:
105            # Bind select_subprotocol then shadow self.select_subprotocol.
106            # Use setattr to work around https://github.com/python/mypy/issues/2427.
107            setattr(
108                self,
109                "select_subprotocol",
110                select_subprotocol.__get__(self, self.__class__),
111            )
112
113    def accept(self, request: Request) -> Response:
114        """
115        Create a handshake response to accept the connection.
116
117        If the handshake request is valid and the handshake successful,
118        :meth:`accept` returns an HTTP response with status code 101.
119
120        Else, it returns an HTTP response with another status code. This rejects
121        the connection, like :meth:`reject` would.
122
123        You must send the handshake response with :meth:`send_response`.
124
125        You may modify the response before sending it, typically by adding HTTP
126        headers.
127
128        Args:
129            request: WebSocket handshake request received from the client.
130
131        Returns:
132            WebSocket handshake response or HTTP response to send to the client.
133
134        """
135        try:
136            (
137                accept_header,
138                extensions_header,
139                protocol_header,
140            ) = self.process_request(request)
141        except InvalidOrigin as exc:
142            request._exception = exc
143            self.handshake_exc = exc
144            if self.debug:
145                self.logger.debug("! invalid origin", exc_info=True)
146            return self.reject(
147                http.HTTPStatus.FORBIDDEN,
148                f"Failed to open a WebSocket connection: {exc}.\n",
149            )
150        except InvalidUpgrade as exc:
151            request._exception = exc
152            self.handshake_exc = exc
153            if self.debug:
154                self.logger.debug("! invalid upgrade", exc_info=True)
155            response = self.reject(
156                http.HTTPStatus.UPGRADE_REQUIRED,
157                (
158                    f"Failed to open a WebSocket connection: {exc}.\n"
159                    f"\n"
160                    f"You cannot access a WebSocket server directly "
161                    f"with a browser. You need a WebSocket client.\n"
162                ),
163            )
164            response.headers["Upgrade"] = "websocket"
165            return response
166        except InvalidHandshake as exc:
167            request._exception = exc
168            self.handshake_exc = exc
169            if self.debug:
170                self.logger.debug("! invalid handshake", exc_info=True)
171            exc_chain = cast(BaseException, exc)
172            exc_str = f"{exc_chain}"
173            while exc_chain.__cause__ is not None:
174                exc_chain = exc_chain.__cause__
175                exc_str += f"; {exc_chain}"
176            return self.reject(
177                http.HTTPStatus.BAD_REQUEST,
178                f"Failed to open a WebSocket connection: {exc_str}.\n",
179            )
180        except Exception as exc:
181            # Handle exceptions raised by user-provided select_subprotocol and
182            # unexpected errors.
183            request._exception = exc
184            self.handshake_exc = exc
185            self.logger.error("opening handshake failed", exc_info=True)
186            return self.reject(
187                http.HTTPStatus.INTERNAL_SERVER_ERROR,
188                (
189                    "Failed to open a WebSocket connection.\n"
190                    "See server log for more information.\n"
191                ),
192            )
193
194        headers = Headers()
195        headers["Date"] = email.utils.formatdate(usegmt=True)
196        headers["Upgrade"] = "websocket"
197        headers["Connection"] = "Upgrade"
198        headers["Sec-WebSocket-Accept"] = accept_header
199        if extensions_header is not None:
200            headers["Sec-WebSocket-Extensions"] = extensions_header
201        if protocol_header is not None:
202            headers["Sec-WebSocket-Protocol"] = protocol_header
203        return Response(101, "Switching Protocols", headers)
204
205    def process_request(
206        self,
207        request: Request,
208    ) -> tuple[str, str | None, str | None]:
209        """
210        Check a handshake request and negotiate extensions and subprotocol.
211
212        This function doesn't verify that the request is an HTTP/1.1 or higher
213        GET request and doesn't check the ``Host`` header. These controls are
214        usually performed earlier in the HTTP request handling code. They're
215        the responsibility of the caller.
216
217        Args:
218            request: WebSocket handshake request received from the client.
219
220        Returns:
221            ``Sec-WebSocket-Accept``, ``Sec-WebSocket-Extensions``, and
222            ``Sec-WebSocket-Protocol`` headers for the handshake response.
223
224        Raises:
225            InvalidHandshake: If the handshake request is invalid;
226                then the server must return 400 Bad Request error.
227
228        """
229        headers = request.headers
230
231        connection: list[ConnectionOption] = sum(
232            [parse_connection(value) for value in headers.get_all("Connection")], []
233        )
234        if not any(value.lower() == "upgrade" for value in connection):
235            raise InvalidUpgrade(
236                "Connection", ", ".join(connection) if connection else None
237            )
238
239        upgrade: list[UpgradeProtocol] = sum(
240            [parse_upgrade(value) for value in headers.get_all("Upgrade")], []
241        )
242        # For compatibility with non-strict implementations, ignore case when
243        # checking the Upgrade header. The RFC always uses "websocket", except
244        # in section 11.2. (IANA registration) where it uses "WebSocket".
245        if not (len(upgrade) == 1 and upgrade[0].lower() == "websocket"):
246            raise InvalidUpgrade("Upgrade", ", ".join(upgrade) if upgrade else None)
247
248        try:
249            key = headers["Sec-WebSocket-Key"]
250        except KeyError:
251            raise InvalidHeader("Sec-WebSocket-Key") from None
252        except MultipleValuesError:
253            raise InvalidHeader("Sec-WebSocket-Key", "multiple values") from None
254        try:
255            raw_key = base64.b64decode(key.encode(), validate=True)
256        except binascii.Error as exc:
257            raise InvalidHeaderValue("Sec-WebSocket-Key", key) from exc
258        if len(raw_key) != 16:
259            raise InvalidHeaderValue("Sec-WebSocket-Key", key)
260        accept_header = accept_key(key)
261
262        try:
263            version = headers["Sec-WebSocket-Version"]
264        except KeyError:
265            raise InvalidHeader("Sec-WebSocket-Version") from None
266        except MultipleValuesError:
267            raise InvalidHeader("Sec-WebSocket-Version", "multiple values") from None
268        if version != "13":
269            raise InvalidHeaderValue("Sec-WebSocket-Version", version)
270
271        self.origin = self.process_origin(headers)
272        extensions_header, self.extensions = self.process_extensions(headers)
273        protocol_header = self.subprotocol = self.process_subprotocol(headers)
274
275        return (accept_header, extensions_header, protocol_header)
276
277    def process_origin(self, headers: Headers) -> Origin | None:
278        """
279        Handle the Origin HTTP request header.
280
281        Args:
282            headers: WebSocket handshake request headers.
283
284        Returns:
285           origin, if it is acceptable.
286
287        Raises:
288            InvalidHandshake: If the Origin header is invalid.
289            InvalidOrigin: If the origin isn't acceptable.
290
291        """
292        # "The user agent MUST NOT include more than one Origin header field"
293        # per https://datatracker.ietf.org/doc/html/rfc6454#section-7.3.
294        try:
295            origin = headers.get("Origin")
296        except MultipleValuesError:
297            raise InvalidHeader("Origin", "multiple values") from None
298        if origin is not None:
299            origin = cast(Origin, origin)
300        if self.origins is not None:
301            for origin_or_regex in self.origins:
302                if origin_or_regex == origin or (
303                    isinstance(origin_or_regex, re.Pattern)
304                    and origin is not None
305                    and origin_or_regex.fullmatch(origin) is not None
306                ):
307                    break
308            else:
309                raise InvalidOrigin(origin)
310        return origin
311
312    def process_extensions(
313        self,
314        headers: Headers,
315    ) -> tuple[str | None, list[Extension]]:
316        """
317        Handle the Sec-WebSocket-Extensions HTTP request header.
318
319        Accept or reject each extension proposed in the client request.
320        Negotiate parameters for accepted extensions.
321
322        Per :rfc:`6455`, negotiation rules are defined by the specification of
323        each extension.
324
325        To provide this level of flexibility, for each extension proposed by
326        the client, we check for a match with each extension available in the
327        server configuration. If no match is found, the extension is ignored.
328
329        If several variants of the same extension are proposed by the client,
330        it may be accepted several times, which won't make sense in general.
331        Extensions must implement their own requirements. For this purpose,
332        the list of previously accepted extensions is provided.
333
334        This process doesn't allow the server to reorder extensions. It can
335        only select a subset of the extensions proposed by the client.
336
337        Other requirements, for example related to mandatory extensions or the
338        order of extensions, may be implemented by overriding this method.
339
340        Args:
341            headers: WebSocket handshake request headers.
342
343        Returns:
344            ``Sec-WebSocket-Extensions`` HTTP response header and list of
345            accepted extensions.
346
347        Raises:
348            InvalidHandshake: If the Sec-WebSocket-Extensions header is invalid.
349
350        """
351        response_header_value: str | None = None
352
353        extension_headers: list[ExtensionHeader] = []
354        accepted_extensions: list[Extension] = []
355
356        header_values = headers.get_all("Sec-WebSocket-Extensions")
357
358        if header_values and self.available_extensions:
359            parsed_header_values: list[ExtensionHeader] = sum(
360                [parse_extension(header_value) for header_value in header_values], []
361            )
362
363            for name, request_params in parsed_header_values:
364                for ext_factory in self.available_extensions:
365                    # Skip non-matching extensions based on their name.
366                    if ext_factory.name != name:
367                        continue
368
369                    # Skip non-matching extensions based on their params.
370                    try:
371                        response_params, extension = ext_factory.process_request_params(
372                            request_params, accepted_extensions
373                        )
374                    except NegotiationError:
375                        continue
376
377                    # Add matching extension to the final list.
378                    extension_headers.append((name, response_params))
379                    accepted_extensions.append(extension)
380
381                    # Break out of the loop once we have a match.
382                    break
383
384                # If we didn't break from the loop, no extension in our list
385                # matched what the client sent. The extension is declined.
386
387        # Serialize extension header.
388        if extension_headers:
389            response_header_value = build_extension(extension_headers)
390
391        return response_header_value, accepted_extensions
392
393    def process_subprotocol(self, headers: Headers) -> Subprotocol | None:
394        """
395        Handle the Sec-WebSocket-Protocol HTTP request header.
396
397        Args:
398            headers: WebSocket handshake request headers.
399
400        Returns:
401           Subprotocol, if one was selected; this is also the value of the
402           ``Sec-WebSocket-Protocol`` response header.
403
404        Raises:
405            InvalidHandshake: If the Sec-WebSocket-Subprotocol header is invalid.
406
407        """
408        subprotocols: Sequence[Subprotocol] = sum(
409            [
410                parse_subprotocol(header_value)
411                for header_value in headers.get_all("Sec-WebSocket-Protocol")
412            ],
413            [],
414        )
415        return self.select_subprotocol(subprotocols)
416
417    def select_subprotocol(
418        self,
419        subprotocols: Sequence[Subprotocol],
420    ) -> Subprotocol | None:
421        """
422        Pick a subprotocol among those offered by the client.
423
424        If several subprotocols are supported by both the client and the server,
425        pick the first one in the list declared the server.
426
427        If the server doesn't support any subprotocols, continue without a
428        subprotocol, regardless of what the client offers.
429
430        If the server supports at least one subprotocol and the client doesn't
431        offer any, abort the handshake with an HTTP 400 error.
432
433        You provide a ``select_subprotocol`` argument to :class:`ServerProtocol`
434        to override this logic. For example, you could accept the connection
435        even if client doesn't offer a subprotocol, rather than reject it.
436
437        Here's how to negotiate the ``chat`` subprotocol if the client supports
438        it and continue without a subprotocol otherwise::
439
440            def select_subprotocol(protocol, subprotocols):
441                if "chat" in subprotocols:
442                    return "chat"
443
444        Args:
445            subprotocols: List of subprotocols offered by the client.
446
447        Returns:
448            Selected subprotocol, if a common subprotocol was found.
449
450            :obj:`None` to continue without a subprotocol.
451
452        Raises:
453            NegotiationError: Custom implementations may raise this exception
454                to abort the handshake with an HTTP 400 error.
455
456        """
457        # Server doesn't offer any subprotocols.
458        if not self.available_subprotocols:  # None or empty list
459            return None
460
461        # Server offers at least one subprotocol but client doesn't offer any.
462        if not subprotocols:
463            raise NegotiationError("missing subprotocol")
464
465        # Server and client both offer subprotocols. Look for a shared one.
466        proposed_subprotocols = set(subprotocols)
467        for subprotocol in self.available_subprotocols:
468            if subprotocol in proposed_subprotocols:
469                return subprotocol
470
471        # No common subprotocol was found.
472        raise NegotiationError(
473            "invalid subprotocol; expected one of "
474            + ", ".join(self.available_subprotocols)
475        )
476
477    def reject(self, status: StatusLike, text: str) -> Response:
478        """
479        Create a handshake response to reject the connection.
480
481        A short plain text response is the best fallback when failing to
482        establish a WebSocket connection.
483
484        You must send the handshake response with :meth:`send_response`.
485
486        You may modify the response before sending it, for example by changing
487        HTTP headers.
488
489        Args:
490            status: HTTP status code.
491            text: HTTP response body; it will be encoded to UTF-8.
492
493        Returns:
494            HTTP response to send to the client.
495
496        """
497        # If status is an int instead of an HTTPStatus, fix it automatically.
498        status = http.HTTPStatus(status)
499        body = text.encode()
500        headers = Headers(
501            [
502                ("Date", email.utils.formatdate(usegmt=True)),
503                ("Connection", "close"),
504                ("Content-Length", str(len(body))),
505                ("Content-Type", "text/plain; charset=utf-8"),
506            ]
507        )
508        return Response(status.value, status.phrase, headers, body)
509
510    def send_response(self, response: Response) -> None:
511        """
512        Send a handshake response to the client.
513
514        Args:
515            response: WebSocket handshake response event to send.
516
517        """
518        if self.debug:
519            code, phrase = response.status_code, response.reason_phrase
520            self.logger.debug("> HTTP/1.1 %d %s", code, phrase)
521            for key, value in response.headers.raw_items():
522                self.logger.debug("> %s: %s", key, value)
523            if response.body:
524                self.logger.debug("> [body] (%d bytes)", len(response.body))
525
526        self.writes.append(response.serialize())
527
528        if response.status_code == 101:
529            assert self.state is CONNECTING
530            self.state = OPEN
531            self.logger.info("connection open")
532
533        else:
534            self.logger.info(
535                "connection rejected (%d %s)",
536                response.status_code,
537                response.reason_phrase,
538            )
539
540            self.send_eof()
541            self.parser = self.discard()
542            next(self.parser)  # start coroutine
543
544    def parse(self) -> Generator[None]:
545        if self.state is CONNECTING:
546            try:
547                request = yield from Request.parse(
548                    self.reader.read_line,
549                )
550            except Exception as exc:
551                self.handshake_exc = InvalidMessage(
552                    "did not receive a valid HTTP request"
553                )
554                self.handshake_exc.__cause__ = exc
555                self.send_eof()
556                self.parser = self.discard()
557                next(self.parser)  # start coroutine
558                yield
559
560            if self.debug:
561                self.logger.debug("< GET %s HTTP/1.1", request.path)
562                for key, value in request.headers.raw_items():
563                    self.logger.debug("< %s: %s", key, value)
564
565            self.events.append(request)
566
567        yield from super().parse()
568
569
570class ServerConnection(ServerProtocol):
571    def __init__(self, *args: Any, **kwargs: Any) -> None:
572        warnings.warn(  # deprecated in 11.0 - 2023-04-02
573            "ServerConnection was renamed to ServerProtocol",
574            DeprecationWarning,
575        )
576        super().__init__(*args, **kwargs)
577
578
579lazy_import(
580    globals(),
581    deprecated_aliases={
582        # deprecated in 14.0 - 2024-11-09
583        "WebSocketServer": ".legacy.server",
584        "WebSocketServerProtocol": ".legacy.server",
585        "broadcast": ".legacy.server",
586        "serve": ".legacy.server",
587        "unix_serve": ".legacy.server",
588    },
589)
590 
codekingpro/portable-devtools · Team Ai