fix: 渲染引擎安全加固P0+P1 补修第二轮
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 16s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 1m11s
CI/CD Pipeline / Production Browser E2E (pull_request) Failing after 1535h8m25s
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 1535h8m27s
CI/CD Pipeline / Deploy Production (pull_request) Failing after 1535h8m27s
CI/CD Pipeline / Build Production Runtime Images (pull_request) Failing after 1535h8m28s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Failing after 1535h8m28s
CI/CD Pipeline / Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 1535h40m5s

- SSRF漏接点补全:TTS 3个下载点 + BGM 3个下载点 全部接入 url_security.py 校验
  - TTS: workflow.py (_transfer_audio_to_oss, _download_and_merge_segments) + streaming_service.py (_download_audio)
  - BGM: generation.py (外部直链, 预设库) + batch_download.py (HTTP回退)
- 重定向防护:手动跟随重定向,每次跳转前重新校验目标 URL(禁用默认自动跟随)
- 文件大小/类型限制:新增 safe_download_file/safe_download_bytes,流式下载 + 大小上限 + MIME 白名单
- 裸 subprocess 补齐:5 处全部改走统一 run_ffmpeg/run_ffprobe
  - asset_analyzer.py: 3处 (ffprobe + 2个ffmpeg)
  - ingest.py: 1处 (ffprobe)
  - voice_extraction.py: 1处 (ffmpeg)
