Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
server.py998 linesDownload Raw Back to asyncio
1from __future__ import annotations
2
3import asyncio
4import hmac
5import http
6import logging
7import re
8import socket
9import sys
10from collections.abc import Awaitable, Generator, Iterable, Sequence
11from types import TracebackType
12from typing import Any, Callable, Mapping, cast
13
14from ..exceptions import InvalidHeader
15from ..extensions.base import ServerExtensionFactory
16from ..extensions.permessage_deflate import enable_server_permessage_deflate
17from ..frames import CloseCode
18from ..headers import (
19    build_www_authenticate_basic,
20    parse_authorization_basic,
21    validate_subprotocols,
22)
23from ..http11 import SERVER, Request, Response
24from ..protocol import CONNECTING, OPEN, Event
25from ..server import ServerProtocol
26from ..typing import LoggerLike, Origin, StatusLike, Subprotocol
27from .compatibility import asyncio_timeout
28from .connection import Connection, broadcast
29
30
31__all__ = [
32    "broadcast",
33    "serve",
34    "unix_serve",
35    "ServerConnection",
36    "Server",
37    "basic_auth",
38]
39
40
41class ServerConnection(Connection):
42    """
43    :mod:`asyncio` implementation of a WebSocket server connection.
44
45    :class:`ServerConnection` provides :meth:`recv` and :meth:`send` methods for
46    receiving and sending messages.
47
48    It supports asynchronous iteration to receive messages::
49
50        async for message in websocket:
51            await process(message)
52
53    The iterator exits normally when the connection is closed with code
54    1000 (OK) or 1001 (going away) or without a close code. It raises a
55    :exc:`~websockets.exceptions.ConnectionClosedError` when the connection is
56    closed with any other code.
57
58    The ``ping_interval``, ``ping_timeout``, ``close_timeout``, ``max_queue``,
59    and ``write_limit`` arguments have the same meaning as in :func:`serve`.
60
61    Args:
62        protocol: Sans-I/O connection.
63        server: Server that manages this connection.
64
65    """
66
67    def __init__(
68        self,
69        protocol: ServerProtocol,
70        server: Server,
71        *,
72        ping_interval: float | None = 20,
73        ping_timeout: float | None = 20,
74        close_timeout: float | None = 10,
75        max_queue: int | None | tuple[int | None, int | None] = 16,
76        write_limit: int | tuple[int, int | None] = 2**15,
77    ) -> None:
78        self.protocol: ServerProtocol
79        super().__init__(
80            protocol,
81            ping_interval=ping_interval,
82            ping_timeout=ping_timeout,
83            close_timeout=close_timeout,
84            max_queue=max_queue,
85            write_limit=write_limit,
86        )
87        self.server = server
88        self.request_rcvd: asyncio.Future[None] = self.loop.create_future()
89        self.username: str  # see basic_auth()
90        self.handler: Callable[[ServerConnection], Awaitable[None]]  # see route()
91        self.handler_kwargs: Mapping[str, Any]  # see route()
92
93    def respond(self, status: StatusLike, text: str) -> Response:
94        """
95        Create a plain text HTTP response.
96
97        ``process_request`` and ``process_response`` may call this method to
98        return an HTTP response instead of performing the WebSocket opening
99        handshake.
100
101        You can modify the response before returning it, for example by changing
102        HTTP headers.
103
104        Args:
105            status: HTTP status code.
106            text: HTTP response body; it will be encoded to UTF-8.
107
108        Returns:
109            HTTP response to send to the client.
110
111        """
112        return self.protocol.reject(status, text)
113
114    async def handshake(
115        self,
116        process_request: (
117            Callable[
118                [ServerConnection, Request],
119                Awaitable[Response | None] | Response | None,
120            ]
121            | None
122        ) = None,
123        process_response: (
124            Callable[
125                [ServerConnection, Request, Response],
126                Awaitable[Response | None] | Response | None,
127            ]
128            | None
129        ) = None,
130        server_header: str | None = SERVER,
131    ) -> None:
132        """
133        Perform the opening handshake.
134
135        """
136        await asyncio.wait(
137            [self.request_rcvd, self.connection_lost_waiter],
138            return_when=asyncio.FIRST_COMPLETED,
139        )
140
141        if self.request is not None:
142            async with self.send_context(expected_state=CONNECTING):
143                response = None
144
145                if process_request is not None:
146                    try:
147                        response = process_request(self, self.request)
148                        if isinstance(response, Awaitable):
149                            response = await response
150                    except Exception as exc:
151                        self.protocol.handshake_exc = exc
152                        response = self.protocol.reject(
153                            http.HTTPStatus.INTERNAL_SERVER_ERROR,
154                            (
155                                "Failed to open a WebSocket connection.\n"
156                                "See server log for more information.\n"
157                            ),
158                        )
159
160                if response is None:
161                    if self.server.is_serving():
162                        self.response = self.protocol.accept(self.request)
163                    else:
164                        self.response = self.protocol.reject(
165                            http.HTTPStatus.SERVICE_UNAVAILABLE,
166                            "Server is shutting down.\n",
167                        )
168                else:
169                    assert isinstance(response, Response)  # help mypy
170                    self.response = response
171
172                if server_header:
173                    self.response.headers["Server"] = server_header
174
175                response = None
176
177                if process_response is not None:
178                    try:
179                        response = process_response(self, self.request, self.response)
180                        if isinstance(response, Awaitable):
181                            response = await response
182                    except Exception as exc:
183                        self.protocol.handshake_exc = exc
184                        response = self.protocol.reject(
185                            http.HTTPStatus.INTERNAL_SERVER_ERROR,
186                            (
187                                "Failed to open a WebSocket connection.\n"
188                                "See server log for more information.\n"
189                            ),
190                        )
191
192                if response is not None:
193                    assert isinstance(response, Response)  # help mypy
194                    self.response = response
195
196                self.protocol.send_response(self.response)
197
198        # self.protocol.handshake_exc is set when the connection is lost before
199        # receiving a request, when the request cannot be parsed, or when the
200        # handshake fails, including when process_request or process_response
201        # raises an exception.
202
203        # It isn't set when process_request or process_response sends an HTTP
204        # response that rejects the handshake.
205
206        if self.protocol.handshake_exc is not None:
207            raise self.protocol.handshake_exc
208
209    def process_event(self, event: Event) -> None:
210        """
211        Process one incoming event.
212
213        """
214        # First event - handshake request.
215        if self.request is None:
216            assert isinstance(event, Request)
217            self.request = event
218            self.request_rcvd.set_result(None)
219        # Later events - frames.
220        else:
221            super().process_event(event)
222
223    def connection_made(self, transport: asyncio.BaseTransport) -> None:
224        super().connection_made(transport)
225        self.server.start_connection_handler(self)
226
227
228class Server:
229    """
230    WebSocket server returned by :func:`serve`.
231
232    This class mirrors the API of :class:`asyncio.Server`.
233
234    It keeps track of WebSocket connections in order to close them properly
235    when shutting down.
236
237    Args:
238        handler: Connection handler. It receives the WebSocket connection,
239            which is a :class:`ServerConnection`, in argument.
240        process_request: Intercept the request during the opening handshake.
241            Return an HTTP response to force the response. Return :obj:`None` to
242            continue normally. When you force an HTTP 101 Continue response, the
243            handshake is successful. Else, the connection is aborted.
244            ``process_request`` may be a function or a coroutine.
245        process_response: Intercept the response during the opening handshake.
246            Modify the response or return a new HTTP response to force the
247            response. Return :obj:`None` to continue normally. When you force an
248            HTTP 101 Continue response, the handshake is successful. Else, the
249            connection is aborted. ``process_response`` may be a function or a
250            coroutine.
251        server_header: Value of  the ``Server`` response header.
252            It defaults to ``"Python/x.y.z websockets/X.Y"``. Setting it to
253            :obj:`None` removes the header.
254        open_timeout: Timeout for opening connections in seconds.
255            :obj:`None` disables the timeout.
256        logger: Logger for this server.
257            It defaults to ``logging.getLogger("websockets.server")``.
258            See the :doc:`logging guide <../../topics/logging>` for details.
259
260    """
261
262    def __init__(
263        self,
264        handler: Callable[[ServerConnection], Awaitable[None]],
265        *,
266        process_request: (
267            Callable[
268                [ServerConnection, Request],
269                Awaitable[Response | None] | Response | None,
270            ]
271            | None
272        ) = None,
273        process_response: (
274            Callable[
275                [ServerConnection, Request, Response],
276                Awaitable[Response | None] | Response | None,
277            ]
278            | None
279        ) = None,
280        server_header: str | None = SERVER,
281        open_timeout: float | None = 10,
282        logger: LoggerLike | None = None,
283    ) -> None:
284        self.loop = asyncio.get_running_loop()
285        self.handler = handler
286        self.process_request = process_request
287        self.process_response = process_response
288        self.server_header = server_header
289        self.open_timeout = open_timeout
290        if logger is None:
291            logger = logging.getLogger("websockets.server")
292        self.logger = logger
293
294        # Keep track of active connections.
295        self.handlers: dict[ServerConnection, asyncio.Task[None]] = {}
296
297        # Task responsible for closing the server and terminating connections.
298        self.close_task: asyncio.Task[None] | None = None
299
300        # Completed when the server is closed and connections are terminated.
301        self.closed_waiter: asyncio.Future[None] = self.loop.create_future()
302
303    @property
304    def connections(self) -> set[ServerConnection]:
305        """
306        Set of active connections.
307
308        This property contains all connections that completed the opening
309        handshake successfully and didn't start the closing handshake yet.
310        It can be useful in combination with :func:`~broadcast`.
311
312        """
313        return {connection for connection in self.handlers if connection.state is OPEN}
314
315    def wrap(self, server: asyncio.Server) -> None:
316        """
317        Attach to a given :class:`asyncio.Server`.
318
319        Since :meth:`~asyncio.loop.create_server` doesn't support injecting a
320        custom ``Server`` class, the easiest solution that doesn't rely on
321        private :mod:`asyncio` APIs is to:
322
323        - instantiate a :class:`Server`
324        - give the protocol factory a reference to that instance
325        - call :meth:`~asyncio.loop.create_server` with the factory
326        - attach the resulting :class:`asyncio.Server` with this method
327
328        """
329        self.server = server
330        for sock in server.sockets:
331            if sock.family == socket.AF_INET:
332                name = "%s:%d" % sock.getsockname()
333            elif sock.family == socket.AF_INET6:
334                name = "[%s]:%d" % sock.getsockname()[:2]
335            elif sock.family == socket.AF_UNIX:
336                name = sock.getsockname()
337            # In the unlikely event that someone runs websockets over a
338            # protocol other than IP or Unix sockets, avoid crashing.
339            else:  # pragma: no cover
340                name = str(sock.getsockname())
341            self.logger.info("server listening on %s", name)
342
343    async def conn_handler(self, connection: ServerConnection) -> None:
344        """
345        Handle the lifecycle of a WebSocket connection.
346
347        Since this method doesn't have a caller that can handle exceptions,
348        it attempts to log relevant ones.
349
350        It guarantees that the TCP connection is closed before exiting.
351
352        """
353        try:
354            async with asyncio_timeout(self.open_timeout):
355                try:
356                    await connection.handshake(
357                        self.process_request,
358                        self.process_response,
359                        self.server_header,
360                    )
361                except asyncio.CancelledError:
362                    connection.transport.abort()
363                    raise
364                except Exception:
365                    connection.logger.error("opening handshake failed", exc_info=True)
366                    connection.transport.abort()
367                    return
368
369            if connection.protocol.state is not OPEN:
370                # process_request or process_response rejected the handshake.
371                connection.transport.abort()
372                return
373
374            try:
375                connection.start_keepalive()
376                await self.handler(connection)
377            except Exception:
378                connection.logger.error("connection handler failed", exc_info=True)
379                await connection.close(CloseCode.INTERNAL_ERROR)
380            else:
381                await connection.close()
382
383        except TimeoutError:
384            # When the opening handshake times out, there's nothing to log.
385            pass
386
387        except Exception:  # pragma: no cover
388            # Don't leak connections on unexpected errors.
389            connection.transport.abort()
390
391        finally:
392            # Registration is tied to the lifecycle of conn_handler() because
393            # the server waits for connection handlers to terminate, even if
394            # all connections are already closed.
395            del self.handlers[connection]
396
397    def start_connection_handler(self, connection: ServerConnection) -> None:
398        """
399        Register a connection with this server.
400
401        """
402        # The connection must be registered in self.handlers immediately.
403        # If it was registered in conn_handler(), a race condition could
404        # happen when closing the server after scheduling conn_handler()
405        # but before it starts executing.
406        self.handlers[connection] = self.loop.create_task(self.conn_handler(connection))
407
408    def close(
409        self,
410        close_connections: bool = True,
411        code: CloseCode | int = CloseCode.GOING_AWAY,
412        reason: str = "",
413    ) -> None:
414        """
415        Close the server.
416
417        * Close the underlying :class:`asyncio.Server`.
418        * When ``close_connections`` is :obj:`True`, which is the default, close
419          existing connections. Specifically:
420
421          * Reject opening WebSocket connections with an HTTP 503 (service
422            unavailable) error. This happens when the server accepted the TCP
423            connection but didn't complete the opening handshake before closing.
424          * Close open WebSocket connections with code 1001 (going away).
425            ``code`` and ``reason`` can be customized, for example to use code
426            1012 (service restart).
427
428        * Wait until all connection handlers terminate.
429
430        :meth:`close` is idempotent.
431
432        """
433        if self.close_task is None:
434            self.close_task = self.get_loop().create_task(
435                self._close(close_connections, code, reason)
436            )
437
438    async def _close(
439        self,
440        close_connections: bool = True,
441        code: CloseCode | int = CloseCode.GOING_AWAY,
442        reason: str = "",
443    ) -> None:
444        """
445        Implementation of :meth:`close`.
446
447        This calls :meth:`~asyncio.Server.close` on the underlying
448        :class:`asyncio.Server` object to stop accepting new connections and
449        then closes open connections.
450
451        """
452        self.logger.info("server closing")
453
454        # Stop accepting new connections.
455        self.server.close()
456
457        # Wait until all accepted connections reach connection_made() and call
458        # register(). See https://github.com/python/cpython/issues/79033 for
459        # details. This workaround can be removed when dropping Python < 3.11.
460        await asyncio.sleep(0)
461
462        # After server.close(), handshake() closes OPENING connections with an
463        # HTTP 503 error.
464
465        if close_connections:
466            # Close OPEN connections with code 1001 by default.
467            close_tasks = [
468                asyncio.create_task(connection.close(code, reason))
469                for connection in self.handlers
470                if connection.protocol.state is not CONNECTING
471            ]
472            # asyncio.wait doesn't accept an empty first argument.
473            if close_tasks:
474                await asyncio.wait(close_tasks)
475
476        # Wait until all TCP connections are closed.
477        await self.server.wait_closed()
478
479        # Wait until all connection handlers terminate.
480        # asyncio.wait doesn't accept an empty first argument.
481        if self.handlers:
482            await asyncio.wait(self.handlers.values())
483
484        # Tell wait_closed() to return.
485        self.closed_waiter.set_result(None)
486
487        self.logger.info("server closed")
488
489    async def wait_closed(self) -> None:
490        """
491        Wait until the server is closed.
492
493        When :meth:`wait_closed` returns, all TCP connections are closed and
494        all connection handlers have returned.
495
496        To ensure a fast shutdown, a connection handler should always be
497        awaiting at least one of:
498
499        * :meth:`~ServerConnection.recv`: when the connection is closed,
500          it raises :exc:`~websockets.exceptions.ConnectionClosedOK`;
501        * :meth:`~ServerConnection.wait_closed`: when the connection is
502          closed, it returns.
503
504        Then the connection handler is immediately notified of the shutdown;
505        it can clean up and exit.
506
507        """
508        await asyncio.shield(self.closed_waiter)
509
510    def get_loop(self) -> asyncio.AbstractEventLoop:
511        """
512        See :meth:`asyncio.Server.get_loop`.
513
514        """
515        return self.server.get_loop()
516
517    def is_serving(self) -> bool:  # pragma: no cover
518        """
519        See :meth:`asyncio.Server.is_serving`.
520
521        """
522        return self.server.is_serving()
523
524    async def start_serving(self) -> None:  # pragma: no cover
525        """
526        See :meth:`asyncio.Server.start_serving`.
527
528        Typical use::
529
530            server = await serve(..., start_serving=False)
531            # perform additional setup here...
532            # ... then start the server
533            await server.start_serving()
534
535        """
536        await self.server.start_serving()
537
538    async def serve_forever(self) -> None:  # pragma: no cover
539        """
540        See :meth:`asyncio.Server.serve_forever`.
541
542        Typical use::
543
544            server = await serve(...)
545            # this coroutine doesn't return
546            # canceling it stops the server
547            await server.serve_forever()
548
549        This is an alternative to using :func:`serve` as an asynchronous context
550        manager. Shutdown is triggered by canceling :meth:`serve_forever`
551        instead of exiting a :func:`serve` context.
552
553        """
554        await self.server.serve_forever()
555
556    @property
557    def sockets(self) -> tuple[socket.socket, ...]:
558        """
559        See :attr:`asyncio.Server.sockets`.
560
561        """
562        return self.server.sockets
563
564    async def __aenter__(self) -> Server:  # pragma: no cover
565        return self
566
567    async def __aexit__(
568        self,
569        exc_type: type[BaseException] | None,
570        exc_value: BaseException | None,
571        traceback: TracebackType | None,
572    ) -> None:  # pragma: no cover
573        self.close()
574        await self.wait_closed()
575
576
577# This is spelled in lower case because it's exposed as a callable in the API.
578class serve:
579    """
580    Create a WebSocket server listening on ``host`` and ``port``.
581
582    Whenever a client connects, the server creates a :class:`ServerConnection`,
583    performs the opening handshake, and delegates to the ``handler`` coroutine.
584
585    The handler receives the :class:`ServerConnection` instance, which you can
586    use to send and receive messages.
587
588    Once the handler completes, either normally or with an exception, the server
589    performs the closing handshake and closes the connection.
590
591    This coroutine returns a :class:`Server` whose API mirrors
592    :class:`asyncio.Server`. Treat it as an asynchronous context manager to
593    ensure that the server will be closed::
594
595        from websockets.asyncio.server import serve
596
597        def handler(websocket):
598            ...
599
600        # set this future to exit the server
601        stop = asyncio.get_running_loop().create_future()
602
603        async with serve(handler, host, port):
604            await stop
605
606    Alternatively, call :meth:`~Server.serve_forever` to serve requests and
607    cancel it to stop the server::
608
609        server = await serve(handler, host, port)
610        await server.serve_forever()
611
612    Args:
613        handler: Connection handler. It receives the WebSocket connection,
614            which is a :class:`ServerConnection`, in argument.
615        host: Network interfaces the server binds to.
616            See :meth:`~asyncio.loop.create_server` for details.
617        port: TCP port the server listens on.
618            See :meth:`~asyncio.loop.create_server` for details.
619        origins: Acceptable values of the ``Origin`` header, for defending
620            against Cross-Site WebSocket Hijacking attacks. Values can be
621            :class:`str` to test for an exact match or regular expressions
622            compiled by :func:`re.compile` to test against a pattern. Include
623            :obj:`None` in the list if the lack of an origin is acceptable.
624        extensions: List of supported extensions, in order in which they
625            should be negotiated and run.
626        subprotocols: List of supported subprotocols, in order of decreasing
627            preference.
628        select_subprotocol: Callback for selecting a subprotocol among
629            those supported by the client and the server. It receives a
630            :class:`ServerConnection` (not a
631            :class:`~websockets.server.ServerProtocol`!) instance and a list of
632            subprotocols offered by the client. Other than the first argument,
633            it has the same behavior as the
634            :meth:`ServerProtocol.select_subprotocol
635            <websockets.server.ServerProtocol.select_subprotocol>` method.
636        compression: The "permessage-deflate" extension is enabled by default.
637            Set ``compression`` to :obj:`None` to disable it. See the
638            :doc:`compression guide <../../topics/compression>` for details.
639        process_request: Intercept the request during the opening handshake.
640            Return an HTTP response to force the response or :obj:`None` to
641            continue normally. When you force an HTTP 101 Continue response, the
642            handshake is successful. Else, the connection is aborted.
643            ``process_request`` may be a function or a coroutine.
644        process_response: Intercept the response during the opening handshake.
645            Return an HTTP response to force the response or :obj:`None` to
646            continue normally. When you force an HTTP 101 Continue response, the
647            handshake is successful. Else, the connection is aborted.
648            ``process_response`` may be a function or a coroutine.
649        server_header: Value of  the ``Server`` response header.
650            It defaults to ``"Python/x.y.z websockets/X.Y"``. Setting it to
651            :obj:`None` removes the header.
652        open_timeout: Timeout for opening connections in seconds.
653            :obj:`None` disables the timeout.
654        ping_interval: Interval between keepalive pings in seconds.
655            :obj:`None` disables keepalive.
656        ping_timeout: Timeout for keepalive pings in seconds.
657            :obj:`None` disables timeouts.
658        close_timeout: Timeout for closing connections in seconds.
659            :obj:`None` disables the timeout.
660        max_size: Maximum size of incoming messages in bytes.
661            :obj:`None` disables the limit. You may pass a ``(max_message_size,
662            max_fragment_size)`` tuple to set different limits for messages and
663            fragments when you expect long messages sent in short fragments.
664        max_queue: High-water mark of the buffer where frames are received.
665            It defaults to 16 frames. The low-water mark defaults to ``max_queue
666            // 4``. You may pass a ``(high, low)`` tuple to set the high-water
667            and low-water marks. If you want to disable flow control entirely,
668            you may set it to ``None``, although that's a bad idea.
669        write_limit: High-water mark of write buffer in bytes. It is passed to
670            :meth:`~asyncio.WriteTransport.set_write_buffer_limits`. It defaults
671            to 32 KiB. You may pass a ``(high, low)`` tuple to set the
672            high-water and low-water marks.
673        logger: Logger for this server.
674            It defaults to ``logging.getLogger("websockets.server")``. See the
675            :doc:`logging guide <../../topics/logging>` for details.
676        create_connection: Factory for the :class:`ServerConnection` managing
677            the connection. Set it to a wrapper or a subclass to customize
678            connection handling.
679
680    Any other keyword arguments are passed to the event loop's
681    :meth:`~asyncio.loop.create_server` method.
682
683    For example:
684
685    * You can set ``ssl`` to a :class:`~ssl.SSLContext` to enable TLS.
686
687    * You can set ``sock`` to provide a preexisting TCP socket. You may call
688      :func:`socket.create_server` (not to be confused with the event loop's
689      :meth:`~asyncio.loop.create_server` method) to create a suitable server
690      socket and customize it.
691
692    * You can set ``start_serving`` to ``False`` to start accepting connections
693      only after you call :meth:`~Server.start_serving()` or
694      :meth:`~Server.serve_forever()`.
695
696    """
697
698    def __init__(
699        self,
700        handler: Callable[[ServerConnection], Awaitable[None]],
701        host: str | None = None,
702        port: int | None = None,
703        *,
704        # WebSocket
705        origins: Sequence[Origin | re.Pattern[str] | None] | None = None,
706        extensions: Sequence[ServerExtensionFactory] | None = None,
707        subprotocols: Sequence[Subprotocol] | None = None,
708        select_subprotocol: (
709            Callable[
710                [ServerConnection, Sequence[Subprotocol]],
711                Subprotocol | None,
712            ]
713            | None
714        ) = None,
715        compression: str | None = "deflate",
716        # HTTP
717        process_request: (
718            Callable[
719                [ServerConnection, Request],
720                Awaitable[Response | None] | Response | None,
721            ]
722            | None
723        ) = None,
724        process_response: (
725            Callable[
726                [ServerConnection, Request, Response],
727                Awaitable[Response | None] | Response | None,
728            ]
729            | None
730        ) = None,
731        server_header: str | None = SERVER,
732        # Timeouts
733        open_timeout: float | None = 10,
734        ping_interval: float | None = 20,
735        ping_timeout: float | None = 20,
736        close_timeout: float | None = 10,
737        # Limits
738        max_size: int | None | tuple[int | None, int | None] = 2**20,
739        max_queue: int | None | tuple[int | None, int | None] = 16,
740        write_limit: int | tuple[int, int | None] = 2**15,
741        # Logging
742        logger: LoggerLike | None = None,
743        # Escape hatch for advanced customization
744        create_connection: type[ServerConnection] | None = None,
745        # Other keyword arguments are passed to loop.create_server
746        **kwargs: Any,
747    ) -> None:
748        if subprotocols is not None:
749            validate_subprotocols(subprotocols)
750
751        if compression == "deflate":
752            extensions = enable_server_permessage_deflate(extensions)
753        elif compression is not None:
754            raise ValueError(f"unsupported compression: {compression}")
755
756        if create_connection is None:
757            create_connection = ServerConnection
758
759        self.server = Server(
760            handler,
761            process_request=process_request,
762            process_response=process_response,
763            server_header=server_header,
764            open_timeout=open_timeout,
765            logger=logger,
766        )
767
768        if kwargs.get("ssl") is not None:
769            kwargs.setdefault("ssl_handshake_timeout", open_timeout)
770            if sys.version_info[:2] >= (3, 11):  # pragma: no branch
771                kwargs.setdefault("ssl_shutdown_timeout", close_timeout)
772
773        def factory() -> ServerConnection:
774            """
775            Create an asyncio protocol for managing a WebSocket connection.
776
777            """
778            # Create a closure to give select_subprotocol access to connection.
779            protocol_select_subprotocol: (
780                Callable[
781                    [ServerProtocol, Sequence[Subprotocol]],
782                    Subprotocol | None,
783                ]
784                | None
785            ) = None
786            if select_subprotocol is not None:
787
788                def protocol_select_subprotocol(
789                    protocol: ServerProtocol,
790                    subprotocols: Sequence[Subprotocol],
791                ) -> Subprotocol | None:
792                    # mypy doesn't know that select_subprotocol is immutable.
793                    assert select_subprotocol is not None
794                    # Ensure this function is only used in the intended context.
795                    assert protocol is connection.protocol
796                    return select_subprotocol(connection, subprotocols)
797
798            # This is a protocol in the Sans-I/O implementation of websockets.
799            protocol = ServerProtocol(
800                origins=origins,
801                extensions=extensions,
802                subprotocols=subprotocols,
803                select_subprotocol=protocol_select_subprotocol,
804                max_size=max_size,
805                logger=logger,
806            )
807            # This is a connection in websockets and a protocol in asyncio.
808            connection = create_connection(
809                protocol,
810                self.server,
811                ping_interval=ping_interval,
812                ping_timeout=ping_timeout,
813                close_timeout=close_timeout,
814                max_queue=max_queue,
815                write_limit=write_limit,
816            )
817            return connection
818
819        loop = asyncio.get_running_loop()
820        if kwargs.pop("unix", False):
821            self.create_server = loop.create_unix_server(factory, **kwargs)
822        else:
823            # mypy cannot tell that kwargs must provide sock when port is None.
824            self.create_server = loop.create_server(factory, host, port, **kwargs)  # type: ignore[arg-type]
825
826    # async with serve(...) as ...: ...
827
828    async def __aenter__(self) -> Server:
829        return await self
830
831    async def __aexit__(
832        self,
833        exc_type: type[BaseException] | None,
834        exc_value: BaseException | None,
835        traceback: TracebackType | None,
836    ) -> None:
837        self.server.close()
838        await self.server.wait_closed()
839
840    # ... = await serve(...)
841
842    def __await__(self) -> Generator[Any, None, Server]:
843        # Create a suitable iterator by calling __await__ on a coroutine.
844        return self.__await_impl__().__await__()
845
846    async def __await_impl__(self) -> Server:
847        server = await self.create_server
848        self.server.wrap(server)
849        return self.server
850
851    # ... = yield from serve(...) - remove when dropping Python < 3.11
852
853    __iter__ = __await__
854
855
856def unix_serve(
857    handler: Callable[[ServerConnection], Awaitable[None]],
858    path: str | None = None,
859    **kwargs: Any,
860) -> Awaitable[Server]:
861    """
862    Create a WebSocket server listening on a Unix socket.
863
864    This function is identical to :func:`serve`, except the ``host`` and
865    ``port`` arguments are replaced by ``path``. It's only available on Unix.
866
867    It's useful for deploying a server behind a reverse proxy such as nginx.
868
869    Args:
870        handler: Connection handler. It receives the WebSocket connection,
871            which is a :class:`ServerConnection`, in argument.
872        path: File system path to the Unix socket.
873
874    """
875    return serve(handler, unix=True, path=path, **kwargs)
876
877
878def is_credentials(credentials: Any) -> bool:
879    try:
880        username, password = credentials
881    except (TypeError, ValueError):
882        return False
883    else:
884        return isinstance(username, str) and isinstance(password, str)
885
886
887def basic_auth(
888    realm: str = "",
889    credentials: tuple[str, str] | Iterable[tuple[str, str]] | None = None,
890    check_credentials: Callable[[str, str], Awaitable[bool] | bool] | None = None,
891) -> Callable[[ServerConnection, Request], Awaitable[Response | None]]:
892    """
893    Factory for ``process_request`` to enforce HTTP Basic Authentication.
894
895    :func:`basic_auth` is designed to integrate with :func:`serve` as follows::
896
897        from websockets.asyncio.server import basic_auth, serve
898
899        async with serve(
900            ...,
901            process_request=basic_auth(
902                realm="my dev server",
903                credentials=("hello", "iloveyou"),
904            ),
905        ):
906
907    If authentication succeeds, the connection's ``username`` attribute is set.
908    If it fails, the server responds with an HTTP 401 Unauthorized status.
909
910    One of ``credentials`` or ``check_credentials`` must be provided; not both.
911
912    Args:
913        realm: Scope of protection. It should contain only ASCII characters
914            because the encoding of non-ASCII characters is undefined. Refer to
915            section 2.2 of :rfc:`7235` for details.
916        credentials: Hard coded authorized credentials. It can be a
917            ``(username, password)`` pair or a list of such pairs.
918        check_credentials: Function or coroutine that verifies credentials.
919            It receives ``username`` and ``password`` arguments and returns
920            whether they're valid.
921    Raises:
922        TypeError: If ``credentials`` or ``check_credentials`` is wrong.
923        ValueError: If ``credentials`` and ``check_credentials`` are both
924            provided or both not provided.
925
926    """
927    if (credentials is None) == (check_credentials is None):
928        raise ValueError("provide either credentials or check_credentials")
929
930    if credentials is not None:
931        if is_credentials(credentials):
932            credentials_list = [cast(tuple[str, str], credentials)]
933        elif isinstance(credentials, Iterable):
934            credentials_list = list(cast(Iterable[tuple[str, str]], credentials))
935            if not all(is_credentials(item) for item in credentials_list):
936                raise TypeError(f"invalid credentials argument: {credentials}")
937        else:
938            raise TypeError(f"invalid credentials argument: {credentials}")
939
940        credentials_dict = dict(credentials_list)
941
942        def check_credentials(username: str, password: str) -> bool:
943            try:
944                expected_password = credentials_dict[username]
945            except KeyError:
946                return False
947            return hmac.compare_digest(expected_password, password)
948
949    assert check_credentials is not None  # help mypy
950
951    async def process_request(
952        connection: ServerConnection,
953        request: Request,
954    ) -> Response | None:
955        """
956        Perform HTTP Basic Authentication.
957
958        If it succeeds, set the connection's ``username`` attribute and return
959        :obj:`None`. If it fails, return an HTTP 401 Unauthorized responss.
960
961        """
962        try:
963            authorization = request.headers["Authorization"]
964        except KeyError:
965            response = connection.respond(
966                http.HTTPStatus.UNAUTHORIZED,
967                "Missing credentials\n",
968            )
969            response.headers["WWW-Authenticate"] = build_www_authenticate_basic(realm)
970            return response
971
972        try:
973            username, password = parse_authorization_basic(authorization)
974        except InvalidHeader:
975            response = connection.respond(
976                http.HTTPStatus.UNAUTHORIZED,
977                "Unsupported credentials\n",
978            )
979            response.headers["WWW-Authenticate"] = build_www_authenticate_basic(realm)
980            return response
981
982        valid_credentials = check_credentials(username, password)
983        if isinstance(valid_credentials, Awaitable):
984            valid_credentials = await valid_credentials
985
986        if not valid_credentials:
987            response = connection.respond(
988                http.HTTPStatus.UNAUTHORIZED,
989                "Invalid credentials\n",
990            )
991            response.headers["WWW-Authenticate"] = build_www_authenticate_basic(realm)
992            return response
993
994        connection.username = username
995        return None
996
997    return process_request
998 
codekingpro/portable-devtools · Team Ai