feat(#1970): 新 API 字段 + 叙事模式 PR3 - assembly_mode/script_id/tts_*/video_ratio
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 1s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 45s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 49s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m15s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 2m59s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 3m10s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 3m31s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 3m44s
AI Code Review / AI Code Review (pull_request) Successful in 6m28s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 7m4s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 9m47s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 1s
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 7m32s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 21s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 38s

- CreateGenerationTaskRequest 新增 assembly_mode(random 默认/narrative)、
  script_id、tts_voice_id、tts_voice_source(preset/clone)、video_ratio;
  模型校验:narrative 必须带 script_id+tts_voice_id,比例枚举白名单
- 叙事模式入队前同步完成 TTS:读文案(归属校验) → 解析音色(preset/clone)
  → 复用 tts_job workflow 同步合成(失败 4xx 不入队,积分按 ai_voice 口径扣退)
  → 转存配音库 audio asset,voice_library_id 指向它,下游渲染零改动
- narrative_match 纯模块:文案 tags 与素材标签归一化求交集,命中池优先
  + smart_match 评分,不足/无匹配自动降级现有随机逻辑(行为与旧版一致)
- video_ratio→输出分辨率映射(9:16/16:9/1:1/3:4/4:3),显式分辨率优先
- writeback 落 assembly_mode/script_id/video_ratio 便于追溯
- one_take/pip/voice_pip 标记 deprecated(不删代码),voice_over 保留
- 新增 53 个测试;全量 tests/unit 15736 passed / 28 skipped
This commit is contained in:
xiaoxia
2026-09-18 07:17:20 +08:00
parent a59a6a588a
commit 797a220d43
9 changed files with 1521 additions and 10 deletions
+145 -4
View File
@@ -16,10 +16,12 @@ from app.core.task_enqueue import (
from app.dependencies import (
get_asset_library_repository,
get_asset_repository,
get_cosyvoice_service,
get_db_session,
get_generated_video_repository,
get_generation_task_repository,
get_project_repository,
get_voice_clone_profile_repository,
)
from app.schemas.generated_video import (
GeneratedVideoResponse,
@@ -132,6 +134,8 @@ def _select_assets_from_library(
mode: str,
count: int,
rng=None,
script_tags: list | None = None,
tag_names_by_id: dict | None = None,
) -> list[str]:
"""根据选取模式从素材库中选取 ready 状态的视频素材 ID。
@@ -141,6 +145,8 @@ def _select_assets_from_library(
count: 选取数量,0 表示全部(仅 smart 模式有效)
rng: 可选随机源(smart 模式排序噪声用),生产环境不传则内部随机;
测试可注入固定种子或零噪声随机源获得确定性结果。
script_tags: #1970 叙事模式文案标签;非空时标签命中素材优先,不足再用其余素材兜底。
tag_names_by_id: asset_id → 素材标签名列表(素材只存 tag_ids 时由调用方查名称注入)。
Returns:
选中的素材 ID 列表
@@ -150,6 +156,20 @@ def _select_assets_from_library(
if not ready_video_assets:
return []
# 叙事模式(#1970 PR3):文案标签命中池优先;无任何命中时完全降级为现有随机逻辑。
if script_tags:
from packages.domain.narrative_match import pick_narrative_assets
limit = count if count > 0 else None
picked = pick_narrative_assets(
ready_video_assets,
script_tags=script_tags,
tag_names_by_id=tag_names_by_id,
limit=limit,
rng=rng,
)
return [a.id for a in picked]
if mode == "smart":
# 智能匹配:统一使用 packages/domain/smart_match.py 的多维评分+多样性选取
# 评分维度:质量分(40%) + 时长适配(30%) + 新鲜度(20%) + 未使用加分(10%)
@@ -162,6 +182,53 @@ def _select_assets_from_library(
return [a.id for a in ready_video_assets]
# #1970 PR3:video_ratio → 默认输出分辨率(显式 output_width/output_height 优先)
_VIDEO_RATIO_DIMENSIONS = {
"9:16": (1080, 1920),
"16:9": (1920, 1080),
"1:1": (1080, 1080),
"3:4": (1080, 1440),
"4:3": (1440, 1080),
}
def _resolve_output_dimensions(request: CreateGenerationTaskRequest) -> tuple[int, int]:
"""解析输出分辨率:显式 output_width/output_height 非旧默认值时优先,否则按 video_ratio。
前端 #1973 总是同时传 video_ratio 与具体分辨率,两者一致;此函数主要服务
只传比例的调用方,并保证旧调用(不传比例)维持 1280x720 行为。
"""
width, height = request.output_width, request.output_height
ratio = (request.video_ratio or "").strip()
if ratio in _VIDEO_RATIO_DIMENSIONS and (width, height) == (1280, 720):
return _VIDEO_RATIO_DIMENSIONS[ratio]
return width, height
def _load_asset_tag_names(db: Session, assets: list, user_id: str) -> dict[str, list[str]]:
"""叙事模式:查 TagModel 名称,构造 asset_id → 标签名列表(失败返回空 dict 降级随机)。"""
try:
from packages.adapters.sqlalchemy_impl.models import AssetTagModel, TagModel
tag_ids = {tid for a in assets for tid in (getattr(a, "tag_ids", None) or [])}
if not tag_ids:
return {}
name_rows = (
db.query(TagModel.id, TagModel.name).filter(TagModel.id.in_(tag_ids), TagModel.user_id == user_id).all()
)
name_by_id = {row.id: row.name for row in name_rows}
links = db.query(AssetTagModel.asset_id, AssetTagModel.tag_id).filter(AssetTagModel.tag_id.in_(tag_ids)).all()
index: dict[str, list[str]] = {}
for asset_id, tag_id in links:
name = name_by_id.get(tag_id)
if name:
index.setdefault(asset_id, []).append(name)
return index
except Exception: # noqa: BLE001 - 标签匹配是加分项,查询失败不阻断生成
logger.warning("[叙事模式] 素材标签查询失败,降级随机选片", exc_info=True)
return {}
def _writeback_edit_plan_config(
plan_id: str,
task_id: str,
@@ -169,12 +236,23 @@ def _writeback_edit_plan_config(
db: Session,
dedup_enabled: bool | None = None,
video_index: int | None = None,
assembly_mode: str | None = None,
script_id: str | None = None,
video_ratio: str | None = None,
) -> None:
"""[已下沉] 路由层兼容别名 → app.services.generation_common.writeback_edit_plan_config。"""
from app.services.generation_common import writeback_edit_plan_config
return writeback_edit_plan_config(
plan_id, task_id, title_config, db, dedup_enabled=dedup_enabled, video_index=video_index
plan_id,
task_id,
title_config,
db,
dedup_enabled=dedup_enabled,
video_index=video_index,
assembly_mode=assembly_mode,
script_id=script_id,
video_ratio=video_ratio,
)
@@ -225,16 +303,63 @@ def create_generation_task(
asset_library_repository: Any = Depends(get_asset_library_repository),
asset_repository: Any = Depends(get_asset_repository),
db: Session = Depends(get_db_session),
cosyvoice_service: Any = Depends(get_cosyvoice_service),
voice_clone_repository: Any = Depends(get_voice_clone_profile_repository),
) -> BatchGenerationTaskResponse:
logger.info(
"[生成任务] 接收请求: user_id=%s, template_id=%s, asset_count=%d, mode=%s, count=%d",
"[生成任务] 接收请求: user_id=%s, template_id=%s, asset_count=%d, mode=%s, assembly=%s, count=%d",
authenticated_user.user.id,
request.template_id,
len(request.asset_ids),
request.asset_select_mode,
request.assembly_mode,
request.count,
)
# video_ratio → 默认分辨率(显式分辨率优先)
request.output_width, request.output_height = _resolve_output_dimensions(request)
# ── #1970 PR3 叙事模式:入队前同步合成配音并落为 audio asset ──
# 合成结果覆盖 voice_library_id(下游按 audio asset id 消费),失败直接 4xx 不入队。
narrative_script_tags: list = []
if request.assembly_mode == "narrative":
from app.config import settings as _settings
from app.services.narrative_service import NarrativeError, prepare_narrative_voice
from packages.adapters.sqlalchemy_impl.tts_job_repository import SQLAlchemyTTSJobRepository
try:
narrative_ctx = prepare_narrative_voice(
db=db,
user_id=authenticated_user.user.id,
script_id=request.script_id,
tts_voice_id=request.tts_voice_id,
tts_voice_source=request.tts_voice_source,
tts_repository=SQLAlchemyTTSJobRepository(db),
cosyvoice_service=cosyvoice_service,
voice_clone_repository=voice_clone_repository,
asset_repository=asset_repository,
asset_library_repository=asset_library_repository,
project_repository=project_repository,
storage_service=get_storage_service(),
points_enabled=bool(getattr(_settings, "points_enabled", False)),
is_member=bool(getattr(authenticated_user.user, "is_member", False)),
member_type=getattr(authenticated_user.user, "member_type", None),
)
except NarrativeError as e:
logger.warning("[叙事模式] 配音前置处理失败: %s", e.message)
raise HTTPException(status_code=e.status_code, detail=e.message) from e
request.voice_library_id = narrative_ctx.voice_asset_id
narrative_script_tags = list(getattr(narrative_ctx.script, "tags", None) or [])
logger.info(
"[叙事模式] 配音已就绪: script_id=%s, tts_job=%s, voice_asset=%s, duration=%.2f",
request.script_id,
narrative_ctx.tts_job_id,
narrative_ctx.voice_asset_id,
narrative_ctx.audio_duration,
)
try:
project_id, asset_library_id = _resolve_project_and_library(
request, project_repository, asset_library_repository, asset_repository, authenticated_user
@@ -260,19 +385,29 @@ def create_generation_task(
# 素材库自动匹配:当未显式指定 asset_ids 时,按模式自动选取
if not resolved_asset_ids:
_tag_index = (
_load_asset_tag_names(db, assets, authenticated_user.user.id) if narrative_script_tags else None
)
resolved_asset_ids = _select_assets_from_library(
assets,
mode=request.asset_select_mode,
count=request.asset_select_count,
script_tags=narrative_script_tags or None,
tag_names_by_id=_tag_index,
)
elif project_id and not resolved_asset_ids and request.asset_select_mode in ("smart",):
# 项目级模式:未指定 asset_ids 且选择了 smart 模式时,也自动选取
elif project_id and not resolved_asset_ids and (request.asset_select_mode in ("smart",) or narrative_script_tags):
# 项目级模式:未指定 asset_ids 且选择了 smart 模式(或叙事模式按标签匹配)时自动选取
assets = asset_repository.find_by_project(project_id)
if assets:
_tag_index = (
_load_asset_tag_names(db, assets, authenticated_user.user.id) if narrative_script_tags else None
)
resolved_asset_ids = _select_assets_from_library(
assets,
mode=request.asset_select_mode,
count=request.asset_select_count,
script_tags=narrative_script_tags or None,
tag_names_by_id=_tag_index,
)
if not resolved_asset_ids:
raise HTTPException(
@@ -337,6 +472,9 @@ def create_generation_task(
title_config=fallback_title_config,
db=db,
dedup_enabled=request.dedup_enabled,
assembly_mode=request.assembly_mode,
script_id=request.script_id or None,
video_ratio=request.video_ratio or None,
)
logger.info(
@@ -685,6 +823,9 @@ def create_generation_task(
db=db,
dedup_enabled=request.dedup_enabled,
video_index=task_index,
assembly_mode=request.assembly_mode,
script_id=request.script_id or None,
video_ratio=request.video_ratio or None,
)
if safe_enqueue_generation_task(
+33
View File
@@ -103,6 +103,19 @@ class CreateGenerationTaskRequest(BaseModel):
# False:跳过 edge_crop、不注入微变换,渲染确定性(固定种子)。
dedup_enabled: bool = Field(default=True, description="智能降重开关,默认开启;关闭后跳过边缘裁切与微变换")
# ── 剪辑组装模式(#1970 PR3)──
# random(默认,完全兼容现有随机混剪)/ narrative(叙事剪辑:文案→TTS 配音→标签匹配画面)
assembly_mode: str = Field(default="random", description="组装模式:random=随机混剪(默认),narrative=叙事剪辑")
# 叙事模式必填:文案库 scripts.id(后端据此读取 content 合成 TTS)
script_id: str = Field(default="", description="叙事模式必填:文案库 ID")
# 叙事模式必填:TTS 音色 ID(preset 为 CosyVoice 音色 id;clone 为克隆档案 id)
tts_voice_id: str = Field(default="", description="叙事模式必填:TTS 音色 ID(系统音色或克隆档案 ID)")
tts_voice_source: str = Field(default="preset", description="TTS 音色来源:preset=系统预设(默认),clone=克隆音色")
# 视频比例:当前前端 9:16/16:9;与 output_width/output_height 并存,传了具体分辨率时以分辨率为准
video_ratio: str = Field(
default="", description="视频比例,如 9:16(默认竖屏)/16:9;与显式分辨率冲突时以分辨率为准"
)
@model_validator(mode="after")
def _check_variant_arrays(self) -> "CreateGenerationTaskRequest":
"""变体数组字段长度校验 + #1749 配音严格守卫。
@@ -132,6 +145,26 @@ class CreateGenerationTaskRequest(BaseModel):
raise ValueError(f"variant_plan_ids 长度({len(self.variant_plan_ids)})必须与 count({self.count})一致")
return self
@model_validator(mode="after")
def _check_assembly_mode(self) -> "CreateGenerationTaskRequest":
"""#1970 组装模式与叙事模式入参校验。"""
if self.assembly_mode not in ("random", "narrative"):
raise ValueError("assembly_mode 仅支持 'random'(默认)或 'narrative'")
if self.tts_voice_source not in ("preset", "clone"):
raise ValueError("tts_voice_source 仅支持 'preset' 或 'clone'")
if self.video_ratio:
parts = self.video_ratio.split(":")
if len(parts) != 2 or not all(p.isdigit() and int(p) > 0 for p in parts):
raise ValueError("video_ratio 格式必须为 '宽:高',如 9:16 或 16:9")
if self.video_ratio not in ("9:16", "16:9", "1:1", "3:4", "4:3"):
raise ValueError("video_ratio 仅支持 9:16 / 16:9 / 1:1 / 3:4 / 4:3")
if self.assembly_mode == "narrative":
if not self.script_id.strip():
raise ValueError("叙事模式(narrative)必须提供 script_id(文案库 ID)")
if not self.tts_voice_id.strip():
raise ValueError("叙事模式(narrative)必须提供 tts_voice_id(TTS 音色 ID)")
return self
@model_validator(mode="after")
def _check_at_least_one_mode(self) -> "CreateGenerationTaskRequest":
has_project = bool(self.project_id.strip())
+11 -1
View File
@@ -63,11 +63,15 @@ def writeback_edit_plan_config(
db: Session,
dedup_enabled: bool | None = None,
video_index: int | None = None,
assembly_mode: str | None = None,
script_id: str | None = None,
video_ratio: str | None = None,
) -> None:
"""任务入队成功后,回写 EditPlan.config:generation_task_id + title_config。
用 merge 方式更新,不整体覆盖 config,避免丢失其他字段。
#1970:dedup_enabled 非 None 时一并写入,worker 据此决定 edge_crop/微变换。
#1970:dedup_enabled 非 None 时一并写入,worker 据此决定 edge_crop/微变换;
PR3 叙事模式再写 assembly_mode/script_id/video_ratio(可追溯,不影响渲染)。
失败只记日志,不影响任务创建。
"""
if not plan_id:
@@ -87,6 +91,12 @@ def writeback_edit_plan_config(
merged["dedup_enabled"] = bool(dedup_enabled)
if video_index is not None:
merged["video_index"] = int(video_index)
if assembly_mode:
merged["assembly_mode"] = assembly_mode
if script_id:
merged["script_id"] = script_id
if video_ratio:
merged["video_ratio"] = video_ratio
if title_config:
# #1901 统一字段名为 "title"(worker sync_configs_to_plan 写的是 "title")
+344
View File
@@ -0,0 +1,344 @@
"""叙事剪辑前置服务 — #1970 PR3.
叙事模式(assembly_mode='narrative')在生成任务入队前同步完成:
1. 按 script_id 读取文案(归属校验);
2. 按 tts_voice_source 解析音色(preset=CosyVoice 音色 id;clone=克隆档案 id,
解析档案归属并取其 CosyVoice voice_id);
3. 同步 TTS 合成(复用 tts_job 现有 workflow:提交即同步返回,未完成则轮询兜底),
失败直接抛 NarrativeError(HTTP 层转 4xx,任务不入队);
4. 把合成音频转存为配音库 audio asset(与 /tts/jobs/{id}/save-to-library 同一套
存储路径与元信息约定),返回 asset_id —— 下游仍以 voice_library_id(实为
audio asset id)消费,渲染链路零改动。
积分扣点与 /tts 合成端点保持一致(ai_voice 场景),失败退费。
"""
from __future__ import annotations
import json
import logging
import math
import subprocess
import tempfile
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import ScriptModel
from packages.application.cosyvoice_service import CosyVoiceService
from packages.application.tts_job.use_cases import CreateTTSJobUseCase
from packages.application.tts_job.workflow import TTSWorkflowService
from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, ClassificationStatus
from packages.domain.points_rules import calculate_points_cost
from packages.domain.points_service import PointsService
from packages.shared.storage import SharedStorageService
logger = logging.getLogger(__name__)
_POINTS_SCENE = "ai_voice"
_SYNTH_TIMEOUT = 180.0 # 叙事配音在 HTTP 请求内同步等待,长文案分段合成时留出余量
_CONTENT_TYPE_MAP = {"mp3": "audio/mpeg", "wav": "audio/wav", "pcm": "audio/pcm", "opus": "audio/opus"}
class NarrativeError(Exception):
"""叙事模式前置处理失败(文案/音色/TTS/落库)。"""
def __init__(self, message: str, *, status_code: int = 400) -> None:
super().__init__(message)
self.message = message
self.status_code = status_code
@dataclass(slots=True)
class NarrativeContext:
"""叙事模式前置处理结果。"""
script: ScriptModel
voice_asset_id: str
tts_job_id: str
audio_duration: float
def _find_or_create_voice_library(
*,
user_id: str,
project_repository: Any,
asset_library_repository: Any,
) -> AssetLibrary:
"""找到(或自动创建)用户 voice 素材库;与 tts.py 保存配音库逻辑一致。"""
projects = project_repository.find_accessible_projects(user_id)
if not projects:
raise NarrativeError("没有可用的项目,无法保存叙事配音", status_code=400)
for project in projects:
for lib in asset_library_repository.find_by_project(project.id):
kind = lib.kind.value if hasattr(lib.kind, "value") else lib.kind
if kind == AssetLibraryKind.VOICE.value:
return lib
project = projects[0]
library = AssetLibrary.create(project_id=project.id, name="配音素材库", kind=AssetLibraryKind.VOICE)
from sqlalchemy.exc import IntegrityError
try:
return asset_library_repository.create(library)
except IntegrityError:
session = getattr(asset_library_repository, "session", None)
if session is not None:
try:
session.rollback()
except Exception: # noqa: BLE001 - 回滚失败不影响重查
logger.warning("IntegrityError 后回滚 session 失败", exc_info=True)
for lib in asset_library_repository.find_by_project(project.id):
kind = lib.kind.value if hasattr(lib.kind, "value") else lib.kind
if kind == AssetLibraryKind.VOICE.value:
return lib
raise NarrativeError("配音素材库创建失败,请重试", status_code=500) from None
def _resolve_voice(
*,
user_id: str,
tts_voice_id: str,
tts_voice_source: str,
voice_clone_repository: Any,
) -> tuple[str, str]:
"""解析音色 → (CosyVoice voice_id, voice_clone_profile_id)。"""
if tts_voice_source == "clone":
profile = voice_clone_repository.get(tts_voice_id)
if profile is None:
raise NarrativeError("克隆音色不存在", status_code=404)
if profile.user_id != user_id:
raise NarrativeError("无权使用该克隆音色", status_code=403)
if not profile.voice_id:
raise NarrativeError("音色克隆尚未完成,请稍后再试", status_code=400)
return profile.voice_id, profile.id
# preset:tts_voice_id 即 CosyVoice 音色 id;与 /tts 端点一致,
# 若前端误传克隆档案 UUID,同样兼容解析。
profile = voice_clone_repository.get(tts_voice_id)
if profile is not None:
if profile.user_id != user_id:
raise NarrativeError("无权使用该音色", status_code=403)
if not profile.voice_id:
raise NarrativeError("音色克隆尚未完成,请稍后再试", status_code=400)
return profile.voice_id, profile.id
return tts_voice_id, ""
def _save_tts_job_as_voice_asset(
*,
job: Any,
user_id: str,
name: str,
project_repository: Any,
asset_library_repository: Any,
asset_repository: Any,
storage_service: SharedStorageService,
) -> Asset:
"""把已完成 TTS job 的音频转存为配音库 audio asset(同 save-to-library 约定)。"""
if not job.output_audio_url and not job.output_audio_key:
raise NarrativeError("TTS 合成缺少输出音频", status_code=502)
library = _find_or_create_voice_library(
user_id=user_id,
project_repository=project_repository,
asset_library_repository=asset_library_repository,
)
audio_format = (job.format or "mp3").strip() or "mp3"
content_type = _CONTENT_TYPE_MAP.get(audio_format, "audio/mpeg")
storage_key = f"uploads/voice/tts/{job.id}.{audio_format}"
tmp_path: Path | None = None
audio_duration: float | None = None
file_size = 0
try:
with tempfile.NamedTemporaryFile(suffix=f".{audio_format}", delete=False) as tmp:
tmp_path = Path(tmp.name)
download_source = job.output_audio_key or job.output_audio_url
downloaded = storage_service.download_asset(download_source, tmp_path)
if not downloaded or not tmp_path.exists() or tmp_path.stat().st_size == 0:
raise NarrativeError("叙事配音音频转存失败", status_code=502)
file_size = tmp_path.stat().st_size
storage_service.upload_file(tmp_path, storage_key, content_type=content_type)
try:
proc = subprocess.run(
[
"ffprobe",
"-v",
"quiet",
"-print_format",
"json",
"-show_format",
str(tmp_path),
],
capture_output=True,
text=True,
timeout=10,
)
if proc.returncode == 0:
dur = float(json.loads(proc.stdout).get("format", {}).get("duration", 0))
if dur > 0:
audio_duration = dur
except Exception: # noqa: BLE001 - ffprobe 仅用于时长兜底
logger.warning("叙事配音 ffprobe 时长提取失败: job_id=%s", job.id, exc_info=True)
except NarrativeError:
raise
except Exception as e: # noqa: BLE001
logger.error("叙事配音转存失败: job_id=%s, error=%s", job.id, e, exc_info=True)
raise NarrativeError("叙事配音音频转存失败", status_code=502) from e
finally:
if tmp_path and tmp_path.exists():
try:
tmp_path.unlink()
except OSError:
pass
metadata_: dict[str, object] = {
"source": "tts_job",
"tts_job_id": job.id,
"narrative": True,
"format": job.format,
"sample_rate": job.sample_rate,
"voice_id": job.voice_id,
"voice_name": job.voice_model or "",
}
if job.metadata:
for key in ("speed", "language"):
if key in job.metadata:
metadata_[key] = job.metadata[key]
asset = Asset.create(
project_id=library.project_id,
library_id=library.id,
name=name or f"叙事配音-{job.id[:8]}",
storage_key=storage_key,
mime_type=content_type,
metadata=metadata_,
file_size=file_size,
duration=job.duration or audio_duration or None,
status=AssetStatus.READY,
classification_status=ClassificationStatus.PENDING,
uploaded_by_user_id=user_id,
)
try:
return asset_repository.create(asset)
except Exception as e: # noqa: BLE001
logger.error("叙事配音 asset 落库失败,清理 OSS: %s, error=%s", storage_key, e, exc_info=True)
try:
storage_service.delete_file(storage_key)
except Exception: # noqa: BLE001
logger.warning("清理孤儿 OSS 文件失败: %s", storage_key, exc_info=True)
raise NarrativeError("叙事配音保存失败,请重试", status_code=502) from e
def prepare_narrative_voice(
*,
db: Session,
user_id: str,
script_id: str,
tts_voice_id: str,
tts_voice_source: str,
tts_repository: Any,
cosyvoice_service: CosyVoiceService,
voice_clone_repository: Any,
asset_repository: Any,
asset_library_repository: Any,
project_repository: Any,
storage_service: SharedStorageService,
points_enabled: bool = False,
is_member: bool = False,
member_type: str | None = None,
) -> NarrativeContext:
"""叙事模式入队前同步合成配音并落为 audio asset。
Raises:
NarrativeError: 文案缺失/归属不符、音色不可用、TTS 失败、转存失败。
"""
script = db.query(ScriptModel).filter(ScriptModel.id == script_id, ScriptModel.user_id == user_id).first()
if script is None:
raise NarrativeError("文案不存在或无权使用", status_code=404)
content = (script.content or "").strip()
if not content:
raise NarrativeError("文案内容为空,无法合成配音", status_code=400)
actual_voice_id, clone_profile_id = _resolve_voice(
user_id=user_id,
tts_voice_id=tts_voice_id,
tts_voice_source=tts_voice_source,
voice_clone_repository=voice_clone_repository,
)
# 积分扣点(与 /tts 合成端点同口径),失败时在合成失败分支退费
points_svc = PointsService() if points_enabled else None
points_deducted = 0
if points_svc is not None:
est_minutes = max(1.0, math.ceil(len(content) / 240))
points_deducted = calculate_points_cost(
_POINTS_SCENE,
is_member=is_member,
duration_minutes=est_minutes,
member_type=member_type,
)
deduct_res = points_svc.deduct_points(user_id, points_deducted, _POINTS_SCENE, db)
if not deduct_res["success"]:
raise NarrativeError(
f"积分不足,需要 {points_deducted} 积分,当前余额 {deduct_res['balance']}",
status_code=402,
)
use_case = CreateTTSJobUseCase(tts_repository)
job = use_case.execute(
user_id=user_id,
input_text=content,
voice_id=actual_voice_id,
voice_clone_profile_id=clone_profile_id,
metadata={"speed": 1.0, "emotion": "", "language": "zh-CN", "narrative": True, "script_id": script_id},
)
workflow = TTSWorkflowService(repository=tts_repository, cosyvoice_service=cosyvoice_service)
try:
job = workflow.start_synthesis(job.id)
if not job.is_completed:
job = workflow.poll_and_process_synthesis(job.id, timeout=_SYNTH_TIMEOUT)
except Exception as e: # noqa: BLE001 - 同步合成异常统一转 NarrativeError
logger.error("叙事配音 TTS 合成失败: job_id=%s, error=%s", job.id, e, exc_info=True)
try:
workflow.process_synthesis_failure(job.id, str(e))
except Exception: # noqa: BLE001
logger.warning("标记叙事 TTS job 失败出错: job_id=%s", job.id, exc_info=True)
if points_deducted and points_svc is not None:
try:
points_svc.refund_points(user_id, points_deducted, _POINTS_SCENE, db, ref_id=job.id)
except Exception: # noqa: BLE001
logger.warning("叙事 TTS 失败退积分异常: job_id=%s", job.id, exc_info=True)
raise NarrativeError(f"配音合成失败:{e}", status_code=502) from e
if not job.is_completed:
if points_deducted and points_svc is not None:
try:
points_svc.refund_points(user_id, points_deducted, _POINTS_SCENE, db, ref_id=job.id)
except Exception: # noqa: BLE001
logger.warning("叙事 TTS 未完成退积分异常: job_id=%s", job.id, exc_info=True)
raise NarrativeError("配音合成未完成,请稍后重试", status_code=504)
asset = _save_tts_job_as_voice_asset(
job=job,
user_id=user_id,
name=(script.title or "叙事配音")[:60],
project_repository=project_repository,
asset_library_repository=asset_library_repository,
asset_repository=asset_repository,
storage_service=storage_service,
)
return NarrativeContext(
script=script,
voice_asset_id=asset.id,
tts_job_id=job.id,
audio_duration=float(job.duration or asset.duration or 0.0),
)
+17 -5
View File
@@ -12,9 +12,21 @@ else:
class EditingMode(StrEnum):
"""剪辑模式枚举"""
"""剪辑模式枚举。
ONE_TAKE = "one_take" # 顺序拼接模式
PIP = "pip" # 画中画模式
VOICE_OVER = "voice_over" # 口播+B-roll模式
VOICE_PIP = "voice_pip" # 口播+画中画组合模式
#1970 智能剪辑流程重构(2026-09)后,剪辑组装模式改由
``CreateGenerationTaskRequest.assembly_mode``('random'/'narrative')表达。
本枚举仅保留模板体系仍在使用的模式;以下三个模式标记 deprecated,
不主动删除代码(pip/voice_pip 在路由入口已统一映射为 one_take),
待确认无存量引用后在技术债清理中移除:
- ONE_TAKE(deprecated):顺序拼接,等同 assembly_mode='random'
- PIP(deprecated):画中画已下线,入口映射 one_take
- VOICE_PIP(deprecated):口播+画中画已下线,入口映射 one_take
- VOICE_OVER:保留,口播+B-roll 模板仍在使用
"""
ONE_TAKE = "one_take" # deprecated(#1970):顺序拼接,等同 assembly_mode='random'
PIP = "pip" # deprecated(#1970):画中画已下线,入口映射 one_take
VOICE_OVER = "voice_over" # 口播+B-roll模式(保留)
VOICE_PIP = "voice_pip" # deprecated(#1970):口播+画中画已下线,入口映射 one_take
+132
View File
@@ -0,0 +1,132 @@
"""叙事剪辑素材标签匹配 — #1970 PR3.
叙事模式下,选片在现有评分(smart_match / atom_clip_selector)之前先做一层
文案标签匹配:
- 文案 tags 与素材 tag 名归一化后求交集;
- 命中任一标签的素材作为「优先候选池」,未命中的作为普通池;
- 调用方对优先池跑现有 smart_select_assets,数量不足时用普通池补足
(无任何匹配 → 完全降级为现有随机逻辑,行为与改造前一致)。
纯函数模块:标签 id→名称映射由调用方查 TagModel 后注入,不直接碰 DB。
"""
from __future__ import annotations
from typing import Any, Iterable
# 标签归一化后仍短于此长度的标签不参与匹配(避免「的」「是」这类噪声短词)
MIN_TAG_LEN = 2
def normalize_tag(tag: Any) -> str:
"""标签归一化:去空白、小写。数字/英文统一小写,中文不受影响。"""
if tag is None:
return ""
return str(tag).strip().lower()
def _normalize_tags(tags: Iterable[Any]) -> set[str]:
out: set[str] = set()
for t in tags or []:
norm = normalize_tag(t)
if len(norm) >= MIN_TAG_LEN:
out.add(norm)
return out
def build_asset_tag_name_index(tag_names_by_id: dict[str, Any]) -> dict[str, set[str]]:
"""构造 asset_id → 归一化标签名集合 的索引。
Args:
tag_names_by_id: {asset_id: [标签名或标签id, ...]},允许混入 None/空值
"""
index: dict[str, set[str]] = {}
for asset_id, names in (tag_names_by_id or {}).items():
index[asset_id] = _normalize_tags(names)
return index
def match_assets_by_script_tags(
assets: list[Any],
*,
script_tags: Iterable[Any],
tag_names_by_id: dict[str, Any] | None = None,
) -> tuple[list[Any], list[Any]]:
"""按文案标签把素材拆成「命中池 / 未命中池」,保持输入相对顺序。
Args:
assets: 候选素材(domain Asset,需有 id 与 tag_ids)。
script_tags: 文案 tags(字符串数组,名称语义)。
tag_names_by_id: asset_id → 素材标签名列表;素材只有 tag_ids 时由调用方
查 TagModel 名称后传入。为空则视为无素材命中。
Returns:
(matched, unmatched):命中任一文案标签的素材 / 其余素材。
文案无有效标签时 matched 为空(调用方直接走随机逻辑)。
"""
wanted = _normalize_tags(script_tags)
if not wanted:
return [], list(assets)
name_index = build_asset_tag_name_index(tag_names_by_id or {})
matched: list[Any] = []
unmatched: list[Any] = []
for asset in assets:
asset_id = str(getattr(asset, "id", "") or "")
names = set(name_index.get(asset_id, set()))
# 兼容素材自身带字符串 tags(旧链路/测试替身)
raw_tags = getattr(asset, "tags", None)
if raw_tags:
names |= _normalize_tags(raw_tags)
if names & wanted:
matched.append(asset)
else:
unmatched.append(asset)
return matched, unmatched
def pick_narrative_assets(
assets: list[Any],
*,
script_tags: Iterable[Any],
tag_names_by_id: dict[str, Any] | None = None,
limit: int | None = None,
rng: Any = None,
) -> list[Any]:
"""叙事模式选片:标签命中池优先,不足部分从未命中池按现有评分补齐。
本函数只负责「标签优先 + 兜底降级」的顺序编排;评分仍复用
smart_match.smart_select_assets(质量/时长/新鲜度/未使用 + 随机噪声),
不重写评分维度。
Args:
assets: ready 视频素材候选(调用方负责状态/类型过滤)。
script_tags / tag_names_by_id: 见 match_assets_by_script_tags。
limit: 需要的素材数量;None 表示全部(命中池 + 全部未命中池)。
rng: 注入 smart_select_assets 的随机源(可复现)。
Returns:
选中的素材列表。无任何标签命中时等价于对全量跑 smart_select_assets。
"""
from packages.domain.smart_match import smart_select_assets
matched, unmatched = match_assets_by_script_tags(
assets,
script_tags=script_tags,
tag_names_by_id=tag_names_by_id,
)
need = limit if (limit is not None and limit > 0) else None
if not matched:
# 完全降级:与改造前随机混剪同一逻辑
return [r.asset for r in smart_select_assets(assets, kind="video", limit=need, rng=rng)]
picked = [r.asset for r in smart_select_assets(matched, kind="video", limit=need, rng=rng)]
if need is not None and len(picked) < need and unmatched:
rest_need = need - len(picked)
picked.extend(r.asset for r in smart_select_assets(unmatched, kind="video", limit=rest_need, rng=rng))
elif need is None:
picked.extend(r.asset for r in smart_select_assets(unmatched, kind="video", rng=rng))
return picked
+218
View File
@@ -0,0 +1,218 @@
"""#1970 PR3 schema 校验 + 路由辅助函数测试。"""
from __future__ import annotations
from dataclasses import dataclass, field
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from app.api.routes import generation_tasks as gt
from app.schemas.generation_task import CreateGenerationTaskRequest
from pydantic import ValidationError
# ── schema ─────────────────────────────────────────────────────────────────
def _base_payload(**overrides):
payload = dict(
template_id="tpl1",
asset_ids=["a1", "a2"],
duration=30,
title_text="t",
editing_mode="voice_over",
)
payload.update(overrides)
return payload
class TestAssemblySchema:
def test_defaults(self):
req = CreateGenerationTaskRequest(**_base_payload())
assert req.assembly_mode == "random"
assert req.script_id == ""
assert req.tts_voice_id == ""
assert req.tts_voice_source == "preset"
assert req.video_ratio == "" # 空串=沿用模板默认(前端新流程显式传 9:16)
assert req.dedup_enabled is True
def test_narrative_accepts_fields(self):
req = CreateGenerationTaskRequest(
**_base_payload(
assembly_mode="narrative",
script_id="s1",
tts_voice_id="longxiaochun",
tts_voice_source="clone",
video_ratio="16:9",
)
)
assert req.assembly_mode == "narrative"
assert req.script_id == "s1"
def test_bad_assembly_mode_rejected(self):
with pytest.raises(ValidationError):
CreateGenerationTaskRequest(**_base_payload(assembly_mode="movie"))
def test_bad_voice_source_rejected(self):
with pytest.raises(ValidationError):
CreateGenerationTaskRequest(**_base_payload(tts_voice_source="elevenlabs"))
def test_bad_video_ratio_rejected(self):
with pytest.raises(ValidationError):
CreateGenerationTaskRequest(**_base_payload(video_ratio="4:5"))
def test_narrative_without_script_rejected(self):
with pytest.raises(ValidationError) as ei:
CreateGenerationTaskRequest(**_base_payload(assembly_mode="narrative"))
assert "script_id" in str(ei.value)
def test_narrative_without_voice_rejected(self):
with pytest.raises(ValidationError) as ei:
CreateGenerationTaskRequest(**_base_payload(assembly_mode="narrative", script_id="s1"))
assert "tts_voice_id" in str(ei.value)
def test_random_mode_ignores_script_absence(self):
req = CreateGenerationTaskRequest(**_base_payload())
assert req.assembly_mode == "random"
# ── _select_assets_from_library 的叙事分支 ─────────────────────────────────
@dataclass
class _Asset:
id: str
status: object = field(default_factory=lambda: SimpleNamespace(value="ready"))
mime_type: str = "video/mp4"
tags: list[str] = field(default_factory=list)
tag_ids: list[str] = field(default_factory=list)
file_type: str = "video"
quality_score: float | None = None
duration: float = 8.0
created_at: object = None
metadata: dict = field(default_factory=dict)
class TestNarrativeSelectInRoute:
def test_narrative_tags_prioritize_matched(self):
assets = [
_Asset("a1", tags=["工厂"]),
_Asset("a2", tags=["旅游"]),
_Asset("a3", tags=["工厂"]),
]
picked = gt._select_assets_from_library(assets, mode="all", count=2, script_tags=["工厂"])
assert set(picked) == {"a1", "a3"}
def test_narrative_no_match_falls_back_to_full_pool(self):
assets = [_Asset("a1", tags=["工厂"]), _Asset("a2", tags=["旅游"])]
picked = gt._select_assets_from_library(assets, mode="all", count=2, script_tags=["美食"])
assert set(picked) == {"a1", "a2"}
def test_tag_ids_via_index(self):
assets = [_Asset("a1", tag_ids=["t1"]), _Asset("a2", tag_ids=["t2"])]
picked = gt._select_assets_from_library(
assets,
mode="all",
count=1,
script_tags=["教程"],
tag_names_by_id={"a1": ["教程"], "a2": ["旅游"]},
)
assert picked == ["a1"]
def test_no_script_tags_smart_path_unchanged(self):
assets = [_Asset("a1"), _Asset("a2")]
picked = gt._select_assets_from_library(assets, mode="smart", count=1)
assert picked # 非空即可,评分逻辑由 smart_match 自己的测试覆盖
# ── _load_asset_tag_names(DB 替身) ────────────────────────────────────────
class _FakeRow:
def __init__(self, **kw):
self.__dict__.update(kw)
class _FakeQuery:
def __init__(self, rows):
self._rows = rows
def filter(self, *a, **k):
return self
def all(self):
return self._rows
class _FakeDb:
def __init__(self, name_rows, link_rows):
self._maps = {
"names": name_rows,
"links": link_rows,
}
def query(self, *cols):
# _load_asset_tag_names 两次查询:第一次取 (id, name),第二次取 (asset_id, tag_id)
keys = tuple(getattr(c, "key", None) for c in cols)
if keys and keys[0] == "id":
return _FakeQuery(self._maps["names"])
return _FakeQuery(self._maps["links"])
@dataclass
class _TagIdAsset:
id: str
tag_ids: list[str]
class TestLoadAssetTagNames:
def test_builds_index(self):
assets = [_TagIdAsset("a1", ["t1", "t2"]), _TagIdAsset("a2", ["t2"])]
db = _FakeDb(
name_rows=[_FakeRow(id="t1", name="工厂"), _FakeRow(id="t2", name="带货")],
link_rows=[
("a1", "t1"),
("a1", "t2"),
("a2", "t2"),
],
)
idx = gt._load_asset_tag_names(db, assets, "u1")
assert idx == {"a1": ["工厂", "带货"], "a2": ["带货"]}
def test_no_tag_ids_returns_empty(self):
assert gt._load_asset_tag_names(_FakeDb([], []), [_TagIdAsset("a1", [])], "u1") == {}
def test_query_failure_degrades_empty(self):
class BoomQuery:
def filter(self, *a, **k):
raise RuntimeError("db down")
class BoomDb:
def query(self, *a):
return BoomQuery()
idx = gt._load_asset_tag_names(BoomDb(), [_TagIdAsset("a1", ["t1"])], "u1")
assert idx == {}
# ── _resolve_output_dimensions ─────────────────────────────────────────────
class TestResolveOutputDimensions:
def _req(self, ratio="", width=1280, height=720):
return CreateGenerationTaskRequest(**_base_payload(video_ratio=ratio, output_width=width, output_height=height))
def test_known_ratios(self):
assert gt._resolve_output_dimensions(self._req("9:16")) == (1080, 1920)
assert gt._resolve_output_dimensions(self._req("16:9")) == (1920, 1080)
assert gt._resolve_output_dimensions(self._req("1:1")) == (1080, 1080)
assert gt._resolve_output_dimensions(self._req("4:3")) == (1440, 1080)
assert gt._resolve_output_dimensions(self._req("3:4")) == (1080, 1440)
def test_old_call_default_kept_when_no_ratio(self):
assert gt._resolve_output_dimensions(self._req("")) == (1280, 720)
def test_explicit_dimensions_take_precedence(self):
# 非旧默认值(720p)的显式分辨率优先于 ratio 映射
req = self._req("9:16", width=1440, height=2560)
assert gt._resolve_output_dimensions(req) == (1440, 2560)
+167
View File
@@ -0,0 +1,167 @@
"""#1970 PR3 叙事模式文案标签匹配纯函数测试。"""
from __future__ import annotations
import random
from dataclasses import dataclass, field
import pytest
from packages.domain.narrative_match import (
build_asset_tag_name_index,
match_assets_by_script_tags,
normalize_tag,
pick_narrative_assets,
)
@dataclass
class FakeAsset:
id: str
tag_ids: list[str] = field(default_factory=list)
tags: list[str] = field(default_factory=list)
status: str = "ready"
file_type: str = "video"
duration: float = 10.0
quality_score: float | None = None
created_at: object = None
metadata: dict = field(default_factory=dict)
# ── normalize_tag ──────────────────────────────────────────────────────────
class TestNormalizeTag:
def test_strip_and_lower(self):
assert normalize_tag(" 带货 ") == "带货"
assert normalize_tag("Factory") == "factory"
def test_none_and_non_string(self):
assert normalize_tag(None) == ""
assert normalize_tag(123) == "123"
def test_short_tag_filtered_by_normalize_set(self):
# 单字噪声标签不参与匹配(_normalize_tags 层过滤)
from packages.domain.narrative_match import _normalize_tags
assert _normalize_tags(["的", " a ", "工厂"]) == {"工厂"}
# ── match_assets_by_script_tags ────────────────────────────────────────────
class TestMatchSplit:
def test_split_by_tag_names(self):
assets = [
FakeAsset("a1", tags=["工厂"]),
FakeAsset("a2", tags=["旅游"]),
FakeAsset("a3", tags=["工厂", "车间"]),
]
matched, unmatched = match_assets_by_script_tags(assets, script_tags=["工厂"])
assert [a.id for a in matched] == ["a1", "a3"]
assert [a.id for a in unmatched] == ["a2"]
def test_case_insensitive(self):
assets = [FakeAsset("a1", tags=["Factory"])]
matched, unmatched = match_assets_by_script_tags(assets, script_tags=["FACTORY"])
assert [a.id for a in matched] == ["a1"]
assert unmatched == []
def test_tag_ids_via_name_index(self):
assets = [FakeAsset("a1", tag_ids=["t1"]), FakeAsset("a2", tag_ids=["t2"])]
index = {"a1": ["测评"], "a2": ["vlog"]}
matched, unmatched = match_assets_by_script_tags(assets, script_tags=["测评"], tag_names_by_id=index)
assert [a.id for a in matched] == ["a1"]
assert [a.id for a in unmatched] == ["a2"]
def test_empty_script_tags_degrades_all_unmatched(self):
assets = [FakeAsset("a1", tags=["工厂"])]
matched, unmatched = match_assets_by_script_tags(assets, script_tags=[])
assert matched == []
assert [a.id for a in unmatched] == ["a1"]
def test_no_match_degrades(self):
assets = [FakeAsset("a1", tags=["工厂"]), FakeAsset("a2", tags=["车间"])]
matched, unmatched = match_assets_by_script_tags(assets, script_tags=["美食"])
assert matched == []
assert {a.id for a in unmatched} == {"a1", "a2"}
def test_order_preserved(self):
assets = [FakeAsset(f"a{i}", tags=["x" if i % 2 else "工厂"]) for i in range(6)]
matched, _ = match_assets_by_script_tags(assets, script_tags=["工厂"])
assert [a.id for a in matched] == ["a0", "a2", "a4"]
def test_build_index_ignores_blank(self):
# 空白/None/单字符噪声标签均不参与匹配
idx = build_asset_tag_name_index({"a1": [" 工厂 ", "", None, "A"]})
assert idx == {"a1": {"工厂"}}
# ── pick_narrative_assets ──────────────────────────────────────────────────
class TestPickNarrativeAssets:
def _assets(self):
# smart_match 需要 created_at(None 走 recency 兜底)
import datetime as dt
old = dt.datetime(2020, 1, 1, tzinfo=dt.UTC)
return [
FakeAsset("match1", tags=["工厂"], created_at=old),
FakeAsset("nomatch1", tags=["旅游"], created_at=old),
FakeAsset("match2", tags=["工厂"], created_at=old),
FakeAsset("nomatch2", tags=["美食"], created_at=old),
]
def test_matched_pool_prioritized(self):
picked = pick_narrative_assets(self._assets(), script_tags=["工厂"], limit=2, rng=random.Random(0))
assert {a.id for a in picked} <= {"match1", "match2"}
assert all(a.id.startswith("match") for a in picked)
def test_fallback_fills_from_unmatched(self):
picked = pick_narrative_assets(self._assets(), script_tags=["工厂"], limit=4, rng=random.Random(0))
ids = {a.id for a in picked}
assert ids == {"match1", "match2", "nomatch1", "nomatch2"}
# 命中池排在前面
assert picked[0].id.startswith("match")
assert picked[1].id.startswith("match")
def test_no_tag_match_equals_random_selection(self):
assets = self._assets()
picked = pick_narrative_assets(assets, script_tags=["不存在"], limit=3, rng=random.Random(42))
assert len(picked) == 3
def test_empty_tags_selects_all_pool(self):
assets = self._assets()
picked = pick_narrative_assets(assets, script_tags=[], limit=None, rng=random.Random(1))
assert len(picked) == 4
def test_limit_none_returns_all_with_matched_first(self):
picked = pick_narrative_assets(self._assets(), script_tags=["工厂"], limit=None, rng=random.Random(1))
assert len(picked) == 4
assert {a.id for a in picked[:2]} == {"match1", "match2"}
def test_tag_ids_index_path(self):
assets = [FakeAsset("a1", tag_ids=["t1"]), FakeAsset("a2", tag_ids=["t2"])]
# 补 created_at
import datetime as dt
for a in assets:
a.created_at = dt.datetime(2020, 1, 1, tzinfo=dt.UTC)
picked = pick_narrative_assets(
assets,
script_tags=["教程"],
tag_names_by_id={"a1": ["教程"], "a2": ["旅游"]},
limit=1,
rng=random.Random(0),
)
assert [a.id for a in picked] == ["a1"]
def test_deterministic_with_seed(self):
r1 = pick_narrative_assets(self._assets(), script_tags=["工厂"], limit=4, rng=random.Random(7))
r2 = pick_narrative_assets(self._assets(), script_tags=["工厂"], limit=4, rng=random.Random(7))
assert [a.id for a in r1] == [a.id for a in r2]
if __name__ == "__main__":
pytest.main([__file__, "-q"])
+454
View File
@@ -0,0 +1,454 @@
"""#1970 PR3 叙事前置服务 narrative_service 单元测试(不依赖真实 PG/OSS/CosyVoice)。"""
from __future__ import annotations
from dataclasses import dataclass, field
from types import SimpleNamespace
from typing import Any
import pytest
from apps.api.app.services import narrative_service as ns
from apps.api.app.services.narrative_service import (
NarrativeError,
_resolve_voice,
_save_tts_job_as_voice_asset,
prepare_narrative_voice,
)
from packages.adapters.sqlalchemy_impl.models import ScriptModel
from packages.domain.tts_job import TTSJob, TTSJobStatus
# ── fakes ──────────────────────────────────────────────────────────────────
@dataclass
class FakeProfile:
id: str = "prof-1"
user_id: str = "u1"
voice_id: str = "cv-voice-1"
class FakeCloneRepo:
def __init__(self, profile: FakeProfile | None = None):
self._profile = profile
def get(self, pid: str) -> FakeProfile | None:
if self._profile and self._profile.id == pid:
return self._profile
return None
class FakeQuery:
def __init__(self, script: ScriptModel | None):
self._script = script
def filter(self, *conditions):
# 服务端写 filter(...).filter(...) 链式调用;归属/ID 已在 FakeDb 构造时过滤
return self
def first(self):
return self._script
class FakeDb:
def __init__(self, script: ScriptModel | None, *, current_user: str = "u1", query_script_id: str = "script-1"):
self._script = script
self._current_user = current_user
self._query_script_id = query_script_id
def query(self, model):
visible = self._script
if visible is not None and (visible.user_id != self._current_user or visible.id != self._query_script_id):
visible = None
return FakeQuery(visible)
def _make_script(*, user_id: str = "u1", content: str = "这是一段口播文案", title: str = "测试文案", tags=None):
return ScriptModel(
id="script-1",
user_id=user_id,
title=title,
content=content,
segments=[],
tags=tags if tags is not None else ["带货"],
)
@dataclass
class FakeLibrary:
id: str = "lib-voice"
project_id: str = "p1"
kind: Any = field(default_factory=lambda: SimpleNamespace(value="voice"))
@dataclass
class FakeProject:
id: str = "p1"
class FakeProjectRepo:
def __init__(self, projects=None):
self._projects = projects if projects is not None else [FakeProject()]
def find_accessible_projects(self, user_id):
return self._projects
class FakeLibraryRepo:
def __init__(self, libs=None):
self._libs = libs if libs is not None else [FakeLibrary()]
self.created: list = []
def find_by_project(self, project_id):
return list(self._libs)
def create(self, library):
self.created.append(library)
return library
@dataclass
class FakeAsset:
id: str = "asset-new"
duration: float | None = 12.0
class FakeAssetRepo:
def __init__(self):
self.created: list = []
def create(self, asset):
wrapped = FakeAsset(id="asset-new", duration=getattr(asset, "duration", None))
self.created.append(asset)
return wrapped
class FakeStorage:
def __init__(self, *, fail_download: bool = False):
self.fail_download = fail_download
self.uploaded: list = []
def download_asset(self, source, dest_path) -> bool:
if self.fail_download:
return False
dest_path.write_bytes(b"FAKEAUDIO")
return True
def upload_file(self, path, key, content_type="", **kwargs):
self.uploaded.append((key, content_type))
def delete_file(self, key):
pass
class FakeTTSRepo:
def __init__(self, job: TTSJob):
self.job = job
self.saved: list[TTSJob] = []
def create(self, job: TTSJob) -> TTSJob:
self.saved.append(job)
self.job = job
return job
def update(self, job: TTSJob) -> TTSJob:
self.job = job
return job
def get(self, job_id: str) -> TTSJob | None:
return self.job if self.job.id == job_id else None
class FakeCosyVoice:
pass
def _make_completed_job() -> TTSJob:
job = TTSJob.create(
user_id="u1",
input_text="这是一段口播文案",
voice_id="cv-voice-1",
voice_clone_profile_id="",
format="mp3",
sample_rate=22050,
)
job.mark_processing()
job.mark_completed(
output_audio_url="https://oss/tts/output/job-1.mp3",
output_audio_key="tts/output/job-1.mp3",
duration=12.5,
)
return job
# ── _resolve_voice ─────────────────────────────────────────────────────────
class TestResolveVoice:
def test_preset_returns_id_directly_when_no_profile(self):
voice_id, clone_id = _resolve_voice(
user_id="u1",
tts_voice_id="longxiaochun",
tts_voice_source="preset",
voice_clone_repository=FakeCloneRepo(None),
)
assert voice_id == "longxiaochun"
assert clone_id == ""
def test_preset_id_that_is_clone_profile_uuid_resolves(self):
repo = FakeCloneRepo(FakeProfile())
voice_id, clone_id = _resolve_voice(
user_id="u1",
tts_voice_id="prof-1",
tts_voice_source="preset",
voice_clone_repository=repo,
)
assert voice_id == "cv-voice-1"
assert clone_id == "prof-1"
def test_clone_source(self):
voice_id, clone_id = _resolve_voice(
user_id="u1",
tts_voice_id="prof-1",
tts_voice_source="clone",
voice_clone_repository=FakeCloneRepo(FakeProfile()),
)
assert voice_id == "cv-voice-1"
assert clone_id == "prof-1"
def test_clone_missing_404(self):
with pytest.raises(NarrativeError) as ei:
_resolve_voice(
user_id="u1",
tts_voice_id="nope",
tts_voice_source="clone",
voice_clone_repository=FakeCloneRepo(None),
)
assert ei.value.status_code == 404
def test_clone_other_user_403(self):
repo = FakeCloneRepo(FakeProfile(user_id="someone-else"))
with pytest.raises(NarrativeError) as ei:
_resolve_voice(
user_id="u1",
tts_voice_id="prof-1",
tts_voice_source="clone",
voice_clone_repository=repo,
)
assert ei.value.status_code == 403
def test_clone_not_ready_400(self):
repo = FakeCloneRepo(FakeProfile(voice_id=""))
with pytest.raises(NarrativeError) as ei:
_resolve_voice(
user_id="u1",
tts_voice_id="prof-1",
tts_voice_source="clone",
voice_clone_repository=repo,
)
assert ei.value.status_code == 400
# ── save asset ─────────────────────────────────────────────────────────────
class TestSaveVoiceAsset:
def _deps(self, **storage_kw):
return dict(
user_id="u1",
name="测试配音",
project_repository=FakeProjectRepo(),
asset_library_repository=FakeLibraryRepo(),
asset_repository=FakeAssetRepo(),
storage_service=FakeStorage(**storage_kw),
)
def test_save_creates_asset(self):
job = _make_completed_job()
deps = self._deps()
asset = _save_tts_job_as_voice_asset(job=job, **deps)
assert asset.id == "asset-new"
assert deps["asset_repository"].created[0].mime_type == "audio/mpeg"
assert deps["storage_service"].uploaded[0][0] == "uploads/voice/tts/" + job.id + ".mp3"
def test_no_project_raises(self):
job = _make_completed_job()
deps = self._deps()
deps["project_repository"] = FakeProjectRepo(projects=[])
with pytest.raises(NarrativeError):
_save_tts_job_as_voice_asset(job=job, **deps)
def test_download_fail_raises_502(self):
job = _make_completed_job()
deps = self._deps(fail_download=True)
with pytest.raises(NarrativeError) as ei:
_save_tts_job_as_voice_asset(job=job, **deps)
assert ei.value.status_code == 502
def test_job_without_output_raises(self):
job = TTSJob.create(user_id="u1", input_text="x", voice_id="v", voice_clone_profile_id="")
with pytest.raises(NarrativeError) as ei:
_save_tts_job_as_voice_asset(job=job, **self._deps())
assert ei.value.status_code == 502
# ── prepare_narrative_voice 主流程(monkeypatch workflow) ─────────────────
class TestPrepareNarrativeVoice:
def _deps(self, db_script=None, *, has_script=True, clone_profile=None, storage_fail=False, points_enabled=False):
job = _make_completed_job()
return dict(
db=FakeDb(db_script if db_script is not None else (_make_script() if has_script else None)),
user_id="u1",
script_id="script-1",
tts_voice_id="longxiaochun",
tts_voice_source="preset",
tts_repository=FakeTTSRepo(job),
cosyvoice_service=FakeCosyVoice(),
voice_clone_repository=FakeCloneRepo(clone_profile),
asset_repository=FakeAssetRepo(),
asset_library_repository=FakeLibraryRepo(),
project_repository=FakeProjectRepo(),
storage_service=FakeStorage(fail_download=storage_fail),
points_enabled=points_enabled,
)
def test_success_returns_context(self, monkeypatch):
captured = {}
class FakeWorkflow:
def __init__(self, *, repository, cosyvoice_service):
captured["repo"] = repository
self._repo = repository
def start_synthesis(self, job_id):
job = self._repo.get(job_id)
job.mark_processing()
job.mark_completed(
output_audio_url="https://oss/tts/output/x.mp3",
output_audio_key="tts/output/x.mp3",
duration=12.5,
)
return job
def poll_and_process_synthesis(self, job_id, timeout=120.0):
return self._repo.get(job_id)
monkeypatch.setattr(ns, "TTSWorkflowService", FakeWorkflow)
ctx = prepare_narrative_voice(**self._deps())
assert ctx.voice_asset_id == "asset-new"
assert ctx.tts_job_id
assert ctx.audio_duration == pytest.approx(12.5)
assert ctx.script.tags == ["带货"]
def test_script_missing_404(self):
deps = self._deps(has_script=False)
with pytest.raises(NarrativeError) as ei:
prepare_narrative_voice(**deps)
assert ei.value.status_code == 404
def test_script_other_user_404(self):
deps = self._deps(db_script=_make_script(user_id="other"))
with pytest.raises(NarrativeError) as ei:
prepare_narrative_voice(**deps)
assert ei.value.status_code == 404
def test_empty_content_400(self):
deps = self._deps(db_script=_make_script(content=" "))
with pytest.raises(NarrativeError) as ei:
prepare_narrative_voice(**deps)
assert ei.value.status_code == 400
def test_synth_failure_raises_502(self, monkeypatch):
class FailingWorkflow:
def __init__(self, *, repository, cosyvoice_service):
self._repo = repository
def start_synthesis(self, job_id):
raise RuntimeError("cosyvoice down")
def process_synthesis_failure(self, job_id, error):
return None
monkeypatch.setattr(ns, "TTSWorkflowService", FailingWorkflow)
with pytest.raises(NarrativeError) as ei:
prepare_narrative_voice(**self._deps())
assert ei.value.status_code == 502
assert "配音合成失败" in ei.value.message
def test_points_insufficient_402(self, monkeypatch):
class FakePoints:
def deduct_points(self, *a, **k):
return {"success": False, "balance": 0}
monkeypatch.setattr(ns, "PointsService", lambda: FakePoints())
deps = self._deps(points_enabled=True)
with pytest.raises(NarrativeError) as ei:
prepare_narrative_voice(**deps)
assert ei.value.status_code == 402
def test_points_refund_on_failure(self, monkeypatch):
class FakePoints:
def __init__(self):
self.refunded = 0
def deduct_points(self, *a, **k):
return {"success": True, "balance": 100}
def refund_points(self, user_id, amount, source, db, ref_id="", **k):
self.refunded += amount
points = FakePoints()
monkeypatch.setattr(ns, "PointsService", lambda: points)
class FailingWorkflow:
def __init__(self, *, repository, cosyvoice_service):
pass
def start_synthesis(self, job_id):
raise RuntimeError("boom")
def process_synthesis_failure(self, job_id, error):
return None
monkeypatch.setattr(ns, "TTSWorkflowService", FailingWorkflow)
deps = self._deps(points_enabled=True)
with pytest.raises(NarrativeError):
prepare_narrative_voice(**deps)
assert points.refunded > 0
def test_clone_source_resolves_profile(self, monkeypatch):
captured = {}
class FakeWorkflow:
def __init__(self, *, repository, cosyvoice_service):
self._repo = repository
captured["cosy"] = cosyvoice_service
def start_synthesis(self, job_id):
job = self._repo.get(job_id)
captured["voice_id"] = job.voice_id
job.mark_processing()
job.mark_completed(
output_audio_url="https://oss/tts/output/x.mp3",
output_audio_key="tts/output/x.mp3",
duration=12.5,
)
return job
def poll_and_process_synthesis(self, job_id, timeout=120.0):
return self._repo.get(job_id)
monkeypatch.setattr(ns, "TTSWorkflowService", FakeWorkflow)
deps = self._deps(clone_profile=FakeProfile())
deps["tts_voice_id"] = "prof-1"
deps["tts_voice_source"] = "clone"
prepare_narrative_voice(**deps)
assert captured["voice_id"] == "cv-voice-1"
if __name__ == "__main__":
import pytest as _pytest
_pytest.main([__file__, "-q"])