"""SSRF 防护工具:URL 协议白名单 + 私网/保留地址拦截。 供 image_fix、webhook_store 等需要外发请求的模块共享,避免每个调用点各写一份。 """ from __future__ import annotations import ipaddress import socket from typing import Optional from urllib.parse import urlparse from ..errors import HttpError # 允许的外发请求协议 ALLOWED_SCHEMES = {"http", "https"} def _is_dangerous_ip(ip: ipaddress._BaseAddress) -> bool: """判断 IP 是否属于私网/回环/链路本地/保留/多播/未指定。""" return ( ip.is_private or ip.is_loopback or ip.is_link_local or ip.is_reserved or ip.is_multicast or ip.is_unspecified ) def _is_dangerous_host(host: str) -> bool: """判断主机名是否解析到危险 IP(含 DNS rebinding 检测)。""" if not host: return True # 直接是 IP 字面量 try: ip = ipaddress.ip_address(host) return _is_dangerous_ip(ip) except ValueError: pass # 域名:解析所有 A/AAAA 记录,任一命中危险段即拒绝 try: resolved = socket.getaddrinfo(host, None) except Exception: # DNS 解析失败:保守按危险处理(防止临时解析失败绕过校验) return True for _, _, _, _, sockaddr in resolved: ip_str = sockaddr[0] try: ip = ipaddress.ip_address(ip_str) except ValueError: continue if _is_dangerous_ip(ip): return True return False def validate_outbound_url(url: str, *, max_length: int = 2048) -> None: """校验外发请求 URL,拒绝 SSRF 目标。 拦截: - 非 http/https 协议(file://、gopher://、dict:// 等) - 私网 / 回环 / 链路本地 / 保留 / 多播 / 未指定 IP - 解析失败或为空的主机 - 过长 URL(防止解析器边界情况) 注意:该校验只在请求发起前做一次。若目标域名后续被 DNS rebinding 改写, 仍可能绕过——因此调用方若使用 follow_redirects,应在每次重定向后重新校验。 """ if not url or not isinstance(url, str): raise HttpError("URL is required", status=400, code="bad_request") if len(url) > max_length: raise HttpError( f"URL too long (>{max_length} chars)", status=400, code="bad_request", ) parsed = urlparse(url) scheme = (parsed.scheme or "").lower() if scheme not in ALLOWED_SCHEMES: raise HttpError( f"unsupported URL scheme: {parsed.scheme!r} (only http/https allowed)", status=400, code="bad_request", ) host = parsed.hostname or "" if not host: raise HttpError("URL missing host", status=400, code="bad_request") if _is_dangerous_host(host): raise HttpError( f"URL host resolves to a forbidden (private/loopback/reserved) address: {host}", status=400, code="bad_request", )