- URL安全模块迁移到 packages/shared/ 作为单一来源,worker 端保留向后兼容 re-export
- 新增 run_ffprobe 统一工具函数到 ffmpeg_utils
- 新增 7 个下载安全单测,累计 33 个 URL 安全测试
This commit is contained in:
CI Bot
2026-07-14 15:34:29 +08:00
parent 3b80edd8c5
commit 334e2b1fc2
11 changed files with 712 additions and 279 deletions
@@ -119,6 +119,55 @@ def run_ffmpeg(
raise
def run_ffprobe(
command: list[str],
*,
capture_output: bool = True,
timeout: int = 30,
) -> tuple[str, str]:
"""执行 FFprobe 命令。
Args:
command: 完整的 ffprobe 命令列表(含 "ffprobe" 本身)
capture_output: 是否捕获 stdout/stderr
timeout: 超时时间(秒),默认 30s;None 表示不设超时
Returns:
(stdout, stderr) 元组
Raises:
subprocess.CalledProcessError: 命令执行失败时抛出
subprocess.TimeoutExpired: 超时未完成时抛出
"""
try:
result = subprocess.run( # nosec B603
command,
check=True,
stdout=subprocess.PIPE if capture_output else None,
stderr=subprocess.PIPE if capture_output else None,
text=True,
timeout=timeout,
)
return (result.stdout or "", result.stderr or "")
except subprocess.TimeoutExpired:
logger.error(
"FFprobe 命令超时 (%ds): command=%s",
timeout or -1,
" ".join(str(c) for c in command[:20]),
)
raise
except subprocess.CalledProcessError as e:
stderr_text = (e.stderr or "").strip()
logger.error(
"FFprobe 命令失败: exit_code=%d command=%s\nstderr:\n%s",
e.returncode,
" ".join(str(c) for c in command[:20]),
stderr_text[:5000],
)
raise
def probe_has_audio(local_path: str | Path) -> bool:
"""探测文件是否包含音频流。
+18 -219
View File
@@ -1,222 +1,21 @@
"""URL 安全校验工具 — SSRF 防护.
"""URL 安全校验工具 — SSRF 防护(向后兼容层).
统一的外部 URL 安全校验方案,覆盖所有渲染管线中的外部下载场景:
- BGM 音频 URL 下载
- TTS 音频结果 URL 下载
- PiP 画中画素材 URL 下载
- 贴纸/封面等图片 URL 下载
- 通用 URL 可访问性校验
防护要点:
1. Scheme 白名单:仅允许 http/https
2. 主机 SSRF 防护:禁止内网 IP、回环地址、链路本地地址、元数据服务
3. 端口白名单:仅允许 80/443(标准 HTTP/HTTPS)
4. 域名校验:禁止 IP 直接访问(除非在白名单中)
5. 重定向防护:下载时禁止跳转到内网
本模块为向后兼容而保留,实际实现已迁移至 packages.shared.url_security。
所有符号均从该模块重新导出,请新代码直接 import packages.shared.url_security。
"""
from __future__ import annotations
import ipaddress
import logging
import socket
from urllib.parse import urlparse
logger = logging.getLogger(__name__)
# 允许的 URL scheme
ALLOWED_SCHEMES = {"http", "https"}
# 允许的端口(标准 HTTP/HTTPS)
ALLOWED_PORTS = {80, 443}
# 可信域名白名单(可根据实际 OSS/CDN 域名配置)
# 从环境变量读取,格式:"oss-cn-hangzhou.aliyuncs.com,cdn.example.com"
# 默认空表示所有公网域名都允许,但仍会做 SSRF 检查
import os
TRUSTED_DOMAINS: set[str] = set()
_env_trusted = os.environ.get("URL_SECURITY_TRUSTED_DOMAINS", "")
if _env_trusted:
TRUSTED_DOMAINS = {d.strip() for d in _env_trusted.split(",") if d.strip()}
# 是否允许 IP 直接访问(默认禁止,防止绕过 DNS 校验)
ALLOW_DIRECT_IP = os.environ.get("URL_SECURITY_ALLOW_DIRECT_IP", "false").lower() == "true"
# 最大 URL 长度
MAX_URL_LENGTH = 2048
class UrlSecurityError(ValueError):
"""URL 安全校验失败."""
pass
def validate_url_safety(url: str, *, purpose: str = "download") -> str:
"""校验 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:
"""对域名做 DNS 解析并检查所有解析结果的 IP 是否安全.
注意:这不能完全防止 DNS rebinding,但能防御大部分 SSRF 场景。
"""
try:
# 解析所有地址
infos = socket.getaddrinfo(hostname, None)
if not infos:
raise UrlSecurityError(f"域名解析失败: {hostname}")
for info in infos:
ip_str = info[4][0]
try:
ip_obj = ipaddress.ip_address(ip_str)
_check_ssrf_ip(ip_obj)
except ValueError:
# 无法解析为 IP,跳过(不应该发生)
continue
except socket.gaierror as e:
raise UrlSecurityError(f"域名解析失败: {hostname} ({e})") from e
def _is_trusted_ip(ip_obj: ipaddress.IPv4Address | ipaddress.IPv6Address) -> bool:
"""检查 IP 是否在可信列表中(目前通过环境变量配置域名,IP 级信任暂不开放)."""
return False
def _is_trusted_domain(hostname: str) -> bool:
"""检查域名是否在可信白名单中(支持子域名匹配)."""
hostname_lower = hostname.lower()
if hostname_lower in TRUSTED_DOMAINS:
return True
# 检查子域名
for domain in TRUSTED_DOMAINS:
if hostname_lower.endswith("." + domain.lower()):
return True
return False
def is_url_safe(url: str, *, purpose: str = "download") -> bool:
"""便捷函数:检查 URL 是否安全,不抛异常."""
try:
validate_url_safety(url, purpose=purpose)
return True
except UrlSecurityError:
return False
from packages.shared.url_security import ( # noqa: F401
ALLOWED_AUDIO_MIME_TYPES,
ALLOWED_IMAGE_MIME_TYPES,
ALLOWED_PORTS,
ALLOWED_SCHEMES,
ALLOWED_VIDEO_MIME_TYPES,
DEFAULT_MAX_DOWNLOAD_SIZE,
MAX_URL_LENGTH,
TRUSTED_DOMAINS,
UrlSecurityError,
is_url_safe,
safe_download_bytes,
safe_download_file,
validate_url_safety,
)
+23 -23
View File
@@ -130,6 +130,8 @@ class AssetAnalyzer:
info = VideoInfo()
try:
from video_processing.ffmpeg_utils import run_ffprobe
cmd = [
"ffprobe",
"-v",
@@ -140,15 +142,8 @@ class AssetAnalyzer:
"-show_streams",
self.video_path,
]
result = subprocess.run(
cmd,
capture_output=True,
text=True,
timeout=30,
)
if result.returncode == 0:
data = json.loads(result.stdout)
stdout, _ = run_ffprobe(cmd, timeout=30)
data = json.loads(stdout)
streams = data.get("streams", [])
format_info = data.get("format", {})
@@ -224,14 +219,14 @@ class AssetAnalyzer:
output_path,
]
result = subprocess.run(
cmd,
capture_output=True,
text=True,
timeout=10,
)
from video_processing.ffmpeg_utils import run_ffmpeg
if result.returncode == 0 and os.path.exists(output_path):
try:
run_ffmpeg(cmd, timeout=10)
except Exception:
continue
if os.path.exists(output_path):
# 读取帧并转换为 numpy 数组
img = self._load_image_as_array(output_path)
if img is not None:
@@ -397,14 +392,19 @@ class AssetAnalyzer:
audio_path,
]
result_audio = subprocess.run(
cmd,
capture_output=True,
text=True,
timeout=30,
)
from video_processing.ffmpeg_utils import run_ffmpeg
if result_audio.returncode == 0 and os.path.exists(audio_path):
try:
run_ffmpeg(cmd, timeout=30)
except Exception:
# 音频提取失败,返回默认分析结果
return AudioAnalysis(
has_speech=False,
speech_ratio=0.0,
avg_volume=0.0,
)
if os.path.exists(audio_path):
# 读取音频数据
import struct
@@ -106,7 +106,12 @@ def _download_video_to_file(url: str, dest_path: str) -> None:
except Exception:
pass
# 回退到 HTTP 下载
import urllib.request
# 回退到 HTTP 下载(含 SSRF 防护 + 大小限制)
from video_processing.url_security import safe_download_file
urllib.request.urlretrieve(url, dest_path) # nosec B310
safe_download_file(
url,
dest_path,
purpose="batch_video_download",
timeout=300.0,
)
+22 -4
View File
@@ -364,10 +364,19 @@ def _prepare_bgm_track(
try:
parsed = urlparse(audio_url)
if parsed.scheme in ("http", "https"):
import urllib.request
from video_processing.url_security import (
ALLOWED_AUDIO_MIME_TYPES,
safe_download_file,
)
logger.info("[task_id=%s] [BGM] 从URL下载: %s", task_id, audio_url[:80])
urllib.request.urlretrieve(audio_url, bgm_file) # nosec B310
safe_download_file(
audio_url,
str(bgm_file),
purpose="bgm_download",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
timeout=60.0,
)
if bgm_file.exists() and bgm_file.stat().st_size > 0:
return str(bgm_file)
except Exception as e:
@@ -401,10 +410,19 @@ def _prepare_bgm_track(
preset = get_preset_bgm(preset_id)
if preset and preset.audio_url:
import urllib.request
from video_processing.url_security import (
ALLOWED_AUDIO_MIME_TYPES,
safe_download_file,
)
logger.info("[task_id=%s] [BGM] 从预设库下载: preset_id=%s", task_id, preset_id)
urllib.request.urlretrieve(preset.audio_url, bgm_file) # nosec B310
safe_download_file(
preset.audio_url,
str(bgm_file),
purpose="bgm_preset_download",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
timeout=60.0,
)
if bgm_file.exists() and bgm_file.stat().st_size > 0:
return str(bgm_file)
except Exception as e:
+8 -8
View File
@@ -45,6 +45,8 @@ def extract_media_metadata(file_url: str, media_type: str) -> dict:
try:
if media_type == "video":
# 使用 ffprobe 提取视频元数据
from video_processing.ffmpeg_utils import run_ffprobe
cmd = [
"ffprobe",
"-v",
@@ -55,16 +57,11 @@ def extract_media_metadata(file_url: str, media_type: str) -> dict:
"-show_streams",
file_url,
]
result = subprocess.run(
cmd,
capture_output=True,
text=True,
timeout=30,
)
if result.returncode == 0:
try:
stdout, _ = run_ffprobe(cmd, timeout=30)
import json as json_lib
probe_data = json_lib.loads(result.stdout)
probe_data = json_lib.loads(stdout)
# 提取视频流信息
for stream in probe_data.get("streams", []):
@@ -83,6 +80,9 @@ def extract_media_metadata(file_url: str, media_type: str) -> dict:
metadata["size_bytes"] = int(format_info.get("size", 0))
metadata["bitrate"] = int(format_info.get("bit_rate", 0))
except Exception as e:
logger.warning("视频元数据提取失败: %s", e)
elif media_type == "image":
# 使用 Pillow 提取图片元数据
try:
@@ -19,14 +19,12 @@ class VoiceExtractor:
"""Extract voice tracks and background music from videos using FFmpeg."""
@staticmethod
def _run_ffmpeg(cmd: list[str]) -> subprocess.CompletedProcess:
"""Run FFmpeg command and return result."""
logger.info(f"Running FFmpeg: {chr(39).join(cmd)}")
result = subprocess.run(cmd, capture_output=True, text=True)
if result.returncode != 0:
logger.error(f"FFmpeg error: {result.stderr}")
raise RuntimeError(f"FFmpeg failed: {result.stderr}")
return result
def _run_ffmpeg(cmd: list[str]) -> None:
"""Run FFmpeg command using 统一 run_ffmpeg 工具."""
from video_processing.ffmpeg_utils import run_ffmpeg
logger.info("Running FFmpeg: %s", " ".join(cmd[:10]))
run_ffmpeg(cmd)
def extract_voice(
self,
@@ -13,6 +13,11 @@ from typing import Any, Optional
import httpx
from packages.shared.url_security import (
ALLOWED_AUDIO_MIME_TYPES,
safe_download_bytes,
)
from packages.application.cosyvoice_service import CosyVoiceError, CosyVoiceService
from packages.application.tts_job.text_splitter import split_text
@@ -224,10 +229,13 @@ class TTSStreamingService:
# ── 工具方法 ────────────────────────────────────────────
def _download_audio(self, url: str) -> bytes:
"""下载音频数据。"""
resp = httpx.get(url, timeout=60.0, follow_redirects=True)
resp.raise_for_status()
return resp.content
"""下载音频数据(含 SSRF 防护 + 大小限制 + 重定向校验)。"""
return safe_download_bytes(
url,
purpose="tts_streaming_download",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
timeout=60.0,
)
async def _stream_audio_chunks(self, websocket: Any, audio_data: bytes) -> int:
"""将音频数据分块通过 WebSocket 推送。
+23 -10
View File
@@ -19,6 +19,14 @@ from typing import Optional
import httpx
from packages.shared.url_security import (
ALLOWED_AUDIO_MIME_TYPES,
UrlSecurityError,
safe_download_bytes,
safe_download_file,
validate_url_safety,
)
from packages.application.cosyvoice_service import (
CosyVoiceAuthError,
CosyVoiceError,
@@ -96,10 +104,13 @@ class TTSWorkflowService:
content_type = content_type_map.get(audio_format, "application/octet-stream")
try:
# 下载临时音频
resp = httpx.get(temp_url, timeout=60.0, follow_redirects=True)
resp.raise_for_status()
audio_data = resp.content
# 安全下载临时音频(SSRF 防护 + 大小限制 + 重定向校验)
audio_data = safe_download_bytes(
temp_url,
purpose="tts_audio_download",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
timeout=60.0,
)
# 上传到 OSS
file_obj = io.BytesIO(audio_data)
@@ -463,13 +474,15 @@ class TTSWorkflowService:
total_duration += result.get("duration", 0.0)
# 下载分段音频到临时文件
resp = httpx.get(audio_url, timeout=60.0, follow_redirects=True)
resp.raise_for_status()
# 安全下载分段音频到临时文件(SSRF 防护 + 大小限制)
seg_path = os.path.join(temp_dir, f"seg_{idx:03d}.{job.format}")
with open(seg_path, "wb") as f:
f.write(resp.content)
safe_download_file(
audio_url,
seg_path,
purpose="tts_segment_download",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
timeout=60.0,
)
audio_paths.append(seg_path)
# 合并
+413
View File
@@ -0,0 +1,413 @@
"""URL 安全校验工具 — SSRF 防护.
统一的外部 URL 安全校验方案,覆盖所有渲染管线和 TTS 中的外部下载场景。
放在 packages/shared/ 作为单一来源,worker 和 application 层都可引用。
防护要点:
1. Scheme 白名单:仅允许 http/https
2. 主机 SSRF 防护:禁止内网 IP、回环地址、链路本地地址、元数据服务
3. 端口白名单:仅允许 80/443(标准 HTTP/HTTPS)
4. 域名校验:禁止 IP 直接访问(除非在白名单中)
5. 重定向防护:手动跟随重定向,每次跳转前重新校验目标 URL
6. 文件大小限制:流式下载,超过上限立即中断
7. MIME 类型白名单:可选的内容类型校验
"""
from __future__ import annotations
import ipaddress
import logging
import os
import socket
import urllib.error
import urllib.request
from urllib.parse import urljoin, urlparse
logger = logging.getLogger(__name__)
# 允许的 URL scheme
ALLOWED_SCHEMES = {"http", "https"}
# 允许的端口(标准 HTTP/HTTPS)
ALLOWED_PORTS = {80, 443}
# 可信域名白名单(可根据实际 OSS/CDN 域名配置)
# 从环境变量读取,格式:"oss-cn-hangzhou.aliyuncs.com,cdn.example.com"
# 默认空表示所有公网域名都允许,但仍会做 SSRF 检查
TRUSTED_DOMAINS: set[str] = set()
_env_trusted = os.environ.get("URL_SECURITY_TRUSTED_DOMAINS", "")
if _env_trusted:
TRUSTED_DOMAINS = {d.strip() for d in _env_trusted.split(",") if d.strip()}
# 是否允许 IP 直接访问(默认禁止,防止绕过 DNS 校验)
ALLOW_DIRECT_IP = os.environ.get("URL_SECURITY_ALLOW_DIRECT_IP", "false").lower() == "true"
# 最大 URL 长度
MAX_URL_LENGTH = 2048
# 单次下载最大文件大小(默认 200MB)
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
# 最大重定向次数
_MAX_REDIRECTS = 5
class UrlSecurityError(ValueError):
"""URL 安全校验失败."""
pass
class NoRedirectHandler(urllib.request.HTTPRedirectHandler):
"""禁止自动重定向的 handler,用于手动控制重定向以做安全校验."""
def redirect_request(self, req, fp, code, msg, headers, newurl): # noqa: N802
return None
def validate_url_safety(url: str, *, purpose: str = "download") -> str:
"""校验 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:
"""对域名做 DNS 解析并检查所有解析结果的 IP 是否安全.
注意:这不能完全防止 DNS rebinding,但能防御大部分 SSRF 场景。
"""
try:
# 解析所有地址
infos = socket.getaddrinfo(hostname, None)
if not infos:
raise UrlSecurityError(f"域名解析失败: {hostname}")
for info in infos:
ip_str = info[4][0]
try:
ip_obj = ipaddress.ip_address(ip_str)
_check_ssrf_ip(ip_obj)
except ValueError:
# 无法解析为 IP,跳过(不应该发生)
continue
except socket.gaierror as e:
raise UrlSecurityError(f"域名解析失败: {hostname} ({e})") from e
def _is_trusted_ip(ip_obj: ipaddress.IPv4Address | ipaddress.IPv6Address) -> bool:
"""检查 IP 是否在可信列表中(目前通过环境变量配置域名,IP 级信任暂不开放)."""
return False
def _is_trusted_domain(hostname: str) -> bool:
"""检查域名是否在可信白名单中(支持子域名匹配)."""
hostname_lower = hostname.lower()
if hostname_lower in TRUSTED_DOMAINS:
return True
# 检查子域名
for domain in TRUSTED_DOMAINS:
if hostname_lower.endswith("." + domain.lower()):
return True
return False
def is_url_safe(url: str, *, purpose: str = "download") -> bool:
"""便捷函数:检查 URL 是否安全,不抛异常."""
try:
validate_url_safety(url, purpose=purpose)
return True
except UrlSecurityError:
return False
# ── 安全下载 ───────────────────────────────────────────────────────────────────
def safe_download_file(
url: str,
dest_path: str,
*,
purpose: str = "download",
max_size: int = DEFAULT_MAX_DOWNLOAD_SIZE,
allowed_mime_types: set[str] | None = None,
timeout: float = 60.0,
) -> int:
"""安全下载 URL 到本地文件。
包含防护:
- SSRF 校验(初始 URL + 每次重定向后都校验)
- 重定向次数限制 + 手动跟随(避免重定向绕过 SSRF)
- 文件大小限制(流式读取,超过立即中断)
- MIME 类型白名单(可选)
Args:
url: 下载 URL
dest_path: 目标文件路径
purpose: 用途描述(日志用)
max_size: 最大下载字节数,超过则中断并抛出 UrlSecurityError
allowed_mime_types: 允许的 Content-Type 集合,None 表示不校验
timeout: 单次请求超时(秒)
Returns:
实际下载的字节数
Raises:
UrlSecurityError: 安全校验失败
"""
current_url = url
redirect_count = 0
total_bytes = 0
# 使用不自动跟随重定向的 opener
no_redirect_opener = urllib.request.build_opener(NoRedirectHandler())
while True:
# 每次请求前都做 SSRF 校验(重定向目标也会校验)
validate_url_safety(current_url, purpose=purpose)
req = urllib.request.Request(current_url, method="GET")
req.add_header("User-Agent", "xiaoxia-saas-worker/1.0")
try:
resp = no_redirect_opener.open(req, timeout=timeout) # nosec B310
except urllib.error.HTTPError as e:
# 3xx 重定向
if 300 <= e.code < 400 and e.headers.get("Location"):
if redirect_count >= _MAX_REDIRECTS:
raise UrlSecurityError(
f"重定向次数超过限制 ({_MAX_REDIRECTS})"
) from e
redirect_count += 1
current_url = urljoin(current_url, e.headers["Location"])
continue
raise UrlSecurityError(f"HTTP 错误: {e.code} {e.reason}") from e
except urllib.error.URLError as e:
raise UrlSecurityError(f"URL 错误: {e.reason}") from e
try:
# Content-Type 校验
if allowed_mime_types is not None:
content_type = (
resp.headers.get("Content-Type", "").split(";")[0].strip().lower()
)
if content_type and content_type not in allowed_mime_types:
raise UrlSecurityError(
f"不允许的 Content-Type: {content_type}, "
f"允许: {sorted(allowed_mime_types)}"
)
# Content-Length 预检
content_length = resp.headers.get("Content-Length")
if content_length and int(content_length) > max_size:
raise UrlSecurityError(
f"文件过大: {content_length} bytes > {max_size} bytes 上限"
)
# 流式下载,实时检查大小
with open(dest_path, "wb") as f:
while True:
chunk = resp.read(_DOWNLOAD_CHUNK_SIZE)
if not chunk:
break
total_bytes += len(chunk)
if total_bytes > max_size:
raise UrlSecurityError(
f"下载超过大小限制: {total_bytes} bytes > {max_size} bytes"
)
f.write(chunk)
return total_bytes
finally:
resp.close()
def safe_download_bytes(
url: str,
*,
purpose: str = "download",
max_size: int = DEFAULT_MAX_DOWNLOAD_SIZE,
allowed_mime_types: set[str] | None = None,
timeout: float = 60.0,
) -> bytes:
"""安全下载 URL 并返回字节内容。
防护同 safe_download_file,但结果返回在内存中(适合小文件)。
"""
import tempfile
fd, tmp_path = tempfile.mkstemp()
os.close(fd)
try:
safe_download_file(
url,
tmp_path,
purpose=purpose,
max_size=max_size,
allowed_mime_types=allowed_mime_types,
timeout=timeout,
)
with open(tmp_path, "rb") as f:
return f.read()
finally:
try:
os.unlink(tmp_path)
except OSError:
pass
+130
View File
@@ -8,9 +8,15 @@ import unittest
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "worker"))
import shutil
import tempfile
from video_processing.url_security import ( # noqa: E402
ALLOWED_AUDIO_MIME_TYPES,
UrlSecurityError,
is_url_safe,
safe_download_bytes,
safe_download_file,
validate_url_safety,
)
@@ -162,3 +168,127 @@ class TestUrlSecurityValidation(unittest.TestCase):
if __name__ == "__main__":
unittest.main()
class TestSafeDownload(unittest.TestCase):
"""安全下载函数测试."""
def setUp(self):
self.temp_dir = tempfile.mkdtemp()
def tearDown(self):
shutil.rmtree(self.temp_dir, ignore_errors=True)
def test_safe_download_file_rejects_ssrf(self):
"""SSRF 风险 URL 应该被拒绝下载."""
dest = os.path.join(self.temp_dir, "test.bin")
with self.assertRaises(UrlSecurityError):
safe_download_file("http://127.0.0.1/test", dest, purpose="test")
def test_safe_download_bytes_rejects_ssrf(self):
"""SSRF 风险 URL 应该被拒绝下载(bytes 版本)."""
with self.assertRaises(UrlSecurityError):
safe_download_bytes("http://localhost/test", purpose="test")
def test_safe_download_file_size_limit(self):
"""超过大小限制应该被拒绝."""
# 用 mock server 测试太大的 content-length
dest = os.path.join(self.temp_dir, "test.bin")
# 直接验证参数:max_size=0 时任何下载都应超限
# (这里用一个可访问的 URL 并设置极小的限制)
# 为避免依赖外部网络,这里只测试函数参数传递
with unittest.mock.patch("urllib.request.build_opener") as mock_opener:
mock_resp = unittest.mock.MagicMock()
mock_resp.headers = {"Content-Length": "1000"}
mock_resp.read.return_value = b""
mock_opener.return_value.open.return_value = mock_resp
# 设置 max_size=500,content-length=1000 应被拒绝
with self.assertRaises(UrlSecurityError):
safe_download_file(
"https://example.com/test",
dest,
purpose="test",
max_size=500,
)
def test_safe_download_file_mime_rejected(self):
"""不允许的 MIME 类型应该被拒绝."""
dest = os.path.join(self.temp_dir, "test.bin")
with unittest.mock.patch("urllib.request.build_opener") as mock_opener:
mock_resp = unittest.mock.MagicMock()
mock_resp.headers = {"Content-Type": "text/html"}
mock_resp.read.return_value = b""
mock_opener.return_value.open.return_value = mock_resp
with self.assertRaises(UrlSecurityError):
safe_download_file(
"https://example.com/test.mp3",
dest,
purpose="test",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
)
def test_safe_download_file_mime_allowed(self):
"""允许的 MIME 类型应该通过."""
dest = os.path.join(self.temp_dir, "test.mp3")
with unittest.mock.patch("urllib.request.build_opener") as mock_opener:
mock_resp = unittest.mock.MagicMock()
mock_resp.headers = {"Content-Type": "audio/mpeg"}
mock_resp.read.side_effect = [b"audio_data", b""]
mock_resp.geturl.return_value = "https://example.com/test.mp3"
mock_opener.return_value.open.return_value = mock_resp
size = safe_download_file(
"https://example.com/test.mp3",
dest,
purpose="test",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
)
self.assertEqual(size, 10)
self.assertTrue(os.path.exists(dest))
def test_safe_download_file_stream_size_limit(self):
"""流式下载时超过大小限制应该中断."""
dest = os.path.join(self.temp_dir, "test.bin")
with unittest.mock.patch("urllib.request.build_opener") as mock_opener:
mock_resp = unittest.mock.MagicMock()
mock_resp.headers = {}
# 每次返回 100 字节,max_size=500,第 6 次读取就超限
mock_resp.read.side_effect = lambda n: b"x" * n if n < 1000 else b"x" * 100
# 改成返回固定 100 字节,直到第 N 次后返回空
call_count = [0]
def mock_read(size):
call_count[0] += 1
if call_count[0] > 10:
return b""
return b"x" * 100
mock_resp.read = mock_read
mock_opener.return_value.open.return_value = mock_resp
with self.assertRaises(UrlSecurityError):
safe_download_file(
"https://example.com/test",
dest,
purpose="test",
max_size=500, # 500 字节上限
)
def test_safe_download_bytes_returns_content(self):
"""safe_download_bytes 应该返回文件内容."""
test_data = b"hello world test audio"
with unittest.mock.patch("urllib.request.build_opener") as mock_opener:
mock_resp = unittest.mock.MagicMock()
mock_resp.headers = {"Content-Type": "audio/mpeg"}
call_count = [0]
def mock_read(size):
call_count[0] += 1
if call_count[0] > 1:
return b""
return test_data
mock_resp.read = mock_read
mock_opener.return_value.open.return_value = mock_resp
result = safe_download_bytes(
"https://example.com/test.mp3",
purpose="test",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
)
self.assertEqual(result, test_data)