Team Ai
Apppublic

booleanbeyond/jobfetch

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes
security.py177 linesDownload Raw Back to app
1"""Egress guard (SSRF protection).2 3This service takes a URL from an untrusted user and fetches it. That is a4textbook SSRF sink, so every outbound request — including every redirect hop and5every URL discovered inside a page — goes through `resolve_and_validate()`.6 7Defence layers:8  1. Scheme allow-list (http/https only). Blocks file://, gopher://, data:, etc.9  2. Port allow-list.10  3. Host deny-list (cloud metadata endpoints).11  4. DNS resolution, then rejection if *any* resolved address is private,12     loopback, link-local, multicast, reserved or unspecified. Resolving first13     also defeats obfuscated literals (0x7f.1, 2130706433, decimal, IPv6-mapped).14  5. DNS pinning: the connection is made to a validated IP with the original15     hostname carried in the Host header and TLS SNI, which closes the16     DNS-rebinding window between validation and connect.17  6. Redirects are followed manually so each hop is re-validated.18"""19 20from __future__ import annotations21 22import asyncio23import ipaddress24import socket25from dataclasses import dataclass26from urllib.parse import urlsplit, urlunsplit27 28from .config import settings29 30ALLOWED_SCHEMES = {"http", "https"}31DEFAULT_PORTS = {"http": 80, "https": 443}32 33 34class UnsafeURLError(ValueError):35    """Raised when a URL must not be fetched."""36 37 38@dataclass(frozen=True)39class ValidatedTarget:40    url: str41    scheme: str42    host: str  # hostname as written (used for Host header + SNI)43    port: int44    ips: tuple[str, ...]  # validated, public addresses45 46    @property47    def pinned_url(self) -> str:48        """URL with the host replaced by a validated literal IP."""49        parts = urlsplit(self.url)50        ip = self.ips[0]51        literal = f"[{ip}]" if ":" in ip else ip52        netloc = f"{literal}:{self.port}"53        return urlunsplit((parts.scheme, netloc, parts.path or "/", parts.query, ""))54 55 56def _is_public(ip: ipaddress._BaseAddress) -> bool:57    if ip.is_private or ip.is_loopback or ip.is_link_local:58        return False59    if ip.is_multicast or ip.is_reserved or ip.is_unspecified:60        return False61    # IPv4-mapped / 6to4 / Teredo tunnels can smuggle private v4 space.62    if isinstance(ip, ipaddress.IPv6Address):63        if ip.ipv4_mapped is not None and not _is_public(ip.ipv4_mapped):64            return False65        if ip.sixtofour is not None and not _is_public(ip.sixtofour):66            return False67        if ip.teredo is not None and not _is_public(ip.teredo[1]):68            return False69    # Carrier-grade NAT.70    if isinstance(ip, ipaddress.IPv4Address) and ip in ipaddress.ip_network("100.64.0.0/10"):71        return False72    return True73 74 75def normalise_url(raw: str) -> str:76    """Trim, add a scheme if the user pasted a bare domain, drop the fragment."""77    if raw is None:78        raise UnsafeURLError("URL is required")79    url = raw.strip()80    if not url:81        raise UnsafeURLError("URL is required")82    if len(url) > 2048:83        raise UnsafeURLError("URL is too long")84    if "\n" in url or "\r" in url or "\t" in url:85        raise UnsafeURLError("URL contains control characters")86    if "://" not in url:87        # Reject things like "javascript:alert(1)" that have a scheme but no //.88        head = url.split(":", 1)[0].lower()89        if ":" in url and head.isalpha() and head not in ALLOWED_SCHEMES:90            raise UnsafeURLError(f"Unsupported URL scheme: {head}")91        url = "https://" + url92    parts = urlsplit(url)93    if parts.scheme.lower() not in ALLOWED_SCHEMES:94        raise UnsafeURLError(f"Unsupported URL scheme: {parts.scheme}")95    return urlunsplit(96        (parts.scheme.lower(), parts.netloc, parts.path or "/", parts.query, "")97    )98 99 100async def _resolve(host: str, port: int) -> list[str]:101    loop = asyncio.get_running_loop()102    try:103        infos = await loop.getaddrinfo(host, port, type=socket.SOCK_STREAM)104    except socket.gaierror as exc:105        raise UnsafeURLError(f"Could not resolve host: {host}") from exc106    seen: list[str] = []107    for info in infos:108        addr = info[4][0]109        if addr not in seen:110            seen.append(addr)111    if not seen:112        raise UnsafeURLError(f"Could not resolve host: {host}")113    return seen114 115 116async def resolve_and_validate(raw_url: str) -> ValidatedTarget:117    """Validate a URL and return the concrete addresses it is safe to hit."""118    url = normalise_url(raw_url)119    parts = urlsplit(url)120 121    host = (parts.hostname or "").strip().rstrip(".")122    if not host:123        raise UnsafeURLError("URL has no host")124 125    lowered = host.lower()126    for blocked in settings.blocked_hosts:127        b = blocked.lower()128        if lowered == b or lowered.endswith("." + b):129            raise UnsafeURLError(f"Host is blocked: {host}")130 131    if parts.username or parts.password:132        # user:pass@host is a classic filter-bypass trick and never needed here.133        raise UnsafeURLError("Credentials in URL are not allowed")134 135    try:136        port = parts.port or DEFAULT_PORTS[parts.scheme]137    except ValueError as exc:138        raise UnsafeURLError("Invalid port") from exc139 140    if not settings.allow_private_hosts and port not in settings.allowed_ports:141        raise UnsafeURLError(f"Port not allowed: {port}")142 143    if settings.allow_private_hosts:144        # Explicit dev/test escape hatch (fixture server on 127.0.0.1).145        ips = await _resolve(host, port)146        return ValidatedTarget(url, parts.scheme, host, port, tuple(ips))147 148    ips = await _resolve(host, port)149    validated: list[str] = []150    for addr in ips:151        try:152            ip = ipaddress.ip_address(addr)153        except ValueError:154            continue155        if not _is_public(ip):156            raise UnsafeURLError(157                f"Refusing to fetch {host}: resolves to non-public address {addr}"158            )159        validated.append(addr)160 161    if not validated:162        raise UnsafeURLError(f"No usable address for host: {host}")163 164    return ValidatedTarget(url, parts.scheme, host, port, tuple(validated))165 166 167def same_registrable_site(a: str, b: str) -> bool:168    """Cheap same-site check (no PSL dependency): compare last two labels."""169    try:170        ha = (urlsplit(a).hostname or "").lower()171        hb = (urlsplit(b).hostname or "").lower()172    except ValueError:173        return False174    if not ha or not hb:175        return False176    return ha.split(".")[-2:] == hb.split(".")[-2:]177