382 lines
14 KiB
Python
382 lines
14 KiB
Python
"""Fetching pages that strangers ask the checker to look at.
|
|
|
|
Every URL reaching this module is author-supplied, so the service must not be
|
|
usable as a probe against the machine it runs on or its neighbours. Addresses
|
|
are vetted before the connection and the connection is pinned to the vetted
|
|
address, so a second DNS answer cannot redirect it (DNS rebinding).
|
|
"""
|
|
|
|
from dataclasses import dataclass, replace
|
|
import ipaddress
|
|
import os
|
|
import re
|
|
import socket
|
|
import time
|
|
from urllib.parse import urlsplit, urlunsplit
|
|
|
|
import httpx
|
|
from publicsuffixlist import PublicSuffixList
|
|
|
|
# Read caps, in bytes. The page cap is above the 256 KB that SPEC.md 4.3 makes a
|
|
# MUST so an oversized page is reported as too large rather than unreachable.
|
|
PAGE_CAP = 288 * 1024
|
|
IMAGE_CAP = 64 * 1024
|
|
CSS_CAP = 256 * 1024
|
|
|
|
MAX_REDIRECTS = 3
|
|
MAX_URL_LENGTH = 2048
|
|
USER_AGENT = "mews.page checker (+https://mews.page/about)"
|
|
|
|
_PSL = PublicSuffixList()
|
|
|
|
# Ranges ipaddress does not classify but which must never be reached: carrier
|
|
# NAT (which is also Tailscale's range), IETF protocol assignment, benchmarking,
|
|
# the documentation ranges, and the IPv6 transition mechanisms that tunnel an
|
|
# IPv4 address inside an address that looks global.
|
|
_BLOCKED_NETS = [
|
|
ipaddress.ip_network(net)
|
|
for net in (
|
|
"100.64.0.0/10",
|
|
"192.0.0.0/24",
|
|
"198.18.0.0/15",
|
|
"192.0.2.0/24",
|
|
"198.51.100.0/24",
|
|
"203.0.113.0/24",
|
|
"240.0.0.0/4",
|
|
"255.255.255.255/32",
|
|
"2001:db8::/32",
|
|
"2002::/16",
|
|
"2001::/32",
|
|
"64:ff9b::/96",
|
|
)
|
|
]
|
|
|
|
# Hosts whose own addresses must be unreachable, as CIDRs. The deployment sets
|
|
# this to the machine's public address so a submission cannot be aimed at a
|
|
# service sharing the box.
|
|
_EXTRA_NETS = [
|
|
ipaddress.ip_network(net.strip())
|
|
for net in os.environ.get("MEWS_BLOCK_NETS", "").split(",")
|
|
if net.strip()
|
|
]
|
|
|
|
_LOCAL_SUFFIXES = (".local", ".internal", ".home.arpa", ".localhost")
|
|
|
|
# What a host name may be made of. Checking this first means a typo gets a
|
|
# clearer answer than a complaint about public suffixes.
|
|
_HOSTNAME = re.compile(
|
|
r"^[a-z0-9]([a-z0-9-]*[a-z0-9])?(\.[a-z0-9]([a-z0-9-]*[a-z0-9])?)*$"
|
|
)
|
|
|
|
|
|
class UrlError(Exception):
|
|
"""A URL the checker will not fetch. The message is shown to the author."""
|
|
|
|
|
|
class FetchError(Exception):
|
|
"""A fetch that produced no page. The message is shown to the author."""
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class Fetched:
|
|
"""One response the checker read, with the body it kept."""
|
|
|
|
url: str
|
|
status: int
|
|
headers: httpx.Headers
|
|
body: bytes
|
|
truncated: bool
|
|
scheme_downgraded: bool = False
|
|
|
|
|
|
def registered_domain(host: str) -> str | None:
|
|
"""Return the registrable domain of host, or None when it has no public suffix.
|
|
|
|
Comparing registrable domains rather than hostnames is what makes the
|
|
same-site rules in SPEC.md 4.2 and 7.4 mean anything: a free subdomain host
|
|
would otherwise let any two unrelated sites count as one.
|
|
"""
|
|
if not host:
|
|
return None
|
|
try:
|
|
name = host.strip().rstrip(".").lower().encode("idna").decode("ascii")
|
|
except (UnicodeError, UnicodeDecodeError):
|
|
return None
|
|
return _PSL.privatesuffix(name)
|
|
|
|
|
|
def same_site(a: str, b: str) -> bool:
|
|
"""Whether two hosts share a registrable domain."""
|
|
first = registered_domain(a)
|
|
return first is not None and first == registered_domain(b)
|
|
|
|
|
|
def normalise_url(raw: str, *, allow_loopback: bool = False) -> str:
|
|
"""Return a URL the checker is willing to fetch, or raise UrlError.
|
|
|
|
This runs before any name lookup, so a rejected URL costs nothing. With
|
|
allow_loopback the rules relax enough to reach a test server on this
|
|
machine: an address literal, localhost, and any port.
|
|
"""
|
|
text = (raw or "").strip()
|
|
if not text:
|
|
raise UrlError("Enter the address of a page to check.")
|
|
if len(text) > MAX_URL_LENGTH:
|
|
raise UrlError("That address is too long.")
|
|
if "://" not in text:
|
|
text = "https://" + text
|
|
|
|
parts = urlsplit(text)
|
|
if parts.scheme not in ("http", "https"):
|
|
raise UrlError("Enter an address that starts with http:// or https://.")
|
|
if "@" in parts.netloc:
|
|
raise UrlError("Enter an address without a user name in it.")
|
|
|
|
try:
|
|
host, port = parts.hostname, parts.port
|
|
except ValueError as error:
|
|
raise UrlError("That address has a port the checker can't read.") from error
|
|
if not host:
|
|
raise UrlError(
|
|
"That doesn't look like a web address. Enter the full address "
|
|
"of a page, like https://example.com/"
|
|
)
|
|
if port is not None and port not in (80, 443) and not allow_loopback:
|
|
raise UrlError("The checker reads pages on the usual web ports only.")
|
|
|
|
host = host.lower().rstrip(".")
|
|
if _is_ip_literal(host):
|
|
if not (allow_loopback and _is_loopback_literal(host)):
|
|
raise UrlError("Enter a domain name rather than an IP address.")
|
|
elif host == "localhost" or host.endswith(_LOCAL_SUFFIXES):
|
|
if not allow_loopback:
|
|
raise UrlError("That address is only reachable on a local network.")
|
|
elif not _HOSTNAME.match(_ascii(host)):
|
|
raise UrlError(
|
|
"That doesn't look like a web address. Enter the full address "
|
|
"of a page, like https://example.com/"
|
|
)
|
|
elif registered_domain(host) is None:
|
|
raise UrlError("That domain name isn't one the checker can reach.")
|
|
|
|
netloc = host if port is None else f"{host}:{port}"
|
|
return urlunsplit((parts.scheme, netloc, parts.path or "/", parts.query, ""))
|
|
|
|
|
|
def _ascii(host: str) -> str:
|
|
"""Return the punycode form of a host name, or the name unchanged."""
|
|
try:
|
|
return host.encode("idna").decode("ascii")
|
|
except (UnicodeError, UnicodeDecodeError):
|
|
return host
|
|
|
|
|
|
def _is_ip_literal(host: str) -> bool:
|
|
"""Whether host is an address literal in any of the forms a parser accepts."""
|
|
candidate = host.strip("[]")
|
|
try:
|
|
ipaddress.ip_address(candidate)
|
|
except ValueError:
|
|
pass
|
|
else:
|
|
return True
|
|
# Decimal, octal and hex integer forms of an IPv4 address, which urlsplit
|
|
# leaves alone but a resolver would accept.
|
|
if host.isdigit():
|
|
return True
|
|
bare = host.replace(".", "")
|
|
return bare.startswith(("0x", "0X")) or (
|
|
host.startswith("0") and len(host) > 1 and bare.isdigit()
|
|
)
|
|
|
|
|
|
def _is_loopback_literal(host: str) -> bool:
|
|
"""Whether host is a literal address on this machine."""
|
|
try:
|
|
return ipaddress.ip_address(host.strip("[]")).is_loopback
|
|
except ValueError:
|
|
return False
|
|
|
|
|
|
def _embedded(address: ipaddress.IPv6Address) -> ipaddress.IPv4Address | None:
|
|
"""Return the IPv4 address a transition mechanism hides inside an IPv6 one."""
|
|
for attribute in ("ipv4_mapped", "sixtofour"):
|
|
value = getattr(address, attribute, None)
|
|
if value is not None:
|
|
return value
|
|
teredo = getattr(address, "teredo", None)
|
|
if teredo:
|
|
return teredo[1]
|
|
if address in ipaddress.ip_network("64:ff9b::/96"):
|
|
return ipaddress.IPv4Address(int(address) & 0xFFFFFFFF)
|
|
return None
|
|
|
|
|
|
def vet_address(address: str, *, allow_loopback: bool = False) -> None:
|
|
"""Raise UrlError unless address is a public one the checker may connect to."""
|
|
ip = ipaddress.ip_address(address)
|
|
if allow_loopback and ip.is_loopback:
|
|
return
|
|
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
|
|
or any(ip in net for net in _BLOCKED_NETS)
|
|
or any(ip in net for net in _EXTRA_NETS)
|
|
):
|
|
raise UrlError("That address isn't on the public internet.")
|
|
if isinstance(ip, ipaddress.IPv6Address):
|
|
inner = _embedded(ip)
|
|
if inner is not None:
|
|
vet_address(str(inner), allow_loopback=allow_loopback)
|
|
|
|
|
|
def resolve(url: str, *, allow_loopback: bool = False) -> str:
|
|
"""Return one vetted address for the URL's host, rejecting the host if any fails.
|
|
|
|
Every answer has to pass: a host that resolves to one public and one private
|
|
address would otherwise be a coin toss.
|
|
"""
|
|
parts = urlsplit(url)
|
|
host = parts.hostname or ""
|
|
port = parts.port or (443 if parts.scheme == "https" else 80)
|
|
try:
|
|
infos = socket.getaddrinfo(host, port, type=socket.SOCK_STREAM)
|
|
except socket.gaierror as error:
|
|
raise FetchError(
|
|
"We couldn't find that domain name. Check the address and try again."
|
|
) from error
|
|
addresses = [info[4][0] for info in infos]
|
|
if not addresses:
|
|
raise FetchError("We couldn't find that domain name.")
|
|
for address in addresses:
|
|
vet_address(address, allow_loopback=allow_loopback)
|
|
return addresses[0]
|
|
|
|
|
|
class Fetcher:
|
|
"""Reads pages and their sub-resources, under a time and byte budget.
|
|
|
|
One Fetcher serves one check, so the whole-check budget is shared across the
|
|
page, its stylesheet and its images.
|
|
"""
|
|
|
|
def __init__(self, *, allow_loopback: bool = False, budget: float = 45.0):
|
|
self.allow_loopback = allow_loopback
|
|
self._deadline = time.monotonic() + budget
|
|
self._client = httpx.Client(
|
|
follow_redirects=False,
|
|
trust_env=False,
|
|
http2=False,
|
|
verify=True,
|
|
timeout=httpx.Timeout(connect=5.0, read=10.0, write=5.0, pool=5.0),
|
|
limits=httpx.Limits(max_connections=4, max_keepalive_connections=2),
|
|
headers={
|
|
"User-Agent": USER_AGENT,
|
|
"Accept": "text/html,application/xhtml+xml",
|
|
"Accept-Encoding": "gzip",
|
|
},
|
|
)
|
|
|
|
def __enter__(self) -> "Fetcher":
|
|
return self
|
|
|
|
def __exit__(self, *exc: object) -> None:
|
|
self.close()
|
|
|
|
def close(self) -> None:
|
|
"""Close the underlying connections."""
|
|
self._client.close()
|
|
|
|
def get(
|
|
self,
|
|
url: str,
|
|
*,
|
|
cap: int = PAGE_CAP,
|
|
headers: dict[str, str] | None = None,
|
|
) -> Fetched:
|
|
"""Fetch one URL, following redirects by hand and vetting each hop."""
|
|
current = normalise_url(url, allow_loopback=self.allow_loopback)
|
|
downgraded = False
|
|
for _ in range(MAX_REDIRECTS + 1):
|
|
response = self._request(current, cap=cap, headers=headers)
|
|
if response.status not in (301, 302, 303, 307, 308):
|
|
if downgraded:
|
|
return replace(response, scheme_downgraded=True)
|
|
return response
|
|
location = response.headers.get("location", "")
|
|
if not location:
|
|
raise FetchError("That page redirects without saying where to.")
|
|
target = normalise_url(
|
|
str(httpx.URL(current).join(location)),
|
|
allow_loopback=self.allow_loopback,
|
|
)
|
|
was_secure = urlsplit(current).scheme == "https"
|
|
if was_secure and urlsplit(target).scheme == "http":
|
|
downgraded = True
|
|
current = target
|
|
raise FetchError("That page redirects too many times.")
|
|
|
|
def _request(
|
|
self,
|
|
url: str,
|
|
*,
|
|
cap: int,
|
|
headers: dict[str, str] | None,
|
|
) -> Fetched:
|
|
"""Make one pinned request, streaming the body up to cap bytes."""
|
|
self._check_deadline()
|
|
parts = urlsplit(url)
|
|
host = parts.hostname or ""
|
|
port = parts.port
|
|
address = resolve(url, allow_loopback=self.allow_loopback)
|
|
literal = f"[{address}]" if ":" in address else address
|
|
pinned = urlunsplit(
|
|
(
|
|
parts.scheme,
|
|
literal if port is None else f"{literal}:{port}",
|
|
parts.path,
|
|
parts.query,
|
|
"",
|
|
)
|
|
)
|
|
request_headers = {"Host": parts.netloc, **(headers or {})}
|
|
extensions = {"sni_hostname": host} if parts.scheme == "https" else {}
|
|
try:
|
|
with self._client.stream(
|
|
"GET", pinned, headers=request_headers, extensions=extensions
|
|
) as response:
|
|
declared = response.headers.get("content-length")
|
|
if declared and declared.isdigit() and int(declared) > cap:
|
|
raise FetchError("That page is too large for the checker to read.")
|
|
body, truncated = self._read(response, cap)
|
|
except httpx.TimeoutException as error:
|
|
raise FetchError("That page took too long to answer.") from error
|
|
except httpx.HTTPError as error:
|
|
raise FetchError("We couldn't connect to that site.") from error
|
|
return Fetched(url, response.status_code, response.headers, body, truncated)
|
|
|
|
def _read(self, response: httpx.Response, cap: int) -> tuple[bytes, bool]:
|
|
"""Read a streamed body, stopping at cap decoded bytes."""
|
|
chunks: list[bytes] = []
|
|
total = 0
|
|
for chunk in response.iter_bytes():
|
|
self._check_deadline()
|
|
chunks.append(chunk)
|
|
total += len(chunk)
|
|
if total > cap:
|
|
return b"".join(chunks)[:cap], True
|
|
body = b"".join(chunks)
|
|
encoded = response.num_bytes_downloaded or len(body)
|
|
# A body that expanded enormously from a small download is a compression
|
|
# bomb, not a page.
|
|
if encoded and len(body) > 100 * encoded and len(body) > 64 * 1024:
|
|
raise FetchError("That page is compressed in a way the checker won't read.")
|
|
return body, False
|
|
|
|
def _check_deadline(self) -> None:
|
|
if time.monotonic() > self._deadline:
|
|
raise FetchError("Checking that page took too long.")
|