Compare commits
25 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| b4b142180c | |||
| 7f767e2dd1 | |||
| de42b1960e | |||
| f377670076 | |||
| ea93387f98 | |||
| 674fc6763d | |||
| 405fb4c8d3 | |||
| 7e172c0907 | |||
| de1816bbd9 | |||
| bbcc8a12bd | |||
| 12862b4e23 | |||
| 3b2bf414eb | |||
| f681904f61 | |||
| 71f7b6795b | |||
| c6c234757a | |||
| 99a6b4d104 | |||
| 745ea06659 | |||
| 6db705a05d | |||
| 2f03e13928 | |||
| 334e2b1fc2 | |||
| 3b80edd8c5 | |||
| 01daa72f2a | |||
| e0f841fcee | |||
| 70a85b1463 | |||
| 9e10e31a29 |
+1625
-1470
File diff suppressed because one or more lines are too long
@@ -24,12 +24,17 @@ from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_duration, probe_video_info, run_ffmpeg
|
||||
from video_processing.path_security import PathSecurityError, is_in_allowed_dirs, safe_resolve_path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── 常量 ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
MAX_CONCAT_SEGMENTS = 50 # 最大拼接段数(安全上限,防止OOM)
|
||||
|
||||
ALLOWED_VIDEO_EXTENSIONS = {".mp4", ".mov", ".avi", ".mkv", ".webm", ".flv", ".wmv"}
|
||||
|
||||
# concat demuxer 要求一致的参数列表
|
||||
CONCAT_DEMUXER_REQUIRED_PARAMS = [
|
||||
"codec_name", # 视频编码
|
||||
@@ -144,6 +149,50 @@ class ConcatConfig:
|
||||
return len([s for s in self.segments if s.video_path])
|
||||
|
||||
|
||||
# ── 路径安全校验 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _validate_video_path(video_path: str, work_dir: Path) -> None:
|
||||
"""校验视频文件路径安全性.
|
||||
|
||||
规则:
|
||||
- local:// schema → 必须在 work_dir 内
|
||||
- 相对路径 → 必须在 work_dir 内
|
||||
- 绝对路径 → 必须在允许目录白名单内
|
||||
- 扩展名必须是视频格式
|
||||
|
||||
Raises:
|
||||
PathSecurityError: 路径不安全
|
||||
"""
|
||||
if not video_path or not isinstance(video_path, str):
|
||||
raise PathSecurityError("视频路径不能为空")
|
||||
|
||||
# 本地路径(local:// 或相对路径 / 绝对路径)
|
||||
if video_path.startswith("local://") or not video_path.startswith(("http://", "https://", "oss://")):
|
||||
is_abs = video_path.startswith("/") and not video_path.startswith("local://")
|
||||
resolved_path = safe_resolve_path(
|
||||
video_path,
|
||||
work_dir,
|
||||
allow_outside=is_abs,
|
||||
allowed_extensions=ALLOWED_VIDEO_EXTENSIONS,
|
||||
)
|
||||
# 绝对路径额外检查白名单目录(用realpath规范化后的真实路径比较,防止 ../ 遍历绕过)
|
||||
if is_abs:
|
||||
resolved_work_dir = work_dir.resolve()
|
||||
try:
|
||||
resolved_path.relative_to(resolved_work_dir)
|
||||
except ValueError:
|
||||
if not is_in_allowed_dirs(resolved_path):
|
||||
raise PathSecurityError(f"视频路径不在允许目录内: {video_path[:80]}")
|
||||
# URL类型路径不做本地路径校验(由下载阶段的SSRF防护负责)
|
||||
# 但检查扩展名
|
||||
else:
|
||||
path_part = video_path.split("?")[0].split("#")[0]
|
||||
ext = Path(path_part).suffix.lower()
|
||||
if ext and ext not in ALLOWED_VIDEO_EXTENSIONS:
|
||||
raise PathSecurityError(f"不允许的视频文件类型: {ext}")
|
||||
|
||||
|
||||
# ── 视频拼接引擎 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -179,6 +228,27 @@ class ConcatEngine:
|
||||
if not valid_segments:
|
||||
raise ValueError("No valid video segments to concat")
|
||||
|
||||
# ── 安全校验:段数上限 ──
|
||||
if len(valid_segments) > MAX_CONCAT_SEGMENTS:
|
||||
raise ValueError(f"Too many concat segments: {len(valid_segments)} > {MAX_CONCAT_SEGMENTS}")
|
||||
|
||||
# ── 安全校验:所有视频路径白名单校验 ──
|
||||
safe_segments = []
|
||||
for seg in valid_segments:
|
||||
try:
|
||||
_validate_video_path(seg.video_path, self.work_dir)
|
||||
safe_segments.append(seg)
|
||||
except PathSecurityError as e:
|
||||
logger.warning("[concat] skip segment: path security check failed: %s", e)
|
||||
|
||||
if len(safe_segments) != len(valid_segments):
|
||||
valid_segments = safe_segments
|
||||
config.segments = safe_segments
|
||||
logger.info("[concat] %d segments passed security check", len(safe_segments))
|
||||
|
||||
if not valid_segments:
|
||||
raise ValueError("No valid video segments after security check")
|
||||
|
||||
if len(valid_segments) == 1:
|
||||
# 只有一段,直接复制
|
||||
import shutil
|
||||
|
||||
@@ -119,6 +119,54 @@ 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:
|
||||
"""探测文件是否包含音频流。
|
||||
|
||||
|
||||
@@ -21,6 +21,7 @@ from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_duration, run_ffmpeg
|
||||
from video_processing.path_security import PathSecurityError, is_in_allowed_dirs, safe_resolve_path
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from video_processing.render_audio import RenderContext
|
||||
@@ -36,6 +37,8 @@ TRACK_TYPE_VOICEOVER = "voiceover" # 配音(TTS/人声)
|
||||
TRACK_TYPE_SFX = "sfx" # 音效
|
||||
TRACK_TYPE_AMBIENT = "ambient" # 环境音
|
||||
|
||||
MAX_AUDIO_TRACKS = 8 # 最大混音轨道数(安全上限,防止资源耗尽)
|
||||
|
||||
# 各轨道默认音量(相对主音频)
|
||||
DEFAULT_VOLUMES = {
|
||||
TRACK_TYPE_MAIN: 1.0,
|
||||
@@ -153,6 +156,56 @@ class MultiTrackMixConfig:
|
||||
return len([t for t in self.tracks if t.enabled and t.audio_path]) > 0
|
||||
|
||||
|
||||
# ── 路径安全校验 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
ALLOWED_AUDIO_EXTENSIONS = {".mp3", ".wav", ".aac", ".ogg", ".flac", ".m4a", ".wma"}
|
||||
|
||||
|
||||
def _validate_audio_path(audio_path: str, work_dir: Path) -> None:
|
||||
"""校验音频文件路径安全性.
|
||||
|
||||
规则:
|
||||
- local:// schema → 必须在 work_dir 内
|
||||
- 相对路径 → 必须在 work_dir 内
|
||||
- 绝对路径 → 必须在允许目录白名单内
|
||||
- 扩展名必须是音频格式
|
||||
|
||||
Raises:
|
||||
PathSecurityError: 路径不安全
|
||||
"""
|
||||
if not audio_path or not isinstance(audio_path, str):
|
||||
raise PathSecurityError("音频路径不能为空")
|
||||
|
||||
# 本地路径(local:// 或相对路径)
|
||||
if audio_path.startswith("local://") or not audio_path.startswith(("http://", "https://", "oss://")):
|
||||
is_abs = audio_path.startswith("/") and not audio_path.startswith("local://")
|
||||
resolved_path = safe_resolve_path(
|
||||
audio_path,
|
||||
work_dir,
|
||||
allow_outside=is_abs,
|
||||
allowed_extensions=ALLOWED_AUDIO_EXTENSIONS,
|
||||
)
|
||||
# 绝对路径额外检查白名单目录(用realpath规范化后的真实路径比较,防止 ../ 遍历绕过)
|
||||
if is_abs:
|
||||
resolved_work_dir = work_dir.resolve()
|
||||
try:
|
||||
resolved_path.relative_to(resolved_work_dir)
|
||||
except ValueError:
|
||||
if not is_in_allowed_dirs(resolved_path):
|
||||
raise PathSecurityError(f"音频路径不在允许目录内: {audio_path[:80]}")
|
||||
# URL类型路径不做本地路径校验(由下载阶段的SSRF防护负责)
|
||||
# 但检查扩展名
|
||||
else:
|
||||
# URL路径,检查扩展名白名单(取 ? 之前的部分)
|
||||
path_part = audio_path.split("?")[0].split("#")[0]
|
||||
from pathlib import Path as _P
|
||||
|
||||
ext = _P(path_part).suffix.lower()
|
||||
if ext and ext not in ALLOWED_AUDIO_EXTENSIONS:
|
||||
raise PathSecurityError(f"不允许的音频文件类型: {ext}")
|
||||
|
||||
|
||||
# ── 单轨道预处理 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -289,6 +342,39 @@ def mix_multi_track(
|
||||
if target_duration <= 0:
|
||||
target_duration = 5.0
|
||||
|
||||
# ── 安全校验:轨道数量上限 ──
|
||||
enabled_tracks = [t for t in config.tracks if t.enabled and t.audio_path]
|
||||
if len(enabled_tracks) > MAX_AUDIO_TRACKS:
|
||||
logger.warning(
|
||||
"[multi-track] too many tracks: %d > %d, truncating to max",
|
||||
len(enabled_tracks),
|
||||
MAX_AUDIO_TRACKS,
|
||||
)
|
||||
enabled_tracks = enabled_tracks[:MAX_AUDIO_TRACKS]
|
||||
# 更新 config.tracks 为截断后的列表
|
||||
config.tracks = enabled_tracks
|
||||
|
||||
# ── 安全校验:所有音频路径白名单校验 ──
|
||||
# 主音频路径
|
||||
try:
|
||||
_validate_audio_path(str(main_audio_path), ctx.work_dir)
|
||||
except PathSecurityError as e:
|
||||
logger.error("[multi-track] main audio path security check failed: %s", e)
|
||||
raise
|
||||
|
||||
# 各轨道音频路径
|
||||
valid_tracks = []
|
||||
for track in enabled_tracks:
|
||||
try:
|
||||
_validate_audio_path(track.audio_path, ctx.work_dir)
|
||||
valid_tracks.append(track)
|
||||
except PathSecurityError as e:
|
||||
logger.warning("[multi-track] skip track %s: path security check failed: %s", track.track_id, e)
|
||||
|
||||
if len(valid_tracks) != len(enabled_tracks):
|
||||
config.tracks = valid_tracks
|
||||
logger.info("[multi-track] %d tracks passed security check", len(valid_tracks))
|
||||
|
||||
# 收集所有有效轨道(已预处理好的)
|
||||
prepared_tracks: list[Path] = []
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -0,0 +1,303 @@
|
||||
"""路径安全校验工具 — 路径遍历防护.
|
||||
|
||||
统一的文件路径安全校验方案,覆盖所有渲染管线中的路径处理场景:
|
||||
- 本地素材路径校验
|
||||
- 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()) # nosec B108
|
||||
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
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -392,10 +392,48 @@ 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 +452,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)
|
||||
|
||||
@@ -28,6 +28,7 @@ from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from video_processing.path_security import PathSecurityError, is_in_allowed_dirs, safe_resolve_path
|
||||
from video_processing.render_subtitles import generate_ass_subtitles
|
||||
from video_processing.subtitle_generator import generate_ass_from_timeline
|
||||
|
||||
@@ -36,6 +37,8 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 常量 ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
ALLOWED_SUBTITLE_EXTENSIONS = {".srt", ".ass", ".vtt", ".sub"}
|
||||
|
||||
# 9宫格位置映射(ASS alignment 编号)
|
||||
POSITION_ALIGNMENT = {
|
||||
"top_left": 7,
|
||||
@@ -614,6 +617,7 @@ def build_subtitle_filter(
|
||||
*,
|
||||
video_input_label: str = "0:v",
|
||||
output_label: str = "subtitled",
|
||||
work_dir: Path | str | None = None,
|
||||
) -> str:
|
||||
"""生成 FFmpeg subtitles 滤镜字符串.
|
||||
|
||||
@@ -621,13 +625,63 @@ def build_subtitle_filter(
|
||||
ass_path: ASS 字幕文件路径
|
||||
video_input_label: 视频输入标签(如 "0:v" 或 "[v_out]")
|
||||
output_label: 输出标签
|
||||
work_dir: 工作目录(必填,用于路径安全校验,防止路径遍历绕过)
|
||||
|
||||
Returns:
|
||||
filter_complex 片段,如 "[0:v]subtitles=xxx.ass[subtitled]"
|
||||
|
||||
Raises:
|
||||
PathSecurityError: 字幕路径不安全或 work_dir 未提供
|
||||
"""
|
||||
# ── 安全校验:字幕文件路径白名单 ──
|
||||
ass_path_str = str(ass_path)
|
||||
if work_dir is None or not str(work_dir).strip():
|
||||
raise PathSecurityError("work_dir 必须提供,不能为 None 或空")
|
||||
|
||||
_validate_subtitle_path(ass_path_str, Path(work_dir))
|
||||
|
||||
# FFmpeg subtitles filter 的路径需要转义:
|
||||
# - Windows 路径的 \ → /
|
||||
# - 冒号 : → \:
|
||||
# - 单引号 ' → '\''
|
||||
safe_path = str(ass_path).replace("\\", "/").replace(":", "\\:").replace("'", "'\\''")
|
||||
safe_path = ass_path_str.replace("\\", "/").replace(":", "\\:").replace("'", "'\\''")
|
||||
return f"{video_input_label}subtitles='{safe_path}'[{output_label}]"
|
||||
|
||||
|
||||
def _validate_subtitle_path(subtitle_path: str, work_dir: Path) -> None:
|
||||
"""校验字幕文件路径安全性.
|
||||
|
||||
规则:
|
||||
- 必须是本地路径(不支持远程URL字幕)
|
||||
- local:// schema → 必须在 work_dir 内
|
||||
- 相对路径 → 必须在 work_dir 内
|
||||
- 绝对路径 → 必须在允许目录白名单内
|
||||
- 扩展名必须是字幕格式
|
||||
|
||||
Raises:
|
||||
PathSecurityError: 路径不安全
|
||||
"""
|
||||
if not subtitle_path or not isinstance(subtitle_path, str):
|
||||
raise PathSecurityError("字幕路径不能为空")
|
||||
|
||||
# 不允许远程URL字幕(subtitles滤镜不支持远程加载,且有SSRF风险)
|
||||
if subtitle_path.startswith(("http://", "https://", "oss://")):
|
||||
raise PathSecurityError("不允许使用远程URL字幕文件")
|
||||
|
||||
is_abs = subtitle_path.startswith("/") and not subtitle_path.startswith("local://")
|
||||
|
||||
resolved_path = safe_resolve_path(
|
||||
subtitle_path,
|
||||
work_dir,
|
||||
allow_outside=is_abs,
|
||||
allowed_extensions=ALLOWED_SUBTITLE_EXTENSIONS,
|
||||
)
|
||||
|
||||
# 绝对路径额外检查白名单目录(用realpath规范化后的真实路径比较,防止 ../ 遍历绕过)
|
||||
if is_abs:
|
||||
resolved_work_dir = work_dir.resolve()
|
||||
try:
|
||||
resolved_path.relative_to(resolved_work_dir)
|
||||
except ValueError:
|
||||
if not is_in_allowed_dirs(resolved_path):
|
||||
raise PathSecurityError(f"字幕路径不在允许目录内: {subtitle_path[:80]}")
|
||||
|
||||
@@ -640,8 +640,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",
|
||||
@@ -658,15 +658,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,
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
"""URL 安全校验工具 — SSRF 防护(向后兼容层).
|
||||
|
||||
本模块为向后兼容而保留,实际实现已迁移至 packages.shared.url_security。
|
||||
所有符号均从该模块重新导出,请新代码直接 import packages.shared.url_security。
|
||||
"""
|
||||
|
||||
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,
|
||||
)
|
||||
@@ -130,6 +130,8 @@ class AssetAnalyzer:
|
||||
info = VideoInfo()
|
||||
|
||||
try:
|
||||
from video_processing.ffmpeg_utils import run_ffprobe
|
||||
|
||||
cmd = [
|
||||
"ffprobe",
|
||||
"-v",
|
||||
@@ -140,38 +142,31 @@ class AssetAnalyzer:
|
||||
"-show_streams",
|
||||
self.video_path,
|
||||
]
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=30,
|
||||
)
|
||||
stdout, _ = run_ffprobe(cmd, timeout=30)
|
||||
data = json.loads(stdout)
|
||||
streams = data.get("streams", [])
|
||||
format_info = data.get("format", {})
|
||||
|
||||
if result.returncode == 0:
|
||||
data = json.loads(result.stdout)
|
||||
streams = data.get("streams", [])
|
||||
format_info = data.get("format", {})
|
||||
for stream in streams:
|
||||
if stream.get("codec_type") == "video":
|
||||
info.width = int(stream.get("width", 0))
|
||||
info.height = int(stream.get("height", 0))
|
||||
info.codec = stream.get("codec_name", "")
|
||||
|
||||
for stream in streams:
|
||||
if stream.get("codec_type") == "video":
|
||||
info.width = int(stream.get("width", 0))
|
||||
info.height = int(stream.get("height", 0))
|
||||
info.codec = stream.get("codec_name", "")
|
||||
# 解析帧率
|
||||
fps_str = stream.get("r_frame_rate", "0/1")
|
||||
if "/" in fps_str:
|
||||
num, denom = fps_str.split("/")
|
||||
info.fps = float(num) / float(denom) if float(denom) != 0 else 0.0
|
||||
else:
|
||||
info.fps = float(fps_str)
|
||||
|
||||
# 解析帧率
|
||||
fps_str = stream.get("r_frame_rate", "0/1")
|
||||
if "/" in fps_str:
|
||||
num, denom = fps_str.split("/")
|
||||
info.fps = float(num) / float(denom) if float(denom) != 0 else 0.0
|
||||
else:
|
||||
info.fps = float(fps_str)
|
||||
elif stream.get("codec_type") == "audio":
|
||||
info.has_audio = True
|
||||
|
||||
elif stream.get("codec_type") == "audio":
|
||||
info.has_audio = True
|
||||
|
||||
info.duration = float(format_info.get("duration", 0))
|
||||
info.bitrate = int(format_info.get("bit_rate", 0))
|
||||
info.file_size = int(format_info.get("size", 0))
|
||||
info.duration = float(format_info.get("duration", 0))
|
||||
info.bitrate = int(format_info.get("bit_rate", 0))
|
||||
info.file_size = int(format_info.get("size", 0))
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to get video info: {e}")
|
||||
@@ -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,16 @@ 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 (
|
||||
ALLOWED_VIDEO_MIME_TYPES,
|
||||
safe_download_file,
|
||||
)
|
||||
|
||||
urllib.request.urlretrieve(url, dest_path) # nosec B310
|
||||
safe_download_file(
|
||||
url,
|
||||
dest_path,
|
||||
purpose="batch_video_download",
|
||||
allowed_mime_types=ALLOWED_VIDEO_MIME_TYPES | {"application/octet-stream"},
|
||||
timeout=300.0,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
# 上传结果
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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:
|
||||
@@ -418,17 +436,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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -15,6 +15,7 @@ import httpx
|
||||
|
||||
from packages.application.cosyvoice_service import CosyVoiceError, CosyVoiceService
|
||||
from packages.application.tts_job.text_splitter import split_text
|
||||
from packages.shared.url_security import ALLOWED_AUDIO_MIME_TYPES, safe_download_bytes
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -224,10 +225,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 推送。
|
||||
|
||||
Executable → Regular
+23
-15
@@ -19,16 +19,19 @@ from typing import Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from packages.application.cosyvoice_service import (
|
||||
CosyVoiceAuthError,
|
||||
CosyVoiceError,
|
||||
CosyVoiceService,
|
||||
)
|
||||
from packages.application.cosyvoice_service import CosyVoiceAuthError, CosyVoiceError, CosyVoiceService
|
||||
from packages.application.tts_job.audio_merger import AudioMerger
|
||||
from packages.application.tts_job.text_splitter import split_text
|
||||
from packages.domain.tts_job import TTSJob, TTSJobStatus
|
||||
from packages.ports.tts_job_repository import TTSJobRepository
|
||||
from packages.shared.storage import SharedStorageService, get_shared_storage_service
|
||||
from packages.shared.url_security import (
|
||||
ALLOWED_AUDIO_MIME_TYPES,
|
||||
UrlSecurityError,
|
||||
safe_download_bytes,
|
||||
safe_download_file,
|
||||
validate_url_safety,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -96,10 +99,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 +469,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)
|
||||
|
||||
# 合并
|
||||
|
||||
@@ -0,0 +1,402 @@
|
||||
"""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
|
||||
+50
-18
@@ -48,23 +48,55 @@ omit = [
|
||||
]
|
||||
branch = true
|
||||
|
||||
[tool.coverage.report]
|
||||
exclude_lines = [
|
||||
"pragma: no cover",
|
||||
"def __repr__",
|
||||
"if __name__ == .__main__.:",
|
||||
"raise NotImplementedError",
|
||||
"pass",
|
||||
"if TYPE_CHECKING:",
|
||||
"class .*Protocol",
|
||||
"@abstractmethod",
|
||||
"raise AssertionError",
|
||||
"raise RuntimeError",
|
||||
"if 0:",
|
||||
"if __debug__:",
|
||||
[tool.ruff]
|
||||
target-version = "py311"
|
||||
line-length = 120
|
||||
exclude = [
|
||||
".git",
|
||||
".cache",
|
||||
"__pycache__",
|
||||
".venv",
|
||||
"venv",
|
||||
"node_modules",
|
||||
"alembic",
|
||||
".gitea",
|
||||
".next",
|
||||
"dist",
|
||||
"build",
|
||||
]
|
||||
show_missing = true
|
||||
skip_covered = false
|
||||
|
||||
[tool.coverage.xml]
|
||||
output = "coverage.xml"
|
||||
[tool.ruff.lint]
|
||||
# 当前阶段:摸底模式,规则集与原flake8对齐
|
||||
# 后续迭代计划:
|
||||
# Phase 1: 修完 bugbear 后正式替换 flake8
|
||||
# Phase 2: 启用 UP(pyupgrade) + SIM(simplify)
|
||||
# Phase 3: 启用 RET(return) + ARG(unused-args)
|
||||
select = [
|
||||
"E", # pycodestyle errors(同flake8)
|
||||
"F", # pyflakes(同flake8)
|
||||
"W", # pycodestyle warnings(同flake8)
|
||||
"B", # flake8-bugbear(新增,摸底用)
|
||||
]
|
||||
# 与原 setup.cfg flake8 配置对齐,确保不新增阻断
|
||||
ignore = [
|
||||
"E203",
|
||||
"W503",
|
||||
"E501", # line-too-long(black管)
|
||||
"E302",
|
||||
"E402", # module-import-not-at-top(循环导入多)
|
||||
"E722", # bare-except
|
||||
"W291",
|
||||
"W293",
|
||||
"F401", # unused-import
|
||||
"F403",
|
||||
"F405",
|
||||
"F841", # unused-variable
|
||||
"B008", # do-not-perform-callback-from-arg(fastapi依赖注入)
|
||||
]
|
||||
|
||||
[tool.ruff.lint.per-file-ignores]
|
||||
"__init__.py" = ["F401", "F403", "F405"]
|
||||
"tests/*" = ["E402", "F401", "F841"]
|
||||
"packages/ports/*" = ["E301", "E704"]
|
||||
"apps/*/migrations/*" = ["ALL"]
|
||||
"alembic/*" = ["ALL"]
|
||||
|
||||
@@ -716,18 +716,29 @@ class TestBuildSubtitlesFromPlan:
|
||||
class TestSubtitleFilter:
|
||||
"""字幕滤镜构建测试."""
|
||||
|
||||
def test_build_subtitle_filter(self):
|
||||
def test_build_subtitle_filter(self, work_dir):
|
||||
from video_processing.subtitle_render_engine import build_subtitle_filter
|
||||
|
||||
result = build_subtitle_filter("/tmp/test.ass", video_input_label="[v_in]", output_label="out")
|
||||
ass_file = work_dir / "test.ass"
|
||||
ass_file.write_text("test", encoding="utf-8")
|
||||
|
||||
result = build_subtitle_filter(
|
||||
str(ass_file),
|
||||
video_input_label="[v_in]",
|
||||
output_label="out",
|
||||
work_dir=work_dir,
|
||||
)
|
||||
assert "subtitles=" in result
|
||||
assert "[v_in]" in result
|
||||
assert "[out]" in result
|
||||
|
||||
def test_default_labels(self):
|
||||
def test_default_labels(self, work_dir):
|
||||
from video_processing.subtitle_render_engine import build_subtitle_filter
|
||||
|
||||
result = build_subtitle_filter("/tmp/sub.ass")
|
||||
ass_file = work_dir / "sub.ass"
|
||||
ass_file.write_text("test", encoding="utf-8")
|
||||
|
||||
result = build_subtitle_filter(str(ass_file), work_dir=work_dir)
|
||||
assert "0:v" in result
|
||||
assert "[subtitled]" in result
|
||||
|
||||
|
||||
Executable
+243
@@ -0,0 +1,243 @@
|
||||
"""路径安全校验工具单元测试 — 路径遍历防护."""
|
||||
|
||||
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<name>:"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()
|
||||
@@ -0,0 +1,484 @@
|
||||
"""PR #312 安全债务修复 单元测试.
|
||||
|
||||
测试4个P1安全修复:
|
||||
1. 多轨道混音:audio_path 路径安全 + 轨道数量上限
|
||||
2. 视频拼接:video_path 路径安全 + 段数上限
|
||||
3. 字幕渲染:字幕文件路径白名单校验
|
||||
"""
|
||||
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker"))
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
import pytest
|
||||
from video_processing.path_security import PathSecurityError
|
||||
|
||||
# ── Fixtures ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def work_dir(tmp_path):
|
||||
return tmp_path
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_audio(work_dir):
|
||||
"""生成一个测试音频文件."""
|
||||
import subprocess
|
||||
|
||||
path = work_dir / "test.aac"
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
"sine=frequency=440:duration=1:sample_rate=44100",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"128k",
|
||||
str(path),
|
||||
],
|
||||
capture_output=True,
|
||||
check=True,
|
||||
timeout=30,
|
||||
)
|
||||
return path
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_video(work_dir):
|
||||
"""生成一个测试视频文件."""
|
||||
import subprocess
|
||||
|
||||
path = work_dir / "test.mp4"
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
"testsrc=duration=1:size=320x240:rate=30",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
"sine=frequency=440:duration=1:sample_rate=44100",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-preset",
|
||||
"ultrafast",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"128k",
|
||||
"-shortest",
|
||||
str(path),
|
||||
],
|
||||
capture_output=True,
|
||||
check=True,
|
||||
timeout=60,
|
||||
)
|
||||
return path
|
||||
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
# 1. 多轨道混音安全测试
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestMultiTrackSecurity:
|
||||
"""多轨道混音安全测试."""
|
||||
|
||||
def test_track_count_limit_exceeded(self, work_dir, sample_audio):
|
||||
"""超过最大轨道数时应截断到上限."""
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from video_processing.multi_track_mixer import (
|
||||
MAX_AUDIO_TRACKS,
|
||||
AudioTrack,
|
||||
MultiTrackMixConfig,
|
||||
mix_multi_track,
|
||||
)
|
||||
|
||||
# 创建超过上限的轨道数
|
||||
tracks = []
|
||||
for i in range(MAX_AUDIO_TRACKS + 5):
|
||||
tracks.append(
|
||||
AudioTrack(
|
||||
track_id=f"track_{i}",
|
||||
track_type="sfx",
|
||||
audio_path=str(sample_audio),
|
||||
volume=0.5,
|
||||
)
|
||||
)
|
||||
|
||||
config = MultiTrackMixConfig(tracks=tracks)
|
||||
|
||||
ctx = MagicMock()
|
||||
ctx.work_dir = work_dir
|
||||
ctx.plan_id = "test_plan"
|
||||
|
||||
# mock _prepare_single_track 避免实际跑ffmpeg
|
||||
with patch("video_processing.multi_track_mixer._prepare_single_track", return_value=True):
|
||||
with patch("video_processing.multi_track_mixer.run_ffmpeg"):
|
||||
import shutil
|
||||
|
||||
with patch("shutil.copy2"):
|
||||
result = mix_multi_track(ctx, sample_audio, config, 10.0)
|
||||
|
||||
# 验证轨道被截断到上限
|
||||
assert len(config.tracks) == MAX_AUDIO_TRACKS
|
||||
assert result is not None
|
||||
|
||||
def test_track_count_within_limit(self, work_dir, sample_audio):
|
||||
"""轨道数在限制内时正常处理."""
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from video_processing.multi_track_mixer import (
|
||||
MAX_AUDIO_TRACKS,
|
||||
AudioTrack,
|
||||
MultiTrackMixConfig,
|
||||
mix_multi_track,
|
||||
)
|
||||
|
||||
tracks = []
|
||||
for i in range(3):
|
||||
tracks.append(
|
||||
AudioTrack(
|
||||
track_id=f"track_{i}",
|
||||
track_type="sfx",
|
||||
audio_path=str(sample_audio),
|
||||
volume=0.5,
|
||||
)
|
||||
)
|
||||
|
||||
config = MultiTrackMixConfig(tracks=tracks)
|
||||
|
||||
ctx = MagicMock()
|
||||
ctx.work_dir = work_dir
|
||||
ctx.plan_id = "test_plan"
|
||||
|
||||
with patch("video_processing.multi_track_mixer._prepare_single_track", return_value=True):
|
||||
with patch("video_processing.multi_track_mixer.run_ffmpeg"):
|
||||
result = mix_multi_track(ctx, sample_audio, config, 10.0)
|
||||
|
||||
assert len(config.tracks) == 3
|
||||
assert result is not None
|
||||
|
||||
def test_audio_path_traversal_attack(self, work_dir, sample_audio):
|
||||
"""路径遍历攻击应被拦截."""
|
||||
from video_processing.multi_track_mixer import _validate_audio_path
|
||||
|
||||
# 路径遍历
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_audio_path("../../../etc/passwd", work_dir)
|
||||
|
||||
# local:// 路径遍历
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_audio_path("local://../../../etc/passwd", work_dir)
|
||||
|
||||
def test_audio_path_allowed_extension(self, work_dir, sample_audio):
|
||||
"""允许的音频扩展名应通过校验."""
|
||||
from video_processing.multi_track_mixer import _validate_audio_path
|
||||
|
||||
# 在work_dir内的音频文件
|
||||
test_file = work_dir / "test.mp3"
|
||||
test_file.touch()
|
||||
_validate_audio_path(str(test_file), work_dir) # 不应抛异常
|
||||
|
||||
test_file2 = work_dir / "test.wav"
|
||||
test_file2.touch()
|
||||
_validate_audio_path(str(test_file2), work_dir) # 不应抛异常
|
||||
|
||||
def test_audio_path_disallowed_extension(self, work_dir):
|
||||
"""不允许的文件扩展名应被拦截."""
|
||||
from video_processing.multi_track_mixer import _validate_audio_path
|
||||
|
||||
test_file = work_dir / "test.exe"
|
||||
test_file.touch()
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_audio_path(str(test_file), work_dir)
|
||||
|
||||
test_file2 = work_dir / "test.php"
|
||||
test_file2.touch()
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_audio_path(str(test_file2), work_dir)
|
||||
|
||||
def test_audio_path_empty(self, work_dir):
|
||||
"""空路径应被拦截."""
|
||||
from video_processing.multi_track_mixer import _validate_audio_path
|
||||
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_audio_path("", work_dir)
|
||||
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_audio_path(None, work_dir)
|
||||
|
||||
def test_audio_path_traversal_bypass_startswith(self, work_dir):
|
||||
"""【P1绕过】用../构造伪work_dir前缀路径,真实路径逃逸,必须被拦截.
|
||||
|
||||
漏洞:旧代码用 startswith(str(work_dir)) 比原始字符串,
|
||||
/tmp/work/../../opt/secret.aac 会通过 startswith 检查,跳过白名单校验。
|
||||
修复:用 realpath 规范化后再比较。
|
||||
"""
|
||||
from video_processing.multi_track_mixer import _validate_audio_path
|
||||
|
||||
evil_path = str(work_dir / "../../../../opt/secret.aac")
|
||||
with pytest.raises(PathSecurityError, match="不在允许目录"):
|
||||
_validate_audio_path(evil_path, work_dir)
|
||||
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
# 2. 视频拼接安全测试
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestConcatSecurity:
|
||||
"""视频拼接安全测试."""
|
||||
|
||||
def test_segment_count_limit_exceeded(self, work_dir, sample_video):
|
||||
"""超过最大段数时应报错."""
|
||||
from video_processing.concat_engine import (
|
||||
MAX_CONCAT_SEGMENTS,
|
||||
ConcatConfig,
|
||||
ConcatEngine,
|
||||
ConcatSegment,
|
||||
)
|
||||
|
||||
# 创建超过上限的段数
|
||||
segments = []
|
||||
for i in range(MAX_CONCAT_SEGMENTS + 5):
|
||||
segments.append(ConcatSegment(video_path=str(sample_video)))
|
||||
|
||||
config = ConcatConfig(segments=segments)
|
||||
engine = ConcatEngine(work_dir)
|
||||
output_path = work_dir / "output.mp4"
|
||||
|
||||
with pytest.raises(ValueError, match="Too many concat segments"):
|
||||
engine.concat_videos(config, output_path)
|
||||
|
||||
def test_segment_count_within_limit(self, work_dir, sample_video):
|
||||
"""段数在限制内时正常处理."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from video_processing.concat_engine import (
|
||||
MAX_CONCAT_SEGMENTS,
|
||||
ConcatConfig,
|
||||
ConcatEngine,
|
||||
ConcatSegment,
|
||||
)
|
||||
|
||||
segments = [
|
||||
ConcatSegment(video_path=str(sample_video)),
|
||||
ConcatSegment(video_path=str(sample_video)),
|
||||
ConcatSegment(video_path=str(sample_video)),
|
||||
]
|
||||
|
||||
config = ConcatConfig(segments=segments)
|
||||
engine = ConcatEngine(work_dir)
|
||||
output_path = work_dir / "output.mp4"
|
||||
|
||||
# mock ffmpeg执行
|
||||
with patch.object(engine, "_concat_filter", return_value=output_path):
|
||||
with patch.object(engine, "_can_use_stream_copy", return_value=False):
|
||||
result = engine.concat_videos(config, output_path)
|
||||
|
||||
assert result == output_path
|
||||
|
||||
def test_video_path_traversal_attack(self, work_dir):
|
||||
"""路径遍历攻击应被拦截."""
|
||||
from video_processing.concat_engine import _validate_video_path
|
||||
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_video_path("../../../etc/passwd", work_dir)
|
||||
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_video_path("local://../../../etc/passwd", work_dir)
|
||||
|
||||
def test_video_path_allowed_extension(self, work_dir):
|
||||
"""允许的视频扩展名应通过校验."""
|
||||
from video_processing.concat_engine import _validate_video_path
|
||||
|
||||
for ext in [".mp4", ".mov", ".avi", ".mkv", ".webm"]:
|
||||
test_file = work_dir / f"test{ext}"
|
||||
test_file.touch()
|
||||
_validate_video_path(str(test_file), work_dir) # 不应抛异常
|
||||
|
||||
def test_video_path_disallowed_extension(self, work_dir):
|
||||
"""不允许的文件扩展名应被拦截."""
|
||||
from video_processing.concat_engine import _validate_video_path
|
||||
|
||||
test_file = work_dir / "test.exe"
|
||||
test_file.touch()
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_video_path(str(test_file), work_dir)
|
||||
|
||||
test_file2 = work_dir / "test.js"
|
||||
test_file2.touch()
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_video_path(str(test_file2), work_dir)
|
||||
|
||||
def test_video_path_empty(self, work_dir):
|
||||
"""空路径应被拦截."""
|
||||
from video_processing.concat_engine import _validate_video_path
|
||||
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_video_path("", work_dir)
|
||||
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_video_path(None, work_dir)
|
||||
|
||||
def test_video_path_traversal_bypass_startswith(self, work_dir):
|
||||
"""【P1绕过】视频路径../遍历绕过startswith检查,必须被拦截.
|
||||
|
||||
漏洞:旧代码用 startswith(str(work_dir)) 比原始字符串,
|
||||
/tmp/work/../../opt/secret.mp4 会通过 startswith 检查,跳过白名单校验。
|
||||
修复:用 realpath 规范化后再比较。
|
||||
"""
|
||||
from video_processing.concat_engine import _validate_video_path
|
||||
|
||||
evil_path = str(work_dir / "../../../../opt/secret.mp4")
|
||||
with pytest.raises(PathSecurityError, match="不在允许目录"):
|
||||
_validate_video_path(evil_path, work_dir)
|
||||
|
||||
def test_invalid_segments_skipped(self, work_dir, sample_video):
|
||||
"""路径不安全的片段应被跳过."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from video_processing.concat_engine import (
|
||||
ConcatConfig,
|
||||
ConcatEngine,
|
||||
ConcatSegment,
|
||||
)
|
||||
|
||||
segments = [
|
||||
ConcatSegment(video_path=str(sample_video)),
|
||||
ConcatSegment(video_path="../../../etc/passwd"), # 不安全路径
|
||||
ConcatSegment(video_path=str(sample_video)),
|
||||
]
|
||||
|
||||
config = ConcatConfig(segments=segments)
|
||||
engine = ConcatEngine(work_dir)
|
||||
output_path = work_dir / "output.mp4"
|
||||
|
||||
with patch.object(engine, "_concat_filter", return_value=output_path):
|
||||
with patch.object(engine, "_can_use_stream_copy", return_value=False):
|
||||
result = engine.concat_videos(config, output_path)
|
||||
|
||||
# 验证只有2个安全片段保留
|
||||
assert len(config.segments) == 2
|
||||
assert result == output_path
|
||||
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
# 3. 字幕渲染安全测试
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TestSubtitleSecurity:
|
||||
"""字幕渲染安全测试."""
|
||||
|
||||
def test_subtitle_path_traversal_attack(self, work_dir):
|
||||
"""路径遍历攻击应被拦截."""
|
||||
from video_processing.subtitle_render_engine import _validate_subtitle_path
|
||||
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_subtitle_path("../../../etc/passwd", work_dir)
|
||||
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_subtitle_path("local://../../../etc/shadow", work_dir)
|
||||
|
||||
def test_subtitle_path_allowed_extension(self, work_dir):
|
||||
"""允许的字幕扩展名应通过校验."""
|
||||
from video_processing.subtitle_render_engine import _validate_subtitle_path
|
||||
|
||||
for ext in [".srt", ".ass", ".vtt", ".sub"]:
|
||||
test_file = work_dir / f"test{ext}"
|
||||
test_file.touch()
|
||||
_validate_subtitle_path(str(test_file), work_dir) # 不应抛异常
|
||||
|
||||
def test_subtitle_path_disallowed_extension(self, work_dir):
|
||||
"""不允许的文件扩展名应被拦截."""
|
||||
from video_processing.subtitle_render_engine import _validate_subtitle_path
|
||||
|
||||
test_file = work_dir / "test.exe"
|
||||
test_file.touch()
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_subtitle_path(str(test_file), work_dir)
|
||||
|
||||
test_file2 = work_dir / "test.mp4"
|
||||
test_file2.touch()
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_subtitle_path(str(test_file2), work_dir)
|
||||
|
||||
def test_subtitle_remote_url_blocked(self, work_dir):
|
||||
"""远程URL字幕应被拦截."""
|
||||
from video_processing.subtitle_render_engine import _validate_subtitle_path
|
||||
|
||||
with pytest.raises(PathSecurityError, match="远程URL"):
|
||||
_validate_subtitle_path("http://evil.com/evil.ass", work_dir)
|
||||
|
||||
with pytest.raises(PathSecurityError, match="远程URL"):
|
||||
_validate_subtitle_path("https://evil.com/evil.srt", work_dir)
|
||||
|
||||
def test_subtitle_path_empty(self, work_dir):
|
||||
"""空路径应被拦截."""
|
||||
from video_processing.subtitle_render_engine import _validate_subtitle_path
|
||||
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_subtitle_path("", work_dir)
|
||||
|
||||
with pytest.raises(PathSecurityError):
|
||||
_validate_subtitle_path(None, work_dir)
|
||||
|
||||
def test_build_filter_with_safe_path(self, work_dir):
|
||||
"""安全路径应正常生成滤镜字符串."""
|
||||
from video_processing.subtitle_render_engine import build_subtitle_filter
|
||||
|
||||
ass_file = work_dir / "subtitle.ass"
|
||||
ass_file.write_text("test", encoding="utf-8")
|
||||
|
||||
result = build_subtitle_filter(ass_file, work_dir=work_dir)
|
||||
assert "subtitles=" in result
|
||||
assert "subtitle.ass" in result
|
||||
assert "[subtitled]" in result
|
||||
|
||||
def test_build_filter_with_unsafe_path_raises(self, work_dir):
|
||||
"""不安全路径应抛出异常."""
|
||||
from video_processing.subtitle_render_engine import build_subtitle_filter
|
||||
|
||||
with pytest.raises(PathSecurityError):
|
||||
build_subtitle_filter("../../../etc/passwd", work_dir=work_dir)
|
||||
|
||||
def test_build_filter_work_dir_required(self, work_dir):
|
||||
"""不传work_dir时必须报错(防止自证清白绕过)."""
|
||||
from video_processing.subtitle_render_engine import build_subtitle_filter
|
||||
|
||||
ass_file = work_dir / "sub.ass"
|
||||
ass_file.write_text("test", encoding="utf-8")
|
||||
|
||||
# 不传 work_dir 必须报错
|
||||
with pytest.raises(PathSecurityError, match="work_dir"):
|
||||
build_subtitle_filter(ass_file) # type: ignore[call-arg]
|
||||
|
||||
# 传 None 也必须报错
|
||||
with pytest.raises(PathSecurityError, match="work_dir"):
|
||||
build_subtitle_filter(ass_file, work_dir=None) # type: ignore[arg-type]
|
||||
|
||||
# 传空字符串也必须报错
|
||||
with pytest.raises(PathSecurityError, match="work_dir"):
|
||||
build_subtitle_filter(ass_file, work_dir="")
|
||||
|
||||
def test_subtitle_path_traversal_bypass_startswith(self, work_dir):
|
||||
"""【P1绕过】字幕路径../遍历绕过startswith检查,必须被拦截."""
|
||||
from video_processing.subtitle_render_engine import _validate_subtitle_path
|
||||
|
||||
evil_path = str(work_dir / "../../../../opt/secret.srt")
|
||||
with pytest.raises(PathSecurityError, match="不在允许目录"):
|
||||
_validate_subtitle_path(evil_path, work_dir)
|
||||
Regular → Executable
+28
-41
@@ -6,6 +6,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
@@ -67,13 +68,10 @@ def _make_workflow(
|
||||
class TestTransferAudioToOSS:
|
||||
"""测试 _transfer_audio_to_oss 方法。"""
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_success_download_and_upload(self, mock_httpx: MagicMock) -> None:
|
||||
@patch("packages.application.tts_job.workflow.safe_download_bytes")
|
||||
def test_success_download_and_upload(self, mock_download: MagicMock) -> None:
|
||||
"""成功下载音频并上传到 OSS,返回永久 URL 和 storage_key。"""
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"fake audio data"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
mock_download.return_value = b"fake audio data"
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/tts-outputs/user_001/job_123.mp3"
|
||||
@@ -89,20 +87,21 @@ class TestTransferAudioToOSS:
|
||||
assert url == "https://oss.example.com/tts-outputs/user_001/job_123.mp3"
|
||||
assert key == "tts-outputs/user_001/job_123.mp3"
|
||||
|
||||
mock_httpx.get.assert_called_once_with(
|
||||
mock_download.assert_called_once_with(
|
||||
"https://cosyvoice-temp.com/audio.mp3",
|
||||
purpose="tts_audio_download",
|
||||
allowed_mime_types=unittest.mock.ANY,
|
||||
timeout=60.0,
|
||||
follow_redirects=True,
|
||||
)
|
||||
storage.upload_file.assert_called_once()
|
||||
call_args = storage.upload_file.call_args
|
||||
assert call_args[0][1] == "tts-outputs/user_001/job_123.mp3"
|
||||
assert call_args[1]["content_type"] == "audio/mpeg"
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_download_failure_fallback(self, mock_httpx: MagicMock) -> None:
|
||||
@patch("packages.application.tts_job.workflow.safe_download_bytes")
|
||||
def test_download_failure_fallback(self, mock_download: MagicMock) -> None:
|
||||
"""下载失败时回退到原始临时 URL,storage_key 为空。"""
|
||||
mock_httpx.get.side_effect = Exception("Network error")
|
||||
mock_download.side_effect = Exception("Network error")
|
||||
|
||||
workflow = _make_workflow()
|
||||
url, key = workflow._transfer_audio_to_oss(
|
||||
@@ -114,13 +113,10 @@ class TestTransferAudioToOSS:
|
||||
assert url == "https://cosyvoice-temp.com/audio.mp3"
|
||||
assert key == ""
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_upload_failure_fallback(self, mock_httpx: MagicMock) -> None:
|
||||
@patch("packages.application.tts_job.workflow.safe_download_bytes")
|
||||
def test_upload_failure_fallback(self, mock_download: MagicMock) -> None:
|
||||
"""上传 OSS 失败时回退到原始临时 URL。"""
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"fake audio data"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
mock_download.return_value = b"fake audio data"
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.side_effect = Exception("OSS bucket error")
|
||||
@@ -135,13 +131,10 @@ class TestTransferAudioToOSS:
|
||||
assert url == "https://cosyvoice-temp.com/audio.mp3"
|
||||
assert key == ""
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_wav_content_type(self, mock_httpx: MagicMock) -> None:
|
||||
@patch("packages.application.tts_job.workflow.safe_download_bytes")
|
||||
def test_wav_content_type(self, mock_download: MagicMock) -> None:
|
||||
"""wav 格式使用正确的 content_type。"""
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"fake wav data"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
mock_download.return_value = b"fake wav data"
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/audio.wav"
|
||||
@@ -162,13 +155,10 @@ class TestTransferAudioToOSS:
|
||||
class TestProcessSynthesisResultWithOSS:
|
||||
"""测试 process_synthesis_result 集成 OSS 转存。"""
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_stores_permanent_url_and_key(self, mock_httpx: MagicMock) -> None:
|
||||
@patch("packages.application.tts_job.workflow.safe_download_bytes")
|
||||
def test_stores_permanent_url_and_key(self, mock_download: MagicMock) -> None:
|
||||
"""合成结果存 OSS 永久 URL 和 storage_key。"""
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"audio bytes"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
mock_download.return_value = b"audio bytes"
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/tts-outputs/user_001/test_job_001.mp3"
|
||||
@@ -192,10 +182,10 @@ class TestProcessSynthesisResultWithOSS:
|
||||
assert result.duration == 5.0
|
||||
assert result.file_size == 50000
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_fallback_to_temp_url_on_oss_failure(self, mock_httpx: MagicMock) -> None:
|
||||
@patch("packages.application.tts_job.workflow.safe_download_bytes")
|
||||
def test_fallback_to_temp_url_on_oss_failure(self, mock_download: MagicMock) -> None:
|
||||
"""OSS 转存失败时,使用 CosyVoice 临时 URL(不阻塞合成流程)。"""
|
||||
mock_httpx.get.side_effect = Exception("Download failed")
|
||||
mock_download.side_effect = Exception("Download failed")
|
||||
|
||||
repo = MagicMock()
|
||||
job = _make_job(status=TTSJobStatus.PROCESSING)
|
||||
@@ -216,13 +206,10 @@ class TestProcessSynthesisResultWithOSS:
|
||||
class TestStartSynthesisSyncWithOSS:
|
||||
"""测试 start_synthesis 同步路径的 OSS 转存。"""
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_sync_path_transfers_to_oss(self, mock_httpx: MagicMock) -> None:
|
||||
@patch("packages.application.tts_job.workflow.safe_download_bytes")
|
||||
def test_sync_path_transfers_to_oss(self, mock_download: MagicMock) -> None:
|
||||
"""CosyVoice 同步返回 audio_url 时,也走 OSS 转存。"""
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"sync audio bytes"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
mock_download.return_value = b"sync audio bytes"
|
||||
|
||||
storage = MagicMock()
|
||||
storage.upload_file.return_value = "https://oss.example.com/tts-outputs/user_001/test_job_001.mp3"
|
||||
@@ -248,10 +235,10 @@ class TestStartSynthesisSyncWithOSS:
|
||||
assert job.output_audio_key == "tts-outputs/user_001/test_job_001.mp3"
|
||||
assert job.duration == 2.0
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_sync_path_oss_failure_stores_temp_url(self, mock_httpx: MagicMock) -> None:
|
||||
@patch("packages.application.tts_job.workflow.safe_download_bytes")
|
||||
def test_sync_path_oss_failure_stores_temp_url(self, mock_download: MagicMock) -> None:
|
||||
"""同步路径 OSS 失败时,降级存储临时 URL。"""
|
||||
mock_httpx.get.side_effect = Exception("Network error")
|
||||
mock_download.side_effect = Exception("Network error")
|
||||
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
service.submit_synthesize_task.return_value = {
|
||||
|
||||
Regular → Executable
+11
-23
@@ -233,14 +233,11 @@ class TestStartSegmentSynthesis:
|
||||
# 短文本走普通路径,不调用分段
|
||||
assert job.status == TTSJobStatus.PROCESSING
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_long_text_sync_segments(self, mock_httpx: MagicMock) -> None:
|
||||
@patch("packages.application.tts_job.workflow.safe_download_file")
|
||||
def test_long_text_sync_segments(self, mock_download: MagicMock) -> None:
|
||||
"""长文本同步分段:所有段立即返回 audio_url,直接合并。"""
|
||||
# Mock 分段音频下载
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"segment audio"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
mock_download.return_value = 1024 # 模拟文件大小
|
||||
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
# 每个分段都同步返回 audio_url
|
||||
@@ -355,14 +352,11 @@ class TestHandleSegmentFailure:
|
||||
class TestPollSegmentTasks:
|
||||
"""测试 _poll_segment_tasks 分段缺失重新合成(适配同步接口)。"""
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_all_segments_done(self, mock_httpx: MagicMock) -> None:
|
||||
@patch("packages.application.tts_job.workflow.safe_download_file")
|
||||
def test_all_segments_done(self, mock_download: MagicMock) -> None:
|
||||
"""所有分段缺少 audio_url 时重新同步合成,合并后标记完成。"""
|
||||
# Mock 下载分段音频
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"seg audio"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
mock_download.return_value = 1024 # 模拟文件大小
|
||||
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
service.submit_synthesize_task.side_effect = [
|
||||
@@ -422,13 +416,10 @@ class TestPollSegmentTasks:
|
||||
|
||||
assert result.status == TTSJobStatus.FAILED
|
||||
|
||||
@patch("packages.application.tts_job.workflow.httpx")
|
||||
def test_partial_audio_urls_reuse_existing(self, mock_httpx: MagicMock) -> None:
|
||||
@patch("packages.application.tts_job.workflow.safe_download_file")
|
||||
def test_partial_audio_urls_reuse_existing(self, mock_download: MagicMock) -> None:
|
||||
"""部分分段已有 audio_url 时直接复用,缺失的重新合成。"""
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"seg audio"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
mock_download.return_value = 1024 # 模拟文件大小
|
||||
|
||||
service = MagicMock(spec=CosyVoiceService)
|
||||
# 只有 1 个分段需要重新合成
|
||||
@@ -515,11 +506,8 @@ class TestPollAndProcessSynthesisSegmentDetection:
|
||||
|
||||
workflow = _make_workflow(cosyvoice_service=service, repo=repo, storage=storage)
|
||||
|
||||
with patch("packages.application.tts_job.workflow.httpx") as mock_httpx:
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"audio"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
with patch("packages.application.tts_job.workflow.safe_download_file") as mock_download:
|
||||
mock_download.return_value = 1024 # 模拟文件大小
|
||||
|
||||
result = workflow.poll_and_process_synthesis("test_job_seg")
|
||||
|
||||
|
||||
Regular → Executable
+4
-7
@@ -221,17 +221,14 @@ class TestTTSStreamingService:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_download_audio(self):
|
||||
"""下载音频数据。"""
|
||||
"""下载音频数据(SSRF防护走safe_download_bytes,mock掉安全层)。"""
|
||||
cosyvoice = MagicMock(spec=CosyVoiceService)
|
||||
service = TTSStreamingService(cosyvoice)
|
||||
|
||||
with patch("packages.application.tts_job.streaming_service.httpx") as mock_httpx:
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.content = b"audio data"
|
||||
mock_resp.raise_for_status.return_value = None
|
||||
mock_httpx.get.return_value = mock_resp
|
||||
with patch("packages.application.tts_job.streaming_service.safe_download_bytes") as mock_download:
|
||||
mock_download.return_value = b"audio data"
|
||||
|
||||
result = service._download_audio("https://example.com/audio.mp3")
|
||||
|
||||
assert result == b"audio data"
|
||||
mock_httpx.get.assert_called_once()
|
||||
mock_download.assert_called_once()
|
||||
|
||||
Executable
+296
@@ -0,0 +1,296 @@
|
||||
"""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"))
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
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()
|
||||
|
||||
|
||||
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)
|
||||
Reference in New Issue
Block a user