client.py 30 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800
  1. from __future__ import annotations
  2. import logging
  3. import os
  4. import ssl as ssl_module
  5. import traceback
  6. import urllib.parse
  7. from collections.abc import AsyncIterator, Generator, Sequence
  8. from types import TracebackType
  9. from typing import Any, Callable, Literal
  10. import trio
  11. from ..client import ClientProtocol, backoff, process_exception
  12. from ..datastructures import Headers, HeadersLike
  13. from ..exceptions import (
  14. InvalidProxyMessage,
  15. InvalidProxyStatus,
  16. InvalidStatus,
  17. ProxyError,
  18. SecurityError,
  19. )
  20. from ..extensions.base import ClientExtensionFactory
  21. from ..extensions.permessage_deflate import enable_client_permessage_deflate
  22. from ..headers import validate_subprotocols
  23. from ..http11 import USER_AGENT, Response
  24. from ..protocol import CONNECTING, Event
  25. from ..proxy import Proxy, get_proxy, parse_proxy, prepare_connect_request
  26. from ..streams import StreamReader
  27. from ..typing import LoggerLike, Origin, PathLike, Subprotocol
  28. from ..uri import WebSocketURI, parse_uri
  29. from .connection import Connection
  30. from .utils import race_events
  31. __all__ = ["connect", "unix_connect", "ClientConnection"]
  32. MAX_REDIRECTS = int(os.environ.get("WEBSOCKETS_MAX_REDIRECTS", "10"))
  33. class ClientConnection(Connection):
  34. """
  35. :mod:`trio` implementation of a WebSocket client connection.
  36. :class:`ClientConnection` provides :meth:`recv` and :meth:`send` coroutines
  37. for receiving and sending messages.
  38. It supports asynchronous iteration to receive messages::
  39. async for message in websocket:
  40. await process(message)
  41. The iterator exits normally when the connection is closed with close code
  42. 1000 (OK) or 1001 (going away) or without a close code. It raises a
  43. :exc:`~websockets.exceptions.ConnectionClosedError` when the connection is
  44. closed with any other code.
  45. The ``ping_interval``, ``ping_timeout``, ``close_timeout``, and
  46. ``max_queue`` arguments have the same meaning as in :func:`connect`.
  47. Args:
  48. nursery: Trio nursery.
  49. stream: Trio stream connected to a WebSocket server.
  50. protocol: Sans-I/O connection.
  51. """
  52. def __init__(
  53. self,
  54. nursery: trio.Nursery,
  55. stream: trio.abc.Stream,
  56. protocol: ClientProtocol,
  57. *,
  58. ping_interval: float | None = 20,
  59. ping_timeout: float | None = 20,
  60. close_timeout: float | None = 10,
  61. max_queue: int | None | tuple[int | None, int | None] = 16,
  62. ) -> None:
  63. self.protocol: ClientProtocol
  64. super().__init__(
  65. nursery,
  66. stream,
  67. protocol,
  68. ping_interval=ping_interval,
  69. ping_timeout=ping_timeout,
  70. close_timeout=close_timeout,
  71. max_queue=max_queue,
  72. )
  73. self.response_rcvd = trio.Event()
  74. async def handshake(
  75. self,
  76. additional_headers: HeadersLike | None = None,
  77. user_agent_header: str | None = USER_AGENT,
  78. ) -> None:
  79. """
  80. Perform the opening handshake.
  81. """
  82. self.request = self.protocol.connect()
  83. if additional_headers is not None:
  84. self.request.headers.update(additional_headers)
  85. if user_agent_header is not None:
  86. self.request.headers.setdefault("User-Agent", user_agent_header)
  87. async with self.send_context(expected_state=CONNECTING):
  88. self.protocol.send_request(self.request)
  89. await race_events(self.response_rcvd, self.stream_closed)
  90. # self.protocol.handshake_exc is set when the connection is lost before
  91. # receiving a response, when the response cannot be parsed, or when the
  92. # response fails the handshake.
  93. if self.protocol.handshake_exc is not None:
  94. raise self.protocol.handshake_exc
  95. def process_event(self, event: Event) -> None:
  96. """
  97. Process one incoming event.
  98. """
  99. # First event - handshake response.
  100. if self.response is None:
  101. assert isinstance(event, Response)
  102. self.response = event
  103. self.response_rcvd.set()
  104. # Later events - frames.
  105. else:
  106. super().process_event(event)
  107. # This is spelled in lower case because it's exposed as a callable in the API.
  108. class connect:
  109. """
  110. Connect to the WebSocket server at ``uri``.
  111. :func:`connect` should be treated as an asynchronous context manager
  112. yielding a :class:`ClientConnection`, which can then receive and send
  113. messages::
  114. from websockets.trio.client import connect
  115. async with connect(...) as websocket:
  116. ...
  117. The connection is closed automatically when exiting the context.
  118. :func:`connect` can also be treated as an infinite asynchronous iterator
  119. to reconnect automatically on errors::
  120. async for websocket in connect(...):
  121. try:
  122. ...
  123. except websockets.exceptions.ConnectionClosed:
  124. continue
  125. If the connection fails with a transient error, it is retried with
  126. exponential backoff. If it fails with a fatal error, the exception is
  127. raised, breaking out of the loop.
  128. The connection is closed automatically after each iteration of the loop.
  129. :func:`connect` cannot be awaited directly. This is because it runs a task
  130. to manage the connection and Trio doesn't support spawning tasks without a
  131. context that ensures completion.
  132. Args:
  133. uri: URI of the WebSocket server.
  134. stream: Preexisting TCP stream. ``stream`` overrides the host and port
  135. from ``uri``. You may call :func:`~trio.open_tcp_stream` to create a
  136. suitable TCP stream.
  137. ssl: Configuration for enabling TLS on the connection.
  138. server_hostname: Host name for the TLS handshake. ``server_hostname``
  139. overrides the host name from ``uri``.
  140. origin: Value of the ``Origin`` header, for servers that require it.
  141. extensions: List of supported extensions, in order in which they
  142. should be negotiated and run.
  143. subprotocols: List of supported subprotocols, in order of decreasing
  144. preference.
  145. compression: The "permessage-deflate" extension is enabled by default.
  146. Set ``compression`` to :obj:`None` to disable it. See the
  147. :doc:`compression guide <../../topics/compression>` for details.
  148. additional_headers: Arbitrary HTTP headers to add to the handshake
  149. request.
  150. user_agent_header: Value of the ``User-Agent`` request header.
  151. It defaults to ``"Python/x.y.z websockets/X.Y"``.
  152. Setting it to :obj:`None` removes the header.
  153. proxy: If a proxy is configured, it is used by default. Set ``proxy``
  154. to :obj:`None` to disable the proxy or to the address of a proxy
  155. to override the system configuration. See the :doc:`proxy docs
  156. <../../topics/proxies>` for details.
  157. proxy_ssl: Configuration for enabling TLS on the proxy connection.
  158. proxy_server_hostname: Host name for the TLS handshake with the proxy.
  159. ``proxy_server_hostname`` overrides the host name from ``proxy``.
  160. process_exception: When reconnecting automatically, tell whether an
  161. error is transient or fatal. The default behavior is defined by
  162. :func:`~websockets.client.process_exception`. Refer to its
  163. documentation for details.
  164. open_timeout: Timeout for opening the connection in seconds.
  165. :obj:`None` disables the timeout.
  166. ping_interval: Interval between keepalive pings in seconds.
  167. :obj:`None` disables keepalive.
  168. ping_timeout: Timeout for keepalive pings in seconds.
  169. :obj:`None` disables timeouts.
  170. close_timeout: Timeout for closing the connection in seconds.
  171. :obj:`None` disables the timeout.
  172. reconnect_delays: Delays in seconds between reconnection attempts.
  173. Default is exponential backoff with 5s jitter, capped at 60s.
  174. max_size: Maximum size of incoming messages in bytes.
  175. :obj:`None` disables the limit. You may pass a ``(max_message_size,
  176. max_fragment_size)`` tuple to set different limits for messages and
  177. fragments when you expect long messages sent in short fragments.
  178. max_queue: High-water mark of the buffer where frames are received.
  179. It defaults to 16 frames. The low-water mark defaults to ``max_queue
  180. // 4``. You may pass a ``(high, low)`` tuple to set the high-water
  181. and low-water marks. If you want to disable flow control entirely,
  182. you may set it to ``None``, although that's a bad idea.
  183. logger: Logger for this client.
  184. It defaults to ``logging.getLogger("websockets.client")``.
  185. See the :doc:`logging guide <../../topics/logging>` for details.
  186. create_connection: Factory for the :class:`ClientConnection` managing
  187. the connection. Set it to a wrapper or a subclass to customize
  188. connection handling.
  189. Any other keyword arguments are passed to :func:`~trio.open_tcp_stream`.
  190. For example, you can set ``host`` and ``port`` to connect to a different
  191. host and port from those found in ``uri``. This only changes the destination
  192. of the TCP connection. The host name from ``uri`` is still used in the TLS
  193. handshake for secure connections and in the ``Host`` header.
  194. Raises:
  195. InvalidURI: If ``uri`` isn't a valid WebSocket URI.
  196. InvalidProxy: If ``proxy`` isn't a valid proxy.
  197. OSError: If the TCP connection fails.
  198. InvalidHandshake: If the opening handshake fails.
  199. TimeoutError: If the opening handshake times out.
  200. """
  201. # Arguments of type SSLContext don't render correctly in the documentation
  202. # because of https://github.com/sphinx-doc/sphinx/issues/13838.
  203. def __init__(
  204. self,
  205. uri: str,
  206. *,
  207. # TCP/TLS
  208. stream: trio.abc.Stream | None = None,
  209. ssl: ssl_module.SSLContext | None = None,
  210. server_hostname: str | None = None,
  211. # WebSocket
  212. origin: Origin | None = None,
  213. extensions: Sequence[ClientExtensionFactory] | None = None,
  214. subprotocols: Sequence[Subprotocol] | None = None,
  215. compression: str | None = "deflate",
  216. # HTTP
  217. additional_headers: HeadersLike | None = None,
  218. user_agent_header: str | None = USER_AGENT,
  219. proxy: str | Literal[True] | None = True,
  220. proxy_ssl: ssl_module.SSLContext | None = None,
  221. proxy_server_hostname: str | None = None,
  222. process_exception: Callable[[Exception], Exception | None] = process_exception,
  223. # Timeouts
  224. open_timeout: float | None = 10,
  225. ping_interval: float | None = 20,
  226. ping_timeout: float | None = 20,
  227. close_timeout: float | None = 10,
  228. reconnect_delays: Callable[[], Generator[float]] = backoff,
  229. # Limits
  230. max_size: int | None | tuple[int | None, int | None] = 2**20,
  231. max_queue: int | None | tuple[int | None, int | None] = 16,
  232. # Logging
  233. logger: LoggerLike | None = None,
  234. # Escape hatch for advanced customization
  235. create_connection: type[ClientConnection] | None = None,
  236. # Other keyword arguments are passed to trio.open_tcp_stream
  237. **kwargs: Any,
  238. ) -> None:
  239. self.uri = uri
  240. self.ws_uri = parse_uri(uri)
  241. if not self.ws_uri.secure and ssl is not None:
  242. raise ValueError("ssl argument is incompatible with a ws:// URI")
  243. if subprotocols is not None:
  244. validate_subprotocols(subprotocols)
  245. if compression == "deflate":
  246. extensions = enable_client_permessage_deflate(extensions)
  247. elif compression is not None:
  248. raise ValueError(f"unsupported compression: {compression}")
  249. if logger is None:
  250. logger = logging.getLogger("websockets.client")
  251. if create_connection is None:
  252. create_connection = ClientConnection
  253. self.stream = stream
  254. self.ssl = ssl
  255. self.server_hostname = server_hostname
  256. self.additional_headers = additional_headers
  257. self.user_agent_header = user_agent_header
  258. self.proxy = proxy
  259. self.proxy_ssl = proxy_ssl
  260. self.proxy_server_hostname = proxy_server_hostname
  261. self.process_exception = process_exception
  262. self.open_timeout = open_timeout
  263. self.reconnect_delays = reconnect_delays
  264. self.logger = logger
  265. self.create_connection = create_connection
  266. self.open_tcp_stream_kwargs = kwargs
  267. self.protocol_kwargs = dict(
  268. origin=origin,
  269. extensions=extensions,
  270. subprotocols=subprotocols,
  271. max_size=max_size,
  272. logger=logger,
  273. )
  274. self.connection_kwargs = dict(
  275. ping_interval=ping_interval,
  276. ping_timeout=ping_timeout,
  277. close_timeout=close_timeout,
  278. max_queue=max_queue,
  279. )
  280. async def open_tcp_stream(self) -> trio.abc.Stream:
  281. """Open a TCP or Unix connection to the server, possibly through a proxy."""
  282. kwargs = self.open_tcp_stream_kwargs.copy()
  283. unix = kwargs.pop("unix", False)
  284. proxy = self.proxy
  285. if unix:
  286. proxy = None
  287. if proxy is True:
  288. proxy = get_proxy(self.ws_uri)
  289. if unix:
  290. return await trio.open_unix_socket(kwargs.pop("path"))
  291. elif proxy is not None:
  292. proxy_parsed = parse_proxy(proxy)
  293. if proxy_parsed.scheme[:5] == "socks":
  294. return await connect_socks_proxy(
  295. proxy_parsed,
  296. self.ws_uri,
  297. # websockets is consistent with trio while python_socks is
  298. # consistent across implementations.
  299. local_addr=kwargs.pop("local_address", None),
  300. )
  301. elif proxy_parsed.scheme[:4] == "http":
  302. if proxy_parsed.scheme != "https" and self.proxy_ssl is not None:
  303. raise ValueError(
  304. "proxy_ssl argument is incompatible with an http:// proxy"
  305. )
  306. return await connect_http_proxy(
  307. proxy_parsed,
  308. self.ws_uri,
  309. user_agent_header=self.user_agent_header,
  310. ssl=self.proxy_ssl,
  311. server_hostname=self.proxy_server_hostname,
  312. **kwargs,
  313. )
  314. else:
  315. raise AssertionError("parse_proxy returned unsupported proxy")
  316. else: # proxy is None
  317. kwargs.setdefault("host", self.ws_uri.host)
  318. kwargs.setdefault("port", self.ws_uri.port)
  319. return await trio.open_tcp_stream(**kwargs)
  320. async def enable_tls(self, stream: trio.abc.Stream) -> trio.abc.Stream:
  321. """Enable TLS on the connection."""
  322. if self.ssl is None:
  323. ssl = ssl_module.create_default_context()
  324. else:
  325. ssl = self.ssl
  326. if self.server_hostname is None:
  327. server_hostname = self.ws_uri.host
  328. else:
  329. server_hostname = self.server_hostname
  330. ssl_stream = trio.SSLStream(
  331. stream,
  332. ssl,
  333. server_hostname=server_hostname,
  334. https_compatible=True,
  335. )
  336. await ssl_stream.do_handshake()
  337. return ssl_stream
  338. async def open_connection(self, nursery: trio.Nursery) -> ClientConnection:
  339. """Create a WebSocket connection."""
  340. if self.stream is None:
  341. stream = await self.open_tcp_stream()
  342. else:
  343. stream = self.stream
  344. try:
  345. if self.ws_uri.secure:
  346. stream = await self.enable_tls(stream)
  347. protocol = ClientProtocol(
  348. self.ws_uri,
  349. **self.protocol_kwargs, # type: ignore
  350. )
  351. # self.create_connection defaults to ClientConnection.
  352. connection = self.create_connection(
  353. nursery,
  354. stream,
  355. protocol,
  356. **self.connection_kwargs, # type: ignore
  357. )
  358. await connection.handshake(
  359. self.additional_headers,
  360. self.user_agent_header,
  361. )
  362. except trio.Cancelled:
  363. await trio.aclose_forcefully(stream)
  364. # The nursery running this coroutine was canceled.
  365. # The next checkpoint raises trio.Cancelled.
  366. # aclose_forcefully() never returns.
  367. raise AssertionError("nursery should be canceled")
  368. except Exception:
  369. # Always close the connection even though keep-alive is the default
  370. # in HTTP/1.1 because the current implementation ties opening the
  371. # TCP/TLS connection with initializing the WebSocket protocol.
  372. await trio.aclose_forcefully(stream)
  373. raise
  374. return connection
  375. def process_redirect(self, exc: Exception) -> Exception | str:
  376. """
  377. Determine whether a connection error is a redirect that can be followed.
  378. Return the new URI if it's a valid redirect. Else, return an exception.
  379. """
  380. if not (
  381. isinstance(exc, InvalidStatus)
  382. and exc.response.status_code
  383. in [
  384. 300, # Multiple Choices
  385. 301, # Moved Permanently
  386. 302, # Found
  387. 303, # See Other
  388. 307, # Temporary Redirect
  389. 308, # Permanent Redirect
  390. ]
  391. and "Location" in exc.response.headers
  392. ):
  393. return exc
  394. old_ws_uri = self.ws_uri
  395. new_uri = urllib.parse.urljoin(self.uri, exc.response.headers["Location"])
  396. new_ws_uri = parse_uri(new_uri)
  397. # If connect() received a stream, it is closed and cannot be reused.
  398. if self.stream is not None:
  399. return ValueError(
  400. f"cannot follow redirect to {new_uri} with a preexisting stream"
  401. )
  402. # TLS downgrade is forbidden.
  403. if old_ws_uri.secure and not new_ws_uri.secure:
  404. return SecurityError(f"cannot follow redirect to non-secure URI {new_uri}")
  405. # Apply restrictions to cross-origin redirects.
  406. if (
  407. old_ws_uri.secure != new_ws_uri.secure
  408. or old_ws_uri.host != new_ws_uri.host
  409. or old_ws_uri.port != new_ws_uri.port
  410. ):
  411. # Cross-origin redirects on Unix sockets don't quite make sense.
  412. if self.open_tcp_stream_kwargs.get("unix", False):
  413. return ValueError(
  414. f"cannot follow cross-origin redirect to {new_uri} "
  415. f"with a Unix socket"
  416. )
  417. # Cross-origin redirects when host and port are overridden are ill-defined.
  418. if (
  419. self.open_tcp_stream_kwargs.get("host") is not None
  420. or self.open_tcp_stream_kwargs.get("port") is not None
  421. ):
  422. return ValueError(
  423. f"cannot follow cross-origin redirect to {new_uri} "
  424. f"with an explicit host or port"
  425. )
  426. # Strip credentials to avoid leaking them to a different origin.
  427. if self.additional_headers is not None:
  428. self.additional_headers = Headers(
  429. (
  430. (key, value)
  431. for key, value in Headers(self.additional_headers).raw_items()
  432. if key.lower()
  433. not in ["authorization", "cookie", "proxy-authorization"]
  434. )
  435. )
  436. return new_uri
  437. async def connect(self, nursery: trio.Nursery) -> ClientConnection:
  438. try:
  439. with (
  440. trio.CancelScope()
  441. if self.open_timeout is None
  442. else trio.fail_after(self.open_timeout)
  443. ):
  444. for _ in range(MAX_REDIRECTS):
  445. try:
  446. connection = await self.open_connection(nursery)
  447. except Exception as exc:
  448. exc_or_uri = self.process_redirect(exc)
  449. if isinstance(exc_or_uri, Exception):
  450. # Response isn't a valid redirect; raise the exception.
  451. if exc_or_uri is exc:
  452. raise
  453. else:
  454. raise exc_or_uri from exc
  455. else:
  456. # Response is a valid redirect; follow it.
  457. self.uri = exc_or_uri
  458. self.ws_uri = parse_uri(exc_or_uri)
  459. continue
  460. else:
  461. connection.start_keepalive()
  462. return connection
  463. else:
  464. raise SecurityError(f"more than {MAX_REDIRECTS} redirects")
  465. except trio.TooSlowError as exc:
  466. # Re-raise exception with an informative error message.
  467. raise TimeoutError("timed out during opening handshake") from exc
  468. # Do not define __await__ for ... = await nursery.start(connect, ...)
  469. # because it doesn't look idiomatic in Trio.
  470. # async with connect(...) as ...: ...
  471. async def __aenter__(self) -> ClientConnection:
  472. await self.__aenter_nursery__()
  473. try:
  474. self.connection = await self.connect(self.nursery)
  475. return self.connection
  476. except BaseException as exc:
  477. await self.__aexit_nursery__(type(exc), exc, exc.__traceback__)
  478. raise AssertionError("expected __aexit_nursery__ to re-raise the exception")
  479. async def __aexit__(
  480. self,
  481. exc_type: type[BaseException] | None,
  482. exc_value: BaseException | None,
  483. traceback: TracebackType | None,
  484. ) -> None:
  485. try:
  486. try:
  487. await self.connection.aclose()
  488. finally:
  489. del self.connection
  490. finally:
  491. await self.__aexit_nursery__(exc_type, exc_value, traceback)
  492. async def __aenter_nursery__(self) -> None:
  493. if hasattr(self, "nursery_manager"):
  494. raise RuntimeError("connect() isn't reentrant")
  495. self.nursery_manager = trio.open_nursery()
  496. self.nursery = await self.nursery_manager.__aenter__()
  497. async def __aexit_nursery__(
  498. self,
  499. exc_type: type[BaseException] | None,
  500. exc_value: BaseException | None,
  501. traceback: TracebackType | None,
  502. ) -> None:
  503. # We need a nursery to start the recv_events and keepalive coroutines.
  504. # They aren't expected to raise exceptions; instead they catch and log
  505. # all unexpected errors. To keep the nursery an implementation detail,
  506. # unwrap exceptions raised by user code — per the second option here:
  507. # https://trio.readthedocs.io/en/stable/reference-core.html#designing-for-multiple-errors
  508. try:
  509. await self.nursery_manager.__aexit__(exc_type, exc_value, traceback)
  510. except BaseException as exc:
  511. assert isinstance(exc, BaseExceptionGroup)
  512. try:
  513. trio._util.raise_single_exception_from_group(exc)
  514. except trio._util.MultipleExceptionError:
  515. raise AssertionError(
  516. "unexpected multiple exceptions; please file a bug report"
  517. ) from exc
  518. finally:
  519. del self.nursery_manager
  520. # async for ... in connect(...):
  521. async def __aiter__(self) -> AsyncIterator[ClientConnection]:
  522. delays: Generator[float] | None = None
  523. while True:
  524. try:
  525. async with self as connection:
  526. yield connection
  527. except Exception as exc:
  528. # Determine whether the exception is retryable or fatal.
  529. # The API of process_exception is "return an exception or None";
  530. # "raise an exception" is also supported because it's a frequent
  531. # mistake. It isn't documented in order to keep the API simple.
  532. try:
  533. new_exc = self.process_exception(exc)
  534. except Exception as raised_exc:
  535. new_exc = raised_exc
  536. # The connection failed with a fatal error.
  537. # Raise the exception and exit the loop.
  538. if new_exc is exc:
  539. raise
  540. if new_exc is not None:
  541. raise new_exc from exc
  542. # The connection failed with a retryable error.
  543. # Start or continue backoff and reconnect.
  544. if delays is None:
  545. delays = self.reconnect_delays()
  546. delay = next(delays)
  547. self.logger.info(
  548. "connect failed; reconnecting in %.1f seconds: %s",
  549. delay,
  550. traceback.format_exception_only(exc)[0].strip(),
  551. )
  552. await trio.sleep(delay)
  553. else:
  554. # The connection succeeded. Reset backoff.
  555. delays = None
  556. def unix_connect(
  557. path: PathLike | None = None,
  558. uri: str | None = None,
  559. **kwargs: Any,
  560. ) -> connect:
  561. """
  562. Connect to a WebSocket server listening on a Unix socket.
  563. This function accepts the same keyword arguments as :func:`connect`.
  564. It's only available on Unix.
  565. It's mainly useful for debugging servers listening on Unix sockets.
  566. Args:
  567. path: File system path to the Unix socket.
  568. uri: URI of the WebSocket server. ``uri`` defaults to
  569. ``ws://localhost/`` or, when a ``ssl`` argument is provided, to
  570. ``wss://localhost/``.
  571. """
  572. stream = kwargs.get("stream")
  573. if path is None and stream is None:
  574. raise ValueError("missing path argument")
  575. elif path is not None and stream is not None:
  576. raise ValueError("path is incompatible with stream")
  577. if uri is None:
  578. if kwargs.get("ssl") is None:
  579. uri = "ws://localhost/"
  580. else:
  581. uri = "wss://localhost/"
  582. return connect(uri=uri, unix=True, path=path, **kwargs)
  583. try:
  584. from python_socks import ProxyType
  585. from python_socks.async_.trio import Proxy as SocksProxy
  586. except ImportError:
  587. async def connect_socks_proxy(
  588. proxy: Proxy,
  589. ws_uri: WebSocketURI,
  590. **kwargs: Any,
  591. ) -> trio.abc.Stream:
  592. raise ImportError("connecting through a SOCKS proxy requires python-socks")
  593. else:
  594. SOCKS_PROXY_TYPES = {
  595. "socks5h": ProxyType.SOCKS5,
  596. "socks5": ProxyType.SOCKS5,
  597. "socks4a": ProxyType.SOCKS4,
  598. "socks4": ProxyType.SOCKS4,
  599. }
  600. SOCKS_PROXY_RDNS = {
  601. "socks5h": True,
  602. "socks5": False,
  603. "socks4a": True,
  604. "socks4": False,
  605. }
  606. async def connect_socks_proxy(
  607. proxy: Proxy,
  608. ws_uri: WebSocketURI,
  609. **kwargs: Any,
  610. ) -> trio.abc.Stream:
  611. """Connect via a SOCKS proxy and return the socket."""
  612. socks_proxy = SocksProxy(
  613. SOCKS_PROXY_TYPES[proxy.scheme],
  614. proxy.host,
  615. proxy.port,
  616. proxy.username,
  617. proxy.password,
  618. SOCKS_PROXY_RDNS[proxy.scheme],
  619. )
  620. # connect() is documented to raise OSError.
  621. # socks_proxy.connect() re-raises trio.TooSlowError as ProxyTimeoutError.
  622. # Wrap other exceptions in ProxyError, a subclass of InvalidHandshake.
  623. try:
  624. return trio.SocketStream(
  625. await socks_proxy.connect(ws_uri.host, ws_uri.port, **kwargs)
  626. )
  627. except OSError:
  628. raise
  629. except Exception as exc:
  630. raise ProxyError("failed to connect to SOCKS proxy") from exc
  631. async def read_connect_response(stream: trio.abc.Stream) -> Response:
  632. reader = StreamReader()
  633. parser = Response.parse(
  634. reader.read_line,
  635. reader.read_exact,
  636. reader.read_to_eof,
  637. proxy=True,
  638. )
  639. try:
  640. while True:
  641. data = await stream.receive_some(4096)
  642. if data:
  643. reader.feed_data(data)
  644. else:
  645. reader.feed_eof()
  646. next(parser)
  647. except StopIteration as exc:
  648. assert isinstance(exc.value, Response) # help mypy
  649. response = exc.value
  650. if 200 <= response.status_code < 300:
  651. return response
  652. else:
  653. raise InvalidProxyStatus(response)
  654. except Exception as exc:
  655. raise InvalidProxyMessage(
  656. "did not receive a valid HTTP response from proxy"
  657. ) from exc
  658. async def connect_http_proxy(
  659. proxy: Proxy,
  660. ws_uri: WebSocketURI,
  661. *,
  662. user_agent_header: str | None = None,
  663. ssl: ssl_module.SSLContext | None = None,
  664. server_hostname: str | None = None,
  665. **kwargs: Any,
  666. ) -> trio.abc.Stream:
  667. stream: trio.abc.Stream
  668. stream = await trio.open_tcp_stream(proxy.host, proxy.port, **kwargs)
  669. try:
  670. # Initialize TLS wrapper and perform TLS handshake
  671. if proxy.scheme == "https":
  672. if ssl is None:
  673. ssl = ssl_module.create_default_context()
  674. if server_hostname is None:
  675. server_hostname = proxy.host
  676. ssl_stream = trio.SSLStream(
  677. stream,
  678. ssl,
  679. server_hostname=server_hostname,
  680. https_compatible=True,
  681. )
  682. await ssl_stream.do_handshake()
  683. stream = ssl_stream
  684. # Send CONNECT request to the proxy and read response.
  685. request = prepare_connect_request(proxy, ws_uri, user_agent_header)
  686. await stream.send_all(request)
  687. await read_connect_response(stream)
  688. except (trio.Cancelled, Exception):
  689. await trio.aclose_forcefully(stream)
  690. raise
  691. return stream