codekingpro/portable-devtools
114k
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 