Spaces:
Running
Running
| """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", | |
| ) | |