108 lines
3.6 KiB
Python
108 lines
3.6 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import hashlib
|
|
import ipaddress
|
|
import re
|
|
from dataclasses import dataclass
|
|
from typing import Protocol
|
|
from urllib.parse import quote, unquote, urlsplit, urlunsplit
|
|
|
|
import idna
|
|
|
|
URL_CANDIDATE = re.compile(r"(?i)\b(?:[a-z][a-z0-9+.-]*://)[^\s<>{}\[\]\"']+")
|
|
METADATA = {
|
|
ipaddress.ip_address("169.254.169.254"),
|
|
ipaddress.ip_address("100.100.100.200"),
|
|
ipaddress.ip_address("fd00:ec2::254"),
|
|
}
|
|
|
|
|
|
class DnsError(RuntimeError):
|
|
pass
|
|
|
|
|
|
class DnsNxDomain(DnsError):
|
|
pass
|
|
|
|
|
|
class Resolver(Protocol):
|
|
async def resolve(
|
|
self, hostname: str
|
|
) -> tuple[ipaddress.IPv4Address | ipaddress.IPv6Address, ...]: ...
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class CanonicalUrl:
|
|
value: str
|
|
digest: bytes
|
|
hostname: str
|
|
literal_ip: ipaddress.IPv4Address | ipaddress.IPv6Address | None
|
|
|
|
|
|
def extract_urls(text: str, *, maximum: int = 5, max_length: int = 2048) -> tuple[str, ...]:
|
|
values = tuple(match.group(0).rstrip(".,;:!?)]") for match in URL_CANDIDATE.finditer(text))
|
|
if len(values) > maximum or any(len(value) > max_length for value in values):
|
|
raise ValueError("URL limits exceeded")
|
|
return values
|
|
|
|
|
|
def canonicalize(raw: str) -> CanonicalUrl:
|
|
parsed = urlsplit(raw)
|
|
if parsed.scheme.lower() not in {"http", "https"}:
|
|
raise PermissionError("url.forbidden_scheme")
|
|
if not parsed.hostname or parsed.username is not None or parsed.password is not None:
|
|
raise PermissionError("url.credentials_present" if parsed.username else "url.malformed")
|
|
try:
|
|
host = idna.encode(parsed.hostname, uts46=True, transitional=False).decode("ascii").lower()
|
|
except idna.IDNAError as exc:
|
|
raise PermissionError("url.confusable_host") from exc
|
|
try:
|
|
literal = ipaddress.ip_address(host)
|
|
if isinstance(literal, ipaddress.IPv6Address) and literal.ipv4_mapped:
|
|
literal = literal.ipv4_mapped
|
|
except ValueError:
|
|
literal = None
|
|
try:
|
|
parsed_port = parsed.port
|
|
except ValueError as exc:
|
|
raise PermissionError("url.malformed") from exc
|
|
port = (
|
|
f":{parsed_port}"
|
|
if parsed_port and parsed_port != (443 if parsed.scheme == "https" else 80)
|
|
else ""
|
|
)
|
|
path = quote(unquote(parsed.path or "/"), safe="/:@-._~!$&'()*+,;=")
|
|
query = quote(unquote(parsed.query), safe="=&/:?@-._~!$'()*+,;")
|
|
canonical = urlunsplit((parsed.scheme.lower(), host + port, path, query, ""))
|
|
return CanonicalUrl(canonical, hashlib.sha256(canonical.encode()).digest(), host, literal)
|
|
|
|
|
|
def classify_ip(address: ipaddress.IPv4Address | ipaddress.IPv6Address) -> str | None:
|
|
if isinstance(address, ipaddress.IPv6Address) and address.ipv4_mapped:
|
|
address = address.ipv4_mapped
|
|
if address in METADATA or address.is_private or address.is_loopback or address.is_link_local:
|
|
return "url.private_destination"
|
|
if address.is_multicast or address.is_unspecified or address.is_reserved:
|
|
return "url.reserved_destination"
|
|
return None
|
|
|
|
|
|
async def check_url(
|
|
raw: str, resolver: Resolver, timeout_sec: float = 1.0
|
|
) -> tuple[CanonicalUrl, str | None]:
|
|
canonical = canonicalize(raw)
|
|
if canonical.literal_ip:
|
|
return canonical, classify_ip(canonical.literal_ip)
|
|
try:
|
|
addresses = await asyncio.wait_for(resolver.resolve(canonical.hostname), timeout_sec)
|
|
except DnsNxDomain:
|
|
return canonical, "url.nxdomain"
|
|
except (TimeoutError, DnsError) as exc:
|
|
raise DnsError("DNS dependency unavailable") from exc
|
|
for address in addresses:
|
|
denied = classify_ip(address)
|
|
if denied:
|
|
return canonical, denied
|
|
return canonical, None
|