diff --git a/apps/worker/video_processing/ffmpeg_utils.py b/apps/worker/video_processing/ffmpeg_utils.py index d40cc2d1f..10ee0e752 100755 --- a/apps/worker/video_processing/ffmpeg_utils.py +++ b/apps/worker/video_processing/ffmpeg_utils.py @@ -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: """探测文件是否包含音频流。 diff --git a/apps/worker/video_processing/oss_helpers.py b/apps/worker/video_processing/oss_helpers.py index 78944886d..3e486ed51 100755 --- a/apps/worker/video_processing/oss_helpers.py +++ b/apps/worker/video_processing/oss_helpers.py @@ -220,25 +220,64 @@ def resolve_asset_path(asset_id: str, work_dir: Path) -> Path | None: """从 asset_id 解析到本地文件路径。 策略(按优先级): - 1. 如果 asset_id 是本地绝对路径(/var/storage/...)→ 直接返回 + 1. 如果 asset_id 是本地绝对路径(/var/storage/...)→ 安全校验后返回 2. 如果 work_dir 下已有缓存文件 → 返回缓存路径 3. 从 OSS 下载到 work_dir/{hash}.mp4 → 返回下载路径 4. 下载失败 → 返回 None 缓存策略:以 asset_id 的 SHA256 前 16 位为文件名,避免重复下载。 - """ - # 1. 本地绝对路径 - if asset_id.startswith("/") and os.path.exists(asset_id): - return Path(asset_id) - # 2. 缓存命中 + 安全: + - 本地绝对路径必须在 ASSET_ALLOWED_DIRS 环境变量指定的目录内 + - 文件名经过 sanitize,防止路径遍历 + - 禁止空字节、控制字符 + """ + from video_processing.path_security import ( + PathSecurityError, + get_allowed_local_dirs, + is_in_allowed_dirs, + sanitize_filename, + ) + + if not asset_id or not isinstance(asset_id, str): + return None + + # 空字节检测 + if "\x00" in asset_id: + logger.warning("asset_id 包含空字节,拒绝: %s", asset_id[:50]) + return None + + # 1. 本地绝对路径 — 必须在允许的目录内 + if asset_id.startswith("/") and os.path.exists(asset_id): + try: + resolved = Path(asset_id).resolve() + if is_in_allowed_dirs(resolved, get_allowed_local_dirs()): + return resolved + else: + logger.warning( + "本地素材路径不在允许目录内,拒绝: %s (allowed=%s)", + asset_id[:80], + get_allowed_local_dirs(), + ) + return None + except (OSError, PathSecurityError): + return None + + # 2. 缓存命中(使用 hash 而非原始 ID,防止路径遍历) cache_hash = hashlib.sha256(asset_id.encode()).hexdigest()[:16] - cached_path = work_dir / f"{cache_hash}.mp4" + safe_name = sanitize_filename(cache_hash) + cached_path = work_dir / f"{safe_name}.mp4" if cached_path.exists() and cached_path.stat().st_size > 0: return cached_path - # 3. 从 OSS 下载 - if download_asset(asset_id, cached_path): + # 3. 从 OSS 下载(先标准化 key,防止路径遍历注入) + safe_key = normalize_storage_key(asset_id) + # 额外校验:存储键不能包含 ../ 或绝对路径 + if ".." in safe_key or safe_key.startswith("/"): + logger.warning("asset_id 包含路径遍历模式,拒绝下载: %s", asset_id[:80]) + return None + + if download_asset(safe_key, cached_path): return cached_path return None diff --git a/apps/worker/video_processing/pip_engine.py b/apps/worker/video_processing/pip_engine.py index cb84550e3..59c26d6ef 100755 --- a/apps/worker/video_processing/pip_engine.py +++ b/apps/worker/video_processing/pip_engine.py @@ -465,18 +465,44 @@ class PiPEngine: layer: PiPLayerConfig, asset_path_map: dict[str, Path], ) -> Path | None: - """验证图层素材是否可用,返回本地路径或None(降级跳过).""" + """验证图层素材是否可用,返回本地路径或None(降级跳过). + + 安全: + - local_path 类型:必须在允许的目录内,防止路径遍历 + - url 类型:必须通过 SSRF 安全校验 + """ + from video_processing.path_security import is_in_allowed_dirs + from video_processing.url_security import UrlSecurityError, validate_url_safety + try: if layer.source_type == "local_path": - path = Path(layer.source) - if path.exists(): - return path + if not layer.source: + return None + # 路径安全校验:必须在允许目录内 + src_path = Path(layer.source) + if not src_path.exists(): + return None + if not is_in_allowed_dirs(src_path): + logger.warning( + "PiP local_path 不在允许目录内,拒绝: %s", + layer.source[:80], + ) + return None + return src_path.resolve() elif layer.source_type == "asset_id": if layer.source in asset_path_map: return asset_path_map[layer.source] + return None elif layer.source_type == "url": - # URL类型由调用者负责下载,这里返回标记 - return None # 暂时不支持直接URL + # URL类型:先做SSRF安全校验,由调用者负责实际下载 + try: + validate_url_safety(layer.source, purpose="pip_source") + logger.info("PiP URL 安全校验通过: %s", layer.source[:80]) + except UrlSecurityError as e: + logger.warning("PiP URL 安全校验失败: %s (error=%s)", layer.source[:80], e) + return None + # 暂时不支持直接URL下载,返回None表示降级跳过 + return None except Exception as e: logger.warning("PiP素材验证失败: %s", e) diff --git a/apps/worker/video_processing/sticker_engine.py b/apps/worker/video_processing/sticker_engine.py index 0cb4068c9..afbdc996e 100755 --- a/apps/worker/video_processing/sticker_engine.py +++ b/apps/worker/video_processing/sticker_engine.py @@ -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) diff --git a/apps/worker/video_processing/unified_render_service.py b/apps/worker/video_processing/unified_render_service.py index 653ef594e..76b751052 100644 --- a/apps/worker/video_processing/unified_render_service.py +++ b/apps/worker/video_processing/unified_render_service.py @@ -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, diff --git a/apps/worker/video_processing/url_security.py b/apps/worker/video_processing/url_security.py new file mode 100644 index 000000000..3ea89d178 --- /dev/null +++ b/apps/worker/video_processing/url_security.py @@ -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, +) diff --git a/apps/worker/worker_app/tasks/asset_analyzer.py b/apps/worker/worker_app/tasks/asset_analyzer.py index ec609aa5b..8c3035a3e 100755 --- a/apps/worker/worker_app/tasks/asset_analyzer.py +++ b/apps/worker/worker_app/tasks/asset_analyzer.py @@ -130,6 +130,8 @@ class AssetAnalyzer: info = VideoInfo() try: + from video_processing.ffmpeg_utils import run_ffprobe + cmd = [ "ffprobe", "-v", @@ -140,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 diff --git a/apps/worker/worker_app/tasks/batch_download.py b/apps/worker/worker_app/tasks/batch_download.py index 3924f8d8b..404ce2766 100755 --- a/apps/worker/worker_app/tasks/batch_download.py +++ b/apps/worker/worker_app/tasks/batch_download.py @@ -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, + ) diff --git a/apps/worker/worker_app/tasks/compose_video.py b/apps/worker/worker_app/tasks/compose_video.py index 4150e0fc2..4509d5e02 100755 --- a/apps/worker/worker_app/tasks/compose_video.py +++ b/apps/worker/worker_app/tasks/compose_video.py @@ -6,7 +6,6 @@ from __future__ import annotations import os -import subprocess import tempfile from pathlib import Path @@ -115,16 +114,12 @@ def _compose_with_legacy_engine(task, job_service, job, plan_id: str, db) -> dic logger.info("Executing FFmpeg for job %s, plan %s", job_id, plan_id) try: - subprocess.run( - compose_cmd.command, - check=True, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - text=True, - timeout=3600, - ) - except subprocess.CalledProcessError as e: - job_service.fail_job(job_id, f"FFmpeg 执行失败: {e.stderr[:500]}") + from video_processing.ffmpeg_utils import run_ffmpeg + + run_ffmpeg(compose_cmd.command, timeout=3600) + except Exception as e: + error_msg = f"FFmpeg 执行失败: {str(e)[:500]}" + job_service.fail_job(job_id, error_msg) raise # 上传结果 diff --git a/apps/worker/worker_app/tasks/edit_plan_generation.py b/apps/worker/worker_app/tasks/edit_plan_generation.py index 5c789d99b..9b422ab57 100755 --- a/apps/worker/worker_app/tasks/edit_plan_generation.py +++ b/apps/worker/worker_app/tasks/edit_plan_generation.py @@ -257,7 +257,6 @@ def _render_with_legacy( ) -> dict: """旧引擎路径(VideoComposeService + FFmpeg filter_complex)。""" import os - import subprocess from apps.api.app.services.video_compose_service import VideoComposeService @@ -278,16 +277,11 @@ def _render_with_legacy( logger.info("执行 FFmpeg (legacy): plan_id=%s", plan_id) try: - subprocess.run( - compose_cmd.command, - check=True, - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - text=True, - timeout=3600, - ) - except subprocess.CalledProcessError as e: - error_msg = f"FFmpeg 执行失败: {e.stderr[:500]}" + from video_processing.ffmpeg_utils import run_ffmpeg + + run_ffmpeg(compose_cmd.command, timeout=3600) + except Exception as e: + error_msg = f"FFmpeg 执行失败: {str(e)[:500]}" logger.error("FFmpeg 执行失败(legacy): %s — %s", plan_id, error_msg) _mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, error_msg) return {"status": "error", "message": error_msg} diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 8db0609a7..1af503185 100644 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -364,10 +364,19 @@ def _prepare_bgm_track( try: parsed = urlparse(audio_url) if parsed.scheme in ("http", "https"): - import urllib.request + from video_processing.url_security import ( + ALLOWED_AUDIO_MIME_TYPES, + safe_download_file, + ) logger.info("[task_id=%s] [BGM] 从URL下载: %s", task_id, audio_url[:80]) - urllib.request.urlretrieve(audio_url, bgm_file) # nosec B310 + safe_download_file( + audio_url, + str(bgm_file), + purpose="bgm_download", + allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES, + timeout=60.0, + ) if bgm_file.exists() and bgm_file.stat().st_size > 0: return str(bgm_file) except Exception as e: @@ -401,10 +410,19 @@ def _prepare_bgm_track( preset = get_preset_bgm(preset_id) if preset and preset.audio_url: - import urllib.request + from video_processing.url_security import ( + ALLOWED_AUDIO_MIME_TYPES, + safe_download_file, + ) logger.info("[task_id=%s] [BGM] 从预设库下载: preset_id=%s", task_id, preset_id) - urllib.request.urlretrieve(preset.audio_url, bgm_file) # nosec B310 + safe_download_file( + preset.audio_url, + str(bgm_file), + purpose="bgm_preset_download", + allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES, + timeout=60.0, + ) if bgm_file.exists() and bgm_file.stat().st_size > 0: return str(bgm_file) except Exception as e: @@ -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: diff --git a/apps/worker/worker_app/tasks/ingest.py b/apps/worker/worker_app/tasks/ingest.py index 0261c6e8c..e701198b9 100755 --- a/apps/worker/worker_app/tasks/ingest.py +++ b/apps/worker/worker_app/tasks/ingest.py @@ -45,6 +45,8 @@ def extract_media_metadata(file_url: str, media_type: str) -> dict: try: if media_type == "video": # 使用 ffprobe 提取视频元数据 + from video_processing.ffmpeg_utils import run_ffprobe + cmd = [ "ffprobe", "-v", @@ -55,16 +57,11 @@ def extract_media_metadata(file_url: str, media_type: str) -> dict: "-show_streams", file_url, ] - result = subprocess.run( - cmd, - capture_output=True, - text=True, - timeout=30, - ) - if result.returncode == 0: + try: + stdout, _ = run_ffprobe(cmd, timeout=30) import json as json_lib - probe_data = json_lib.loads(result.stdout) + probe_data = json_lib.loads(stdout) # 提取视频流信息 for stream in probe_data.get("streams", []): @@ -83,6 +80,9 @@ def extract_media_metadata(file_url: str, media_type: str) -> dict: metadata["size_bytes"] = int(format_info.get("size", 0)) metadata["bitrate"] = int(format_info.get("bit_rate", 0)) + except Exception as e: + logger.warning("视频元数据提取失败: %s", e) + elif media_type == "image": # 使用 Pillow 提取图片元数据 try: diff --git a/apps/worker/worker_app/tasks/voice_extraction.py b/apps/worker/worker_app/tasks/voice_extraction.py index b08335809..efbc652a0 100644 --- a/apps/worker/worker_app/tasks/voice_extraction.py +++ b/apps/worker/worker_app/tasks/voice_extraction.py @@ -19,14 +19,12 @@ class VoiceExtractor: """Extract voice tracks and background music from videos using FFmpeg.""" @staticmethod - def _run_ffmpeg(cmd: list[str]) -> subprocess.CompletedProcess: - """Run FFmpeg command and return result.""" - logger.info(f"Running FFmpeg: {chr(39).join(cmd)}") - result = subprocess.run(cmd, capture_output=True, text=True) - if result.returncode != 0: - logger.error(f"FFmpeg error: {result.stderr}") - raise RuntimeError(f"FFmpeg failed: {result.stderr}") - return result + def _run_ffmpeg(cmd: list[str]) -> None: + """Run FFmpeg command using 统一 run_ffmpeg 工具.""" + from video_processing.ffmpeg_utils import run_ffmpeg + + logger.info("Running FFmpeg: %s", " ".join(cmd[:10])) + run_ffmpeg(cmd) def extract_voice( self, diff --git a/packages/application/tts_job/streaming_service.py b/packages/application/tts_job/streaming_service.py index ffb4674f3..d509cd5d5 100644 --- a/packages/application/tts_job/streaming_service.py +++ b/packages/application/tts_job/streaming_service.py @@ -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 推送。 diff --git a/packages/application/tts_job/workflow.py b/packages/application/tts_job/workflow.py old mode 100755 new mode 100644 index 4b9a472cc..0415a60cc --- a/packages/application/tts_job/workflow.py +++ b/packages/application/tts_job/workflow.py @@ -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) # 合并 diff --git a/tests/unit/test_path_security.py b/tests/unit/test_path_security.py new file mode 100755 index 000000000..b388130c4 --- /dev/null +++ b/tests/unit/test_path_security.py @@ -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:"test|?*.mp4') + self.assertNotIn("<", result) + self.assertNotIn(">", result) + self.assertNotIn(":", result) + self.assertNotIn('"', result) + self.assertNotIn("|", result) + self.assertNotIn("?", result) + self.assertNotIn("*", result) + + def test_chinese_filename_preserved(self): + """中文文件名应该保留.""" + result = sanitize_filename("视频素材.mp4") + self.assertIn("视频素材", result) + + def test_long_filename_truncated(self): + """超长文件名应该被截断.""" + long_name = "a" * 300 + ".mp4" + result = sanitize_filename(long_name) + self.assertLessEqual(len(result), 255) + self.assertTrue(result.endswith(".mp4")) + + +class TestAllowedDirs(unittest.TestCase): + """允许目录配置测试.""" + + def test_get_allowed_dirs_returns_list(self): + """get_allowed_local_dirs 应该返回列表.""" + dirs = get_allowed_local_dirs() + self.assertIsInstance(dirs, list) + + def test_is_in_allowed_dirs_tmp(self): + """/tmp 应该在默认允许目录内.""" + self.assertTrue(is_in_allowed_dirs("/tmp/test.mp4")) + + def test_is_path_safe_convenience(self): + """is_path_safe 便捷函数应该正常工作.""" + with tempfile.TemporaryDirectory() as tmpdir: + self.assertTrue(is_path_safe("test.mp4", tmpdir)) + self.assertFalse(is_path_safe("../etc/passwd", tmpdir)) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit/test_tts_oss_transfer.py b/tests/unit/test_tts_oss_transfer.py old mode 100644 new mode 100755 index 99d5fc292..63203c2d4 --- a/tests/unit/test_tts_oss_transfer.py +++ b/tests/unit/test_tts_oss_transfer.py @@ -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 = { diff --git a/tests/unit/test_tts_segment_synthesis.py b/tests/unit/test_tts_segment_synthesis.py old mode 100644 new mode 100755 index 583b3a2ba..6e9a4c2c3 --- a/tests/unit/test_tts_segment_synthesis.py +++ b/tests/unit/test_tts_segment_synthesis.py @@ -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") diff --git a/tests/unit/test_tts_streaming.py b/tests/unit/test_tts_streaming.py old mode 100644 new mode 100755 index 8b2bd29ca..1f587aedd --- a/tests/unit/test_tts_streaming.py +++ b/tests/unit/test_tts_streaming.py @@ -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() diff --git a/tests/unit/test_url_security.py b/tests/unit/test_url_security.py new file mode 100755 index 000000000..2c7ac13d6 --- /dev/null +++ b/tests/unit/test_url_security.py @@ -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)