"""SSRF protection. The service fetches arbitrary URLs on request, which is exactly the shape of an SSRF vulnerability. Every URL is validated after DNS resolution, not on the string alone, and the check is repeated on every redirect hop. """ from __future__ import annotations import ipaddress import logging import socket from urllib.parse import urlparse from ..config import Settings from ..errors import BlockedTargetError logger = logging.getLogger(__name__) ALLOWED_SCHEMES = ("http", "https") # Ranges that must never be reachable from a user supplied URL. PRIVATE_NETWORKS = [ ipaddress.ip_network("127.0.0.0/8"), ipaddress.ip_network("10.0.0.0/8"), ipaddress.ip_network("172.16.0.0/12"), ipaddress.ip_network("192.168.0.0/16"), ipaddress.ip_network("169.254.0.0/16"), ipaddress.ip_network("0.0.0.0/8"), ipaddress.ip_network("100.64.0.0/10"), ipaddress.ip_network("192.0.0.0/24"), ipaddress.ip_network("198.18.0.0/15"), ipaddress.ip_network("224.0.0.0/4"), ipaddress.ip_network("240.0.0.0/4"), ipaddress.ip_network("::1/128"), ipaddress.ip_network("fc00::/7"), ipaddress.ip_network("fe80::/10"), ipaddress.ip_network("::/128"), ] class UrlGuard: """Validates URLs against the configured policy.""" def __init__(self, settings: Settings) -> None: self._block_private = settings.ssrf_block_private self._allowed_hosts = {host.lower() for host in settings.ssrf_allowed_hosts} self._extra_blocked: list[ipaddress._BaseNetwork] = [] for cidr in settings.ssrf_extra_blocked_cidrs: try: self._extra_blocked.append(ipaddress.ip_network(cidr, strict=False)) except ValueError: logger.warning("Ignoring invalid CIDR in SSRF_EXTRA_BLOCKED_CIDRS", extra={"cidr": cidr}) def check(self, url: str) -> str: """Raise BlockedTargetError when the URL must not be fetched. Returns the hostname so callers can reuse it without parsing again. """ parsed = urlparse(url) scheme = (parsed.scheme or "").lower() if scheme not in ALLOWED_SCHEMES: raise BlockedTargetError( "Povolena jsou pouze schemata http a https.", {"url": url, "scheme": scheme or None}, ) host = parsed.hostname if not host: raise BlockedTargetError("Adresa neobsahuje hostname.", {"url": url}) if host.lower() in self._allowed_hosts: logger.info("Host explicitly allowlisted", extra={"host": host}) return host for address in self._resolve(host, url): self._check_address(address, host, url) return host def _resolve(self, host: str, url: str) -> list[ipaddress.IPv4Address | ipaddress.IPv6Address]: # A literal IP address needs no lookup. try: return [ipaddress.ip_address(host)] except ValueError: pass try: infos = socket.getaddrinfo(host, None, proto=socket.IPPROTO_TCP) except socket.gaierror as exc: raise BlockedTargetError( f"Hostname {host} se nepodarilo prelozit na IP adresu.", {"url": url, "reason": str(exc)}, ) from exc addresses = [] for info in infos: try: addresses.append(ipaddress.ip_address(info[4][0])) except ValueError: continue if not addresses: raise BlockedTargetError(f"Hostname {host} nema zadnou pouzitelnou IP adresu.", {"url": url}) return addresses def _check_address(self, address, host: str, url: str) -> None: if self._block_private: for network in PRIVATE_NETWORKS: if address.version == network.version and address in network: logger.warning( "Blocked request to private address", extra={"host": host, "address": str(address), "network": str(network)}, ) raise BlockedTargetError( "Cilova adresa smeruje do privatniho nebo vyhrazeneho rozsahu a je zablokovana.", {"url": url, "host": host, "address": str(address)}, ) for network in self._extra_blocked: if address.version == network.version and address in network: logger.warning( "Blocked request by configured CIDR", extra={"host": host, "address": str(address), "network": str(network)}, ) raise BlockedTargetError( "Cilova adresa je v konfigurovanem seznamu blokovanych rozsahu.", {"url": url, "host": host, "address": str(address)}, )