"""#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)