http_proxy.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367
  1. from __future__ import annotations
  2. import base64
  3. import logging
  4. import ssl
  5. import typing
  6. from .._backends.base import SOCKET_OPTION, AsyncNetworkBackend
  7. from .._exceptions import ProxyError
  8. from .._models import (
  9. URL,
  10. Origin,
  11. Request,
  12. Response,
  13. enforce_bytes,
  14. enforce_headers,
  15. enforce_url,
  16. )
  17. from .._ssl import default_ssl_context
  18. from .._synchronization import AsyncLock
  19. from .._trace import Trace
  20. from .connection import AsyncHTTPConnection
  21. from .connection_pool import AsyncConnectionPool
  22. from .http11 import AsyncHTTP11Connection
  23. from .interfaces import AsyncConnectionInterface
  24. ByteOrStr = typing.Union[bytes, str]
  25. HeadersAsSequence = typing.Sequence[typing.Tuple[ByteOrStr, ByteOrStr]]
  26. HeadersAsMapping = typing.Mapping[ByteOrStr, ByteOrStr]
  27. logger = logging.getLogger("httpcore.proxy")
  28. def merge_headers(
  29. default_headers: typing.Sequence[tuple[bytes, bytes]] | None = None,
  30. override_headers: typing.Sequence[tuple[bytes, bytes]] | None = None,
  31. ) -> list[tuple[bytes, bytes]]:
  32. """
  33. Append default_headers and override_headers, de-duplicating if a key exists
  34. in both cases.
  35. """
  36. default_headers = [] if default_headers is None else list(default_headers)
  37. override_headers = [] if override_headers is None else list(override_headers)
  38. has_override = set(key.lower() for key, value in override_headers)
  39. default_headers = [
  40. (key, value)
  41. for key, value in default_headers
  42. if key.lower() not in has_override
  43. ]
  44. return default_headers + override_headers
  45. class AsyncHTTPProxy(AsyncConnectionPool): # pragma: nocover
  46. """
  47. A connection pool that sends requests via an HTTP proxy.
  48. """
  49. def __init__(
  50. self,
  51. proxy_url: URL | bytes | str,
  52. proxy_auth: tuple[bytes | str, bytes | str] | None = None,
  53. proxy_headers: HeadersAsMapping | HeadersAsSequence | None = None,
  54. ssl_context: ssl.SSLContext | None = None,
  55. proxy_ssl_context: ssl.SSLContext | None = None,
  56. max_connections: int | None = 10,
  57. max_keepalive_connections: int | None = None,
  58. keepalive_expiry: float | None = None,
  59. http1: bool = True,
  60. http2: bool = False,
  61. retries: int = 0,
  62. local_address: str | None = None,
  63. uds: str | None = None,
  64. network_backend: AsyncNetworkBackend | None = None,
  65. socket_options: typing.Iterable[SOCKET_OPTION] | None = None,
  66. ) -> None:
  67. """
  68. A connection pool for making HTTP requests.
  69. Parameters:
  70. proxy_url: The URL to use when connecting to the proxy server.
  71. For example `"http://127.0.0.1:8080/"`.
  72. proxy_auth: Any proxy authentication as a two-tuple of
  73. (username, password). May be either bytes or ascii-only str.
  74. proxy_headers: Any HTTP headers to use for the proxy requests.
  75. For example `{"Proxy-Authorization": "Basic <username>:<password>"}`.
  76. ssl_context: An SSL context to use for verifying connections.
  77. If not specified, the default `httpcore.default_ssl_context()`
  78. will be used.
  79. proxy_ssl_context: The same as `ssl_context`, but for a proxy server rather than a remote origin.
  80. max_connections: The maximum number of concurrent HTTP connections that
  81. the pool should allow. Any attempt to send a request on a pool that
  82. would exceed this amount will block until a connection is available.
  83. max_keepalive_connections: The maximum number of idle HTTP connections
  84. that will be maintained in the pool.
  85. keepalive_expiry: The duration in seconds that an idle HTTP connection
  86. may be maintained for before being expired from the pool.
  87. http1: A boolean indicating if HTTP/1.1 requests should be supported
  88. by the connection pool. Defaults to True.
  89. http2: A boolean indicating if HTTP/2 requests should be supported by
  90. the connection pool. Defaults to False.
  91. retries: The maximum number of retries when trying to establish
  92. a connection.
  93. local_address: Local address to connect from. Can also be used to
  94. connect using a particular address family. Using
  95. `local_address="0.0.0.0"` will connect using an `AF_INET` address
  96. (IPv4), while using `local_address="::"` will connect using an
  97. `AF_INET6` address (IPv6).
  98. uds: Path to a Unix Domain Socket to use instead of TCP sockets.
  99. network_backend: A backend instance to use for handling network I/O.
  100. """
  101. super().__init__(
  102. ssl_context=ssl_context,
  103. max_connections=max_connections,
  104. max_keepalive_connections=max_keepalive_connections,
  105. keepalive_expiry=keepalive_expiry,
  106. http1=http1,
  107. http2=http2,
  108. network_backend=network_backend,
  109. retries=retries,
  110. local_address=local_address,
  111. uds=uds,
  112. socket_options=socket_options,
  113. )
  114. self._proxy_url = enforce_url(proxy_url, name="proxy_url")
  115. if (
  116. self._proxy_url.scheme == b"http" and proxy_ssl_context is not None
  117. ): # pragma: no cover
  118. raise RuntimeError(
  119. "The `proxy_ssl_context` argument is not allowed for the http scheme"
  120. )
  121. self._ssl_context = ssl_context
  122. self._proxy_ssl_context = proxy_ssl_context
  123. self._proxy_headers = enforce_headers(proxy_headers, name="proxy_headers")
  124. if proxy_auth is not None:
  125. username = enforce_bytes(proxy_auth[0], name="proxy_auth")
  126. password = enforce_bytes(proxy_auth[1], name="proxy_auth")
  127. userpass = username + b":" + password
  128. authorization = b"Basic " + base64.b64encode(userpass)
  129. self._proxy_headers = [
  130. (b"Proxy-Authorization", authorization)
  131. ] + self._proxy_headers
  132. def create_connection(self, origin: Origin) -> AsyncConnectionInterface:
  133. if origin.scheme == b"http":
  134. return AsyncForwardHTTPConnection(
  135. proxy_origin=self._proxy_url.origin,
  136. proxy_headers=self._proxy_headers,
  137. remote_origin=origin,
  138. keepalive_expiry=self._keepalive_expiry,
  139. network_backend=self._network_backend,
  140. proxy_ssl_context=self._proxy_ssl_context,
  141. )
  142. return AsyncTunnelHTTPConnection(
  143. proxy_origin=self._proxy_url.origin,
  144. proxy_headers=self._proxy_headers,
  145. remote_origin=origin,
  146. ssl_context=self._ssl_context,
  147. proxy_ssl_context=self._proxy_ssl_context,
  148. keepalive_expiry=self._keepalive_expiry,
  149. http1=self._http1,
  150. http2=self._http2,
  151. network_backend=self._network_backend,
  152. )
  153. class AsyncForwardHTTPConnection(AsyncConnectionInterface):
  154. def __init__(
  155. self,
  156. proxy_origin: Origin,
  157. remote_origin: Origin,
  158. proxy_headers: HeadersAsMapping | HeadersAsSequence | None = None,
  159. keepalive_expiry: float | None = None,
  160. network_backend: AsyncNetworkBackend | None = None,
  161. socket_options: typing.Iterable[SOCKET_OPTION] | None = None,
  162. proxy_ssl_context: ssl.SSLContext | None = None,
  163. ) -> None:
  164. self._connection = AsyncHTTPConnection(
  165. origin=proxy_origin,
  166. keepalive_expiry=keepalive_expiry,
  167. network_backend=network_backend,
  168. socket_options=socket_options,
  169. ssl_context=proxy_ssl_context,
  170. )
  171. self._proxy_origin = proxy_origin
  172. self._proxy_headers = enforce_headers(proxy_headers, name="proxy_headers")
  173. self._remote_origin = remote_origin
  174. async def handle_async_request(self, request: Request) -> Response:
  175. headers = merge_headers(self._proxy_headers, request.headers)
  176. url = URL(
  177. scheme=self._proxy_origin.scheme,
  178. host=self._proxy_origin.host,
  179. port=self._proxy_origin.port,
  180. target=bytes(request.url),
  181. )
  182. proxy_request = Request(
  183. method=request.method,
  184. url=url,
  185. headers=headers,
  186. content=request.stream,
  187. extensions=request.extensions,
  188. )
  189. return await self._connection.handle_async_request(proxy_request)
  190. def can_handle_request(self, origin: Origin) -> bool:
  191. return origin == self._remote_origin
  192. async def aclose(self) -> None:
  193. await self._connection.aclose()
  194. def info(self) -> str:
  195. return self._connection.info()
  196. def is_available(self) -> bool:
  197. return self._connection.is_available()
  198. def has_expired(self) -> bool:
  199. return self._connection.has_expired()
  200. def is_idle(self) -> bool:
  201. return self._connection.is_idle()
  202. def is_closed(self) -> bool:
  203. return self._connection.is_closed()
  204. def __repr__(self) -> str:
  205. return f"<{self.__class__.__name__} [{self.info()}]>"
  206. class AsyncTunnelHTTPConnection(AsyncConnectionInterface):
  207. def __init__(
  208. self,
  209. proxy_origin: Origin,
  210. remote_origin: Origin,
  211. ssl_context: ssl.SSLContext | None = None,
  212. proxy_ssl_context: ssl.SSLContext | None = None,
  213. proxy_headers: typing.Sequence[tuple[bytes, bytes]] | None = None,
  214. keepalive_expiry: float | None = None,
  215. http1: bool = True,
  216. http2: bool = False,
  217. network_backend: AsyncNetworkBackend | None = None,
  218. socket_options: typing.Iterable[SOCKET_OPTION] | None = None,
  219. ) -> None:
  220. self._connection: AsyncConnectionInterface = AsyncHTTPConnection(
  221. origin=proxy_origin,
  222. keepalive_expiry=keepalive_expiry,
  223. network_backend=network_backend,
  224. socket_options=socket_options,
  225. ssl_context=proxy_ssl_context,
  226. )
  227. self._proxy_origin = proxy_origin
  228. self._remote_origin = remote_origin
  229. self._ssl_context = ssl_context
  230. self._proxy_ssl_context = proxy_ssl_context
  231. self._proxy_headers = enforce_headers(proxy_headers, name="proxy_headers")
  232. self._keepalive_expiry = keepalive_expiry
  233. self._http1 = http1
  234. self._http2 = http2
  235. self._connect_lock = AsyncLock()
  236. self._connected = False
  237. async def handle_async_request(self, request: Request) -> Response:
  238. timeouts = request.extensions.get("timeout", {})
  239. timeout = timeouts.get("connect", None)
  240. async with self._connect_lock:
  241. if not self._connected:
  242. target = b"%b:%d" % (self._remote_origin.host, self._remote_origin.port)
  243. connect_url = URL(
  244. scheme=self._proxy_origin.scheme,
  245. host=self._proxy_origin.host,
  246. port=self._proxy_origin.port,
  247. target=target,
  248. )
  249. connect_headers = merge_headers(
  250. [(b"Host", target), (b"Accept", b"*/*")], self._proxy_headers
  251. )
  252. connect_request = Request(
  253. method=b"CONNECT",
  254. url=connect_url,
  255. headers=connect_headers,
  256. extensions=request.extensions,
  257. )
  258. connect_response = await self._connection.handle_async_request(
  259. connect_request
  260. )
  261. if connect_response.status < 200 or connect_response.status > 299:
  262. reason_bytes = connect_response.extensions.get("reason_phrase", b"")
  263. reason_str = reason_bytes.decode("ascii", errors="ignore")
  264. msg = "%d %s" % (connect_response.status, reason_str)
  265. await self._connection.aclose()
  266. raise ProxyError(msg)
  267. stream = connect_response.extensions["network_stream"]
  268. # Upgrade the stream to SSL
  269. ssl_context = (
  270. default_ssl_context()
  271. if self._ssl_context is None
  272. else self._ssl_context
  273. )
  274. alpn_protocols = ["http/1.1", "h2"] if self._http2 else ["http/1.1"]
  275. ssl_context.set_alpn_protocols(alpn_protocols)
  276. kwargs = {
  277. "ssl_context": ssl_context,
  278. "server_hostname": self._remote_origin.host.decode("ascii"),
  279. "timeout": timeout,
  280. }
  281. async with Trace("start_tls", logger, request, kwargs) as trace:
  282. stream = await stream.start_tls(**kwargs)
  283. trace.return_value = stream
  284. # Determine if we should be using HTTP/1.1 or HTTP/2
  285. ssl_object = stream.get_extra_info("ssl_object")
  286. http2_negotiated = (
  287. ssl_object is not None
  288. and ssl_object.selected_alpn_protocol() == "h2"
  289. )
  290. # Create the HTTP/1.1 or HTTP/2 connection
  291. if http2_negotiated or (self._http2 and not self._http1):
  292. from .http2 import AsyncHTTP2Connection
  293. self._connection = AsyncHTTP2Connection(
  294. origin=self._remote_origin,
  295. stream=stream,
  296. keepalive_expiry=self._keepalive_expiry,
  297. )
  298. else:
  299. self._connection = AsyncHTTP11Connection(
  300. origin=self._remote_origin,
  301. stream=stream,
  302. keepalive_expiry=self._keepalive_expiry,
  303. )
  304. self._connected = True
  305. return await self._connection.handle_async_request(request)
  306. def can_handle_request(self, origin: Origin) -> bool:
  307. return origin == self._remote_origin
  308. async def aclose(self) -> None:
  309. await self._connection.aclose()
  310. def info(self) -> str:
  311. return self._connection.info()
  312. def is_available(self) -> bool:
  313. return self._connection.is_available()
  314. def has_expired(self) -> bool:
  315. return self._connection.has_expired()
  316. def is_idle(self) -> bool:
  317. return self._connection.is_idle()
  318. def is_closed(self) -> bool:
  319. return self._connection.is_closed()
  320. def __repr__(self) -> str:
  321. return f"<{self.__class__.__name__} [{self.info()}]>"