fix: 封面选帧输入源改为最终成片,预览片段作为兼容回退 #1534

Merged
xiaoxia merged 7 commits from fix/cover-from-final-video into develop 2026-08-29 01:19:02 +08:00
2 changed files with 839 additions and 81 deletions
+326 -76
View File
@@ -1,15 +1,18 @@
"""封面生成路由 — Generation 模块.
端点:
- POST /generate-cover AI 生成封面(从预览视频中抽帧)
- POST /generate-cover AI 生成封面(从最终成片视频中抽帧,兼容预览片段回退
挂载路径: /api/v1/generation/generate-cover
"""
from __future__ import annotations
import ipaddress
import logging
import re
from typing import Any, List, Optional
from urllib.parse import urlparse
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session, get_generated_video_repository
@@ -24,6 +27,7 @@ from packages.adapters.sqlalchemy_impl.generation_task_repository import (
)
from packages.application import ListGeneratedVideosByTaskUseCase
from packages.domain.config_schemas import normalize_plan_config
from packages.shared.storage import get_shared_storage_service
from .templates_editor.dependencies import get_draft_plan_id, get_editor_services
@@ -51,6 +55,14 @@ class GenerateCoverRequest(BaseModel):
default=None,
description="上传的封面图片 URL,仅 cover_type=upload 时有效",
)
generated_video_id: Optional[str] = Field(
default=None,
description="确认生成产出的最终视频 ID。传入后封面从该视频文件抽帧,而非预览片段。",
)
video_url: Optional[str] = Field(
default=None,
description="最终视频 URL(兜底)。当 generated_video_id 不可用时,直接从此 URL 对应的视频抽帧。",
)
class GenerateCoverResponse(BaseModel):
@@ -121,8 +133,6 @@ def _persist_cover_frame(
exc_info=True,
)
from packages.shared.storage import get_shared_storage_service
storage = get_shared_storage_service()
cover_key = f"covers/{plan_id}/cover_{uuid.uuid4().hex[:8]}.jpg"
storage.upload_file(
@@ -140,6 +150,106 @@ def _persist_cover_frame(
Path(tmp_path).unlink(missing_ok=True)
def _get_task_video_url(db: Session, task_id: str) -> Optional[str]:
"""从 GenerationTask 关联的 GeneratedVideo 中获取视频 storage_key / URL."""
try:
video_repo = get_generated_video_repository(db)
use_case = ListGeneratedVideosByTaskUseCase(video_repo)
videos = use_case.execute(task_id)
if videos:
return getattr(videos[0], "file_url", "") or ""
except Exception:
logger.warning("[封面生成] 获取任务视频失败: task_id=%s", task_id, exc_info=True)
return None
def _resolve_storage_key_to_url(storage_key: str) -> Optional[str]:
"""将 storage_key 或完整 URL 转换为可访问的裸 URL。"""
if not storage_key:
return None
try:
if storage_key.startswith("http"):
url = storage_key
else:
storage_svc = get_shared_storage_service()
url = storage_svc.get_url(storage_key)
if url:
url = re.sub(r"(?<!:)//", "/", url)
return url
except Exception as e:
logger.warning("[封面生成] storage_key 转 URL 失败: key=%s err=%s", storage_key, e)
return None
def _endpoint_host(value: str) -> str:
"""从 endpoint / URL 字符串中安全提取主机名(兼容有无 scheme 两种配置)。"""
v = (value or "").strip().lower()
if not v:
return ""
if "://" in v:
return (urlparse(v).hostname or "").lower()
# 无 scheme:去掉可能的端口(host:port),urlparse 补 // 以正确解析
return (urlparse("//" + v).hostname or "").lower()
def _is_private_or_reserved_host(host: str) -> bool:
"""判断主机名是否为内网/回环/链路本地/保留地址(IPv4 与 IPv6 统一处理)。
使用标准库 ipaddress 判定;非 IP 主机名(如 localhost)单独处理。
"""
h = host.strip().lower()
if h in {"localhost", "0.0.0.0", "::", "::1"}:
return True
try:
addr = ipaddress.ip_address(h)
# is_private 覆盖 10/8、172.16/12、192.168/16、127/8、169.254/16、
# ::1、fc00::/7、fe80::/10 等全部私有/保留段
return bool(addr.is_private or addr.is_loopback or addr.is_link_local or addr.is_reserved)
except ValueError:
return False
def _is_trusted_media_url(url: str) -> bool:
"""校验 URL 是否指向受信任的存储域名(OSS bucket / 本地存储),防止 SSRF。
用户可通过 video_url 传入视频地址,但服务端(MediaKit)会主动请求该 URL
因此必须限制为自家存储域名,拒绝内网地址、元数据地址等任意主机。
"""
if not url:
return False
try:
parsed = urlparse(url.strip())
if parsed.scheme not in ("http", "https"):
return False
host = (parsed.hostname or "").lower()
if not host:
return False
# 拒绝一切内网/回环/链路本地/保留地址(IPv4 + IPv6,标准库判定)
if _is_private_or_reserved_host(host):
return False
# 允许:自家 OSS bucket 域名(<bucket>.<endpoint>)或 endpoint 自身及其子域
try:
storage_svc = get_shared_storage_service()
trusted_hosts = set()
public_base = getattr(storage_svc, "public_url", "") or ""
h1 = _endpoint_host(public_base)
if h1:
trusted_hosts.add(h1)
h2 = _endpoint_host(getattr(storage_svc, "endpoint", "") or "")
if h2:
trusted_hosts.add(h2)
for trusted in trusted_hosts:
if host == trusted or host.endswith("." + trusted):
return True
except Exception:
logger.warning("[封面生成] 存储域名白名单初始化失败,URL 校验从严拒绝", exc_info=True)
return False
return False
except Exception:
logger.warning("[封面生成] video_url 白名单校验异常,从严拒绝: url=%s", url[:80], exc_info=True)
return False
@router.post("/generate-cover", response_model=GenerateCoverResponse)
def generate_cover(
body: GenerateCoverRequest,
@@ -149,12 +259,16 @@ def generate_cover(
db: Session = Depends(get_db_session),
current_user: AuthenticatedUser = Depends(get_current_user),
) -> GenerateCoverResponse:
"""AI 生成封面 — 从预览视频中抽帧.
"""AI 生成封面 — 优先从最终成片视频中抽帧,回退到预览片段.
流程(串行):
1. 预览视频已渲染完成(通过 3 步查找获取 URL)
2. 用裸 URL 让 MediaKit 下载视频并抽帧
3. 帧图下载后上传到 OSS covers/ 路径
1. 优先使用前端传入的 generation_task_id 定位最终成片任务,
或自动查找 plan 关联的已完成最终成片任务(is_preview=False
2. 回退:从预览片段获取视频 URL(兼容旧流程)
3. 用裸 URL 让 MediaKit 下载视频并抽帧
4. 帧图下载后上传到 OSS covers/ 路径
MediaKit 的调用方式(strategy / max_frames / 轮询 / 重试 / 降级)不变。
"""
_, plan_svc = services
plan = plan_svc.get_plan_or_raise(plan_id)
@@ -182,27 +296,111 @@ def generate_cover(
)
return GenerateCoverResponse(plan_id=plan_id, cover=cover_data)
# ── 3 步查找预览视频 URL ──────────────────────────────────────────
# 第一步:从 plan.config 读取
# ── 查找用于抽帧的视频 URL ────────────────────────────────────────
# 优先级:
# 0. 请求体显式传入的 generation_task_id(最终成片任务)
# 1. plan.config.rendered_storage_key
# 2. plan.config.generation_task_id 对应的任务
# 3. source_edit_plan_id 关联的已完成「最终成片」任务(is_preview=False
# 4. source_edit_plan_id 关联的已完成预览任务(is_preview=True,兼容回退)
# 5. user + template 最近的已完成预览任务(兜底)
logger.info("[封面生成] 步骤1: 从 plan.config 查找 rendered_storage_key: plan_id=%s", plan_id)
rendered_storage_key = (plan.config or {}).get("rendered_storage_key", "")
# 第二步:如果还没有,通过 generation_task_id 查找预览任务的产物
# 步骤 0:请求体传入最终视频标识(generated_video_id 或 video_url
if not rendered_storage_key:
# 0a:通过 generated_video_id 查找最终成片视频
if body.generated_video_id:
logger.info(
"[封面生成] 步骤0a: 使用 generated_video_id: plan_id=%s video_id=%s",
plan_id,
body.generated_video_id,
)
try:
gv_repo = get_generated_video_repository(db)
gv = gv_repo.get(body.generated_video_id)
if gv:
file_url = getattr(gv, "file_url", "") or ""
if file_url:
# 权限校验(双重,任何一层确认归属不符即拒绝):
# 1) GeneratedVideo.user_id 直接归属(老数据可能为空,为空时不据此放行)
gv_owner = (getattr(gv, "user_id", "") or "").strip()
if gv_owner and gv_owner != current_user.user.id:
raise HTTPException(status_code=403, detail="无权访问该视频")
# 2) 关联 generation_task 归属校验;关联任务缺失时不可静默放行:
# 若 video 自身无 owner 信息且关联任务也查不到,拒绝访问
gv_task_id = getattr(gv, "generation_task_id", "") or ""
task0 = None
if gv_task_id:
try:
task0 = SQLAlchemyGenerationTaskRepository(db).get(gv_task_id)
except Exception:
logger.warning(
"[封面生成] 步骤0a关联任务查询异常: plan_id=%s task_id=%s",
plan_id,
gv_task_id,
exc_info=True,
)
if task0 is not None:
task_owner = (getattr(task0, "created_by_user_id", "") or "").strip()
if task_owner and task_owner != current_user.user.id:
raise HTTPException(status_code=403, detail="无权访问该视频")
elif not gv_owner:
# video 无 owner 且关联任务不存在/无法确认归属 → 拒绝,防止越权
logger.warning(
"[封面生成] 步骤0a视频归属无法确认,拒绝访问: plan_id=%s video_id=%s",
plan_id,
body.generated_video_id,
)
raise HTTPException(status_code=403, detail="无权访问该视频")
rendered_storage_key = file_url
logger.info(
"[封面生成] ✅ 步骤0a找到最终成片: plan_id=%s video_id=%s url=%s",
plan_id,
body.generated_video_id,
file_url[:80],
)
except HTTPException:
raise
except Exception:
logger.warning(
"[封面生成] 步骤0a查找视频失败: plan_id=%s video_id=%s",
plan_id,
body.generated_video_id,
exc_info=True,
)
# 0b:直接使用 video_url(兜底)— 必须通过存储域名白名单校验,防止 SSRF
if not rendered_storage_key and body.video_url:
if _is_trusted_media_url(body.video_url):
logger.info(
"[封面生成] 步骤0b: 使用请求体传入的 video_url(白名单通过): plan_id=%s url=%s",
plan_id,
body.video_url[:80],
)
rendered_storage_key = body.video_url
else:
logger.warning(
"[封面生成] 步骤0b: video_url 不在受信任存储域名白名单内,已忽略: plan_id=%s url=%s",
plan_id,
body.video_url[:80],
)
# 步骤 2:通过 plan.config.generation_task_id 查找
if not rendered_storage_key:
generation_task_id = (plan.config or {}).get("generation_task_id", "")
logger.info(
"[封面生成] 步骤2: 通过 generation_task_id 查找: plan_id=%s task_id=%s", plan_id, generation_task_id
)
if generation_task_id:
logger.info(
"[封面生成] 步骤2: 通过 plan.config.generation_task_id 查找: plan_id=%s task_id=%s",
plan_id,
generation_task_id,
)
try:
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
task = gen_task_repo.get(generation_task_id)
_repo = SQLAlchemyGenerationTaskRepository(db)
task = _repo.get(generation_task_id)
if task:
video_repo = get_generated_video_repository(db)
use_case = ListGeneratedVideosByTaskUseCase(video_repo)
videos = use_case.execute(task.id)
if videos:
rendered_storage_key = getattr(videos[0], "file_url", "") or ""
rendered_storage_key = _get_task_video_url(db, task.id) or ""
if rendered_storage_key:
logger.info(
"[封面生成] ✅ 步骤2找到视频: plan_id=%s task_id=%s url=%s",
plan_id,
@@ -211,26 +409,23 @@ def generate_cover(
)
except Exception:
logger.warning(
"封面生成: 通过 generation_task_id 查找视频失败: plan_id=%s",
"[封面生成] 步骤2查找失败: plan_id=%s",
plan_id,
exc_info=True,
)
# 第 2.5 步:通过 plan_id 作为 source_edit_plan_id 查找关联的已完成预览任务
# 步骤 3:通过 source_edit_plan_id 查找已完成「最终成片」任务(is_preview=False
if not rendered_storage_key:
try:
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
logger.info("[封面生成] 步骤2.5: 通过 source_edit_plan_id 查找: plan_id=%s", plan_id)
preview_tasks = gen_task_repo.list_by_source_edit_plan(plan_id)
for pt in preview_tasks:
if getattr(pt, "status", "") == "completed" and getattr(pt, "is_preview", False):
video_repo = get_generated_video_repository(db)
use_case = ListGeneratedVideosByTaskUseCase(video_repo)
videos = use_case.execute(pt.id)
if videos:
rendered_storage_key = getattr(videos[0], "file_url", "") or ""
_repo = SQLAlchemyGenerationTaskRepository(db)
logger.info("[封面生成] 步骤3: 查找最终成片任务(is_preview=False): plan_id=%s", plan_id)
all_tasks = _repo.list_by_source_edit_plan(plan_id)
for pt in all_tasks:
if getattr(pt, "status", "") == "completed" and not getattr(pt, "is_preview", False):
rendered_storage_key = _get_task_video_url(db, pt.id) or ""
if rendered_storage_key:
logger.info(
"[封面生成] ✅ 步骤2.5找到视频: plan_id=%s task_id=%s url=%s",
"[封面生成] ✅ 步骤3找到最终成片: plan_id=%s task_id=%s url=%s",
plan_id,
pt.id,
rendered_storage_key[:80],
@@ -238,66 +433,74 @@ def generate_cover(
break
except Exception:
logger.warning(
"封面生成: 通过 source_edit_plan_id 查找预览任务失败: plan_id=%s",
"[封面生成] 步骤3查找最终成片失败: plan_id=%s",
plan_id,
exc_info=True,
)
# 第三步:按 user + template 查找最近的已完成预览任务(兜底)
# 步骤 4:兼容回退 — 通过 source_edit_plan_id 查找已完成预览任务
if not rendered_storage_key:
try:
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
logger.info("[封面生成] 步骤3: 通过 user+template 查找: plan_id=%s template_id=%s", plan_id, template_id)
preview_tasks = gen_task_repo.list_latest_completed_preview(
_repo = SQLAlchemyGenerationTaskRepository(db)
logger.info("[封面生成] 步骤4: 回退查找预览任务(is_preview=True): plan_id=%s", plan_id)
preview_tasks = _repo.list_by_source_edit_plan(plan_id)
for pt in preview_tasks:
if getattr(pt, "status", "") == "completed" and getattr(pt, "is_preview", False):
rendered_storage_key = _get_task_video_url(db, pt.id) or ""
if rendered_storage_key:
logger.info(
"[封面生成] ✅ 步骤4找到预览视频: plan_id=%s task_id=%s url=%s",
plan_id,
pt.id,
rendered_storage_key[:80],
)
break
except Exception:
logger.warning(
"[封面生成] 步骤4查找预览任务失败: plan_id=%s",
plan_id,
exc_info=True,
)
# 步骤 5:按 user + template 查找最近的已完成预览任务(兜底)
if not rendered_storage_key:
try:
_repo = SQLAlchemyGenerationTaskRepository(db)
logger.info(
"[封面生成] 步骤5: 通过 user+template 查找预览任务: plan_id=%s template_id=%s",
plan_id,
template_id,
)
preview_tasks = _repo.list_latest_completed_preview(
user_id=str(current_user.user.id),
template_id=template_id,
)
if preview_tasks:
completed_preview = preview_tasks[0]
video_repo = get_generated_video_repository(db)
use_case = ListGeneratedVideosByTaskUseCase(video_repo)
videos = use_case.execute(completed_preview.id)
if videos:
rendered_storage_key = getattr(videos[0], "file_url", "") or ""
rendered_storage_key = _get_task_video_url(db, preview_tasks[0].id) or ""
if rendered_storage_key:
logger.info(
"封面视频: 通过 user+template 找到预览任务: plan_id=%s template_id=%s task_id=%s",
"[封面生成] ✅ 步骤5找到预览视频: plan_id=%s task_id=%s",
plan_id,
template_id,
completed_preview.id,
preview_tasks[0].id,
)
except Exception:
logger.warning(
"封面警告: user+template 查找预览任务失败: plan_id=%s template_id=%s",
"[封面生成] 步骤5 user+template 查找失败: plan_id=%s",
plan_id,
template_id,
exc_info=True,
)
# 使用裸 URLrendered/* 已配置公开读);找不到渲染视频时不立即报错,
# 因为步骤 E 可以直接从源素材抽帧(历史数据或 Worker 抽帧失败时的兜底)
# 将 storage_key 转换为可访问 URL;找不到视频时不立即报错,
# 因为步骤 E2 可以直接从源素材抽帧(历史数据或 Worker 抽帧失败时的兜底)
primary_video_url = None
if rendered_storage_key:
plan_svc.update_plan_config(plan_id, {"rendered_storage_key": rendered_storage_key})
try:
if rendered_storage_key.startswith("http"):
primary_video_url = rendered_storage_key
else:
from packages.shared.storage import get_shared_storage_service
storage_svc = get_shared_storage_service()
primary_video_url = storage_svc.get_url(rendered_storage_key)
if primary_video_url:
import re as _re
primary_video_url = _re.sub(r"(?<!:)//", "/", primary_video_url)
logger.info(
"获取预览视频URL用于封面生成: plan_id=%s url=%s",
plan_id,
primary_video_url[:80] if primary_video_url else "",
)
except Exception as e:
logger.warning("获取预览视频URL失败: plan_id=%s err=%s", plan_id, e)
primary_video_url = None
primary_video_url = _resolve_storage_key_to_url(rendered_storage_key)
logger.info(
"[封面生成] 封面抽帧视频URL: plan_id=%s url=%s",
plan_id,
primary_video_url[:80] if primary_video_url else "",
)
# 统一封面管道:优先从 GenerationTask.cover_url 读取渲染后视频抽帧的封面
# 多步查找 cover_url,和查找视频 URL 一样的 fallback 逻辑
@@ -326,20 +529,67 @@ def generate_cover(
exc_info=True,
)
# 步骤 B:通过 source_edit_plan_id 查找关联预览任务的 cover_url
# 步骤 A2:通过 generated_video_id 查找关联任务的 cover_url
if not cover_url_from_task and body.generated_video_id:
try:
gv_repo = get_generated_video_repository(db)
gv = gv_repo.get(body.generated_video_id)
if gv:
gv_task_id = getattr(gv, "generation_task_id", "") or ""
if gv_task_id:
task_a2 = gen_task_repo.get(gv_task_id)
if task_a2 and getattr(task_a2, "cover_url", ""):
cover_url_from_task = task_a2.cover_url
logger.info(
"[封面生成] 封面(步骤A2-video-task): plan_id=%s video_id=%s url=%s",
plan_id,
body.generated_video_id,
cover_url_from_task[:80],
)
except Exception:
logger.warning(
"[封面生成] 步骤A2读取 cover_url 失败: plan_id=%s video_id=%s",
plan_id,
body.generated_video_id,
exc_info=True,
)
# 步骤 B:通过 source_edit_plan_id 查找关联任务的 cover_url
# 优先最终成片任务(is_preview=False),其次预览任务
if not cover_url_from_task:
try:
preview_tasks = gen_task_repo.list_by_source_edit_plan(plan_id)
for pt in preview_tasks:
if getattr(pt, "status", "") == "completed" and getattr(pt, "cover_url", ""):
all_tasks = gen_task_repo.list_by_source_edit_plan(plan_id)
# 先找最终成片
for pt in all_tasks:
if (
getattr(pt, "status", "") == "completed"
and not getattr(pt, "is_preview", False)
and getattr(pt, "cover_url", "")
):
cover_url_from_task = pt.cover_url
logger.info(
"[封面生成] 统一管道封面(步骤B-source_plan): plan_id=%s task_id=%s url=%s",
"[封面生成] 封面(步骤B-final): plan_id=%s task_id=%s url=%s",
plan_id,
pt.id,
cover_url_from_task[:80],
)
break
# 再找预览
if not cover_url_from_task:
for pt in all_tasks:
if (
getattr(pt, "status", "") == "completed"
and getattr(pt, "is_preview", False)
and getattr(pt, "cover_url", "")
):
cover_url_from_task = pt.cover_url
logger.info(
"[封面生成] 封面(步骤B-preview): plan_id=%s task_id=%s url=%s",
plan_id,
pt.id,
cover_url_from_task[:80],
)
break
except Exception:
logger.warning(
"[封面生成] 步骤B查找 cover_url 失败: plan_id=%s",
+513 -5
View File
@@ -181,7 +181,7 @@ class TestUnifiedCoverPipelineEndpoint:
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("packages.shared.storage.get_shared_storage_service") as mock_storage_getter,
patch("app.api.routes.generation_cover.get_shared_storage_service") as mock_storage_getter,
patch("packages.shared.mediakit_client.get_mediakit_client") as mock_mk_getter,
):
mock_repo = MagicMock()
@@ -449,7 +449,7 @@ class TestUnifiedCoverPipelineEndpoint:
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("packages.shared.storage.get_shared_storage_service") as mock_storage_getter,
patch("app.api.routes.generation_cover.get_shared_storage_service") as mock_storage_getter,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
@@ -590,7 +590,7 @@ class TestSourceEditPlanFallback:
patch("app.api.routes.generation_cover.get_generated_video_repository") as mock_video_repo,
patch("app.api.routes.generation_cover.ListGeneratedVideosByTaskUseCase") as mock_usecase_cls,
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
patch("packages.shared.storage.get_shared_storage_service") as mock_storage_getter,
patch("app.api.routes.generation_cover.get_shared_storage_service") as mock_storage_getter,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
@@ -657,7 +657,7 @@ class TestSourceEditPlanFallback:
patch("app.api.routes.generation_cover.get_generated_video_repository") as mock_video_repo,
patch("app.api.routes.generation_cover.ListGeneratedVideosByTaskUseCase") as mock_usecase_cls,
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
patch("packages.shared.storage.get_shared_storage_service") as mock_storage_getter,
patch("app.api.routes.generation_cover.get_shared_storage_service") as mock_storage_getter,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
@@ -1177,7 +1177,7 @@ class TestUploadCoverType:
return_value=mock_asset_repo,
),
patch("packages.shared.mediakit_client.get_mediakit_client", return_value=mock_mk),
patch("packages.shared.storage.get_shared_storage_service") as mock_storage_getter,
patch("app.api.routes.generation_cover.get_shared_storage_service") as mock_storage_getter,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
@@ -1200,3 +1200,511 @@ class TestUploadCoverType:
assert exc_info.value.status_code == 400
mock_mk.extract_frames.assert_not_called()
class TestCoverFromFinalVideo:
"""测试封面从最终成片任务(is_preview=False)获取视频源。"""
def test_generated_video_fields_in_schema(self):
"""请求体支持 generated_video_id 和 video_url 字段。"""
from app.api.routes.generation_cover import GenerateCoverRequest
req = GenerateCoverRequest(
generated_video_id="gv-001",
video_url="https://example.com/final.mp4",
)
assert req.generated_video_id == "gv-001"
assert req.video_url == "https://example.com/final.mp4"
# 默认 None
req_default = GenerateCoverRequest()
assert req_default.generated_video_id is None
assert req_default.video_url is None
def test_cover_uses_final_video_when_generated_video_id_provided(self):
"""传 generated_video_id 时,从该最终成片视频抽帧。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
mock_plan.config = {}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
# Generated video
mock_gv = MagicMock()
mock_gv.file_url = "rendered/final/video.mp4"
mock_gv.generation_task_id = "task-final-001"
mock_gv.user_id = "user-1"
# 最终成片任务
mock_final_task = MagicMock()
mock_final_task.id = "task-final-001"
mock_final_task.created_by_user_id = "user-1"
mock_final_task.cover_url = ""
mock_gv_repo = MagicMock()
mock_gv_repo.get.return_value = mock_gv
mock_db = MagicMock()
mock_current_user = MagicMock()
mock_current_user.user.id = "user-1"
body = GenerateCoverRequest(
cover_type="ai_frame",
generated_video_id="gv-final-001",
)
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch(
"app.api.routes.generation_cover.get_generated_video_repository",
return_value=mock_gv_repo,
),
patch("app.api.routes.generation_cover.get_shared_storage_service") as mock_storage_getter,
patch("packages.shared.mediakit_client.get_mediakit_client") as mock_mk_getter,
patch(
"app.api.routes.generation_cover._persist_cover_frame",
return_value="https://oss.example.com/covers/final-cover.jpg",
) as mock_persist,
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = mock_final_task
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_storage_svc = MagicMock()
mock_storage_svc.get_url.return_value = "https://oss.example.com/rendered/final/video.mp4"
mock_storage_svc.public_url = "https://oss.example.com"
mock_storage_svc.endpoint = "oss.example.com"
mock_storage_getter.return_value = mock_storage_svc
mock_mk = MagicMock()
mock_mk.is_available = True
mock_mk.extract_frames.return_value = [{"image_url": "https://mk/frame.jpg"}]
mock_mk_getter.return_value = mock_mk
mock_normalize.return_value = {
"cover": {"type": "ai_frame", "image_url": "https://oss.example.com/covers/final-cover.jpg"}
}
from app.api.routes.generation_cover import generate_cover
result = generate_cover(
body=body,
template_id="tpl-1",
plan_id="plan-final",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=mock_current_user,
)
assert result.cover["image_url"] == "https://oss.example.com/covers/final-cover.jpg"
call_kwargs = mock_mk.extract_frames.call_args.kwargs
assert "rendered/final/video.mp4" in call_kwargs["video_url"]
assert call_kwargs["strategy"] == "SpecifiedFrames"
assert call_kwargs["max_frames"] == 1
assert call_kwargs["max_retries"] == 0
mock_persist.assert_called_once()
def test_cover_uses_video_url_directly(self):
"""传 video_url 时,直接从该 URL 抽帧。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
mock_plan.config = {}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
mock_current_user = MagicMock()
mock_current_user.user.id = "user-1"
body = GenerateCoverRequest(
cover_type="ai_frame",
video_url="https://oss.example.com/rendered/final/video.mp4",
)
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("app.api.routes.generation_cover.get_shared_storage_service") as mock_storage_getter,
patch("packages.shared.mediakit_client.get_mediakit_client") as mock_mk_getter,
patch(
"app.api.routes.generation_cover._persist_cover_frame",
return_value="https://oss.example.com/covers/c.jpg",
),
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_storage_svc = MagicMock()
mock_storage_svc.public_url = "https://oss.example.com"
mock_storage_svc.endpoint = "oss.example.com"
mock_storage_getter.return_value = mock_storage_svc
mock_mk = MagicMock()
mock_mk.is_available = True
mock_mk.extract_frames.return_value = [{"image_url": "https://mk/f.jpg"}]
mock_mk_getter.return_value = mock_mk
mock_normalize.return_value = {
"cover": {"type": "ai_frame", "image_url": "https://oss.example.com/covers/c.jpg"}
}
from app.api.routes.generation_cover import generate_cover
generate_cover(
body=body,
template_id="tpl-1",
plan_id="plan-url",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=mock_current_user,
)
call_kwargs = mock_mk.extract_frames.call_args.kwargs
assert "rendered/final/video.mp4" in call_kwargs["video_url"]
def test_cover_prefers_final_task_over_preview_in_source_plan(self):
"""步骤3source_edit_plan 关联任务中,优先使用 is_preview=False 的最终成片。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
mock_plan.config = {} # 无 rendered_storage_key / generation_task_id
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
# 一个预览任务 + 一个最终成片任务
mock_preview = MagicMock()
mock_preview.id = "task-preview"
mock_preview.status = "completed"
mock_preview.is_preview = True
mock_preview.cover_url = ""
mock_final = MagicMock()
mock_final.id = "task-final"
mock_final.status = "completed"
mock_final.is_preview = False
mock_final.cover_url = ""
mock_video_preview = MagicMock()
mock_video_preview.file_url = "rendered/preview/video.mp4"
mock_video_final = MagicMock()
mock_video_final.file_url = "rendered/final/video.mp4"
mock_db = MagicMock()
mock_current_user = MagicMock()
mock_current_user.user.id = "user-1"
body = GenerateCoverRequest(cover_type="ai_frame")
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("app.api.routes.generation_cover.ListGeneratedVideosByTaskUseCase") as mock_list_videos,
patch("app.api.routes.generation_cover.get_shared_storage_service") as mock_storage_getter,
patch("packages.shared.mediakit_client.get_mediakit_client") as mock_mk_getter,
patch(
"app.api.routes.generation_cover._persist_cover_frame",
return_value="https://oss.example.com/covers/c.jpg",
),
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
# list_by_source_edit_plan 返回 [preview, final],最终成片排在后面
mock_repo.list_by_source_edit_plan.return_value = [mock_preview, mock_final]
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
# 根据 task id 返回不同 video
def get_videos(task_id):
if task_id == "task-final":
return [mock_video_final]
return [mock_video_preview]
mock_use_case = MagicMock()
mock_use_case.execute.side_effect = get_videos
mock_list_videos.return_value = mock_use_case
mock_storage_svc = MagicMock()
mock_storage_svc.get_url.side_effect = lambda key: f"https://oss.example.com/{key}"
mock_storage_svc.public_url = "https://oss.example.com"
mock_storage_svc.endpoint = "oss.example.com"
mock_storage_getter.return_value = mock_storage_svc
mock_mk = MagicMock()
mock_mk.is_available = True
mock_mk.extract_frames.return_value = [{"image_url": "https://mk/f.jpg"}]
mock_mk_getter.return_value = mock_mk
mock_normalize.return_value = {
"cover": {"type": "ai_frame", "image_url": "https://oss.example.com/covers/c.jpg"}
}
from app.api.routes.generation_cover import generate_cover
generate_cover(
body=body,
template_id="tpl-1",
plan_id="plan-priority",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=mock_current_user,
)
# 应该使用 final video URL 抽帧,而非 preview
call_kwargs = mock_mk.extract_frames.call_args.kwargs
assert "rendered/final/video.mp4" in call_kwargs["video_url"]
assert "rendered/preview" not in call_kwargs["video_url"]
def test_cover_generated_video_permission_denied(self):
"""generated_video_id 关联任务属于其他用户时,返回 403。"""
from unittest.mock import MagicMock, patch
import pytest
from app.api.routes.generation_cover import GenerateCoverRequest
from fastapi import HTTPException
mock_plan = MagicMock()
mock_plan.config = {}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_gv = MagicMock()
mock_gv.file_url = "rendered/other/video.mp4"
mock_gv.generation_task_id = "task-other"
mock_gv.user_id = "" # 老数据无 user_id,走关联任务归属校验
mock_other_task = MagicMock()
mock_other_task.created_by_user_id = "other-user"
mock_gv_repo = MagicMock()
mock_gv_repo.get.return_value = mock_gv
mock_db = MagicMock()
mock_current_user = MagicMock()
mock_current_user.user.id = "user-1"
body = GenerateCoverRequest(
cover_type="ai_frame",
generated_video_id="gv-other",
)
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch(
"app.api.routes.generation_cover.get_generated_video_repository",
return_value=mock_gv_repo,
),
):
mock_repo = MagicMock()
mock_repo.get.return_value = mock_other_task
mock_repo_cls.return_value = mock_repo
from app.api.routes.generation_cover import generate_cover
with pytest.raises(HTTPException) as exc_info:
generate_cover(
body=body,
template_id="tpl-1",
plan_id="plan-perm",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=mock_current_user,
)
assert exc_info.value.status_code == 403
def test_cover_video_url_ssrf_blocked(self):
"""video_url 指向内网/非白名单域名时被忽略,不向其发起抽帧请求。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import GenerateCoverRequest
mock_plan = MagicMock()
mock_plan.config = {}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_db = MagicMock()
mock_current_user = MagicMock()
mock_current_user.user.id = "user-1"
# SSRF 攻击载荷:内网元数据地址
body = GenerateCoverRequest(
cover_type="ai_frame",
video_url="http://100.100.100.200/latest/meta-data/",
)
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch("app.api.routes.generation_cover.get_shared_storage_service") as mock_storage_getter,
patch("packages.shared.mediakit_client.get_mediakit_client") as mock_mk_getter,
patch("app.api.routes.generation_cover._persist_cover_frame") as mock_persist,
patch("app.api.routes.generation_cover.normalize_plan_config") as mock_normalize,
):
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_repo.list_by_source_edit_plan.return_value = []
mock_repo.list_latest_completed_preview.return_value = []
mock_repo_cls.return_value = mock_repo
mock_storage_svc = MagicMock()
mock_storage_svc.public_url = "https://xiaoxia-media.oss-cn-hangzhou.aliyuncs.com"
mock_storage_svc.endpoint = "oss-cn-hangzhou.aliyuncs.com"
mock_storage_svc.get_url.side_effect = lambda k: f"https://xiaoxia-media.oss-cn-hangzhou.aliyuncs.com/{k}"
mock_storage_getter.return_value = mock_storage_svc
mock_mk = MagicMock()
mock_mk.is_available = True
mock_mk.extract_frames.return_value = [{"image_url": "https://mk/f.jpg"}]
mock_mk_getter.return_value = mock_mk
mock_normalize.return_value = {"cover": {"type": "ai_frame", "image_url": "https://mk/f.jpg"}}
from app.api.routes.generation_cover import generate_cover
from fastapi import HTTPException
# 内网 URL 被白名单拦截后,无任何可用视频源 → 400(而不是向内网发请求)
with pytest.raises(HTTPException) as exc_info:
generate_cover(
body=body,
template_id="tpl-1",
plan_id="plan-ssrf",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=mock_current_user,
)
assert exc_info.value.status_code == 400
# MediaKit 从未被要求抽帧该内网地址
if mock_mk.extract_frames.called:
called_url = mock_mk.extract_frames.call_args.kwargs.get("video_url", "")
assert "100.100.100.200" not in called_url
assert "meta-data" not in called_url
def test_cover_generated_video_ownership_unverifiable_denied(self):
"""video 无 user_id 且关联任务不存在时,归属无法确认 → 403(防权限绕过)。"""
from unittest.mock import MagicMock, patch
import pytest
from app.api.routes.generation_cover import GenerateCoverRequest
from fastapi import HTTPException
mock_plan = MagicMock()
mock_plan.config = {}
mock_plan_svc = MagicMock()
mock_plan_svc.get_plan_or_raise.return_value = mock_plan
mock_template_svc = MagicMock()
mock_gv = MagicMock()
mock_gv.file_url = "rendered/mystery/video.mp4"
mock_gv.generation_task_id = "task-gone" # 关联任务已删除
mock_gv.user_id = "" # 老数据无 owner
mock_gv_repo = MagicMock()
mock_gv_repo.get.return_value = mock_gv
mock_db = MagicMock()
mock_current_user = MagicMock()
mock_current_user.user.id = "user-1"
body = GenerateCoverRequest(
cover_type="ai_frame",
generated_video_id="gv-mystery",
)
with (
patch("app.api.routes.generation_cover.SQLAlchemyGenerationTaskRepository") as mock_repo_cls,
patch(
"app.api.routes.generation_cover.get_generated_video_repository",
return_value=mock_gv_repo,
),
):
mock_repo = MagicMock()
mock_repo.get.return_value = None # 关联任务查不到
mock_repo_cls.return_value = mock_repo
from app.api.routes.generation_cover import generate_cover
with pytest.raises(HTTPException) as exc_info:
generate_cover(
body=body,
template_id="tpl-1",
plan_id="plan-orphan",
services=(mock_template_svc, mock_plan_svc),
db=mock_db,
current_user=mock_current_user,
)
assert exc_info.value.status_code == 403
def test_is_trusted_media_url_blocks_internal_and_ipv6(self):
"""白名单函数:内网 IPv4/IPv6/元数据地址一律拒绝,自家 OSS 域名放行。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import _is_trusted_media_url
mock_storage = MagicMock()
mock_storage.public_url = "https://xiaoxia-media.oss-cn-hangzhou.aliyuncs.com"
mock_storage.endpoint = "oss-cn-hangzhou.aliyuncs.com"
with patch(
"app.api.routes.generation_cover.get_shared_storage_service",
return_value=mock_storage,
):
# 内网 / 元数据 / IPv6 本地地址全部拒绝
for bad in [
"http://127.0.0.1/admin",
"http://10.0.0.5/video.mp4",
"http://192.168.1.1/video.mp4",
"http://172.16.0.1/video.mp4",
"http://169.254.169.254/latest/meta-data/",
"http://[::1]:8080/video.mp4",
"http://[fe80::1]/video.mp4",
"http://[fc00::1]/video.mp4",
"http://localhost/x",
"ftp://oss-cn-hangzhou.aliyuncs.com/a.mp4",
"",
]:
assert _is_trusted_media_url(bad) is False, f"应拒绝: {bad}"
# 自家 OSS 域名(含签名 URL 子路径、bucket 域名)放行
for good in [
"https://xiaoxia-media.oss-cn-hangzhou.aliyuncs.com/rendered/final/v.mp4",
"https://xiaoxia-media.oss-cn-hangzhou.aliyuncs.com/rendered/v.mp4?Expires=123&Signature=abc",
]:
assert _is_trusted_media_url(good) is True, f"应放行: {good}"
def test_is_trusted_media_url_endpoint_with_scheme_parsed(self):
"""endpoint 配置带 http:// 前缀时也能正确提取主机名,不出现 .http 后缀绕过。"""
from unittest.mock import MagicMock, patch
from app.api.routes.generation_cover import _is_trusted_media_url
mock_storage = MagicMock()
mock_storage.public_url = "http://oss.internal.example.com:9000"
mock_storage.endpoint = "http://oss.internal.example.com:9000"
with patch(
"app.api.routes.generation_cover.get_shared_storage_service",
return_value=mock_storage,
):
# 正确域名放行
assert _is_trusted_media_url("http://oss.internal.example.com:9000/a/b.mp4") is True
# 伪造后缀域名必须拒绝(修复前 split(':')[0] 会取到 'http' 导致绕过)
assert _is_trusted_media_url("http://evil-http.com/x.mp4") is False
assert _is_trusted_media_url("http://evil.http/x.mp4") is False