import asyncio import ipaddress import socket from collections.abc import Awaitable, Callable from dataclasses import dataclass from urllib.parse import SplitResult, urlsplit, urlunsplit class TargetPolicyError(Exception): pass @dataclass(frozen=True) class ValidatedTarget: url: str hostname: str port: int addresses: tuple[str, ...] Resolver = Callable[[str, int], Awaitable[list[tuple[object, ...]]]] def redact_url(value: str) -> str: parsed = urlsplit(value) host = parsed.hostname or "invalid-host" if ":" in host: host = f"[{host}]" port = f":{parsed.port}" if parsed.port else "" return urlunsplit((parsed.scheme.lower(), f"{host}{port}", parsed.path or "/", "", "")) def _public_address(value: str) -> bool: address = ipaddress.ip_address(value) return bool( address.is_global and not address.is_private and not address.is_loopback and not address.is_link_local and not address.is_multicast and not address.is_reserved and not address.is_unspecified ) async def system_resolver(host: str, port: int) -> list[tuple[object, ...]]: loop = asyncio.get_running_loop() return await loop.getaddrinfo(host, port, type=socket.SOCK_STREAM, proto=socket.IPPROTO_TCP) async def validate_target(url: str, resolver: Resolver = system_resolver) -> ValidatedTarget: try: parsed: SplitResult = urlsplit(url) if parsed.scheme.lower() not in {"http", "https"}: raise TargetPolicyError("only HTTP and HTTPS targets are allowed") if parsed.username is not None or parsed.password is not None: raise TargetPolicyError("target userinfo is not allowed") if not parsed.hostname: raise TargetPolicyError("target hostname is required") hostname = parsed.hostname.rstrip(".").lower() port = parsed.port or (443 if parsed.scheme.lower() == "https" else 80) except ValueError as exc: raise TargetPolicyError("target URL is invalid") from exc addresses: set[str] = set() try: literal = ipaddress.ip_address(hostname.split("%", 1)[0]) addresses.add(str(literal)) except ValueError: try: answers = await resolver(hostname, port) for answer in answers: sockaddr = answer[4] if isinstance(sockaddr, tuple) and sockaddr: addresses.add(str(ipaddress.ip_address(str(sockaddr[0]).split("%", 1)[0]))) except (OSError, ValueError) as exc: raise TargetPolicyError("target DNS resolution failed") from exc if not addresses: raise TargetPolicyError("target DNS returned no addresses") if len(addresses) > 32: raise TargetPolicyError("target DNS returned too many addresses") if not all(_public_address(address) for address in addresses): raise TargetPolicyError("target resolves to a non-public address") return ValidatedTarget(url=url, hostname=hostname, port=port, addresses=tuple(sorted(addresses)))