diff --git a/tests/unit/test_classification_jobs.py b/tests/unit/test_classification_jobs.py new file mode 100755 index 000000000..edd3deef3 --- /dev/null +++ b/tests/unit/test_classification_jobs.py @@ -0,0 +1,96 @@ +"""AI分类任务 UseCase 单元测试.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest + +from packages.application.classification_jobs import ( + SubmitClassificationJobCommand, + SubmitClassificationJobUseCase, +) +from packages.domain import ClassificationJob + + +@pytest.fixture +def mock_repo(): + return MagicMock() + + +class TestSubmitClassificationJobUseCase: + """SubmitClassificationJobUseCase 测试""" + + def test_submit_job_success(self, mock_repo): + """正常提交分类任务""" + mock_repo.create.side_effect = lambda j: j + use_case = SubmitClassificationJobUseCase(mock_repo) + + command = SubmitClassificationJobCommand( + project_id="proj_001", + asset_id="asset_001", + ) + result = use_case.execute(command) + + assert isinstance(result, ClassificationJob) + assert result.project_id == "proj_001" + assert result.asset_id == "asset_001" + assert result.status == "pending" + assert result.confidence == 0.0 + assert result.error_message == "" + mock_repo.create.assert_called_once() + + def test_submit_job_generates_id(self, mock_repo): + """提交任务时生成 id""" + mock_repo.create.side_effect = lambda j: j + use_case = SubmitClassificationJobUseCase(mock_repo) + + command = SubmitClassificationJobCommand( + project_id="proj_001", + asset_id="asset_001", + ) + result = use_case.execute(command) + + assert result.id is not None + assert len(result.id) > 0 + + def test_submit_job_two_different_ids(self, mock_repo): + """两次提交生成不同的 id""" + mock_repo.create.side_effect = lambda j: j + use_case = SubmitClassificationJobUseCase(mock_repo) + + command = SubmitClassificationJobCommand( + project_id="proj_001", + asset_id="asset_001", + ) + r1 = use_case.execute(command) + r2 = use_case.execute(command) + + assert r1.id != r2.id + + def test_submit_job_initial_classification_empty(self, mock_repo): + """初始 classification 为空""" + mock_repo.create.side_effect = lambda j: j + use_case = SubmitClassificationJobUseCase(mock_repo) + + command = SubmitClassificationJobCommand( + project_id="proj_001", + asset_id="asset_001", + ) + result = use_case.execute(command) + + assert result.classification == "" + + def test_submit_job_returns_repo_result(self, mock_repo): + """返回 repository.create 的结果""" + expected = MagicMock(spec=ClassificationJob) + mock_repo.create.return_value = expected + + use_case = SubmitClassificationJobUseCase(mock_repo) + command = SubmitClassificationJobCommand( + project_id="proj_001", + asset_id="asset_001", + ) + result = use_case.execute(command) + + assert result is expected diff --git a/tests/unit/test_generation_tasks.py b/tests/unit/test_generation_tasks.py index bd74fd9e5..44339c401 100755 --- a/tests/unit/test_generation_tasks.py +++ b/tests/unit/test_generation_tasks.py @@ -1,13 +1,6 @@ -""" -生成任务应用层用例单元测试(第十九波) +"""生成任务 UseCase 单元测试.""" -覆盖: -- CreateGenerationTaskUseCase -- GetGenerationTaskUseCase -- ListUserTasksFilteredUseCase -- RetryGenerationTaskUseCase -- Command / Filter / Result 对象 -""" +from __future__ import annotations from unittest.mock import MagicMock @@ -22,7 +15,7 @@ from packages.application.generation_tasks import ( ListUserTasksFilteredUseCase, RetryGenerationTaskUseCase, ) -from packages.domain.generation_task import GenerationTask, GenerationTaskStatus +from packages.domain import GenerationTask @pytest.fixture @@ -30,291 +23,238 @@ def mock_repo(): return MagicMock() -def make_task(status=GenerationTaskStatus.PENDING, **kwargs): - task = GenerationTask( - id="task-1", - project_id="proj-1", - asset_library_id="lib-1", - strategy_id="strat-1", - template_id="tmpl-1", - asset_ids=["asset-1"], - title_ids=["title-1"], - voice_ids=["voice-1"], - created_by_user_id="user-1", - video_title="测试标题", - ) - if status != GenerationTaskStatus.PENDING: - object.__setattr__(task, "status", status) - # 应用额外 kwargs - for k, v in kwargs.items(): - object.__setattr__(task, k, v) +@pytest.fixture +def sample_task(): + task = MagicMock(spec=GenerationTask) + task.id = "task_001" + task.project_id = "proj_001" + task.status = "pending" return task -# ============================================================ -# CreateGenerationTaskUseCase -# ============================================================ - - class TestCreateGenerationTaskUseCase: - """CreateGenerationTaskUseCase 创建生成任务""" + """CreateGenerationTaskUseCase 测试""" - def test_create_success(self, mock_repo): - """正常创建任务""" + def test_create_task_success(self, mock_repo): + """正常创建生成任务""" mock_repo.create.side_effect = lambda t: t + use_case = CreateGenerationTaskUseCase(mock_repo) - cmd = CreateGenerationTaskCommand( - project_id="proj-1", - asset_library_id="lib-1", - strategy_id="strat-1", - voice_library_id="vlib-1", - template_id="tmpl-1", - asset_ids=["a1", "a2"], - title_ids=["t1"], - voice_ids=["v1"], - created_by_user_id="user-1", - source_edit_plan_id="plan-1", - asset_select_mode="auto", - batch_id="batch-1", - video_title="我的视频", + command = CreateGenerationTaskCommand( + project_id="proj_001", + template_id="tpl_001", + asset_library_id="lib_001", + voice_library_id="voice_lib_001", + created_by_user_id="user_001", + ) + result = use_case.execute(command) + + assert isinstance(result, GenerationTask) + assert result.project_id == "proj_001" + assert result.template_id == "tpl_001" + assert result.status == "pending" + assert result.progress == 0.0 + assert result.result_count == 0 + mock_repo.create.assert_called_once() + + def test_create_task_generates_id(self, mock_repo): + """创建任务时生成 id""" + mock_repo.create.side_effect = lambda t: t + use_case = CreateGenerationTaskUseCase(mock_repo) + + command = CreateGenerationTaskCommand(project_id="proj_001") + result = use_case.execute(command) + + assert result.id is not None + assert len(result.id) > 0 + + def test_create_task_with_asset_ids(self, mock_repo): + """创建带 asset_ids 的任务""" + mock_repo.create.side_effect = lambda t: t + use_case = CreateGenerationTaskUseCase(mock_repo) + + command = CreateGenerationTaskCommand( + project_id="proj_001", + asset_ids=["asset_1", "asset_2", "asset_3"], + title_ids=["title_1", "title_2"], + voice_ids=["voice_1"], + ) + result = use_case.execute(command) + + assert len(result.asset_ids) == 3 + assert len(result.title_ids) == 2 + assert len(result.voice_ids) == 1 + + def test_create_task_with_auto_retry(self, mock_repo): + """创建带自动重试配置的任务""" + mock_repo.create.side_effect = lambda t: t + use_case = CreateGenerationTaskUseCase(mock_repo) + + command = CreateGenerationTaskCommand( + project_id="proj_001", auto_retry_enabled=True, auto_retry_max=3, ) - uc = CreateGenerationTaskUseCase(mock_repo) - task = uc.execute(cmd) + result = use_case.execute(command) - assert task.project_id == "proj-1" - assert task.asset_library_id == "lib-1" - assert task.strategy_id == "strat-1" - assert task.voice_library_id == "vlib-1" - assert task.template_id == "tmpl-1" - assert task.asset_ids == ["a1", "a2"] - assert task.title_ids == ["t1"] - assert task.voice_ids == ["v1"] - assert task.created_by_user_id == "user-1" - assert task.source_edit_plan_id == "plan-1" - assert task.asset_select_mode == "auto" - assert task.batch_id == "batch-1" - assert task.video_title == "我的视频" - assert task.auto_retry_enabled is True - assert task.auto_retry_max == 3 - assert task.status == GenerationTaskStatus.PENDING - assert task.progress == 0.0 - assert task.result_count == 0 - mock_repo.create.assert_called_once() + assert result.auto_retry_enabled is True + assert result.auto_retry_max == 3 - def test_create_default_values(self, mock_repo): - """默认参数值""" + def test_create_task_with_bgm_config(self, mock_repo): + """创建带 BGM 配置的任务""" mock_repo.create.side_effect = lambda t: t + use_case = CreateGenerationTaskUseCase(mock_repo) - cmd = CreateGenerationTaskCommand( - project_id="proj-1", - asset_library_id="lib-1", + bgm = {"enabled": True, "volume": 0.5, "library_id": "bgm_lib"} + command = CreateGenerationTaskCommand( + project_id="proj_001", + bgm_config=bgm, + resolution="1080p", + video_title="测试视频", ) - uc = CreateGenerationTaskUseCase(mock_repo) - task = uc.execute(cmd) + result = use_case.execute(command) - assert task.asset_ids == [] - assert task.title_ids == [] - assert task.voice_ids == [] - assert task.created_by_user_id == "" - assert task.video_title == "" - assert task.auto_retry_enabled is False - assert task.auto_retry_max == 0 + assert result.bgm_config == bgm + assert result.resolution == "1080p" + assert result.video_title == "测试视频" - def test_create_id_is_generated(self, mock_repo): - """ID 会自动生成""" + def test_create_task_defaults(self, mock_repo): + """默认参数的任务""" mock_repo.create.side_effect = lambda t: t + use_case = CreateGenerationTaskUseCase(mock_repo) - cmd = CreateGenerationTaskCommand(project_id="proj-1", asset_library_id="lib-1") - uc = CreateGenerationTaskUseCase(mock_repo) - task = uc.execute(cmd) + command = CreateGenerationTaskCommand() + result = use_case.execute(command) - assert task.id - assert isinstance(task.id, str) - assert len(task.id) > 10 # uuid hex - - -# ============================================================ -# GetGenerationTaskUseCase -# ============================================================ + assert result.project_id == "" + assert result.asset_ids == [] + assert result.auto_retry_enabled is False + assert result.auto_retry_max == 0 class TestGetGenerationTaskUseCase: - """GetGenerationTaskUseCase 获取任务""" + """GetGenerationTaskUseCase 测试""" - def test_get_existing(self, mock_repo): - """获取存在的任务""" - task = make_task() - mock_repo.get.return_value = task + def test_get_task_success(self, mock_repo, sample_task): + """获取任务成功""" + mock_repo.get.return_value = sample_task - uc = GetGenerationTaskUseCase(mock_repo) - result = uc.execute("task-1") + use_case = GetGenerationTaskUseCase(mock_repo) + result = use_case.execute("task_001") - assert result is task - mock_repo.get.assert_called_once_with("task-1") + assert result is sample_task + mock_repo.get.assert_called_once_with("task_001") - def test_get_not_found(self, mock_repo): - """获取不存在的任务返回 None""" + def test_get_task_not_found(self, mock_repo): + """任务不存在返回 None""" mock_repo.get.return_value = None - uc = GetGenerationTaskUseCase(mock_repo) - result = uc.execute("nonexistent") + use_case = GetGenerationTaskUseCase(mock_repo) + result = use_case.execute("nonexistent") assert result is None -# ============================================================ -# ListUserTasksFilteredUseCase -# ============================================================ - - class TestListUserTasksFilteredUseCase: - """ListUserTasksFilteredUseCase 按用户筛选任务""" + """ListUserTasksFilteredUseCase 测试""" - def test_list_without_filters(self, mock_repo): - """无筛选条件查询""" - tasks = [make_task(), make_task()] - mock_repo.list_by_user_filtered.return_value = tasks - mock_repo.count_by_user_filtered.return_value = 2 + def test_list_without_filter(self, mock_repo, sample_task): + """不带筛选条件查询""" + mock_repo.list_by_user_filtered.return_value = [sample_task] + mock_repo.count_by_user_filtered.return_value = 1 - uc = ListUserTasksFilteredUseCase(mock_repo) - result = uc.execute("user-1") + use_case = ListUserTasksFilteredUseCase(mock_repo) + result = use_case.execute("user_001") assert isinstance(result, ListGenerationTasksResult) - assert len(result.items) == 2 - assert result.total == 2 - mock_repo.list_by_user_filtered.assert_called_once_with("user-1", status=None, limit=None, offset=0) - mock_repo.count_by_user_filtered.assert_called_once_with("user-1", status=None) + assert len(result.items) == 1 + assert result.total == 1 + mock_repo.list_by_user_filtered.assert_called_once_with( + "user_001", status=None, limit=None, offset=0 + ) def test_list_with_status_filter(self, mock_repo): """按状态筛选""" mock_repo.list_by_user_filtered.return_value = [] mock_repo.count_by_user_filtered.return_value = 0 - uc = ListUserTasksFilteredUseCase(mock_repo) - uc.execute("user-1", status="running") + use_case = ListUserTasksFilteredUseCase(mock_repo) + result = use_case.execute("user_001", status="completed") - mock_repo.list_by_user_filtered.assert_called_once_with("user-1", status="running", limit=None, offset=0) - mock_repo.count_by_user_filtered.assert_called_once_with("user-1", status="running") + assert result.total == 0 + mock_repo.list_by_user_filtered.assert_called_once_with( + "user_001", status="completed", limit=None, offset=0 + ) def test_list_with_pagination(self, mock_repo): - """分页查询""" + """带分页参数查询""" mock_repo.list_by_user_filtered.return_value = [] - mock_repo.count_by_user_filtered.return_value = 100 + mock_repo.count_by_user_filtered.return_value = 50 - uc = ListUserTasksFilteredUseCase(mock_repo) - result = uc.execute("user-1", limit=10, offset=20) + use_case = ListUserTasksFilteredUseCase(mock_repo) + result = use_case.execute("user_001", limit=10, offset=20) - assert result.total == 100 - mock_repo.list_by_user_filtered.assert_called_once_with("user-1", status=None, limit=10, offset=20) + assert result.total == 50 + mock_repo.list_by_user_filtered.assert_called_once_with( + "user_001", status=None, limit=10, offset=20 + ) - def test_list_empty_result(self, mock_repo): - """空结果""" + def test_list_with_all_params(self, mock_repo): + """带所有筛选和分页参数""" mock_repo.list_by_user_filtered.return_value = [] - mock_repo.count_by_user_filtered.return_value = 0 + mock_repo.count_by_user_filtered.return_value = 5 - uc = ListUserTasksFilteredUseCase(mock_repo) - result = uc.execute("user-1", status="failed") + use_case = ListUserTasksFilteredUseCase(mock_repo) + use_case.execute("user_001", status="failed", limit=20, offset=0) - assert result.items == [] - assert result.total == 0 - - -# ============================================================ -# RetryGenerationTaskUseCase -# ============================================================ + mock_repo.list_by_user_filtered.assert_called_once_with( + "user_001", status="failed", limit=20, offset=0 + ) + mock_repo.count_by_user_filtered.assert_called_once_with( + "user_001", status="failed" + ) class TestRetryGenerationTaskUseCase: - """RetryGenerationTaskUseCase 重试失败任务""" + """RetryGenerationTaskUseCase 测试""" - def test_retry_success(self, mock_repo): - """失败任务重试成功""" - task = make_task( - status=GenerationTaskStatus.FAILED, - error_message="网络超时", - retry_count=0, - ) + def test_retry_failed_task(self, mock_repo): + """重试失败的任务""" + task = MagicMock(spec=GenerationTask) + task.is_failed = True mock_repo.get.return_value = task - mock_repo.update.side_effect = lambda t: t + mock_repo.update.return_value = task - uc = RetryGenerationTaskUseCase(mock_repo) - result = uc.execute("task-1") + use_case = RetryGenerationTaskUseCase(mock_repo) + result = use_case.execute("task_001") - assert result.status == GenerationTaskStatus.PENDING - assert result.retry_count == 1 - assert result.error_message == "" - assert result.error_info == {} - assert result.progress == 0.0 - assert result.result_count == 0 - assert result.started_at is None - assert result.completed_at is None - mock_repo.update.assert_called_once() + task.mark_pending_from_failed.assert_called_once() + mock_repo.update.assert_called_once_with(task) + assert result is task def test_retry_not_found(self, mock_repo): - """任务不存在""" + """任务不存在抛出 ValueError""" mock_repo.get.return_value = None - uc = RetryGenerationTaskUseCase(mock_repo) + use_case = RetryGenerationTaskUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): - uc.execute("nonexistent") + use_case.execute("nonexistent") - def test_retry_not_failed(self, mock_repo): - """非失败状态不能重试""" - task = make_task(status=GenerationTaskStatus.RUNNING) + mock_repo.update.assert_not_called() + + def test_retry_non_failed_task(self, mock_repo): + """非失败状态的任务不能重试""" + task = MagicMock(spec=GenerationTask) + task.is_failed = False + task.status = MagicMock() + task.status.value = "running" mock_repo.get.return_value = task - uc = RetryGenerationTaskUseCase(mock_repo) - with pytest.raises(ValueError, match="只有失败状态"): - uc.execute("task-1") + use_case = RetryGenerationTaskUseCase(mock_repo) - def test_retry_pending_not_allowed(self, mock_repo): - """pending 状态不能重试""" - task = make_task(status=GenerationTaskStatus.PENDING) - mock_repo.get.return_value = task + with pytest.raises(ValueError, match="只有失败状态的任务才能重试"): + use_case.execute("task_001") - uc = RetryGenerationTaskUseCase(mock_repo) - with pytest.raises(ValueError, match="只有失败状态"): - uc.execute("task-1") - - def test_retry_preserves_id(self, mock_repo): - """重试复用同一个 task_id""" - task = make_task(status=GenerationTaskStatus.FAILED) - original_id = task.id - mock_repo.get.return_value = task - mock_repo.update.side_effect = lambda t: t - - uc = RetryGenerationTaskUseCase(mock_repo) - result = uc.execute("task-1") - - assert result.id == original_id - - -# ============================================================ -# Command / Filter / Result 对象 -# ============================================================ - - -class TestCommandAndDataObjects: - """命令对象和数据对象""" - - def test_create_command_defaults(self): - cmd = CreateGenerationTaskCommand() - assert cmd.project_id == "" - assert cmd.asset_library_id == "" - assert cmd.asset_ids == [] - assert cmd.title_ids == [] - assert cmd.voice_ids == [] - assert cmd.auto_retry_enabled is False - assert cmd.auto_retry_max == 0 - - def test_list_filter_defaults(self): - f = ListTasksFilter() - assert f.status is None - - def test_list_result(self): - task = make_task() - r = ListGenerationTasksResult(items=[task], total=1) - assert len(r.items) == 1 - assert r.total == 1 + mock_repo.update.assert_not_called() + task.mark_pending_from_failed.assert_not_called() diff --git a/tests/unit/test_ingest_jobs.py b/tests/unit/test_ingest_jobs.py new file mode 100755 index 000000000..72f95964c --- /dev/null +++ b/tests/unit/test_ingest_jobs.py @@ -0,0 +1,72 @@ +"""素材入库任务 UseCase 单元测试.""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest + +from packages.application.ingest_jobs import ( + SubmitIngestJobCommand, + SubmitIngestJobUseCase, +) +from packages.domain import IngestJob + + +@pytest.fixture +def mock_repo(): + return MagicMock() + + +class TestSubmitIngestJobUseCase: + """SubmitIngestJobUseCase 测试""" + + def test_submit_job_success(self, mock_repo): + """正常提交入库任务""" + mock_repo.create.side_effect = lambda j: j + use_case = SubmitIngestJobUseCase(mock_repo) + + command = SubmitIngestJobCommand( + project_id="proj_001", + library_id="lib_001", + storage_key="videos/test.mp4", + file_hash="abc123def", + ) + result = use_case.execute(command) + + assert isinstance(result, IngestJob) + assert result.project_id == "proj_001" + assert result.library_id == "lib_001" + assert result.storage_key == "videos/test.mp4" + assert result.file_hash == "abc123def" + mock_repo.create.assert_called_once() + + def test_submit_job_without_hash(self, mock_repo): + """不传 file_hash 时默认为空""" + mock_repo.create.side_effect = lambda j: j + use_case = SubmitIngestJobUseCase(mock_repo) + + command = SubmitIngestJobCommand( + project_id="proj_001", + library_id="lib_001", + storage_key="images/test.png", + ) + result = use_case.execute(command) + + assert result.file_hash == "" + mock_repo.create.assert_called_once() + + def test_submit_job_returns_repo_result(self, mock_repo): + """返回 repository.create 的结果""" + expected_job = MagicMock(spec=IngestJob) + mock_repo.create.return_value = expected_job + + use_case = SubmitIngestJobUseCase(mock_repo) + command = SubmitIngestJobCommand( + project_id="proj_001", + library_id="lib_001", + storage_key="test.mp4", + ) + result = use_case.execute(command) + + assert result is expected_job diff --git a/tests/unit/test_password_reset_use_case.py b/tests/unit/test_password_reset_use_case.py old mode 100644 new mode 100755 index 930281b30..db6c30230 --- a/tests/unit/test_password_reset_use_case.py +++ b/tests/unit/test_password_reset_use_case.py @@ -1,9 +1,9 @@ -""" -密码重置 Use Case 测试 -""" +"""密码重置 UseCase 单元测试.""" + +from __future__ import annotations from datetime import datetime, timedelta, timezone -from unittest.mock import Mock +from unittest.mock import MagicMock, patch import pytest @@ -16,196 +16,284 @@ from packages.application.auth.password_reset_use_case import ( from packages.domain.entities import User +@pytest.fixture +def mock_user_repo(): + return MagicMock() + + +@pytest.fixture +def mock_email_service(): + svc = MagicMock() + svc.send_password_reset_email.return_value = (True, None) + return svc + + +@pytest.fixture +def sample_user(): + user = User( + id="user_001", + email="user@example.com", + display_name="测试用户", + username="testuser", + password_hash="old_hash", + ) + user.password_reset_token = None + user.password_reset_expires_at = None + return user + + +class TestRequestPasswordResetRequest: + """RequestPasswordResetRequest 测试""" + + def test_email_lowercased_and_stripped(self): + """邮箱转小写并去空格""" + req = RequestPasswordResetRequest(" User@Example.COM ") + assert req.email == "user@example.com" + + def test_empty_email(self): + """空邮箱""" + req = RequestPasswordResetRequest("") + assert req.email == "" + + class TestRequestPasswordResetUseCase: - """请求密码重置测试""" + """RequestPasswordResetUseCase 测试""" - @pytest.fixture - def mock_user_repo(self): - repo = Mock() - repo.find_by_email = Mock(return_value=None) - repo.save = Mock() - return repo + def test_request_success(self, mock_user_repo, mock_email_service, sample_user): + """请求重置成功""" + mock_user_repo.find_by_email.return_value = sample_user + mock_user_repo.save.return_value = sample_user - @pytest.fixture - def use_case(self, mock_user_repo): - email_service = Mock() - email_service.send_password_reset_email.return_value = (True, None) - return RequestPasswordResetUseCase( - user_repository=mock_user_repo, - base_url="https://test.com", - token_expire_hours=1, - email_service=email_service, + use_case = RequestPasswordResetUseCase( + mock_user_repo, + base_url="https://example.com", + email_service=mock_email_service, ) + request = RequestPasswordResetRequest("user@example.com") + success, error = use_case.execute(request) - @pytest.fixture - def test_user(self): - return User( - id="user-123", - email="test@example.com", - username="testuser", - display_name="Test User", - password_hash="hash", + assert success is True + assert error is None + assert sample_user.password_reset_token is not None + assert len(sample_user.password_reset_token) > 0 + assert sample_user.password_reset_expires_at is not None + mock_user_repo.save.assert_called_once() + mock_email_service.send_password_reset_email.assert_called_once() + + def test_request_user_not_found_returns_success(self, mock_user_repo, mock_email_service): + """用户不存在也返回成功(安全考虑,不暴露用户存在性)""" + mock_user_repo.find_by_email.return_value = None + + use_case = RequestPasswordResetUseCase( + mock_user_repo, + base_url="https://example.com", + email_service=mock_email_service, ) + request = RequestPasswordResetRequest("nonexistent@example.com") + success, error = use_case.execute(request) - def test_request_reset_success(self, use_case, mock_user_repo, test_user): - """测试请求重置成功""" - mock_user_repo.find_by_email.return_value = test_user + assert success is True + assert error is None + mock_user_repo.save.assert_not_called() + mock_email_service.send_password_reset_email.assert_not_called() - request = RequestPasswordResetRequest(email="test@example.com") + def test_request_empty_email_returns_error(self, mock_user_repo, mock_email_service): + """空邮箱返回错误""" + use_case = RequestPasswordResetUseCase( + mock_user_repo, + base_url="https://example.com", + email_service=mock_email_service, + ) + request = RequestPasswordResetRequest("") + success, error = use_case.execute(request) + + assert success is False + assert "Email is required" in error + + def test_reset_token_expiry_set(self, mock_user_repo, mock_email_service, sample_user): + """重置令牌过期时间正确设置""" + mock_user_repo.find_by_email.return_value = sample_user + mock_user_repo.save.return_value = sample_user + + use_case = RequestPasswordResetUseCase( + mock_user_repo, + base_url="https://example.com", + token_expire_hours=2, + email_service=mock_email_service, + ) + request = RequestPasswordResetRequest("user@example.com") + use_case.execute(request) + + assert sample_user.password_reset_expires_at is not None + # 过期时间应该在约2小时后 + expected = datetime.now(timezone.utc) + timedelta(hours=2) + diff = abs((sample_user.password_reset_expires_at - expected).total_seconds()) + assert diff < 10 # 允许10秒误差 + + def test_email_contains_reset_url(self, mock_user_repo, mock_email_service, sample_user): + """重置邮件包含正确的重置链接""" + mock_user_repo.find_by_email.return_value = sample_user + mock_user_repo.save.return_value = sample_user + + use_case = RequestPasswordResetUseCase( + mock_user_repo, + base_url="https://app.example.com", + email_service=mock_email_service, + ) + request = RequestPasswordResetRequest("user@example.com") + use_case.execute(request) + + call_args = mock_email_service.send_password_reset_email.call_args + reset_url = call_args[1]["reset_url"] if "reset_url" in call_args[1] else call_args[0][2] + assert "https://app.example.com/reset-password?token=" in reset_url + + def test_email_failure_does_not_affect_result(self, mock_user_repo, mock_email_service, sample_user): + """邮件发送失败不影响返回结果(安全考虑)""" + mock_user_repo.find_by_email.return_value = sample_user + mock_user_repo.save.return_value = sample_user + mock_email_service.send_password_reset_email.return_value = (False, "SMTP error") + + use_case = RequestPasswordResetUseCase( + mock_user_repo, + base_url="https://example.com", + email_service=mock_email_service, + ) + request = RequestPasswordResetRequest("user@example.com") success, error = use_case.execute(request) assert success is True assert error is None - # 验证保存了用户 - mock_user_repo.save.assert_called_once() - saved_user = mock_user_repo.save.call_args[0][0] - assert saved_user.password_reset_token is not None - assert saved_user.password_reset_expires_at is not None + def test_different_tokens_each_time(self, mock_user_repo, mock_email_service, sample_user): + """每次请求生成不同的 token""" + mock_user_repo.find_by_email.return_value = sample_user + mock_user_repo.save.return_value = sample_user - # 验证发送了邮件 - use_case.email_service.send_password_reset_email.assert_called_once() + use_case = RequestPasswordResetUseCase( + mock_user_repo, + base_url="https://example.com", + email_service=mock_email_service, + ) + request = RequestPasswordResetRequest("user@example.com") - def test_request_reset_user_not_exists(self, use_case, mock_user_repo): - """测试用户不存在(仍返回成功,避免暴露)""" - mock_user_repo.find_by_email.return_value = None + use_case.execute(request) + token1 = sample_user.password_reset_token - request = RequestPasswordResetRequest(email="nonexistent@example.com") - success, error = use_case.execute(request) + use_case.execute(request) + token2 = sample_user.password_reset_token - assert success is True # 安全考虑,仍返回成功 - assert error is None + assert token1 != token2 - # 不发送邮件 - use_case.email_service.send_password_reset_email.assert_not_called() - def test_request_reset_missing_email(self, use_case): - """测试缺少邮箱""" - request = RequestPasswordResetRequest(email="") - success, error = use_case.execute(request) +class TestResetPasswordRequest: + """ResetPasswordRequest 测试""" - assert success is False - assert error == "Email is required" + def test_stores_token_and_password(self): + """正确存储 token 和新密码""" + req = ResetPasswordRequest(token="abc123", new_password="NewPass1!") + assert req.token == "abc123" + assert req.new_password == "NewPass1!" class TestResetPasswordUseCase: - """重置密码测试""" + """ResetPasswordUseCase 测试""" - @pytest.fixture - def mock_user_repo(self): - repo = Mock() - repo.find_by_password_reset_token = Mock(return_value=None) - repo.save = Mock() - return repo + def test_reset_success(self, mock_user_repo, sample_user): + """重置密码成功""" + sample_user.password_reset_token = "valid_token" + sample_user.password_reset_expires_at = datetime.now(timezone.utc) + timedelta(hours=1) + mock_user_repo.find_by_password_reset_token.return_value = sample_user + mock_user_repo.save.return_value = sample_user - @pytest.fixture - def use_case(self, mock_user_repo): - return ResetPasswordUseCase(user_repository=mock_user_repo) - - @pytest.fixture - def test_user(self): - return User( - id="user-123", - email="test@example.com", - username="testuser", - display_name="Test User", - password_hash="old-hash", - password_reset_token="valid-token", - password_reset_expires_at=datetime.now(timezone.utc) + timedelta(hours=1), - ) - - def test_reset_password_success(self, use_case, mock_user_repo, test_user): - """测试重置密码成功""" - mock_user_repo.find_by_password_reset_token.return_value = test_user - - request = ResetPasswordRequest( - token="valid-token", - new_password="NewSecurePass123", - ) + use_case = ResetPasswordUseCase(mock_user_repo) + request = ResetPasswordRequest(token="valid_token", new_password="NewSecurePass1!") success, error = use_case.execute(request) assert success is True assert error is None - - # 验证密码已更新 - assert test_user.password_hash != "old-hash" - assert test_user.password_reset_token is None - assert test_user.password_reset_expires_at is None - - # 验证保存了用户 + assert sample_user.password_reset_token is None + assert sample_user.password_reset_expires_at is None + assert sample_user.password_hash != "old_hash" mock_user_repo.save.assert_called_once() - def test_reset_password_success_with_naive_database_datetime(self, use_case, mock_user_repo, test_user): - """测试数据库返回 naive datetime 时仍可重置密码""" - test_user.password_reset_expires_at = (datetime.now(timezone.utc) + timedelta(hours=1)).replace(tzinfo=None) - mock_user_repo.find_by_password_reset_token.return_value = test_user - - success, error = use_case.execute(ResetPasswordRequest(token="valid-token", new_password="NewSecurePass123")) - - assert success is True - assert error is None - mock_user_repo.save.assert_called_once() - - def test_reset_password_weak_password(self, use_case, mock_user_repo, test_user): - """测试弱密码""" - mock_user_repo.find_by_password_reset_token.return_value = test_user - - request = ResetPasswordRequest( - token="valid-token", - new_password="weak", - ) + def test_reset_empty_token(self, mock_user_repo): + """空 token 返回错误""" + use_case = ResetPasswordUseCase(mock_user_repo) + request = ResetPasswordRequest(token="", new_password="NewPass1!") success, error = use_case.execute(request) assert success is False - assert "at least 8 characters" in error + assert "Reset token is required" in error + mock_user_repo.save.assert_not_called() - def test_reset_password_invalid_token(self, use_case, mock_user_repo): - """测试无效令牌""" + def test_reset_empty_password(self, mock_user_repo): + """空密码返回错误""" + use_case = ResetPasswordUseCase(mock_user_repo) + request = ResetPasswordRequest(token="sometoken", new_password="") + success, error = use_case.execute(request) + + assert success is False + assert "New password is required" in error + mock_user_repo.save.assert_not_called() + + def test_reset_weak_password(self, mock_user_repo): + """弱密码返回错误""" + use_case = ResetPasswordUseCase(mock_user_repo) + request = ResetPasswordRequest(token="sometoken", new_password="weak") + success, error = use_case.execute(request) + + assert success is False + assert error is not None + mock_user_repo.save.assert_not_called() + + def test_reset_invalid_token(self, mock_user_repo): + """无效 token 返回错误""" mock_user_repo.find_by_password_reset_token.return_value = None - request = ResetPasswordRequest( - token="invalid-token", - new_password="NewSecurePass123", - ) + use_case = ResetPasswordUseCase(mock_user_repo) + request = ResetPasswordRequest(token="invalid_token", new_password="NewPass1!") success, error = use_case.execute(request) assert success is False - assert error == "Invalid or expired reset token" + assert "Invalid or expired" in error + mock_user_repo.save.assert_not_called() - def test_reset_password_expired_token(self, use_case, mock_user_repo, test_user): - """测试过期令牌""" - test_user.password_reset_expires_at = datetime.now(timezone.utc) - timedelta(hours=1) - mock_user_repo.find_by_password_reset_token.return_value = test_user + def test_reset_expired_token(self, mock_user_repo, sample_user): + """过期 token 返回错误""" + sample_user.password_reset_token = "expired_token" + sample_user.password_reset_expires_at = datetime.now(timezone.utc) - timedelta(hours=1) + mock_user_repo.find_by_password_reset_token.return_value = sample_user - request = ResetPasswordRequest( - token="valid-token", - new_password="NewSecurePass123", - ) + use_case = ResetPasswordUseCase(mock_user_repo) + request = ResetPasswordRequest(token="expired_token", new_password="NewPass1!") success, error = use_case.execute(request) assert success is False - assert error == "Reset token has expired" + assert "expired" in error.lower() + mock_user_repo.save.assert_not_called() - def test_reset_password_missing_token(self, use_case): - """测试缺少令牌""" - request = ResetPasswordRequest( - token="", - new_password="NewSecurePass123", - ) + def test_reset_naive_datetime_treated_as_utc(self, mock_user_repo, sample_user): + """无时区的过期时间按 UTC 处理""" + sample_user.password_reset_token = "naive_token" + # 用无时区的时间,设置为过去 + sample_user.password_reset_expires_at = datetime.utcnow() - timedelta(hours=1) + mock_user_repo.find_by_password_reset_token.return_value = sample_user + + use_case = ResetPasswordUseCase(mock_user_repo) + request = ResetPasswordRequest(token="naive_token", new_password="NewPass1!") success, error = use_case.execute(request) assert success is False - assert error == "Reset token is required" + assert "expired" in error.lower() - def test_reset_password_missing_password(self, use_case, mock_user_repo, test_user): - """测试缺少新密码""" - mock_user_repo.find_by_password_reset_token.return_value = test_user + def test_reset_no_expiry_set(self, mock_user_repo, sample_user): + """没有设置过期时间的 token 可以使用""" + sample_user.password_reset_token = "no_expiry_token" + sample_user.password_reset_expires_at = None + mock_user_repo.find_by_password_reset_token.return_value = sample_user - request = ResetPasswordRequest( - token="valid-token", - new_password="", - ) + use_case = ResetPasswordUseCase(mock_user_repo) + request = ResetPasswordRequest(token="no_expiry_token", new_password="NewPass1!") success, error = use_case.execute(request) - assert success is False - assert error == "New password is required" + assert success is True