diff --git a/tests/unit/test_config_base.py b/tests/unit/test_config_base.py index dee4a698e..082cbbed1 100755 --- a/tests/unit/test_config_base.py +++ b/tests/unit/test_config_base.py @@ -1,178 +1,181 @@ -"""Config Base 单元测试""" - +"""config/base.py 单测 — SharedSettings + 单例缓存管理.""" from __future__ import annotations +import os +from unittest.mock import patch + import pytest from packages.config.base import ( SharedSettings, + _settings_cache, get_cached_settings, get_shared_settings, reload_settings_cache, + _get_env_file, ) +# ── SharedSettings 基本配置 ──────────────────────────────────────────────── + + class TestSharedSettingsDefaults: - """SharedSettings 默认值测试""" - - @pytest.fixture(autouse=True) - def clean_env(self, monkeypatch): - """清除所有可能影响的环境变量,确保测的是代码默认值""" - env_vars = [ - "ENVIRONMENT", - "DEBUG", - "AUTO_CREATE_SCHEMA", - "DATABASE_URL", - "DATABASE_POOL_SIZE", - "DATABASE_MAX_OVERFLOW", - "DATABASE_POOL_TIMEOUT", - "DATABASE_POOL_RECYCLE", - "REDIS_URL", - "CELERY_BROKER_URL", - "CELERY_RESULT_BACKEND", - "OSS_ENDPOINT", - "OSS_ACCESS_KEY_ID", - "OSS_ACCESS_KEY_SECRET", - "OSS_BUCKET_NAME", - "OSS_DIRECT_UPLOAD_MAX_MB", - "OSS_DIRECT_UPLOAD_EXPIRE_SECONDS", - "COSYVOICE_API_KEY", - "COSYVOICE_BASE_URL", - "COSYVOICE_MODEL", - "COSYVOICE_VOICE", - "COSYVOICE_SAMPLE_RATE", - "COSYVOICE_FORMAT", - "COSYVOICE_CLONE_MODEL", - "DOUBAO_API_KEY", - "DOUBAO_MODEL", - "DOUBAO_BASE_URL", - "DOUBAO_TIMEOUT", - "DOUBAO_MAX_RETRIES", - ] - for var in env_vars: - monkeypatch.delenv(var, raising=False) - reload_settings_cache() - yield - reload_settings_cache() - - def _make_settings(self): - """构造不读 env 文件的纯净 settings""" - return SharedSettings(_env_file="/dev/null") + """SharedSettings 默认值验证.""" def test_default_environment(self): - """默认环境为 development""" - s = self._make_settings() - assert s.environment == "development" + settings = SharedSettings() + assert settings.environment == "development" def test_default_debug(self): - """默认开启 debug""" - s = self._make_settings() - assert s.debug is True + settings = SharedSettings() + assert settings.debug is True - def test_default_database_config(self): - """数据库默认配置""" - s = self._make_settings() - assert "postgresql" in s.database_url - assert s.database_pool_size == 20 - assert s.database_max_overflow == 10 - assert s.database_pool_timeout == 30 - assert s.database_pool_recycle == 3600 + def test_default_database_url(self): + settings = SharedSettings() + assert "postgresql" in settings.database_url - def test_default_redis_config(self): - """Redis 默认配置""" - s = self._make_settings() - assert s.redis_url.startswith("redis://") + def test_default_database_pool_size(self): + settings = SharedSettings() + assert settings.database_pool_size == 20 + assert settings.database_max_overflow == 10 + + def test_default_redis_url(self): + settings = SharedSettings() + assert settings.redis_url.startswith("redis://") def test_default_celery_config(self): - """Celery 默认配置""" - s = self._make_settings() - assert s.celery_broker_url.startswith("redis://") - assert s.celery_result_backend.startswith("redis://") + settings = SharedSettings() + assert "redis://" in settings.celery_broker_url + assert "redis://" in settings.celery_result_backend def test_default_oss_config(self): - """OSS 默认配置""" - s = self._make_settings() - assert s.oss_endpoint.endswith("aliyuncs.com") - assert s.oss_bucket_name == "xiaoxia-autocut" - assert s.oss_direct_upload_max_mb == 2000 - assert s.oss_direct_upload_expire_seconds == 900 + settings = SharedSettings() + assert settings.oss_bucket_name == "xiaoxia-autocut" + assert settings.oss_direct_upload_max_mb == 2000 + assert settings.oss_direct_upload_expire_seconds == 900 def test_default_cosyvoice_config(self): - """CosyVoice 默认配置""" - s = self._make_settings() - assert s.cosyvoice_model == "cosyvoice-v3-flash" - assert s.cosyvoice_sample_rate == 22050 - assert s.cosyvoice_format == "mp3" - assert s.cosyvoice_clone_model == "voice-enrollment" + settings = SharedSettings() + assert settings.cosyvoice_model == "cosyvoice-v3-flash" + assert settings.cosyvoice_sample_rate == 22050 + assert settings.cosyvoice_format == "mp3" def test_default_doubao_config(self): - """豆包默认配置""" - s = self._make_settings() - assert s.doubao_timeout == 30 - assert s.doubao_max_retries == 2 - assert "volces.com" in s.doubao_base_url + settings = SharedSettings() + assert settings.doubao_timeout == 30 + assert settings.doubao_max_retries == 2 - def test_default_empty_api_keys(self): - """API Key 默认空字符串""" - s = self._make_settings() - assert s.oss_access_key_id == "" - assert s.oss_access_key_secret == "" - assert s.cosyvoice_api_key == "" - assert s.doubao_api_key == "" + def test_env_override(self): + """环境变量可以覆盖默认值.""" + with patch.dict(os.environ, {"DEBUG": "false", "ENVIRONMENT": "production"}): + settings = SharedSettings() + assert settings.debug is False + assert settings.environment == "production" - def test_auto_create_schema_default(self): - """auto_create_schema 默认 False""" - s = self._make_settings() - assert s.auto_create_schema is False + def test_extra_env_ignored(self): + """model_config extra=ignore,未定义字段忽略.""" + with patch.dict(os.environ, {"RANDOM_UNKNOWN_VAR": "value"}): + # 不抛异常就是通过 + settings = SharedSettings() + assert not hasattr(settings, "random_unknown_var") -class TestSettingsSingleton: - """单例管理测试""" +# ── 单例缓存机制 ────────────────────────────────────────────────────────── + + +class TestSettingsCache: + """get_cached_settings / reload_settings_cache 单例机制.""" def setup_method(self): - """每个测试前清空缓存""" + """每个测试前清空缓存.""" reload_settings_cache() def teardown_method(self): - """每个测试后清空缓存""" reload_settings_cache() - def test_get_cached_settings_same_instance(self): - """同一类两次调用返回同一实例""" + def test_first_call_creates_instance(self): + """第一次调用创建实例并缓存.""" + settings = get_cached_settings(SharedSettings) + assert isinstance(settings, SharedSettings) + assert "SharedSettings" in _settings_cache + + def test_second_call_returns_same_instance(self): + """第二次调用返回同一实例(单例).""" s1 = get_cached_settings(SharedSettings) s2 = get_cached_settings(SharedSettings) assert s1 is s2 - def test_get_shared_settings_returns_shared_settings(self): - """get_shared_settings 返回 SharedSettings 实例""" - s = get_shared_settings() - assert isinstance(s, SharedSettings) - - def test_get_shared_settings_singleton(self): - """get_shared_settings 是单例""" - s1 = get_shared_settings() - s2 = get_shared_settings() - assert s1 is s2 - - def test_reload_settings_cache_clears(self): - """reload 后获取新实例""" + def test_reload_clears_cache(self): + """reload 后再次调用会创建新实例.""" s1 = get_cached_settings(SharedSettings) reload_settings_cache() s2 = get_cached_settings(SharedSettings) assert s1 is not s2 def test_custom_cache_key(self): - """自定义 cache_key 分开缓存""" - s1 = get_cached_settings(SharedSettings, cache_key="key_a") - s2 = get_cached_settings(SharedSettings, cache_key="key_b") + """支持自定义缓存 key.""" + s1 = get_cached_settings(SharedSettings, cache_key="custom_key") + assert "custom_key" in _settings_cache + assert "custom_key" not in ["SharedSettings"] or "SharedSettings" in _settings_cache + + # 不同 key 是不同实例 + s2 = get_cached_settings(SharedSettings, cache_key="another") assert s1 is not s2 - # 但值相同 - assert s1.database_url == s2.database_url - def test_different_classes_separate_cache(self): - """不同类使用不同缓存""" - from packages.config.api_settings import APISettings + def test_get_shared_settings_returns_singleton(self): + """get_shared_settings 是 SharedSettings 的便捷入口.""" + s1 = get_shared_settings() + s2 = get_shared_settings() + assert s1 is s2 + assert isinstance(s1, SharedSettings) - shared = get_shared_settings() - api = get_cached_settings(APISettings) - assert shared is not api + +# ── _get_env_file ────────────────────────────────────────────────────────── + + +class TestGetEnvFile: + """_get_env_file 环境文件选择逻辑.""" + + def test_default_development_uses_dot_env(self): + """默认 development 环境用 .env.""" + with patch.dict(os.environ, {}, clear=True): + # 没有 APP_ENV 时默认 development + result = _get_env_file() + assert result == ".env" + + def test_explicit_development_uses_dot_env(self): + """显式指定 development 也用 .env.""" + with patch.dict(os.environ, {"APP_ENV": "development"}): + result = _get_env_file() + assert result == ".env" + + def test_production_env_file(self, tmp_path): + """非 development 环境用 .env.{env},文件存在时返回它.""" + env_file = tmp_path / ".env.production" + env_file.write_text("DEBUG=false") + + with patch.dict(os.environ, {"APP_ENV": "production"}): + # 用 tmp_path 作为工作目录 + import os as _os + original_cwd = _os.getcwd() + _os.chdir(tmp_path) + try: + result = _get_env_file() + assert result == ".env.production" + finally: + _os.chdir(original_cwd) + + def test_env_file_not_found_falls_back_to_dot_env(self, tmp_path): + """环境文件不存在时回退到 .env.""" + env_file = tmp_path / ".env" + env_file.write_text("DEBUG=true") + + with patch.dict(os.environ, {"APP_ENV": "staging"}): + import os as _os + original_cwd = _os.getcwd() + _os.chdir(tmp_path) + try: + result = _get_env_file() + assert result == ".env" + finally: + _os.chdir(original_cwd) diff --git a/tests/unit/test_recipe_use_cases.py b/tests/unit/test_recipe_use_cases.py index 55720e95d..1f18ad945 100755 --- a/tests/unit/test_recipe_use_cases.py +++ b/tests/unit/test_recipe_use_cases.py @@ -1,22 +1,17 @@ -"""配方 Recipe UseCase 单元测试.""" - +"""Recipe Use Cases 单测 — 配方业务逻辑.""" from __future__ import annotations from unittest.mock import MagicMock, patch import pytest -from packages.application.recipe.commands import ( - CreateRecipeCommand, - RecipeItemCommand, - UpdateRecipeCommand, -) from packages.application.recipe.use_cases import ( CreateRecipeUseCase, DeleteRecipeUseCase, FeatureDisabledError, GetRecipeUseCase, ListRecipesUseCase, + MissingAssetWarning, UpdateRecipeUseCase, UseRecipeResult, UseRecipeUseCase, @@ -25,28 +20,35 @@ from packages.domain.exceptions import NotFoundError from packages.domain.recipe import Recipe, RecipeItem -def _make_recipe(id: str, name: str, user_id: str = "user_1", item_count: int = 0) -> Recipe: - items = [ - RecipeItem( - id=f"item_{i}", - recipe_id=id, - item_type="asset", - item_id=f"asset_{i}", - position=i, - ) - for i in range(item_count) - ] - return Recipe( - id=id, - user_id=user_id, - name=name, - description="测试配方", - template_id="tmpl_1", - generation_params={"resolution": "1080p"}, - items=items, - is_active=True, +# ── Fixtures / Helpers ───────────────────────────────────────────────────── + + +def make_recipe(**kwargs) -> Recipe: + defaults = dict( + id="recipe-1", + user_id="user-1", + name="Test Recipe", + description="Test description", + template_id="tpl-1", + generation_params={}, + items=[], metadata_={}, ) + defaults.update(kwargs) + return Recipe(**defaults) + + +def make_recipe_item(**kwargs) -> RecipeItem: + defaults = dict( + id="item-1", + recipe_id="recipe-1", + item_type="video", + item_id="asset-1", + position=0, + metadata_={}, + ) + defaults.update(kwargs) + return RecipeItem(**defaults) @pytest.fixture @@ -54,348 +56,303 @@ def mock_repo(): return MagicMock() -class TestListRecipesUseCase: - """ListRecipesUseCase 测试""" - - def test_list_returns_results(self, mock_repo): - """正常返回配方列表""" - recipe = _make_recipe("r1", "配方1") - mock_repo.list_by_user.return_value = [recipe] - use_case = ListRecipesUseCase(mock_repo) - - result = use_case.execute("user_1") - - assert len(result) == 1 - assert result[0].id == "r1" - mock_repo.list_by_user.assert_called_once_with("user_1", skip=0, limit=50) - - def test_list_with_pagination(self, mock_repo): - """带分页参数""" - mock_repo.list_by_user.return_value = [] - use_case = ListRecipesUseCase(mock_repo) - - use_case.execute("user_1", skip=5, limit=10) - - mock_repo.list_by_user.assert_called_once_with("user_1", skip=5, limit=10) - - def test_empty_list(self, mock_repo): - """空列表""" - mock_repo.list_by_user.return_value = [] - use_case = ListRecipesUseCase(mock_repo) - - result = use_case.execute("user_1") - - assert result == [] - - -class TestGetRecipeUseCase: - """GetRecipeUseCase 测试""" - - def test_get_existing(self, mock_repo): - """获取存在的配方""" - recipe = _make_recipe("r1", "配方1", item_count=3) - mock_repo.get.return_value = recipe - use_case = GetRecipeUseCase(mock_repo) - - result = use_case.execute("r1", "user_1") - - assert result is not None - assert result.id == "r1" - assert len(result.items) == 3 - mock_repo.get.assert_called_once_with("r1", "user_1") - - def test_get_nonexistent_returns_none(self, mock_repo): - """获取不存在的配方返回 None""" - mock_repo.get.return_value = None - use_case = GetRecipeUseCase(mock_repo) - - result = use_case.execute("noexist", "user_1") - - assert result is None +# ── CreateRecipeUseCase ──────────────────────────────────────────────────── class TestCreateRecipeUseCase: - """CreateRecipeUseCase 测试""" + """创建配方.""" def test_create_without_items(self, mock_repo): - """创建不带items的配方""" - mock_repo.create.side_effect = lambda x: x - use_case = CreateRecipeUseCase(mock_repo) + mock_repo.create.return_value = make_recipe() + uc = CreateRecipeUseCase(mock_repo) - command = CreateRecipeCommand( - user_id="user_1", - name="新配方", - description="测试", - template_id="tmpl_1", - generation_params={"key": "value"}, - items=[], + from packages.application.recipe.commands import CreateRecipeCommand + + cmd = CreateRecipeCommand( + user_id="user-1", + name="My Recipe", + description="My description", + template_id="tpl-1", + generation_params={"key": "val"}, ) - result = use_case.execute(command) + result = uc.execute(cmd) - assert isinstance(result, Recipe) - assert result.name == "新配方" - assert result.user_id == "user_1" - assert result.template_id == "tmpl_1" - assert result.generation_params == {"key": "value"} - assert result.items == [] + assert result.id == "recipe-1" mock_repo.create.assert_called_once() mock_repo.create_items.assert_not_called() def test_create_with_items(self, mock_repo): - """创建带items的配方""" - mock_repo.create.side_effect = lambda x: x - mock_repo.create_items.side_effect = lambda items: items - use_case = CreateRecipeUseCase(mock_repo) + recipe = make_recipe() + mock_repo.create.return_value = recipe + uc = CreateRecipeUseCase(mock_repo) - command = CreateRecipeCommand( - user_id="user_1", - name="带素材配方", + from packages.application.recipe.commands import ( + CreateRecipeCommand, + RecipeItemCommand, + ) + + cmd = CreateRecipeCommand( + user_id="user-1", + name="My Recipe", + description="", + template_id="tpl-1", items=[ - RecipeItemCommand(item_type="asset", item_id="a1", position=0), - RecipeItemCommand(item_type="title", item_id="t1", position=1), - RecipeItemCommand(item_type="voice", item_id="v1", position=2), + RecipeItemCommand(item_type="video", item_id="v1", position=0), + RecipeItemCommand(item_type="audio", item_id="a1", position=1), ], ) - result = use_case.execute(command) + result = uc.execute(cmd) - assert len(result.items) == 3 - assert result.items[0].item_type == "asset" - assert result.items[1].item_type == "title" - assert result.items[2].item_type == "voice" mock_repo.create.assert_called_once() mock_repo.create_items.assert_called_once() - created_items = mock_repo.create_items.call_args[0][0] - assert len(created_items) == 3 + items_arg = mock_repo.create_items.call_args[0][0] + assert len(items_arg) == 2 + assert items_arg[0].item_type == "video" + assert items_arg[1].item_type == "audio" + # items 被赋值到 recipe + assert len(result.items) == 2 - def test_create_with_default_values(self, mock_repo): - """使用默认值创建""" - mock_repo.create.side_effect = lambda x: x - use_case = CreateRecipeUseCase(mock_repo) - command = CreateRecipeCommand(user_id="user_1", name="极简配方") - result = use_case.execute(command) +# ── ListRecipesUseCase ───────────────────────────────────────────────────── - assert result.description == "" - assert result.template_id == "" - assert result.generation_params == {} - assert result.items == [] - assert result.metadata_ == {} + +class TestListRecipesUseCase: + """列表查询.""" + + def test_list_passes_params(self, mock_repo): + mock_repo.list_by_user.return_value = [make_recipe()] + uc = ListRecipesUseCase(mock_repo) + + result = uc.execute("user-1", skip=10, limit=20) + + mock_repo.list_by_user.assert_called_once_with( + "user-1", skip=10, limit=20 + ) + assert len(result) == 1 + + def test_list_default_params(self, mock_repo): + mock_repo.list_by_user.return_value = [] + uc = ListRecipesUseCase(mock_repo) + + uc.execute("user-1") + + mock_repo.list_by_user.assert_called_once_with( + "user-1", skip=0, limit=50 + ) + + +# ── GetRecipeUseCase ─────────────────────────────────────────────────────── + + +class TestGetRecipeUseCase: + """单个查询.""" + + def test_get_found(self, mock_repo): + mock_repo.get.return_value = make_recipe() + uc = GetRecipeUseCase(mock_repo) + + result = uc.execute("recipe-1", "user-1") + assert result.id == "recipe-1" + mock_repo.get.assert_called_once_with("recipe-1", "user-1") + + def test_get_not_found(self, mock_repo): + mock_repo.get.return_value = None + uc = GetRecipeUseCase(mock_repo) + + result = uc.execute("nonexistent", "user-1") + assert result is None + + +# ── UpdateRecipeUseCase ──────────────────────────────────────────────────── class TestUpdateRecipeUseCase: - """UpdateRecipeUseCase 测试""" + """更新配方.""" - def test_update_name(self, mock_repo): - """更新配方名称""" - recipe = _make_recipe("r1", "旧名称") - mock_repo.get.return_value = recipe - mock_repo.update.side_effect = lambda x: x - mock_repo.list_items.return_value = [] - use_case = UpdateRecipeUseCase(mock_repo) + def test_update_basic_fields(self, mock_repo): + existing = make_recipe(name="Old Name", description="Old desc") + mock_repo.get.return_value = existing + mock_repo.update.return_value = existing + uc = UpdateRecipeUseCase(mock_repo) - command = UpdateRecipeCommand(recipe_id="r1", user_id="user_1", name="新名称") - result = use_case.execute(command) + from packages.application.recipe.commands import UpdateRecipeCommand - assert result.name == "新名称" - # 其他不变 - assert result.description == "测试配方" - assert result.template_id == "tmpl_1" - mock_repo.get.assert_called_once_with("r1", "user_1") - mock_repo.update.assert_called_once() - # 没传items时从repository加载 - mock_repo.list_items.assert_called_once_with("r1") - - def test_update_multiple_fields(self, mock_repo): - """同时更新多个字段""" - recipe = _make_recipe("r1", "旧") - mock_repo.get.return_value = recipe - mock_repo.update.side_effect = lambda x: x - mock_repo.list_items.return_value = [] - use_case = UpdateRecipeUseCase(mock_repo) - - command = UpdateRecipeCommand( - recipe_id="r1", - user_id="user_1", - description="新描述", - template_id="tmpl_new", - generation_params={"new": "params"}, + cmd = UpdateRecipeCommand( + recipe_id="recipe-1", + user_id="user-1", + name="New Name", + description="New desc", ) - result = use_case.execute(command) + result = uc.execute(cmd) - assert result.description == "新描述" - assert result.template_id == "tmpl_new" - assert result.generation_params == {"new": "params"} + assert result.name == "New Name" + assert result.description == "New desc" + mock_repo.update.assert_called_once() - def test_update_items_replaces_old(self, mock_repo): - """更新items时删除旧的并创建新的""" - recipe = _make_recipe("r1", "配方", item_count=2) - mock_repo.get.return_value = recipe - mock_repo.update.side_effect = lambda x: x - mock_repo.create_items.side_effect = lambda items: items - use_case = UpdateRecipeUseCase(mock_repo) + def test_update_not_found_raises(self, mock_repo): + mock_repo.get.return_value = None + uc = UpdateRecipeUseCase(mock_repo) - command = UpdateRecipeCommand( - recipe_id="r1", - user_id="user_1", + from packages.application.recipe.commands import UpdateRecipeCommand + + cmd = UpdateRecipeCommand( + recipe_id="nonexistent", user_id="user-1", name="X" + ) + with pytest.raises(NotFoundError): + uc.execute(cmd) + mock_repo.update.assert_not_called() + + def test_update_template_and_params(self, mock_repo): + existing = make_recipe(template_id="old-tpl", generation_params={"a": 1}) + mock_repo.get.return_value = existing + mock_repo.update.return_value = existing + uc = UpdateRecipeUseCase(mock_repo) + + from packages.application.recipe.commands import UpdateRecipeCommand + + cmd = UpdateRecipeCommand( + recipe_id="recipe-1", + user_id="user-1", + template_id="new-tpl", + generation_params={"b": 2}, + ) + result = uc.execute(cmd) + + assert result.template_id == "new-tpl" + assert result.generation_params == {"b": 2} + + def test_update_replaces_items(self, mock_repo): + """提供 items 时,删除旧的并创建新的.""" + existing = make_recipe(items=[make_recipe_item(id="old-item")]) + mock_repo.get.return_value = existing + mock_repo.update.return_value = existing + uc = UpdateRecipeUseCase(mock_repo) + + from packages.application.recipe.commands import ( + UpdateRecipeCommand, + RecipeItemCommand, + ) + + cmd = UpdateRecipeCommand( + recipe_id="recipe-1", + user_id="user-1", items=[ - RecipeItemCommand(item_type="asset", item_id="new_a", position=0), - RecipeItemCommand(item_type="title", item_id="new_t", position=1), + RecipeItemCommand(item_type="video", item_id="v1", position=0), + RecipeItemCommand(item_type="audio", item_id="a1", position=1), ], ) - result = use_case.execute(command) + result = uc.execute(cmd) - mock_repo.delete_items_by_recipe.assert_called_once_with("r1") + mock_repo.delete_items_by_recipe.assert_called_once_with("recipe-1") mock_repo.create_items.assert_called_once() + items_arg = mock_repo.create_items.call_args[0][0] + assert len(items_arg) == 2 assert len(result.items) == 2 - assert result.items[0].item_id == "new_a" - def test_update_empty_items_list(self, mock_repo): - """更新为空items列表也会替换""" - recipe = _make_recipe("r1", "配方", item_count=3) - mock_repo.get.return_value = recipe - mock_repo.update.side_effect = lambda x: x - mock_repo.create_items.return_value = [] - use_case = UpdateRecipeUseCase(mock_repo) + def test_update_without_items_reloads_from_repo(self, mock_repo): + """不提供 items 时,从 repository 加载.""" + existing = make_recipe() + mock_repo.get.return_value = existing + mock_repo.update.return_value = existing + mock_repo.list_items.return_value = [make_recipe_item()] + uc = UpdateRecipeUseCase(mock_repo) - command = UpdateRecipeCommand(recipe_id="r1", user_id="user_1", items=[]) - result = use_case.execute(command) + from packages.application.recipe.commands import UpdateRecipeCommand - mock_repo.delete_items_by_recipe.assert_called_once() - mock_repo.create_items.assert_called_once_with([]) - assert result.items == [] + cmd = UpdateRecipeCommand( + recipe_id="recipe-1", user_id="user-1", name="New Name" + ) + result = uc.execute(cmd) - def test_update_nonexistent_raises(self, mock_repo): - """更新不存在的配方抛出 NotFoundError""" - mock_repo.get.return_value = None - use_case = UpdateRecipeUseCase(mock_repo) + mock_repo.list_items.assert_called_once_with("recipe-1") + assert len(result.items) == 1 + mock_repo.delete_items_by_recipe.assert_not_called() - command = UpdateRecipeCommand(recipe_id="noexist", user_id="user_1", name="新名称") - with pytest.raises(NotFoundError, match="not found"): - use_case.execute(command) - mock_repo.update.assert_not_called() +# ── DeleteRecipeUseCase ──────────────────────────────────────────────────── class TestDeleteRecipeUseCase: - """DeleteRecipeUseCase 测试""" + """删除配方.""" def test_delete_success(self, mock_repo): - """删除成功""" mock_repo.delete.return_value = True - use_case = DeleteRecipeUseCase(mock_repo) - - result = use_case.execute("r1", "user_1") + uc = DeleteRecipeUseCase(mock_repo) + result = uc.execute("recipe-1", "user-1") assert result is True - mock_repo.delete.assert_called_once_with("r1", "user_1") + mock_repo.delete.assert_called_once_with("recipe-1", "user-1") - def test_delete_nonexistent_returns_false(self, mock_repo): - """删除不存在的返回 False""" + def test_delete_not_found(self, mock_repo): mock_repo.delete.return_value = False - use_case = DeleteRecipeUseCase(mock_repo) - - result = use_case.execute("noexist", "user_1") + uc = DeleteRecipeUseCase(mock_repo) + result = uc.execute("nonexistent", "user-1") assert result is False +# ── UseRecipeUseCase ─────────────────────────────────────────────────────── + + class TestUseRecipeUseCase: - """UseRecipeUseCase 使用配方测试""" + """使用配方(feature flag + 校验).""" - def test_use_recipe_premium_enabled(self, mock_repo): - """premium用户可以使用配方""" - recipe = _make_recipe("r1", "配方1", item_count=2) - mock_repo.get.return_value = recipe - use_case = UseRecipeUseCase(mock_repo) + def test_feature_disabled_for_free_users(self, mock_repo): + """free 套餐没有配方复用功能.""" + uc = UseRecipeUseCase(mock_repo) - # 用 patch mock feature_flags - with patch("packages.application.recipe.use_cases.feature_flags") as mock_ff: - mock_ff.is_enabled.return_value = True - result = use_case.execute("r1", "user_1", user_plan="premium") - - assert isinstance(result, UseRecipeResult) - assert result.recipe.id == "r1" - assert isinstance(result.warnings, list) - mock_repo.get.assert_called_once_with("r1", "user_1") - - def test_use_recipe_basic_enabled(self, mock_repo): - """basic用户可以使用配方""" - recipe = _make_recipe("r1", "配方1") - mock_repo.get.return_value = recipe - use_case = UseRecipeUseCase(mock_repo) - - with patch("packages.application.recipe.use_cases.feature_flags") as mock_ff: - mock_ff.is_enabled.return_value = True - result = use_case.execute("r1", "user_1", user_plan="basic") - - assert result.recipe.id == "r1" - - def test_use_recipe_feature_disabled(self, mock_repo): - """功能未启用时抛出 FeatureDisabledError""" - use_case = UseRecipeUseCase(mock_repo) - - with patch("packages.application.recipe.use_cases.feature_flags") as mock_ff: - mock_ff.is_enabled.return_value = False - with pytest.raises(FeatureDisabledError, match="仅对基础版和高级版"): - use_case.execute("r1", "user_1", user_plan="free") + with pytest.raises(FeatureDisabledError) as exc_info: + uc.execute("recipe-1", "user-1", user_plan="free") + assert "基础版" in str(exc_info.value) or "仅对" in str(exc_info.value) mock_repo.get.assert_not_called() - def test_use_recipe_not_found(self, mock_repo): - """配方不存在时抛出 NotFoundError""" + def test_success_for_premium_users(self, mock_repo): + """premium 套餐可以使用.""" + recipe = make_recipe(items=[make_recipe_item()]) + mock_repo.get.return_value = recipe + uc = UseRecipeUseCase(mock_repo) + + result = uc.execute("recipe-1", "user-1", user_plan="premium") + + assert isinstance(result, UseRecipeResult) + assert result.recipe.id == "recipe-1" + assert isinstance(result.warnings, list) + + def test_success_for_basic_users(self, mock_repo): + """basic 套餐也可以使用.""" + recipe = make_recipe() + mock_repo.get.return_value = recipe + uc = UseRecipeUseCase(mock_repo) + + result = uc.execute("recipe-1", "user-1", user_plan="basic") + assert result.recipe.id == "recipe-1" + + def test_recipe_not_found(self, mock_repo): mock_repo.get.return_value = None - use_case = UseRecipeUseCase(mock_repo) + uc = UseRecipeUseCase(mock_repo) - with patch("packages.application.recipe.use_cases.feature_flags") as mock_ff: - mock_ff.is_enabled.return_value = True - with pytest.raises(NotFoundError, match="not found"): - use_case.execute("noexist", "user_1", user_plan="premium") + with pytest.raises(NotFoundError): + uc.execute("nonexistent", "user-1", user_plan="premium") + + def test_warnings_returns_list(self, mock_repo): + """返回的 warnings 是列表(即使为空).""" + recipe = make_recipe() + mock_repo.get.return_value = recipe + uc = UseRecipeUseCase(mock_repo) + + result = uc.execute("recipe-1", "user-1", user_plan="basic") + assert isinstance(result.warnings, list) -class TestRecipeCommands: - """命令数据类测试""" +# ── MissingAssetWarning ──────────────────────────────────────────────────── - def test_create_recipe_command_fields(self): - """CreateRecipeCommand 字段""" - cmd = CreateRecipeCommand( - user_id="u1", - name="测试", - description="desc", - template_id="t1", - generation_params={"a": 1}, - items=[RecipeItemCommand(item_type="asset", item_id="a1", position=0)], - metadata_={"key": "val"}, - ) - assert cmd.user_id == "u1" - assert cmd.name == "测试" - assert cmd.description == "desc" - assert cmd.template_id == "t1" - assert cmd.generation_params == {"a": 1} - assert len(cmd.items) == 1 - assert cmd.items[0].item_type == "asset" - assert cmd.metadata_ == {"key": "val"} - def test_recipe_item_command_defaults(self): - """RecipeItemCommand 默认值""" - cmd = RecipeItemCommand(item_type="asset", item_id="a1") - assert cmd.position == 0 - assert cmd.metadata_ == {} +class TestMissingAssetWarning: + """MissingAssetWarning 数据类.""" - def test_update_recipe_command_defaults_none(self): - """UpdateRecipeCommand 字段默认None""" - cmd = UpdateRecipeCommand(recipe_id="r1", user_id="u1") - assert cmd.name is None - assert cmd.description is None - assert cmd.template_id is None - assert cmd.generation_params is None - assert cmd.items is None - assert cmd.metadata_ is None - - def test_commands_are_dataclasses(self): - """都是 dataclass""" - from dataclasses import is_dataclass - - assert is_dataclass(CreateRecipeCommand) - assert is_dataclass(UpdateRecipeCommand) - assert is_dataclass(RecipeItemCommand) - assert is_dataclass(UseRecipeResult) + def test_creation(self): + w = MissingAssetWarning(item_type="video", item_id="v1", position=0) + assert w.item_type == "video" + assert w.item_id == "v1" + assert w.position == 0 diff --git a/tests/unit/test_title_library_use_cases.py b/tests/unit/test_title_library_use_cases.py index 5843f149c..62807b40b 100755 --- a/tests/unit/test_title_library_use_cases.py +++ b/tests/unit/test_title_library_use_cases.py @@ -1,17 +1,10 @@ -"""标题库 UseCase 单元测试.""" - +"""Title Library Use Cases 单测 — 标题库业务逻辑.""" from __future__ import annotations -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock import pytest -from packages.application.title_library.commands import ( - CreateTitleLibraryCommand, - IncrementTitleUsageCommand, - PickTitleCommand, - UpdateTitleLibraryCommand, -) from packages.application.title_library.use_cases import ( CreateTitleLibraryUseCase, DeleteTitleLibraryUseCase, @@ -25,19 +18,23 @@ from packages.domain.exceptions import NotFoundError, QuotaExceededError from packages.domain.title_library import TitleLibraryItem -def _make_item(id: str, name: str, text: str, usage_count: int = 0, category: str = "default") -> TitleLibraryItem: - return TitleLibraryItem( - id=id, - user_id="user_1", - name=name, - text=text, - category=category, +# ── Fixtures / Helpers ───────────────────────────────────────────────────── + + +def make_item(**kwargs) -> TitleLibraryItem: + defaults = dict( + id="title-1", + user_id="user-1", + name="Test Title", + text="This is a test title", + category="default", description="", tags=[], - usage_count=usage_count, + usage_count=0, is_active=True, - metadata_={}, ) + defaults.update(kwargs) + return TitleLibraryItem(**defaults) @pytest.fixture @@ -45,363 +42,414 @@ def mock_repo(): return MagicMock() -@pytest.fixture -def sample_item(): - return _make_item("title_1", "爆款标题", "这是一个爆款标题文案", usage_count=5) +# ── ListTitleLibraryUseCase ──────────────────────────────────────────────── class TestListTitleLibraryUseCase: - """ListTitleLibraryUseCase 测试""" + """列表查询.""" - def test_list_returns_results(self, mock_repo, sample_item): - """正常返回标题列表""" - mock_repo.list_by_user.return_value = [sample_item] - use_case = ListTitleLibraryUseCase(mock_repo) + def test_list_passes_params(self, mock_repo): + mock_repo.list_by_user.return_value = [make_item()] + uc = ListTitleLibraryUseCase(mock_repo) - result = use_case.execute("user_1") + result = uc.execute("user-1", category="vlog", skip=5, limit=10) + mock_repo.list_by_user.assert_called_once_with( + "user-1", category="vlog", skip=5, limit=10 + ) assert len(result) == 1 - assert result[0].id == "title_1" - mock_repo.list_by_user.assert_called_once_with("user_1", category=None, skip=0, limit=50) - def test_list_with_category(self, mock_repo, sample_item): - """按分类过滤""" - mock_repo.list_by_user.return_value = [sample_item] - use_case = ListTitleLibraryUseCase(mock_repo) - - use_case.execute("user_1", category="电商") - - mock_repo.list_by_user.assert_called_once_with("user_1", category="电商", skip=0, limit=50) - - def test_list_with_pagination(self, mock_repo, sample_item): - """带分页参数""" - mock_repo.list_by_user.return_value = [sample_item] - use_case = ListTitleLibraryUseCase(mock_repo) - - use_case.execute("user_1", skip=10, limit=20) - - mock_repo.list_by_user.assert_called_once_with("user_1", category=None, skip=10, limit=20) - - def test_empty_list(self, mock_repo): - """空列表""" + def test_list_default_params(self, mock_repo): mock_repo.list_by_user.return_value = [] - use_case = ListTitleLibraryUseCase(mock_repo) + uc = ListTitleLibraryUseCase(mock_repo) - result = use_case.execute("user_1") + uc.execute("user-1") - assert result == [] + mock_repo.list_by_user.assert_called_once_with( + "user-1", category=None, skip=0, limit=50 + ) + + +# ── GetTitleLibraryUseCase ───────────────────────────────────────────────── class TestGetTitleLibraryUseCase: - """GetTitleLibraryUseCase 测试""" + """单个查询.""" - def test_get_existing(self, mock_repo, sample_item): - """获取存在的标题""" - mock_repo.get.return_value = sample_item - use_case = GetTitleLibraryUseCase(mock_repo) + def test_get_found(self, mock_repo): + item = make_item() + mock_repo.get.return_value = item + uc = GetTitleLibraryUseCase(mock_repo) - result = use_case.execute("title_1", "user_1") + result = uc.execute("title-1", "user-1") - assert result is not None - assert result.id == "title_1" - mock_repo.get.assert_called_once_with("title_1", "user_1") + mock_repo.get.assert_called_once_with("title-1", "user-1") + assert result.id == "title-1" - def test_get_nonexistent_returns_none(self, mock_repo): - """获取不存在的标题返回 None""" + def test_get_not_found(self, mock_repo): mock_repo.get.return_value = None - use_case = GetTitleLibraryUseCase(mock_repo) - - result = use_case.execute("nonexistent", "user_1") + uc = GetTitleLibraryUseCase(mock_repo) + result = uc.execute("nonexistent", "user-1") assert result is None +# ── CreateTitleLibraryUseCase ────────────────────────────────────────────── + + class TestCreateTitleLibraryUseCase: - """CreateTitleLibraryUseCase 测试""" + """创建标题.""" - def test_create_success(self, mock_repo, sample_item): - """创建成功""" - mock_repo.count_by_user.return_value = 0 - mock_repo.create.return_value = sample_item - use_case = CreateTitleLibraryUseCase(mock_repo) + def test_create_success_within_quota(self, mock_repo): + mock_repo.count_by_user.return_value = 2 + mock_repo.create.return_value = make_item(id="new-id") + uc = CreateTitleLibraryUseCase(mock_repo) - command = CreateTitleLibraryCommand( - user_id="user_1", - name="新标题", - text="新标题文案", - category="default", - description="", - tags=[], - metadata_={}, + from packages.application.title_library.commands import CreateTitleLibraryCommand + + cmd = CreateTitleLibraryCommand( + user_id="user-1", + name="New Title", + text="New title text", + category="vlog", ) - result = use_case.execute(command, plan_name="free") + result = uc.execute(cmd, plan_name="free") - assert result.id == "title_1" - mock_repo.count_by_user.assert_called_once_with("user_1") + assert result.id == "new-id" + mock_repo.count_by_user.assert_called_once_with("user-1") mock_repo.create.assert_called_once() def test_create_quota_exceeded(self, mock_repo): - """超过配额时抛出 QuotaExceededError""" + """超过配额时抛 QuotaExceededError.""" + # free 计划 MAX_TITLES 假设很小,或者 count 很大 mock_repo.count_by_user.return_value = 9999 - use_case = CreateTitleLibraryUseCase(mock_repo) + uc = CreateTitleLibraryUseCase(mock_repo) - command = CreateTitleLibraryCommand( - user_id="user_1", - name="新标题", - text="文案", - category="default", - description="", - tags=[], - metadata_={}, + from packages.application.title_library.commands import CreateTitleLibraryCommand + + cmd = CreateTitleLibraryCommand( + user_id="user-1", name="Title", text="Text", category="default" ) with pytest.raises(QuotaExceededError): - use_case.execute(command, plan_name="free") + uc.execute(cmd, plan_name="free") mock_repo.create.assert_not_called() - def test_create_with_tags_and_metadata(self, mock_repo, sample_item): - """创建时带 tags 和 metadata_""" + def test_create_item_fields(self, mock_repo): + """创建时所有字段正确传递.""" mock_repo.count_by_user.return_value = 0 - mock_repo.create.return_value = sample_item - use_case = CreateTitleLibraryUseCase(mock_repo) + mock_repo.create.return_value = make_item() + uc = CreateTitleLibraryUseCase(mock_repo) - command = CreateTitleLibraryCommand( - user_id="user_1", - name="带标签标题", - text="文案", - category="电商", - description="测试描述", - tags=["爆款", "促销"], - metadata_={"source": "manual"}, + from packages.application.title_library.commands import CreateTitleLibraryCommand + + cmd = CreateTitleLibraryCommand( + user_id="user-1", + name="My Title", + text="Title text content", + category="food", + description="A food title", + tags=["t1", "t2"], + metadata_={"source": "import"}, ) - use_case.execute(command, plan_name="premium") + uc.execute(cmd, plan_name="premium") created = mock_repo.create.call_args[0][0] - assert isinstance(created, TitleLibraryItem) - assert created.name == "带标签标题" - assert created.category == "电商" - assert created.tags == ["爆款", "促销"] - assert created.metadata_ == {"source": "manual"} + assert created.name == "My Title" + assert created.text == "Title text content" + assert created.category == "food" + assert created.description == "A food title" + assert created.tags == ["t1", "t2"] + assert created.metadata_ == {"source": "import"} + assert created.user_id == "user-1" + + +# ── UpdateTitleLibraryUseCase ────────────────────────────────────────────── class TestUpdateTitleLibraryUseCase: - """UpdateTitleLibraryUseCase 测试""" + """更新标题.""" - def test_update_name(self, mock_repo, sample_item): - """更新标题名称""" - mock_repo.get.return_value = sample_item - mock_repo.update.side_effect = lambda x: x - use_case = UpdateTitleLibraryUseCase(mock_repo) + def test_update_success(self, mock_repo): + existing = make_item(name="Old Name", text="Old text") + mock_repo.get.return_value = existing + mock_repo.update.return_value = existing + uc = UpdateTitleLibraryUseCase(mock_repo) - command = UpdateTitleLibraryCommand(title_id="title_1", user_id="user_1", name="新名称") - result = use_case.execute(command) + from packages.application.title_library.commands import UpdateTitleLibraryCommand - assert result.name == "新名称" - # 其他字段不变 - assert result.text == "这是一个爆款标题文案" - mock_repo.get.assert_called_once_with("title_1", "user_1") + cmd = UpdateTitleLibraryCommand( + title_id="title-1", + user_id="user-1", + name="New Name", + text="New text", + ) + result = uc.execute(cmd) + + assert result.name == "New Name" + assert result.text == "New text" mock_repo.update.assert_called_once() - def test_update_multiple_fields(self, mock_repo, sample_item): - """同时更新多个字段""" - mock_repo.get.return_value = sample_item - mock_repo.update.side_effect = lambda x: x - use_case = UpdateTitleLibraryUseCase(mock_repo) + def test_update_not_found_raises(self, mock_repo): + mock_repo.get.return_value = None + uc = UpdateTitleLibraryUseCase(mock_repo) - command = UpdateTitleLibraryCommand( - title_id="title_1", - user_id="user_1", - text="新文案内容", - category="美食", + from packages.application.title_library.commands import UpdateTitleLibraryCommand + + cmd = UpdateTitleLibraryCommand( + title_id="nonexistent", user_id="user-1" + ) + with pytest.raises(NotFoundError): + uc.execute(cmd) + mock_repo.update.assert_not_called() + + def test_update_partial_fields(self, mock_repo): + """只更新传了的字段,其他保持不变.""" + existing = make_item( + name="Original", + text="Original text", + category="default", + tags=["old"], + is_active=True, + ) + mock_repo.get.return_value = existing + mock_repo.update.return_value = existing + uc = UpdateTitleLibraryUseCase(mock_repo) + + from packages.application.title_library.commands import UpdateTitleLibraryCommand + + # 只更新 name 和 is_active + cmd = UpdateTitleLibraryCommand( + title_id="title-1", + user_id="user-1", + name="New Name", is_active=False, ) - result = use_case.execute(command) + result = uc.execute(cmd) - assert result.text == "新文案内容" - assert result.category == "美食" + # 更新了的字段 + assert result.name == "New Name" assert result.is_active is False + # 没更新的保持原样 + assert result.text == "Original text" + assert result.category == "default" + assert result.tags == ["old"] - def test_update_nonexistent_raises(self, mock_repo): - """更新不存在的标题抛出 NotFoundError""" - mock_repo.get.return_value = None - use_case = UpdateTitleLibraryUseCase(mock_repo) + def test_update_metadata(self, mock_repo): + existing = make_item(metadata_={"old_key": "old_val"}) + mock_repo.get.return_value = existing + mock_repo.update.return_value = existing + uc = UpdateTitleLibraryUseCase(mock_repo) - command = UpdateTitleLibraryCommand(title_id="noexist", user_id="user_1", name="新名称") - with pytest.raises(NotFoundError, match="not found"): - use_case.execute(command) + from packages.application.title_library.commands import UpdateTitleLibraryCommand - mock_repo.update.assert_not_called() + cmd = UpdateTitleLibraryCommand( + title_id="title-1", + user_id="user-1", + metadata_={"new_key": "new_val"}, + ) + result = uc.execute(cmd) + assert result.metadata_ == {"new_key": "new_val"} + + +# ── DeleteTitleLibraryUseCase ────────────────────────────────────────────── class TestDeleteTitleLibraryUseCase: - """DeleteTitleLibraryUseCase 测试""" + """删除标题.""" def test_delete_success(self, mock_repo): - """删除成功""" mock_repo.delete.return_value = True - use_case = DeleteTitleLibraryUseCase(mock_repo) - - result = use_case.execute("title_1", "user_1") + uc = DeleteTitleLibraryUseCase(mock_repo) + result = uc.execute("title-1", "user-1") assert result is True - mock_repo.delete.assert_called_once_with("title_1", "user_1") + mock_repo.delete.assert_called_once_with("title-1", "user-1") - def test_delete_nonexistent_returns_false(self, mock_repo): - """删除不存在的返回 False""" + def test_delete_not_found(self, mock_repo): mock_repo.delete.return_value = False - use_case = DeleteTitleLibraryUseCase(mock_repo) - - result = use_case.execute("noexist", "user_1") + uc = DeleteTitleLibraryUseCase(mock_repo) + result = uc.execute("nonexistent", "user-1") assert result is False +# ── IncrementTitleUsageUseCase ───────────────────────────────────────────── + + class TestIncrementTitleUsageUseCase: - """IncrementTitleUsageUseCase 测试""" + """递增使用次数.""" def test_increment_positive(self, mock_repo): - """正增量时调用 repository""" mock_repo.increment_usage_count.return_value = True - use_case = IncrementTitleUsageUseCase(mock_repo) + uc = IncrementTitleUsageUseCase(mock_repo) - command = IncrementTitleUsageCommand(title_id="title_1", user_id="user_1", increment=1) - result = use_case.execute(command) + from packages.application.title_library.commands import IncrementTitleUsageCommand + cmd = IncrementTitleUsageCommand( + title_id="title-1", user_id="user-1", increment=1 + ) + result = uc.execute(cmd) assert result is True - mock_repo.increment_usage_count.assert_called_once_with("title_1", "user_1", increment=1) + mock_repo.increment_usage_count.assert_called_once_with( + "title-1", "user-1", increment=1 + ) def test_increment_zero_returns_false(self, mock_repo): - """增量为0返回False,不调用repository""" - use_case = IncrementTitleUsageUseCase(mock_repo) + """increment <= 0 直接返回 False,不调 repository.""" + uc = IncrementTitleUsageUseCase(mock_repo) - command = IncrementTitleUsageCommand(title_id="title_1", user_id="user_1", increment=0) - result = use_case.execute(command) + from packages.application.title_library.commands import IncrementTitleUsageCommand + cmd = IncrementTitleUsageCommand( + title_id="title-1", user_id="user-1", increment=0 + ) + result = uc.execute(cmd) assert result is False mock_repo.increment_usage_count.assert_not_called() def test_increment_negative_returns_false(self, mock_repo): - """负增量返回False""" - use_case = IncrementTitleUsageUseCase(mock_repo) + uc = IncrementTitleUsageUseCase(mock_repo) - command = IncrementTitleUsageCommand(title_id="title_1", user_id="user_1", increment=-1) - result = use_case.execute(command) + from packages.application.title_library.commands import IncrementTitleUsageCommand + cmd = IncrementTitleUsageCommand( + title_id="title-1", user_id="user-1", increment=-5 + ) + result = uc.execute(cmd) assert result is False mock_repo.increment_usage_count.assert_not_called() - def test_increment_large_number(self, mock_repo): - """大增量值""" + def test_increment_larger_number(self, mock_repo): mock_repo.increment_usage_count.return_value = True - use_case = IncrementTitleUsageUseCase(mock_repo) + uc = IncrementTitleUsageUseCase(mock_repo) - command = IncrementTitleUsageCommand(title_id="title_1", user_id="user_1", increment=10) - use_case.execute(command) + from packages.application.title_library.commands import IncrementTitleUsageCommand - mock_repo.increment_usage_count.assert_called_once_with("title_1", "user_1", increment=10) + cmd = IncrementTitleUsageCommand( + title_id="title-1", user_id="user-1", increment=10 + ) + result = uc.execute(cmd) + assert result is True + mock_repo.increment_usage_count.assert_called_once_with( + "title-1", "user-1", increment=10 + ) + + +# ── PickTitleUseCase ─────────────────────────────────────────────────────── class TestPickTitleUseCase: - """PickTitleUseCase 智能选标题测试""" + """智能选择标题.""" - def test_pick_from_multiple(self, mock_repo): - """从多个标题中选一个(最少使用的前5个中随机)""" - items = [_make_item(f"t{i}", f"标题{i}", f"文案{i}", usage_count=i) for i in range(10)] - mock_repo.list_by_user.return_value = items - use_case = PickTitleUseCase(mock_repo) - - command = PickTitleCommand(user_id="user_1") - result = use_case.execute(command) - - assert result is not None - assert isinstance(result, TitleLibraryItem) - # 选出的应该是使用次数最少的前5个之一(0-4) - assert result.usage_count <= 4 - mock_repo.list_by_user.assert_called_once() - - def test_pick_empty_returns_none(self, mock_repo): - """空标题库返回 None""" + def test_empty_list_returns_none(self, mock_repo): mock_repo.list_by_user.return_value = [] - use_case = PickTitleUseCase(mock_repo) + uc = PickTitleUseCase(mock_repo) - command = PickTitleCommand(user_id="user_1") - result = use_case.execute(command) + from packages.application.title_library.commands import PickTitleCommand + cmd = PickTitleCommand(user_id="user-1") + result = uc.execute(cmd) assert result is None - def test_pick_with_category(self, mock_repo): - """按分类选标题""" - items = [_make_item("t1", "标题1", "文案1", category="美食")] + def test_picks_from_available(self, mock_repo): + items = [make_item(id=f"t{i}", usage_count=i) for i in range(3)] mock_repo.list_by_user.return_value = items - use_case = PickTitleUseCase(mock_repo) + uc = PickTitleUseCase(mock_repo) - command = PickTitleCommand(user_id="user_1", category="美食") - result = use_case.execute(command) + from packages.application.title_library.commands import PickTitleCommand + cmd = PickTitleCommand(user_id="user-1") + result = uc.execute(cmd) + + # 结果应该是候选池中之一(最少使用的前5个) + assert result in items + assert result.usage_count <= 2 # 肯定是前3个里的 + + def test_exclude_ids(self, mock_repo): + """排除指定ID后从剩余中选.""" + items = [make_item(id=f"t{i}", usage_count=i) for i in range(10)] + mock_repo.list_by_user.return_value = items + uc = PickTitleUseCase(mock_repo) + + from packages.application.title_library.commands import PickTitleCommand + + # 排除前5个 + cmd = PickTitleCommand( + user_id="user-1", exclude_ids=[f"t{i}" for i in range(5)] + ) + result = uc.execute(cmd) + + # 应该从后5个里选 + assert result.id in [f"t{i}" for i in range(5, 10)] + + def test_exclude_all_falls_back_to_all(self, mock_repo): + """排除全部时,回退到从全部里选.""" + items = [make_item(id=f"t{i}", usage_count=i) for i in range(3)] + mock_repo.list_by_user.return_value = items + uc = PickTitleUseCase(mock_repo) + + from packages.application.title_library.commands import PickTitleCommand + + cmd = PickTitleCommand( + user_id="user-1", exclude_ids=["t0", "t1", "t2"] + ) + result = uc.execute(cmd) + # 回退到全部,还是能选出一个 assert result is not None + assert result.id in ["t0", "t1", "t2"] + + def test_category_filter(self, mock_repo): + """按分类过滤.""" + items = [make_item(id="t1", category="vlog"), make_item(id="t2", category="food")] + mock_repo.list_by_user.return_value = items + uc = PickTitleUseCase(mock_repo) + + from packages.application.title_library.commands import PickTitleCommand + + cmd = PickTitleCommand(user_id="user-1", category="vlog") + uc.execute(cmd) + + # 传给 repo 的参数里带了 category + call_kwargs = mock_repo.list_by_user.call_args[1] + assert call_kwargs["category"] == "vlog" + + def test_only_active_titles(self, mock_repo): + """只从活跃标题中选.""" + items = [make_item(id="t1"), make_item(id="t2")] + mock_repo.list_by_user.return_value = items + uc = PickTitleUseCase(mock_repo) + + from packages.application.title_library.commands import PickTitleCommand + + cmd = PickTitleCommand(user_id="user-1") + uc.execute(cmd) + call_kwargs = mock_repo.list_by_user.call_args[1] - assert call_kwargs["category"] == "美食" assert call_kwargs["is_active"] is True - def test_pick_exclude_ids(self, mock_repo): - """排除指定ID""" - items = [ - _make_item("t1", "标题1", "文案1", usage_count=1), - _make_item("t2", "标题2", "文案2", usage_count=2), - _make_item("t3", "标题3", "文案3", usage_count=3), - ] + def test_candidate_pool_size(self, mock_repo): + """候选池大小限制为5个最少使用的.""" + items = [make_item(id=f"t{i}", usage_count=10 - i) for i in range(20)] mock_repo.list_by_user.return_value = items - use_case = PickTitleUseCase(mock_repo) + uc = PickTitleUseCase(mock_repo) - command = PickTitleCommand(user_id="user_1", exclude_ids=["t1", "t2"]) - result = use_case.execute(command) + from packages.application.title_library.commands import PickTitleCommand - # 排除两个后只剩t3 - assert result.id == "t3" + # 多次运行,确保选的都是使用次数最少的 + selected_ids = set() + for _ in range(50): + result = uc.execute(PickTitleCommand(user_id="user-1")) + selected_ids.add(result.id) - def test_pick_exclude_all_falls_back(self, mock_repo): - """排除全部时从所有标题中选""" - items = [ - _make_item("t1", "标题1", "文案1", usage_count=1), - _make_item("t2", "标题2", "文案2", usage_count=2), - ] - mock_repo.list_by_user.return_value = items - use_case = PickTitleUseCase(mock_repo) - - command = PickTitleCommand(user_id="user_1", exclude_ids=["t1", "t2"]) - result = use_case.execute(command) - - # 排除全部后fallback到全部,所以还是能选出一个 - assert result is not None - assert result.id in ("t1", "t2") - - def test_pick_single_item(self, mock_repo): - """只有一个标题时选它""" - item = _make_item("only", "唯一标题", "唯一文案", usage_count=10) - mock_repo.list_by_user.return_value = [item] - use_case = PickTitleUseCase(mock_repo) - - command = PickTitleCommand(user_id="user_1") - result = use_case.execute(command) - - assert result.id == "only" - - def test_pick_prefers_less_used(self, mock_repo): - """倾向于选择使用次数少的""" - items = [ - _make_item("t_used", "常用", "常用", usage_count=100), - _make_item("t_fresh", "新的", "新的", usage_count=0), - ] - mock_repo.list_by_user.return_value = items - use_case = PickTitleUseCase(mock_repo) - - # 跑多次,验证使用少的出现在候选池里 - results = set() - for _ in range(20): - command = PickTitleCommand(user_id="user_1") - r = use_case.execute(command) - if r: - results.add(r.id) - - # 两个都在候选池(少于5个),所以都可能被选中 - assert "t_used" in results or "t_fresh" in results + # 选出来的都应该是 usage_count 最高的那5个(因为升序取前5, + # items[0]有usage_count=10,items[1]有9...items[4]有6, + # 都是"使用次数最少"的前5个) + # 不对,items[i] 的 usage_count = 10 - i + # items[0]=10, items[1]=9, ... items[9]=1, items[10]=0, items[11]=-1... + # 升序排列的话,items[19]=-9 最小,items[18]=-8 次之 ... + # 前5个最小的是 items[19], items[18], items[17], items[16], items[15] + # 即id为 t19, t18, t17, t16, t15 + expected_pool = {f"t{i}" for i in range(15, 20)} + assert selected_ids.issubset(expected_pool) + assert len(selected_ids) > 0