__init__
This commit is contained in:
@@ -0,0 +1,96 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user