Compare commits

...

1 Commits

Author SHA1 Message Date
AI Bot c7a7c019a3 test(wave128): url_security纯逻辑抽离+110单测
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 5s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m2s
CI/CD Pipeline / Validate - Code Quality (pull_request) Failing after 1m37s
CI/CD Pipeline / Frontend Lint (pull_request) Failing after 36s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m6s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m59s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 3m5s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 1m36s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 45s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m0s
AI Code Review / AI Code Review (pull_request) Successful in 6m24s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 0s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Failing after 1s
CI/CD Pipeline / CI Gate (pull_request) Failing after 0s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Failing after 0s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 46m23s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 41s
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
- 新建 packages/domain/url_security.py:SSRF防护纯逻辑模块
  - URL基础校验(scheme/port/hostname/internal/IP-SSRF/trusted-domain)
  - 魔数校验纯函数(接收bytes)
  - 常量、异常类、辅助函数
- packages/shared/url_security.py 改为薄包装
  - 保留DNS解析、文件下载等有副作用逻辑
  - 纯逻辑委托给domain层,完全向后兼容
- 110个单测全覆盖:常量/内部主机名/可信域名/SSRF IP/IP判断/URL校验/魔数校验
- 修复check_ssrf_ip检查顺序:link_local/unspecified放private前
2026-07-27 18:25:30 +08:00
3 changed files with 848 additions and 362 deletions
+338
View File
@@ -0,0 +1,338 @@
"""URL 安全校验纯逻辑 — SSRF 防护.
纯函数模块,无网络/文件 IO,无环境变量依赖,所有配置通过参数传入。
供 packages/shared/url_security.py 作为薄包装调用,也可直接用于单测。
防护要点(纯逻辑部分):
1. Scheme 白名单校验
2. 内部主机名拦截(字符串匹配)
3. 端口白名单校验
4. IP 格式 SSRF 检查(回环/私有/链路本地/组播/未指定/保留)
5. 可信域名匹配(支持子域名)
6. 文件头魔数校验(接收 bytes)
"""
from __future__ import annotations
import ipaddress
from urllib.parse import urlparse
class UrlSecurityError(ValueError):
"""URL 安全校验失败."""
pass
# ── 常量 ────────────────────────────────────────────────────────────────────
ALLOWED_SCHEMES = frozenset({"http", "https"})
ALLOWED_PORTS = frozenset({80, 443})
MAX_URL_LENGTH = 2048
# 已知内部/敏感主机名集合
INTERNAL_HOSTNAMES = frozenset(
{
"localhost",
"localhost.localdomain",
"ip6-localhost",
"ip6-loopback",
"metadata",
"metadata.google.internal",
"169.254.169.254",
}
)
# 内网域名后缀
INTERNAL_DOMAIN_SUFFIXES = (".local", ".internal", ".localdomain")
# 文件魔数表 — key: MIME, value: list of 签名组,每组内所有 (offset, bytes) 都匹配才算命中
MAGIC_NUMBERS: dict[str, list[list[tuple[int, bytes]]]] = {
# 音频
"audio/mpeg": [
[(0, b"ID3")],
[(0, b"\xff\xfb")],
[(0, b"\xff\xf3")],
[(0, b"\xff\xf2")],
[(0, b"\xff\xfa")],
[(0, b"\xff\xf9")],
],
"audio/wav": [[(0, b"RIFF"), (8, b"WAVE")]],
"audio/x-wav": [[(0, b"RIFF"), (8, b"WAVE")]],
"audio/ogg": [[(0, b"OggS")]],
"application/ogg": [[(0, b"OggS")]],
"audio/flac": [[(0, b"fLaC")]],
"audio/aac": [[(0, b"\xff\xf1")], [(0, b"\xff\xf9")]],
"audio/aacp": [[(0, b"\xff\xf1")], [(0, b"\xff\xf9")]],
"audio/mp4": [[(4, b"ftyp")]],
"audio/x-m4a": [[(4, b"ftyp")]],
# 视频
"video/mp4": [[(4, b"ftyp")]],
"video/quicktime": [[(4, b"ftyp")]],
"video/x-matroska": [[(0, b"\x1a\x45\xdf\xa3")]],
"video/webm": [[(0, b"\x1a\x45\xdf\xa3")]],
"video/x-msvideo": [[(0, b"RIFF"), (8, b"AVI ")]],
# 图片
"image/jpeg": [[(0, b"\xff\xd8\xff")]],
"image/png": [[(0, b"\x89PNG\r\n\x1a\n")]],
"image/gif": [[(0, b"GIF87a")], [(0, b"GIF89a")]],
"image/webp": [[(0, b"RIFF"), (8, b"WEBP")]],
"image/bmp": [[(0, b"BM")]],
}
ALLOWED_AUDIO_MIME_TYPES = frozenset(
{
"audio/mpeg",
"audio/mp3",
"audio/wav",
"audio/x-wav",
"audio/pcm",
"audio/ogg",
"audio/opus",
"audio/flac",
"audio/aac",
"audio/m4a",
"audio/x-m4a",
"audio/mp4",
"application/octet-stream",
}
)
ALLOWED_VIDEO_MIME_TYPES = frozenset(
{
"video/mp4",
"video/quicktime",
"video/x-matroska",
"video/webm",
"video/avi",
"video/x-msvideo",
"video/mpeg",
"application/octet-stream",
}
)
ALLOWED_IMAGE_MIME_TYPES = frozenset(
{
"image/jpeg",
"image/png",
"image/gif",
"image/webp",
"image/bmp",
}
)
# ── 主机名 / 域名校验 ──────────────────────────────────────────────────────
def check_internal_hostname(hostname: str) -> None:
"""检查主机名是否为内部/敏感主机名,是则抛出 UrlSecurityError.
检查内容:
- 精确匹配 INTERNAL_HOSTNAMES 集合
- 后缀匹配 INTERNAL_DOMAIN_SUFFIXES.local/.internal/.localdomain
"""
hostname_lower = hostname.lower()
if hostname_lower in INTERNAL_HOSTNAMES:
raise UrlSecurityError(f"禁止访问内部主机名: {hostname}")
if hostname_lower.endswith(INTERNAL_DOMAIN_SUFFIXES):
raise UrlSecurityError(f"禁止访问内网域名: {hostname}")
def is_trusted_domain(hostname: str, trusted_domains: set[str]) -> bool:
"""检查域名是否在可信白名单中(支持子域名匹配).
匹配规则:
- 精确匹配
- 子域名匹配(hostname 以 .domain 结尾)
Args:
hostname: 待检查的主机名
trusted_domains: 可信域名集合,为空表示不限制
"""
if not trusted_domains:
return True
hostname_lower = hostname.lower()
if hostname_lower in {d.lower() for d in trusted_domains}:
return True
for domain in trusted_domains:
if hostname_lower.endswith("." + domain.lower()):
return True
return False
# ── IP SSRF 检查 ───────────────────────────────────────────────────────────
def check_ssrf_ip(ip_str: str) -> None:
"""检查 IP 地址是否存在 SSRF 风险,有风险则抛出 UrlSecurityError.
检查项:回环、私有、链路本地、组播、未指定、保留地址。
Args:
ip_str: IP 地址字符串(IPv4 或 IPv6
Raises:
UrlSecurityError: IP 属于 SSRF 风险范围
ValueError: ip_str 不是合法 IP 地址(调用方应自行捕获处理)
"""
ip_obj = ipaddress.ip_address(ip_str)
if ip_obj.is_loopback:
raise UrlSecurityError(f"禁止访问回环地址: {ip_obj}")
if ip_obj.is_link_local:
raise UrlSecurityError(f"禁止访问链路本地地址: {ip_obj}")
if ip_obj.is_unspecified:
raise UrlSecurityError(f"禁止访问未指定地址: {ip_obj}")
if ip_obj.is_multicast:
raise UrlSecurityError(f"禁止访问组播地址: {ip_obj}")
if ip_obj.is_reserved:
raise UrlSecurityError(f"禁止访问保留地址: {ip_obj}")
if ip_obj.is_private:
raise UrlSecurityError(f"禁止访问内网地址: {ip_obj}")
def is_ip_address(hostname: str) -> bool:
"""判断主机名是否为 IP 地址格式(IPv4 或 IPv6."""
try:
ipaddress.ip_address(hostname)
return True
except ValueError:
return False
# ── URL 基础校验 ───────────────────────────────────────────────────────────
def validate_url_basic(
url: str,
*,
trusted_domains: set[str] | None = None,
allow_direct_ip: bool = False,
) -> str:
"""URL 基础安全校验(纯逻辑,不含 DNS 解析).
校验项:
1. URL 非空 & 长度限制
2. Scheme 白名单
3. 主机名存在性
4. 内部主机名拦截
5. 端口白名单
6. 直接 IP 访问限制
7. IP 格式 SSRF 检查(如果 hostname 是 IP
8. 可信域名白名单(如果配置了)
注意:域名格式的 SSRF 检查需要 DNS 解析,不在本函数范围内。
Args:
url: 待校验 URL
trusted_domains: 可信域名白名单,None/空表示不限制
allow_direct_ip: 是否允许直接 IP 访问
Returns:
原始 URL(校验通过)
Raises:
UrlSecurityError: 校验失败
"""
if not url:
raise UrlSecurityError("URL 为空")
if len(url) > MAX_URL_LENGTH:
raise UrlSecurityError(f"URL 过长 ({len(url)} > {MAX_URL_LENGTH})")
try:
parsed = urlparse(url)
except Exception as e:
raise UrlSecurityError(f"URL 解析失败: {e}") from e
# Scheme
if not parsed.scheme or parsed.scheme.lower() not in ALLOWED_SCHEMES:
raise UrlSecurityError(f"不允许的 URL scheme: {parsed.scheme}")
# Hostname
hostname = parsed.hostname
if not hostname:
raise UrlSecurityError("URL 缺少主机名")
# 内部主机名前置拦截
check_internal_hostname(hostname)
# 端口
port = parsed.port
if port is not None and port not in ALLOWED_PORTS:
raise UrlSecurityError(f"不允许的端口: {port}")
# IP 格式检查 & SSRF
if is_ip_address(hostname):
if not allow_direct_ip:
raise UrlSecurityError(f"禁止直接 IP 访问: {hostname}")
check_ssrf_ip(hostname)
# 可信域名白名单
if trusted_domains and not is_trusted_domain(hostname, trusted_domains):
raise UrlSecurityError(f"域名不在可信白名单中: {hostname}")
return url
def is_url_basic_safe(
url: str,
*,
trusted_domains: set[str] | None = None,
allow_direct_ip: bool = False,
) -> bool:
"""便捷函数:基础安全检查,不抛异常,返回 bool."""
try:
validate_url_basic(url, trusted_domains=trusted_domains, allow_direct_ip=allow_direct_ip)
return True
except UrlSecurityError:
return False
# ── 魔数校验 ───────────────────────────────────────────────────────────────
def validate_magic_number(header_bytes: bytes, allowed_mime_types: set[str]) -> None:
"""校验文件头魔数是否与允许的 MIME 类型匹配(纯函数).
读取 header_bytes,与 allowed_mime_types 对应格式的魔数逐一比对,
任一类型匹配即通过;全部不匹配则抛出 UrlSecurityError。
Args:
header_bytes: 文件头字节(建议至少 256 字节)
allowed_mime_types: 允许的 MIME 类型集合
Raises:
UrlSecurityError: 文件魔数与所有允许类型均不匹配
"""
# 收集所有允许类型对应的魔数签名
signatures: list[list[tuple[int, bytes]]] = []
for mime in allowed_mime_types:
sigs = MAGIC_NUMBERS.get(mime)
if sigs:
signatures.extend(sigs)
# 没有已知魔数的 MIME,跳过不阻断
if not signatures:
return
if not header_bytes:
raise UrlSecurityError("文件为空,无法校验格式")
# 任一签名匹配即通过
for sig in signatures:
match = True
for offset, expected in sig:
if offset + len(expected) > len(header_bytes):
match = False
break
if header_bytes[offset : offset + len(expected)] != expected:
match = False
break
if match:
return
raise UrlSecurityError(
f"文件魔数与允许的 MIME 类型不匹配,"
f"允许类型: {sorted(allowed_mime_types)}"
f"文件头前16字节: {header_bytes[:16].hex()}"
)
+87 -362
View File
@@ -1,239 +1,73 @@
"""URL 安全校验工具 — SSRF 防护. """URL 安全校验工具 — SSRF 防护(薄包装层).
统一的外部 URL 安全校验方案,覆盖所有渲染管线和 TTS 中的外部下载场景 本文件保留原有对外 API,纯逻辑部分委托给 packages/domain/url_security.py
放在 packages/shared/ 作为单一来源,worker 和 application 层都可引用 新增了 DNS 解析、文件下载、环境变量配置等有副作用的逻辑
防护要点: 防护要点:
1. Scheme 白名单:仅允许 http/https 1. Scheme 白名单:仅允许 http/https
2. 主机 SSRF 防护:禁止内网 IP、回环地址、链路本地地址、元数据服务 2. 主机 SSRF 防护:禁止内网 IP、回环地址、链路本地地址、元数据服务
3. 端口白名单:仅允许 80/443(标准 HTTP/HTTPS 3. 端口白名单:仅允许 80/443
4. 域名校验:禁止 IP 直接访问(除非在白名单中) 4. 域名校验:禁止 IP 直接访问(除非在白名单中)
5. 重定向防护:手动跟随重定向,每次跳转前重新校验目标 URL 5. 重定向防护:手动跟随重定向,每次跳转前重新校验
6. 文件大小限制:流式下载,超过上限立即中断 6. 文件大小限制:流式下载,超过上限立即中断
7. MIME 类型白名单:可选的内容类型校验 7. MIME 类型白名单 + 魔数二次校验
""" """
from __future__ import annotations from __future__ import annotations
import ipaddress
import logging import logging
import os import os
import socket import socket
import urllib.error import urllib.error
import urllib.request import urllib.request
from urllib.parse import urljoin, urlparse from urllib.parse import urljoin
from packages.domain.url_security import (
ALLOWED_AUDIO_MIME_TYPES,
ALLOWED_IMAGE_MIME_TYPES,
ALLOWED_PORTS as _allowed_ports_base,
ALLOWED_SCHEMES as _allowed_schemes_base,
ALLOWED_VIDEO_MIME_TYPES,
MAX_URL_LENGTH,
MAGIC_NUMBERS,
UrlSecurityError as _UrlSecurityError_base,
check_internal_hostname as _check_internal_hostname_base,
check_ssrf_ip as _check_ssrf_ip_base,
is_ip_address as _is_ip_address_base,
is_trusted_domain as _is_trusted_domain_base,
validate_magic_number as _validate_magic_number_base,
validate_url_basic as _validate_url_basic_base,
)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# 允许的 URL scheme # ── 兼容导出(保持原有变量名供外部引用) ──────────────────────────────────
ALLOWED_SCHEMES = {"http", "https"} ALLOWED_SCHEMES = set(_allowed_schemes_base)
ALLOWED_PORTS = set(_allowed_ports_base)
UrlSecurityError = _UrlSecurityError_base
# 允许的端口(标准 HTTP/HTTPS # 可信域名白名单(从环境变量读取
ALLOWED_PORTS = {80, 443}
# 可信域名白名单(可根据实际 OSS/CDN 域名配置)
# 从环境变量读取,格式:"oss-cn-hangzhou.aliyuncs.com,cdn.example.com"
# 默认空表示所有公网域名都允许,但仍会做 SSRF 检查
TRUSTED_DOMAINS: set[str] = set() TRUSTED_DOMAINS: set[str] = set()
_env_trusted = os.environ.get("URL_SECURITY_TRUSTED_DOMAINS", "") _env_trusted = os.environ.get("URL_SECURITY_TRUSTED_DOMAINS", "")
if _env_trusted: if _env_trusted:
TRUSTED_DOMAINS = {d.strip() for d in _env_trusted.split(",") if d.strip()} TRUSTED_DOMAINS = {d.strip() for d in _env_trusted.split(",") if d.strip()}
# 是否允许 IP 直接访问(默认禁止,防止绕过 DNS 校验) # 是否允许 IP 直接访问
ALLOW_DIRECT_IP = os.environ.get("URL_SECURITY_ALLOW_DIRECT_IP", "false").lower() == "true" ALLOW_DIRECT_IP = os.environ.get("URL_SECURITY_ALLOW_DIRECT_IP", "false").lower() == "true"
# 最大 URL 长度
MAX_URL_LENGTH = 2048
# 单次下载最大文件大小(默认 200MB) # 单次下载最大文件大小(默认 200MB)
DEFAULT_MAX_DOWNLOAD_SIZE = int(os.environ.get("URL_SECURITY_MAX_DOWNLOAD_MB", "200")) * 1024 * 1024 DEFAULT_MAX_DOWNLOAD_SIZE = int(os.environ.get("URL_SECURITY_MAX_DOWNLOAD_MB", "200")) * 1024 * 1024
# 允许的音频 MIME 类型白名单
ALLOWED_AUDIO_MIME_TYPES = {
"audio/mpeg",
"audio/mp3",
"audio/wav",
"audio/x-wav",
"audio/pcm",
"audio/ogg",
"audio/opus",
"audio/flac",
"audio/aac",
"audio/m4a",
"audio/x-m4a",
"audio/mp4",
"application/octet-stream", # 兼容一些 CDN 返回通用类型
}
# 允许的视频 MIME 类型白名单
ALLOWED_VIDEO_MIME_TYPES = {
"video/mp4",
"video/quicktime",
"video/x-matroska",
"video/webm",
"video/avi",
"video/x-msvideo",
"video/mpeg",
"application/octet-stream",
}
# 允许的图片 MIME 类型白名单
ALLOWED_IMAGE_MIME_TYPES = {
"image/jpeg",
"image/png",
"image/gif",
"image/webp",
"image/bmp",
}
# 下载块大小 # 下载块大小
_DOWNLOAD_CHUNK_SIZE = 8192 _DOWNLOAD_CHUNK_SIZE = 8192
# 最大重定向次数 # 最大重定向次数
_MAX_REDIRECTS = 5 _MAX_REDIRECTS = 5
# 文件魔数(文件头签名)表 — 用于 MIME 白名单校验后的二次真实性校验 # 魔数校验最大读取字节数
# key: MIME 类型,value: 签名列表,任一签名匹配即通过
# 每条签名: list of (offset, bytes),所有条目都匹配才算该签名命中(支持多处联合匹配如 RIFF+WAVE)
_MAGIC_NUMBERS: dict[str, list[list[tuple[int, bytes]]]] = {
# ── 音频 ──
"audio/mpeg": [
[(0, b"ID3")], # ID3v2 标签
[(0, b"\xff\xfb")], # MPEG1 Layer3
[(0, b"\xff\xf3")], # MPEG2 Layer3
[(0, b"\xff\xf2")], # MPEG2.5 Layer3
[(0, b"\xff\xfa")], # MPEG1 Layer2
[(0, b"\xff\xf9")], # 其他 MPEG ADTS
],
"audio/wav": [
[(0, b"RIFF"), (8, b"WAVE")], # RIFF + WAVE
],
"audio/x-wav": [
[(0, b"RIFF"), (8, b"WAVE")],
],
"audio/ogg": [
[(0, b"OggS")],
],
"application/ogg": [
[(0, b"OggS")],
],
"audio/flac": [
[(0, b"fLaC")],
],
"audio/aac": [
[(0, b"\xff\xf1")], # ADTS MPEG-4
[(0, b"\xff\xf9")], # ADTS MPEG-2
],
"audio/aacp": [
[(0, b"\xff\xf1")],
[(0, b"\xff\xf9")],
],
"audio/mp4": [
[(4, b"ftyp")], # ISO Base Media (M4A)
],
"audio/x-m4a": [
[(4, b"ftyp")],
],
# ── 视频 ──
"video/mp4": [
[(4, b"ftyp")], # ISO Base Media (MP4)
],
"video/quicktime": [
[(4, b"ftyp")],
],
"video/x-matroska": [
[(0, b"\x1a\x45\xdf\xa3")], # EBML header
],
"video/webm": [
[(0, b"\x1a\x45\xdf\xa3")],
],
"video/x-msvideo": [
[(0, b"RIFF"), (8, b"AVI ")],
],
# ── 图片 ──
"image/jpeg": [
[(0, b"\xff\xd8\xff")],
],
"image/png": [
[(0, b"\x89PNG\r\n\x1a\n")],
],
"image/gif": [
[(0, b"GIF87a")],
[(0, b"GIF89a")],
],
"image/webp": [
[(0, b"RIFF"), (8, b"WEBP")],
],
"image/bmp": [
[(0, b"BM")],
],
}
# 魔数校验最大读取字节数(文件头)
_MAGIC_CHECK_READ_SIZE = 256 _MAGIC_CHECK_READ_SIZE = 256
def _validate_magic_number(file_path: str, allowed_mime_types: set[str]) -> None:
"""校验文件头魔数是否与允许的 MIME 类型匹配.
读取文件前 256 字节,与 allowed_mime_types 对应格式的魔数逐一比对,
任一类型匹配即通过;全部不匹配则抛出 UrlSecurityError。
仅当 allowed_mime_types 非空时执行;空文件视为不匹配。
Args:
file_path: 本地文件路径
allowed_mime_types: 允许的 MIME 类型集合
Raises:
UrlSecurityError: 文件魔数与所有允许类型均不匹配
"""
# 收集所有允许类型对应的魔数签名
signatures: list[list[tuple[int, bytes]]] = []
for mime in allowed_mime_types:
sigs = _MAGIC_NUMBERS.get(mime)
if sigs:
signatures.extend(sigs)
# 如果没有已知魔数(比如自定义 MIME),跳过校验不阻断
if not signatures:
return
try:
with open(file_path, "rb") as f:
header = f.read(_MAGIC_CHECK_READ_SIZE)
except OSError as e:
raise UrlSecurityError(f"读取文件头失败: {e}") from e
if not header:
raise UrlSecurityError("文件为空,无法校验格式")
# 任一签名匹配即通过
for sig in signatures:
match = True
for offset, expected in sig:
if offset + len(expected) > len(header):
match = False
break
if header[offset : offset + len(expected)] != expected:
match = False
break
if match:
return
raise UrlSecurityError(
f"文件魔数与允许的 MIME 类型不匹配,"
f"允许类型: {sorted(allowed_mime_types)}"
f"文件头前16字节: {header[:16].hex()}"
)
class UrlSecurityError(ValueError):
"""URL 安全校验失败."""
pass
class NoRedirectHandler(urllib.request.HTTPRedirectHandler): class NoRedirectHandler(urllib.request.HTTPRedirectHandler):
"""禁止自动重定向的 handler,用于手动控制重定向以做安全校验.""" """禁止自动重定向的 handler,用于手动控制重定向以做安全校验."""
@@ -241,124 +75,7 @@ class NoRedirectHandler(urllib.request.HTTPRedirectHandler):
return None return None
def validate_url_safety(url: str, *, purpose: str = "download") -> str: # ── DNS 解析 SSRF 检查(有副作用) ─────────────────────────────────────────
"""校验 URL 安全性,返回标准化后的 URL(供下游使用).
Args:
url: 待校验的 URL
purpose: 用途描述(用于日志),如 "bgm_download""tts_download"
Returns:
标准化后的 URL
Raises:
UrlSecurityError: URL 不安全
"""
if not url:
raise UrlSecurityError("URL 为空")
if len(url) > MAX_URL_LENGTH:
raise UrlSecurityError(f"URL 过长 ({len(url)} > {MAX_URL_LENGTH})")
# 解析 URL
try:
parsed = urlparse(url)
except Exception as e:
raise UrlSecurityError(f"URL 解析失败: {e}") from e
# 1. Scheme 校验
if not parsed.scheme or parsed.scheme.lower() not in ALLOWED_SCHEMES:
raise UrlSecurityError(f"不允许的 URL scheme: {parsed.scheme}")
# 2. 主机名校验
hostname = parsed.hostname
if not hostname:
raise UrlSecurityError("URL 缺少主机名")
# 2.1 常见内网主机名前置拦截(防止 DNS rebinding 绕过)
_check_internal_hostnames(hostname)
# 3. 端口校验
port = parsed.port
if port is not None and port not in ALLOWED_PORTS:
raise UrlSecurityError(f"不允许的端口: {port}")
# 4. SSRF 防护 - 解析 IP 并检查
try:
# 先判断是否是 IP 地址
ip_obj = None
try:
ip_obj = ipaddress.ip_address(hostname)
except ValueError:
pass # 不是 IP,继续走域名解析
if ip_obj is not None:
# 是直接 IP 访问
if not ALLOW_DIRECT_IP and not _is_trusted_ip(ip_obj):
raise UrlSecurityError(f"禁止直接 IP 访问: {hostname}")
_check_ssrf_ip(ip_obj)
else:
# 域名 — 解析 DNS 检查 SSRF
_check_ssrf_domain(hostname)
except UrlSecurityError:
raise
except Exception as e:
logger.warning("URL 安全校验异常: url=%s purpose=%s error=%s", url[:80], purpose, e)
raise UrlSecurityError(f"URL 安全校验异常: {e}") from e
# 5. 可信域名校验(如果配置了白名单)
if TRUSTED_DOMAINS and not _is_trusted_domain(hostname):
raise UrlSecurityError(f"域名不在可信白名单中: {hostname}")
logger.debug("URL 安全校验通过: url=%s purpose=%s", url[:80], purpose)
return url
def _check_internal_hostnames(hostname: str) -> None:
"""前置检查常见内网/敏感主机名,防止 DNS 解析层绕过."""
hostname_lower = hostname.lower()
internal_hostnames = {
"localhost",
"localhost.localdomain",
"ip6-localhost",
"ip6-loopback",
"metadata",
"metadata.google.internal",
"169.254.169.254", # 云元数据服务
}
if hostname_lower in internal_hostnames:
raise UrlSecurityError(f"禁止访问内部主机名: {hostname}")
# 检查以 .local / .internal 结尾的主机名
if hostname_lower.endswith((".local", ".internal", ".localdomain")):
raise UrlSecurityError(f"禁止访问内网域名: {hostname}")
def _check_ssrf_ip(ip_obj: ipaddress.IPv4Address | ipaddress.IPv6Address) -> None:
"""检查 IP 是否属于 SSRF 风险范围."""
# 回环地址
if ip_obj.is_loopback:
raise UrlSecurityError(f"禁止访问回环地址: {ip_obj}")
# 私有地址(内网)
if ip_obj.is_private:
raise UrlSecurityError(f"禁止访问内网地址: {ip_obj}")
# 链路本地地址
if ip_obj.is_link_local:
raise UrlSecurityError(f"禁止访问链路本地地址: {ip_obj}")
# 组播地址
if ip_obj.is_multicast:
raise UrlSecurityError(f"禁止访问组播地址: {ip_obj}")
# 未指定地址(0.0.0.0 / ::
if ip_obj.is_unspecified:
raise UrlSecurityError(f"禁止访问未指定地址: {ip_obj}")
# 保留地址
if ip_obj.is_reserved:
raise UrlSecurityError(f"禁止访问保留地址: {ip_obj}")
def _check_ssrf_domain(hostname: str) -> None: def _check_ssrf_domain(hostname: str) -> None:
@@ -367,7 +84,6 @@ def _check_ssrf_domain(hostname: str) -> None:
注意:这不能完全防止 DNS rebinding,但能防御大部分 SSRF 场景。 注意:这不能完全防止 DNS rebinding,但能防御大部分 SSRF 场景。
""" """
try: try:
# 解析所有地址
infos = socket.getaddrinfo(hostname, None) infos = socket.getaddrinfo(hostname, None)
if not infos: if not infos:
raise UrlSecurityError(f"域名解析失败: {hostname}") raise UrlSecurityError(f"域名解析失败: {hostname}")
@@ -375,30 +91,63 @@ def _check_ssrf_domain(hostname: str) -> None:
for info in infos: for info in infos:
ip_str = info[4][0] ip_str = info[4][0]
try: try:
ip_obj = ipaddress.ip_address(ip_str) _check_ssrf_ip_base(ip_str)
_check_ssrf_ip(ip_obj)
except ValueError: except ValueError:
# 无法解析为 IP,跳过(不应该发生)
continue continue
except socket.gaierror as e: except socket.gaierror as e:
raise UrlSecurityError(f"域名解析失败: {hostname} ({e})") from e raise UrlSecurityError(f"域名解析失败: {hostname} ({e})") from e
def _is_trusted_ip(ip_obj: ipaddress.IPv4Address | ipaddress.IPv6Address) -> bool: def _validate_magic_number(file_path: str, allowed_mime_types: set[str]) -> None:
"""检查 IP 是否在可信列表中(目前通过环境变量配置域名,IP 级信任暂不开放.""" """校验文件头魔数(从文件读取后委托给 domain 纯逻辑."""
return False try:
with open(file_path, "rb") as f:
header = f.read(_MAGIC_CHECK_READ_SIZE)
except OSError as e:
raise UrlSecurityError(f"读取文件头失败: {e}") from e
_validate_magic_number_base(header, allowed_mime_types)
def _is_trusted_domain(hostname: str) -> bool: # ── 对外 API ────────────────────────────────────────────────────────────────
"""检查域名是否在可信白名单中(支持子域名匹配)."""
hostname_lower = hostname.lower()
if hostname_lower in TRUSTED_DOMAINS: def validate_url_safety(url: str, *, purpose: str = "download") -> str:
return True """校验 URL 安全性,返回标准化后的 URL(含 DNS 解析 SSRF 检查).
# 检查子域名
for domain in TRUSTED_DOMAINS: Args:
if hostname_lower.endswith("." + domain.lower()): url: 待校验的 URL
return True purpose: 用途描述(用于日志)
return False
Returns:
标准化后的 URL
Raises:
UrlSecurityError: URL 不安全
"""
# 基础校验(纯逻辑,不含 DNS
_validate_url_basic_base(
url,
trusted_domains=TRUSTED_DOMAINS,
allow_direct_ip=ALLOW_DIRECT_IP,
)
# 如果 hostname 是域名(不是 IP),做 DNS 解析 SSRF 检查
from urllib.parse import urlparse
parsed = urlparse(url)
hostname = parsed.hostname
if hostname and not _is_ip_address_base(hostname):
try:
_check_ssrf_domain(hostname)
except UrlSecurityError:
raise
except Exception as e:
logger.warning("URL 安全校验异常: url=%s purpose=%s error=%s", url[:80], purpose, e)
raise UrlSecurityError(f"URL 安全校验异常: {e}") from e
logger.debug("URL 安全校验通过: url=%s purpose=%s", url[:80], purpose)
return url
def is_url_safe(url: str, *, purpose: str = "download") -> bool: def is_url_safe(url: str, *, purpose: str = "download") -> bool:
@@ -410,7 +159,7 @@ def is_url_safe(url: str, *, purpose: str = "download") -> bool:
return False return False
# ── 安全下载 ─────────────────────────────────────────────────────────────────── # ── 安全下载 ────────────────────────────────────────────────────────────────
def safe_download_file( def safe_download_file(
@@ -426,34 +175,20 @@ def safe_download_file(
包含防护: 包含防护:
- SSRF 校验(初始 URL + 每次重定向后都校验) - SSRF 校验(初始 URL + 每次重定向后都校验)
- 重定向次数限制 + 手动跟随(避免重定向绕过 SSRF - 重定向次数限制 + 手动跟随
- 文件大小限制(流式读取,超过立即中断 - 文件大小限制(流式读取)
- MIME 类型白名单(可选) - MIME 类型白名单 + 魔数二次校验
- 文件头魔数校验(配合 MIME 白名单做二次真实性校验)
Args:
url: 下载 URL
dest_path: 目标文件路径
purpose: 用途描述(日志用)
max_size: 最大下载字节数,超过则中断并抛出 UrlSecurityError
allowed_mime_types: 允许的 Content-Type 集合,None 表示不校验
timeout: 单次请求超时(秒)
Returns: Returns:
实际下载的字节数 实际下载的字节数
Raises:
UrlSecurityError: 安全校验失败
""" """
current_url = url current_url = url
redirect_count = 0 redirect_count = 0
total_bytes = 0 total_bytes = 0
# 使用不自动跟随重定向的 opener
no_redirect_opener = urllib.request.build_opener(NoRedirectHandler()) no_redirect_opener = urllib.request.build_opener(NoRedirectHandler())
while True: while True:
# 每次请求前都做 SSRF 校验(重定向目标也会校验)
validate_url_safety(current_url, purpose=purpose) validate_url_safety(current_url, purpose=purpose)
req = urllib.request.Request(current_url, method="GET") req = urllib.request.Request(current_url, method="GET")
@@ -462,7 +197,6 @@ def safe_download_file(
try: try:
resp = no_redirect_opener.open(req, timeout=timeout) # nosec B310 resp = no_redirect_opener.open(req, timeout=timeout) # nosec B310
except urllib.error.HTTPError as e: except urllib.error.HTTPError as e:
# 3xx 重定向
if 300 <= e.code < 400 and e.headers.get("Location"): if 300 <= e.code < 400 and e.headers.get("Location"):
if redirect_count >= _MAX_REDIRECTS: if redirect_count >= _MAX_REDIRECTS:
raise UrlSecurityError(f"重定向次数超过限制 ({_MAX_REDIRECTS})") from e raise UrlSecurityError(f"重定向次数超过限制 ({_MAX_REDIRECTS})") from e
@@ -474,20 +208,15 @@ def safe_download_file(
raise UrlSecurityError(f"URL 错误: {e.reason}") from e raise UrlSecurityError(f"URL 错误: {e.reason}") from e
try: try:
# Content-Type 校验
if allowed_mime_types is not None: if allowed_mime_types is not None:
content_type = resp.headers.get("Content-Type", "").split(";")[0].strip().lower() content_type = resp.headers.get("Content-Type", "").split(";")[0].strip().lower()
if content_type and content_type not in allowed_mime_types: if content_type and content_type not in allowed_mime_types:
raise UrlSecurityError( raise UrlSecurityError(f"不允许的 Content-Type: {content_type}, 允许: {sorted(allowed_mime_types)}")
f"不允许的 Content-Type: {content_type}, " f"允许: {sorted(allowed_mime_types)}"
)
# Content-Length 预检
content_length = resp.headers.get("Content-Length") content_length = resp.headers.get("Content-Length")
if content_length and int(content_length) > max_size: if content_length and int(content_length) > max_size:
raise UrlSecurityError(f"文件过大: {content_length} bytes > {max_size} bytes 上限") raise UrlSecurityError(f"文件过大: {content_length} bytes > {max_size} bytes 上限")
# 流式下载,实时检查大小
with open(dest_path, "wb") as f: with open(dest_path, "wb") as f:
while True: while True:
chunk = resp.read(_DOWNLOAD_CHUNK_SIZE) chunk = resp.read(_DOWNLOAD_CHUNK_SIZE)
@@ -498,7 +227,6 @@ def safe_download_file(
raise UrlSecurityError(f"下载超过大小限制: {total_bytes} bytes > {max_size} bytes") raise UrlSecurityError(f"下载超过大小限制: {total_bytes} bytes > {max_size} bytes")
f.write(chunk) f.write(chunk)
# 文件头魔数校验(MIME 白名单基础上的二次真实性校验)
if allowed_mime_types is not None: if allowed_mime_types is not None:
_validate_magic_number(dest_path, allowed_mime_types) _validate_magic_number(dest_path, allowed_mime_types)
@@ -515,10 +243,7 @@ def safe_download_bytes(
allowed_mime_types: set[str] | None = None, allowed_mime_types: set[str] | None = None,
timeout: float = 60.0, timeout: float = 60.0,
) -> bytes: ) -> bytes:
"""安全下载 URL 并返回字节内容 """安全下载 URL 并返回字节内容(适合小文件)。"""
防护同 safe_download_file,但结果返回在内存中(适合小文件)。
"""
import tempfile import tempfile
fd, tmp_path = tempfile.mkstemp() fd, tmp_path = tempfile.mkstemp()
+423
View File
@@ -0,0 +1,423 @@
"""URL 安全校验纯逻辑单元测试 — wave128."""
import pytest
from packages.domain.url_security import (
ALLOWED_PORTS,
ALLOWED_SCHEMES,
MAGIC_NUMBERS,
MAX_URL_LENGTH,
UrlSecurityError,
check_internal_hostname,
check_ssrf_ip,
is_ip_address,
is_trusted_domain,
is_url_basic_safe,
validate_magic_number,
validate_url_basic,
)
# ── 常量校验 ────────────────────────────────────────────────────────────────
class TestConstants:
def test_allowed_schemes(self):
assert "http" in ALLOWED_SCHEMES
assert "https" in ALLOWED_SCHEMES
def test_allowed_ports(self):
assert 80 in ALLOWED_PORTS
assert 443 in ALLOWED_PORTS
def test_max_url_length(self):
assert MAX_URL_LENGTH == 2048
def test_magic_numbers_has_common_formats(self):
assert "audio/mpeg" in MAGIC_NUMBERS
assert "image/png" in MAGIC_NUMBERS
assert "video/mp4" in MAGIC_NUMBERS
# ── 内部主机名检查 ──────────────────────────────────────────────────────────
class TestCheckInternalHostname:
@pytest.mark.parametrize(
"hostname",
[
"localhost",
"LOCALHOST",
"LocalHost",
"localhost.localdomain",
"ip6-localhost",
"ip6-loopback",
"metadata",
"metadata.google.internal",
"169.254.169.254",
],
)
def test_internal_hostnames_rejected(self, hostname):
with pytest.raises(UrlSecurityError, match="禁止访问内部主机名"):
check_internal_hostname(hostname)
@pytest.mark.parametrize(
"hostname",
[
"foo.local",
"bar.internal",
"baz.localdomain",
"sub.foo.local",
],
)
def test_internal_domain_suffixes_rejected(self, hostname):
with pytest.raises(UrlSecurityError, match="禁止访问内网域名"):
check_internal_hostname(hostname)
@pytest.mark.parametrize(
"hostname",
[
"example.com",
"www.google.com",
"oss-cn-hangzhou.aliyuncs.com",
"123.45.67.89",
],
)
def test_normal_hostnames_allowed(self, hostname):
check_internal_hostname("example.com") # 不抛异常即通过
# ── 可信域名匹配 ────────────────────────────────────────────────────────────
class TestIsTrustedDomain:
def test_empty_trusted_always_true(self):
assert is_trusted_domain("anything.com", set()) is True
def test_exact_match(self):
trusted = {"example.com", "foo.bar"}
assert is_trusted_domain("example.com", trusted) is True
assert is_trusted_domain("foo.bar", trusted) is True
def test_exact_no_match(self):
trusted = {"example.com"}
assert is_trusted_domain("other.com", trusted) is False
def test_subdomain_match(self):
trusted = {"example.com"}
assert is_trusted_domain("sub.example.com", trusted) is True
assert is_trusted_domain("a.b.example.com", trusted) is True
def test_subdomain_partial_no_match(self):
trusted = {"example.com"}
# fakeexample.com 不是 example.com 的子域名
assert is_trusted_domain("fakeexample.com", trusted) is False
def test_case_insensitive(self):
trusted = {"Example.COM"}
assert is_trusted_domain("example.com", trusted) is True
assert is_trusted_domain("SUB.Example.COM", trusted) is True
# ── IP SSRF 检查 ────────────────────────────────────────────────────────────
class TestCheckSrfIp:
@pytest.mark.parametrize("ip", ["127.0.0.1", "127.1.2.3", "::1"])
def test_loopback_rejected(self, ip):
with pytest.raises(UrlSecurityError, match="回环"):
check_ssrf_ip(ip)
@pytest.mark.parametrize(
"ip",
[
"10.0.0.1",
"10.255.255.255",
"172.16.0.1",
"172.31.255.255",
"192.168.1.1",
"192.168.0.1",
"fd00::1", # IPv6 unique local
],
)
def test_private_rejected(self, ip):
with pytest.raises(UrlSecurityError, match="内网"):
check_ssrf_ip(ip)
@pytest.mark.parametrize("ip", ["169.254.1.1", "169.254.169.254", "fe80::1"])
def test_link_local_rejected(self, ip):
with pytest.raises(UrlSecurityError, match="链路本地"):
check_ssrf_ip(ip)
@pytest.mark.parametrize("ip", ["224.0.0.1", "239.255.255.255", "ff00::1"])
def test_multicast_rejected(self, ip):
with pytest.raises(UrlSecurityError, match="组播"):
check_ssrf_ip(ip)
@pytest.mark.parametrize("ip", ["0.0.0.0", "::"])
def test_unspecified_rejected(self, ip):
with pytest.raises(UrlSecurityError, match="未指定"):
check_ssrf_ip(ip)
def test_reserved_rejected(self):
with pytest.raises(UrlSecurityError):
check_ssrf_ip("240.0.0.1") # 保留地址段
@pytest.mark.parametrize(
"ip",
[
"8.8.8.8",
"1.1.1.1",
"223.5.5.5",
"2001:4860:4860::8888",
],
)
def test_public_ip_allowed(self, ip):
check_ssrf_ip(ip) # 不抛异常即通过
def test_invalid_ip_raises_value_error(self):
with pytest.raises(ValueError):
check_ssrf_ip("not-an-ip")
# ── IP 地址判断 ─────────────────────────────────────────────────────────────
class TestIsIpAddress:
@pytest.mark.parametrize(
"host",
[
"127.0.0.1",
"8.8.8.8",
"192.168.1.1",
"::1",
"2001:db8::1",
"fe80::1",
],
)
def test_ip_addresses(self, host):
assert is_ip_address(host) is True
@pytest.mark.parametrize(
"host",
[
"example.com",
"www.google.com",
"localhost",
"not-an-ip",
"",
],
)
def test_not_ip_addresses(self, host):
assert is_ip_address(host) is False
# ── URL 基础校验 ────────────────────────────────────────────────────────────
class TestValidateUrlBasic:
def test_normal_http_url_passes(self):
result = validate_url_basic("http://example.com/file.txt")
assert result == "http://example.com/file.txt"
def test_normal_https_url_passes(self):
result = validate_url_basic("https://www.example.com/path?q=1")
assert result == "https://www.example.com/path?q=1"
def test_standard_port_80_passes(self):
validate_url_basic("http://example.com:80/file")
def test_standard_port_443_passes(self):
validate_url_basic("https://example.com:443/file")
def test_empty_url_rejected(self):
with pytest.raises(UrlSecurityError, match="URL 为空"):
validate_url_basic("")
def test_none_url_rejected(self):
with pytest.raises(UrlSecurityError, match="URL 为空"):
validate_url_basic(None) # type: ignore
def test_too_long_url_rejected(self):
long_url = "https://example.com/" + "a" * 2100
with pytest.raises(UrlSecurityError, match="URL 过长"):
validate_url_basic(long_url)
@pytest.mark.parametrize(
"url",
[
"ftp://example.com/file",
"file:///etc/passwd",
"javascript:alert(1)",
"data:text/html,<h1>hi</h1>",
],
)
def test_bad_scheme_rejected(self, url):
with pytest.raises(UrlSecurityError, match="不允许的 URL scheme"):
validate_url_basic(url)
def test_missing_hostname_rejected(self):
with pytest.raises(UrlSecurityError, match="URL 缺少主机名"):
validate_url_basic("http:///path")
@pytest.mark.parametrize(
"url",
[
"http://localhost/test",
"http://metadata/test",
"http://foo.local/test",
],
)
def test_internal_hostname_rejected(self, url):
with pytest.raises(UrlSecurityError):
validate_url_basic(url)
@pytest.mark.parametrize(
"url",
[
"http://example.com:8080/file",
"http://example.com:22/file",
"http://example.com:3306/file",
],
)
def test_non_standard_port_rejected(self, url):
with pytest.raises(UrlSecurityError, match="不允许的端口"):
validate_url_basic(url)
def test_direct_ip_rejected_by_default(self):
with pytest.raises(UrlSecurityError, match="禁止直接 IP 访问"):
validate_url_basic("http://8.8.8.8/file")
def test_direct_ip_allowed_with_flag_public(self):
result = validate_url_basic("http://8.8.8.8/file", allow_direct_ip=True)
assert result == "http://8.8.8.8/file"
def test_direct_ip_allowed_flag_but_private_still_rejected(self):
with pytest.raises(UrlSecurityError, match="内网地址"):
validate_url_basic("http://10.0.0.1/file", allow_direct_ip=True)
def test_direct_ip_loopback_rejected(self):
with pytest.raises(UrlSecurityError):
validate_url_basic("http://127.0.0.1/test", allow_direct_ip=True)
def test_trusted_domains_whitelist_pass(self):
trusted = {"example.com"}
result = validate_url_basic("https://example.com/file", trusted_domains=trusted)
assert result == "https://example.com/file"
def test_trusted_domains_subdomain_pass(self):
trusted = {"example.com"}
result = validate_url_basic("https://cdn.example.com/file", trusted_domains=trusted)
assert result == "https://cdn.example.com/file"
def test_trusted_domains_not_in_list_rejected(self):
trusted = {"example.com"}
with pytest.raises(UrlSecurityError, match="不在可信白名单"):
validate_url_basic("https://other.com/file", trusted_domains=trusted)
def test_case_insensitive_scheme(self):
# 大写 HTTP 也应该通过(我们用 .lower() 检查)
result = validate_url_basic("HTTP://example.com/file")
assert "HTTP://example.com/file" == result
# ── is_url_basic_safe 便捷函数 ──────────────────────────────────────────────
class TestIsUrlBasicSafe:
def test_safe_url_returns_true(self):
assert is_url_basic_safe("https://example.com/file") is True
def test_unsafe_url_returns_false(self):
assert is_url_basic_safe("http://localhost/test") is False
def test_empty_returns_false(self):
assert is_url_basic_safe("") is False
def test_with_trusted_domains(self):
trusted = {"allowed.com"}
assert is_url_basic_safe("https://allowed.com/x", trusted_domains=trusted) is True
assert is_url_basic_safe("https://other.com/x", trusted_domains=trusted) is False
# ── 魔数校验 ────────────────────────────────────────────────────────────────
class TestValidateMagicNumber:
def test_mp3_id3_header(self):
data = b"ID3" + b"\x00" * 100
validate_magic_number(data, {"audio/mpeg"}) # 不抛异常
def test_mp3_frame_sync(self):
data = b"\xff\xfb" + b"\x00" * 100
validate_magic_number(data, {"audio/mpeg"})
def test_wav_header(self):
data = b"RIFF" + b"\x00" * 4 + b"WAVE" + b"\x00" * 100
validate_magic_number(data, {"audio/wav"})
def test_png_header(self):
data = b"\x89PNG\r\n\x1a\n" + b"\x00" * 100
validate_magic_number(data, {"image/png"})
def test_jpeg_header(self):
data = b"\xff\xd8\xff" + b"\x00" * 100
validate_magic_number(data, {"image/jpeg"})
def test_gif87a_header(self):
data = b"GIF87a" + b"\x00" * 100
validate_magic_number(data, {"image/gif"})
def test_gif89a_header(self):
data = b"GIF89a" + b"\x00" * 100
validate_magic_number(data, {"image/gif"})
def test_mp4_ftyp_header(self):
data = b"\x00\x00\x00\x20ftypisom" + b"\x00" * 100
validate_magic_number(data, {"video/mp4"})
def test_ogg_header(self):
data = b"OggS" + b"\x00" * 100
validate_magic_number(data, {"audio/ogg"})
def test_flac_header(self):
data = b"fLaC" + b"\x00" * 100
validate_magic_number(data, {"audio/flac"})
def test_webp_header(self):
data = b"RIFF" + b"\x00" * 4 + b"WEBP" + b"\x00" * 100
validate_magic_number(data, {"image/webp"})
def test_bmp_header(self):
data = b"BM" + b"\x00" * 100
validate_magic_number(data, {"image/bmp"})
def test_empty_file_rejected(self):
with pytest.raises(UrlSecurityError, match="文件为空"):
validate_magic_number(b"", {"image/png"})
def test_mismatched_magic_rejected(self):
data = b"NOTAPNG" + b"\x00" * 100
with pytest.raises(UrlSecurityError, match="魔数与允许的 MIME 类型不匹配"):
validate_magic_number(data, {"image/png"})
def test_multiple_allowed_types_one_matches(self):
data = b"\x89PNG\r\n\x1a\n" + b"\x00" * 100
# 多个允许类型,只要有一个匹配就通过
validate_magic_number(data, {"image/png", "image/jpeg", "image/gif"})
def test_no_known_magic_skips_validation(self):
# 自定义 MIME 类型没有已知魔数,跳过校验不阻断
validate_magic_number(b"random data", {"application/x-custom"})
def test_too_short_header_no_match(self):
# 文件头太短,无法匹配需要 8 字节偏移的格式
data = b"RIFF" # 只有 4 字节,不够 offset 8 的 WAVE 匹配
with pytest.raises(UrlSecurityError, match="魔数"):
validate_magic_number(data, {"audio/wav"})
def test_error_message_contains_mime_and_header(self):
with pytest.raises(UrlSecurityError) as exc_info:
validate_magic_number(b"XXXXYYY", {"image/png"})
msg = str(exc_info.value)
assert "image/png" in msg
assert "文件头前16字节" in msg