"""Template use cases 单元测试.""" from __future__ import annotations from typing import List, Optional from unittest.mock import MagicMock import pytest from packages.application.template.commands import ( CopyTemplateCommand, CreateCategoryCommand, CreateTemplateCommand, ListTemplatesFilter, SegmentCommand, UpdateTemplateCommand, ValidateTemplateCommand, ) from packages.application.template.use_cases import ( CopyTemplateUseCase, CountTemplatesUseCase, CreateCategoryUseCase, CreateTemplateUseCase, DeleteCategoryUseCase, DeleteTemplateUseCase, GenerateWarning, GetTemplateUsageUseCase, GetTemplateUseCase, ListCategoriesUseCase, ListTagsUseCase, ListTemplatesUseCase, NotFoundError, UpdateTemplateUseCase, ValidateResult, ValidateTemplateUseCase, ValidationError, ) from packages.domain.editing_mode import EditingMode from packages.domain.template import Template, TemplateCategory, TemplateSegment def _make_template( template_id: str = "tpl_001", user_id: str = "user_001", name: str = "测试模板", mode: str = "one_take", segments: Optional[List[TemplateSegment]] = None, estimated_duration: float = 60.0, ) -> Template: tpl = Template( id=template_id, user_id=user_id, name=name, mode=mode, category="测试分类", tags=["tag1", "tag2"], title_config={"enabled": True}, subtitle_config={"enabled": False}, bgm_config={"enabled": True}, estimated_duration=estimated_duration, ) if segments is not None: tpl.segments = segments return tpl def _make_segments(count: int = 1, material_type: Optional[str] = None) -> List[TemplateSegment]: return [ TemplateSegment( id=f"seg_{i}", template_id="tpl_001", segment_order=i, duration_min=3.0, duration_max=8.0, material_type=material_type, ) for i in range(count) ] class TestCreateTemplateUseCase: def test_creates_template_with_segments(self) -> None: repo = MagicMock() repo.create.side_effect = lambda t: t # 返回传入的template repo.create_segments.return_value = None use_case = CreateTemplateUseCase(repo) cmd = CreateTemplateCommand( user_id="user_001", name="新模板", mode=EditingMode.ONE_TAKE.value, category="分类A", tags=["t1", "t2"], segments=[ SegmentCommand(segment_order=0, duration_min=2.0, duration_max=5.0), SegmentCommand(segment_order=1, duration_min=3.0, duration_max=6.0), ], ) result = use_case.execute(cmd) assert result.name == "新模板" assert result.mode == EditingMode.ONE_TAKE.value assert len(result.segments) == 2 assert result.segments[0].segment_order == 0 assert result.segments[1].segment_order == 1 repo.create.assert_called_once() repo.create_segments.assert_called_once() def test_invalid_mode_raises_validation_error(self) -> None: repo = MagicMock() use_case = CreateTemplateUseCase(repo) cmd = CreateTemplateCommand( user_id="user_001", name="测试", mode="invalid_mode", segments=[], ) with pytest.raises(ValidationError, match="无效的剪辑模式"): use_case.execute(cmd) def test_creates_without_segments(self) -> None: repo = MagicMock() repo.create.side_effect = lambda t: t repo.create_segments.return_value = None use_case = CreateTemplateUseCase(repo) cmd = CreateTemplateCommand( user_id="user_001", name="空片段模板", mode=EditingMode.ONE_TAKE.value, segments=[], ) result = use_case.execute(cmd) assert len(result.segments) == 0 repo.create_segments.assert_called_once_with([]) def test_generates_uuid_for_template_and_segments(self) -> None: repo = MagicMock() repo.create.side_effect = lambda t: t repo.create_segments.return_value = None use_case = CreateTemplateUseCase(repo) cmd = CreateTemplateCommand( user_id="user_001", name="UUID测试", mode=EditingMode.VOICE_OVER.value, segments=[ SegmentCommand(segment_order=0, duration_min=1.0, duration_max=3.0, material_type="人物"), ], ) result = use_case.execute(cmd) assert len(result.id) == 32 # uuid hex assert len(result.segments[0].id) == 32 assert result.segments[0].template_id == result.id class TestListTemplatesUseCase: def test_list_without_filter(self) -> None: repo = MagicMock() expected = [_make_template("t1"), _make_template("t2")] repo.list_by_user.return_value = expected use_case = ListTemplatesUseCase(repo) result = use_case.execute("user_001", skip=0, limit=10) assert len(result) == 2 repo.list_by_user.assert_called_once_with("user_001", skip=0, limit=10) def test_list_with_filter(self) -> None: repo = MagicMock() expected = [_make_template("t1")] repo.list_by_user.return_value = expected use_case = ListTemplatesUseCase(repo) f = ListTemplatesFilter(category="分类A", tag="t1", keyword="测试", mode="one_take") result = use_case.execute("user_001", skip=0, limit=10, filter=f) assert len(result) == 1 repo.list_by_user.assert_called_once_with( "user_001", skip=0, limit=10, category="分类A", tag="t1", keyword="测试", mode="one_take", ) class TestCountTemplatesUseCase: def test_count_without_filter(self) -> None: repo = MagicMock() repo.count_by_user.return_value = 42 use_case = CountTemplatesUseCase(repo) result = use_case.execute("user_001") assert result == 42 repo.count_by_user.assert_called_once_with("user_001") def test_count_with_filter(self) -> None: repo = MagicMock() repo.count_by_user.return_value = 5 use_case = CountTemplatesUseCase(repo) f = ListTemplatesFilter(category="分类A") result = use_case.execute("user_001", filter=f) assert result == 5 repo.count_by_user.assert_called_once_with( "user_001", category="分类A", tag=None, keyword=None, mode=None, ) class TestGetTemplateUseCase: def test_returns_template_when_found(self) -> None: repo = MagicMock() expected = _make_template() repo.get.return_value = expected use_case = GetTemplateUseCase(repo) result = use_case.execute("tpl_001", "user_001") assert result is expected repo.get.assert_called_once_with("tpl_001", "user_001") def test_returns_none_when_not_found(self) -> None: repo = MagicMock() repo.get.return_value = None use_case = GetTemplateUseCase(repo) result = use_case.execute("nonexistent", "user_001") assert result is None class TestUpdateTemplateUseCase: def test_updates_name_and_tags(self) -> None: repo = MagicMock() existing = _make_template() existing.segments = _make_segments(2) repo.get.return_value = existing repo.update.side_effect = lambda t: t repo.list_segments.return_value = existing.segments use_case = UpdateTemplateUseCase(repo) cmd = UpdateTemplateCommand( template_id="tpl_001", user_id="user_001", name="新名字", tags=["new_tag"], ) result = use_case.execute(cmd) assert result.name == "新名字" assert result.tags == ["new_tag"] # mode没变 assert result.mode == EditingMode.ONE_TAKE.value repo.update.assert_called_once() def test_not_found_raises(self) -> None: repo = MagicMock() repo.get.return_value = None use_case = UpdateTemplateUseCase(repo) cmd = UpdateTemplateCommand(template_id="nonexistent", user_id="user_001", name="x") with pytest.raises(NotFoundError): use_case.execute(cmd) def test_invalid_mode_raises(self) -> None: repo = MagicMock() repo.get.return_value = _make_template() use_case = UpdateTemplateUseCase(repo) cmd = UpdateTemplateCommand( template_id="tpl_001", user_id="user_001", mode="invalid", ) with pytest.raises(ValidationError, match="无效的剪辑模式"): use_case.execute(cmd) def test_replaces_segments_when_provided(self) -> None: repo = MagicMock() existing = _make_template() existing.segments = _make_segments(2) repo.get.return_value = existing repo.update.side_effect = lambda t: t repo.delete_segments_by_template.return_value = None repo.create_segments.return_value = None use_case = UpdateTemplateUseCase(repo) cmd = UpdateTemplateCommand( template_id="tpl_001", user_id="user_001", segments=[ SegmentCommand(segment_order=0, duration_min=1.0, duration_max=2.0), SegmentCommand(segment_order=1, duration_min=3.0, duration_max=4.0), SegmentCommand(segment_order=2, duration_min=5.0, duration_max=6.0), ], ) result = use_case.execute(cmd) assert len(result.segments) == 3 repo.delete_segments_by_template.assert_called_once_with("tpl_001") repo.create_segments.assert_called_once() def test_no_segments_keeps_existing(self) -> None: repo = MagicMock() existing = _make_template() existing.segments = _make_segments(3) repo.get.return_value = existing repo.update.side_effect = lambda t: t repo.list_segments.return_value = existing.segments use_case = UpdateTemplateUseCase(repo) cmd = UpdateTemplateCommand( template_id="tpl_001", user_id="user_001", name="只改名字", ) result = use_case.execute(cmd) assert len(result.segments) == 3 repo.delete_segments_by_template.assert_not_called() repo.create_segments.assert_not_called() repo.list_segments.assert_called_once_with("tpl_001") class TestDeleteTemplateUseCase: def test_delete_success(self) -> None: repo = MagicMock() repo.delete.return_value = True use_case = DeleteTemplateUseCase(repo) result = use_case.execute("tpl_001", "user_001") assert result is True repo.delete.assert_called_once_with("tpl_001", "user_001") def test_delete_not_found(self) -> None: repo = MagicMock() repo.delete.return_value = False use_case = DeleteTemplateUseCase(repo) result = use_case.execute("nonexistent", "user_001") assert result is False class TestCopyTemplateUseCase: def test_copy_success(self) -> None: repo = MagicMock() original = _make_template(name="原模板") repo.get.return_value = original copied = _make_template(template_id="copied_001", name="原模板 副本") repo.copy_template.return_value = copied use_case = CopyTemplateUseCase(repo) cmd = CopyTemplateCommand( template_id="tpl_001", user_id="user_001", new_name="原模板 副本", ) result = use_case.execute(cmd) assert result.name == "原模板 副本" repo.copy_template.assert_called_once_with("tpl_001", "user_001", "原模板 副本") def test_not_found_raises(self) -> None: repo = MagicMock() repo.get.return_value = None use_case = CopyTemplateUseCase(repo) cmd = CopyTemplateCommand(template_id="no", user_id="u1", new_name="x") with pytest.raises(NotFoundError): use_case.execute(cmd) def test_empty_name_raises(self) -> None: repo = MagicMock() repo.get.return_value = _make_template() use_case = CopyTemplateUseCase(repo) cmd = CopyTemplateCommand(template_id="t1", user_id="u1", new_name=" ") with pytest.raises(ValidationError, match="名称不能为空"): use_case.execute(cmd) def test_name_stripped(self) -> None: repo = MagicMock() repo.get.return_value = _make_template() repo.copy_template.return_value = _make_template(name="新名字") use_case = CopyTemplateUseCase(repo) cmd = CopyTemplateCommand(template_id="t1", user_id="u1", new_name=" 新名字 ") use_case.execute(cmd) repo.copy_template.assert_called_once_with("t1", "u1", "新名字") class TestValidateTemplateUseCase: def test_one_take_with_one_segment_passes(self) -> None: repo = MagicMock() tpl = _make_template(mode=EditingMode.ONE_TAKE.value) tpl.segments = _make_segments(1) repo.get.return_value = tpl use_case = ValidateTemplateUseCase(repo) cmd = ValidateTemplateCommand(template_id="tpl_001", user_id="user_001") result = use_case.execute(cmd) assert isinstance(result, ValidateResult) assert result.template is tpl assert len(result.warnings) == 0 def test_one_take_with_multiple_segments_raises(self) -> None: repo = MagicMock() tpl = _make_template(mode=EditingMode.ONE_TAKE.value) tpl.segments = _make_segments(3) repo.get.return_value = tpl use_case = ValidateTemplateUseCase(repo) cmd = ValidateTemplateCommand(template_id="tpl_001", user_id="user_001") with pytest.raises(ValidationError, match="恰好有 1 个片段"): use_case.execute(cmd) def test_voice_over_with_valid_material_types_passes(self) -> None: repo = MagicMock() tpl = _make_template(mode=EditingMode.VOICE_OVER.value, estimated_duration=30.0) tpl.segments = [ TemplateSegment( id="s1", template_id="t1", segment_order=0, duration_min=3, duration_max=5, material_type="人物" ), TemplateSegment( id="s2", template_id="t1", segment_order=1, duration_min=3, duration_max=5, material_type="场景" ), ] repo.get.return_value = tpl use_case = ValidateTemplateUseCase(repo) cmd = ValidateTemplateCommand(template_id="tpl_001", user_id="user_001") result = use_case.execute(cmd) assert len(result.warnings) == 0 def test_voice_over_missing_material_type_raises(self) -> None: repo = MagicMock() tpl = _make_template(mode=EditingMode.VOICE_OVER.value) tpl.segments = [ TemplateSegment( id="s1", template_id="t1", segment_order=0, duration_min=3, duration_max=5, material_type=None ), ] repo.get.return_value = tpl use_case = ValidateTemplateUseCase(repo) cmd = ValidateTemplateCommand(template_id="tpl_001", user_id="user_001") with pytest.raises(ValidationError, match="material_type"): use_case.execute(cmd) def test_voice_over_invalid_material_type_raises(self) -> None: repo = MagicMock() tpl = _make_template(mode=EditingMode.VOICE_OVER.value) tpl.segments = [ TemplateSegment( id="s1", template_id="t1", segment_order=0, duration_min=3, duration_max=5, material_type="动物" ), ] repo.get.return_value = tpl use_case = ValidateTemplateUseCase(repo) cmd = ValidateTemplateCommand(template_id="tpl_001", user_id="user_001") with pytest.raises(ValidationError, match="material_type"): use_case.execute(cmd) def test_voiceover_duration_within_range_no_warning(self) -> None: repo = MagicMock() tpl = _make_template(mode=EditingMode.ONE_TAKE.value, estimated_duration=60.0) tpl.segments = _make_segments(1) repo.get.return_value = tpl use_case = ValidateTemplateUseCase(repo) # 65s vs 60s = 1.08 ratio,在±30%内 cmd = ValidateTemplateCommand(template_id="t1", user_id="u1", voiceover_duration=65.0) result = use_case.execute(cmd) assert len(result.warnings) == 0 def test_voiceover_duration_too_short_warns(self) -> None: repo = MagicMock() tpl = _make_template(mode=EditingMode.ONE_TAKE.value, estimated_duration=60.0) tpl.segments = _make_segments(1) repo.get.return_value = tpl use_case = ValidateTemplateUseCase(repo) # 20s vs 60s = 0.33 ratio,超过±30% cmd = ValidateTemplateCommand(template_id="t1", user_id="u1", voiceover_duration=20.0) result = use_case.execute(cmd) assert len(result.warnings) == 1 assert result.warnings[0].code == "voiceover_duration_mismatch" assert "偏差超过" in result.warnings[0].message assert result.warnings[0].details["ratio"] < 0.7 def test_voiceover_duration_too_long_warns(self) -> None: repo = MagicMock() tpl = _make_template(mode=EditingMode.ONE_TAKE.value, estimated_duration=60.0) tpl.segments = _make_segments(1) repo.get.return_value = tpl use_case = ValidateTemplateUseCase(repo) # 100s vs 60s = 1.67 ratio,超过±30% cmd = ValidateTemplateCommand(template_id="t1", user_id="u1", voiceover_duration=100.0) result = use_case.execute(cmd) assert len(result.warnings) == 1 assert result.warnings[0].code == "voiceover_duration_mismatch" assert result.warnings[0].details["ratio"] > 1.3 def test_zero_estimated_duration_no_warning(self) -> None: repo = MagicMock() tpl = _make_template(mode=EditingMode.ONE_TAKE.value, estimated_duration=0.0) tpl.segments = _make_segments(1) repo.get.return_value = tpl use_case = ValidateTemplateUseCase(repo) cmd = ValidateTemplateCommand(template_id="t1", user_id="u1", voiceover_duration=50.0) result = use_case.execute(cmd) # estimated_duration=0不做偏差检查 assert len(result.warnings) == 0 def test_no_voiceover_duration_no_warning(self) -> None: repo = MagicMock() tpl = _make_template(estimated_duration=60.0) tpl.segments = _make_segments(1) repo.get.return_value = tpl use_case = ValidateTemplateUseCase(repo) cmd = ValidateTemplateCommand(template_id="t1", user_id="u1") # 不传voiceover_duration result = use_case.execute(cmd) assert len(result.warnings) == 0 def test_not_found_raises(self) -> None: repo = MagicMock() repo.get.return_value = None use_case = ValidateTemplateUseCase(repo) cmd = ValidateTemplateCommand(template_id="no", user_id="u1") with pytest.raises(NotFoundError): use_case.execute(cmd) class TestCategoryUseCases: def test_create_category(self) -> None: repo = MagicMock() cat = TemplateCategory(id="cat_001", user_id="u1", name="新分类") repo.create_category.return_value = cat use_case = CreateCategoryUseCase(repo) cmd = CreateCategoryCommand(user_id="u1", name="新分类") result = use_case.execute(cmd) assert result.name == "新分类" repo.create_category.assert_called_once() def test_list_categories(self) -> None: repo = MagicMock() expected = [TemplateCategory(id="c1", user_id="u1", name="A")] repo.list_categories.return_value = expected use_case = ListCategoriesUseCase(repo) result = use_case.execute("u1") assert result == expected repo.list_categories.assert_called_once_with("u1") def test_delete_category(self) -> None: repo = MagicMock() repo.delete_category.return_value = True use_case = DeleteCategoryUseCase(repo) result = use_case.execute("cat_001", "u1") assert result is True repo.delete_category.assert_called_once_with("cat_001", "u1") class TestListTagsUseCase: def test_returns_tags_list(self) -> None: repo = MagicMock() repo.list_tags.return_value = ["tag1", "tag2", "tag3"] use_case = ListTagsUseCase(repo) result = use_case.execute("u1") assert result == ["tag1", "tag2", "tag3"] repo.list_tags.assert_called_once_with("u1") class TestGetTemplateUsageUseCase: def test_returns_usage_count(self) -> None: repo = MagicMock() repo.get_usage_count.return_value = 15 use_case = GetTemplateUsageUseCase(repo) result = use_case.execute("tpl_001") assert result == 15 repo.get_usage_count.assert_called_once_with("tpl_001") class TestGenerateWarning: def test_warning_default_details(self) -> None: w = GenerateWarning(code="test_code", message="test message") assert w.code == "test_code" assert w.message == "test message" assert w.details == {} def test_warning_with_details(self) -> None: w = GenerateWarning(code="test", message="msg", details={"key": "value"}) assert w.details == {"key": "value"}