__init__
This commit is contained in:
@@ -0,0 +1,109 @@
|
||||
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
|
||||
Reference in New Issue
Block a user