diff --git a/apps/worker/video_processing/oss_helpers.py b/apps/worker/video_processing/oss_helpers.py index 78944886d..3e486ed51 100755 --- a/apps/worker/video_processing/oss_helpers.py +++ b/apps/worker/video_processing/oss_helpers.py @@ -220,25 +220,64 @@ def resolve_asset_path(asset_id: str, work_dir: Path) -> Path | None: """从 asset_id 解析到本地文件路径。 策略(按优先级): - 1. 如果 asset_id 是本地绝对路径(/var/storage/...)→ 直接返回 + 1. 如果 asset_id 是本地绝对路径(/var/storage/...)→ 安全校验后返回 2. 如果 work_dir 下已有缓存文件 → 返回缓存路径 3. 从 OSS 下载到 work_dir/{hash}.mp4 → 返回下载路径 4. 下载失败 → 返回 None 缓存策略:以 asset_id 的 SHA256 前 16 位为文件名,避免重复下载。 - """ - # 1. 本地绝对路径 - if asset_id.startswith("/") and os.path.exists(asset_id): - return Path(asset_id) - # 2. 缓存命中 + 安全: + - 本地绝对路径必须在 ASSET_ALLOWED_DIRS 环境变量指定的目录内 + - 文件名经过 sanitize,防止路径遍历 + - 禁止空字节、控制字符 + """ + from video_processing.path_security import ( + PathSecurityError, + get_allowed_local_dirs, + is_in_allowed_dirs, + sanitize_filename, + ) + + if not asset_id or not isinstance(asset_id, str): + return None + + # 空字节检测 + if "\x00" in asset_id: + logger.warning("asset_id 包含空字节,拒绝: %s", asset_id[:50]) + return None + + # 1. 本地绝对路径 — 必须在允许的目录内 + if asset_id.startswith("/") and os.path.exists(asset_id): + try: + resolved = Path(asset_id).resolve() + if is_in_allowed_dirs(resolved, get_allowed_local_dirs()): + return resolved + else: + logger.warning( + "本地素材路径不在允许目录内,拒绝: %s (allowed=%s)", + asset_id[:80], + get_allowed_local_dirs(), + ) + return None + except (OSError, PathSecurityError): + return None + + # 2. 缓存命中(使用 hash 而非原始 ID,防止路径遍历) cache_hash = hashlib.sha256(asset_id.encode()).hexdigest()[:16] - cached_path = work_dir / f"{cache_hash}.mp4" + safe_name = sanitize_filename(cache_hash) + cached_path = work_dir / f"{safe_name}.mp4" if cached_path.exists() and cached_path.stat().st_size > 0: return cached_path - # 3. 从 OSS 下载 - if download_asset(asset_id, cached_path): + # 3. 从 OSS 下载(先标准化 key,防止路径遍历注入) + safe_key = normalize_storage_key(asset_id) + # 额外校验:存储键不能包含 ../ 或绝对路径 + if ".." in safe_key or safe_key.startswith("/"): + logger.warning("asset_id 包含路径遍历模式,拒绝下载: %s", asset_id[:80]) + return None + + if download_asset(safe_key, cached_path): return cached_path return None diff --git a/apps/worker/video_processing/path_security.py b/apps/worker/video_processing/path_security.py new file mode 100755 index 000000000..548a93910 --- /dev/null +++ b/apps/worker/video_processing/path_security.py @@ -0,0 +1,277 @@ +"""路径安全校验工具 — 路径遍历防护. + +统一的文件路径安全校验方案,覆盖所有渲染管线中的路径处理场景: +- 本地素材路径校验 +- local:// 路径 schema 校验 +- 工作目录内路径安全约束 +- 防止路径遍历攻击 (../) + +防护要点: +1. 所有用户可控路径必须在允许的目录内 +2. 解析符号链接后的真实路径仍需在允许目录内 +3. 禁止空路径、相对路径遍历、绝对路径逃逸 +4. 路径字符限制与规范化 +""" + +from __future__ import annotations + +import logging +import os +from pathlib import Path + +logger = logging.getLogger(__name__) + +# 最大路径长度 +MAX_PATH_LENGTH = 4096 + +# 允许的文件扩展名(渲染相关) +ALLOWED_MEDIA_EXTENSIONS = { + ".mp4", ".mov", ".avi", ".mkv", ".webm", ".flv", ".wmv", # 视频 + ".mp3", ".wav", ".aac", ".ogg", ".flac", ".m4a", ".wma", # 音频 + ".jpg", ".jpeg", ".png", ".gif", ".bmp", ".webp", ".tiff", # 图片 + ".srt", ".ass", ".vtt", ".sub", # 字幕 + ".txt", ".json", # 文本/配置 +} + +# local:// schema 前缀 +LOCAL_SCHEMA_PREFIX = "local://" + + +class PathSecurityError(ValueError): + """路径安全校验失败.""" + + pass + + +def safe_resolve_path( + input_path: str | Path, + base_dir: str | Path, + *, + allow_outside: bool = False, + allowed_extensions: set[str] | None = None, +) -> Path: + """安全解析路径,确保最终路径在 base_dir 内. + + Args: + input_path: 输入路径(相对或绝对) + base_dir: 基路径目录,解析后的路径必须在此目录内 + allow_outside: 是否允许路径在 base_dir 外(默认禁止) + allowed_extensions: 允许的文件扩展名集合(None 表示不限制) + + Returns: + 解析后的绝对路径 Path 对象 + + Raises: + PathSecurityError: 路径不安全 + """ + if input_path is None: + raise PathSecurityError("路径不能为空") + + path_str = str(input_path).strip() + if not path_str: + raise PathSecurityError("路径不能为空") + + if len(path_str) > MAX_PATH_LENGTH: + raise PathSecurityError(f"路径过长 ({len(path_str)} > {MAX_PATH_LENGTH})") + + # 空字节检测(必须在 Path() 之前) + if "\x00" in path_str: + raise PathSecurityError("路径包含空字节") + + # 处理 local:// schema + if path_str.startswith(LOCAL_SCHEMA_PREFIX): + path_str = path_str[len(LOCAL_SCHEMA_PREFIX):] + # local:// 后必须是相对路径(相对于 base_dir),不能是绝对路径 + if os.path.isabs(path_str): + raise PathSecurityError("local:// 路径不能是绝对路径") + + # 规范化 base_dir + base_dir = Path(base_dir).resolve() + if not base_dir.is_dir(): + raise PathSecurityError(f"基路径不是有效目录: {base_dir}") + + # 解析输入路径 + input_path_obj = Path(path_str) + + # 如果是绝对路径且不允许外部路径 + if input_path_obj.is_absolute() and not allow_outside: + raise PathSecurityError("禁止使用绝对路径(需在工作目录内)") + + # 组合并解析为绝对路径 + if input_path_obj.is_absolute(): + full_path = input_path_obj.resolve() + else: + full_path = (base_dir / input_path_obj).resolve() + + # 检查路径遍历 — 确保最终路径在 base_dir 内 + if not allow_outside: + try: + full_path.relative_to(base_dir) + except ValueError: + raise PathSecurityError( + f"路径遍历检测:路径 '{path_str}' 超出基路径 '{base_dir}' 范围" + ) + + # 扩展名校验 + if allowed_extensions is not None: + ext = full_path.suffix.lower() + if ext and ext not in allowed_extensions: + raise PathSecurityError(f"不允许的文件类型: {ext}") + + # 检查危险路径模式 + _check_dangerous_patterns(full_path) + + return full_path + + +def _check_dangerous_patterns(path: Path) -> None: + """检查危险路径模式.""" + path_str = str(path) + + # 检查空字节 + if "\x00" in path_str: + raise PathSecurityError("路径包含空字节") + + # 检查特殊设备文件(Linux) + dangerous_prefixes = [ + "/proc/", "/sys/", "/dev/", "/etc/passwd", "/etc/shadow", + "/root/", "/boot/", "/var/run/", + ] + for prefix in dangerous_prefixes: + if path_str.startswith(prefix): + raise PathSecurityError(f"禁止访问系统路径: {prefix}") + + +def is_path_safe( + input_path: str | Path, + base_dir: str | Path, + *, + allow_outside: bool = False, +) -> bool: + """便捷函数:检查路径是否安全,不抛异常.""" + try: + safe_resolve_path(input_path, base_dir, allow_outside=allow_outside) + return True + except PathSecurityError: + return False + + +def validate_local_schema_path( + schema_path: str, + work_dir: str | Path, +) -> Path: + """校验 local:// schema 路径,返回安全的本地路径. + + local:// 路径规则: + - 必须以 local:// 开头 + - 后面必须是相对路径 + - 最终解析后必须在 work_dir 内 + - 不允许 ../ 遍历 + + Args: + schema_path: local:// 开头的路径 + work_dir: 工作目录 + + Returns: + 解析后的安全路径 + + Raises: + PathSecurityError: 路径不安全 + """ + if not schema_path.startswith(LOCAL_SCHEMA_PREFIX): + raise PathSecurityError(f"路径必须以 {LOCAL_SCHEMA_PREFIX} 开头") + + return safe_resolve_path(schema_path, work_dir, allow_outside=False) + + +def sanitize_filename(filename: str) -> str: + """清理文件名,移除危险字符. + + 保留:字母、数字、下划线、连字符、点、中文字符 + 移除:路径分隔符、控制字符、特殊符号等 + """ + import re + + if not filename: + return "unnamed" + + # 移除路径分隔符和危险字符 + # 保留: 字母数字、中文字符、下划线、连字符、点、空格 + sanitized = re.sub(r'[\\/\x00-\x1f\x7f<>:"|?*]', "_", filename) + + # 移除开头的点和连续的点(防止隐藏文件和路径遍历) + while sanitized.startswith("."): + sanitized = sanitized[1:] + + # 限制长度 + if len(sanitized) > 255: + name, ext = os.path.splitext(sanitized) + sanitized = name[: 255 - len(ext)] + ext + + # 空文件名兜底 + if not sanitized or sanitized == ".": + sanitized = "unnamed" + + return sanitized + + +# ── 允许目录配置 ────────────────────────────────────────────────────────────── + + +def get_allowed_local_dirs() -> list[Path]: + """获取允许的本地素材目录列表(从环境变量读取). + + 环境变量 ASSET_ALLOWED_DIRS,多个目录用冒号分隔(Linux)或分号分隔(Windows)。 + 默认包含 /tmp。 + + 用于: + - resolve_asset_path 本地绝对路径白名单 + - PiP local_path 类型白名单 + - 贴纸本地路径白名单 + """ + env_dirs = os.environ.get("ASSET_ALLOWED_DIRS", "") + dirs: list[Path] = [] + if env_dirs: + import re + + sep = ";" if os.name == "nt" else ":" + for d in re.split(f"[{sep}]", env_dirs): + d = d.strip() + if d: + try: + dirs.append(Path(d).resolve()) + except OSError: + pass + # 默认允许 /tmp + if not dirs: + try: + dirs.append(Path("/tmp").resolve()) + except OSError: + pass + return dirs + + +def is_in_allowed_dirs(path: str | Path, allowed_dirs: list[Path] | None = None) -> bool: + """检查路径是否在允许的目录列表内. + + Args: + path: 待检查的路径 + allowed_dirs: 允许的目录列表,None 则使用默认配置 + + Returns: + True 表示在允许目录内 + """ + if allowed_dirs is None: + allowed_dirs = get_allowed_local_dirs() + + try: + resolved = Path(path).resolve() + for allowed in allowed_dirs: + try: + resolved.relative_to(allowed) + return True + except ValueError: + continue + return False + except OSError: + return False diff --git a/apps/worker/video_processing/pip_engine.py b/apps/worker/video_processing/pip_engine.py index cb84550e3..59c26d6ef 100755 --- a/apps/worker/video_processing/pip_engine.py +++ b/apps/worker/video_processing/pip_engine.py @@ -465,18 +465,44 @@ class PiPEngine: layer: PiPLayerConfig, asset_path_map: dict[str, Path], ) -> Path | None: - """验证图层素材是否可用,返回本地路径或None(降级跳过).""" + """验证图层素材是否可用,返回本地路径或None(降级跳过). + + 安全: + - local_path 类型:必须在允许的目录内,防止路径遍历 + - url 类型:必须通过 SSRF 安全校验 + """ + from video_processing.path_security import is_in_allowed_dirs + from video_processing.url_security import UrlSecurityError, validate_url_safety + try: if layer.source_type == "local_path": - path = Path(layer.source) - if path.exists(): - return path + if not layer.source: + return None + # 路径安全校验:必须在允许目录内 + src_path = Path(layer.source) + if not src_path.exists(): + return None + if not is_in_allowed_dirs(src_path): + logger.warning( + "PiP local_path 不在允许目录内,拒绝: %s", + layer.source[:80], + ) + return None + return src_path.resolve() elif layer.source_type == "asset_id": if layer.source in asset_path_map: return asset_path_map[layer.source] + return None elif layer.source_type == "url": - # URL类型由调用者负责下载,这里返回标记 - return None # 暂时不支持直接URL + # URL类型:先做SSRF安全校验,由调用者负责实际下载 + try: + validate_url_safety(layer.source, purpose="pip_source") + logger.info("PiP URL 安全校验通过: %s", layer.source[:80]) + except UrlSecurityError as e: + logger.warning("PiP URL 安全校验失败: %s (error=%s)", layer.source[:80], e) + return None + # 暂时不支持直接URL下载,返回None表示降级跳过 + return None except Exception as e: logger.warning("PiP素材验证失败: %s", e) diff --git a/apps/worker/video_processing/sticker_engine.py b/apps/worker/video_processing/sticker_engine.py index 0cb4068c9..eacf05565 100755 --- a/apps/worker/video_processing/sticker_engine.py +++ b/apps/worker/video_processing/sticker_engine.py @@ -392,10 +392,50 @@ class StickerEngine: ) parsed_stickers.append((z, config)) else: - # 图片贴纸 - image_path = s.get("image_path", "") or s.get("image_url", "") - if not image_path or not Path(image_path).exists(): - logger.warning("贴纸素材不存在,跳过: %s", image_path) + # 图片贴纸 — 安全校验:区分本地路径和URL + image_path = s.get("image_path", "") + image_url = s.get("image_url", "") + + safe_image_path: Path | None = None + + if image_path: + # 本地路径:路径遍历防护 + from video_processing.path_security import is_in_allowed_dirs + + try: + p = Path(image_path) + if not p.exists(): + logger.warning("贴纸素材不存在,跳过: %s", image_path[:80]) + continue + if not is_in_allowed_dirs(p): + logger.warning( + "贴纸路径不在允许目录内,拒绝: %s", image_path[:80] + ) + continue + safe_image_path = p.resolve() + except Exception as e: + logger.warning("贴纸路径校验失败,跳过: %s error=%s", image_path[:80], e) + continue + elif image_url: + # URL:SSRF 安全校验(暂不自动下载,仅校验安全性) + from video_processing.url_security import ( + UrlSecurityError, + validate_url_safety, + ) + + try: + validate_url_safety(image_url, purpose="sticker_image") + except UrlSecurityError as e: + logger.warning("贴纸URL安全校验失败,跳过: %s error=%s", image_url[:80], e) + continue + # URL 类型暂不支持自动下载,跳过 + logger.info("贴纸URL类型暂不支持自动下载,跳过: %s", image_url[:80]) + continue + else: + logger.warning("贴纸缺少 image_path 和 image_url,跳过") + continue + + if safe_image_path is None: continue config = ImageStickerConfig( @@ -414,11 +454,11 @@ class StickerEngine: fade_in=float(s.get("fade_in", 0)), fade_out=float(s.get("fade_out", 0)), z_index=z, - image_url=str(s.get("image_url", "")), + image_url=image_url, ) parsed_stickers.append((z, config)) image_stickers.append(config) - image_paths.append(image_path) + image_paths.append(str(safe_image_path)) except Exception as e: logger.warning("贴纸配置解析失败,跳过: %s", e) diff --git a/apps/worker/video_processing/unified_render_service.py b/apps/worker/video_processing/unified_render_service.py index c465277b0..b78c008b9 100755 --- a/apps/worker/video_processing/unified_render_service.py +++ b/apps/worker/video_processing/unified_render_service.py @@ -639,8 +639,8 @@ class UnifiedRenderService: return timeline def _extract_audio(self, video_path: Path, output_path: Path) -> None: - """从视频中提取音频为16kHz单声道wav(ASR友好格式)。""" - import subprocess + """从视频中提取音频为16kHz单声道wav(ASR友好格式).""" + from video_processing.ffmpeg_utils import run_ffmpeg cmd = [ "ffmpeg", @@ -657,15 +657,10 @@ class UnifiedRenderService: str(output_path), ] - result = subprocess.run( - cmd, - capture_output=True, - text=True, - timeout=120, - ) - - if result.returncode != 0: - raise RuntimeError(f"音频提取失败: {result.stderr[:200]}") + try: + run_ffmpeg(cmd, timeout=120) + except Exception as e: + raise RuntimeError(f"音频提取失败: {str(e)[:200]}") from e def _maybe_add_voiceover_layer( self, diff --git a/apps/worker/video_processing/url_security.py b/apps/worker/video_processing/url_security.py new file mode 100755 index 000000000..0afa46978 --- /dev/null +++ b/apps/worker/video_processing/url_security.py @@ -0,0 +1,222 @@ +"""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. 重定向防护:下载时禁止跳转到内网 +""" + +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 diff --git a/apps/worker/worker_app/tasks/compose_video.py b/apps/worker/worker_app/tasks/compose_video.py index 4150e0fc2..4509d5e02 100755 --- a/apps/worker/worker_app/tasks/compose_video.py +++ b/apps/worker/worker_app/tasks/compose_video.py @@ -6,7 +6,6 @@ from __future__ import annotations import os -import subprocess import tempfile from pathlib import Path @@ -115,16 +114,12 @@ def _compose_with_legacy_engine(task, job_service, job, plan_id: str, db) -> dic logger.info("Executing FFmpeg for job %s, plan %s", job_id, plan_id) try: - subprocess.run( - compose_cmd.command, - check=True, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - text=True, - timeout=3600, - ) - except subprocess.CalledProcessError as e: - job_service.fail_job(job_id, f"FFmpeg 执行失败: {e.stderr[:500]}") + from video_processing.ffmpeg_utils import run_ffmpeg + + run_ffmpeg(compose_cmd.command, timeout=3600) + except Exception as e: + error_msg = f"FFmpeg 执行失败: {str(e)[:500]}" + job_service.fail_job(job_id, error_msg) raise # 上传结果 diff --git a/apps/worker/worker_app/tasks/edit_plan_generation.py b/apps/worker/worker_app/tasks/edit_plan_generation.py index 5c789d99b..9b422ab57 100755 --- a/apps/worker/worker_app/tasks/edit_plan_generation.py +++ b/apps/worker/worker_app/tasks/edit_plan_generation.py @@ -257,7 +257,6 @@ def _render_with_legacy( ) -> dict: """旧引擎路径(VideoComposeService + FFmpeg filter_complex)。""" import os - import subprocess from apps.api.app.services.video_compose_service import VideoComposeService @@ -278,16 +277,11 @@ def _render_with_legacy( logger.info("执行 FFmpeg (legacy): plan_id=%s", plan_id) try: - subprocess.run( - compose_cmd.command, - check=True, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - text=True, - timeout=3600, - ) - except subprocess.CalledProcessError as e: - error_msg = f"FFmpeg 执行失败: {e.stderr[:500]}" + from video_processing.ffmpeg_utils import run_ffmpeg + + run_ffmpeg(compose_cmd.command, timeout=3600) + except Exception as e: + error_msg = f"FFmpeg 执行失败: {str(e)[:500]}" logger.error("FFmpeg 执行失败(legacy): %s — %s", plan_id, error_msg) _mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, error_msg) return {"status": "error", "message": error_msg} diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 04973fd1a..433fa07a5 100755 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -418,17 +418,31 @@ def _prepare_bgm_track( def _verify_url_accessible(url: str, timeout: float = 10.0, retries: int = 2) -> bool: """HEAD 请求校验 URL 可访问(含重试,防止 OSS 抖动误报)。 + 安全: + - 请求前先做 SSRF 安全校验(内网IP/回环地址/链路本地地址等) + - scheme 仅允许 http/https + - 端口仅允许 80/443 + Args: url: 待校验的 URL timeout: 单次请求超时时间(秒) retries: 最大重试次数(默认 2 次,首次失败后间隔 1s 重试) Returns: - True 表示 URL 可访问(HTTP 2xx/3xx),False 表示所有尝试均失败。 + True 表示 URL 可访问(HTTP 2xx/3xx),False 表示所有尝试均失败或安全校验不通过。 """ import time import urllib.request + from video_processing.url_security import UrlSecurityError, validate_url_safety + + # P0-1 SSRF 防护:请求前先校验 URL 安全性 + try: + validate_url_safety(url, purpose="url_verify") + except UrlSecurityError as e: + logger.warning("URL 安全校验失败,拒绝访问: url=%s error=%s", url[:80], e) + return False + last_error: Exception | None = None for attempt in range(1 + retries): try: diff --git a/tests/unit/test_path_security.py b/tests/unit/test_path_security.py new file mode 100755 index 000000000..78bbebcaf --- /dev/null +++ b/tests/unit/test_path_security.py @@ -0,0 +1,241 @@ +"""路径安全校验工具单元测试 — 路径遍历防护.""" + +from __future__ import annotations + +import os +import sys +import tempfile +import unittest +from pathlib import Path + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "worker")) + +from video_processing.path_security import ( # noqa: E402 + PathSecurityError, + get_allowed_local_dirs, + is_in_allowed_dirs, + is_path_safe, + safe_resolve_path, + sanitize_filename, + validate_local_schema_path, +) + + +class TestSafeResolvePath(unittest.TestCase): + """安全路径解析测试.""" + + def setUp(self): + self.tmpdir = tempfile.mkdtemp() + + def tearDown(self): + import shutil + + shutil.rmtree(self.tmpdir, ignore_errors=True) + + # ── 正常路径 ───────────────────────────────────────────────────────── + + def test_simple_relative_path(self): + """简单相对路径应该正常解析.""" + result = safe_resolve_path("test.mp4", self.tmpdir) + self.assertEqual(result.name, "test.mp4") + self.assertTrue(str(result).startswith(self.tmpdir)) + + def test_subdirectory_path(self): + """子目录路径应该正常解析.""" + result = safe_resolve_path("sub/dir/file.mp4", self.tmpdir) + self.assertTrue(str(result).startswith(self.tmpdir)) + self.assertIn("sub/dir/file.mp4", str(result).replace("\\", "/")) + + def test_dot_slash_path(self): + """./ 开头的路径应该正常解析.""" + result = safe_resolve_path("./test.mp4", self.tmpdir) + self.assertEqual(result.name, "test.mp4") + + # ── 路径遍历防护 ───────────────────────────────────────────────────── + + def test_parent_traversal_rejected(self): + """../ 路径遍历应该被拒绝.""" + with self.assertRaises(PathSecurityError): + safe_resolve_path("../etc/passwd", self.tmpdir) + + def test_multiple_parent_traversal_rejected(self): + """多级 ../ 遍历应该被拒绝.""" + with self.assertRaises(PathSecurityError): + safe_resolve_path("../../etc/passwd", self.tmpdir) + + def test_mixed_traversal_rejected(self): + """混合路径遍历应该被拒绝.""" + with self.assertRaises(PathSecurityError): + safe_resolve_path("./sub/../../etc/shadow", self.tmpdir) + + def test_absolute_path_rejected(self): + """绝对路径(超出基目录)应该被拒绝.""" + with self.assertRaises(PathSecurityError): + safe_resolve_path("/etc/passwd", self.tmpdir) + + # ── 空字节注入 ─────────────────────────────────────────────────────── + + def test_null_byte_rejected(self): + """空字节注入应该被拒绝.""" + with self.assertRaises(PathSecurityError): + safe_resolve_path("test\x00.mp4", self.tmpdir) + + # ── 空路径 ────────────────────────────────────────────────────────── + + def test_empty_path_rejected(self): + """空路径应该被拒绝.""" + with self.assertRaises(PathSecurityError): + safe_resolve_path("", self.tmpdir) + + def test_none_path_rejected(self): + """None 路径应该被拒绝.""" + with self.assertRaises(PathSecurityError): + safe_resolve_path(None, self.tmpdir) # type: ignore + + def test_whitespace_path_rejected(self): + """空白路径应该被拒绝.""" + with self.assertRaises(PathSecurityError): + safe_resolve_path(" ", self.tmpdir) + + # ── 路径长度 ──────────────────────────────────────────────────────── + + def test_too_long_path_rejected(self): + """超长路径应该被拒绝.""" + long_path = "a" * 5000 + ".mp4" + with self.assertRaises(PathSecurityError): + safe_resolve_path(long_path, self.tmpdir) + + # ── 系统路径防护 ───────────────────────────────────────────────────── + + def test_proc_path_rejected_when_absolute(self): + """/proc/ 路径在绝对路径模式下应该被拒绝(因为超出基目录).""" + with self.assertRaises(PathSecurityError): + safe_resolve_path("/proc/self/environ", self.tmpdir) + + # ── 扩展名校验 ─────────────────────────────────────────────────────── + + def test_extension_whitelist_pass(self): + """白名单内的扩展名应该通过.""" + result = safe_resolve_path( + "test.mp4", self.tmpdir, + allowed_extensions={".mp4", ".mov"}, + ) + self.assertEqual(result.suffix.lower(), ".mp4") + + def test_extension_whitelist_reject(self): + """白名单外的扩展名应该被拒绝.""" + with self.assertRaises(PathSecurityError): + safe_resolve_path( + "test.exe", self.tmpdir, + allowed_extensions={".mp4", ".mov"}, + ) + + +class TestLocalSchemaPath(unittest.TestCase): + """local:// schema 路径测试.""" + + def setUp(self): + self.tmpdir = tempfile.mkdtemp() + + def tearDown(self): + import shutil + + shutil.rmtree(self.tmpdir, ignore_errors=True) + + def test_valid_local_schema(self): + """有效的 local:// 相对路径应该通过.""" + # 创建测试文件 + test_file = Path(self.tmpdir) / "test.mp4" + test_file.touch() + + result = validate_local_schema_path("local://test.mp4", self.tmpdir) + self.assertTrue(result.exists()) + + def test_local_schema_absolute_rejected(self): + """local:// + 绝对路径应该被拒绝.""" + with self.assertRaises(PathSecurityError): + validate_local_schema_path("local:///etc/passwd", self.tmpdir) + + def test_local_schema_traversal_rejected(self): + """local:// + 路径遍历应该被拒绝.""" + with self.assertRaises(PathSecurityError): + validate_local_schema_path("local://../etc/passwd", self.tmpdir) + + def test_non_local_schema_rejected(self): + """非 local:// 开头的路径应该被拒绝.""" + with self.assertRaises(PathSecurityError): + validate_local_schema_path("http://example.com/test", self.tmpdir) + + +class TestSanitizeFilename(unittest.TestCase): + """文件名清理测试.""" + + def test_normal_filename(self): + """正常文件名应该保持不变.""" + self.assertEqual(sanitize_filename("video.mp4"), "video.mp4") + + def test_path_separators_removed(self): + """路径分隔符应该被替换.""" + self.assertNotIn("/", sanitize_filename("../path/to/file.mp4")) + self.assertNotIn("\\", sanitize_filename("..\\path\\file.mp4")) + + def test_leading_dots_removed(self): + """开头的点应该被移除.""" + result = sanitize_filename(".hidden") + self.assertFalse(result.startswith(".")) + self.assertEqual(result, "hidden") + + def test_multiple_leading_dots_removed(self): + """多个开头的点应该全部被移除.""" + result = sanitize_filename("...hidden") + self.assertFalse(result.startswith(".")) + + def test_empty_filename_default(self): + """空文件名应该返回 unnamed.""" + self.assertEqual(sanitize_filename(""), "unnamed") + + def test_special_chars_removed(self): + """特殊字符应该被替换.""" + result = sanitize_filename('file:"test|?*.mp4') + self.assertNotIn("<", result) + self.assertNotIn(">", result) + self.assertNotIn(":", result) + self.assertNotIn('"', result) + self.assertNotIn("|", result) + self.assertNotIn("?", result) + self.assertNotIn("*", result) + + def test_chinese_filename_preserved(self): + """中文文件名应该保留.""" + result = sanitize_filename("视频素材.mp4") + self.assertIn("视频素材", result) + + def test_long_filename_truncated(self): + """超长文件名应该被截断.""" + long_name = "a" * 300 + ".mp4" + result = sanitize_filename(long_name) + self.assertLessEqual(len(result), 255) + self.assertTrue(result.endswith(".mp4")) + + +class TestAllowedDirs(unittest.TestCase): + """允许目录配置测试.""" + + def test_get_allowed_dirs_returns_list(self): + """get_allowed_local_dirs 应该返回列表.""" + dirs = get_allowed_local_dirs() + self.assertIsInstance(dirs, list) + + def test_is_in_allowed_dirs_tmp(self): + """/tmp 应该在默认允许目录内.""" + self.assertTrue(is_in_allowed_dirs("/tmp/test.mp4")) + + def test_is_path_safe_convenience(self): + """is_path_safe 便捷函数应该正常工作.""" + with tempfile.TemporaryDirectory() as tmpdir: + self.assertTrue(is_path_safe("test.mp4", tmpdir)) + self.assertFalse(is_path_safe("../etc/passwd", tmpdir)) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit/test_url_security.py b/tests/unit/test_url_security.py new file mode 100755 index 000000000..fa76105b3 --- /dev/null +++ b/tests/unit/test_url_security.py @@ -0,0 +1,164 @@ +"""URL 安全校验工具单元测试 — SSRF 防护.""" + +from __future__ import annotations + +import os +import sys +import unittest + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "worker")) + +from video_processing.url_security import ( # noqa: E402 + UrlSecurityError, + is_url_safe, + validate_url_safety, +) + + +class TestUrlSecurityValidation(unittest.TestCase): + """URL 安全校验测试.""" + + # ── Scheme 白名单 ────────────────────────────────────────────────────── + + def test_http_scheme_allowed(self): + """HTTP scheme 应该被允许.""" + result = validate_url_safety("http://example.com/test", purpose="test") + self.assertEqual(result, "http://example.com/test") + + def test_https_scheme_allowed(self): + """HTTPS scheme 应该被允许.""" + result = validate_url_safety("https://example.com/test", purpose="test") + self.assertEqual(result, "https://example.com/test") + + def test_file_scheme_rejected(self): + """file:// scheme 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("file:///etc/passwd", purpose="test") + + def test_ftp_scheme_rejected(self): + """ftp:// scheme 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("ftp://example.com/test", purpose="test") + + def test_empty_scheme_rejected(self): + """空 scheme 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("example.com/test", purpose="test") + + # ── 端口白名单 ──────────────────────────────────────────────────────── + + def test_port_80_allowed(self): + """端口 80 应该被允许.""" + # 80端口是默认HTTP端口,不显式指定也可以 + result = validate_url_safety("http://example.com:80/test", purpose="test") + self.assertIn("example.com", result) + + def test_port_443_allowed(self): + """端口 443 应该被允许.""" + result = validate_url_safety("https://example.com:443/test", purpose="test") + self.assertIn("example.com", result) + + def test_port_8080_rejected(self): + """非标准端口 8080 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http://example.com:8080/test", purpose="test") + + def test_port_22_rejected(self): + """SSH 端口 22 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http://example.com:22/test", purpose="test") + + # ── SSRF: 直接 IP 访问 ─────────────────────────────────────────────── + + def test_loopback_ip_rejected(self): + """回环地址 127.0.0.1 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http://127.0.0.1/test", purpose="test") + + def test_private_ip_192_rejected(self): + """内网地址 192.168.x.x 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http://192.168.1.1/test", purpose="test") + + def test_private_ip_10_rejected(self): + """内网地址 10.x.x.x 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http://10.0.0.1/test", purpose="test") + + def test_private_ip_172_rejected(self): + """内网地址 172.16.x.x 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http://172.16.0.1/test", purpose="test") + + def test_unspecified_ip_rejected(self): + """未指定地址 0.0.0.0 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http://0.0.0.0/test", purpose="test") + + def test_ipv6_loopback_rejected(self): + """IPv6 回环 ::1 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http://[::1]/test", purpose="test") + + def test_ipv6_link_local_rejected(self): + """IPv6 链路本地地址应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http://[fe80::1]/test", purpose="test") + + # ── SSRF: 内网主机名 ───────────────────────────────────────────────── + + def test_localhost_rejected(self): + """localhost 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http://localhost/test", purpose="test") + + def test_local_domain_rejected(self): + """.local 域名应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http://printer.local/test", purpose="test") + + def test_internal_domain_rejected(self): + """.internal 域名应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http://db.internal/test", purpose="test") + + # ── URL 格式校验 ───────────────────────────────────────────────────── + + def test_empty_url_rejected(self): + """空 URL 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("", purpose="test") + + def test_none_url_rejected(self): + """None URL 应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety(None, purpose="test") # type: ignore + + def test_url_too_long_rejected(self): + """超长 URL 应该被拒绝.""" + long_url = "https://example.com/" + "a" * 3000 + with self.assertRaises(UrlSecurityError): + validate_url_safety(long_url, purpose="test") + + def test_no_hostname_rejected(self): + """缺少主机名应该被拒绝.""" + with self.assertRaises(UrlSecurityError): + validate_url_safety("http:///test", purpose="test") + + # ── is_url_safe 便捷函数 ───────────────────────────────────────────── + + def test_is_url_safe_true(self): + """安全 URL 应该返回 True.""" + self.assertTrue(is_url_safe("https://example.com/test", purpose="test")) + + def test_is_url_safe_false(self): + """不安全 URL 应该返回 False.""" + self.assertFalse(is_url_safe("http://127.0.0.1/test", purpose="test")) + + def test_is_url_safe_empty(self): + """空 URL 应该返回 False.""" + self.assertFalse(is_url_safe("", purpose="test")) + + +if __name__ == "__main__": + unittest.main()