diff --git a/tests/unit/test_job_use_cases.py b/tests/unit/test_job_use_cases.py new file mode 100755 index 000000000..bbe51e7f4 --- /dev/null +++ b/tests/unit/test_job_use_cases.py @@ -0,0 +1,318 @@ +"""Job Use Cases 单元测试""" + +from __future__ import annotations + +from unittest.mock import MagicMock + +import pytest + +from packages.application.jobs import ( + CancelJobUseCase, + CompleteJobCommand, + CompleteJobUseCase, + CreateJobCommand, + CreateJobUseCase, + FailJobCommand, + FailJobUseCase, + GetJobStatisticsUseCase, + GetJobUseCase, + ListJobsUseCase, + RetryJobUseCase, + SubmitJobUseCase, + UpdateJobProgressCommand, + UpdateJobProgressUseCase, +) +from packages.domain.job import Job, JobStatus, JobType + + +@pytest.fixture +def mock_repo(): + return MagicMock() + + +@pytest.fixture +def sample_job(): + return Job.create( + project_id="proj_001", + job_type=JobType.VIDEO_COMPOSE, + payload={"template_id": "tpl_001"}, + source_id="src_001", + created_by_user_id="user_001", + max_retries=3, + ) + + +class TestCreateJobCommand: + """CreateJobCommand 测试""" + + def test_default_values(self): + cmd = CreateJobCommand(project_id="p1", job_type=JobType.VIDEO_COMPOSE) + assert cmd.project_id == "p1" + assert cmd.payload == {} + assert cmd.source_id == "" + assert cmd.created_by_user_id == "" + assert cmd.max_retries == 3 + + +class TestCreateJobUseCase: + """CreateJobUseCase 测试""" + + def test_create_success(self, mock_repo, sample_job): + mock_repo.create.return_value = sample_job + use_case = CreateJobUseCase(mock_repo) + cmd = CreateJobCommand( + project_id="proj_001", + job_type=JobType.VIDEO_COMPOSE, + payload={"template_id": "tpl_001"}, + source_id="src_001", + created_by_user_id="user_001", + max_retries=3, + ) + result = use_case.execute(cmd) + assert result.status == JobStatus.PENDING + assert result.project_id == "proj_001" + mock_repo.create.assert_called_once() + + def test_create_with_string_job_type(self, mock_repo): + mock_repo.create.side_effect = lambda x: x + use_case = CreateJobUseCase(mock_repo) + cmd = CreateJobCommand(project_id="p1", job_type="video_compose") + result = use_case.execute(cmd) + assert result.job_type == JobType.VIDEO_COMPOSE + + +class TestSubmitJobUseCase: + """SubmitJobUseCase 测试""" + + def test_submit_success(self, mock_repo, sample_job): + mock_repo.get.return_value = sample_job + mock_repo.update.side_effect = lambda x: x + use_case = SubmitJobUseCase(mock_repo) + result = use_case.execute(sample_job.id, celery_task_id="celery_123") + assert result.status == JobStatus.RUNNING + assert result.celery_task_id == "celery_123" + + def test_submit_not_found(self, mock_repo): + mock_repo.get.return_value = None + use_case = SubmitJobUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + use_case.execute("nonexistent") + + def test_submit_wrong_status(self, mock_repo, sample_job): + sample_job.status = JobStatus.RUNNING + mock_repo.get.return_value = sample_job + use_case = SubmitJobUseCase(mock_repo) + with pytest.raises(ValueError, match="只有 pending"): + use_case.execute(sample_job.id) + + +class TestUpdateJobProgressUseCase: + """UpdateJobProgressUseCase 测试""" + + def test_update_progress_success(self, mock_repo, sample_job): + sample_job.status = JobStatus.RUNNING + mock_repo.get.return_value = sample_job + mock_repo.update.side_effect = lambda x: x + use_case = UpdateJobProgressUseCase(mock_repo) + cmd = UpdateJobProgressCommand( + job_id=sample_job.id, progress=50.0, current_stage="渲染中" + ) + result = use_case.execute(cmd) + assert result.progress == 50.0 + assert result.current_stage == "渲染中" + + def test_update_progress_not_found(self, mock_repo): + mock_repo.get.return_value = None + use_case = UpdateJobProgressUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + use_case.execute(UpdateJobProgressCommand(job_id="x", progress=10)) + + def test_update_progress_wrong_status(self, mock_repo, sample_job): + sample_job.status = JobStatus.PENDING + mock_repo.get.return_value = sample_job + use_case = UpdateJobProgressUseCase(mock_repo) + with pytest.raises(ValueError, match="只有 running"): + use_case.execute(UpdateJobProgressCommand(job_id=sample_job.id, progress=10)) + + +class TestCompleteJobUseCase: + """CompleteJobUseCase 测试""" + + def test_complete_from_running(self, mock_repo, sample_job): + sample_job.status = JobStatus.RUNNING + mock_repo.get.return_value = sample_job + mock_repo.update.side_effect = lambda x: x + use_case = CompleteJobUseCase(mock_repo) + cmd = CompleteJobCommand(job_id=sample_job.id, result={"url": "http://..."}) + result = use_case.execute(cmd) + assert result.status == JobStatus.SUCCESS + assert result.result["url"] == "http://..." + + def test_complete_from_pending(self, mock_repo, sample_job): + sample_job.status = JobStatus.PENDING + mock_repo.get.return_value = sample_job + mock_repo.update.side_effect = lambda x: x + use_case = CompleteJobUseCase(mock_repo) + result = use_case.execute(CompleteJobCommand(job_id=sample_job.id)) + assert result.status == JobStatus.SUCCESS + + def test_complete_not_found(self, mock_repo): + mock_repo.get.return_value = None + use_case = CompleteJobUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + use_case.execute(CompleteJobCommand(job_id="x")) + + def test_complete_failed_status_raises(self, mock_repo, sample_job): + sample_job.status = JobStatus.FAILED + mock_repo.get.return_value = sample_job + use_case = CompleteJobUseCase(mock_repo) + with pytest.raises(ValueError, match="只有 running/pending"): + use_case.execute(CompleteJobCommand(job_id=sample_job.id)) + + +class TestFailJobUseCase: + """FailJobUseCase 测试""" + + def test_fail_success(self, mock_repo, sample_job): + sample_job.status = JobStatus.RUNNING + mock_repo.get.return_value = sample_job + mock_repo.update.side_effect = lambda x: x + use_case = FailJobUseCase(mock_repo) + cmd = FailJobCommand(job_id=sample_job.id, error_message="渲染失败") + result = use_case.execute(cmd) + assert result.status == JobStatus.FAILED + assert "渲染失败" in result.error_message + + def test_fail_not_found(self, mock_repo): + mock_repo.get.return_value = None + use_case = FailJobUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + use_case.execute(FailJobCommand(job_id="x", error_message="err")) + + def test_fail_updates_error_message(self, mock_repo, sample_job): + sample_job.status = JobStatus.RUNNING + mock_repo.get.return_value = sample_job + mock_repo.update.side_effect = lambda x: x + use_case = FailJobUseCase(mock_repo) + result = use_case.execute(FailJobCommand(job_id=sample_job.id, error_message="连接超时")) + assert result.error_message == "连接超时" + + +class TestRetryJobUseCase: + """RetryJobUseCase 测试""" + + def test_retry_success(self, mock_repo, sample_job): + sample_job.status = JobStatus.FAILED + sample_job.retry_count = 1 + mock_repo.get.return_value = sample_job + mock_repo.update.side_effect = lambda x: x + use_case = RetryJobUseCase(mock_repo) + result = use_case.execute(sample_job.id) + assert result.status == JobStatus.PENDING + assert result.retry_count == 2 + + def test_retry_not_found(self, mock_repo): + mock_repo.get.return_value = None + use_case = RetryJobUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + use_case.execute("nonexistent") + + +class TestCancelJobUseCase: + """CancelJobUseCase 测试""" + + def test_cancel_pending(self, mock_repo, sample_job): + mock_repo.get.return_value = sample_job + mock_repo.update.side_effect = lambda x: x + use_case = CancelJobUseCase(mock_repo) + result = use_case.execute(sample_job.id) + assert result.status == JobStatus.CANCELLED + + def test_cancel_running(self, mock_repo, sample_job): + sample_job.status = JobStatus.RUNNING + mock_repo.get.return_value = sample_job + mock_repo.update.side_effect = lambda x: x + use_case = CancelJobUseCase(mock_repo) + result = use_case.execute(sample_job.id) + assert result.status == JobStatus.CANCELLED + + def test_cancel_terminal_raises(self, mock_repo, sample_job): + sample_job.status = JobStatus.SUCCESS + mock_repo.get.return_value = sample_job + use_case = CancelJobUseCase(mock_repo) + with pytest.raises(ValueError, match="终态"): + use_case.execute(sample_job.id) + + def test_cancel_not_found(self, mock_repo): + mock_repo.get.return_value = None + use_case = CancelJobUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + use_case.execute("nonexistent") + + +class TestGetJobUseCase: + """GetJobUseCase 测试""" + + def test_get_exists(self, mock_repo, sample_job): + mock_repo.get.return_value = sample_job + use_case = GetJobUseCase(mock_repo) + result = use_case.execute(sample_job.id) + assert result.id == sample_job.id + + def test_get_not_found(self, mock_repo): + mock_repo.get.return_value = None + use_case = GetJobUseCase(mock_repo) + result = use_case.execute("nonexistent") + assert result is None + + +class TestListJobsUseCase: + """ListJobsUseCase 测试""" + + def test_list_by_project(self, mock_repo, sample_job): + mock_repo.list_by_project.return_value = [sample_job] + use_case = ListJobsUseCase(mock_repo) + result = use_case.execute(project_id="proj_001") + assert len(result) == 1 + mock_repo.list_by_project.assert_called_once() + + def test_list_by_user(self, mock_repo, sample_job): + mock_repo.list_by_user.return_value = [sample_job] + use_case = ListJobsUseCase(mock_repo) + result = use_case.execute(user_id="user_001") + assert len(result) == 1 + mock_repo.list_by_user.assert_called_once() + + def test_list_no_filter_raises(self, mock_repo): + use_case = ListJobsUseCase(mock_repo) + with pytest.raises(ValueError, match="必须指定"): + use_case.execute() + + def test_list_with_filters(self, mock_repo, sample_job): + mock_repo.list_by_project.return_value = [sample_job] + use_case = ListJobsUseCase(mock_repo) + use_case.execute( + project_id="p1", + job_type=JobType.VIDEO_COMPOSE, + status=JobStatus.RUNNING, + limit=20, + offset=10, + ) + mock_repo.list_by_project.assert_called_once_with( + "p1", job_type=JobType.VIDEO_COMPOSE, status=JobStatus.RUNNING, limit=20, offset=10 + ) + + +class TestGetJobStatisticsUseCase: + """GetJobStatisticsUseCase 测试""" + + def test_statistics(self, mock_repo): + mock_repo.count_by_project.side_effect = [10, 2, 3, 4, 1] + use_case = GetJobStatisticsUseCase(mock_repo) + stats = use_case.execute("proj_001") + assert stats["project_id"] == "proj_001" + assert stats["total"] == 10 + assert stats["pending"] == 2 + assert stats["running"] == 3 + assert stats["success"] == 4 + assert stats["failed"] == 1 diff --git a/tests/unit/test_schema_guard.py b/tests/unit/test_schema_guard.py old mode 100644 new mode 100755 index 12545ac14..5bed36fac --- a/tests/unit/test_schema_guard.py +++ b/tests/unit/test_schema_guard.py @@ -1,19 +1,84 @@ +"""Schema Guard 单元测试""" + +from __future__ import annotations + import pytest -from packages.adapters.sqlalchemy_impl.schema_guard import assert_auto_create_schema_allowed +from packages.adapters.sqlalchemy_impl.schema_guard import ( + BLOCKED_AUTO_CREATE_ENVIRONMENTS, + assert_auto_create_schema_allowed, + normalize_environment, +) -@pytest.mark.parametrize("environment", ["staging", "production", " STAGING ", "Production"]) -def test_auto_create_schema_is_forbidden_in_deployed_environments(environment): - with pytest.raises(RuntimeError, match="AUTO_CREATE_SCHEMA is forbidden"): - assert_auto_create_schema_allowed(environment, enabled=True) +class TestNormalizeEnvironment: + """normalize_environment 测试""" + + def test_development(self): + assert normalize_environment("development") == "development" + + def test_staging(self): + assert normalize_environment("staging") == "staging" + + def test_production(self): + assert normalize_environment("production") == "production" + + def test_none_returns_development(self): + assert normalize_environment(None) == "development" + + def test_empty_string_returns_development(self): + assert normalize_environment("") == "development" + + def test_case_insensitive(self): + assert normalize_environment("PRODUCTION") == "production" + assert normalize_environment("Staging") == "staging" + + def test_strips_whitespace(self): + assert normalize_environment(" production ") == "production" -@pytest.mark.parametrize("environment", ["development", "test", "local", ""]) -def test_auto_create_schema_is_allowed_only_for_local_environments(environment): - assert_auto_create_schema_allowed(environment, enabled=True) +class TestAssertAutoCreateSchemaAllowed: + """assert_auto_create_schema_allowed 测试""" + def test_development_enabled_ok(self): + # development 环境允许 auto_create + assert_auto_create_schema_allowed("development", True) -@pytest.mark.parametrize("environment", ["staging", "production"]) -def test_disabled_auto_create_schema_is_allowed_everywhere(environment): - assert_auto_create_schema_allowed(environment, enabled=False) + def test_development_disabled_ok(self): + assert_auto_create_schema_allowed("development", False) + + def test_staging_disabled_ok(self): + # staging 禁用时没问题 + assert_auto_create_schema_allowed("staging", False) + + def test_production_disabled_ok(self): + assert_auto_create_schema_allowed("production", False) + + def test_staging_enabled_raises(self): + with pytest.raises(RuntimeError, match="AUTO_CREATE_SCHEMA"): + assert_auto_create_schema_allowed("staging", True) + + def test_production_enabled_raises(self): + with pytest.raises(RuntimeError, match="AUTO_CREATE_SCHEMA"): + assert_auto_create_schema_allowed("production", True) + + def test_case_insensitive_blocked(self): + with pytest.raises(RuntimeError): + assert_auto_create_schema_allowed("PRODUCTION", True) + with pytest.raises(RuntimeError): + assert_auto_create_schema_allowed("Staging", True) + + def test_none_environment_enabled_ok(self): + # None 视为 development,允许 + assert_auto_create_schema_allowed(None, True) + + def test_custom_env_enabled_ok(self): + # 其他环境不受限制 + assert_auto_create_schema_allowed("test", True) + assert_auto_create_schema_allowed("qa", True) + + def test_blocked_environments_count(self): + # 确认只有 staging 和 production 被阻止 + assert "staging" in BLOCKED_AUTO_CREATE_ENVIRONMENTS + assert "production" in BLOCKED_AUTO_CREATE_ENVIRONMENTS + assert len(BLOCKED_AUTO_CREATE_ENVIRONMENTS) == 2 diff --git a/tests/unit/test_sms_service.py b/tests/unit/test_sms_service.py new file mode 100755 index 000000000..9bc2db0c0 --- /dev/null +++ b/tests/unit/test_sms_service.py @@ -0,0 +1,178 @@ +"""SMS Service 单元测试""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from packages.adapters.sms.sms_service import ( + AliyunSmsService, + NoopSmsService, + get_sms_service, +) + + +class TestNoopSmsService: + """NoopSmsService 测试""" + + def test_send_verification_code_returns_true(self): + svc = NoopSmsService() + assert svc.send_verification_code("13800138000", "123456") is True + + def test_send_template_sms_returns_true(self): + svc = NoopSmsService() + assert svc.send_template_sms( + "13800138000", "SMS_123", {"code": "123456"} + ) is True + + def test_send_verification_code_empty_code(self): + svc = NoopSmsService() + assert svc.send_verification_code("13800138000", "") is True + + +class TestAliyunSmsServiceInit: + """AliyunSmsService 初始化测试""" + + def test_default_values_from_env(self, monkeypatch): + monkeypatch.setenv("ALIYUN_SMS_ACCESS_KEY_ID", "env_key") + monkeypatch.setenv("ALIYUN_SMS_ACCESS_KEY_SECRET", "env_secret") + monkeypatch.setenv("ALIYUN_SMS_SIGN_NAME", "env_sign") + monkeypatch.setenv("ALIYUN_SMS_VERIFY_TEMPLATE_ID", "env_tpl") + + svc = AliyunSmsService() + assert svc.access_key_id == "env_key" + assert svc.access_key_secret == "env_secret" + assert svc.sign_name == "env_sign" + assert svc.verify_template_id == "env_tpl" + + def test_explicit_params_override_env(self, monkeypatch): + monkeypatch.setenv("ALIYUN_SMS_ACCESS_KEY_ID", "env_key") + + svc = AliyunSmsService(access_key_id="explicit_key") + assert svc.access_key_id == "explicit_key" + + def test_default_sign_name(self, monkeypatch): + monkeypatch.delenv("ALIYUN_SMS_SIGN_NAME", raising=False) + svc = AliyunSmsService() + assert svc.sign_name == "小应剪辑" + + def test_default_template_id(self, monkeypatch): + monkeypatch.delenv("ALIYUN_SMS_VERIFY_TEMPLATE_ID", raising=False) + svc = AliyunSmsService() + assert svc.verify_template_id == "SMS_123456789" + + +class TestAliyunSmsServiceSend: + """发送短信测试(mock SDK)""" + + @pytest.fixture + def svc(self): + return AliyunSmsService( + access_key_id="key", + access_key_secret="secret", + sign_name="测试签名", + verify_template_id="SMS_VERIFY", + ) + + def test_send_verification_code_delegates_to_template(self, svc): + """验证码调用 send_template_sms""" + with patch.object(svc, "send_template_sms", return_value=True) as mock_send: + result = svc.send_verification_code("13800138000", "654321") + assert result is True + mock_send.assert_called_once_with( + "13800138000", "SMS_VERIFY", {"code": "654321"} + ) + + def test_send_template_sms_success(self, svc): + """发送成功返回 True""" + mock_body = MagicMock() + mock_body.code = "OK" + mock_body.message = "OK" + mock_response = MagicMock() + mock_response.body = mock_body + + with patch.dict("sys.modules"): + # mock 整个 alibabacloud 模块 + mock_client_cls = MagicMock() + mock_client_cls.return_value.send_sms.return_value = mock_response + + mock_dysms_models = MagicMock() + mock_dysms_models.SendSmsRequest = MagicMock(return_value=MagicMock()) + + mock_openapi_models = MagicMock() + mock_openapi_models.Config = MagicMock() + + with patch.object(svc, "_AliyunSmsService__import_sdk", create=True): + pass + + # 直接 patch 模块名来模拟 SDK 存在 + import sys + sys.modules["alibabacloud_dysmsapi20170525"] = MagicMock() + sys.modules["alibabacloud_dysmsapi20170525.models"] = mock_dysms_models + sys.modules["alibabacloud_dysmsapi20170525.client"] = MagicMock( + Client=mock_client_cls + ) + sys.modules["alibabacloud_tea_openapi"] = MagicMock() + sys.modules["alibabacloud_tea_openapi.models"] = mock_openapi_models + + try: + result = svc.send_template_sms( + "13800138000", "SMS_TPL", {"code": "123"} + ) + assert result is True + finally: + for key in [ + "alibabacloud_dysmsapi20170525", + "alibabacloud_dysmsapi20170525.models", + "alibabacloud_dysmsapi20170525.client", + "alibabacloud_tea_openapi", + "alibabacloud_tea_openapi.models", + ]: + sys.modules.pop(key, None) + + def test_send_template_sms_sdk_not_installed(self, svc): + """SDK 未安装返回 False""" + with patch.object(svc, "send_template_sms"): + pass + # 确保没有 SDK 时返回 False + import sys + saved_modules = {} + for key in list(sys.modules.keys()): + if "alibabacloud" in key: + saved_modules[key] = sys.modules.pop(key) + + try: + result = svc.send_template_sms("13800138000", "tpl", {}) + assert result is False + finally: + sys.modules.update(saved_modules) + + +class TestGetSmsService: + """工厂函数测试""" + + def test_default_noop(self, monkeypatch): + monkeypatch.delenv("SMS_PROVIDER", raising=False) + svc = get_sms_service() + assert isinstance(svc, NoopSmsService) + + def test_noop_provider(self, monkeypatch): + monkeypatch.setenv("SMS_PROVIDER", "noop") + svc = get_sms_service() + assert isinstance(svc, NoopSmsService) + + def test_aliyun_provider(self, monkeypatch): + monkeypatch.setenv("SMS_PROVIDER", "aliyun") + svc = get_sms_service() + assert isinstance(svc, AliyunSmsService) + + def test_case_insensitive_provider(self, monkeypatch): + monkeypatch.setenv("SMS_PROVIDER", "AliYun") + svc = get_sms_service() + assert isinstance(svc, AliyunSmsService) + + def test_unknown_provider_falls_back_to_noop(self, monkeypatch): + monkeypatch.setenv("SMS_PROVIDER", "unknown") + svc = get_sms_service() + assert isinstance(svc, NoopSmsService)