From 4fa3e4eb924f652e8b28f174c761d34e462e5433 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Fri, 18 Sep 2026 07:30:43 +0800 Subject: [PATCH] =?UTF-8?q?feat(#1970):=20=E6=96=B0=20API=20=E5=AD=97?= =?UTF-8?q?=E6=AE=B5=20+=20=E5=8F=99=E4=BA=8B=E6=A8=A1=E5=BC=8F=20PR3=20-?= =?UTF-8?q?=20assembly=5Fmode/script=5Fid/tts=5F*/video=5Fratio=20(#1976)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- apps/api/app/api/routes/generation_tasks.py | 149 ++++++- apps/api/app/schemas/generation_task.py | 33 ++ apps/api/app/services/generation_common.py | 12 +- apps/api/app/services/narrative_service.py | 344 +++++++++++++++ packages/domain/editing_mode.py | 22 +- packages/domain/narrative_match.py | 132 ++++++ tests/unit/test_1970_assembly_schema.py | 218 ++++++++++ tests/unit/test_1970_narrative_match.py | 167 +++++++ tests/unit/test_1970_narrative_service.py | 454 ++++++++++++++++++++ 9 files changed, 1521 insertions(+), 10 deletions(-) create mode 100644 apps/api/app/services/narrative_service.py create mode 100644 packages/domain/narrative_match.py create mode 100644 tests/unit/test_1970_assembly_schema.py create mode 100644 tests/unit/test_1970_narrative_match.py create mode 100644 tests/unit/test_1970_narrative_service.py diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index a63a6c725..560a81750 100755 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -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( diff --git a/apps/api/app/schemas/generation_task.py b/apps/api/app/schemas/generation_task.py index a54e2d0e0..afab39fa8 100755 --- a/apps/api/app/schemas/generation_task.py +++ b/apps/api/app/schemas/generation_task.py @@ -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()) diff --git a/apps/api/app/services/generation_common.py b/apps/api/app/services/generation_common.py index 8ad6eb623..9cd0d54ee 100644 --- a/apps/api/app/services/generation_common.py +++ b/apps/api/app/services/generation_common.py @@ -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") diff --git a/apps/api/app/services/narrative_service.py b/apps/api/app/services/narrative_service.py new file mode 100644 index 000000000..71ea39e6b --- /dev/null +++ b/apps/api/app/services/narrative_service.py @@ -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), + ) diff --git a/packages/domain/editing_mode.py b/packages/domain/editing_mode.py index 69a34d318..ce66e6a5f 100644 --- a/packages/domain/editing_mode.py +++ b/packages/domain/editing_mode.py @@ -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 diff --git a/packages/domain/narrative_match.py b/packages/domain/narrative_match.py new file mode 100644 index 000000000..5d72c6766 --- /dev/null +++ b/packages/domain/narrative_match.py @@ -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 diff --git a/tests/unit/test_1970_assembly_schema.py b/tests/unit/test_1970_assembly_schema.py new file mode 100644 index 000000000..ce38f3b3d --- /dev/null +++ b/tests/unit/test_1970_assembly_schema.py @@ -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) diff --git a/tests/unit/test_1970_narrative_match.py b/tests/unit/test_1970_narrative_match.py new file mode 100644 index 000000000..e3ce7d5e2 --- /dev/null +++ b/tests/unit/test_1970_narrative_match.py @@ -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"]) diff --git a/tests/unit/test_1970_narrative_service.py b/tests/unit/test_1970_narrative_service.py new file mode 100644 index 000000000..5a00df0a0 --- /dev/null +++ b/tests/unit/test_1970_narrative_service.py @@ -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"])