diff --git a/app/security.py b/app/security.py new file mode 100644 index 0000000..38259d7 --- /dev/null +++ b/app/security.py @@ -0,0 +1,54 @@ +import asyncio +import ipaddress +import socket +from collections.abc import Awaitable, Callable +from urllib.parse import SplitResult, urlsplit, urlunsplit + +Resolver = Callable[[str, int], Awaitable[list[str]]] + + +class UnsafeTargetError(ValueError): + pass + + +async def resolve_addresses(host: str, port: int) -> list[str]: + def lookup() -> list[str]: + records = socket.getaddrinfo(host, port, type=socket.SOCK_STREAM) + return list({record[4][0] for record in records}) + + return await asyncio.to_thread(lookup) + + +def redacted_url(url: str) -> str: + parsed = urlsplit(url) + host = parsed.hostname or "invalid-host" + if ":" in host: + host = f"[{host}]" + port = f":{parsed.port}" if parsed.port else "" + return urlunsplit((parsed.scheme, f"{host}{port}", parsed.path or "/", "", "")) + + +async def validate_target(url: str, resolver: Resolver) -> SplitResult: + parsed = urlsplit(url) + if parsed.scheme not in {"http", "https"} or not parsed.hostname: + raise UnsafeTargetError("only absolute HTTP(S) URLs are allowed") + if parsed.username is not None or parsed.password is not None: + raise UnsafeTargetError("URL credentials are not allowed") + try: + port = parsed.port or (443 if parsed.scheme == "https" else 80) + except ValueError as exc: + raise UnsafeTargetError("invalid port") from exc + try: + addresses = await resolver(parsed.hostname, port) + except (OSError, socket.gaierror) as exc: + raise UnsafeTargetError("DNS resolution failed") from exc + if not addresses: + raise UnsafeTargetError("DNS resolution returned no addresses") + for address in addresses: + try: + ip = ipaddress.ip_address(address) + except ValueError as exc: + raise UnsafeTargetError("DNS returned an invalid address") from exc + if not ip.is_global: + raise UnsafeTargetError("target resolves to a non-public address") + return parsed