97 lines
3.2 KiB
Python
97 lines
3.2 KiB
Python
from __future__ import annotations
|
|
|
|
import threading
|
|
from collections.abc import Iterable
|
|
|
|
import httpcore
|
|
import httpx
|
|
|
|
|
|
class PinnedAddressBook:
|
|
"""Thread-safe, short-lived host-to-validated-IP mapping for direct HTTP connections."""
|
|
|
|
def __init__(self) -> None:
|
|
self._addresses: dict[tuple[str, int], tuple[str, ...]] = {}
|
|
self._lock = threading.Lock()
|
|
|
|
def pin(self, host: str, port: int, addresses: Iterable[str]) -> None:
|
|
values = tuple(addresses)
|
|
if not values:
|
|
raise ValueError("At least one validated address is required.")
|
|
with self._lock:
|
|
self._addresses[(host.rstrip(".").lower(), port)] = values
|
|
|
|
def addresses_for(self, host: str, port: int) -> tuple[str, ...]:
|
|
with self._lock:
|
|
addresses = self._addresses.get((host.rstrip(".").lower(), port))
|
|
if not addresses:
|
|
raise httpcore.ConnectError(f"No validated address is available for {host}:{port}")
|
|
return addresses
|
|
|
|
|
|
class PinnedNetworkBackend(httpcore.NetworkBackend):
|
|
"""Dials only validated numeric addresses while httpcore retains the logical hostname."""
|
|
|
|
def __init__(self, address_book: PinnedAddressBook) -> None:
|
|
self._address_book = address_book
|
|
self._backend = httpcore.SyncBackend()
|
|
|
|
def connect_tcp(
|
|
self,
|
|
host: str,
|
|
port: int,
|
|
timeout: float | None = None,
|
|
local_address: str | None = None,
|
|
socket_options: Iterable[httpcore.SOCKET_OPTION] | None = None,
|
|
) -> httpcore.NetworkStream:
|
|
last_error: Exception | None = None
|
|
for address in self._address_book.addresses_for(host, port):
|
|
try:
|
|
return self._backend.connect_tcp(
|
|
address,
|
|
port,
|
|
timeout=timeout,
|
|
local_address=local_address,
|
|
socket_options=socket_options,
|
|
)
|
|
except (httpcore.ConnectError, httpcore.ConnectTimeout) as error:
|
|
last_error = error
|
|
if last_error:
|
|
raise last_error
|
|
raise httpcore.ConnectError(f"Could not connect to a validated address for {host}:{port}")
|
|
|
|
def connect_unix_socket(
|
|
self,
|
|
path: str,
|
|
timeout: float | None = None,
|
|
socket_options: Iterable[httpcore.SOCKET_OPTION] | None = None,
|
|
) -> httpcore.NetworkStream:
|
|
return self._backend.connect_unix_socket(
|
|
path, timeout=timeout, socket_options=socket_options
|
|
)
|
|
|
|
def sleep(self, seconds: float) -> None:
|
|
self._backend.sleep(seconds)
|
|
|
|
|
|
class PinnedHTTPTransport(httpx.HTTPTransport):
|
|
"""HTTPX transport that prevents a second, uncontrolled DNS lookup at connect time."""
|
|
|
|
def __init__(
|
|
self,
|
|
address_book: PinnedAddressBook,
|
|
limits: httpx.Limits,
|
|
proxy: str | None = None,
|
|
) -> None:
|
|
super().__init__(
|
|
verify=True,
|
|
trust_env=False,
|
|
http1=True,
|
|
http2=False,
|
|
limits=limits,
|
|
proxy=proxy,
|
|
retries=0,
|
|
)
|
|
# HTTPX constructs a direct httpcore ConnectionPool for this transport.
|
|
self._pool._network_backend = PinnedNetworkBackend(address_book)
|