diff --git a/apps/worker/video_processing/ffmpeg_utils.py b/apps/worker/video_processing/ffmpeg_utils.py index 3fd7b6ac0..f1a68377d 100644 --- a/apps/worker/video_processing/ffmpeg_utils.py +++ b/apps/worker/video_processing/ffmpeg_utils.py @@ -58,16 +58,28 @@ def run_ffmpeg( (stdout, stderr) 元组 Raises: - subprocess.CalledProcessError: 命令执行失败时抛出 + subprocess.CalledProcessError: 命令执行失败时抛出, + 异常信息包含完整 stderr 以便排查。 """ - 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, - ) - return (result.stdout or "", result.stderr or "") + 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, + ) + return (result.stdout or "", result.stderr or "") + except subprocess.CalledProcessError as e: + # 把完整 stderr 打到日志,方便排查 exit code 183 等问题 + stderr_text = (e.stderr or "").strip() + logger.error( + "FFmpeg 命令失败: exit_code=%d command=%s\nstderr:\n%s", + e.returncode, + " ".join(str(c) for c in command[:20]), # 截断过长的命令 + stderr_text[:5000], # 截断过长的 stderr + ) + raise def probe_duration(local_path: str | Path) -> float: diff --git a/apps/worker/video_processing/unified_render_service.py b/apps/worker/video_processing/unified_render_service.py index 6b6a15287..8852ceb67 100644 --- a/apps/worker/video_processing/unified_render_service.py +++ b/apps/worker/video_processing/unified_render_service.py @@ -23,6 +23,7 @@ from __future__ import annotations import logging import os +import subprocess from dataclasses import dataclass, field from pathlib import Path from typing import Any @@ -417,7 +418,10 @@ class UnifiedRenderService: input_args: list[str], output_path: Path, ) -> None: - """执行 FFmpeg 渲染命令。""" + """执行 FFmpeg 渲染命令。 + + 失败时记录完整 filter_complex 以便排查(如 exit code 183)。 + """ command = [ FFMPEG_BIN, "-y", @@ -445,7 +449,17 @@ class UnifiedRenderService: input_args.count("-i"), output_path, ) - run_ffmpeg(command) + try: + run_ffmpeg(command) + except subprocess.CalledProcessError as e: + # 额外记录 filter_complex,方便排查滤镜链构建问题 + logger.error( + "渲染失败: plan_id=%s exit_code=%d\nfilter_complex:\n%s", + self.plan.id, + e.returncode, + filter_complex[:5000], + ) + raise def _probe_output(self, output_path: Path) -> tuple[float, int, int, int]: """探测输出文件的时长、大小、宽高。 diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 9a1c1e3ee..1e6716aa1 100755 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -28,8 +28,6 @@ OUTPUT_HEIGHT = 720 OUTPUT_FPS = 25.0 OUTPUT_DURATION_SECONDS = 5.0 GENERATED_FILES_DIR = Path(os.getenv("GENERATED_FILES_DIR", "/app/generated")) -GENERATED_FILES_URL_PREFIX = os.getenv("GENERATED_FILES_URL_PREFIX", "/generated-files") -PUBLIC_API_BASE_URL = os.getenv("PUBLIC_API_BASE_URL", "https://api.xiaoxiajianji.com").rstrip("/") logger = logging.getLogger(__name__) @@ -89,7 +87,6 @@ from video_processing.dedup_helpers import create_video_record_and_dedup from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_duration, run_ffmpeg from video_processing.oss_helpers import ( download_asset, - oss_bucket, upload_to_oss, ) from video_processing.unified_render_service import UnifiedRenderService @@ -274,45 +271,156 @@ def _download_voice_asset(voice_library_id: str, local_path: Path) -> bool: return download_asset(storage_key, local_path) +def _verify_url_accessible(url: str, timeout: float = 10.0, retries: int = 2) -> bool: + """HEAD 请求校验 URL 可访问(含重试,防止 OSS 抖动误报)。 + + Args: + url: 待校验的 URL + timeout: 单次请求超时时间(秒) + retries: 最大重试次数(默认 2 次,首次失败后间隔 1s 重试) + + Returns: + True 表示 URL 可访问(HTTP 2xx/3xx),False 表示所有尝试均失败。 + """ + import time + import urllib.request + + last_error: Exception | None = None + for attempt in range(1 + retries): + try: + req = urllib.request.Request(url, method="HEAD") + req.add_header("User-Agent", "xiaoxia-saas-worker/1.0") + with urllib.request.urlopen(req, timeout=timeout) as resp: # nosec B310 + if resp.status < 400: + return True + last_error = Exception(f"HTTP {resp.status}") + except Exception as e: + last_error = e + + if attempt < retries: + logger.warning( + "URL 校验失败,1s 后重试: url=%s attempt=%d/%d error=%s", + url, + attempt + 1, + retries, + last_error, + ) + time.sleep(1) + + logger.warning("URL 可访问性校验最终失败: url=%s error=%s", url, last_error) + return False + + def _download_library_assets( - asset_library_id: str, temp_path: Path, + *, + asset_library_id: str = "", + project_id: str = "", + asset_ids: list[str] | None = None, video_extensions: tuple = (".mp4", ".mov", ".avi", ".mkv", ".webm"), - asset_ids: list[str] | None = None, + strict: bool = True, ) -> list[Path]: - """从素材库下载视频素材。 + """下载视频素材 — 同时支持素材库模式和项目级模式。 + + 两种查询路径: + - 素材库模式:asset_library_id 非空时,按 asset_library_id + asset_ids 查 + - 项目级模式:project_id 非空时,按 project_id + asset_ids 查 + - 两者都提供时优先素材库模式;两者都为空时抛异常 + + 归属校验与下载在同一 DB session 中完成,避免多次连接开销(P3-2)。 Args: - asset_library_id: 素材库 ID temp_path: 临时目录路径 - video_extensions: 支持的视频扩展名 + asset_library_id: 素材库 ID(可选,与 project_id 二选一) + project_id: 项目 ID(可选,与 asset_library_id 二选一) asset_ids: 指定素材 ID 列表,为空则下载全部 ready 视频素材 + video_extensions: 支持的视频扩展名(保留兼容,当前按 file_type 过滤) + strict: 严格模式(默认 True)。 + True — 任何素材下载失败立即抛 RuntimeError; + False — 跳过失败素材,返回成功列表(调用方可通过日志感知失败)。 Returns: 下载成功的视频文件 Path 列表 + + Raises: + ValueError: 当 asset_library_id 和 project_id 都为空时 + RuntimeError: strict=True 时任何下载失败;或指定了 asset_ids 但全部下载失败 """ + if not asset_library_id and not project_id: + raise ValueError("asset_library_id 和 project_id 至少需要提供一个") + try: from packages.adapters.sqlalchemy_impl.models import AssetModel session = SessionLocal() try: + # 构建查询:根据模式选择不同的过滤条件 query = session.query(AssetModel).filter( - AssetModel.asset_library_id == asset_library_id, AssetModel.status == "ready", AssetModel.file_type.in_(["video", "video/mp4", "video/quicktime"]), ) + + if asset_library_id: + # 素材库模式 + query = query.filter(AssetModel.asset_library_id == asset_library_id) + logger.info( + "下载素材库视频: asset_library_id=%s asset_ids=%s", + asset_library_id, + asset_ids or "all", + ) + else: + # 项目级模式 + query = query.filter(AssetModel.project_id == project_id) + logger.info( + "下载项目级视频: project_id=%s asset_ids=%s", + project_id, + asset_ids or "all", + ) + if asset_ids: query = query.filter(AssetModel.id.in_(asset_ids)) + assets = query.order_by(AssetModel.created_at).all() if not assets: - logger.info("No video assets found in library %s", asset_library_id) - return [] + mode_desc = f"素材库 {asset_library_id}" if asset_library_id else f"项目 {project_id}" + msg = f"未找到视频素材: {mode_desc}, asset_ids={asset_ids or 'all'}" + logger.error(msg) + raise RuntimeError(msg) + + # P3-2: 归属校验合并到同一 session + if asset_ids: + found_ids = {a.id for a in assets} + missing_ids = set(asset_ids) - found_ids + if missing_ids: + raise ValueError(f"素材不存在: asset_ids={sorted(missing_ids)}") + for asset in assets: + if asset_library_id and asset.asset_library_id != asset_library_id: + raise ValueError( + f"素材不属于指定素材库: asset_id={asset.id}, " + f"expected_asset_library_id={asset_library_id}, " + f"actual_asset_library_id={asset.asset_library_id}" + ) + if not asset_library_id and project_id and asset.project_id != project_id: + raise ValueError( + f"素材不属于指定项目: asset_id={asset.id}, " + f"expected_project_id={project_id}, " + f"actual_project_id={asset.project_id}" + ) + logger.info( + "素材归属校验通过 (同 session): %d 个 asset_ids", + len(asset_ids), + ) downloaded: list[Path] = [] + failed_assets: list[str] = [] for i, asset in enumerate(assets): storage_key = asset.file_url if asset.file_url else None if not storage_key: + failed_assets.append(f"{asset.name}({asset.id})") + logger.warning("素材缺少 file_url, 跳过: asset_id=%s name=%s", asset.id, asset.name) + if strict: + raise RuntimeError(f"素材缺少 file_url: asset_id={asset.id}, name={asset.name}") continue ext = Path(storage_key).suffix or ".mp4" @@ -321,14 +429,61 @@ def _download_library_assets( downloaded.append(local_file) logger.info("Downloaded asset: %s -> %s", asset.name, local_file) else: - logger.warning("Failed to download asset: %s", asset.name) + failed_assets.append(f"{asset.name}({asset.id})") + logger.warning("Failed to download asset: %s (id=%s)", asset.name, asset.id) + if strict: + raise RuntimeError(f"素材下载失败: asset_id={asset.id}, name={asset.name}") + + # 指定了 asset_ids 但全部下载失败 → 无论 strict 与否都报错 + if asset_ids and not downloaded: + msg = f"指定的 {len(asset_ids)} 个素材全部下载失败, failed={failed_assets}" + logger.error(msg) + raise RuntimeError(msg) + + # 非严格模式有部分失败,记录警告 + if failed_assets and not strict: + logger.warning( + "素材下载部分失败 (非严格模式): failed=%s, succeeded=%d", + failed_assets, + len(downloaded), + ) return downloaded finally: session.close() + except (ValueError, RuntimeError): + raise except Exception as e: - logger.error("Error downloading library assets: %s", e) - return [] + logger.error("Error downloading library assets: %s", e, exc_info=True) + raise RuntimeError(f"素材下载异常: {e}") from e + + +# ── P1 校验函数 ────────────────────────────────────────────────────────────── + + +def _validate_template_exists(template_id: str) -> None: + """校验 template_id 是否存在且可用。 + + Raises: + ValueError: template_id 不存在或已禁用时抛出 + """ + from packages.adapters.sqlalchemy_impl.models import TemplateModel + + session = SessionLocal() + try: + template = ( + session.query(TemplateModel) + .filter( + TemplateModel.id == template_id, + TemplateModel.is_active.is_(True), + ) + .first() + ) + if template is None: + raise ValueError(f"模板不存在或已禁用: template_id={template_id}") + logger.info("模板校验通过: template_id=%s name=%s", template_id, template.name) + finally: + session.close() # ── Celery Task ────────────────────────────────────────────────────────────── @@ -372,6 +527,7 @@ def generate_video(self, task_id: str) -> dict: project_id = gen_task.project_id asset_library_id = gen_task.asset_library_id voice_library_id = gen_task.voice_library_id or "" + template_id = getattr(gen_task, "template_id", "") or "" mode = gen_task.strategy_id or "one_take" task_asset_ids = list(gen_task.asset_ids or []) batch_id = getattr(gen_task, "batch_id", "") or "" @@ -390,12 +546,23 @@ def generate_video(self, task_id: str) -> dict: storage_key = f"generated/projects/{project_id}/tasks/{task_id}/{output_name}" try: + # P1: template_id 存在性校验 + if template_id: + _validate_template_exists(template_id) + + # P1: asset_ids 归属校验 — 已合并到 _download_library_assets 同一 session(P3-2) + with tempfile.TemporaryDirectory(prefix="xiaoxia-generation-") as temp_dir: temp_path = Path(temp_dir) output_path = temp_path / output_name - # 1. 从素材库下载视频素材 - downloaded_videos = _download_library_assets(asset_library_id, temp_path, asset_ids=task_asset_ids or None) + # 1. 从素材库/项目下载视频素材 + downloaded_videos = _download_library_assets( + temp_path, + asset_library_id=asset_library_id, + project_id=project_id, + asset_ids=task_asset_ids or None, + ) # 2. 下载配音(如有) audio_path: str | None = None @@ -405,58 +572,62 @@ def generate_video(self, task_id: str) -> dict: audio_path = str(local_audio) # 3. 渲染 - if downloaded_videos: - # 构建虚拟 plan + clips + asset_path_map - virtual_plan, virtual_clips, asset_path_map = _build_plan_and_clips_from_task( - task_id=task_id, - downloaded_paths=downloaded_videos, - mode=editing_mode.value, + if not downloaded_videos: + # 素材下载为空(不应到达此处,_download_library_assets 已做校验) + raise RuntimeError( + f"素材下载结果为空: task_id={task_id}, " + f"asset_library_id={asset_library_id}, project_id={project_id}, " + f"asset_ids={task_asset_ids}" ) - # 使用 UnifiedRenderService 渲染 - render_service = UnifiedRenderService( - plan=virtual_plan, - clips=virtual_clips, - asset_path_map=asset_path_map, - work_dir=temp_path, - output_width=OUTPUT_WIDTH, - output_height=OUTPUT_HEIGHT, - output_fps=int(OUTPUT_FPS), - ) - render_result = render_service.render() + # 构建虚拟 plan + clips + asset_path_map + virtual_plan, virtual_clips, asset_path_map = _build_plan_and_clips_from_task( + task_id=task_id, + downloaded_paths=downloaded_videos, + mode=editing_mode.value, + ) - # 4. 如有配音,后处理混音 - if audio_path: - final_path = temp_path / f"final-{task_id}.mp4" - try: - _mux_audio_track(render_result.output_path, audio_path, final_path) - # 混音成功,使用混音后的文件 - output_path = final_path - except Exception as mux_err: - logger.warning("音频混合失败,使用无音频版本: %s", mux_err) - output_path = render_result.output_path - else: + # 使用 UnifiedRenderService 渲染 + render_service = UnifiedRenderService( + plan=virtual_plan, + clips=virtual_clips, + asset_path_map=asset_path_map, + work_dir=temp_path, + output_width=OUTPUT_WIDTH, + output_height=OUTPUT_HEIGHT, + output_fps=int(OUTPUT_FPS), + ) + render_result = render_service.render() + + # 4. 如有配音,后处理混音 + if audio_path: + final_path = temp_path / f"final-{task_id}.mp4" + try: + _mux_audio_track(render_result.output_path, audio_path, final_path) + # 混音成功,使用混音后的文件 + output_path = final_path + except Exception as mux_err: + logger.warning("音频混合失败,使用无音频版本: %s", mux_err) output_path = render_result.output_path else: - # 无素材,生成 fallback 视频 - _create_fallback_clip(output_path, f"Generated Video {task_id[:8]}") + output_path = render_result.output_path file_size = output_path.stat().st_size duration = probe_duration(output_path) - # 5. 上传到 OSS - bucket = oss_bucket() - if bucket: - try: - bucket.put_object_from_file(storage_key, str(output_path)) - except Exception as oss_err: - logger.warning("OSS upload failed: %s", oss_err) + # 5. 上传到 OSS — 失败必须抛异常,不能静默忽略 + file_url = upload_to_oss(output_path, storage_key) + if not file_url: + # OSS 未配置或上传失败 + raise RuntimeError( + f"OSS 上传失败: task_id={task_id}, storage_key={storage_key}, " f"output_path={output_path}" + ) - # 构建视频 URL - if bucket: - file_url = f"{PUBLIC_API_BASE_URL}/{storage_key}" - else: - file_url = f"{GENERATED_FILES_URL_PREFIX}/{task_id}/{output_name}" + # HEAD 校验 URL 可访问 + if not _verify_url_accessible(file_url): + raise RuntimeError(f"OSS 上传后 URL 不可访问: file_url={file_url}, " f"storage_key={storage_key}") + + logger.info("OSS 上传成功: file_url=%s", file_url) # 6. 创建 GeneratedVideo 记录 + 查重 dedup_session = SessionLocal() diff --git a/tests/unit/test_generation_p3_optimizations.py b/tests/unit/test_generation_p3_optimizations.py new file mode 100644 index 000000000..d4c379b2a --- /dev/null +++ b/tests/unit/test_generation_p3_optimizations.py @@ -0,0 +1,380 @@ +"""P3 优化单元测试 — generation.py 三项优化. + +覆盖: + P3-1: _download_library_assets strict 模式 + P3-2: 归属校验合并到同一 DB session + P3-3: _verify_url_accessible HEAD 重试 +""" + +from __future__ import annotations + +import sys +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker")) + + +# ── P3-3: _verify_url_accessible 重试 ─────────────────────────────────────── + + +class TestVerifyUrlAccessibleRetry: + """_verify_url_accessible 重试逻辑.""" + + @patch("time.sleep") + @patch("urllib.request.urlopen") + def test_first_attempt_success(self, mock_urlopen, mock_sleep): + """首次成功,不重试.""" + from worker_app.tasks.generation import _verify_url_accessible + + mock_resp = MagicMock() + mock_resp.status = 200 + mock_resp.__enter__ = MagicMock(return_value=mock_resp) + mock_resp.__exit__ = MagicMock(return_value=False) + mock_urlopen.return_value = mock_resp + + assert _verify_url_accessible("https://example.com/file.mp4") is True + assert mock_urlopen.call_count == 1 + mock_sleep.assert_not_called() + + @patch("time.sleep") + @patch("urllib.request.urlopen") + def test_retry_then_success(self, mock_urlopen, mock_sleep): + """首次失败,重试后成功.""" + from worker_app.tasks.generation import _verify_url_accessible + + # 第一次失败(网络异常),第二次成功 + mock_resp_ok = MagicMock() + mock_resp_ok.status = 200 + mock_resp_ok.__enter__ = MagicMock(return_value=mock_resp_ok) + mock_resp_ok.__exit__ = MagicMock(return_value=False) + + mock_urlopen.side_effect = [ + OSError("connection reset"), + mock_resp_ok, + ] + + assert _verify_url_accessible("https://example.com/file.mp4") is True + assert mock_urlopen.call_count == 2 + mock_sleep.assert_called_once_with(1) + + @patch("time.sleep") + @patch("urllib.request.urlopen") + def test_all_retries_exhausted(self, mock_urlopen, mock_sleep): + """全部重试耗尽,返回 False.""" + from worker_app.tasks.generation import _verify_url_accessible + + mock_urlopen.side_effect = OSError("connection refused") + + assert _verify_url_accessible("https://example.com/file.mp4") is False + # 1 首次 + 2 重试 = 3 次 + assert mock_urlopen.call_count == 3 + assert mock_sleep.call_count == 2 + + @patch("time.sleep") + @patch("urllib.request.urlopen") + def test_http_500_then_success(self, mock_urlopen, mock_sleep): + """HTTP 500 后重试成功.""" + from worker_app.tasks.generation import _verify_url_accessible + + mock_resp_500 = MagicMock() + mock_resp_500.status = 500 + mock_resp_500.__enter__ = MagicMock(return_value=mock_resp_500) + mock_resp_500.__exit__ = MagicMock(return_value=False) + + mock_resp_200 = MagicMock() + mock_resp_200.status = 200 + mock_resp_200.__enter__ = MagicMock(return_value=mock_resp_200) + mock_resp_200.__exit__ = MagicMock(return_value=False) + + mock_urlopen.side_effect = [mock_resp_500, mock_resp_200] + + assert _verify_url_accessible("https://example.com/file.mp4") is True + assert mock_urlopen.call_count == 2 + + @patch("time.sleep") + @patch("urllib.request.urlopen") + def test_custom_retries_zero(self, mock_urlopen, mock_sleep): + """retries=0 时不重试.""" + from worker_app.tasks.generation import _verify_url_accessible + + mock_urlopen.side_effect = OSError("timeout") + + assert _verify_url_accessible("https://example.com/file.mp4", retries=0) is False + assert mock_urlopen.call_count == 1 + mock_sleep.assert_not_called() + + +# ── P3-1: _download_library_assets strict 模式 ────────────────────────────── + + +def _make_mock_asset( + asset_id: str, + name: str, + file_url: str | None, + asset_library_id: str = "lib-1", + project_id: str = "", +): + """构造 mock AssetModel 实例.""" + return SimpleNamespace( + id=asset_id, + name=name, + file_url=file_url, + asset_library_id=asset_library_id, + project_id=project_id, + status="ready", + file_type="video", + created_at="2026-01-01", + ) + + +def _setup_mock_session(assets): + """构造 mock session,返回 (mock_session, mock_query_chain).""" + mock_session = MagicMock() + mock_query = MagicMock() + + # chain: session.query().filter().filter().order_by().all() + mock_session.query.return_value = mock_query + mock_query.filter.return_value = mock_query + mock_query.order_by.return_value = mock_query + mock_query.all.return_value = assets + + return mock_session + + +class TestDownloadLibraryAssetsStrictMode: + """_download_library_assets strict 模式.""" + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_strict_mode_raises_on_download_failure(self, mock_download, mock_session_factory): + """strict=True 时,单个素材下载失败立即抛 RuntimeError.""" + from worker_app.tasks.generation import _download_library_assets + + assets = [ + _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4"), + _make_mock_asset("a2", "video2.mp4", "uploads/video2.mp4"), + ] + mock_session = _setup_mock_session(assets) + mock_session_factory.return_value = mock_session + + # 第一个成功,第二个失败 + mock_download.side_effect = [True, False] + + with pytest.raises(RuntimeError, match="素材下载失败"): + _download_library_assets( + Path("/tmp"), + asset_library_id="lib-1", + asset_ids=["a1", "a2"], + strict=True, + ) + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_non_strict_mode_returns_partial_results(self, mock_download, mock_session_factory): + """strict=False 时,跳过失败素材,返回成功列表.""" + from worker_app.tasks.generation import _download_library_assets + + assets = [ + _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4"), + _make_mock_asset("a2", "video2.mp4", "uploads/video2.mp4"), + _make_mock_asset("a3", "video3.mp4", "uploads/video3.mp4"), + ] + mock_session = _setup_mock_session(assets) + mock_session_factory.return_value = mock_session + + # 第一个成功,第二个失败,第三个成功 + mock_download.side_effect = [True, False, True] + + result = _download_library_assets( + Path("/tmp"), + asset_library_id="lib-1", + asset_ids=["a1", "a2", "a3"], + strict=False, + ) + + assert len(result) == 2 + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_non_strict_all_fail_raises(self, mock_download, mock_session_factory): + """strict=False 但全部失败时仍抛 RuntimeError.""" + from worker_app.tasks.generation import _download_library_assets + + assets = [ + _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4"), + ] + mock_session = _setup_mock_session(assets) + mock_session_factory.return_value = mock_session + mock_download.return_value = False + + with pytest.raises(RuntimeError, match="全部下载失败"): + _download_library_assets( + Path("/tmp"), + asset_library_id="lib-1", + asset_ids=["a1"], + strict=False, + ) + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_strict_mode_raises_on_missing_file_url(self, mock_download, mock_session_factory): + """strict=True 时,素材缺少 file_url 立即抛异常.""" + from worker_app.tasks.generation import _download_library_assets + + assets = [ + _make_mock_asset("a1", "video1.mp4", None), # file_url 为空 + ] + mock_session = _setup_mock_session(assets) + mock_session_factory.return_value = mock_session + + with pytest.raises(RuntimeError, match="素材缺少 file_url"): + _download_library_assets( + Path("/tmp"), + asset_library_id="lib-1", + strict=True, + ) + + # download_asset 不应被调用 + mock_download.assert_not_called() + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_default_is_strict(self, mock_download, mock_session_factory): + """默认 strict=True.""" + from worker_app.tasks.generation import _download_library_assets + + assets = [ + _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4"), + ] + mock_session = _setup_mock_session(assets) + mock_session_factory.return_value = mock_session + mock_download.return_value = False + + # 不传 strict 参数,默认严格模式 + with pytest.raises(RuntimeError, match="素材下载失败"): + _download_library_assets( + Path("/tmp"), + asset_library_id="lib-1", + ) + + +# ── P3-2: 归属校验合并到同一 session ──────────────────────────────────────── + + +class TestDownloadLibraryAssetsOwnershipValidation: + """归属校验合并到 _download_library_assets 同一 session.""" + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_ownership_mismatch_raises_value_error(self, mock_download, mock_session_factory): + """asset_ids 不属于指定素材库时抛 ValueError.""" + from worker_app.tasks.generation import _download_library_assets + + # asset 属于 lib-2,但请求的是 lib-1 + assets = [ + _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4", asset_library_id="lib-2"), + ] + mock_session = _setup_mock_session(assets) + mock_session_factory.return_value = mock_session + + with pytest.raises(ValueError, match="素材不属于指定素材库"): + _download_library_assets( + Path("/tmp"), + asset_library_id="lib-1", + asset_ids=["a1"], + ) + + # 不应调用 download_asset(校验在下载前) + mock_download.assert_not_called() + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_missing_asset_ids_raises_value_error(self, mock_download, mock_session_factory): + """指定的 asset_ids 不存在时抛 ValueError.""" + from worker_app.tasks.generation import _download_library_assets + + # DB 返回空(asset_ids 不存在,query 过滤后无结果) + mock_session = _setup_mock_session([]) + mock_session_factory.return_value = mock_session + + with pytest.raises(RuntimeError, match="未找到视频素材"): + _download_library_assets( + Path("/tmp"), + asset_library_id="lib-1", + asset_ids=["nonexistent-id"], + ) + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_project_ownership_mismatch_raises(self, mock_download, mock_session_factory): + """项目级模式下归属不匹配抛 ValueError.""" + from worker_app.tasks.generation import _download_library_assets + + assets = [ + _make_mock_asset( + "a1", + "video1.mp4", + "uploads/video1.mp4", + asset_library_id="", + project_id="proj-2", + ), + ] + mock_session = _setup_mock_session(assets) + mock_session_factory.return_value = mock_session + + with pytest.raises(ValueError, match="素材不属于指定项目"): + _download_library_assets( + Path("/tmp"), + project_id="proj-1", + asset_ids=["a1"], + ) + + mock_download.assert_not_called() + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_ownership_pass_then_download(self, mock_download, mock_session_factory): + """归属校验通过后正常下载.""" + from worker_app.tasks.generation import _download_library_assets + + assets = [ + _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4", asset_library_id="lib-1"), + ] + mock_session = _setup_mock_session(assets) + mock_session_factory.return_value = mock_session + mock_download.return_value = True + + result = _download_library_assets( + Path("/tmp"), + asset_library_id="lib-1", + asset_ids=["a1"], + ) + + assert len(result) == 1 + mock_download.assert_called_once() + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_single_session_used(self, mock_download, mock_session_factory): + """验证只创建了一个 DB session(P3-2 核心).""" + from worker_app.tasks.generation import _download_library_assets + + assets = [ + _make_mock_asset("a1", "video1.mp4", "uploads/video1.mp4", asset_library_id="lib-1"), + ] + mock_session = _setup_mock_session(assets) + mock_session_factory.return_value = mock_session + mock_download.return_value = True + + _download_library_assets( + Path("/tmp"), + asset_library_id="lib-1", + asset_ids=["a1"], + ) + + # SessionLocal 只调用一次(合并前会调用两次:校验 + 下载) + assert mock_session_factory.call_count == 1 diff --git a/tests/unit/test_oneclick_gen_p0_fixes.py b/tests/unit/test_oneclick_gen_p0_fixes.py new file mode 100644 index 000000000..7bfd58ff0 --- /dev/null +++ b/tests/unit/test_oneclick_gen_p0_fixes.py @@ -0,0 +1,268 @@ +"""P0/P1 修复单元测试 — 一键生成 P0 问题 + P1 校验. + +覆盖: + P0-1: _download_library_assets 双模式查询(asset_library_id / project_id) + P0-2: OSS 上传失败抛异常 + URL 可访问性校验 + P0-3: FFmpeg 失败时完整 stderr 日志 + P1: template_id 存在性校验 + asset_ids 归属校验 +""" + +from __future__ import annotations + +import logging +import subprocess +import sys +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + +# 添加 worker app 到 sys.path +_WORKER_ROOT = Path(__file__).resolve().parents[2] / "apps" / "worker" +if str(_WORKER_ROOT) not in sys.path: + sys.path.insert(0, str(_WORKER_ROOT)) + + +# ── P0-1: _download_library_assets ──────────────────────────────────────────── + + +class TestDownloadLibraryAssets: + """P0-1: 素材下载双模式 + 错误处理.""" + + def _make_asset(self, id_: str, file_url: str, project_id: str = "p1", library_id: str = "lib1"): + mock = MagicMock() + mock.id = id_ + mock.file_url = file_url + mock.name = f"asset_{id_}" + mock.project_id = project_id + mock.asset_library_id = library_id + return mock + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_asset_library_mode(self, mock_download, mock_session_factory): + """素材库模式:按 asset_library_id 查询.""" + from worker_app.tasks.generation import _download_library_assets + + session = MagicMock() + mock_session_factory.return_value = session + query = MagicMock() + session.query.return_value = query + filter_result = MagicMock() + query.filter.return_value = filter_result + in_filter = MagicMock() + filter_result.filter.return_value = in_filter + assets = [self._make_asset("a1", "video/a1.mp4")] + in_filter.order_by.return_value.all.return_value = assets + + mock_download.return_value = True + + with patch("worker_app.tasks.generation.AssetModel", create=True): + result = _download_library_assets( + Path("/tmp/test"), + asset_library_id="lib1", + ) + + assert len(result) == 1 + mock_download.assert_called_once() + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_project_mode(self, mock_download, mock_session_factory): + """项目级模式:asset_library_id 为空时按 project_id 查询.""" + from worker_app.tasks.generation import _download_library_assets + + session = MagicMock() + mock_session_factory.return_value = session + query = MagicMock() + session.query.return_value = query + filter_result = MagicMock() + query.filter.return_value = filter_result + proj_filter = MagicMock() + filter_result.filter.return_value = proj_filter + assets = [self._make_asset("a1", "video/a1.mp4", project_id="proj1")] + proj_filter.order_by.return_value.all.return_value = assets + + mock_download.return_value = True + + result = _download_library_assets( + Path("/tmp/test"), + project_id="proj1", + ) + + assert len(result) == 1 + + def test_both_empty_raises(self): + """asset_library_id 和 project_id 都为空时抛 ValueError.""" + from worker_app.tasks.generation import _download_library_assets + + with pytest.raises(ValueError, match="至少需要提供一个"): + _download_library_assets(Path("/tmp/test")) + + @patch("worker_app.tasks.generation.SessionLocal") + def test_no_assets_found_raises(self, mock_session_factory): + """查不到素材时抛 RuntimeError.""" + from worker_app.tasks.generation import _download_library_assets + + session = MagicMock() + mock_session_factory.return_value = session + query = MagicMock() + session.query.return_value = query + filter_result = MagicMock() + query.filter.return_value = filter_result + in_filter = MagicMock() + filter_result.filter.return_value = in_filter + in_filter.order_by.return_value.all.return_value = [] + + with pytest.raises(RuntimeError, match="未找到视频素材"): + _download_library_assets( + Path("/tmp/test"), + asset_library_id="lib1", + ) + + @patch("worker_app.tasks.generation.SessionLocal") + @patch("worker_app.tasks.generation.download_asset") + def test_all_asset_ids_fail_raises(self, mock_download, mock_session_factory): + """指定 asset_ids 但全部下载失败时抛 RuntimeError.""" + from worker_app.tasks.generation import _download_library_assets + + session = MagicMock() + mock_session_factory.return_value = session + query = MagicMock() + session.query.return_value = query + filter_result = MagicMock() + query.filter.return_value = filter_result + in_filter = MagicMock() + filter_result.filter.return_value = in_filter + id_filter = MagicMock() + in_filter.filter.return_value = id_filter + assets = [self._make_asset("a1", "video/a1.mp4")] + id_filter.order_by.return_value.all.return_value = assets + + mock_download.return_value = False # 全部下载失败 + + with pytest.raises(RuntimeError, match="素材下载失败"): + _download_library_assets( + Path("/tmp/test"), + asset_library_id="lib1", + asset_ids=["a1"], + ) + + +# ── P0-2: OSS 上传 + URL 校验 ──────────────────────────────────────────────── + + +class TestOSSUploadAndVerify: + """P0-2: OSS 上传失败抛异常 + URL 可访问性校验.""" + + def test_verify_url_accessible_success(self): + """URL 可访问时返回 True.""" + from worker_app.tasks.generation import _verify_url_accessible + + mock_response = MagicMock() + mock_response.status = 200 + mock_response.__enter__ = MagicMock(return_value=mock_response) + mock_response.__exit__ = MagicMock(return_value=False) + + with patch("urllib.request.urlopen", return_value=mock_response): + assert _verify_url_accessible("https://example.com/test.mp4") is True + + def test_verify_url_accessible_failure(self): + """URL 不可访问时返回 False.""" + from worker_app.tasks.generation import _verify_url_accessible + + with patch("urllib.request.urlopen", side_effect=Exception("connection refused")): + assert _verify_url_accessible("https://example.com/test.mp4") is False + + def test_verify_url_404(self): + """URL 返回 404 时返回 False.""" + from worker_app.tasks.generation import _verify_url_accessible + + mock_response = MagicMock() + mock_response.status = 404 + mock_response.__enter__ = MagicMock(return_value=mock_response) + mock_response.__exit__ = MagicMock(return_value=False) + + with patch("urllib.request.urlopen", return_value=mock_response): + assert _verify_url_accessible("https://example.com/test.mp4") is False + + +# ── P0-3: FFmpeg stderr 日志 ───────────────────────────────────────────────── + + +class TestFFmpegStderrLogging: + """P0-3: FFmpeg 失败时完整 stderr 打到日志.""" + + def test_run_ffmpeg_logs_stderr_on_failure(self, caplog): + """run_ffmpeg 失败时记录 stderr 到日志.""" + from video_processing.ffmpeg_utils import run_ffmpeg + + error = subprocess.CalledProcessError( + returncode=183, + cmd=["ffmpeg", "-y", "-i", "input.mp4", "output.mp4"], + output="", + stderr="Error message from ffmpeg: filter graph error details here", + ) + + with patch("subprocess.run", side_effect=error): + with caplog.at_level(logging.ERROR): + with pytest.raises(subprocess.CalledProcessError): + run_ffmpeg(["ffmpeg", "-y", "-i", "input.mp4", "output.mp4"]) + + assert "FFmpeg 命令失败" in caplog.text + assert "exit_code=183" in caplog.text + assert "filter graph error" in caplog.text + + def test_run_ffmpeg_success(self): + """run_ffmpeg 成功时正常返回.""" + from video_processing.ffmpeg_utils import run_ffmpeg + + mock_result = MagicMock() + mock_result.stdout = "output" + mock_result.stderr = "" + + with patch("subprocess.run", return_value=mock_result): + stdout, stderr = run_ffmpeg(["ffmpeg", "-version"]) + assert stdout == "output" + assert stderr == "" + + +# ── P1: template_id + asset_ids 校验 ───────────────────────────────────────── + + +class TestP1Validations: + """P1: template_id 存在性校验 + asset_ids 归属校验.""" + + def test_validate_template_exists_success(self): + """模板存在时不抛异常.""" + from worker_app.tasks.generation import _validate_template_exists + + mock_template = MagicMock() + mock_template.id = "tmpl_001" + mock_template.name = "Test Template" + mock_template.is_active = True + + session = MagicMock() + mock_session = MagicMock() + session.query.return_value = mock_session + filter_result = MagicMock() + mock_session.filter.return_value = filter_result + filter_result.first.return_value = mock_template + + with patch("worker_app.tasks.generation.SessionLocal", return_value=session): + _validate_template_exists("tmpl_001") # 不抛异常 + + def test_validate_template_exists_not_found(self): + """模板不存在时抛 ValueError.""" + from worker_app.tasks.generation import _validate_template_exists + + session = MagicMock() + mock_session = MagicMock() + session.query.return_value = mock_session + filter_result = MagicMock() + mock_session.filter.return_value = filter_result + filter_result.first.return_value = None + + with patch("worker_app.tasks.generation.SessionLocal", return_value=session): + with pytest.raises(ValueError, match="模板不存在"): + _validate_template_exists("tmpl_nonexistent")