fix: 一键生成 P0 修复 + P1 校验 #203
@@ -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:
|
||||
|
||||
@@ -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]:
|
||||
"""探测输出文件的时长、大小、宽高。
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
@@ -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")
|
||||
Reference in New Issue
Block a user