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