"""安全工具 — 文件上传校验和 URL 安全校验。""" import ipaddress import io import logging import os import socket logger = logging.getLogger(__name__) # PDF 文件头魔数 PDF_MAGIC = b"%PDF" # 文件大小限制(默认 30MB) MAX_FILE_SIZE = int(os.environ.get("MAX_UPLOAD_SIZE_MB", "30")) * 1024 * 1024 # PDF 页数限制(默认 200 页) MAX_PDF_PAGES = int(os.environ.get("MAX_PDF_PAGES", "200")) # 私网/保留 IP 前缀 _PRIVATE_NETWORKS = [ ipaddress.ip_network("127.0.0.0/8"), # loopback ipaddress.ip_network("10.0.0.0/8"), # 私网 A ipaddress.ip_network("172.16.0.0/12"), # 私网 B ipaddress.ip_network("192.168.0.0/16"), # 私网 C ipaddress.ip_network("169.254.0.0/16"), # 链路本地 ipaddress.ip_network("0.0.0.0/8"), # 当前网络 ipaddress.ip_network("::1/128"), # IPv6 loopback ipaddress.ip_network("fc00::/7"), # IPv6 私网 ipaddress.ip_network("fe80::/10"), # IPv6 链路本地 ] def prepare_pdf_upload(file_storage, password: str = "") -> tuple[bool, str, bytes | None]: """校验 PDF;加密文件解密后返回可供解析的 PDF 字节。 密码不会写入数据库或日志。空密码可自动处理仅限制复制/打印的 PDF。 返回 (is_valid, error_message, pdf_bytes)。 """ file_storage.seek(0) original_bytes = file_storage.read() file_storage.seek(0) # 1. 检查文件大小 size = len(original_bytes) if size == 0: return False, "文件为空", None if size > MAX_FILE_SIZE: return False, f"文件大小 {size // (1024*1024)}MB 超过限制 {MAX_FILE_SIZE // (1024*1024)}MB", None # 2. 检查 PDF 魔数 if not original_bytes[:8].startswith(PDF_MAGIC): return False, "文件不是有效的 PDF 格式", None # 3. 解密并检查页数 try: import pypdf reader = pypdf.PdfReader(io.BytesIO(original_bytes)) pdf_bytes = original_bytes if reader.is_encrypted: decrypt_result = reader.decrypt(password or "") if decrypt_result == 0: message = "PDF 密码错误,请重新输入" if password else "该 PDF 需要打开密码,请填写密码后重试" return False, message, None writer = pypdf.PdfWriter() writer.append_pages_from_reader(reader) output = io.BytesIO() writer.write(output) pdf_bytes = output.getvalue() # 4. 检查页数 if len(reader.pages) > MAX_PDF_PAGES: return False, f"PDF 页数 {len(reader.pages)} 超过限制 {MAX_PDF_PAGES} 页", None except ImportError: # pypdf 不可用时跳过加密和页数检查 pdf_bytes = original_bytes except Exception as e: logger.warning(f"PDF 解析检查失败: {e}") return False, "PDF 文件内容损坏、密码错误或格式无效", None finally: file_storage.seek(0) return True, "", pdf_bytes def validate_pdf_upload(file_storage, password: str = "") -> tuple[bool, str]: """兼容旧调用:只返回 PDF 校验结果。""" is_valid, error_message, _ = prepare_pdf_upload(file_storage, password) return is_valid, error_message def is_safe_base_url(url: str) -> tuple[bool, str]: """检查 Base URL 是否安全(拒绝私网、环回、元数据地址)。 返回 (is_safe, error_message)。 """ if not url: return True, "" # 只允许 http/https if not url.startswith(("http://", "https://")): return False, "仅支持 http/https 协议" # 解析主机名 from urllib.parse import urlparse try: parsed = urlparse(url) except Exception: return False, "URL 格式无效" hostname = parsed.hostname if not hostname: return False, "URL 缺少主机名" # 拒绝明确的 localhost if hostname in ("localhost", "0.0.0.0"): return False, f"不允许访问 {hostname}" # DNS 解析后检查 IP try: for info in socket.getaddrinfo(hostname, None, socket.AF_UNSPEC, socket.SOCK_STREAM): ip_str = info[4][0] ip = ipaddress.ip_address(ip_str) for network in _PRIVATE_NETWORKS: if ip in network: return False, f"不允许访问私网/保留地址 ({ip_str})" except socket.gaierror: pass # DNS 解析失败不阻止,让后续调用报错 except Exception as e: logger.warning(f"URL 安全检查异常: {e}") return True, "" def _is_private_ip(ip_str: str) -> bool: """检查 IP 是否属于私网/保留地址。""" try: ip = ipaddress.ip_address(ip_str) for network in _PRIVATE_NETWORKS: if ip in network: return True except ValueError: pass return False class _SafeHTTPTransport: """httpcore HTTP 传输层,在连接时验证目标 IP 不是私网地址(防止 DNS rebinding SSRF)。""" def __init__(self, **kwargs): import httpcore self._transport = httpcore.HTTPTransport(**kwargs) def handle_request(self, request): url = request.url host = url.host port = url.port or (443 if url.scheme == b"https" else 80) try: for info in socket.getaddrinfo(host, port, socket.AF_UNSPEC, socket.SOCK_STREAM): ip_str = info[4][0] if _is_private_ip(ip_str): raise ValueError(f"SSRF 阻止: 目标 {host} 解析到私网地址 {ip_str}") except socket.gaierror: pass # DNS 失败时让底层处理 return self._transport.handle_request(request) class _SafeHTTPAsyncTransport: """httpcore 异步 HTTP 传输层,连接时验证 IP。""" def __init__(self, **kwargs): import httpcore self._transport = httpcore.AsyncHTTPTransport(**kwargs) async def handle_async_request(self, request): url = request.url host = url.host port = url.port or (443 if url.scheme == b"https" else 80) try: for info in socket.getaddrinfo(host, port, socket.AF_UNSPEC, socket.SOCK_STREAM): ip_str = info[4][0] if _is_private_ip(ip_str): raise ValueError(f"SSRF 阻止: 目标 {host} 解析到私网地址 {ip_str}") except socket.gaierror: pass return await self._transport.handle_async_request(request) def safe_httpx_client(**kwargs): """创建带有 SSRF 防护的 httpx.Client(每次连接时重新验证 DNS)。""" import httpx transport = _SafeHTTPTransport(retries=kwargs.pop("retries", 0)) return httpx.Client(transport=transport, **kwargs) def safe_httpx_async_client(**kwargs): """创建带有 SSRF 防护的 httpx.AsyncClient。""" import httpx transport = _SafeHTTPAsyncTransport(retries=kwargs.pop("retries", 0)) return httpx.AsyncClient(transport=transport, **kwargs)