"""Shared SSRF guard: resolve a URL's host and confirm it is publicly routable. Used by the ``http_request`` tool (which additionally pins the connection and re-validates every redirect hop) and by server-side image fetching, so an untrusted URL can't reach internal services or the cloud metadata endpoint. """ from __future__ import annotations import ipaddress import socket from typing import Any from urllib.parse import urljoin, urlparse, urlunparse import httpx DEFAULT_MAX_REDIRECTS = 5 _REDIRECT_CODES = {301, 302, 303, 307, 308} def resolve_and_validate(url: str) -> tuple[bool, str, str | None, list | None]: """Resolve a URL's hostname and check every address is safe to contact. Returns (is_safe, reason, hostname, addr_infos). When safe, the caller pins the connection to one of ``addr_infos`` so the request cannot pick up a different (e.g. DNS-rebound) address after validation. """ try: parsed = urlparse(url) if parsed.scheme not in {"http", "https"}: return False, f"Unsupported URL scheme: {parsed.scheme or ''}", None, None hostname = parsed.hostname if not hostname: return False, "Could not parse hostname from URL", None, None try: addr_infos = socket.getaddrinfo(hostname, None) except socket.gaierror: return False, f"Could not resolve hostname: {hostname}", hostname, None if not addr_infos: return False, f"Could not resolve hostname: {hostname}", hostname, None for addr_info in addr_infos: ip_str = addr_info[4][0] try: ip = ipaddress.ip_address(ip_str) except ValueError: return False, f"Could not parse resolved address: {ip_str}", hostname, None # Unwrap IPv4-mapped IPv6 (e.g. ::ffff:127.0.0.1) so a mapped private # address can't slip past the check, then block anything that isn't # publicly routable (covers private/loopback/link-local/reserved/ # unspecified/multicast and the cloud metadata 169.254.0.0/16 range). if isinstance(ip, ipaddress.IPv6Address) and ip.ipv4_mapped is not None: ip = ip.ipv4_mapped if not ip.is_global: return False, f"URL resolves to blocked address: {ip_str}", hostname, None return True, "", hostname, addr_infos except Exception as e: # noqa: BLE001 return False, f"URL validation error: {e}", None, None def is_url_safe(url: str) -> tuple[bool, str]: """Check if a URL is safe to request (not targeting private/internal networks).""" is_safe, reason, _, _ = resolve_and_validate(url) return is_safe, reason def pinned_url(url: str, ip: str) -> str: """Rewrite ``url`` so the connection targets ``ip`` while keeping the path/query. The original hostname is preserved separately for the ``Host`` header and TLS SNI/cert verification (via httpx's ``sni_hostname`` request extension). """ parsed = urlparse(url) host_literal = f"[{ip}]" if ":" in ip else ip netloc = f"{host_literal}:{parsed.port}" if parsed.port else host_literal return urlunparse(parsed._replace(netloc=netloc)) async def request_with_safe_redirects( client: httpx.AsyncClient, method: str, url: str, *, max_redirects: int = DEFAULT_MAX_REDIRECTS, strip_auth_on_redirect: bool = False, **kwargs: Any, ) -> tuple[httpx.Response | None, tuple[str, str] | None]: """Issue a request, validating every redirect target before following it. The hostname is resolved once per hop and the connection is pinned to the validated IP, closing the DNS-rebinding race where a controlled resolver returns a public IP at validation time and a private IP at connect time. Returns ``(response, None)`` on success, or ``(None, (blocked_url, reason))`` when a hop fails validation or the redirect budget is exhausted. When ``strip_auth_on_redirect`` is set, the caller's ``Authorization`` header is dropped once the request leaves the original URL, so a bearer token can't be replayed to a redirect target the caller never chose to authenticate to. """ current_method = method.upper() current_url = url request_kwargs = dict(kwargs) # Pop caller headers/extensions ONCE so they're reused on every redirect hop # (the per-hop Host + SNI are layered on top each time). Popping inside the # loop dropped the caller's Authorization/Accept/etc. on the first redirect. caller_headers = dict(request_kwargs.pop("headers", None) or {}) caller_extensions = dict(request_kwargs.pop("extensions", None) or {}) for redirect_count in range(max_redirects + 1): is_safe, reason, hostname, addr_infos = resolve_and_validate(current_url) if not is_safe or hostname is None or addr_infos is None: return None, (current_url, reason) pinned_ip = addr_infos[0][4][0] parsed = urlparse(current_url) headers = {**caller_headers, "Host": parsed.netloc} extensions = {**caller_extensions, "sni_hostname": hostname} response = await client.request( current_method, pinned_url(current_url, pinned_ip), follow_redirects=False, headers=headers, extensions=extensions, **request_kwargs, ) if response.status_code not in _REDIRECT_CODES: return response, None location = response.headers.get("Location") if not location: return response, None if redirect_count == max_redirects: return None, (current_url, "Too many redirects") current_url = urljoin(current_url, location) if strip_auth_on_redirect: caller_headers = { k: v for k, v in caller_headers.items() if k.lower() != "authorization" } if response.status_code == 303 or ( response.status_code in {301, 302} and current_method not in {"GET", "HEAD"} ): current_method = "GET" request_kwargs.pop("data", None) request_kwargs.pop("content", None) request_kwargs.pop("json", None) return None, (current_url, "Too many redirects")