from __future__ import annotations import ipaddress import re import socket from dataclasses import dataclass from urllib.parse import quote, urlsplit, urlunsplit class UrlPolicyError(ValueError): """Raised when a URL is structurally unsafe or points at a blocked network.""" @dataclass(frozen=True, slots=True) class NetworkRules: allow_private_networks: bool = False allow_nonstandard_ports: bool = False def normalize_url(value: str, rules: NetworkRules) -> str: raw = value.strip() if not raw: raise UrlPolicyError("Enter a website address.") if "://" not in raw: raw = f"https://{raw}" try: parsed = urlsplit(raw) port = parsed.port except ValueError as error: raise UrlPolicyError("The website address has an invalid port.") from error if parsed.scheme.lower() not in {"http", "https"}: raise UrlPolicyError("Only HTTP and HTTPS websites can be captured.") if not parsed.hostname: raise UrlPolicyError("The website address needs a hostname.") if parsed.username or parsed.password: raise UrlPolicyError("Website addresses with embedded credentials are not allowed.") host = parsed.hostname.rstrip(".").lower() try: host = host.encode("idna").decode("ascii") except UnicodeError as error: raise UrlPolicyError("The website hostname is not valid.") from error default_port = 443 if parsed.scheme.lower() == "https" else 80 if port and port != default_port and not rules.allow_nonstandard_ports: raise UrlPolicyError("Only ports 80 and 443 are allowed by this SiteHarbor instance.") host_for_netloc = f"[{host}]" if ":" in host else host netloc = host_for_netloc if not port or port == default_port else f"{host_for_netloc}:{port}" path_value = parsed.path or "/" if re.search(r"%(?![0-9A-Fa-f]{2})", path_value): raise UrlPolicyError("The website address has an invalid percent escape.") # Keep encoded delimiters such as %2F intact: decoding them changes resource identity. path = quote(path_value, safe="/%:@!$&'()*+,;=-._~") return urlunsplit((parsed.scheme.lower(), netloc, path, parsed.query, "")) def canonical_url(value: str, rules: NetworkRules) -> str: """Return a stable URL key with its fragment removed.""" return normalize_url(value, rules) def hostname(value: str) -> str: host = urlsplit(value).hostname if not host: raise UrlPolicyError("The URL has no hostname.") return host.rstrip(".").lower() def is_public_address(address: ipaddress.IPv4Address | ipaddress.IPv6Address) -> bool: return bool( address.is_global and not address.is_loopback and not address.is_link_local and not address.is_multicast and not address.is_unspecified and not address.is_reserved and not address.is_private ) def resolve_and_validate(value: str, rules: NetworkRules) -> tuple[str, ...]: """Resolve a host and reject non-public destinations unless explicitly in dev mode.""" host = hostname(value) try: results = socket.getaddrinfo(host, None, type=socket.SOCK_STREAM) except socket.gaierror as error: raise UrlPolicyError(f"Could not resolve {host}.") from error addresses: set[str] = set() for _, _, _, _, sockaddr in results: address = ipaddress.ip_address(sockaddr[0]) addresses.add(str(address)) if not rules.allow_private_networks and not is_public_address(address): raise UrlPolicyError("Private, local, and reserved network targets are blocked.") if not addresses: raise UrlPolicyError(f"Could not resolve {host}.") return tuple(sorted(addresses)) def same_site_host(left: str, right: str) -> bool: """Allow an exact host or the common example.com/www.example.com redirect pair.""" left = left.lower().removeprefix("www.") right = right.lower().removeprefix("www.") return left == right