test: P3-1 第十八波 jobs应用层用例单元测试 37个 #724

Merged
xiaoxia merged 4 commits from feat/p3-1-jobs-tests into develop 2026-07-22 23:57:27 +08:00
3 changed files with 1094 additions and 58 deletions
+69 -58
View File
@@ -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):
"""合并PRsquash 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}")
+576
View File
@@ -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"
+449
View File
@@ -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