FastAPI + Celery + Next.js + Postgres/Redis app with company monitoring, source collection, LLM-based change analysis, enrichment, and account security (Turnstile, escalating lockout, email verification).
177 lines
5.5 KiB
Python
177 lines
5.5 KiB
Python
"""SSRF-safe HTTP fetching. Every collector and the custom-URL feature must
|
|
route network requests through `safe_fetch` / `fetch_with_retries` - see
|
|
SECURITY.md for the full threat model this defends against.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import ipaddress
|
|
import socket
|
|
import time
|
|
from dataclasses import dataclass
|
|
from urllib.parse import urljoin, urlparse
|
|
|
|
import httpx
|
|
import tenacity
|
|
|
|
from app.core.config import Settings, get_settings
|
|
from app.core.logging import get_logger
|
|
|
|
logger = get_logger(__name__)
|
|
|
|
_ALLOWED_SCHEMES = {"http", "https"}
|
|
_MAX_REDIRECTS = 5
|
|
_METADATA_IPS = {"169.254.169.254", "fd00:ec2::254"}
|
|
|
|
# Per-hostname request spacing. In-process only - a multi-worker Celery
|
|
# deployment would need a shared store (e.g. Redis) for this to be a true
|
|
# global rate limit across workers; see KNOWN_LIMITATIONS.md.
|
|
_last_request_at: dict[str, float] = {}
|
|
_domain_locks: dict[str, asyncio.Lock] = {}
|
|
|
|
|
|
class SsrfBlockedError(Exception):
|
|
"""Raised when a URL resolves to, or points at, a disallowed network target."""
|
|
|
|
|
|
class FetchError(Exception):
|
|
"""Raised for network-level failures after retries are exhausted."""
|
|
|
|
|
|
def _is_blocked_ip(ip: ipaddress.IPv4Address | ipaddress.IPv6Address) -> bool:
|
|
return (
|
|
ip.is_private
|
|
or ip.is_loopback
|
|
or ip.is_link_local
|
|
or ip.is_multicast
|
|
or ip.is_reserved
|
|
or ip.is_unspecified
|
|
or str(ip) in _METADATA_IPS
|
|
)
|
|
|
|
|
|
def _resolve_and_validate(hostname: str) -> None:
|
|
try:
|
|
infos = socket.getaddrinfo(hostname, None)
|
|
except socket.gaierror as exc:
|
|
raise SsrfBlockedError(f"Could not resolve host: {hostname}") from exc
|
|
|
|
if not infos:
|
|
raise SsrfBlockedError(f"Could not resolve host: {hostname}")
|
|
|
|
for info in infos:
|
|
raw_ip = info[4][0]
|
|
ip = ipaddress.ip_address(raw_ip.split("%")[0])
|
|
if _is_blocked_ip(ip):
|
|
raise SsrfBlockedError(f"Resolved address for {hostname} is not a public address: {ip}")
|
|
|
|
|
|
def validate_url(url: str) -> str:
|
|
"""Raises SsrfBlockedError if `url` is unsafe to fetch. Returns the
|
|
hostname."""
|
|
parsed = urlparse(url)
|
|
if parsed.scheme not in _ALLOWED_SCHEMES:
|
|
raise SsrfBlockedError(f"Unsupported URL scheme: {parsed.scheme!r}")
|
|
if not parsed.hostname:
|
|
raise SsrfBlockedError("URL has no hostname")
|
|
_resolve_and_validate(parsed.hostname)
|
|
return parsed.hostname
|
|
|
|
|
|
async def _respect_domain_delay(hostname: str, delay_seconds: float) -> None:
|
|
if delay_seconds <= 0:
|
|
return
|
|
lock = _domain_locks.setdefault(hostname, asyncio.Lock())
|
|
async with lock:
|
|
now = time.monotonic()
|
|
last = _last_request_at.get(hostname)
|
|
if last is not None:
|
|
elapsed = now - last
|
|
if elapsed < delay_seconds:
|
|
await asyncio.sleep(delay_seconds - elapsed)
|
|
_last_request_at[hostname] = time.monotonic()
|
|
|
|
|
|
@dataclass
|
|
class SafeFetchResult:
|
|
status_code: int
|
|
text: str
|
|
content: bytes
|
|
headers: dict[str, str]
|
|
final_url: str
|
|
|
|
|
|
async def safe_fetch(
|
|
url: str,
|
|
*,
|
|
settings: Settings | None = None,
|
|
method: str = "GET",
|
|
max_redirects: int = _MAX_REDIRECTS,
|
|
extra_headers: dict[str, str] | None = None,
|
|
) -> SafeFetchResult:
|
|
"""Fetch `url` with SSRF validation applied to the initial URL and every
|
|
redirect hop. Never follows a redirect without re-validating it."""
|
|
settings = settings or get_settings()
|
|
current_url = url
|
|
|
|
for _ in range(max_redirects + 1):
|
|
hostname = validate_url(current_url)
|
|
await _respect_domain_delay(hostname, settings.scraper_domain_delay_seconds)
|
|
|
|
headers = {"User-Agent": settings.scraper_user_agent, **(extra_headers or {})}
|
|
async with httpx.AsyncClient(
|
|
follow_redirects=False,
|
|
timeout=settings.scraper_request_timeout_seconds,
|
|
headers=headers,
|
|
) as client:
|
|
response = await client.request(method, current_url)
|
|
|
|
if response.status_code in (301, 302, 303, 307, 308) and "location" in response.headers:
|
|
current_url = urljoin(current_url, response.headers["location"])
|
|
continue
|
|
|
|
return SafeFetchResult(
|
|
status_code=response.status_code,
|
|
text=response.text,
|
|
content=response.content,
|
|
headers=dict(response.headers),
|
|
final_url=current_url,
|
|
)
|
|
|
|
raise SsrfBlockedError(f"Too many redirects starting from {url}")
|
|
|
|
|
|
def _is_retryable(exc: BaseException) -> bool:
|
|
if isinstance(exc, SsrfBlockedError):
|
|
return False
|
|
if isinstance(exc, httpx.HTTPError):
|
|
return True
|
|
return False
|
|
|
|
|
|
async def fetch_with_retries(
|
|
url: str,
|
|
*,
|
|
settings: Settings | None = None,
|
|
max_attempts: int = 3,
|
|
**kwargs,
|
|
) -> SafeFetchResult:
|
|
"""`safe_fetch` wrapped with exponential-backoff retry for transient
|
|
network errors only - SSRF blocks and 4xx responses are not retried."""
|
|
settings = settings or get_settings()
|
|
|
|
async for attempt in tenacity.AsyncRetrying(
|
|
stop=tenacity.stop_after_attempt(max_attempts),
|
|
wait=tenacity.wait_exponential(multiplier=1, min=1, max=10),
|
|
retry=tenacity.retry_if_exception(_is_retryable),
|
|
reraise=True,
|
|
):
|
|
with attempt:
|
|
result = await safe_fetch(url, settings=settings, **kwargs)
|
|
if result.status_code >= 500:
|
|
raise FetchError(f"Server error {result.status_code} fetching {url}")
|
|
return result
|
|
|
|
raise FetchError(f"Exhausted retries fetching {url}") # pragma: no cover
|