"""Rejects user-supplied URLs that point inside the network.

Any endpoint that fetches a URL a user chose is an SSRF sink: pointed at cloud
instance metadata it leaks credentials, pointed at an RFC1918 address it reaches
services that were never meant to be public.
"""
import ipaddress
import logging
import socket
from typing import Any
from urllib.parse import urlparse

logger = logging.getLogger(__name__)

# https only: these requests carry the merchant's API key in an Authorization
# header, so http is not merely a weaker transport -- it is a credential
# disclosure to anyone on the network path.
ALLOWED_SCHEMES = ("https",)


class UnsafeUrlError(Exception):
    """The URL points somewhere this service must not fetch from."""


def _resolve(host: str) -> list:
    return [info[4][0] for info in socket.getaddrinfo(host, None)]


def assert_safe_url(url: Any) -> str:
    if not url or not isinstance(url, str):
        raise UnsafeUrlError("url must be a non-empty string")

    # urlparse itself can raise (e.g. a malformed bracketed IPv6 authority).
    # A URL we cannot even parse is definitionally not safe to fetch, so
    # treat a parse failure as a rejection rather than letting it escape.
    try:
        parsed = urlparse(url)
        hostname = parsed.hostname
    except ValueError as ex:
        raise UnsafeUrlError("url could not be parsed") from ex

    if parsed.scheme not in ALLOWED_SCHEMES:
        raise UnsafeUrlError(f"scheme {parsed.scheme!r} is not allowed")
    if not hostname:
        raise UnsafeUrlError("url has no host")

    try:
        addresses = _resolve(hostname)
    except (OSError, UnicodeError) as ex:
        # getaddrinfo raises UnicodeError (not an OSError) on an over-long
        # IDNA label, so both must be caught here or a malformed host 500s
        # instead of being rejected like every other bad input.
        raise UnsafeUrlError(f"host {hostname!r} does not resolve") from ex

    for raw in addresses:
        # Check the RESOLVED address, never the hostname string: a public-looking
        # name can resolve to 127.0.0.1 and would otherwise pass every check.
        ip = ipaddress.ip_address(raw)
        if (ip.is_private or ip.is_loopback or ip.is_link_local
                or ip.is_reserved or ip.is_multicast or ip.is_unspecified):
            # Name the rule, not the address: echoing the resolved IP back to
            # the caller would confirm what exists on the internal network.
            raise UnsafeUrlError(
                f"host {hostname!r} resolves to a non-public address")

    return url
