""" Template Use Cases 单元测试 — 剪辑计划模板 CRUD + 业务规则校验 """ from unittest.mock import Mock 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, DeleteTemplateUseCase, GetTemplateUsageUseCase, GetTemplateUseCase, ListCategoriesUseCase, ListTagsUseCase, ListTemplatesUseCase, NotFoundError, UpdateTemplateUseCase, ValidateTemplateUseCase, ValidationError, ) from packages.domain.template import Template, TemplateCategory, TemplateSegment def _make_repo(): """创建一个 mock repository.""" repo = Mock() repo.list_by_user = Mock(return_value=[]) repo.get = Mock(return_value=None) repo.create = Mock() repo.update = Mock() repo.delete = Mock(return_value=False) repo.count_by_user = Mock(return_value=0) repo.list_segments = Mock(return_value=[]) repo.create_segments = Mock() repo.delete_segments_by_template = Mock(return_value=0) repo.list_categories = Mock(return_value=[]) repo.create_category = Mock() repo.get_category = Mock(return_value=None) repo.delete_category = Mock(return_value=False) repo.copy_template = Mock() repo.list_tags = Mock(return_value=[]) repo.get_usage_count = Mock(return_value=0) return repo def _make_template(**kwargs) -> Template: defaults = dict( id="tmpl-001", user_id="user-001", name="测试模板", mode="pip", category="default", tags=["test"], title_config={"ai_auto_select": True}, subtitle_config={"enabled": True}, bgm_config={"enabled": False}, estimated_duration=60.0, segments=[], ) defaults.update(kwargs) return Template(**defaults) # ── CreateTemplateUseCase ── class TestCreateTemplateUseCase: @pytest.fixture def repo(self): return _make_repo() @pytest.fixture def use_case(self, repo): return CreateTemplateUseCase(repo) def test_create_basic_template(self, use_case, repo): """创建基础模板(无片段).""" repo.create.side_effect = lambda t: t # 返回传入的 template command = CreateTemplateCommand( user_id="user-001", name="画中画模板", mode="pip", category="vlog", tags=["vlog", "pip"], estimated_duration=90.0, ) result = use_case.execute(command) assert result.name == "画中画模板" assert result.mode == "pip" assert result.user_id == "user-001" repo.create.assert_called_once() def test_create_with_segments(self, use_case, repo): """创建模板并附带片段.""" repo.create.side_effect = lambda t: t repo.create_segments.side_effect = lambda segs: segs command = CreateTemplateCommand( user_id="user-001", name="口播混剪模板", mode="voice_over", segments=[ SegmentCommand(segment_order=1, duration_min=5, duration_max=15, material_type="人物"), SegmentCommand(segment_order=2, duration_min=10, duration_max=30, material_type="场景"), ], ) result = use_case.execute(command) assert len(result.segments) == 2 assert result.segments[0].material_type == "人物" repo.create_segments.assert_called_once() def test_create_invalid_mode_raises(self, use_case): """无效剪辑模式应抛出 ValidationError.""" command = CreateTemplateCommand( user_id="user-001", name="无效模板", mode="invalid_mode", ) with pytest.raises(ValidationError, match="无效的剪辑模式"): use_case.execute(command) # ── UpdateTemplateUseCase ── class TestUpdateTemplateUseCase: @pytest.fixture def repo(self): return _make_repo() @pytest.fixture def use_case(self, repo): return UpdateTemplateUseCase(repo) def test_update_name(self, use_case, repo): """更新模板名称.""" existing = _make_template() repo.get.return_value = existing repo.update.side_effect = lambda t: t command = UpdateTemplateCommand( template_id="tmpl-001", user_id="user-001", name="新名称", ) result = use_case.execute(command) assert result.name == "新名称" repo.update.assert_called_once() def test_update_not_found_raises(self, use_case, repo): """模板不存在时抛出 NotFoundError.""" repo.get.return_value = None command = UpdateTemplateCommand( template_id="nonexistent", user_id="user-001", name="新名称", ) with pytest.raises(NotFoundError): use_case.execute(command) def test_update_invalid_mode_raises(self, use_case, repo): """更新为无效模式时抛出 ValidationError.""" existing = _make_template() repo.get.return_value = existing command = UpdateTemplateCommand( template_id="tmpl-001", user_id="user-001", mode="bad_mode", ) with pytest.raises(ValidationError, match="无效的剪辑模式"): use_case.execute(command) def test_replace_segments(self, use_case, repo): """替换片段列表.""" existing = _make_template() repo.get.return_value = existing repo.update.side_effect = lambda t: t repo.create_segments.side_effect = lambda segs: segs command = UpdateTemplateCommand( template_id="tmpl-001", user_id="user-001", segments=[ SegmentCommand(segment_order=1, duration_min=5, duration_max=20, material_type=None), ], ) result = use_case.execute(command) repo.delete_segments_by_template.assert_called_once_with("tmpl-001") repo.create_segments.assert_called_once() assert len(result.segments) == 1 # ── ValidateTemplateUseCase — 业务规则校验 ── class TestValidateTemplateUseCase: @pytest.fixture def repo(self): return _make_repo() @pytest.fixture def use_case(self, repo): return ValidateTemplateUseCase(repo) def test_one_take_with_one_segment_ok(self, use_case, repo): """一镜到底 + 恰好 1 个片段 → 通过.""" seg = TemplateSegment( id="seg-001", template_id="tmpl-001", segment_order=1, duration_min=0, duration_max=60, ) template = _make_template(mode="one_take", segments=[seg]) repo.get.return_value = template command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", ) result = use_case.execute(command) assert result.template.mode == "one_take" assert result.warnings == [] def test_one_take_with_two_segments_raises(self, use_case, repo): """一镜到底 + 2 个片段 → ValidationError.""" segs = [ TemplateSegment(id=f"seg-{i}", template_id="tmpl-001", segment_order=i, duration_min=0, duration_max=30) for i in (1, 2) ] template = _make_template(mode="one_take", segments=segs) repo.get.return_value = template command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", ) with pytest.raises(ValidationError, match="一镜到底模式必须恰好有 1 个片段"): use_case.execute(command) def test_voice_over_all_segments_have_material_type_ok(self, use_case, repo): """口播+B-roll + 所有片段都有 material_type → 通过.""" segs = [ TemplateSegment( id="seg-1", template_id="tmpl-001", segment_order=1, duration_min=5, duration_max=15, material_type="人物", ), TemplateSegment( id="seg-2", template_id="tmpl-001", segment_order=2, duration_min=10, duration_max=30, material_type="场景", ), ] template = _make_template(mode="voice_over", segments=segs) repo.get.return_value = template command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", ) result = use_case.execute(command) assert result.warnings == [] def test_voice_over_missing_material_type_raises(self, use_case, repo): """口播+B-roll + 某片段缺少 material_type → ValidationError.""" segs = [ TemplateSegment( id="seg-1", template_id="tmpl-001", segment_order=1, duration_min=5, duration_max=15, material_type="人物", ), TemplateSegment( id="seg-2", template_id="tmpl-001", segment_order=2, duration_min=10, duration_max=30, material_type=None, ), # 缺失 ] template = _make_template(mode="voice_over", segments=segs) repo.get.return_value = template command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", ) with pytest.raises(ValidationError, match="material_type"): use_case.execute(command) def test_voiceover_duration_within_tolerance_no_warning(self, use_case, repo): """配音时长在 ±30% 以内 → 无警告.""" template = _make_template(estimated_duration=60.0) repo.get.return_value = template command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", voiceover_duration=70.0, # 70/60 = 1.167, within ±30% ) result = use_case.execute(command) assert result.warnings == [] def test_voiceover_duration_exceeds_tolerance_warning(self, use_case, repo): """配音时长超过 ±30% → 警告.""" template = _make_template(estimated_duration=60.0) repo.get.return_value = template command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", voiceover_duration=100.0, # 100/60 = 1.667, exceeds +30% ) result = use_case.execute(command) assert len(result.warnings) == 1 assert result.warnings[0].code == "voiceover_duration_mismatch" def test_voiceover_duration_too_short_warning(self, use_case, repo): """配音时长过短(< 70%)→ 警告.""" template = _make_template(estimated_duration=60.0) repo.get.return_value = template command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", voiceover_duration=30.0, # 30/60 = 0.5, below -30% ) result = use_case.execute(command) assert len(result.warnings) == 1 assert result.warnings[0].code == "voiceover_duration_mismatch" def test_template_not_found_raises(self, use_case, repo): """模板不存在 → NotFoundError.""" repo.get.return_value = None command = ValidateTemplateCommand( template_id="nonexistent", user_id="user-001", ) with pytest.raises(NotFoundError): use_case.execute(command) # ── Category Use Cases ── class TestCategoryUseCases: @pytest.fixture def repo(self): return _make_repo() def test_create_category(self, repo): repo.create_category.side_effect = lambda c: c use_case = CreateCategoryUseCase(repo) command = CreateCategoryCommand(user_id="user-001", name="Vlog") result = use_case.execute(command) assert result.name == "Vlog" repo.create_category.assert_called_once() def test_list_categories(self, repo): categories = [ TemplateCategory(id="cat-1", user_id="user-001", name="Vlog"), TemplateCategory(id="cat-2", user_id="user-001", name="教程"), ] repo.list_categories.return_value = categories use_case = ListCategoriesUseCase(repo) result = use_case.execute("user-001") assert len(result) == 2 assert result[0].name == "Vlog" def test_delete_category_not_found(self, repo): repo.delete_category.return_value = False use_case = DeleteTemplateUseCase(repo) result = use_case.execute("nonexistent", "user-001") assert result is False # ── ListTemplatesUseCase ── class TestListTemplatesUseCase: def test_list_returns_templates(self): repo = _make_repo() templates = [_make_template(id=f"t-{i}") for i in range(3)] repo.list_by_user.return_value = templates use_case = ListTemplatesUseCase(repo) result = use_case.execute("user-001", skip=0, limit=50) assert len(result) == 3 repo.list_by_user.assert_called_once_with("user-001", skip=0, limit=50) # ── GetTemplateUseCase ── class TestGetTemplateUseCase: def test_get_existing(self): repo = _make_repo() template = _make_template() repo.get.return_value = template use_case = GetTemplateUseCase(repo) result = use_case.execute("tmpl-001", "user-001") assert result.id == "tmpl-001" def test_get_nonexistent_returns_none(self): repo = _make_repo() repo.get.return_value = None use_case = GetTemplateUseCase(repo) result = use_case.execute("nonexistent", "user-001") assert result is None # ── CopyTemplateUseCase ── class TestCopyTemplateUseCase: @pytest.fixture def repo(self): repo = _make_repo() source = _make_template( id="tmpl-src", name="源模板", segments=[ TemplateSegment( id="seg-1", template_id="tmpl-src", segment_order=0, duration_min=5.0, duration_max=10.0, material_type=None, ), ], ) repo.get = Mock(return_value=source) def _copy_side_effect(template_id, user_id, new_name): return _make_template( id="tmpl-copied", user_id=user_id, name=new_name, segments=[ TemplateSegment( id="seg-copied", template_id="tmpl-copied", segment_order=0, duration_min=5.0, duration_max=10.0, material_type=None, ) ], ) repo.copy_template = Mock(side_effect=_copy_side_effect) return repo @pytest.fixture def use_case(self, repo): return CopyTemplateUseCase(repo) def test_copy_success(self, use_case, repo): command = CopyTemplateCommand( template_id="tmpl-src", user_id="user-001", new_name="复制的模板", ) result = use_case.execute(command) assert result.id == "tmpl-copied" assert result.name == "复制的模板" assert len(result.segments) == 1 repo.copy_template.assert_called_once_with( "tmpl-src", "user-001", "复制的模板", ) def test_copy_not_found_raises(self, use_case, repo): repo.get = Mock(return_value=None) command = CopyTemplateCommand( template_id="tmpl-nonexist", user_id="user-001", new_name="新名字", ) with pytest.raises(NotFoundError): use_case.execute(command) def test_copy_empty_name_raises(self, use_case, repo): command = CopyTemplateCommand( template_id="tmpl-src", user_id="user-001", new_name=" ", ) with pytest.raises(ValidationError): use_case.execute(command) # ── ListTemplatesUseCase (filter) ── class TestListTemplatesUseCaseWithFilter: def test_list_with_category_filter(self): repo = _make_repo() repo.list_by_user = Mock(return_value=[]) use_case = ListTemplatesUseCase(repo) f = ListTemplatesFilter(category="vlog") use_case.execute("user-001", skip=0, limit=10, filter=f) repo.list_by_user.assert_called_once() call_kwargs = repo.list_by_user.call_args assert call_kwargs[1]["category"] == "vlog" def test_list_with_tag_filter(self): repo = _make_repo() repo.list_by_user = Mock(return_value=[]) use_case = ListTemplatesUseCase(repo) f = ListTemplatesFilter(tag="热门") use_case.execute("user-001", filter=f) call_kwargs = repo.list_by_user.call_args assert call_kwargs[1]["tag"] == "热门" def test_list_with_keyword_filter(self): repo = _make_repo() repo.list_by_user = Mock(return_value=[]) use_case = ListTemplatesUseCase(repo) f = ListTemplatesFilter(keyword="vlog") use_case.execute("user-001", filter=f) call_kwargs = repo.list_by_user.call_args assert call_kwargs[1]["keyword"] == "vlog" def test_list_with_mode_filter(self): repo = _make_repo() repo.list_by_user = Mock(return_value=[]) use_case = ListTemplatesUseCase(repo) f = ListTemplatesFilter(mode="one_take") use_case.execute("user-001", filter=f) call_kwargs = repo.list_by_user.call_args assert call_kwargs[1]["mode"] == "one_take" def test_list_without_filter_uses_defaults(self): repo = _make_repo() repo.list_by_user = Mock(return_value=[]) use_case = ListTemplatesUseCase(repo) use_case.execute("user-001", skip=0, limit=50) call_args = repo.list_by_user.call_args assert call_args[0][0] == "user-001" assert call_args[1]["skip"] == 0 assert call_args[1]["limit"] == 50 # ── CountTemplatesUseCase ── class TestCountTemplatesUseCase: def test_count_without_filter(self): repo = _make_repo() repo.count_by_user = Mock(return_value=5) use_case = CountTemplatesUseCase(repo) result = use_case.execute("user-001") assert result == 5 repo.count_by_user.assert_called_once_with("user-001") def test_count_with_filter(self): repo = _make_repo() repo.count_by_user = Mock(return_value=2) use_case = CountTemplatesUseCase(repo) f = ListTemplatesFilter(category="vlog", tag="热门") result = use_case.execute("user-001", filter=f) assert result == 2 call_kwargs = repo.count_by_user.call_args assert call_kwargs[1]["category"] == "vlog" assert call_kwargs[1]["tag"] == "热门" # ── ListTagsUseCase ── class TestListTagsUseCase: def test_list_tags_returns_sorted(self): repo = _make_repo() repo.list_tags = Mock(return_value=["vlog", "热门", "教程"]) use_case = ListTagsUseCase(repo) result = use_case.execute("user-001") assert result == ["vlog", "热门", "教程"] repo.list_tags.assert_called_once_with("user-001") def test_list_tags_empty(self): repo = _make_repo() repo.list_tags = Mock(return_value=[]) use_case = ListTagsUseCase(repo) result = use_case.execute("user-001") assert result == [] # ── GetTemplateUsageUseCase ── class TestGetTemplateUsageUseCase: def test_get_usage_count(self): repo = _make_repo() repo.get_usage_count = Mock(return_value=3) use_case = GetTemplateUsageUseCase(repo) result = use_case.execute("tmpl-001") assert result == 3 repo.get_usage_count.assert_called_once_with("tmpl-001") def test_get_usage_zero(self): repo = _make_repo() repo.get_usage_count = Mock(return_value=0) use_case = GetTemplateUsageUseCase(repo) result = use_case.execute("tmpl-001") assert result == 0