from __future__ import annotations import io import hashlib import ipaddress import socket from urllib.parse import urljoin, urlparse import requests from PIL import Image DEFAULT_MAX_BYTES = 50 * 1024 * 1024 # 50MB DEFAULT_MAX_REDIRECTS = 3 class ImageLoadError(Exception): """Image failure with stable internal diagnostics and unchanged message text.""" def __init__(self, message: str, *, category: str = "unknown", stage: str = "fetch", original_exception_class: str | None = None, metadata: dict[str, object] | None = None): super().__init__(message) self.category = category self.stage = stage self.original_exception_class = original_exception_class self.metadata = metadata or {} def image_url_log_context(url: str) -> dict[str, str]: """Safe URL fields for logs; query strings, credentials and full paths stay private.""" try: parsed = urlparse(url) result = {"host": parsed.hostname or "unknown", "path_hash": hashlib.sha256(parsed.path.encode("utf-8", "replace")).hexdigest()[:12]} if parsed.scheme: result["scheme"] = parsed.scheme.lower() if parsed.port: result["port"] = str(parsed.port) return result except Exception: return {} def _error(message: str, *, category: str, stage: str, cause: BaseException | None = None, metadata: dict[str, object] | None = None) -> ImageLoadError: return ImageLoadError(message, category=category, stage=stage, original_exception_class=type(cause).__name__ if cause else None, metadata=metadata) def _validate_public_url(url: str) -> str: try: parsed = urlparse(url) except ValueError as e: raise _error(str(e), category="invalid_url", stage="validate", cause=e) from e if parsed.scheme not in ("http", "https"): raise _error("Only http and https URLs are allowed", category="invalid_scheme", stage="validate") if not parsed.hostname: raise _error("URL must include a hostname", category="missing_hostname", stage="validate") hostname = parsed.hostname.strip().lower() if hostname in {"localhost", "127.0.0.1", "::1"}: raise _error("Localhost URLs are not allowed", category="blocked_address", stage="validate") try: port = parsed.port resolved = socket.getaddrinfo(hostname, port or (443 if parsed.scheme == "https" else 80), type=socket.SOCK_STREAM) except ValueError as e: raise _error(str(e), category="invalid_url", stage="validate", cause=e) from e except socket.gaierror as e: raise _error(f"Cannot resolve host: {e}", category="dns_failure", stage="validate", cause=e) from e except (TimeoutError, socket.timeout) as e: raise _error("Host resolution timed out", category="connect_timeout", stage="validate", cause=e) from e for entry in resolved: address = entry[4][0] ip = ipaddress.ip_address(address) if ( ip.is_private or ip.is_loopback or ip.is_link_local or ip.is_multicast or ip.is_reserved or ip.is_unspecified ): raise _error("URLs resolving to private or reserved addresses are not allowed", category="blocked_address", stage="validate") return url def fetch_url_bytes(url: str, timeout: float = 10.0, max_bytes: int = DEFAULT_MAX_BYTES) -> bytes: current_url = _validate_public_url(url) try: for _ in range(DEFAULT_MAX_REDIRECTS + 1): with requests.get(current_url, stream=True, timeout=timeout, allow_redirects=False) as r: if 300 <= r.status_code < 400: location = r.headers.get("location") if not location: raise _error("Redirect response missing location header", category="redirect_missing_location", stage="fetch", metadata={"remote_status": r.status_code}) current_url = _validate_public_url(urljoin(current_url, location)) continue try: r.raise_for_status() except requests.HTTPError as e: category = "remote_http_4xx" if 400 <= r.status_code < 500 else "remote_http_5xx" if r.status_code >= 500 else "fetch_failure" raise _error(str(e), category=category, stage="fetch", cause=e, metadata={"remote_status": r.status_code}) from e content_type = (r.headers.get("content-type") or "").lower() if content_type and not content_type.startswith("image/"): raise _error( f"URL does not point to an image content type: {content_type}", category="invalid_content_type", stage="fetch", metadata={"content_type": content_type, "remote_status": r.status_code}, ) buf = io.BytesIO() total = 0 for chunk in r.iter_content(chunk_size=1024 * 64): if not chunk: continue total += len(chunk) if total > max_bytes: raise _error(f"Image exceeds max_bytes={max_bytes}", category="image_too_large", stage="fetch", metadata={"remote_status": r.status_code}) buf.write(chunk) return buf.getvalue() raise _error(f"Too many redirects (>{DEFAULT_MAX_REDIRECTS})", category="too_many_redirects", stage="fetch") except ImageLoadError: raise except requests.exceptions.ConnectTimeout as e: raise _error(f"Cannot fetch image url: {e}", category="connect_timeout", stage="fetch", cause=e) from e except requests.exceptions.ReadTimeout as e: raise _error(f"Cannot fetch image url: {e}", category="read_timeout", stage="fetch", cause=e) from e except requests.exceptions.Timeout as e: raise _error(f"Cannot fetch image url: {e}", category="read_timeout", stage="fetch", cause=e) from e except requests.exceptions.TooManyRedirects as e: raise _error(f"Cannot fetch image url: {e}", category="too_many_redirects", stage="fetch", cause=e) from e except requests.exceptions.ConnectionError as e: raise _error(f"Cannot fetch image url: {e}", category="connection_error", stage="fetch", cause=e) from e except Exception as e: raise _error(f"Cannot fetch image url: {e}", category="fetch_failure", stage="fetch", cause=e) from e def bytes_to_pil(data: bytes) -> Image.Image: try: img = Image.open(io.BytesIO(data)).convert("RGB") return img except Exception as e: raise _error(f"Cannot decode image: {e}", category="decode_failure", stage="decode", cause=e)