Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
client.py392 linesDownload Raw Back to websockets
1from __future__ import annotations
2
3import os
4import random
5import warnings
6from collections.abc import Generator, Sequence
7from typing import Any
8
9from .datastructures import Headers, MultipleValuesError
10from .exceptions import (
11    InvalidHandshake,
12    InvalidHeader,
13    InvalidHeaderValue,
14    InvalidMessage,
15    InvalidStatus,
16    InvalidUpgrade,
17    NegotiationError,
18)
19from .extensions import ClientExtensionFactory, Extension
20from .headers import (
21    build_authorization_basic,
22    build_extension,
23    build_host,
24    build_subprotocol,
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 CLIENT, CONNECTING, OPEN, Protocol, State
33from .typing import (
34    ConnectionOption,
35    ExtensionHeader,
36    LoggerLike,
37    Origin,
38    Subprotocol,
39    UpgradeProtocol,
40)
41from .uri import WebSocketURI
42from .utils import accept_key, generate_key
43
44
45__all__ = ["ClientProtocol"]
46
47
48class ClientProtocol(Protocol):
49    """
50    Sans-I/O implementation of a WebSocket client connection.
51
52    Args:
53        uri: URI of the WebSocket server, parsed
54            with :func:`~websockets.uri.parse_uri`.
55        origin: Value of the ``Origin`` header. This is useful when connecting
56            to a server that validates the ``Origin`` header to defend against
57            Cross-Site WebSocket Hijacking attacks.
58        extensions: List of supported extensions, in order in which they
59            should be tried.
60        subprotocols: List of supported subprotocols, in order of decreasing
61            preference.
62        state: Initial state of the WebSocket connection.
63        max_size: Maximum size of incoming messages in bytes.
64            :obj:`None` disables the limit. You may pass a ``(max_message_size,
65            max_fragment_size)`` tuple to set different limits for messages and
66            fragments when you expect long messages sent in short fragments.
67        logger: Logger for this connection;
68            defaults to ``logging.getLogger("websockets.client")``;
69            see the :doc:`logging guide <../../topics/logging>` for details.
70
71    """
72
73    def __init__(
74        self,
75        uri: WebSocketURI,
76        *,
77        origin: Origin | None = None,
78        extensions: Sequence[ClientExtensionFactory] | None = None,
79        subprotocols: Sequence[Subprotocol] | None = None,
80        state: State = CONNECTING,
81        max_size: int | None | tuple[int | None, int | None] = 2**20,
82        logger: LoggerLike | None = None,
83    ) -> None:
84        super().__init__(
85            side=CLIENT,
86            state=state,
87            max_size=max_size,
88            logger=logger,
89        )
90        self.uri = uri
91        self.origin = origin
92        self.available_extensions = extensions
93        self.available_subprotocols = subprotocols
94        self.key = generate_key()
95
96    def connect(self) -> Request:
97        """
98        Create a handshake request to open a connection.
99
100        You must send the handshake request with :meth:`send_request`.
101
102        You can modify it before sending it, for example to add HTTP headers.
103
104        Returns:
105            WebSocket handshake request event to send to the server.
106
107        """
108        headers = Headers()
109        headers["Host"] = build_host(self.uri.host, self.uri.port, self.uri.secure)
110        if self.uri.user_info:
111            headers["Authorization"] = build_authorization_basic(*self.uri.user_info)
112        if self.origin is not None:
113            headers["Origin"] = self.origin
114        headers["Upgrade"] = "websocket"
115        headers["Connection"] = "Upgrade"
116        headers["Sec-WebSocket-Key"] = self.key
117        headers["Sec-WebSocket-Version"] = "13"
118        if self.available_extensions is not None:
119            headers["Sec-WebSocket-Extensions"] = build_extension(
120                [
121                    (extension_factory.name, extension_factory.get_request_params())
122                    for extension_factory in self.available_extensions
123                ]
124            )
125        if self.available_subprotocols is not None:
126            headers["Sec-WebSocket-Protocol"] = build_subprotocol(
127                self.available_subprotocols
128            )
129        return Request(self.uri.resource_name, headers)
130
131    def process_response(self, response: Response) -> None:
132        """
133        Check a handshake response.
134
135        Args:
136            request: WebSocket handshake response received from the server.
137
138        Raises:
139            InvalidHandshake: If the handshake response is invalid.
140
141        """
142
143        if response.status_code != 101:
144            raise InvalidStatus(response)
145
146        headers = response.headers
147
148        connection: list[ConnectionOption] = sum(
149            [parse_connection(value) for value in headers.get_all("Connection")], []
150        )
151        if not any(value.lower() == "upgrade" for value in connection):
152            raise InvalidUpgrade(
153                "Connection", ", ".join(connection) if connection else None
154            )
155
156        upgrade: list[UpgradeProtocol] = sum(
157            [parse_upgrade(value) for value in headers.get_all("Upgrade")], []
158        )
159        # For compatibility with non-strict implementations, ignore case when
160        # checking the Upgrade header. It's supposed to be 'WebSocket'.
161        if not (len(upgrade) == 1 and upgrade[0].lower() == "websocket"):
162            raise InvalidUpgrade("Upgrade", ", ".join(upgrade) if upgrade else None)
163
164        try:
165            s_w_accept = headers["Sec-WebSocket-Accept"]
166        except KeyError:
167            raise InvalidHeader("Sec-WebSocket-Accept") from None
168        except MultipleValuesError:
169            raise InvalidHeader("Sec-WebSocket-Accept", "multiple values") from None
170        if s_w_accept != accept_key(self.key):
171            raise InvalidHeaderValue("Sec-WebSocket-Accept", s_w_accept)
172
173        self.extensions = self.process_extensions(headers)
174        self.subprotocol = self.process_subprotocol(headers)
175
176    def process_extensions(self, headers: Headers) -> list[Extension]:
177        """
178        Handle the Sec-WebSocket-Extensions HTTP response header.
179
180        Check that each extension is supported, as well as its parameters.
181
182        :rfc:`6455` leaves the rules up to the specification of each
183        extension.
184
185        To provide this level of flexibility, for each extension accepted by
186        the server, we check for a match with each extension available in the
187        client configuration. If no match is found, an exception is raised.
188
189        If several variants of the same extension are accepted by the server,
190        it may be configured several times, which won't make sense in general.
191        Extensions must implement their own requirements. For this purpose,
192        the list of previously accepted extensions is provided.
193
194        Other requirements, for example related to mandatory extensions or the
195        order of extensions, may be implemented by overriding this method.
196
197        Args:
198            headers: WebSocket handshake response headers.
199
200        Returns:
201            List of accepted extensions.
202
203        Raises:
204            InvalidHandshake: To abort the handshake.
205
206        """
207        accepted_extensions: list[Extension] = []
208
209        extensions = headers.get_all("Sec-WebSocket-Extensions")
210
211        if extensions:
212            if self.available_extensions is None:
213                raise NegotiationError("no extensions supported")
214
215            parsed_extensions: list[ExtensionHeader] = sum(
216                [parse_extension(header_value) for header_value in extensions], []
217            )
218
219            for name, response_params in parsed_extensions:
220                for extension_factory in self.available_extensions:
221                    # Skip non-matching extensions based on their name.
222                    if extension_factory.name != name:
223                        continue
224
225                    # Skip non-matching extensions based on their params.
226                    try:
227                        extension = extension_factory.process_response_params(
228                            response_params, accepted_extensions
229                        )
230                    except NegotiationError:
231                        continue
232
233                    # Add matching extension to the final list.
234                    accepted_extensions.append(extension)
235
236                    # Break out of the loop once we have a match.
237                    break
238
239                # If we didn't break from the loop, no extension in our list
240                # matched what the server sent. Fail the connection.
241                else:
242                    raise NegotiationError(
243                        f"Unsupported extension: "
244                        f"name = {name}, params = {response_params}"
245                    )
246
247        return accepted_extensions
248
249    def process_subprotocol(self, headers: Headers) -> Subprotocol | None:
250        """
251        Handle the Sec-WebSocket-Protocol HTTP response header.
252
253        If provided, check that it contains exactly one supported subprotocol.
254
255        Args:
256            headers: WebSocket handshake response headers.
257
258        Returns:
259           Subprotocol, if one was selected.
260
261        """
262        subprotocol: Subprotocol | None = None
263
264        subprotocols = headers.get_all("Sec-WebSocket-Protocol")
265
266        if subprotocols:
267            if self.available_subprotocols is None:
268                raise NegotiationError("no subprotocols supported")
269
270            parsed_subprotocols: Sequence[Subprotocol] = sum(
271                [parse_subprotocol(header_value) for header_value in subprotocols], []
272            )
273            if len(parsed_subprotocols) > 1:
274                raise InvalidHeader(
275                    "Sec-WebSocket-Protocol",
276                    f"multiple values: {', '.join(parsed_subprotocols)}",
277                )
278
279            subprotocol = parsed_subprotocols[0]
280            if subprotocol not in self.available_subprotocols:
281                raise NegotiationError(f"unsupported subprotocol: {subprotocol}")
282
283        return subprotocol
284
285    def send_request(self, request: Request) -> None:
286        """
287        Send a handshake request to the server.
288
289        Args:
290            request: WebSocket handshake request event.
291
292        """
293        if self.debug:
294            self.logger.debug("> GET %s HTTP/1.1", request.path)
295            for key, value in request.headers.raw_items():
296                self.logger.debug("> %s: %s", key, value)
297
298        self.writes.append(request.serialize())
299
300    def parse(self) -> Generator[None]:
301        if self.state is CONNECTING:
302            try:
303                response = yield from Response.parse(
304                    self.reader.read_line,
305                    self.reader.read_exact,
306                    self.reader.read_to_eof,
307                )
308            except Exception as exc:
309                self.handshake_exc = InvalidMessage(
310                    "did not receive a valid HTTP response"
311                )
312                self.handshake_exc.__cause__ = exc
313                self.send_eof()
314                self.parser = self.discard()
315                next(self.parser)  # start coroutine
316                yield
317
318            if self.debug:
319                code, phrase = response.status_code, response.reason_phrase
320                self.logger.debug("< HTTP/1.1 %d %s", code, phrase)
321                for key, value in response.headers.raw_items():
322                    self.logger.debug("< %s: %s", key, value)
323                if response.body:
324                    self.logger.debug("< [body] (%d bytes)", len(response.body))
325
326            try:
327                self.process_response(response)
328            except InvalidHandshake as exc:
329                response._exception = exc
330                self.events.append(response)
331                self.handshake_exc = exc
332                self.send_eof()
333                self.parser = self.discard()
334                next(self.parser)  # start coroutine
335                yield
336
337            assert self.state is CONNECTING
338            self.state = OPEN
339            self.events.append(response)
340
341        yield from super().parse()
342
343
344class ClientConnection(ClientProtocol):
345    def __init__(self, *args: Any, **kwargs: Any) -> None:
346        warnings.warn(  # deprecated in 11.0 - 2023-04-02
347            "ClientConnection was renamed to ClientProtocol",
348            DeprecationWarning,
349        )
350        super().__init__(*args, **kwargs)
351
352
353BACKOFF_INITIAL_DELAY = float(os.environ.get("WEBSOCKETS_BACKOFF_INITIAL_DELAY", "5"))
354BACKOFF_MIN_DELAY = float(os.environ.get("WEBSOCKETS_BACKOFF_MIN_DELAY", "3.1"))
355BACKOFF_MAX_DELAY = float(os.environ.get("WEBSOCKETS_BACKOFF_MAX_DELAY", "90.0"))
356BACKOFF_FACTOR = float(os.environ.get("WEBSOCKETS_BACKOFF_FACTOR", "1.618"))
357
358
359def backoff(
360    initial_delay: float = BACKOFF_INITIAL_DELAY,
361    min_delay: float = BACKOFF_MIN_DELAY,
362    max_delay: float = BACKOFF_MAX_DELAY,
363    factor: float = BACKOFF_FACTOR,
364) -> Generator[float]:
365    """
366    Generate a series of backoff delays between reconnection attempts.
367
368    Yields:
369        How many seconds to wait before retrying to connect.
370
371    """
372    # Add a random initial delay between 0 and 5 seconds.
373    # See 7.2.3. Recovering from Abnormal Closure in RFC 6455.
374    yield random.random() * initial_delay
375    delay = min_delay
376    while delay < max_delay:
377        yield delay
378        delay *= factor
379    while True:
380        yield max_delay
381
382
383lazy_import(
384    globals(),
385    deprecated_aliases={
386        # deprecated in 14.0 - 2024-11-09
387        "WebSocketClientProtocol": ".legacy.client",
388        "connect": ".legacy.client",
389        "unix_connect": ".legacy.client",
390    },
391)
392 
codekingpro/portable-devtools · Team Ai