diff --git a/src/paperless/network.py b/src/paperless/network.py index 431b6c59a..cabdb1519 100644 --- a/src/paperless/network.py +++ b/src/paperless/network.py @@ -498,31 +498,6 @@ def create_guarded_async_httpx_client( ) -def resolve_hostname_ips(hostname: str) -> list[str]: - try: - addr_info = socket.getaddrinfo(hostname, None) - except socket.gaierror as e: - raise ValueError(f"Could not resolve hostname: {hostname}") from e - - ips = [info[4][0] for info in addr_info if info and info[4]] - if not ips: - raise ValueError(f"Could not resolve hostname: {hostname}") - return ips - - -def format_host_for_url(host: str) -> str: - """ - Format IP address for URL use (wrap IPv6 in brackets). - """ - try: - ip_obj = ipaddress.ip_address(host) - if ip_obj.version == 6: - return f"[{host}]" - return host - except ValueError: - return host - - # urllib3 treats a backslash as ending the authority while urlparse and httpx do # not, so the host checked here could differ from the one that is dialled. # Control and whitespace characters are refused for the same reason. @@ -576,121 +551,3 @@ def validate_outbound_http_url( raise ValueError(blocked_message(e)) from e return parsed - - -def _rewrite_request_to_pinned_ip( - request: httpx.Request, - *, - allow_internal: bool, -) -> httpx.Request: - hostname = request.url.host - - if not hostname: - raise httpx.ConnectError("No hostname in request URL") - - try: - ips = resolve_hostname_ips(hostname) - except ValueError as e: - raise httpx.ConnectError(str(e)) from e - - if not allow_internal: - for ip_str in ips: - if not is_public_ip(ipaddress.ip_address(ip_str)): - raise httpx.ConnectError( - f"Connection blocked: {hostname} resolves to a non-public address", - ) - - ip_str = ips[0] - formatted_ip = format_host_for_url(ip_str) - - new_headers = httpx.Headers(request.headers) - if "host" in new_headers: - del new_headers["host"] - host_header = format_host_for_url(hostname) - default_port = 443 if request.url.scheme == "https" else 80 - if request.url.port and request.url.port != default_port: - host_header = f"{host_header}:{request.url.port}" - new_headers["Host"] = host_header - new_url = request.url.copy_with(host=formatted_ip) - - rewritten_request = httpx.Request( - method=request.method, - url=new_url, - headers=new_headers, - stream=request.stream, - extensions=request.extensions, - ) - rewritten_request.extensions["sni_hostname"] = hostname - - return rewritten_request - - -class PinnedHostHTTPTransport(httpx.HTTPTransport): - """ - HTTP transport that resolves/validates hostnames per request and connects to - a vetted IP while preserving the original Host header and TLS SNI hostname. - """ - - def __init__( - self, - *args, - allow_internal: bool = False, - **kwargs, - ) -> None: - super().__init__(*args, **kwargs) - self.allow_internal = allow_internal - - def handle_request(self, request: httpx.Request) -> httpx.Response: - request = _rewrite_request_to_pinned_ip( - request, - allow_internal=self.allow_internal, - ) - return super().handle_request(request) - - -class PinnedHostAsyncHTTPTransport(httpx.AsyncHTTPTransport): - """ - Async variant of PinnedHostHTTPTransport. - """ - - def __init__( - self, - *args, - allow_internal: bool = False, - **kwargs, - ) -> None: - super().__init__(*args, **kwargs) - self.allow_internal = allow_internal - - async def handle_async_request(self, request: httpx.Request) -> httpx.Response: - request = _rewrite_request_to_pinned_ip( - request, - allow_internal=self.allow_internal, - ) - return await super().handle_async_request(request) - - -def create_pinned_httpx_client( - url: str, - *, - allow_internal: bool = False, - **kwargs, -) -> httpx.Client: - validate_outbound_http_url(url, allow_internal=allow_internal) - return httpx.Client( - transport=PinnedHostHTTPTransport(allow_internal=allow_internal), - **kwargs, - ) - - -def create_pinned_async_httpx_client( - url: str, - *, - allow_internal: bool = False, - **kwargs, -) -> httpx.AsyncClient: - validate_outbound_http_url(url, allow_internal=allow_internal) - return httpx.AsyncClient( - transport=PinnedHostAsyncHTTPTransport(allow_internal=allow_internal), - **kwargs, - ) diff --git a/src/paperless/tests/test_network.py b/src/paperless/tests/test_network.py index bc9b21e12..ed606ce01 100644 --- a/src/paperless/tests/test_network.py +++ b/src/paperless/tests/test_network.py @@ -4,7 +4,6 @@ import pickle import socket from collections.abc import Iterable from typing import Any -from unittest import mock from unittest.mock import MagicMock import anyio @@ -22,7 +21,6 @@ from paperless.network import GuardedAsyncHTTPTransport from paperless.network import GuardedHTTPTransport from paperless.network import HostResolutionError from paperless.network import OutboundRequestBlockedError -from paperless.network import PinnedHostHTTPTransport from paperless.network import _GuardedAsyncBackend from paperless.network import _GuardedSyncBackend from paperless.network import aresolve_public_addresses @@ -34,50 +32,6 @@ from paperless.network import resolve_public_addresses from paperless.network import validate_outbound_http_url -def test_pinned_host_transport_blocks_internal_rebinding(): - transport = PinnedHostHTTPTransport(allow_internal=False) - request = httpx.Request("GET", "http://example.com/test") - - with ( - mock.patch( - "paperless.network.resolve_hostname_ips", - return_value=["127.0.0.1"], - ), - pytest.raises(httpx.ConnectError, match="non-public address"), - ): - transport.handle_request(request) - - -def test_pinned_host_transport_rewrites_to_vetted_ip(): - transport = PinnedHostHTTPTransport(allow_internal=False) - request = httpx.Request("GET", "https://example.com:8443/test") - - def assert_rewritten_request( - self, - rewritten_request, - ): - assert str(rewritten_request.url) == "https://93.184.216.34:8443/test" - assert rewritten_request.headers["Host"] == "example.com:8443" - assert rewritten_request.extensions["sni_hostname"] == "example.com" - return httpx.Response(200, request=rewritten_request) - - with ( - mock.patch( - "paperless.network.resolve_hostname_ips", - return_value=["93.184.216.34"], - ), - mock.patch.object( - httpx.HTTPTransport, - "handle_request", - autospec=True, - side_effect=assert_rewritten_request, - ), - ): - response = transport.handle_request(request) - - assert response.status_code == 200 - - class TestIsPublicIp: @pytest.mark.parametrize( "address",