From 334e2b1fc2c7da482eb986cc30b5418dabde9494 Mon Sep 17 00:00:00 2001 From: CI Bot Date: Tue, 14 Jul 2026 15:34:29 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E6=B8=B2=E6=9F=93=E5=BC=95=E6=93=8E?= =?UTF-8?q?=E5=AE=89=E5=85=A8=E5=8A=A0=E5=9B=BAP0+P1=20=E8=A1=A5=E4=BF=AE?= =?UTF-8?q?=E7=AC=AC=E4=BA=8C=E8=BD=AE?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 安全测试 --- apps/worker/video_processing/ffmpeg_utils.py | 49 +++ apps/worker/video_processing/url_security.py | 237 +--------- .../worker/worker_app/tasks/asset_analyzer.py | 46 +- .../worker/worker_app/tasks/batch_download.py | 11 +- apps/worker/worker_app/tasks/generation.py | 26 +- apps/worker/worker_app/tasks/ingest.py | 16 +- .../worker_app/tasks/voice_extraction.py | 14 +- .../application/tts_job/streaming_service.py | 16 +- packages/application/tts_job/workflow.py | 33 +- packages/shared/url_security.py | 413 ++++++++++++++++++ tests/unit/test_url_security.py | 130 ++++++ 11 files changed, 712 insertions(+), 279 deletions(-) mode change 100755 => 100644 apps/worker/video_processing/url_security.py mode change 100755 => 100644 packages/application/tts_job/workflow.py create mode 100644 packages/shared/url_security.py diff --git a/apps/worker/video_processing/ffmpeg_utils.py b/apps/worker/video_processing/ffmpeg_utils.py index d40cc2d1f..2e61f9075 100755 --- a/apps/worker/video_processing/ffmpeg_utils.py +++ b/apps/worker/video_processing/ffmpeg_utils.py @@ -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: """探测文件是否包含音频流。 diff --git a/apps/worker/video_processing/url_security.py b/apps/worker/video_processing/url_security.py old mode 100755 new mode 100644 index 0afa46978..3ea89d178 --- a/apps/worker/video_processing/url_security.py +++ b/apps/worker/video_processing/url_security.py @@ -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, +) diff --git a/apps/worker/worker_app/tasks/asset_analyzer.py b/apps/worker/worker_app/tasks/asset_analyzer.py index c38a26304..424799841 100755 --- a/apps/worker/worker_app/tasks/asset_analyzer.py +++ b/apps/worker/worker_app/tasks/asset_analyzer.py @@ -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 diff --git a/apps/worker/worker_app/tasks/batch_download.py b/apps/worker/worker_app/tasks/batch_download.py index 3924f8d8b..0b1e5606b 100755 --- a/apps/worker/worker_app/tasks/batch_download.py +++ b/apps/worker/worker_app/tasks/batch_download.py @@ -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, + ) diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 433fa07a5..0bca4f66e 100755 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -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: diff --git a/apps/worker/worker_app/tasks/ingest.py b/apps/worker/worker_app/tasks/ingest.py index 0261c6e8c..e701198b9 100755 --- a/apps/worker/worker_app/tasks/ingest.py +++ b/apps/worker/worker_app/tasks/ingest.py @@ -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: diff --git a/apps/worker/worker_app/tasks/voice_extraction.py b/apps/worker/worker_app/tasks/voice_extraction.py index b08335809..efbc652a0 100644 --- a/apps/worker/worker_app/tasks/voice_extraction.py +++ b/apps/worker/worker_app/tasks/voice_extraction.py @@ -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, diff --git a/packages/application/tts_job/streaming_service.py b/packages/application/tts_job/streaming_service.py index ffb4674f3..97e048472 100644 --- a/packages/application/tts_job/streaming_service.py +++ b/packages/application/tts_job/streaming_service.py @@ -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 推送。 diff --git a/packages/application/tts_job/workflow.py b/packages/application/tts_job/workflow.py old mode 100755 new mode 100644 index 4b9a472cc..5d374e05b --- a/packages/application/tts_job/workflow.py +++ b/packages/application/tts_job/workflow.py @@ -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) # 合并 diff --git a/packages/shared/url_security.py b/packages/shared/url_security.py new file mode 100644 index 000000000..c96a43e3f --- /dev/null +++ b/packages/shared/url_security.py @@ -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 diff --git a/tests/unit/test_url_security.py b/tests/unit/test_url_security.py index fa76105b3..4c5d739a7 100755 --- a/tests/unit/test_url_security.py +++ b/tests/unit/test_url_security.py @@ -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) + +