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
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:
@@ -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(
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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"])
|
||||
@@ -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"])
|
||||
Reference in New Issue
Block a user