From d213b9a3d3d69a8f92f63da62ecaa1d553f2a4f8 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Wed, 22 Jul 2026 18:23:47 +0800 Subject: [PATCH] style: auto-fix black formatting --- scripts/ci/pr_auto_scan.py | 119 ++++++++++--------- tests/unit/test_generation_tasks.py | 24 +--- tests/unit/test_jobs.py | 4 +- tests/unit/test_verification_code_service.py | 50 +++----- 4 files changed, 86 insertions(+), 111 deletions(-) diff --git a/scripts/ci/pr_auto_scan.py b/scripts/ci/pr_auto_scan.py index 483f7819a..9dd7efc93 100644 --- a/scripts/ci/pr_auto_scan.py +++ b/scripts/ci/pr_auto_scan.py @@ -3,6 +3,7 @@ PR自动扫描器:扫描所有open PR,对CI全绿的进行自动审批/合并 作为短作业模式的兜底机制,每5分钟运行一次 """ + import argparse import json import sys @@ -12,22 +13,20 @@ import time import os -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_generation_tasks.py b/tests/unit/test_generation_tasks.py index a4c23206f..bd74fd9e5 100755 --- a/tests/unit/test_generation_tasks.py +++ b/tests/unit/test_generation_tasks.py @@ -126,9 +126,7 @@ class TestCreateGenerationTaskUseCase: """ID 会自动生成""" mock_repo.create.side_effect = lambda t: t - cmd = CreateGenerationTaskCommand( - project_id="proj-1", asset_library_id="lib-1" - ) + cmd = CreateGenerationTaskCommand(project_id="proj-1", asset_library_id="lib-1") uc = CreateGenerationTaskUseCase(mock_repo) task = uc.execute(cmd) @@ -186,12 +184,8 @@ class TestListUserTasksFilteredUseCase: 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 - ) + 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) def test_list_with_status_filter(self, mock_repo): """按状态筛选""" @@ -201,12 +195,8 @@ class TestListUserTasksFilteredUseCase: uc = ListUserTasksFilteredUseCase(mock_repo) uc.execute("user-1", status="running") - 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" - ) + 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") def test_list_with_pagination(self, mock_repo): """分页查询""" @@ -217,9 +207,7 @@ class TestListUserTasksFilteredUseCase: result = uc.execute("user-1", 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 - ) + mock_repo.list_by_user_filtered.assert_called_once_with("user-1", status=None, limit=10, offset=20) def test_list_empty_result(self, mock_repo): """空结果""" diff --git a/tests/unit/test_jobs.py b/tests/unit/test_jobs.py index 61ba82790..c9d178c47 100755 --- a/tests/unit/test_jobs.py +++ b/tests/unit/test_jobs.py @@ -183,9 +183,7 @@ class TestUpdateJobProgressUseCase: 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="处理中" - ) + cmd = UpdateJobProgressCommand(job_id=job.id, progress=50.0, current_stage="处理中") uc = UpdateJobProgressUseCase(mock_repo) result = uc.execute(cmd) diff --git a/tests/unit/test_verification_code_service.py b/tests/unit/test_verification_code_service.py index 978473dc0..87549bb0f 100755 --- a/tests/unit/test_verification_code_service.py +++ b/tests/unit/test_verification_code_service.py @@ -120,9 +120,7 @@ class TestGenerateNormal: 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" - ) + code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND, custom_code="888888") assert err is None assert code.code == "888888" @@ -132,9 +130,7 @@ class TestGenerateNormal: 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 - ) + code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND, ttl_seconds=60) assert err is None # 过期时间 - 创建时间 ≈ 60 秒 @@ -174,9 +170,7 @@ class TestGenerateRateLimit: def test_resend_cooldown_blocked(self, service, mock_repo): """冷却期内发送被拒绝""" # 10 秒前刚发过一条 - recent = make_code( - created_at=datetime.now(timezone.utc) - timedelta(seconds=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 @@ -188,6 +182,7 @@ class TestGenerateRateLimit: # 等待时间应接近 50 秒(60-10) # 提取数字验证范围 import re + match = re.search(r"(\d+)\s*秒", err) assert match wait = int(match.group(1)) @@ -196,9 +191,7 @@ class TestGenerateRateLimit: def test_resend_after_cooldown_ok(self, service, mock_repo): """超过冷却期可以重发""" # 2 分钟前发的,已过冷却 - old = make_code( - created_at=datetime.now(timezone.utc) - timedelta(seconds=120) - ) + 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 @@ -210,9 +203,7 @@ class TestGenerateRateLimit: def test_daily_limit_reached(self, service, mock_repo): """达到每日上限""" # 没有最近的(过了冷却),但今日已达上限 - old = make_code( - created_at=datetime.now(timezone.utc) - timedelta(hours=2) - ) + 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 @@ -253,13 +244,9 @@ class TestGenerateCustomRateLimitParams: def test_custom_cooldown(self, mock_repo): """自定义冷却时间""" - svc = VerificationCodeService( - repo=mock_repo, resend_cooldown=300, daily_limit=5 - ) + 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) - ) + 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 @@ -270,9 +257,7 @@ class TestGenerateCustomRateLimitParams: def test_custom_daily_limit(self, mock_repo): """自定义每日上限""" - svc = VerificationCodeService( - repo=mock_repo, resend_cooldown=60, daily_limit=3 - ) + 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 @@ -308,9 +293,7 @@ class TestVerifyParamValidation: 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 " - ) + ok, err = service.verify(" test@example.com ", CODE_TYPE_EMAIL_BIND, " 123456 ") assert ok assert err is None @@ -329,9 +312,7 @@ class TestVerifyNormal: 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 - ) + ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456", consume=True) assert ok assert err is None @@ -344,9 +325,7 @@ class TestVerifyNormal: 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 - ) + ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456", consume=False) assert ok assert err is None @@ -438,9 +417,7 @@ class TestVerifyCodeTypeIsolation: mock_repo.find_latest.side_effect = side_effect - ok, err = service.verify( - "test@example.com", CODE_TYPE_EMAIL_BIND, "123456" - ) + ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456") assert not ok assert "不存在或已过期" in err @@ -468,4 +445,5 @@ class TestConstants: 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