Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
server.py766 linesDownload Raw Back to sync
1from __future__ import annotations
2
3import hmac
4import http
5import logging
6import os
7import re
8import selectors
9import socket
10import ssl as ssl_module
11import sys
12import threading
13import warnings
14from collections.abc import Iterable, Sequence
15from types import TracebackType
16from typing import Any, Callable, Mapping, cast
17
18from ..exceptions import InvalidHeader
19from ..extensions.base import ServerExtensionFactory
20from ..extensions.permessage_deflate import enable_server_permessage_deflate
21from ..frames import CloseCode
22from ..headers import (
23    build_www_authenticate_basic,
24    parse_authorization_basic,
25    validate_subprotocols,
26)
27from ..http11 import SERVER, Request, Response
28from ..protocol import CONNECTING, OPEN, Event
29from ..server import ServerProtocol
30from ..typing import LoggerLike, Origin, StatusLike, Subprotocol
31from .connection import Connection
32from .utils import Deadline
33
34
35__all__ = ["serve", "unix_serve", "ServerConnection", "Server", "basic_auth"]
36
37
38class ServerConnection(Connection):
39    """
40    :mod:`threading` implementation of a WebSocket server connection.
41
42    :class:`ServerConnection` provides :meth:`recv` and :meth:`send` methods for
43    receiving and sending messages.
44
45    It supports iteration to receive messages::
46
47        for message in websocket:
48            process(message)
49
50    The iterator exits normally when the connection is closed with code
51    1000 (OK) or 1001 (going away) or without a close code. It raises a
52    :exc:`~websockets.exceptions.ConnectionClosedError` when the connection is
53    closed with any other code.
54
55    The ``ping_interval``, ``ping_timeout``, ``close_timeout``, and
56    ``max_queue`` arguments have the same meaning as in :func:`serve`.
57
58    Args:
59        socket: Socket connected to a WebSocket client.
60        protocol: Sans-I/O connection.
61
62    """
63
64    def __init__(
65        self,
66        socket: socket.socket,
67        protocol: ServerProtocol,
68        *,
69        ping_interval: float | None = 20,
70        ping_timeout: float | None = 20,
71        close_timeout: float | None = 10,
72        max_queue: int | None | tuple[int | None, int | None] = 16,
73    ) -> None:
74        self.protocol: ServerProtocol
75        self.request_rcvd = threading.Event()
76        super().__init__(
77            socket,
78            protocol,
79            ping_interval=ping_interval,
80            ping_timeout=ping_timeout,
81            close_timeout=close_timeout,
82            max_queue=max_queue,
83        )
84        self.username: str  # see basic_auth()
85        self.handler: Callable[[ServerConnection], None]  # see route()
86        self.handler_kwargs: Mapping[str, Any]  # see route()
87
88    def respond(self, status: StatusLike, text: str) -> Response:
89        """
90        Create a plain text HTTP response.
91
92        ``process_request`` and ``process_response`` may call this method to
93        return an HTTP response instead of performing the WebSocket opening
94        handshake.
95
96        You can modify the response before returning it, for example by changing
97        HTTP headers.
98
99        Args:
100            status: HTTP status code.
101            text: HTTP response body; it will be encoded to UTF-8.
102
103        Returns:
104            HTTP response to send to the client.
105
106        """
107        return self.protocol.reject(status, text)
108
109    def handshake(
110        self,
111        process_request: (
112            Callable[
113                [ServerConnection, Request],
114                Response | None,
115            ]
116            | None
117        ) = None,
118        process_response: (
119            Callable[
120                [ServerConnection, Request, Response],
121                Response | None,
122            ]
123            | None
124        ) = None,
125        server_header: str | None = SERVER,
126        timeout: float | None = None,
127    ) -> None:
128        """
129        Perform the opening handshake.
130
131        """
132        if not self.request_rcvd.wait(timeout):
133            raise TimeoutError("timed out while waiting for handshake request")
134
135        if self.request is not None:
136            with self.send_context(expected_state=CONNECTING):
137                response = None
138
139                if process_request is not None:
140                    try:
141                        response = process_request(self, self.request)
142                    except Exception as exc:
143                        self.protocol.handshake_exc = exc
144                        response = self.protocol.reject(
145                            http.HTTPStatus.INTERNAL_SERVER_ERROR,
146                            (
147                                "Failed to open a WebSocket connection.\n"
148                                "See server log for more information.\n"
149                            ),
150                        )
151
152                if response is None:
153                    self.response = self.protocol.accept(self.request)
154                else:
155                    self.response = response
156
157                if server_header:
158                    self.response.headers["Server"] = server_header
159
160                response = None
161
162                if process_response is not None:
163                    try:
164                        response = process_response(self, self.request, self.response)
165                    except Exception as exc:
166                        self.protocol.handshake_exc = exc
167                        response = self.protocol.reject(
168                            http.HTTPStatus.INTERNAL_SERVER_ERROR,
169                            (
170                                "Failed to open a WebSocket connection.\n"
171                                "See server log for more information.\n"
172                            ),
173                        )
174
175                    if response is not None:
176                        self.response = response
177
178                self.protocol.send_response(self.response)
179
180        # self.protocol.handshake_exc is set when the connection is lost before
181        # receiving a request, when the request cannot be parsed, or when the
182        # handshake fails, including when process_request or process_response
183        # raises an exception.
184
185        # It isn't set when process_request or process_response sends an HTTP
186        # response that rejects the handshake.
187
188        if self.protocol.handshake_exc is not None:
189            raise self.protocol.handshake_exc
190
191    def process_event(self, event: Event) -> None:
192        """
193        Process one incoming event.
194
195        """
196        # First event - handshake request.
197        if self.request is None:
198            assert isinstance(event, Request)
199            self.request = event
200            self.request_rcvd.set()
201        # Later events - frames.
202        else:
203            super().process_event(event)
204
205    def recv_events(self) -> None:
206        """
207        Read incoming data from the socket and process events.
208
209        """
210        try:
211            super().recv_events()
212        finally:
213            # If the connection is closed during the handshake, unblock it.
214            self.request_rcvd.set()
215
216
217class Server:
218    """
219    WebSocket server returned by :func:`serve`.
220
221    This class mirrors the API of :class:`~socketserver.BaseServer`, notably the
222    :meth:`~socketserver.BaseServer.serve_forever` and
223    :meth:`~socketserver.BaseServer.shutdown` methods, as well as the context
224    manager protocol.
225
226    Args:
227        socket: Server socket listening for new connections.
228        handler: Handler for one connection. Receives the socket and address
229            returned by :meth:`~socket.socket.accept`.
230        logger: Logger for this server.
231            It defaults to ``logging.getLogger("websockets.server")``.
232            See the :doc:`logging guide <../../topics/logging>` for details.
233
234    """
235
236    def __init__(
237        self,
238        socket: socket.socket,
239        handler: Callable[[socket.socket, Any], None],
240        logger: LoggerLike | None = None,
241    ) -> None:
242        self.socket = socket
243        self.handler = handler
244        if logger is None:
245            logger = logging.getLogger("websockets.server")
246        self.logger = logger
247        if sys.platform != "win32":
248            self.shutdown_watcher, self.shutdown_notifier = os.pipe()
249
250    def serve_forever(self) -> None:
251        """
252        See :meth:`socketserver.BaseServer.serve_forever`.
253
254        This method doesn't return. Calling :meth:`shutdown` from another thread
255        stops the server.
256
257        Typical use::
258
259            with serve(...) as server:
260                server.serve_forever()
261
262        """
263        poller = selectors.DefaultSelector()
264        try:
265            poller.register(self.socket, selectors.EVENT_READ)
266        except ValueError:  # pragma: no cover
267            # If shutdown() is called before poller.register(),
268            # the socket is closed and poller.register() raises
269            # ValueError: Invalid file descriptor: -1
270            return
271        if sys.platform != "win32":
272            poller.register(self.shutdown_watcher, selectors.EVENT_READ)
273
274        while True:
275            poller.select()
276            try:
277                # If the socket is closed, this will raise an exception and exit
278                # the loop. So we don't need to check the return value of select().
279                sock, addr = self.socket.accept()
280            except OSError:
281                break
282            # Since there isn't a mechanism for tracking connections and waiting
283            # for them to terminate, we cannot use daemon threads, or else all
284            # connections would be terminate brutally when closing the server.
285            thread = threading.Thread(target=self.handler, args=(sock, addr))
286            thread.start()
287
288    def shutdown(self) -> None:
289        """
290        See :meth:`socketserver.BaseServer.shutdown`.
291
292        """
293        self.socket.close()
294        if sys.platform != "win32":
295            os.write(self.shutdown_notifier, b"x")
296
297    def fileno(self) -> int:
298        """
299        See :meth:`socketserver.BaseServer.fileno`.
300
301        """
302        return self.socket.fileno()
303
304    def __enter__(self) -> Server:
305        return self
306
307    def __exit__(
308        self,
309        exc_type: type[BaseException] | None,
310        exc_value: BaseException | None,
311        traceback: TracebackType | None,
312    ) -> None:
313        self.shutdown()
314
315
316def __getattr__(name: str) -> Any:
317    if name == "WebSocketServer":
318        warnings.warn(  # deprecated in 13.0 - 2024-08-20
319            "WebSocketServer was renamed to Server",
320            DeprecationWarning,
321        )
322        return Server
323    raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
324
325
326def serve(
327    handler: Callable[[ServerConnection], None],
328    host: str | None = None,
329    port: int | None = None,
330    *,
331    # TCP/TLS
332    sock: socket.socket | None = None,
333    ssl: ssl_module.SSLContext | None = None,
334    # WebSocket
335    origins: Sequence[Origin | re.Pattern[str] | None] | None = None,
336    extensions: Sequence[ServerExtensionFactory] | None = None,
337    subprotocols: Sequence[Subprotocol] | None = None,
338    select_subprotocol: (
339        Callable[
340            [ServerConnection, Sequence[Subprotocol]],
341            Subprotocol | None,
342        ]
343        | None
344    ) = None,
345    compression: str | None = "deflate",
346    # HTTP
347    process_request: (
348        Callable[
349            [ServerConnection, Request],
350            Response | None,
351        ]
352        | None
353    ) = None,
354    process_response: (
355        Callable[
356            [ServerConnection, Request, Response],
357            Response | None,
358        ]
359        | None
360    ) = None,
361    server_header: str | None = SERVER,
362    # Timeouts
363    open_timeout: float | None = 10,
364    ping_interval: float | None = 20,
365    ping_timeout: float | None = 20,
366    close_timeout: float | None = 10,
367    # Limits
368    max_size: int | None | tuple[int | None, int | None] = 2**20,
369    max_queue: int | None | tuple[int | None, int | None] = 16,
370    # Logging
371    logger: LoggerLike | None = None,
372    # Escape hatch for advanced customization
373    create_connection: type[ServerConnection] | None = None,
374    **kwargs: Any,
375) -> Server:
376    """
377    Create a WebSocket server listening on ``host`` and ``port``.
378
379    Whenever a client connects, the server creates a :class:`ServerConnection`,
380    performs the opening handshake, and delegates to the ``handler``.
381
382    The handler receives the :class:`ServerConnection` instance, which you can
383    use to send and receive messages.
384
385    Once the handler completes, either normally or with an exception, the server
386    performs the closing handshake and closes the connection.
387
388    This function returns a :class:`Server` whose API mirrors
389    :class:`~socketserver.BaseServer`. Treat it as a context manager to ensure
390    that it will be closed and call :meth:`~Server.serve_forever` to serve
391    requests::
392
393        from websockets.sync.server import serve
394
395        def handler(websocket):
396            ...
397
398        with serve(handler, ...) as server:
399            server.serve_forever()
400
401    Args:
402        handler: Connection handler. It receives the WebSocket connection,
403            which is a :class:`ServerConnection`, in argument.
404        host: Network interfaces the server binds to.
405            See :func:`~socket.create_server` for details.
406        port: TCP port the server listens on.
407            See :func:`~socket.create_server` for details.
408        sock: Preexisting TCP socket. ``sock`` replaces ``host`` and ``port``.
409            You may call :func:`socket.create_server` to create a suitable TCP
410            socket.
411        ssl: Configuration for enabling TLS on the connection.
412        origins: Acceptable values of the ``Origin`` header, for defending
413            against Cross-Site WebSocket Hijacking attacks. Values can be
414            :class:`str` to test for an exact match or regular expressions
415            compiled by :func:`re.compile` to test against a pattern. Include
416            :obj:`None` in the list if the lack of an origin is acceptable.
417        extensions: List of supported extensions, in order in which they
418            should be negotiated and run.
419        subprotocols: List of supported subprotocols, in order of decreasing
420            preference.
421        select_subprotocol: Callback for selecting a subprotocol among
422            those supported by the client and the server. It receives a
423            :class:`ServerConnection` (not a
424            :class:`~websockets.server.ServerProtocol`!) instance and a list of
425            subprotocols offered by the client. Other than the first argument,
426            it has the same behavior as the
427            :meth:`ServerProtocol.select_subprotocol
428            <websockets.server.ServerProtocol.select_subprotocol>` method.
429        compression: The "permessage-deflate" extension is enabled by default.
430            Set ``compression`` to :obj:`None` to disable it. See the
431            :doc:`compression guide <../../topics/compression>` for details.
432        process_request: Intercept the request during the opening handshake.
433            Return an HTTP response to force the response. Return :obj:`None` to
434            continue normally. When you force an HTTP 101 Continue response, the
435            handshake is successful. Else, the connection is aborted.
436        process_response: Intercept the response during the opening handshake.
437            Modify the response or return a new HTTP response to force the
438            response. Return :obj:`None` to continue normally. When you force an
439            HTTP 101 Continue response, the handshake is successful. Else, the
440            connection is aborted.
441        server_header: Value of  the ``Server`` response header.
442            It defaults to ``"Python/x.y.z websockets/X.Y"``. Setting it to
443            :obj:`None` removes the header.
444        open_timeout: Timeout for opening connections in seconds.
445            :obj:`None` disables the timeout.
446        ping_interval: Interval between keepalive pings in seconds.
447            :obj:`None` disables keepalive.
448        ping_timeout: Timeout for keepalive pings in seconds.
449            :obj:`None` disables timeouts.
450        close_timeout: Timeout for closing connections in seconds.
451            :obj:`None` disables the timeout.
452        max_size: Maximum size of incoming messages in bytes.
453            :obj:`None` disables the limit. You may pass a ``(max_message_size,
454            max_fragment_size)`` tuple to set different limits for messages and
455            fragments when you expect long messages sent in short fragments.
456        max_queue: High-water mark of the buffer where frames are received.
457            It defaults to 16 frames. The low-water mark defaults to ``max_queue
458            // 4``. You may pass a ``(high, low)`` tuple to set the high-water
459            and low-water marks. If you want to disable flow control entirely,
460            you may set it to ``None``, although that's a bad idea.
461        logger: Logger for this server.
462            It defaults to ``logging.getLogger("websockets.server")``. See the
463            :doc:`logging guide <../../topics/logging>` for details.
464        create_connection: Factory for the :class:`ServerConnection` managing
465            the connection. Set it to a wrapper or a subclass to customize
466            connection handling.
467
468    Any other keyword arguments are passed to :func:`~socket.create_server`.
469
470    """
471
472    # Process parameters
473
474    # Backwards compatibility: ssl used to be called ssl_context.
475    if ssl is None and "ssl_context" in kwargs:
476        ssl = kwargs.pop("ssl_context")
477        warnings.warn(  # deprecated in 13.0 - 2024-08-20
478            "ssl_context was renamed to ssl",
479            DeprecationWarning,
480        )
481
482    if subprotocols is not None:
483        validate_subprotocols(subprotocols)
484
485    if compression == "deflate":
486        extensions = enable_server_permessage_deflate(extensions)
487    elif compression is not None:
488        raise ValueError(f"unsupported compression: {compression}")
489
490    if create_connection is None:
491        create_connection = ServerConnection
492
493    # Bind socket and listen
494
495    # Private APIs for unix_connect()
496    unix: bool = kwargs.pop("unix", False)
497    path: str | None = kwargs.pop("path", None)
498
499    if sock is None:
500        if unix:
501            if path is None:
502                raise ValueError("missing path argument")
503            kwargs.setdefault("family", socket.AF_UNIX)
504            sock = socket.create_server(path, **kwargs)
505        else:
506            sock = socket.create_server((host, port), **kwargs)
507    else:
508        if path is not None:
509            raise ValueError("path and sock arguments are incompatible")
510
511    # Initialize TLS wrapper
512
513    if ssl is not None:
514        sock = ssl.wrap_socket(
515            sock,
516            server_side=True,
517            # Delay TLS handshake until after we set a timeout on the socket.
518            do_handshake_on_connect=False,
519        )
520
521    # Define request handler
522
523    def conn_handler(sock: socket.socket, addr: Any) -> None:
524        # Calculate timeouts on the TLS and WebSocket handshakes.
525        # The TLS timeout must be set on the socket, then removed
526        # to avoid conflicting with the WebSocket timeout in handshake().
527        deadline = Deadline(open_timeout)
528
529        try:
530            # Disable Nagle algorithm
531
532            if not unix:
533                sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, True)
534
535            # Perform TLS handshake
536
537            if ssl is not None:
538                sock.settimeout(deadline.timeout())
539                # mypy cannot figure this out
540                assert isinstance(sock, ssl_module.SSLSocket)
541                sock.do_handshake()
542                sock.settimeout(None)
543
544            # Create a closure to give select_subprotocol access to connection.
545            protocol_select_subprotocol: (
546                Callable[
547                    [ServerProtocol, Sequence[Subprotocol]],
548                    Subprotocol | None,
549                ]
550                | None
551            ) = None
552            if select_subprotocol is not None:
553
554                def protocol_select_subprotocol(
555                    protocol: ServerProtocol,
556                    subprotocols: Sequence[Subprotocol],
557                ) -> Subprotocol | None:
558                    # mypy doesn't know that select_subprotocol is immutable.
559                    assert select_subprotocol is not None
560                    # Ensure this function is only used in the intended context.
561                    assert protocol is connection.protocol
562                    return select_subprotocol(connection, subprotocols)
563
564            # Initialize WebSocket protocol
565
566            protocol = ServerProtocol(
567                origins=origins,
568                extensions=extensions,
569                subprotocols=subprotocols,
570                select_subprotocol=protocol_select_subprotocol,
571                max_size=max_size,
572                logger=logger,
573            )
574
575            # Initialize WebSocket connection
576
577            assert create_connection is not None  # help mypy
578            connection = create_connection(
579                sock,
580                protocol,
581                ping_interval=ping_interval,
582                ping_timeout=ping_timeout,
583                close_timeout=close_timeout,
584                max_queue=max_queue,
585            )
586        except Exception:
587            sock.close()
588            return
589
590        try:
591            try:
592                connection.handshake(
593                    process_request,
594                    process_response,
595                    server_header,
596                    deadline.timeout(),
597                )
598            except TimeoutError:
599                connection.close_socket()
600                connection.recv_events_thread.join()
601                return
602            except Exception:
603                connection.logger.error("opening handshake failed", exc_info=True)
604                connection.close_socket()
605                connection.recv_events_thread.join()
606                return
607
608            assert connection.protocol.state is OPEN
609            try:
610                connection.start_keepalive()
611                handler(connection)
612            except Exception:
613                connection.logger.error("connection handler failed", exc_info=True)
614                connection.close(CloseCode.INTERNAL_ERROR)
615            else:
616                connection.close()
617
618        except Exception:  # pragma: no cover
619            # Don't leak sockets on unexpected errors.
620            sock.close()
621
622    # Initialize server
623
624    return Server(sock, conn_handler, logger)
625
626
627def unix_serve(
628    handler: Callable[[ServerConnection], None],
629    path: str | None = None,
630    **kwargs: Any,
631) -> Server:
632    """
633    Create a WebSocket server listening on a Unix socket.
634
635    This function accepts the same keyword arguments as :func:`serve`.
636
637    It's only available on Unix.
638
639    It's useful for deploying a server behind a reverse proxy such as nginx.
640
641    Args:
642        handler: Connection handler. It receives the WebSocket connection,
643            which is a :class:`ServerConnection`, in argument.
644        path: File system path to the Unix socket.
645
646    """
647    return serve(handler, unix=True, path=path, **kwargs)
648
649
650def is_credentials(credentials: Any) -> bool:
651    try:
652        username, password = credentials
653    except (TypeError, ValueError):
654        return False
655    else:
656        return isinstance(username, str) and isinstance(password, str)
657
658
659def basic_auth(
660    realm: str = "",
661    credentials: tuple[str, str] | Iterable[tuple[str, str]] | None = None,
662    check_credentials: Callable[[str, str], bool] | None = None,
663) -> Callable[[ServerConnection, Request], Response | None]:
664    """
665    Factory for ``process_request`` to enforce HTTP Basic Authentication.
666
667    :func:`basic_auth` is designed to integrate with :func:`serve` as follows::
668
669        from websockets.sync.server import basic_auth, serve
670
671        with serve(
672            ...,
673            process_request=basic_auth(
674                realm="my dev server",
675                credentials=("hello", "iloveyou"),
676            ),
677        ):
678
679    If authentication succeeds, the connection's ``username`` attribute is set.
680    If it fails, the server responds with an HTTP 401 Unauthorized status.
681
682    One of ``credentials`` or ``check_credentials`` must be provided; not both.
683
684    Args:
685        realm: Scope of protection. It should contain only ASCII characters
686            because the encoding of non-ASCII characters is undefined. Refer to
687            section 2.2 of :rfc:`7235` for details.
688        credentials: Hard coded authorized credentials. It can be a
689            ``(username, password)`` pair or a list of such pairs.
690        check_credentials: Function that verifies credentials.
691            It receives ``username`` and ``password`` arguments and returns
692            whether they're valid.
693    Raises:
694        TypeError: If ``credentials`` or ``check_credentials`` is wrong.
695        ValueError: If ``credentials`` and ``check_credentials`` are both
696            provided or both not provided.
697
698    """
699    if (credentials is None) == (check_credentials is None):
700        raise ValueError("provide either credentials or check_credentials")
701
702    if credentials is not None:
703        if is_credentials(credentials):
704            credentials_list = [cast(tuple[str, str], credentials)]
705        elif isinstance(credentials, Iterable):
706            credentials_list = list(cast(Iterable[tuple[str, str]], credentials))
707            if not all(is_credentials(item) for item in credentials_list):
708                raise TypeError(f"invalid credentials argument: {credentials}")
709        else:
710            raise TypeError(f"invalid credentials argument: {credentials}")
711
712        credentials_dict = dict(credentials_list)
713
714        def check_credentials(username: str, password: str) -> bool:
715            try:
716                expected_password = credentials_dict[username]
717            except KeyError:
718                return False
719            return hmac.compare_digest(expected_password, password)
720
721    assert check_credentials is not None  # help mypy
722
723    def process_request(
724        connection: ServerConnection,
725        request: Request,
726    ) -> Response | None:
727        """
728        Perform HTTP Basic Authentication.
729
730        If it succeeds, set the connection's ``username`` attribute and return
731        :obj:`None`. If it fails, return an HTTP 401 Unauthorized responss.
732
733        """
734        try:
735            authorization = request.headers["Authorization"]
736        except KeyError:
737            response = connection.respond(
738                http.HTTPStatus.UNAUTHORIZED,
739                "Missing credentials\n",
740            )
741            response.headers["WWW-Authenticate"] = build_www_authenticate_basic(realm)
742            return response
743
744        try:
745            username, password = parse_authorization_basic(authorization)
746        except InvalidHeader:
747            response = connection.respond(
748                http.HTTPStatus.UNAUTHORIZED,
749                "Unsupported credentials\n",
750            )
751            response.headers["WWW-Authenticate"] = build_www_authenticate_basic(realm)
752            return response
753
754        if not check_credentials(username, password):
755            response = connection.respond(
756                http.HTTPStatus.UNAUTHORIZED,
757                "Invalid credentials\n",
758            )
759            response.headers["WWW-Authenticate"] = build_www_authenticate_basic(realm)
760            return response
761
762        connection.username = username
763        return None
764
765    return process_request
766 
codekingpro/portable-devtools · Team Ai