diff --git a/scripts/ci/pr_auto_scan.py b/scripts/ci/pr_auto_scan.py index 483f7819a..9ddfc936f 100644 --- a/scripts/ci/pr_auto_scan.py +++ b/scripts/ci/pr_auto_scan.py @@ -3,31 +3,30 @@ PR自动扫描器:扫描所有open PR,对CI全绿的进行自动审批/合并 作为短作业模式的兜底机制,每5分钟运行一次 """ + import argparse import json -import sys -import urllib.request -import urllib.error -import time import os +import sys +import time +import urllib.error +import urllib.request -def api_request(token, repo, endpoint, method='GET', data=None): +def api_request(token, repo, endpoint, method="GET", data=None): """Gitea API请求""" url = f"https://git.xiaoxiajianji.com/api/v1/repos/{repo}/{endpoint}" - headers = { - "Authorization": f"token {token}", - "Content-Type": "application/json" - } + headers = {"Authorization": f"token {token}", "Content-Type": "application/json"} body = json.dumps(data).encode() if data else None req = urllib.request.Request(url, data=body, headers=headers, method=method) - + # 跳过SSL验证 import ssl + ctx = ssl.create_default_context() ctx.check_hostname = False ctx.verify_mode = ssl.CERT_NONE - + try: resp = urllib.request.urlopen(req, context=ctx) return json.loads(resp.read().decode()), resp.status @@ -35,13 +34,12 @@ def api_request(token, repo, endpoint, method='GET', data=None): return json.loads(e.read().decode()) if e.read() else {"error": str(e)}, e.code -def get_open_prs(token, repo, base='develop'): +def get_open_prs(token, repo, base="develop"): """获取所有open的PR""" prs = [] page = 1 while True: - data, code = api_request(token, repo, - f"pulls?state=open&base={base}&sort=recentupdate&per_page=50&page={page}") + data, code = api_request(token, repo, f"pulls?state=open&base={base}&sort=recentupdate&per_page=50&page={page}") if code != 200 or not isinstance(data, list) or len(data) == 0: break prs.extend(data) @@ -63,11 +61,11 @@ def check_required_contexts(token, repo, sha, contexts): """检查指定的context是否都通过""" data, _ = get_commit_status(token, repo, sha) statuses = {s["context"]: s["status"] for s in data.get("statuses", [])} - + all_success = True any_pending = False any_failed = False - + for ctx in contexts: state = statuses.get(ctx, "pending") if state != "success": @@ -76,7 +74,7 @@ def check_required_contexts(token, repo, sha, contexts): any_pending = True if state in ("failure", "error"): any_failed = True - + return all_success, any_pending, any_failed, statuses @@ -85,8 +83,7 @@ def get_pr_files(token, repo, pr_number): files = [] page = 1 while True: - data, code = api_request(token, repo, - f"pulls/{pr_number}/files?per_page=300&page={page}") + data, code = api_request(token, repo, f"pulls/{pr_number}/files?per_page=300&page={page}") if code != 200 or not isinstance(data, list) or len(data) == 0: break files.extend(data) @@ -116,34 +113,44 @@ def has_approval(token, repo, pr_number): def approve_pr(token, repo, pr_number): """审批PR""" # 创建review - data, code = api_request(token, repo, f"pulls/{pr_number}/reviews", - method="POST", - data={"event": "PENDING", "body": "CI全绿,自动审批通过。"}) - + data, code = api_request( + token, + repo, + f"pulls/{pr_number}/reviews", + method="POST", + data={"event": "PENDING", "body": "CI全绿,自动审批通过。"}, + ) + if code not in (200, 201): return False, f"创建review失败: HTTP {code}" - + review_id = data.get("id") if data.get("state") == "APPROVED": return True, "直接创建APPROVED成功" - + if not review_id: return False, "未获取到review ID" - + # submit为APPROVED - data2, code2 = api_request(token, repo, + data2, code2 = api_request( + token, + repo, f"pulls/{pr_number}/reviews/{review_id}/events", method="POST", - data={"event": "APPROVED", "body": "CI全绿,自动审批通过。"}) - + data={"event": "APPROVED", "body": "CI全绿,自动审批通过。"}, + ) + if code2 in (200, 201): return True, "审批提交成功" else: # 尝试另一个端点 - data3, code3 = api_request(token, repo, + data3, code3 = api_request( + token, + repo, f"pulls/{pr_number}/reviews/{review_id}", method="POST", - data={"event": "APPROVED", "body": "CI全绿,自动审批通过。"}) + data={"event": "APPROVED", "body": "CI全绿,自动审批通过。"}, + ) if code3 in (200, 201): return True, "审批提交成功(备用端点)" return False, f"审批提交失败: HTTP {code2}/{code3}" @@ -153,25 +160,29 @@ def merge_pr(token, repo, pr_number): """合并PR(squash merge)""" # 等待几秒让状态同步 time.sleep(30) - + # 检查PR状态 pr_data, code = api_request(token, repo, f"pulls/{pr_number}") if code != 200: return False, f"获取PR状态失败: HTTP {code}" if pr_data.get("state") != "open": return False, f"PR状态不是open: {pr_data.get('state')}" - + # 执行squash merge - data, code = api_request(token, repo, f"pulls/{pr_number}/merge", + data, code = api_request( + token, + repo, + f"pulls/{pr_number}/merge", method="POST", data={ "do": "squash", "merge_title_field": "", "merge_message_field": "", "delete_branch_after_merge": True, - "force_merge": False - }) - + "force_merge": False, + }, + ) + if code == 200: return True, "合并成功" elif code == 405: @@ -189,11 +200,11 @@ def main(): parser.add_argument("--merge", action="store_true", help="执行自动合并") parser.add_argument("--dry-run", default="false", help="试运行模式") parser.add_argument("--max-prs", type=int, default=20, help="最多处理的PR数") - + args = parser.parse_args() - + dry_run = args.dry_run.lower() == "true" - + # required contexts(与分支保护一致) REQUIRED_CONTEXTS_FULL = [ "CI/CD Pipeline / Validate - Code Quality (pull_request)", @@ -213,39 +224,39 @@ def main(): FRONTEND_ONLY_CONTEXT = [ "CI/CD Pipeline / Frontend Lint (pull_request)", ] - + # 获取所有open PR print(f"获取 {args.base} 分支的open PR...") prs = get_open_prs(args.token, args.repo, args.base) print(f"找到 {len(prs)} 个open PR") - + approved_count = 0 merged_count = 0 skipped_count = 0 - - for pr in prs[:args.max_prs]: + + for pr in prs[: args.max_prs]: pr_num = pr["number"] pr_title = pr["title"] head_sha = pr["head"]["sha"] base_ref = pr.get("base", {}).get("ref", "") - + # 跳过draft if pr.get("draft"): print(f"\n⏭️ #{pr_num} {pr_title[:50]} - draft,跳过") skipped_count += 1 continue - + # 跳过目标分支不对的 if base_ref != args.base: skipped_count += 1 continue - + print(f"\n--- #{pr_num} {pr_title[:60]} ---") - + # 判断是否纯前端 files = get_pr_files(args.token, args.repo, pr_num) frontend_only = is_frontend_only(files) - + if frontend_only: approve_contexts = FRONTEND_ONLY_CONTEXT merge_contexts = FRONTEND_ONLY_CONTEXT @@ -254,11 +265,10 @@ def main(): approve_contexts = REQUIRED_CONTEXTS_APPROVE merge_contexts = REQUIRED_CONTEXTS_FULL print(f" 类型: 全栈/后端改动 ({len(files)}个文件)") - + # 检查审批用的CI状态 - all_ok, pending, failed, _ = check_required_contexts( - args.token, args.repo, head_sha, approve_contexts) - + all_ok, pending, failed, _ = check_required_contexts(args.token, args.repo, head_sha, approve_contexts) + # === 自动审批 === if args.approve and all_ok and not failed: if has_approval(args.token, args.repo, pr_num): @@ -278,16 +288,17 @@ def main(): print(f" ❌ CI有失败项,跳过审批") elif pending: print(f" ⏳ CI仍在运行,跳过") - + # === 自动合并 === if args.merge: # 检查合并用的CI状态 merge_ok, merge_pending, merge_failed, _ = check_required_contexts( - args.token, args.repo, head_sha, merge_contexts) - + args.token, args.repo, head_sha, merge_contexts + ) + # 检查审批 approved = has_approval(args.token, args.repo, pr_num) - + if merge_ok and approved and not merge_failed: if dry_run: print(f" 🎯 [DRY-RUN] 将自动合并") @@ -305,7 +316,7 @@ def main(): print(f" ❌ 合并条件未满足: CI有失败") elif not approved: print(f" ⏳ 合并条件未满足: 无审批") - + print(f"\n=== 扫描结果 ===") print(f" 处理PR数: {min(len(prs), args.max_prs)}") print(f" 自动审批: {approved_count} 个") diff --git a/tests/unit/test_jobs.py b/tests/unit/test_jobs.py new file mode 100755 index 000000000..c9d178c47 --- /dev/null +++ b/tests/unit/test_jobs.py @@ -0,0 +1,576 @@ +""" +Job 应用层用例单元测试(第十八波) + +覆盖: +- CreateJobUseCase +- SubmitJobUseCase +- UpdateJobProgressUseCase +- CompleteJobUseCase +- FailJobUseCase +- RetryJobUseCase +- CancelJobUseCase +- GetJobUseCase +- ListJobsUseCase +- GetJobStatisticsUseCase +""" + +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() + + +def make_job( + status=JobStatus.PENDING, + job_type=JobType.VIDEO_COMPOSE, + project_id="proj-1", + **kwargs, +): + job = Job.create( + project_id=project_id, + job_type=job_type, + **kwargs, + ) + # 绕过状态机直接设置状态(测试构造用) + if status != JobStatus.PENDING: + object.__setattr__(job, "status", status) + return job + + +# ============================================================ +# CreateJobUseCase +# ============================================================ + + +class TestCreateJobUseCase: + """CreateJobUseCase 创建任务""" + + def test_create_success(self, mock_repo): + """正常创建任务""" + mock_repo.create.side_effect = lambda j: j + + cmd = CreateJobCommand( + project_id="proj-1", + job_type=JobType.VIDEO_COMPOSE, + payload={"key": "val"}, + source_id="src-1", + created_by_user_id="user-1", + max_retries=5, + ) + uc = CreateJobUseCase(mock_repo) + job = uc.execute(cmd) + + assert job.project_id == "proj-1" + assert job.job_type == JobType.VIDEO_COMPOSE + assert job.payload == {"key": "val"} + assert job.source_id == "src-1" + assert job.created_by_user_id == "user-1" + assert job.max_retries == 5 + assert job.status == JobStatus.PENDING + assert job.progress == 0.0 + mock_repo.create.assert_called_once() + + def test_create_default_values(self, mock_repo): + """默认参数""" + mock_repo.create.side_effect = lambda j: j + + cmd = CreateJobCommand(project_id="proj-1", job_type="video_compose") + uc = CreateJobUseCase(mock_repo) + job = uc.execute(cmd) + + assert job.payload == {} + assert job.source_id == "" + assert job.created_by_user_id == "" + assert job.max_retries == 3 + + def test_create_string_job_type(self, mock_repo): + """字符串类型的 job_type 也支持""" + mock_repo.create.side_effect = lambda j: j + + cmd = CreateJobCommand(project_id="proj-1", job_type="asset_ingest") + uc = CreateJobUseCase(mock_repo) + job = uc.execute(cmd) + + assert job.job_type == JobType.ASSET_INGEST + + +# ============================================================ +# SubmitJobUseCase +# ============================================================ + + +class TestSubmitJobUseCase: + """SubmitJobUseCase 提交任务""" + + def test_submit_success(self, mock_repo): + """正常提交 pending 任务""" + job = make_job(status=JobStatus.PENDING) + mock_repo.get.return_value = job + mock_repo.update.side_effect = lambda j: j + + uc = SubmitJobUseCase(mock_repo) + result = uc.execute(job.id, celery_task_id="celery-123") + + assert result.status == JobStatus.RUNNING + assert result.celery_task_id == "celery-123" + assert result.current_stage == "已提交,等待执行" + assert result.started_at is not None + mock_repo.update.assert_called_once() + + def test_submit_job_not_found(self, mock_repo): + """任务不存在""" + mock_repo.get.return_value = None + + uc = SubmitJobUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + uc.execute("nonexistent") + + def test_submit_already_running(self, mock_repo): + """已经是 running 状态不能再提交""" + job = make_job(status=JobStatus.RUNNING) + mock_repo.get.return_value = job + + uc = SubmitJobUseCase(mock_repo) + with pytest.raises(ValueError, match="只有 pending 状态"): + uc.execute(job.id) + + def test_submit_without_celery_id(self, mock_repo): + """不传 celery_task_id 也可以""" + job = make_job(status=JobStatus.PENDING) + mock_repo.get.return_value = job + mock_repo.update.side_effect = lambda j: j + + uc = SubmitJobUseCase(mock_repo) + result = uc.execute(job.id) + + assert result.status == JobStatus.RUNNING + assert result.celery_task_id == "" + + +# ============================================================ +# UpdateJobProgressUseCase +# ============================================================ + + +class TestUpdateJobProgressUseCase: + """UpdateJobProgressUseCase 更新进度""" + + def test_update_progress_success(self, mock_repo): + """正常更新进度""" + job = make_job(status=JobStatus.RUNNING) + mock_repo.get.return_value = job + mock_repo.update.side_effect = lambda j: j + + cmd = UpdateJobProgressCommand(job_id=job.id, progress=50.0, current_stage="处理中") + uc = UpdateJobProgressUseCase(mock_repo) + result = uc.execute(cmd) + + assert result.progress == 50.0 + assert result.current_stage == "处理中" + mock_repo.update.assert_called_once() + + def test_update_progress_job_not_found(self, mock_repo): + """任务不存在""" + mock_repo.get.return_value = None + + cmd = UpdateJobProgressCommand(job_id="nope", progress=50.0) + uc = UpdateJobProgressUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + uc.execute(cmd) + + def test_update_progress_not_running(self, mock_repo): + """非 running 状态不能更新进度""" + job = make_job(status=JobStatus.PENDING) + mock_repo.get.return_value = job + + cmd = UpdateJobProgressCommand(job_id=job.id, progress=50.0) + uc = UpdateJobProgressUseCase(mock_repo) + with pytest.raises(ValueError, match="只有 running 状态"): + uc.execute(cmd) + + +# ============================================================ +# CompleteJobUseCase +# ============================================================ + + +class TestCompleteJobUseCase: + """CompleteJobUseCase 完成任务""" + + def test_complete_from_running(self, mock_repo): + """从 running 状态完成""" + job = make_job(status=JobStatus.RUNNING) + mock_repo.get.return_value = job + mock_repo.update.side_effect = lambda j: j + + cmd = CompleteJobCommand(job_id=job.id, result={"output": "ok"}) + uc = CompleteJobUseCase(mock_repo) + result = uc.execute(cmd) + + assert result.status == JobStatus.SUCCESS + assert result.progress == 100.0 + assert result.result == {"output": "ok"} + assert result.completed_at is not None + mock_repo.update.assert_called_once() + + def test_complete_from_pending(self, mock_repo): + """从 pending 状态也可以直接完成""" + job = make_job(status=JobStatus.PENDING) + mock_repo.get.return_value = job + mock_repo.update.side_effect = lambda j: j + + cmd = CompleteJobCommand(job_id=job.id) + uc = CompleteJobUseCase(mock_repo) + result = uc.execute(cmd) + + assert result.status == JobStatus.SUCCESS + assert result.progress == 100.0 + + def test_complete_job_not_found(self, mock_repo): + """任务不存在""" + mock_repo.get.return_value = None + + cmd = CompleteJobCommand(job_id="nope") + uc = CompleteJobUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + uc.execute(cmd) + + def test_complete_already_failed(self, mock_repo): + """已失败的任务不能直接标记完成""" + job = make_job(status=JobStatus.FAILED) + job.error_message = "some error" + mock_repo.get.return_value = job + + cmd = CompleteJobCommand(job_id=job.id) + uc = CompleteJobUseCase(mock_repo) + with pytest.raises(ValueError, match="只有 running/pending"): + uc.execute(cmd) + + +# ============================================================ +# FailJobUseCase +# ============================================================ + + +class TestFailJobUseCase: + """FailJobUseCase 失败任务""" + + def test_fail_from_running(self, mock_repo): + """从 running 状态失败""" + job = make_job(status=JobStatus.RUNNING) + mock_repo.get.return_value = job + mock_repo.update.side_effect = lambda j: j + + cmd = FailJobCommand(job_id=job.id, error_message="网络超时") + uc = FailJobUseCase(mock_repo) + result = uc.execute(cmd) + + assert result.status == JobStatus.FAILED + assert result.error_message == "网络超时" + assert result.current_stage == "失败" + assert result.completed_at is not None + mock_repo.update.assert_called_once() + + def test_fail_pending_rejected_by_domain(self, mock_repo): + """pending 状态不能直接失败(领域状态机约束)""" + job = make_job(status=JobStatus.PENDING) + mock_repo.get.return_value = job + + cmd = FailJobCommand(job_id=job.id, error_message="资源不足") + uc = FailJobUseCase(mock_repo) + with pytest.raises(ValueError, match="非法状态转换"): + uc.execute(cmd) + + def test_fail_job_not_found(self, mock_repo): + """任务不存在""" + mock_repo.get.return_value = None + + cmd = FailJobCommand(job_id="nope", error_message="err") + uc = FailJobUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + uc.execute(cmd) + + +# ============================================================ +# RetryJobUseCase +# ============================================================ + + +class TestRetryJobUseCase: + """RetryJobUseCase 重试任务""" + + def test_retry_success(self, mock_repo): + """失败任务重试成功""" + job = make_job(status=JobStatus.FAILED, max_retries=3) + job.retry_count = 0 + job.error_message = "timeout" + mock_repo.get.return_value = job + mock_repo.update.side_effect = lambda j: j + + uc = RetryJobUseCase(mock_repo) + result = uc.execute(job.id) + + assert result.status == JobStatus.PENDING + assert result.retry_count == 1 + assert result.progress == 0.0 + assert result.error_message == "" + assert result.started_at is None + assert result.completed_at is None + assert result.celery_task_id == "" + mock_repo.update.assert_called_once() + + def test_retry_job_not_found(self, mock_repo): + """任务不存在""" + mock_repo.get.return_value = None + + uc = RetryJobUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + uc.execute("nope") + + def test_retry_exceeds_max_retries(self, mock_repo): + """超过最大重试次数不可重试""" + job = make_job(status=JobStatus.FAILED, max_retries=3) + job.retry_count = 3 + mock_repo.get.return_value = job + + uc = RetryJobUseCase(mock_repo) + with pytest.raises(ValueError, match="不可重试"): + uc.execute(job.id) + + +# ============================================================ +# CancelJobUseCase +# ============================================================ + + +class TestCancelJobUseCase: + """CancelJobUseCase 取消任务""" + + def test_cancel_pending(self, mock_repo): + """取消 pending 任务""" + job = make_job(status=JobStatus.PENDING) + mock_repo.get.return_value = job + mock_repo.update.side_effect = lambda j: j + + uc = CancelJobUseCase(mock_repo) + result = uc.execute(job.id) + + assert result.status == JobStatus.CANCELLED + assert result.current_stage == "已取消" + mock_repo.update.assert_called_once() + + def test_cancel_running(self, mock_repo): + """取消 running 任务""" + job = make_job(status=JobStatus.RUNNING) + mock_repo.get.return_value = job + mock_repo.update.side_effect = lambda j: j + + uc = CancelJobUseCase(mock_repo) + result = uc.execute(job.id) + + assert result.status == JobStatus.CANCELLED + + def test_cancel_job_not_found(self, mock_repo): + """任务不存在""" + mock_repo.get.return_value = None + + uc = CancelJobUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + uc.execute("nope") + + def test_cancel_already_success(self, mock_repo): + """已成功的任务不能取消""" + job = make_job(status=JobStatus.SUCCESS) + mock_repo.get.return_value = job + + uc = CancelJobUseCase(mock_repo) + with pytest.raises(ValueError, match="终态"): + uc.execute(job.id) + + def test_cancel_already_failed(self, mock_repo): + """已失败的任务不能取消(走重试)""" + job = make_job(status=JobStatus.FAILED) + job.error_message = "err" + mock_repo.get.return_value = job + + uc = CancelJobUseCase(mock_repo) + with pytest.raises(ValueError, match="终态"): + uc.execute(job.id) + + +# ============================================================ +# GetJobUseCase +# ============================================================ + + +class TestGetJobUseCase: + """GetJobUseCase 获取任务""" + + def test_get_existing(self, mock_repo): + """获取存在的任务""" + job = make_job() + mock_repo.get.return_value = job + + uc = GetJobUseCase(mock_repo) + result = uc.execute(job.id) + + assert result is job + mock_repo.get.assert_called_once_with(job.id) + + def test_get_not_found(self, mock_repo): + """获取不存在的任务返回 None""" + mock_repo.get.return_value = None + + uc = GetJobUseCase(mock_repo) + result = uc.execute("nope") + + assert result is None + + +# ============================================================ +# ListJobsUseCase +# ============================================================ + + +class TestListJobsUseCase: + """ListJobsUseCase 列出任务""" + + def test_list_by_project(self, mock_repo): + """按项目列出""" + jobs = [make_job(), make_job()] + mock_repo.list_by_project.return_value = jobs + + uc = ListJobsUseCase(mock_repo) + result = uc.execute(project_id="proj-1") + + assert len(result) == 2 + mock_repo.list_by_project.assert_called_once() + + def test_list_by_project_with_filters(self, mock_repo): + """按项目 + 类型 + 状态过滤""" + mock_repo.list_by_project.return_value = [] + + uc = ListJobsUseCase(mock_repo) + uc.execute( + project_id="proj-1", + job_type=JobType.VIDEO_COMPOSE, + status=JobStatus.RUNNING, + limit=20, + offset=10, + ) + + mock_repo.list_by_project.assert_called_once_with( + "proj-1", + job_type=JobType.VIDEO_COMPOSE, + status=JobStatus.RUNNING, + limit=20, + offset=10, + ) + + def test_list_by_user(self, mock_repo): + """按用户列出""" + jobs = [make_job()] + mock_repo.list_by_user.return_value = jobs + + uc = ListJobsUseCase(mock_repo) + result = uc.execute(user_id="user-1") + + assert len(result) == 1 + mock_repo.list_by_user.assert_called_once() + + def test_list_no_filter_raises(self, mock_repo): + """不指定 project_id 或 user_id 报错""" + uc = ListJobsUseCase(mock_repo) + with pytest.raises(ValueError, match="必须指定"): + uc.execute() + + def test_list_project_takes_precedence(self, mock_repo): + """同时传 project_id 和 user_id,优先按项目查""" + mock_repo.list_by_project.return_value = [] + + uc = ListJobsUseCase(mock_repo) + uc.execute(project_id="proj-1", user_id="user-1") + + mock_repo.list_by_project.assert_called_once() + mock_repo.list_by_user.assert_not_called() + + +# ============================================================ +# GetJobStatisticsUseCase +# ============================================================ + + +class TestGetJobStatisticsUseCase: + """GetJobStatisticsUseCase 任务统计""" + + def test_stats_counts(self, mock_repo): + """统计各状态数量""" + mock_repo.count_by_project.side_effect = lambda pid, status=None: { + None: 10, # total + JobStatus.PENDING: 2, + JobStatus.RUNNING: 3, + JobStatus.SUCCESS: 4, + JobStatus.FAILED: 1, + }[status] + + uc = GetJobStatisticsUseCase(mock_repo) + stats = uc.execute("proj-1") + + assert stats["project_id"] == "proj-1" + assert stats["total"] == 10 + assert stats["pending"] == 2 + assert stats["running"] == 3 + assert stats["success"] == 4 + assert stats["failed"] == 1 + # 总共调用 5 次 count_by_project + assert mock_repo.count_by_project.call_count == 5 + + +# ============================================================ +# Command 对象 +# ============================================================ + + +class TestCommandObjects: + """命令对象基本属性""" + + def test_create_job_command_defaults(self): + cmd = CreateJobCommand(project_id="p1", job_type=JobType.VIDEO_COMPOSE) + assert cmd.payload == {} + assert cmd.source_id == "" + assert cmd.created_by_user_id == "" + assert cmd.max_retries == 3 + + def test_update_progress_command_defaults(self): + cmd = UpdateJobProgressCommand(job_id="j1", progress=50.0) + assert cmd.current_stage == "" + + def test_complete_job_command_defaults(self): + cmd = CompleteJobCommand(job_id="j1") + assert cmd.result == {} + + def test_fail_job_command(self): + cmd = FailJobCommand(job_id="j1", error_message="err") + assert cmd.error_message == "err" diff --git a/tests/unit/test_verification_code_service.py b/tests/unit/test_verification_code_service.py new file mode 100755 index 000000000..87549bb0f --- /dev/null +++ b/tests/unit/test_verification_code_service.py @@ -0,0 +1,449 @@ +""" +验证码服务单元测试(第十七波) + +覆盖: +- VerificationCodeService.generate +- VerificationCodeService.verify +- 频控逻辑(冷却 + 每日上限) +""" + +from datetime import datetime, timedelta, timezone +from unittest.mock import MagicMock + +import pytest + +from packages.application.auth.verification_code_service import ( + CODE_TYPE_EMAIL_BIND, + CODE_TYPE_EMAIL_LOGIN, + CODE_TYPE_PHONE_BIND, + DAILY_LIMIT, + DEFAULT_TTL_SECONDS, + MAX_ATTEMPTS, + RESEND_COOLDOWN_SECONDS, + VerificationCodeService, +) +from packages.domain.verification_code import VerificationCode + + +@pytest.fixture +def mock_repo(): + """mock 验证码仓储""" + return MagicMock() + + +@pytest.fixture +def service(mock_repo): + """验证码服务实例""" + return VerificationCodeService(repo=mock_repo) + + +def make_code( + recipient="test@example.com", + code_type=CODE_TYPE_EMAIL_BIND, + code="123456", + ttl=300, + used=False, + attempts=0, + created_at=None, +): + """构造一个验证码实体""" + now = created_at or datetime.now(timezone.utc) + return VerificationCode( + id="test-code-id", + recipient=recipient, + code=code, + code_type=code_type, + expires_at=now + timedelta(seconds=ttl), + used_at=now if used else None, + attempts=attempts, + created_at=now, + ) + + +# ============================================================ +# generate - 参数校验 +# ============================================================ + + +class TestGenerateParamValidation: + """generate 参数校验""" + + def test_empty_recipient(self, service): + """空接收方""" + code, err = service.generate("", CODE_TYPE_EMAIL_BIND) + assert code is None + assert "不能为空" in err + + def test_whitespace_recipient_stripped(self, service, mock_repo): + """前后空格会被 strip 掉,正常生成""" + mock_repo.find_latest.return_value = None + mock_repo.count_today.return_value = 0 + code, err = service.generate(" test@example.com ", CODE_TYPE_EMAIL_BIND) + assert err is None + assert code is not None + assert code.recipient == "test@example.com" + + def test_invalid_code_type(self, service): + """无效验证码类型""" + code, err = service.generate("test@example.com", "invalid_type") + assert code is None + assert "无效的验证码类型" in err + + +# ============================================================ +# generate - 正常生成 +# ============================================================ + + +class TestGenerateNormal: + """generate 正常生成场景""" + + def test_generate_success(self, service, mock_repo): + """正常生成验证码""" + mock_repo.find_latest.return_value = None + mock_repo.count_today.return_value = 0 + + code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND) + + assert err is None + assert code is not None + assert code.recipient == "test@example.com" + assert code.code_type == CODE_TYPE_EMAIL_BIND + assert len(code.code) == 6 + assert code.code.isdigit() + assert not code.is_used + assert not code.is_expired + mock_repo.save.assert_called_once() + + def test_custom_code(self, service, mock_repo): + """自定义验证码""" + mock_repo.find_latest.return_value = None + mock_repo.count_today.return_value = 0 + + code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND, custom_code="888888") + + assert err is None + assert code.code == "888888" + + def test_custom_ttl(self, service, mock_repo): + """自定义有效期""" + mock_repo.find_latest.return_value = None + mock_repo.count_today.return_value = 0 + + code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND, ttl_seconds=60) + + assert err is None + # 过期时间 - 创建时间 ≈ 60 秒 + delta = (code.expires_at - code.created_at).total_seconds() + assert delta == 60 + + def test_default_ttl_used_when_not_specified(self, service, mock_repo): + """未指定 ttl 时使用默认值""" + mock_repo.find_latest.return_value = None + mock_repo.count_today.return_value = 0 + + code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND) + + assert err is None + delta = (code.expires_at - code.created_at).total_seconds() + assert delta == DEFAULT_TTL_SECONDS + + def test_phone_bind_type(self, service, mock_repo): + """手机号绑定类型也支持""" + mock_repo.find_latest.return_value = None + mock_repo.count_today.return_value = 0 + + code, err = service.generate("13800138000", CODE_TYPE_PHONE_BIND) + + assert err is None + assert code.code_type == CODE_TYPE_PHONE_BIND + + +# ============================================================ +# generate - 频控 +# ============================================================ + + +class TestGenerateRateLimit: + """generate 频控逻辑""" + + def test_resend_cooldown_blocked(self, service, mock_repo): + """冷却期内发送被拒绝""" + # 10 秒前刚发过一条 + recent = make_code(created_at=datetime.now(timezone.utc) - timedelta(seconds=10)) + mock_repo.find_latest.return_value = recent + mock_repo.count_today.return_value = 1 + + code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND) + + assert code is None + assert "发送太频繁" in err + assert "秒后再试" in err + # 等待时间应接近 50 秒(60-10) + # 提取数字验证范围 + import re + + match = re.search(r"(\d+)\s*秒", err) + assert match + wait = int(match.group(1)) + assert 45 <= wait <= 55 + + def test_resend_after_cooldown_ok(self, service, mock_repo): + """超过冷却期可以重发""" + # 2 分钟前发的,已过冷却 + old = make_code(created_at=datetime.now(timezone.utc) - timedelta(seconds=120)) + mock_repo.find_latest.return_value = old + mock_repo.count_today.return_value = 1 + + code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND) + + assert err is None + assert code is not None + + def test_daily_limit_reached(self, service, mock_repo): + """达到每日上限""" + # 没有最近的(过了冷却),但今日已达上限 + old = make_code(created_at=datetime.now(timezone.utc) - timedelta(hours=2)) + mock_repo.find_latest.return_value = old + mock_repo.count_today.return_value = DAILY_LIMIT + + code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND) + + assert code is None + assert "今日发送次数已达上限" in err + + def test_daily_limit_not_reached(self, service, mock_repo): + """未达每日上限可以发""" + mock_repo.find_latest.return_value = None + mock_repo.count_today.return_value = DAILY_LIMIT - 1 + + code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND) + + assert err is None + assert code is not None + + def test_no_history_first_time_ok(self, service, mock_repo): + """首次发送,无历史记录""" + mock_repo.find_latest.return_value = None + mock_repo.count_today.return_value = 0 + + code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND) + + assert err is None + assert code is not None + mock_repo.save.assert_called_once() + + +# ============================================================ +# generate - 自定义频控参数 +# ============================================================ + + +class TestGenerateCustomRateLimitParams: + """自定义频控参数""" + + def test_custom_cooldown(self, mock_repo): + """自定义冷却时间""" + svc = VerificationCodeService(repo=mock_repo, resend_cooldown=300, daily_limit=5) + # 60 秒前发的,默认冷却 60 秒就够了,但这里设了 300 秒 + recent = make_code(created_at=datetime.now(timezone.utc) - timedelta(seconds=60)) + mock_repo.find_latest.return_value = recent + mock_repo.count_today.return_value = 1 + + code, err = svc.generate("test@example.com", CODE_TYPE_EMAIL_BIND) + + assert code is None + assert "发送太频繁" in err + + def test_custom_daily_limit(self, mock_repo): + """自定义每日上限""" + svc = VerificationCodeService(repo=mock_repo, resend_cooldown=60, daily_limit=3) + mock_repo.find_latest.return_value = None + mock_repo.count_today.return_value = 3 + + code, err = svc.generate("test@example.com", CODE_TYPE_EMAIL_BIND) + + assert code is None + assert "今日发送次数已达上限" in err + + +# ============================================================ +# verify - 参数校验 +# ============================================================ + + +class TestVerifyParamValidation: + """verify 参数校验""" + + def test_empty_recipient(self, service): + """空接收方""" + ok, err = service.verify("", CODE_TYPE_EMAIL_BIND, "123456") + assert not ok + assert "参数不完整" in err + + def test_empty_code(self, service): + """空验证码""" + ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "") + assert not ok + assert "参数不完整" in err + + def test_whitespace_stripped(self, service, mock_repo): + """前后空格会被 strip""" + code = make_code(code="123456") + mock_repo.find_latest.return_value = code + mock_repo.count_today.return_value = 0 + + ok, err = service.verify(" test@example.com ", CODE_TYPE_EMAIL_BIND, " 123456 ") + + assert ok + assert err is None + + +# ============================================================ +# verify - 正常验证 +# ============================================================ + + +class TestVerifyNormal: + """verify 正常验证场景""" + + def test_verify_success_consume(self, service, mock_repo): + """验证成功并消耗""" + code = make_code(code="123456") + mock_repo.find_latest.return_value = code + + ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456", consume=True) + + assert ok + assert err is None + assert code.is_used # 被标记为已使用 + # save 被调用了两次:一次 increment_attempts 后,一次 mark_used 后 + assert mock_repo.save.call_count >= 2 + + def test_verify_success_no_consume(self, service, mock_repo): + """验证成功但不消耗""" + code = make_code(code="123456") + mock_repo.find_latest.return_value = code + + ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456", consume=False) + + assert ok + assert err is None + assert not code.is_used # 未被标记 + + def test_verify_code_not_found(self, service, mock_repo): + """验证码不存在""" + mock_repo.find_latest.return_value = None + + ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456") + + assert not ok + assert "不存在或已过期" in err + + def test_verify_wrong_code(self, service, mock_repo): + """验证码错误""" + code = make_code(code="123456") + mock_repo.find_latest.return_value = code + + ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "999999") + + assert not ok + assert "验证码错误" in err + # 尝试次数增加了 + assert code.attempts == 1 + + def test_verify_already_used(self, service, mock_repo): + """验证码已使用""" + code = make_code(code="123456", used=True) + mock_repo.find_latest.return_value = code + + ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456") + + assert not ok + assert "已使用" in err + + def test_verify_expired(self, service, mock_repo): + """验证码已过期""" + code = make_code(code="123456", ttl=-60) # 已过期 60 秒 + mock_repo.find_latest.return_value = code + + ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456") + + assert not ok + assert "已过期" in err + + def test_verify_attempts_exceeded(self, service, mock_repo): + """超过最大尝试次数""" + code = make_code(code="123456", attempts=MAX_ATTEMPTS) + mock_repo.find_latest.return_value = code + + ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456") + + assert not ok + assert "验证次数过多" in err + # verify 里先 increment_attempts 再判断,所以这里 attempts 应该是 MAX_ATTEMPTS + 1 + assert code.attempts == MAX_ATTEMPTS + 1 + + def test_attempts_increment_on_wrong_code(self, service, mock_repo): + """错误验证码会增加尝试次数""" + code = make_code(code="123456", attempts=0) + mock_repo.find_latest.return_value = code + + service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "000000") + assert code.attempts == 1 + + service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "000001") + assert code.attempts == 2 + + +# ============================================================ +# verify - 不同 code_type 互不干扰 +# ============================================================ + + +class TestVerifyCodeTypeIsolation: + """不同验证码类型互不干扰""" + + def test_email_bind_vs_email_login(self, service, mock_repo): + """用 email_login 类型的验证码去验证 email_bind 应该失败""" + code = make_code(code_type=CODE_TYPE_EMAIL_LOGIN, code="123456") + mock_repo.find_latest.return_value = None # 按 email_bind 查不到 + + # find_latest 按 code_type 查询,传 email_bind 返回 None + def side_effect(recipient, ct): + if ct == CODE_TYPE_EMAIL_LOGIN: + return code + return None + + mock_repo.find_latest.side_effect = side_effect + + ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456") + assert not ok + assert "不存在或已过期" in err + + +# ============================================================ +# 常量值检查 +# ============================================================ + + +class TestConstants: + """常量默认值校验""" + + def test_default_cooldown_60(self): + assert RESEND_COOLDOWN_SECONDS == 60 + + def test_default_daily_limit_10(self): + assert DAILY_LIMIT == 10 + + def test_default_max_attempts_5(self): + assert MAX_ATTEMPTS == 5 + + def test_default_ttl_300(self): + assert DEFAULT_TTL_SECONDS == 300 + + def test_valid_code_types_count(self): + """5 种验证码类型""" + from packages.application.auth.verification_code_service import VALID_CODE_TYPES + + assert len(VALID_CODE_TYPES) == 5