diff --git a/apps/worker/video_processing/render_adapter.py b/apps/worker/video_processing/render_adapter.py index b839a646d..c6611d814 100755 --- a/apps/worker/video_processing/render_adapter.py +++ b/apps/worker/video_processing/render_adapter.py @@ -1,14 +1,16 @@ """统一渲染引擎适配层 — Phase 2. 将 EditPlan + EditPlanClips(来自 DB)适配为 UnifiedRenderService 的输入格式, -封装素材下载、渲染执行、结果上传的完整流程。 +封装素材下载、BGM 准备、ASR 字幕、渲染执行、结果上传的完整流程。 职责: 1. 从 DB 读取 EditPlan + EditPlanClips 2. 下载素材到本地,构建 asset_path_map - 3. 调用 UnifiedRenderService 执行渲染 - 4. 上传渲染结果到 OSS - 5. 支持进度回调(对接 JobService) + 3. 准备 BGM 音频(URL / 素材库 / 预设库) + 4. 初始化 ASR 服务(自动字幕) + 5. 调用 UnifiedRenderService 执行渲染 + 6. 上传渲染结果到 OSS + 7. 支持进度回调(对接 JobService) """ from __future__ import annotations @@ -17,7 +19,7 @@ import logging import tempfile from dataclasses import dataclass from pathlib import Path -from typing import Callable +from typing import Any, Callable from sqlalchemy.orm import Session from video_processing.oss_helpers import download_asset, upload_to_oss @@ -46,8 +48,16 @@ class RenderAdapterResult: width: int = 0 height: int = 0 clip_count: int = 0 + rendered_clip_ids: list[str] = None # 成功渲染的 clip id 列表 + failed_clip_ids: list[str] = None # 失败的 clip id 列表 error_message: str = "" + def __post_init__(self): + if self.rendered_clip_ids is None: + self.rendered_clip_ids = [] + if self.failed_clip_ids is None: + self.failed_clip_ids = [] + ProgressCallback = Callable[[float, str], None] """进度回调:(progress_0_100, stage_description) → None""" @@ -142,22 +152,34 @@ class RenderAdapter: self._report_progress(progress_cb, 15.0, f"下载素材({len(ready_clips)} 个)") # 2. 下载素材 - asset_path_map = self._download_assets(ready_clips, work_dir) + asset_path_map, rendered_clip_ids, failed_clip_ids = self._download_assets(ready_clips, work_dir) if not asset_path_map: return RenderAdapterResult( success=False, error_message="所有素材下载失败", clip_count=len(ready_clips), + rendered_clip_ids=[], + failed_clip_ids=failed_clip_ids, ) + self._report_progress(progress_cb, 35.0, "准备 BGM 音频") + + # 3. 准备 BGM(从 plan.config.bgm 读取配置) + bgm_path = self._prepare_bgm(plan, work_dir, plan_id) + self._report_progress(progress_cb, 40.0, "执行视频渲染") - # 3. 执行统一渲染 + # 4. 初始化 ASR 服务(用于自动字幕) + asr_service = self._get_asr_service() + + # 5. 执行统一渲染 render_svc = UnifiedRenderService( plan=plan, clips=ready_clips, asset_path_map=asset_path_map, work_dir=work_dir, + bgm_path=bgm_path, + asr_service=asr_service, ) result = render_svc.render() @@ -190,6 +212,8 @@ class RenderAdapter: width=result.width, height=result.height, clip_count=len(ready_clips), + rendered_clip_ids=rendered_clip_ids, + failed_clip_ids=failed_clip_ids, ) except Exception as exc: @@ -271,29 +295,153 @@ class RenderAdapter: logger.exception("进度回调失败") @staticmethod - def _download_assets(clips: list[EditPlanClip], work_dir: Path) -> dict[str, Path]: - """下载片段素材到本地,返回 asset_id → local_path 映射。 + def _download_assets(clips: list[EditPlanClip], work_dir: Path) -> tuple[dict[str, Path], list[str], list[str]]: + """下载片段素材到本地。 - 只保留下载成功的素材。 + Returns: + (asset_path_map, rendered_clip_ids, failed_clip_ids) + - asset_path_map: asset_id → local_path 映射(下载成功的) + - rendered_clip_ids: 下载成功的 clip id 列表 + - failed_clip_ids: 下载失败的 clip id 列表 """ asset_dir = work_dir / "assets" asset_dir.mkdir(exist_ok=True) asset_path_map: dict[str, Path] = {} + rendered_clip_ids: list[str] = [] + failed_clip_ids: list[str] = [] + seen_asset_ids: set[str] = set() for clip in clips: asset_id = clip.asset_id if not asset_id: + failed_clip_ids.append(clip.id) continue + # 同一素材已下载过(多个 clip 共享同一素材) + if asset_id in seen_asset_ids: + if asset_id in asset_path_map: + rendered_clip_ids.append(clip.id) + else: + failed_clip_ids.append(clip.id) + continue + + seen_asset_ids.add(asset_id) + # 生成安全的本地文件名 safe_name = f"clip_{clip.order:04d}_{abs(hash(asset_id)) % 100000:05d}.mp4" local_path = asset_dir / safe_name if download_asset(asset_id, local_path): asset_path_map[asset_id] = local_path + rendered_clip_ids.append(clip.id) logger.debug("素材下载成功: clip_id=%s asset_id=%s", clip.id, asset_id[:60]) else: + failed_clip_ids.append(clip.id) logger.warning("素材下载失败: clip_id=%s asset_id=%s", clip.id, asset_id[:60]) - return asset_path_map + return asset_path_map, rendered_clip_ids, failed_clip_ids + + def _prepare_bgm(self, plan, work_dir: Path, plan_id: str) -> str | None: + """准备 BGM 音频文件(从 plan.config.bgm 读取配置)。 + + 支持 3 种来源(按优先级): + 1. audio_url — 外部直链 URL + 2. asset_id — 素材库中的音频素材 + 3. preset_id — 预设 BGM 库 + + 失败不阻断主流程,返回 None。 + """ + from urllib.parse import urlparse + + plan_config = plan.config or {} + bgm_config = plan_config.get("bgm", {}) or {} + + if not bgm_config.get("enabled", False): + return None + + audio_url = bgm_config.get("audio_url", "") or "" + asset_id = bgm_config.get("asset_id", "") or "" + preset_id = bgm_config.get("preset_id", "") or "" + + bgm_file = work_dir / "bgm.mp3" + + # 优先级1:外部直链 URL + if audio_url: + try: + parsed = urlparse(audio_url) + if parsed.scheme in ("http", "https"): + from video_processing.url_security import ( + ALLOWED_AUDIO_MIME_TYPES, + safe_download_file, + ) + + logger.info("[plan_id=%s] [BGM] 从URL下载: %s", plan_id, audio_url[:80]) + 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: + logger.warning("[plan_id=%s] [BGM] URL下载失败: %s", plan_id, e) + + # 优先级2:素材库素材 + if asset_id: + try: + from packages.adapters.sqlalchemy_impl.models import AssetModel + + model = self._db.query(AssetModel).filter(AssetModel.id == asset_id).first() + if model and model.file_url: + storage_key = model.file_url + logger.info("[plan_id=%s] [BGM] 从素材库下载: asset_id=%s", plan_id, asset_id) + ok = download_asset(storage_key, bgm_file) + if ok and bgm_file.exists() and bgm_file.stat().st_size > 0: + return str(bgm_file) + except Exception as e: + logger.warning("[plan_id=%s] [BGM] 素材库下载失败: %s", plan_id, e) + + # 优先级3:预设 BGM 库 + if preset_id: + try: + from packages.domain.preset_bgm import get_preset_bgm + + preset = get_preset_bgm(preset_id) + if preset and preset.audio_url: + from video_processing.url_security import ( + ALLOWED_AUDIO_MIME_TYPES, + safe_download_file, + ) + + logger.info("[plan_id=%s] [BGM] 从预设库下载: preset_id=%s", plan_id, preset_id) + 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: + logger.warning("[plan_id=%s] [BGM] 预设库下载失败: %s", plan_id, e) + + logger.warning("[plan_id=%s] [BGM] 所有来源都无法获取BGM,跳过", plan_id) + return None + + @staticmethod + def _get_asr_service() -> Any | None: + """获取 ASR 服务实例(用于自动生成字幕)。 + + 失败不阻断主流程,返回 None。 + """ + try: + from services.asr_service_factory import get_asr_service + + return get_asr_service() + except Exception as e: + logger.warning("ASR 服务初始化失败,自动字幕将不可用: %s", e) + return None diff --git a/apps/worker/worker_app/tasks/edit_plan_generation.py b/apps/worker/worker_app/tasks/edit_plan_generation.py index b671df4b5..c879a6632 100755 --- a/apps/worker/worker_app/tasks/edit_plan_generation.py +++ b/apps/worker/worker_app/tasks/edit_plan_generation.py @@ -39,7 +39,6 @@ from video_processing.oss_helpers import ( download_asset, upload_to_oss, ) -from video_processing.unified_render_service import UnifiedRenderService # ── Repository imports (延迟导入避免循环依赖) ───────────────────────────────── @@ -218,9 +217,6 @@ def _finalize_render_success( def _render_with_unified( plan, clips, - asset_path_map: dict[str, Path], - tmpdir_path: Path, - rendered_clip_ids: list[str], plan_id: str, generation_task_id: str, plan_repo, @@ -228,31 +224,59 @@ def _render_with_unified( gen_task_repo, db, ) -> dict: - """统一渲染引擎路径(UnifiedRenderService 图层架构)。""" - render_service = UnifiedRenderService( - plan=plan, - clips=clips, - asset_path_map=asset_path_map, - work_dir=tmpdir_path, - output_width=OUTPUT_WIDTH, - output_height=OUTPUT_HEIGHT, - output_fps=int(OUTPUT_FPS), - ) + """统一渲染引擎路径(通过 RenderAdapter 调用 UnifiedRenderService)。 + + RenderAdapter 内部处理:素材下载、BGM 准备、ASR 自动字幕、渲染执行、OSS 上传。 + 本函数只负责:业务状态更新、查重、收尾。 + """ + from video_processing.render_adapter import RenderAdapter + + adapter = RenderAdapter(db) + + # 进度回调:更新 GenerationTask 进度 + def _progress_cb(progress: float, stage: str): + if not generation_task_id: + return + try: + gen_task = gen_task_repo.get(generation_task_id) + if gen_task: + # 映射到 30%~90% 区间(素材下载前已到 30%) + mapped_progress = 30.0 + progress * 0.6 + gen_task.progress = min(mapped_progress, 95.0) + gen_task.append_log( + stage="render_progress", + message=stage, + level="INFO", + progress=mapped_progress, + ) + gen_task_repo.update(gen_task) + except Exception: + pass try: - render_result = render_service.render() + result = adapter.render_plan( + plan_id=plan_id, + job_id=generation_task_id or plan_id, + progress_cb=_progress_cb, + ) except Exception as render_err: logger.error("渲染失败(unified): %s — %s", plan_id, render_err) _mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, f"渲染失败: {render_err}") return {"status": "error", "message": f"渲染失败: {render_err}"} - output_path = render_result.output_path + if not result.success: + logger.error("渲染失败(unified): %s — %s", plan_id, result.error_message) + _mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, result.error_message or "渲染失败") + return {"status": "error", "message": result.error_message or "渲染失败"} - # 上传到 OSS + output_path = result.output_path or Path("") + output_url = result.output_url storage_key = f"rendered/{plan_id}/output.mp4" - output_url = upload_to_oss(output_path, storage_key) - failed_clip_ids: list[str] = [] + # 用 adapter 返回的 clip 明细(以 adapter 的结果为准) + rendered_clip_ids = result.rendered_clip_ids or [] + failed_clip_ids = result.failed_clip_ids or [] + return _finalize_render_success( plan=plan, plan_repo=plan_repo, @@ -262,10 +286,10 @@ def _render_with_unified( plan_id=plan_id, output_url=output_url or "", storage_key=storage_key, - duration=render_result.duration, - file_size=render_result.file_size, - width=render_result.width, - height=render_result.height, + duration=result.duration, + file_size=result.file_size, + width=result.width, + height=result.height, rendered_clip_ids=rendered_clip_ids, failed_clip_ids=failed_clip_ids, generation_task_id=generation_task_id, @@ -406,128 +430,123 @@ def render_edit_plan(self, plan_id: str) -> dict: ) gen_task_repo.update(gen_task) - # 3. 下载素材并构建 asset_path_map - with tempfile.TemporaryDirectory(prefix="edit_plan_") as tmpdir: - tmpdir_path = Path(tmpdir) - asset_path_map: dict[str, Path] = {} - rendered_clip_ids: list[str] = [] - failed_clip_ids: list[str] = [] + # 3. 渲染前取消检查 + if generation_task_id: + current_task = gen_task_repo.get(generation_task_id) + if current_task: + task_status = ( + current_task.status.value if hasattr(current_task.status, "value") else str(current_task.status) + ) + if task_status == "cancelled": + logger.info("任务已被取消,中止渲染: plan_id=%s task_id=%s", plan_id, generation_task_id) + from packages.domain.edit_plan import EditPlanStatus - # 预先批量查询所有素材的 storage_key(file_url) - from packages.adapters.sqlalchemy_impl.models import AssetModel + if plan.status.value == "rendering": + try: + plan.resume_editing() + plan_repo.update(plan) + except ValueError: + pass + return {"status": "cancelled", "plan_id": plan_id, "message": "任务已取消"} - clip_asset_ids = [c.asset_id for c in clips if c.asset_id] - asset_storage_map: dict[str, str] = {} - if clip_asset_ids: - assets = db.query(AssetModel).filter(AssetModel.id.in_(clip_asset_ids)).all() - asset_storage_map = {a.id: a.file_url for a in assets if a.file_url} + # 4. 根据引擎选择渲染方式 + if engine == "unified": + # ── unified 路径:RenderAdapter 统一处理(下载 + BGM + ASR + 渲染 + 上传) + result = _render_with_unified( + plan=plan, + clips=clips, + plan_id=plan_id, + generation_task_id=generation_task_id, + plan_repo=plan_repo, + clip_repo=clip_repo, + gen_task_repo=gen_task_repo, + db=db, + ) + else: + # ── legacy 路径:原有的素材下载 + VideoComposeService + with tempfile.TemporaryDirectory(prefix="edit_plan_") as tmpdir: + tmpdir_path = Path(tmpdir) + asset_path_map: dict[str, Path] = {} + rendered_clip_ids: list[str] = [] + failed_clip_ids: list[str] = [] - for clip in clips: - if not clip.asset_id: - # 没有素材的片段跳过,标记为失败 - clip.mark_failed() - clip_repo.update(clip) - failed_clip_ids.append(clip.id) - continue + # 预先批量查询所有素材的 storage_key(file_url) + from packages.adapters.sqlalchemy_impl.models import AssetModel - if clip.asset_id in asset_path_map: - # 同一素材已下载(多个 clip 共享同一素材) - rendered_clip_ids.append(clip.id) - continue + clip_asset_ids = [c.asset_id for c in clips if c.asset_id] + asset_storage_map: dict[str, str] = {} + if clip_asset_ids: + assets = db.query(AssetModel).filter(AssetModel.id.in_(clip_asset_ids)).all() + asset_storage_map = {a.id: a.file_url for a in assets if a.file_url} - storage_key = asset_storage_map.get(clip.asset_id) - if not storage_key: - logger.warning( - "片段素材无 storage_key,跳过: clip_id=%s asset_id=%s", - clip.id, - clip.asset_id, - ) - clip.mark_failed() - clip_repo.update(clip) - failed_clip_ids.append(clip.id) - continue + for clip in clips: + if not clip.asset_id: + # 没有素材的片段跳过,标记为失败 + clip.mark_failed() + clip_repo.update(clip) + failed_clip_ids.append(clip.id) + continue - # 下载素材 - ext = Path(storage_key).suffix or ".mp4" - local_path = tmpdir_path / f"clip_{clip.order:04d}{ext}" - if download_asset(storage_key, local_path): - asset_path_map[clip.asset_id] = local_path - rendered_clip_ids.append(clip.id) - else: - clip.mark_failed() - clip_repo.update(clip) - failed_clip_ids.append(clip.id) + if clip.asset_id in asset_path_map: + # 同一素材已下载(多个 clip 共享同一素材) + rendered_clip_ids.append(clip.id) + continue - if not asset_path_map: - logger.error("所有片段素材下载失败: %s", plan_id) - plan.mark_failed() - plan_repo.update(plan) + storage_key = asset_storage_map.get(clip.asset_id) + if not storage_key: + logger.warning( + "片段素材无 storage_key,跳过: clip_id=%s asset_id=%s", + clip.id, + clip.asset_id, + ) + clip.mark_failed() + clip_repo.update(clip) + failed_clip_ids.append(clip.id) + continue + + # 下载素材 + ext = Path(storage_key).suffix or ".mp4" + local_path = tmpdir_path / f"clip_{clip.order:04d}{ext}" + if download_asset(storage_key, local_path): + asset_path_map[clip.asset_id] = local_path + rendered_clip_ids.append(clip.id) + else: + clip.mark_failed() + clip_repo.update(clip) + failed_clip_ids.append(clip.id) + + if not asset_path_map: + logger.error("所有片段素材下载失败: %s", plan_id) + plan.mark_failed() + plan_repo.update(plan) + if generation_task_id: + gen_task = gen_task_repo.get(generation_task_id) + if gen_task: + gen_task.status = "failed" + gen_task.error_message = "所有片段素材下载失败" + gen_task.completed_at = datetime.now(timezone.utc) + gen_task.append_log( + stage="download_failed", + message="所有片段素材下载失败", + level="ERROR", + ) + gen_task_repo.update(gen_task) + return {"status": "error", "message": "所有片段素材下载失败"} + + # 素材下载完成,记录日志 if generation_task_id: gen_task = gen_task_repo.get(generation_task_id) if gen_task: - gen_task.status = "failed" - gen_task.error_message = "所有片段素材下载失败" - gen_task.completed_at = datetime.now(timezone.utc) gen_task.append_log( - stage="download_failed", - message="所有片段素材下载失败", - level="ERROR", + stage="download_done", + message=f"素材下载完成,成功 {len(asset_path_map)} 个,失败 {len(failed_clip_ids)} 个", + level="INFO", + success_count=len(asset_path_map), + failed_count=len(failed_clip_ids), ) + gen_task.progress = 30.0 gen_task_repo.update(gen_task) - return {"status": "error", "message": "所有片段素材下载失败"} - # 素材下载完成,记录日志 - if generation_task_id: - gen_task = gen_task_repo.get(generation_task_id) - if gen_task: - gen_task.append_log( - stage="download_done", - message=f"素材下载完成,成功 {len(asset_path_map)} 个,失败 {len(failed_clip_ids)} 个", - level="INFO", - success_count=len(asset_path_map), - failed_count=len(failed_clip_ids), - ) - gen_task.progress = 30.0 - gen_task_repo.update(gen_task) - - # 4. 根据引擎选择渲染方式 - # 取消检查:素材下载完后,确认任务没有被用户取消 - if generation_task_id: - current_task = gen_task_repo.get(generation_task_id) - if current_task: - task_status = ( - current_task.status.value - if hasattr(current_task.status, "value") - else str(current_task.status) - ) - if task_status == "cancelled": - logger.info("任务已被取消,中止渲染: plan_id=%s task_id=%s", plan_id, generation_task_id) - # 计划回到 editing 状态,用户可以继续编辑 - from packages.domain.edit_plan import EditPlanStatus - - if plan.status.value == "rendering": - try: - plan.resume_editing() - plan_repo.update(plan) - except ValueError: - pass - return {"status": "cancelled", "plan_id": plan_id, "message": "任务已取消"} - - if engine == "unified": - result = _render_with_unified( - plan=plan, - clips=clips, - asset_path_map=asset_path_map, - tmpdir_path=tmpdir_path, - rendered_clip_ids=rendered_clip_ids, - plan_id=plan_id, - generation_task_id=generation_task_id, - plan_repo=plan_repo, - clip_repo=clip_repo, - gen_task_repo=gen_task_repo, - db=db, - ) - else: result = _render_with_legacy( plan=plan, clips=clips, @@ -542,8 +561,8 @@ def render_edit_plan(self, plan_id: str) -> dict: db=db, ) - result["engine"] = engine - return result + result["engine"] = engine + return result except Exception as exc: logger.exception("渲染剪辑计划异常: %s", plan_id) diff --git a/tests/unit/test_render_adapter.py b/tests/unit/test_render_adapter.py index 7d6ff9736..e8dfd8d9d 100755 --- a/tests/unit/test_render_adapter.py +++ b/tests/unit/test_render_adapter.py @@ -6,6 +6,7 @@ from __future__ import annotations from dataclasses import dataclass, field +from pathlib import Path from typing import Any from unittest.mock import MagicMock, patch @@ -386,11 +387,13 @@ class TestDownloadAssets: _make_clip("c2", order=1, asset_id="key2.mp4"), ] - result = RenderAdapter._download_assets(clips, tmp_path) + asset_path_map, rendered_ids, failed_ids = RenderAdapter._download_assets(clips, tmp_path) - assert len(result) == 2 - assert "key1.mp4" in result - assert "key2.mp4" in result + assert len(asset_path_map) == 2 + assert "key1.mp4" in asset_path_map + assert "key2.mp4" in asset_path_map + assert len(rendered_ids) == 2 + assert len(failed_ids) == 0 assert mock_download.call_count == 2 @patch("video_processing.render_adapter.download_asset") @@ -402,10 +405,12 @@ class TestDownloadAssets: ] mock_download.return_value = True - result = RenderAdapter._download_assets(clips, tmp_path) + asset_path_map, rendered_ids, failed_ids = RenderAdapter._download_assets(clips, tmp_path) - assert len(result) == 1 - assert "key2.mp4" in result + assert len(asset_path_map) == 1 + assert "key2.mp4" in asset_path_map + assert "c1" in failed_ids + assert "c2" in rendered_ids assert mock_download.call_count == 1 # 只调用了一次下载 @patch("video_processing.render_adapter.download_asset") @@ -417,6 +422,392 @@ class TestDownloadAssets: _make_clip("c1", order=0, asset_id="key1.mp4"), ] - result = RenderAdapter._download_assets(clips, tmp_path) + asset_path_map, rendered_ids, failed_ids = RenderAdapter._download_assets(clips, tmp_path) - assert len(result) == 0 + assert len(asset_path_map) == 0 + assert len(rendered_ids) == 0 + assert "c1" in failed_ids + + @patch("video_processing.render_adapter.download_asset") + def test_partial_download_failure(self, mock_download, tmp_path): + """部分下载失败时正确区分成功/失败。""" + results = {"key1.mp4": True, "key2.mp4": False, "key3.mp4": True} + + def _fake_download(asset_id, local_path): + local_path.parent.mkdir(parents=True, exist_ok=True) + local_path.write_bytes(b"fake") + return results.get(asset_id, False) + + mock_download.side_effect = _fake_download + + clips = [ + _make_clip("c1", order=0, asset_id="key1.mp4"), + _make_clip("c2", order=1, asset_id="key2.mp4"), + _make_clip("c3", order=2, asset_id="key3.mp4"), + ] + + asset_path_map, rendered_ids, failed_ids = RenderAdapter._download_assets(clips, tmp_path) + + assert len(asset_path_map) == 2 + assert "c1" in rendered_ids + assert "c3" in rendered_ids + assert "c2" in failed_ids + + @patch("video_processing.render_adapter.download_asset") + def test_duplicate_asset_downloaded_once(self, mock_download, tmp_path): + """同一素材被多个 clip 引用时只下载一次。""" + mock_download.return_value = True + + clips = [ + _make_clip("c1", order=0, asset_id="shared.mp4"), + _make_clip("c2", order=1, asset_id="shared.mp4"), + ] + + asset_path_map, rendered_ids, failed_ids = RenderAdapter._download_assets(clips, tmp_path) + + assert len(asset_path_map) == 1 + assert mock_download.call_count == 1 + assert "c1" in rendered_ids + assert "c2" in rendered_ids + + +# ── _prepare_bgm 测试 ──────────────────────────────────────────────────────── + + +class TestPrepareBgm: + """BGM 音频准备逻辑测试。""" + + def test_bgm_disabled_returns_none(self, tmp_path): + """BGM 未启用时返回 None。""" + plan = FakePlan(config={"bgm": {"enabled": False}}) + adapter = RenderAdapter(MagicMock()) + result = adapter._prepare_bgm(plan, tmp_path, "plan_001") + assert result is None + + def test_bgm_no_config_returns_none(self, tmp_path): + """没有 BGM 配置时返回 None。""" + plan = FakePlan(config={}) + adapter = RenderAdapter(MagicMock()) + result = adapter._prepare_bgm(plan, tmp_path, "plan_001") + assert result is None + + def test_bgm_from_url(self, tmp_path): + """从 URL 下载 BGM。""" + plan = FakePlan( + config={ + "bgm": { + "enabled": True, + "audio_url": "https://example.com/bgm.mp3", + } + } + ) + adapter = RenderAdapter(MagicMock()) + + with patch("video_processing.url_security.safe_download_file") as mock_download: + + def _fake_download(url, path, **kwargs): + Path(path).parent.mkdir(parents=True, exist_ok=True) + Path(path).write_bytes(b"fake mp3 data") + + mock_download.side_effect = _fake_download + + result = adapter._prepare_bgm(plan, tmp_path, "plan_001") + + assert result is not None + assert Path(result).exists() + + def test_bgm_from_url_failure_falls_through(self, tmp_path): + """URL 下载失败时不抛异常,继续尝试其他来源。""" + plan = FakePlan( + config={ + "bgm": { + "enabled": True, + "audio_url": "https://example.com/bgm.mp3", + "preset_id": "preset_001", + } + } + ) + adapter = RenderAdapter(MagicMock()) + + with patch("video_processing.url_security.safe_download_file") as mock_download: + mock_download.side_effect = Exception("download failed") + + # 预设库返回 None(没有这个预设) + with patch("packages.domain.preset_bgm.get_preset_bgm") as mock_preset: + mock_preset.return_value = None + result = adapter._prepare_bgm(plan, tmp_path, "plan_001") + + assert result is None # 所有来源都失败时返回 None + + def test_bgm_from_asset_library(self, tmp_path): + """从素材库下载 BGM。""" + plan = FakePlan( + config={ + "bgm": { + "enabled": True, + "asset_id": "asset_bgm_001", + } + } + ) + + mock_db = MagicMock() + mock_asset = MagicMock() + mock_asset.file_url = "oss://bgm/sample.mp3" + mock_query = MagicMock() + mock_query.first.return_value = mock_asset + mock_db.query.return_value = mock_query + + adapter = RenderAdapter(mock_db) + + with patch("video_processing.render_adapter.download_asset") as mock_download: + + def _fake_download(storage_key, local_path): + Path(local_path).parent.mkdir(parents=True, exist_ok=True) + Path(local_path).write_bytes(b"fake bgm") + return True + + mock_download.side_effect = _fake_download + + result = adapter._prepare_bgm(plan, tmp_path, "plan_001") + + assert result is not None + assert Path(result).exists() + + def test_bgm_from_preset_library(self, tmp_path): + """从预设 BGM 库下载。""" + plan = FakePlan( + config={ + "bgm": { + "enabled": True, + "preset_id": "preset_calm", + } + } + ) + adapter = RenderAdapter(MagicMock()) + + with patch("packages.domain.preset_bgm.get_preset_bgm") as mock_preset: + mock_preset.return_value = MagicMock(audio_url="https://cdn.example.com/preset_calm.mp3") + + with patch("video_processing.url_security.safe_download_file") as mock_download: + + def _fake_download(url, path, **kwargs): + Path(path).parent.mkdir(parents=True, exist_ok=True) + Path(path).write_bytes(b"preset bgm data") + + mock_download.side_effect = _fake_download + + result = adapter._prepare_bgm(plan, tmp_path, "plan_001") + + assert result is not None + assert Path(result).exists() + + def test_bgm_url_takes_priority_over_asset(self, tmp_path): + """URL 优先级高于素材库。""" + plan = FakePlan( + config={ + "bgm": { + "enabled": True, + "audio_url": "https://example.com/bgm.mp3", + "asset_id": "asset_bgm_001", + } + } + ) + adapter = RenderAdapter(MagicMock()) + + with patch("video_processing.url_security.safe_download_file") as mock_url_download: + + def _fake_url_download(url, path, **kwargs): + Path(path).parent.mkdir(parents=True, exist_ok=True) + Path(path).write_bytes(b"url bgm") + + mock_url_download.side_effect = _fake_url_download + + with patch("video_processing.render_adapter.download_asset") as mock_asset_download: + result = adapter._prepare_bgm(plan, tmp_path, "plan_001") + + assert result is not None + mock_url_download.assert_called_once() + mock_asset_download.assert_not_called() # URL 成功后不会再走素材库 + + +# ── _get_asr_service 测试 ──────────────────────────────────────────────────── + + +class TestGetAsrService: + """ASR 服务初始化测试。""" + + def test_returns_service_when_available(self): + """ASR 服务可用时返回实例。""" + mock_service = MagicMock() + with patch("services.asr_service_factory.get_asr_service") as mock_get: + mock_get.return_value = mock_service + result = RenderAdapter._get_asr_service() + + assert result is mock_service + + def test_returns_none_when_import_fails(self): + """ASR 服务导入失败时返回 None(不阻断主流程)。""" + with patch("services.asr_service_factory.get_asr_service") as mock_get: + mock_get.side_effect = ImportError("asr module not found") + result = RenderAdapter._get_asr_service() + + assert result is None + + def test_returns_none_when_init_fails(self): + """ASR 服务初始化失败时返回 None(不阻断主流程)。""" + with patch("services.asr_service_factory.get_asr_service") as mock_get: + mock_get.side_effect = RuntimeError("ASR init failed") + result = RenderAdapter._get_asr_service() + + assert result is None + + +# ── render_plan BGM/ASR 集成测试 ───────────────────────────────────────────── + + +class TestRenderPlanWithBgmAsr: + """渲染流程中 BGM 和 ASR 的集成测试。""" + + @patch("video_processing.render_adapter.upload_to_oss") + @patch("video_processing.render_adapter.UnifiedRenderService") + @patch("video_processing.render_adapter.download_asset") + def test_bgm_passed_to_render_service(self, mock_download, mock_render_cls, mock_upload, tmp_path): + """BGM 路径被正确传递给 UnifiedRenderService。""" + + def _fake_download(asset_id, local_path): + local_path.parent.mkdir(parents=True, exist_ok=True) + local_path.write_bytes(b"fake video") + return True + + mock_download.side_effect = _fake_download + + mock_render = MagicMock() + mock_render.render.return_value = MagicMock( + output_path=tmp_path / "out.mp4", + duration=5.0, + file_size=1024, + width=1280, + height=720, + ) + mock_render_cls.return_value = mock_render + mock_upload.return_value = "https://oss.example.com/out.mp4" + + plan = FakePlan( + id="plan_bgm", + config={ + "bgm": { + "enabled": True, + "preset_id": "preset_001", + } + }, + ) + clips = [_make_clip("c1", order=0, duration=5.0)] + + adapter, _, _ = _make_adapter(plan=plan, clips=clips) + + # Mock BGM 下载 + with patch.object(adapter, "_prepare_bgm", return_value=str(tmp_path / "bgm.mp3")): + with patch.object(adapter, "_get_asr_service", return_value=None): + result = adapter.render_plan( + "plan_bgm", + work_dir=tmp_path / "work", + ) + + assert result.success + # 验证 UnifiedRenderService 收到了 bgm_path + call_kwargs = mock_render_cls.call_args + assert call_kwargs.kwargs["bgm_path"] == str(tmp_path / "bgm.mp3") + + @patch("video_processing.render_adapter.upload_to_oss") + @patch("video_processing.render_adapter.UnifiedRenderService") + @patch("video_processing.render_adapter.download_asset") + def test_asr_service_passed_to_render_service(self, mock_download, mock_render_cls, mock_upload, tmp_path): + """ASR 服务被正确传递给 UnifiedRenderService。""" + + def _fake_download(asset_id, local_path): + local_path.parent.mkdir(parents=True, exist_ok=True) + local_path.write_bytes(b"fake video") + return True + + mock_download.side_effect = _fake_download + + mock_render = MagicMock() + mock_render.render.return_value = MagicMock( + output_path=tmp_path / "out.mp4", + duration=5.0, + file_size=1024, + width=1280, + height=720, + ) + mock_render_cls.return_value = mock_render + mock_upload.return_value = "https://oss.example.com/out.mp4" + + mock_asr = MagicMock() + plan = FakePlan(id="plan_asr") + clips = [_make_clip("c1", order=0, duration=5.0)] + + adapter, _, _ = _make_adapter(plan=plan, clips=clips) + + with patch.object(adapter, "_prepare_bgm", return_value=None): + with patch.object(adapter, "_get_asr_service", return_value=mock_asr): + result = adapter.render_plan( + "plan_asr", + work_dir=tmp_path / "work", + ) + + assert result.success + call_kwargs = mock_render_cls.call_args + assert call_kwargs.kwargs["asr_service"] is mock_asr + + @patch("video_processing.render_adapter.upload_to_oss") + @patch("video_processing.render_adapter.UnifiedRenderService") + @patch("video_processing.render_adapter.download_asset") + def test_rendered_and_failed_clip_ids_in_result(self, mock_download, mock_render_cls, mock_upload, tmp_path): + """渲染结果中包含成功和失败的 clip id 列表。""" + # 成功/失败映射,通过 asset_id 区分 + download_map = { + "good_001.mp4": True, + "bad_002.mp4": False, + "good_003.mp4": True, + } + + def _fake_download(asset_id, local_path): + local_path.parent.mkdir(parents=True, exist_ok=True) + ok = download_map.get(asset_id, False) + if ok: + local_path.write_bytes(b"fake video") + return ok + + mock_download.side_effect = _fake_download + + mock_render = MagicMock() + mock_render.render.return_value = MagicMock( + output_path=tmp_path / "out.mp4", + duration=3.0, + file_size=512, + width=1280, + height=720, + ) + mock_render_cls.return_value = mock_render + mock_upload.return_value = "https://oss.example.com/out.mp4" + + plan = FakePlan(id="plan_mixed") + clips = [ + _make_clip("c_good1", order=0, asset_id="good_001.mp4"), + _make_clip("c_bad", order=1, asset_id="bad_002.mp4"), + _make_clip("c_good2", order=2, asset_id="good_003.mp4"), + ] + + adapter, _, _ = _make_adapter(plan=plan, clips=clips) + + result = adapter.render_plan( + "plan_mixed", + work_dir=tmp_path / "work", + ) + + assert result.success + assert "c_good1" in result.rendered_clip_ids + assert "c_good2" in result.rendered_clip_ids + assert "c_bad" in result.failed_clip_ids + assert len(result.rendered_clip_ids) == 2 + assert len(result.failed_clip_ids) == 1