From a14ce113cf6287b7b0383345972b37cf28279a68 Mon Sep 17 00:00:00 2001 From: CI Bot Date: Wed, 22 Jul 2026 17:34:14 +0800 Subject: [PATCH 01/11] =?UTF-8?q?test:=20P3-1=20=E7=AC=AC=E5=8D=81?= =?UTF-8?q?=E4=B8=83=E6=B3=A2=20verification=5Fcode=5Fservice=20=E5=8D=95?= =?UTF-8?q?=E5=85=83=E6=B5=8B=E8=AF=95=2032=E4=B8=AA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - generate 参数校验/正常生成/自定义参数 8个 - generate 频控(冷却+每日上限+自定义参数)10个 - verify 参数校验/正常验证/异常场景 11个 - 验证码类型隔离 + 常量校验 3个 - 合计 32 个测试全部通过 --- tests/unit/test_verification_code_service.py | 471 +++++++++++++++++++ 1 file changed, 471 insertions(+) create mode 100755 tests/unit/test_verification_code_service.py diff --git a/tests/unit/test_verification_code_service.py b/tests/unit/test_verification_code_service.py new file mode 100755 index 000000000..978473dc0 --- /dev/null +++ b/tests/unit/test_verification_code_service.py @@ -0,0 +1,471 @@ +""" +验证码服务单元测试(第十七波) + +覆盖: +- 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 -- 2.54.0 From 7ab6b8a2c2a2756739ce6c5cf816a62835075ca6 Mon Sep 17 00:00:00 2001 From: CI Bot Date: Wed, 22 Jul 2026 17:36:49 +0800 Subject: [PATCH 02/11] =?UTF-8?q?test:=20P3-1=20=E7=AC=AC=E5=8D=81?= =?UTF-8?q?=E5=85=AB=E6=B3=A2=20jobs=E5=BA=94=E7=94=A8=E5=B1=82=E7=94=A8?= =?UTF-8?q?=E4=BE=8B=E5=8D=95=E5=85=83=E6=B5=8B=E8=AF=95=2037=E4=B8=AA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - CreateJobUseCase 3个 - SubmitJobUseCase 4个 - UpdateJobProgressUseCase 3个 - CompleteJobUseCase 4个 - FailJobUseCase 3个 - RetryJobUseCase 3个 - CancelJobUseCase 5个 - GetJobUseCase 2个 - ListJobsUseCase 5个 - GetJobStatisticsUseCase 1个 - Command对象 4个 - 合计 37 个测试全部通过 --- tests/unit/test_jobs.py | 578 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 578 insertions(+) create mode 100755 tests/unit/test_jobs.py diff --git a/tests/unit/test_jobs.py b/tests/unit/test_jobs.py new file mode 100755 index 000000000..61ba82790 --- /dev/null +++ b/tests/unit/test_jobs.py @@ -0,0 +1,578 @@ +""" +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" -- 2.54.0 From f82527aa82e93348878647e46a752fd39b7317aa Mon Sep 17 00:00:00 2001 From: CI Bot Date: Wed, 22 Jul 2026 17:38:31 +0800 Subject: [PATCH 03/11] =?UTF-8?q?test:=20P3-1=20=E7=AC=AC=E5=8D=81?= =?UTF-8?q?=E4=B9=9D=E6=B3=A2=20generation=5Ftasks=E5=BA=94=E7=94=A8?= =?UTF-8?q?=E5=B1=82=E7=94=A8=E4=BE=8B=E5=8D=95=E5=85=83=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=2017=E4=B8=AA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - CreateGenerationTaskUseCase 3个 - GetGenerationTaskUseCase 2个 - ListUserTasksFilteredUseCase 4个 - RetryGenerationTaskUseCase 5个 - Command/Filter/Result对象 3个 - 合计 17 个测试全部通过 --- tests/unit/test_generation_tasks.py | 332 ++++++++++++++++++++++++++++ 1 file changed, 332 insertions(+) create mode 100755 tests/unit/test_generation_tasks.py diff --git a/tests/unit/test_generation_tasks.py b/tests/unit/test_generation_tasks.py new file mode 100755 index 000000000..a4c23206f --- /dev/null +++ b/tests/unit/test_generation_tasks.py @@ -0,0 +1,332 @@ +""" +生成任务应用层用例单元测试(第十九波) + +覆盖: +- CreateGenerationTaskUseCase +- GetGenerationTaskUseCase +- ListUserTasksFilteredUseCase +- RetryGenerationTaskUseCase +- Command / Filter / Result 对象 +""" + +from unittest.mock import MagicMock + +import pytest + +from packages.application.generation_tasks import ( + CreateGenerationTaskCommand, + CreateGenerationTaskUseCase, + GetGenerationTaskUseCase, + ListGenerationTasksResult, + ListTasksFilter, + ListUserTasksFilteredUseCase, + RetryGenerationTaskUseCase, +) +from packages.domain.generation_task import GenerationTask, GenerationTaskStatus + + +@pytest.fixture +def mock_repo(): + return MagicMock() + + +def make_task(status=GenerationTaskStatus.PENDING, **kwargs): + task = GenerationTask( + id="task-1", + project_id="proj-1", + asset_library_id="lib-1", + strategy_id="strat-1", + template_id="tmpl-1", + asset_ids=["asset-1"], + title_ids=["title-1"], + voice_ids=["voice-1"], + created_by_user_id="user-1", + video_title="测试标题", + ) + if status != GenerationTaskStatus.PENDING: + object.__setattr__(task, "status", status) + # 应用额外 kwargs + for k, v in kwargs.items(): + object.__setattr__(task, k, v) + return task + + +# ============================================================ +# CreateGenerationTaskUseCase +# ============================================================ + + +class TestCreateGenerationTaskUseCase: + """CreateGenerationTaskUseCase 创建生成任务""" + + def test_create_success(self, mock_repo): + """正常创建任务""" + mock_repo.create.side_effect = lambda t: t + + cmd = CreateGenerationTaskCommand( + project_id="proj-1", + asset_library_id="lib-1", + strategy_id="strat-1", + voice_library_id="vlib-1", + template_id="tmpl-1", + asset_ids=["a1", "a2"], + title_ids=["t1"], + voice_ids=["v1"], + created_by_user_id="user-1", + source_edit_plan_id="plan-1", + asset_select_mode="auto", + batch_id="batch-1", + video_title="我的视频", + auto_retry_enabled=True, + auto_retry_max=3, + ) + uc = CreateGenerationTaskUseCase(mock_repo) + task = uc.execute(cmd) + + assert task.project_id == "proj-1" + assert task.asset_library_id == "lib-1" + assert task.strategy_id == "strat-1" + assert task.voice_library_id == "vlib-1" + assert task.template_id == "tmpl-1" + assert task.asset_ids == ["a1", "a2"] + assert task.title_ids == ["t1"] + assert task.voice_ids == ["v1"] + assert task.created_by_user_id == "user-1" + assert task.source_edit_plan_id == "plan-1" + assert task.asset_select_mode == "auto" + assert task.batch_id == "batch-1" + assert task.video_title == "我的视频" + assert task.auto_retry_enabled is True + assert task.auto_retry_max == 3 + assert task.status == GenerationTaskStatus.PENDING + assert task.progress == 0.0 + assert task.result_count == 0 + mock_repo.create.assert_called_once() + + def test_create_default_values(self, mock_repo): + """默认参数值""" + mock_repo.create.side_effect = lambda t: t + + cmd = CreateGenerationTaskCommand( + project_id="proj-1", + asset_library_id="lib-1", + ) + uc = CreateGenerationTaskUseCase(mock_repo) + task = uc.execute(cmd) + + assert task.asset_ids == [] + assert task.title_ids == [] + assert task.voice_ids == [] + assert task.created_by_user_id == "" + assert task.video_title == "" + assert task.auto_retry_enabled is False + assert task.auto_retry_max == 0 + + def test_create_id_is_generated(self, mock_repo): + """ID 会自动生成""" + mock_repo.create.side_effect = lambda t: t + + cmd = CreateGenerationTaskCommand( + project_id="proj-1", asset_library_id="lib-1" + ) + uc = CreateGenerationTaskUseCase(mock_repo) + task = uc.execute(cmd) + + assert task.id + assert isinstance(task.id, str) + assert len(task.id) > 10 # uuid hex + + +# ============================================================ +# GetGenerationTaskUseCase +# ============================================================ + + +class TestGetGenerationTaskUseCase: + """GetGenerationTaskUseCase 获取任务""" + + def test_get_existing(self, mock_repo): + """获取存在的任务""" + task = make_task() + mock_repo.get.return_value = task + + uc = GetGenerationTaskUseCase(mock_repo) + result = uc.execute("task-1") + + assert result is task + mock_repo.get.assert_called_once_with("task-1") + + def test_get_not_found(self, mock_repo): + """获取不存在的任务返回 None""" + mock_repo.get.return_value = None + + uc = GetGenerationTaskUseCase(mock_repo) + result = uc.execute("nonexistent") + + assert result is None + + +# ============================================================ +# ListUserTasksFilteredUseCase +# ============================================================ + + +class TestListUserTasksFilteredUseCase: + """ListUserTasksFilteredUseCase 按用户筛选任务""" + + def test_list_without_filters(self, mock_repo): + """无筛选条件查询""" + tasks = [make_task(), make_task()] + mock_repo.list_by_user_filtered.return_value = tasks + mock_repo.count_by_user_filtered.return_value = 2 + + uc = ListUserTasksFilteredUseCase(mock_repo) + result = uc.execute("user-1") + + 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 + ) + + def test_list_with_status_filter(self, mock_repo): + """按状态筛选""" + mock_repo.list_by_user_filtered.return_value = [] + mock_repo.count_by_user_filtered.return_value = 0 + + uc = ListUserTasksFilteredUseCase(mock_repo) + uc.execute("user-1", status="running") + + 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): + """分页查询""" + mock_repo.list_by_user_filtered.return_value = [] + mock_repo.count_by_user_filtered.return_value = 100 + + uc = ListUserTasksFilteredUseCase(mock_repo) + 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 + ) + + def test_list_empty_result(self, mock_repo): + """空结果""" + mock_repo.list_by_user_filtered.return_value = [] + mock_repo.count_by_user_filtered.return_value = 0 + + uc = ListUserTasksFilteredUseCase(mock_repo) + result = uc.execute("user-1", status="failed") + + assert result.items == [] + assert result.total == 0 + + +# ============================================================ +# RetryGenerationTaskUseCase +# ============================================================ + + +class TestRetryGenerationTaskUseCase: + """RetryGenerationTaskUseCase 重试失败任务""" + + def test_retry_success(self, mock_repo): + """失败任务重试成功""" + task = make_task( + status=GenerationTaskStatus.FAILED, + error_message="网络超时", + retry_count=0, + ) + mock_repo.get.return_value = task + mock_repo.update.side_effect = lambda t: t + + uc = RetryGenerationTaskUseCase(mock_repo) + result = uc.execute("task-1") + + assert result.status == GenerationTaskStatus.PENDING + assert result.retry_count == 1 + assert result.error_message == "" + assert result.error_info == {} + assert result.progress == 0.0 + assert result.result_count == 0 + assert result.started_at is None + assert result.completed_at is None + mock_repo.update.assert_called_once() + + def test_retry_not_found(self, mock_repo): + """任务不存在""" + mock_repo.get.return_value = None + + uc = RetryGenerationTaskUseCase(mock_repo) + with pytest.raises(ValueError, match="任务不存在"): + uc.execute("nonexistent") + + def test_retry_not_failed(self, mock_repo): + """非失败状态不能重试""" + task = make_task(status=GenerationTaskStatus.RUNNING) + mock_repo.get.return_value = task + + uc = RetryGenerationTaskUseCase(mock_repo) + with pytest.raises(ValueError, match="只有失败状态"): + uc.execute("task-1") + + def test_retry_pending_not_allowed(self, mock_repo): + """pending 状态不能重试""" + task = make_task(status=GenerationTaskStatus.PENDING) + mock_repo.get.return_value = task + + uc = RetryGenerationTaskUseCase(mock_repo) + with pytest.raises(ValueError, match="只有失败状态"): + uc.execute("task-1") + + def test_retry_preserves_id(self, mock_repo): + """重试复用同一个 task_id""" + task = make_task(status=GenerationTaskStatus.FAILED) + original_id = task.id + mock_repo.get.return_value = task + mock_repo.update.side_effect = lambda t: t + + uc = RetryGenerationTaskUseCase(mock_repo) + result = uc.execute("task-1") + + assert result.id == original_id + + +# ============================================================ +# Command / Filter / Result 对象 +# ============================================================ + + +class TestCommandAndDataObjects: + """命令对象和数据对象""" + + def test_create_command_defaults(self): + cmd = CreateGenerationTaskCommand() + assert cmd.project_id == "" + assert cmd.asset_library_id == "" + assert cmd.asset_ids == [] + assert cmd.title_ids == [] + assert cmd.voice_ids == [] + assert cmd.auto_retry_enabled is False + assert cmd.auto_retry_max == 0 + + def test_list_filter_defaults(self): + f = ListTasksFilter() + assert f.status is None + + def test_list_result(self): + task = make_task() + r = ListGenerationTasksResult(items=[task], total=1) + assert len(r.items) == 1 + assert r.total == 1 -- 2.54.0 From 1219a972f60aea95d96477624a8df6caabdaed56 Mon Sep 17 00:00:00 2001 From: CI Bot Date: Wed, 22 Jul 2026 17:40:19 +0800 Subject: [PATCH 04/11] =?UTF-8?q?test:=20P3-1=20=E7=AC=AC=E4=BA=8C?= =?UTF-8?q?=E5=8D=81=E6=B3=A2=20wechat=5Foauth=5Fservice=E5=8D=95=E5=85=83?= =?UTF-8?q?=E6=B5=8B=E8=AF=95=2028=E4=B8=AA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - MemoryStateStore 6个 - is_configured 5个 - generate_auth_url 5个 - handle_callback 9个 - WechatUserInfo + 工厂函数 3个 - 合计 28 个测试全部通过 --- tests/unit/test_wechat_oauth_service.py | 383 ++++++++++++++++++++++++ 1 file changed, 383 insertions(+) create mode 100755 tests/unit/test_wechat_oauth_service.py diff --git a/tests/unit/test_wechat_oauth_service.py b/tests/unit/test_wechat_oauth_service.py new file mode 100755 index 000000000..f4f67e392 --- /dev/null +++ b/tests/unit/test_wechat_oauth_service.py @@ -0,0 +1,383 @@ +""" +微信 OAuth 服务单元测试(第二十波) + +覆盖: +- MemoryStateStore (put / verify_and_consume / 过期清理) +- WechatOAuthService.is_configured +- WechatOAuthService.generate_auth_url (正常模式 + mock模式) +- WechatOAuthService.handle_callback (正常 / 缺code / state无效 / mock模式 / access_token失败 / userinfo失败 / 网络异常) +""" + +import time +from unittest.mock import MagicMock, patch + +import pytest + +from packages.application.auth.wechat_oauth_service import ( + STATE_TTL_SECONDS, + MemoryStateStore, + WechatOAuthService, + WechatUserInfo, + get_wechat_oauth_service, +) + + +# ============================================================ +# MemoryStateStore +# ============================================================ + + +class TestMemoryStateStore: + """MemoryStateStore 内存 state 存储""" + + def test_put_and_verify(self): + """放入并验证成功""" + store = MemoryStateStore() + store.put("state-1") + assert store.verify_and_consume("state-1") is True + + def test_verify_consumes_once(self): + """state 是一次性的,验证后即消费""" + store = MemoryStateStore() + store.put("state-1") + assert store.verify_and_consume("state-1") is True + assert store.verify_and_consume("state-1") is False + + def test_verify_nonexistent(self): + """验证不存在的 state""" + store = MemoryStateStore() + assert store.verify_and_consume("nonexistent") is False + + def test_expired_state_is_cleaned(self): + """过期的 state 会被清理""" + store = MemoryStateStore(ttl_seconds=1) # 1秒过期 + store.put("state-1") + time.sleep(1.1) + assert store.verify_and_consume("state-1") is False + + def test_put_cleans_expired(self): + """put 时会清理过期的""" + store = MemoryStateStore(ttl_seconds=1) + store.put("state-1") + time.sleep(1.1) + store.put("state-2") + # state-1 应该被清理掉了 + assert len(store._states) == 1 + assert "state-2" in store._states + + def test_default_ttl(self): + """默认 TTL 是 10 分钟""" + store = MemoryStateStore() + assert store._ttl == STATE_TTL_SECONDS + + +# ============================================================ +# WechatOAuthService - is_configured +# ============================================================ + + +class TestIsConfigured: + """is_configured 配置检查""" + + def test_fully_configured(self): + """三项都配置了""" + svc = WechatOAuthService( + app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" + ) + assert svc.is_configured() is True + + def test_missing_app_id(self): + """缺 app_id""" + svc = WechatOAuthService( + app_id="", app_secret="secret", redirect_uri="https://example.com/cb" + ) + assert svc.is_configured() is False + + def test_missing_app_secret(self): + """缺 app_secret""" + svc = WechatOAuthService( + app_id="wx123", app_secret="", redirect_uri="https://example.com/cb" + ) + assert svc.is_configured() is False + + def test_missing_redirect_uri(self): + """缺 redirect_uri""" + svc = WechatOAuthService( + app_id="wx123", app_secret="secret", redirect_uri="" + ) + assert svc.is_configured() is False + + def test_none_configured(self): + """全没配置""" + svc = WechatOAuthService(app_id="", app_secret="", redirect_uri="") + assert svc.is_configured() is False + + +# ============================================================ +# WechatOAuthService - generate_auth_url +# ============================================================ + + +class TestGenerateAuthUrl: + """generate_auth_url 生成授权链接""" + + def test_configured_mode(self): + """配置完整时生成正式微信授权链接""" + svc = WechatOAuthService( + app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" + ) + url, state = svc.generate_auth_url() + + assert "open.weixin.qq.com" in url + assert "appid=wx123" in url + assert "redirect_uri=" in url + assert "response_type=code" in url + assert "scope=snsapi_login" in url + assert f"state={state}" in url + assert "#wechat_redirect" in url + assert state # state 非空 + + def test_mock_mode(self): + """未配置时返回 mock URL""" + svc = WechatOAuthService(app_id="", app_secret="", redirect_uri="") + url, state = svc.generate_auth_url() + + assert "/mock/wechat/auth" in url + assert "app_id=mock" in url + assert f"state={state}" in url + assert state + + def test_custom_scope(self): + """自定义 scope""" + svc = WechatOAuthService( + app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" + ) + url, _ = svc.generate_auth_url(scope="snsapi_userinfo") + assert "scope=snsapi_userinfo" in url + + def test_state_is_unique(self): + """每次生成的 state 不同""" + svc = WechatOAuthService( + app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" + ) + _, state1 = svc.generate_auth_url() + _, state2 = svc.generate_auth_url() + assert state1 != state2 + + def test_state_stored_in_store(self): + """生成的 state 会存入 store,可被 callback 验证""" + store = MemoryStateStore() + svc = WechatOAuthService( + app_id="wx123", + app_secret="secret", + redirect_uri="https://example.com/cb", + state_store=store, + ) + _, state = svc.generate_auth_url() + assert store.verify_and_consume(state) is True + + +# ============================================================ +# WechatOAuthService - handle_callback +# ============================================================ + + +class TestHandleCallback: + """handle_callback 处理微信回调""" + + def test_missing_code(self): + """缺少授权码""" + svc = WechatOAuthService( + app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" + ) + user_info, err = svc.handle_callback("", "some-state") + assert user_info is None + assert "缺少授权码" in err + + def test_invalid_state(self): + """state 无效或已过期""" + svc = WechatOAuthService( + app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" + ) + user_info, err = svc.handle_callback("code123", "invalid-state") + assert user_info is None + assert "state" in err + + def test_empty_state(self): + """空 state""" + svc = WechatOAuthService( + app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" + ) + user_info, err = svc.handle_callback("code123", "") + assert user_info is None + assert "state" in err + + def test_mock_mode_success(self): + """mock 模式下返回模拟用户信息""" + svc = WechatOAuthService(app_id="", app_secret="", redirect_uri="") + # 先生成一个有效的 state + _, state = svc.generate_auth_url() + + user_info, err = svc.handle_callback("mock_code_123456", state) + + assert err is None + assert user_info is not None + assert user_info.openid.startswith("mock_") + assert user_info.unionid.startswith("mock_union_") + assert user_info.nickname == "微信测试用户" + + def test_configured_mode_success(self): + """配置完整时正常调用微信 API""" + svc = WechatOAuthService( + app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" + ) + _, state = svc.generate_auth_url() + + with patch("packages.application.auth.wechat_oauth_service.requests.get") as mock_get: + # access_token 响应 + token_resp = MagicMock() + token_resp.json.return_value = { + "access_token": "at_123", + "openid": "openid_abc", + "unionid": "unionid_xyz", + "expires_in": 7200, + } + # userinfo 响应 + user_resp = MagicMock() + user_resp.json.return_value = { + "openid": "openid_abc", + "nickname": "测试用户", + "headimgurl": "https://wx.qq.com/avatar.jpg", + "sex": 1, + } + mock_get.side_effect = [token_resp, user_resp] + + user_info, err = svc.handle_callback("code_abc", state) + + assert err is None + assert user_info is not None + assert user_info.openid == "openid_abc" + assert user_info.unionid == "unionid_xyz" + assert user_info.nickname == "测试用户" + assert user_info.avatar_url == "https://wx.qq.com/avatar.jpg" + # 应该调用了两次 get + assert mock_get.call_count == 2 + + def test_access_token_failed(self): + """access_token 接口返回错误""" + svc = WechatOAuthService( + app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" + ) + _, state = svc.generate_auth_url() + + with patch("packages.application.auth.wechat_oauth_service.requests.get") as mock_get: + err_resp = MagicMock() + err_resp.json.return_value = { + "errcode": 40029, + "errmsg": "invalid code", + } + mock_get.return_value = err_resp + + user_info, err = svc.handle_callback("bad_code", state) + + assert user_info is None + assert "微信授权失败" in err + assert "invalid code" in err + + def test_userinfo_failed(self): + """userinfo 接口返回错误""" + svc = WechatOAuthService( + app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" + ) + _, state = svc.generate_auth_url() + + with patch("packages.application.auth.wechat_oauth_service.requests.get") as mock_get: + token_resp = MagicMock() + token_resp.json.return_value = { + "access_token": "at_123", + "openid": "openid_abc", + } + err_resp = MagicMock() + err_resp.json.return_value = { + "errcode": 40001, + "errmsg": "invalid credential", + } + mock_get.side_effect = [token_resp, err_resp] + + user_info, err = svc.handle_callback("code_abc", state) + + assert user_info is None + assert "获取用户信息失败" in err + + def test_network_error(self): + """网络异常""" + import requests + + svc = WechatOAuthService( + app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" + ) + _, state = svc.generate_auth_url() + + with patch("packages.application.auth.wechat_oauth_service.requests.get") as mock_get: + mock_get.side_effect = requests.ConnectionError("timeout") + + user_info, err = svc.handle_callback("code_abc", state) + + assert user_info is None + assert "微信服务暂不可用" in err + + def test_state_one_time_use(self): + """state 一次性使用,重复使用会失败""" + svc = WechatOAuthService(app_id="", app_secret="", redirect_uri="") + _, state = svc.generate_auth_url() + + # 第一次成功 + user_info1, err1 = svc.handle_callback("code1", state) + assert err1 is None + assert user_info1 is not None + + # 第二次用同一个 state 失败 + user_info2, err2 = svc.handle_callback("code2", state) + assert user_info2 is None + assert "state" in err2 + + +# ============================================================ +# WechatUserInfo +# ============================================================ + + +class TestWechatUserInfo: + """WechatUserInfo 数据类""" + + def test_minimal_fields(self): + info = WechatUserInfo(openid="abc") + assert info.openid == "abc" + assert info.unionid == "" + assert info.nickname == "" + assert info.avatar_url == "" + + def test_full_fields(self): + info = WechatUserInfo( + openid="abc", + unionid="def", + nickname="测试", + avatar_url="https://example.com/avatar.jpg", + ) + assert info.openid == "abc" + assert info.unionid == "def" + assert info.nickname == "测试" + assert info.avatar_url == "https://example.com/avatar.jpg" + + +# ============================================================ +# get_wechat_oauth_service +# ============================================================ + + +class TestGetWechatOAuthService: + """工厂函数""" + + def test_returns_service_instance(self): + svc = get_wechat_oauth_service() + assert isinstance(svc, WechatOAuthService) -- 2.54.0 From 65e1695531d78d844be2086ad6ae5698ec2c77f3 Mon Sep 17 00:00:00 2001 From: CI Bot Date: Wed, 22 Jul 2026 17:55:43 +0800 Subject: [PATCH 05/11] =?UTF-8?q?test:=20P3-1=20=E7=AC=AC=E4=BA=8C?= =?UTF-8?q?=E5=8D=81=E4=B8=80=E6=B3=A2=20JWT+Password=E5=A7=94=E6=89=98?= =?UTF-8?q?=E5=B1=82=E5=8D=95=E5=85=83=E6=B5=8B=E8=AF=95=2024=E4=B8=AA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - JWTHandler 8个(create/verify/additional_claims/不同密钥/自定义配置) - JWT全局配置 2个 - PasswordHandler 12个(hash/verify/needs_rehash/强度校验/自定义rounds) - Password全局配置 3个 - 合计 24 个测试全部通过 --- tests/unit/test_auth_handlers.py | 247 +++++++++++++++++++++++++++++++ 1 file changed, 247 insertions(+) create mode 100755 tests/unit/test_auth_handlers.py diff --git a/tests/unit/test_auth_handlers.py b/tests/unit/test_auth_handlers.py new file mode 100755 index 000000000..0a804aa0e --- /dev/null +++ b/tests/unit/test_auth_handlers.py @@ -0,0 +1,247 @@ +""" +JWT + Password 委托层单元测试(第二十一波) + +覆盖: +- JWTHandler (create/verify/configure/get) +- PasswordHandler (hash/verify/needs_rehash/validate_strength/configure/get) +""" + +from unittest.mock import patch + +import pytest + +from packages.application.auth.jwt_handler import ( + JWTHandler, + configure_jwt_handler, + get_jwt_handler, +) +from packages.application.auth.password_handler import ( + PasswordHandler, + configure_password_handler, + get_password_handler, +) + +SECRET_KEY = "test-secret-key-for-unit-testing-only-not-for-production" + + +# ============================================================ +# JWTHandler +# ============================================================ + + +class TestJWTHandler: + """JWTHandler JWT 委托层""" + + def test_create_and_verify_access_token(self): + """创建并验证 access_token""" + handler = JWTHandler(secret_key=SECRET_KEY) + token = handler.create_access_token(user_id="user-123", role="admin") + + assert isinstance(token, str) + assert len(token) > 0 + + payload = handler.verify_access_token(token) + assert payload["sub"] == "user-123" + assert payload["role"] == "admin" + assert "exp" in payload + assert "type" in payload + assert payload["type"] == "access" + + def test_create_token_with_additional_claims(self): + """携带额外 claims""" + handler = JWTHandler(secret_key=SECRET_KEY) + token = handler.create_access_token( + user_id="user-1", + role="user", + additional_claims={"email": "a@b.com", "org_id": "org-1"}, + ) + payload = handler.verify_access_token(token) + assert payload["email"] == "a@b.com" + assert payload["org_id"] == "org-1" + + def test_create_token_default_role(self): + """默认 role 为空字符串""" + handler = JWTHandler(secret_key=SECRET_KEY) + token = handler.create_access_token(user_id="user-1") + payload = handler.verify_access_token(token) + assert payload["role"] == "" + + def test_verify_generic_token(self): + """verify_token 通用验证方法""" + handler = JWTHandler(secret_key=SECRET_KEY) + token = handler.create_access_token(user_id="user-1") + payload = handler.verify_token(token) + assert payload["sub"] == "user-1" + + def test_verify_invalid_token_raises(self): + """无效 token 验证失败""" + handler = JWTHandler(secret_key=SECRET_KEY) + with pytest.raises(Exception): + handler.verify_access_token("invalid-token") + + def test_verify_wrong_secret(self): + """用不同密钥签名的 token 验证失败""" + handler1 = JWTHandler(secret_key="key-a") + handler2 = JWTHandler(secret_key="key-b") + + token = handler1.create_access_token(user_id="user-1") + with pytest.raises(Exception): + handler2.verify_access_token(token) + + def test_custom_algorithm(self): + """自定义算法""" + handler = JWTHandler(secret_key=SECRET_KEY, algorithm="HS256") + token = handler.create_access_token(user_id="user-1") + payload = handler.verify_access_token(token) + assert payload["sub"] == "user-1" + + def test_custom_expire_minutes(self): + """自定义过期时间""" + handler = JWTHandler( + secret_key=SECRET_KEY, access_token_expire_minutes=60 + ) + token = handler.create_access_token(user_id="user-1") + payload = handler.verify_access_token(token) + assert payload["sub"] == "user-1" + + +# ============================================================ +# JWTHandler - 全局配置 +# ============================================================ + + +class TestJWTGlobalConfig: + """JWT 全局配置与获取""" + + def test_configure_and_get(self): + """配置后可以获取""" + handler = configure_jwt_handler( + secret_key=SECRET_KEY, access_token_expire_minutes=15 + ) + assert isinstance(handler, JWTHandler) + + got = get_jwt_handler() + assert got is handler + + def test_reconfigure_replaces(self): + """重新配置会替换""" + h1 = configure_jwt_handler(secret_key="key-a") + h2 = configure_jwt_handler(secret_key="key-b") + assert h1 is not h2 + assert get_jwt_handler() is h2 + + +# ============================================================ +# PasswordHandler +# ============================================================ + + +class TestPasswordHandler: + """PasswordHandler 密码委托层""" + + def test_hash_and_verify_correct(self): + """哈希并验证正确密码""" + handler = PasswordHandler() + hashed = handler.hash_password("MySecurePass123") + + assert isinstance(hashed, str) + assert hashed != "MySecurePass123" + assert handler.verify_password("MySecurePass123", hashed) is True + + def test_verify_wrong_password(self): + """验证错误密码""" + handler = PasswordHandler() + hashed = handler.hash_password("CorrectPass123") + assert handler.verify_password("WrongPass456", hashed) is False + + def test_hash_is_unique_each_time(self): + """同密码每次哈希不同(salt)""" + handler = PasswordHandler() + h1 = handler.hash_password("SamePass123") + h2 = handler.hash_password("SamePass123") + assert h1 != h2 + # 但都能验证通过 + assert handler.verify_password("SamePass123", h1) + assert handler.verify_password("SamePass123", h2) + + def test_needs_rehash_new_hash(self): + """新生成的哈希不需要重新计算""" + handler = PasswordHandler() + hashed = handler.hash_password("TestPass123") + assert handler.needs_rehash(hashed) is False + + def test_validate_strength_strong(self): + """强密码校验通过""" + handler = PasswordHandler() + ok, err = handler.validate_strength("StrongPass123") + assert ok is True + assert err is None or err == "" + + def test_validate_strength_too_short(self): + """密码太短""" + handler = PasswordHandler() + ok, err = handler.validate_strength("Ab1") + assert ok is False + assert err is not None + + def test_validate_strength_no_uppercase(self): + """缺少大写字母""" + handler = PasswordHandler() + ok, err = handler.validate_strength("lowercase123") + assert ok is False + assert err is not None + + def test_validate_strength_no_lowercase(self): + """缺少小写字母""" + handler = PasswordHandler() + ok, err = handler.validate_strength("UPPERCASE123") + assert ok is False + assert err is not None + + def test_validate_strength_no_digit(self): + """缺少数字""" + handler = PasswordHandler() + ok, err = handler.validate_strength("NoDigitHere") + assert ok is False + assert err is not None + + def test_hash_empty_password(self): + """空密码哈希报错""" + handler = PasswordHandler() + with pytest.raises((ValueError, Exception)): + handler.hash_password("") + + def test_custom_rounds(self): + """自定义 rounds(用低轮次测试更快)""" + handler = PasswordHandler(rounds=4) + hashed = handler.hash_password("TestPass123") + assert handler.verify_password("TestPass123", hashed) + + +# ============================================================ +# PasswordHandler - 全局配置 +# ============================================================ + + +class TestPasswordGlobalConfig: + """Password 全局配置与获取""" + + def test_get_default_handler(self): + """未配置时 get 返回默认实例""" + # 重置默认实例 + with patch("packages.application.auth.password_handler._default_handler", None): + handler = get_password_handler() + assert isinstance(handler, PasswordHandler) + + def test_configure_and_get(self): + """配置后可以获取""" + handler = configure_password_handler(rounds=4) + assert isinstance(handler, PasswordHandler) + got = get_password_handler() + assert got is handler + + def test_reconfigure_replaces(self): + """重新配置会替换""" + h1 = configure_password_handler(rounds=4) + h2 = configure_password_handler(rounds=6) + assert h1 is not h2 -- 2.54.0 From 7051e21204191f26cf0f4e23c74c8e44ebc8fdb8 Mon Sep 17 00:00:00 2001 From: CI Bot Date: Wed, 22 Jul 2026 17:58:55 +0800 Subject: [PATCH 06/11] =?UTF-8?q?test:=20P3-1=20=E7=AC=AC=E4=BA=8C?= =?UTF-8?q?=E5=8D=81=E4=BA=8C=E6=B3=A2=20TTS=E9=9F=B3=E9=A2=91=E5=90=88?= =?UTF-8?q?=E5=B9=B6+=E6=B5=81=E5=BC=8F=E6=9C=8D=E5=8A=A1=E5=8D=95?= =?UTF-8?q?=E5=85=83=E6=B5=8B=E8=AF=95=2014=E4=B8=AA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - AudioMerger 7个(空列表/单文件/多文件/wav格式/异常/临时目录清理) - TTSStreamingService 7个(空文本/过长文本/短文本成功/合成失败/无audio_url/分块推送/send_json容错) - 合计 14 个测试全部通过 --- tests/unit/test_tts_audio_merger.py | 296 ++++++++++++++++++++++++++++ 1 file changed, 296 insertions(+) create mode 100755 tests/unit/test_tts_audio_merger.py diff --git a/tests/unit/test_tts_audio_merger.py b/tests/unit/test_tts_audio_merger.py new file mode 100755 index 000000000..0c847ca2e --- /dev/null +++ b/tests/unit/test_tts_audio_merger.py @@ -0,0 +1,296 @@ +""" +TTS 相关单元测试(第二十二波) + +覆盖: +- AudioMerger (空列表/单文件/多文件合并/格式/异常) +""" + +import os +import subprocess +import tempfile + +import pytest + +from packages.application.tts_job.audio_merger import ( + AudioMergeError, + AudioMerger, +) + + +def _make_silence(duration: float = 0.5, sample_rate: int = 22050, fmt: str = "mp3") -> str: + """生成一段静音音频文件,返回路径。""" + tmp = tempfile.NamedTemporaryFile(suffix=f".{fmt}", delete=False) + tmp.close() + cmd = [ + "ffmpeg", "-y", "-f", "lavfi", + "-i", f"anullsrc=r={sample_rate}:cl=mono", + "-t", str(duration), + "-q:a", "9", + tmp.name, + ] + subprocess.run(cmd, capture_output=True, check=True) + return tmp.name + + +class TestAudioMerger: + """AudioMerger 音频合并器""" + + def test_empty_list_raises(self): + """空列表抛错""" + merger = AudioMerger() + with pytest.raises(AudioMergeError, match="没有可合并"): + merger.merge([]) + + def test_single_file_returns_content(self): + """单个文件直接返回内容""" + path = _make_silence(duration=0.3) + try: + merger = AudioMerger() + result = merger.merge([path]) + assert isinstance(result, bytes) + assert len(result) > 100 # 应该有有效数据 + # 应该和文件本身一致 + with open(path, "rb") as f: + original = f.read() + assert result == original + finally: + os.unlink(path) + + def test_two_files_merged(self): + """两个文件合并""" + p1 = _make_silence(duration=0.3) + p2 = _make_silence(duration=0.4) + try: + merger = AudioMerger() + result = merger.merge([p1, p2]) + assert isinstance(result, bytes) + assert len(result) > 200 # 合并后应该有数据 + # 写出来用 ffprobe 验证时长 + tmp = tempfile.NamedTemporaryFile(suffix=".mp3", delete=False) + tmp.write(result) + tmp.close() + try: + probe = subprocess.run( + ["ffprobe", "-v", "error", "-show_entries", "format=duration", + "-of", "default=noprint_wrappers=1:nokey=1", tmp.name], + capture_output=True, text=True, check=True, + ) + duration = float(probe.stdout.strip()) + # 0.3 + 0.4 = 0.7 秒左右,允许一定误差 + assert 0.5 < duration < 1.0 + finally: + os.unlink(tmp.name) + finally: + os.unlink(p1) + os.unlink(p2) + + def test_three_files_merged(self): + """三个文件合并""" + paths = [_make_silence(duration=0.2) for _ in range(3)] + try: + merger = AudioMerger() + result = merger.merge(paths) + assert isinstance(result, bytes) + assert len(result) > 200 + finally: + for p in paths: + os.unlink(p) + + def test_wav_format(self): + """wav 格式合并""" + p1 = _make_silence(duration=0.2, fmt="wav") + p2 = _make_silence(duration=0.2, fmt="wav") + try: + merger = AudioMerger() + result = merger.merge([p1, p2], output_format="wav") + assert isinstance(result, bytes) + # WAV 头部以 RIFF 开头 + assert result[:4] == b"RIFF" + finally: + os.unlink(p1) + os.unlink(p2) + + def test_nonexistent_file_raises(self): + """不存在的文件会抛错""" + merger = AudioMerger() + with pytest.raises(AudioMergeError): + merger.merge(["/nonexistent/path/a.mp3", "/nonexistent/path/b.mp3"]) + + def test_cleanup_temp_dir(self): + """临时目录会被清理""" + p1 = _make_silence(duration=0.2) + p2 = _make_silence(duration=0.2) + try: + import tempfile as _tf + before = set(os.listdir(_tf.gettempdir())) + merger = AudioMerger() + merger.merge([p1, p2]) + after = set(os.listdir(_tf.gettempdir())) + # 不应该残留 tts_merge_ 前缀的目录 + new_items = after - before + tts_items = [i for i in new_items if i.startswith("tts_merge_")] + assert len(tts_items) == 0, f"残留临时目录: {tts_items}" + finally: + os.unlink(p1) + os.unlink(p2) + + +# ============================================================ +# TTSStreamingService - 入口路由与边界 +# ============================================================ + + +class TestTTSStreamingServiceRouting: + """TTSStreamingService 入口路由与边界条件""" + + @pytest.mark.asyncio + async def test_empty_text_returns_error(self): + """空文本返回错误""" + from unittest.mock import AsyncMock, MagicMock + + from packages.application.tts_job.streaming_service import TTSStreamingService + + mock_cosy = MagicMock() + svc = TTSStreamingService(cosyvoice_service=mock_cosy) + ws = AsyncMock() + + await svc.synthesize_and_stream(ws, {"text": ""}) + + ws.send_json.assert_called_once() + call_args = ws.send_json.call_args[0][0] + assert call_args["type"] == "error" + assert "不能为空" in call_args["message"] + # 不应该调用 cosyvoice + mock_cosy.submit_synthesize_task.assert_not_called() + + @pytest.mark.asyncio + async def test_text_too_long_returns_error(self): + """文本过长返回错误""" + from unittest.mock import AsyncMock, MagicMock + + from packages.application.tts_job.streaming_service import ( + _MAX_TEXT_LENGTH, + TTSStreamingService, + ) + + mock_cosy = MagicMock() + svc = TTSStreamingService(cosyvoice_service=mock_cosy) + ws = AsyncMock() + + long_text = "a" * (_MAX_TEXT_LENGTH + 1) + await svc.synthesize_and_stream(ws, {"text": long_text}) + + ws.send_json.assert_called_once() + call_args = ws.send_json.call_args[0][0] + assert call_args["type"] == "error" + assert "过长" in call_args["message"] + mock_cosy.submit_synthesize_task.assert_not_called() + + @pytest.mark.asyncio + async def test_short_text_routes_to_short_path(self): + """短文本走短文本路径(单段合成)""" + from unittest.mock import AsyncMock, MagicMock, patch + + from packages.application.tts_job.streaming_service import TTSStreamingService + + mock_cosy = MagicMock() + mock_cosy.submit_synthesize_task.return_value = { + "audio_url": "https://example.com/audio.mp3", + "duration": 3.5, + } + svc = TTSStreamingService(cosyvoice_service=mock_cosy) + ws = AsyncMock() + + # mock 掉音频下载 + fake_audio = b"fake_audio_data" * 100 + with patch.object(svc, "_download_audio", return_value=fake_audio): + await svc.synthesize_and_stream(ws, {"text": "你好世界", "voice_id": "v1"}) + + # 应该调用了 cosy + mock_cosy.submit_synthesize_task.assert_called_once() + # 应该有 started 和 done 消息 + msg_types = [c[0][0]["type"] for c in ws.send_json.call_args_list] + assert "started" in msg_types + assert "done" in msg_types + # 应该有音频分块发送 + assert ws.send_bytes.call_count > 0 + + @pytest.mark.asyncio + async def test_short_text_cosy_error(self): + """短文本合成失败返回错误""" + from unittest.mock import AsyncMock, MagicMock + + from packages.application.cosyvoice_service import CosyVoiceError + from packages.application.tts_job.streaming_service import TTSStreamingService + + mock_cosy = MagicMock() + mock_cosy.submit_synthesize_task.side_effect = CosyVoiceError("音色不存在") + svc = TTSStreamingService(cosyvoice_service=mock_cosy) + ws = AsyncMock() + + await svc.synthesize_and_stream(ws, {"text": "你好", "voice_id": "v-bad"}) + + # 最后一条消息应该是 error + last_msg = ws.send_json.call_args_list[-1][0][0] + assert last_msg["type"] == "error" + assert "音色不存在" in last_msg["message"] + + @pytest.mark.asyncio + async def test_short_text_no_audio_url(self): + """合成结果没有 audio_url 返回错误""" + from unittest.mock import AsyncMock, MagicMock + + from packages.application.tts_job.streaming_service import TTSStreamingService + + mock_cosy = MagicMock() + mock_cosy.submit_synthesize_task.return_value = {"duration": 1.0} # 没有 audio_url + svc = TTSStreamingService(cosyvoice_service=mock_cosy) + ws = AsyncMock() + + await svc.synthesize_and_stream(ws, {"text": "你好"}) + + last_msg = ws.send_json.call_args_list[-1][0][0] + assert last_msg["type"] == "error" + assert "音频 URL" in last_msg["message"] + + @pytest.mark.asyncio + async def test_stream_audio_chunks_returns_total(self): + """_stream_audio_chunks 返回正确字节数,分块正确""" + from unittest.mock import AsyncMock, MagicMock + + from packages.application.tts_job.streaming_service import ( + _AUDIO_CHUNK_SIZE, + TTSStreamingService, + ) + + mock_cosy = MagicMock() + svc = TTSStreamingService(cosyvoice_service=mock_cosy) + ws = AsyncMock() + + # 生成 10000 字节的假音频 + audio_data = b"x" * 10000 + total = await svc._stream_audio_chunks(ws, audio_data) + + assert total == 10000 + # 应该分 ceil(10000/4096) = 3 块 + expected_chunks = (10000 + _AUDIO_CHUNK_SIZE - 1) // _AUDIO_CHUNK_SIZE + assert ws.send_bytes.call_count == expected_chunks + # 验证所有块拼接起来等于原数据 + all_bytes = b"".join(c[0][0] for c in ws.send_bytes.call_args_list) + assert all_bytes == audio_data + + @pytest.mark.asyncio + async def test_send_json_handles_error(self): + """_send_json 发送失败不抛出异常""" + from unittest.mock import AsyncMock, MagicMock + + from packages.application.tts_job.streaming_service import TTSStreamingService + + mock_cosy = MagicMock() + svc = TTSStreamingService(cosyvoice_service=mock_cosy) + ws = AsyncMock() + ws.send_json.side_effect = Exception("连接已断开") + + # 不应该抛异常 + await svc._send_json(ws, {"type": "done"}) + ws.send_json.assert_called_once() -- 2.54.0 From 8b7510e2bb0ad3d92a887b04481be5ce1c203eb2 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Wed, 22 Jul 2026 18:23:53 +0800 Subject: [PATCH 07/11] style: auto-fix black formatting --- scripts/ci/pr_auto_scan.py | 119 ++++++++++--------- tests/unit/test_auth_handlers.py | 8 +- tests/unit/test_generation_tasks.py | 24 +--- tests/unit/test_jobs.py | 4 +- tests/unit/test_tts_audio_merger.py | 31 +++-- tests/unit/test_verification_code_service.py | 50 +++----- tests/unit/test_wechat_oauth_service.py | 57 +++------ 7 files changed, 126 insertions(+), 167 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_auth_handlers.py b/tests/unit/test_auth_handlers.py index 0a804aa0e..0a31fd0a1 100755 --- a/tests/unit/test_auth_handlers.py +++ b/tests/unit/test_auth_handlers.py @@ -97,9 +97,7 @@ class TestJWTHandler: def test_custom_expire_minutes(self): """自定义过期时间""" - handler = JWTHandler( - secret_key=SECRET_KEY, access_token_expire_minutes=60 - ) + handler = JWTHandler(secret_key=SECRET_KEY, access_token_expire_minutes=60) token = handler.create_access_token(user_id="user-1") payload = handler.verify_access_token(token) assert payload["sub"] == "user-1" @@ -115,9 +113,7 @@ class TestJWTGlobalConfig: def test_configure_and_get(self): """配置后可以获取""" - handler = configure_jwt_handler( - secret_key=SECRET_KEY, access_token_expire_minutes=15 - ) + handler = configure_jwt_handler(secret_key=SECRET_KEY, access_token_expire_minutes=15) assert isinstance(handler, JWTHandler) got = get_jwt_handler() 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_tts_audio_merger.py b/tests/unit/test_tts_audio_merger.py index 0c847ca2e..051c8dcad 100755 --- a/tests/unit/test_tts_audio_merger.py +++ b/tests/unit/test_tts_audio_merger.py @@ -22,10 +22,16 @@ def _make_silence(duration: float = 0.5, sample_rate: int = 22050, fmt: str = "m tmp = tempfile.NamedTemporaryFile(suffix=f".{fmt}", delete=False) tmp.close() cmd = [ - "ffmpeg", "-y", "-f", "lavfi", - "-i", f"anullsrc=r={sample_rate}:cl=mono", - "-t", str(duration), - "-q:a", "9", + "ffmpeg", + "-y", + "-f", + "lavfi", + "-i", + f"anullsrc=r={sample_rate}:cl=mono", + "-t", + str(duration), + "-q:a", + "9", tmp.name, ] subprocess.run(cmd, capture_output=True, check=True) @@ -71,9 +77,19 @@ class TestAudioMerger: tmp.close() try: probe = subprocess.run( - ["ffprobe", "-v", "error", "-show_entries", "format=duration", - "-of", "default=noprint_wrappers=1:nokey=1", tmp.name], - capture_output=True, text=True, check=True, + [ + "ffprobe", + "-v", + "error", + "-show_entries", + "format=duration", + "-of", + "default=noprint_wrappers=1:nokey=1", + tmp.name, + ], + capture_output=True, + text=True, + check=True, ) duration = float(probe.stdout.strip()) # 0.3 + 0.4 = 0.7 秒左右,允许一定误差 @@ -122,6 +138,7 @@ class TestAudioMerger: p2 = _make_silence(duration=0.2) try: import tempfile as _tf + before = set(os.listdir(_tf.gettempdir())) merger = AudioMerger() merger.merge([p1, p2]) 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 diff --git a/tests/unit/test_wechat_oauth_service.py b/tests/unit/test_wechat_oauth_service.py index f4f67e392..c9783662f 100755 --- a/tests/unit/test_wechat_oauth_service.py +++ b/tests/unit/test_wechat_oauth_service.py @@ -21,7 +21,6 @@ from packages.application.auth.wechat_oauth_service import ( get_wechat_oauth_service, ) - # ============================================================ # MemoryStateStore # ============================================================ @@ -81,30 +80,22 @@ class TestIsConfigured: def test_fully_configured(self): """三项都配置了""" - svc = WechatOAuthService( - app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb") assert svc.is_configured() is True def test_missing_app_id(self): """缺 app_id""" - svc = WechatOAuthService( - app_id="", app_secret="secret", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="", app_secret="secret", redirect_uri="https://example.com/cb") assert svc.is_configured() is False def test_missing_app_secret(self): """缺 app_secret""" - svc = WechatOAuthService( - app_id="wx123", app_secret="", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="", redirect_uri="https://example.com/cb") assert svc.is_configured() is False def test_missing_redirect_uri(self): """缺 redirect_uri""" - svc = WechatOAuthService( - app_id="wx123", app_secret="secret", redirect_uri="" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="secret", redirect_uri="") assert svc.is_configured() is False def test_none_configured(self): @@ -123,9 +114,7 @@ class TestGenerateAuthUrl: def test_configured_mode(self): """配置完整时生成正式微信授权链接""" - svc = WechatOAuthService( - app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb") url, state = svc.generate_auth_url() assert "open.weixin.qq.com" in url @@ -149,17 +138,13 @@ class TestGenerateAuthUrl: def test_custom_scope(self): """自定义 scope""" - svc = WechatOAuthService( - app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb") url, _ = svc.generate_auth_url(scope="snsapi_userinfo") assert "scope=snsapi_userinfo" in url def test_state_is_unique(self): """每次生成的 state 不同""" - svc = WechatOAuthService( - app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb") _, state1 = svc.generate_auth_url() _, state2 = svc.generate_auth_url() assert state1 != state2 @@ -187,27 +172,21 @@ class TestHandleCallback: def test_missing_code(self): """缺少授权码""" - svc = WechatOAuthService( - app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb") user_info, err = svc.handle_callback("", "some-state") assert user_info is None assert "缺少授权码" in err def test_invalid_state(self): """state 无效或已过期""" - svc = WechatOAuthService( - app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb") user_info, err = svc.handle_callback("code123", "invalid-state") assert user_info is None assert "state" in err def test_empty_state(self): """空 state""" - svc = WechatOAuthService( - app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb") user_info, err = svc.handle_callback("code123", "") assert user_info is None assert "state" in err @@ -228,9 +207,7 @@ class TestHandleCallback: def test_configured_mode_success(self): """配置完整时正常调用微信 API""" - svc = WechatOAuthService( - app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb") _, state = svc.generate_auth_url() with patch("packages.application.auth.wechat_oauth_service.requests.get") as mock_get: @@ -265,9 +242,7 @@ class TestHandleCallback: def test_access_token_failed(self): """access_token 接口返回错误""" - svc = WechatOAuthService( - app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb") _, state = svc.generate_auth_url() with patch("packages.application.auth.wechat_oauth_service.requests.get") as mock_get: @@ -286,9 +261,7 @@ class TestHandleCallback: def test_userinfo_failed(self): """userinfo 接口返回错误""" - svc = WechatOAuthService( - app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb") _, state = svc.generate_auth_url() with patch("packages.application.auth.wechat_oauth_service.requests.get") as mock_get: @@ -313,9 +286,7 @@ class TestHandleCallback: """网络异常""" import requests - svc = WechatOAuthService( - app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb" - ) + svc = WechatOAuthService(app_id="wx123", app_secret="secret", redirect_uri="https://example.com/cb") _, state = svc.generate_auth_url() with patch("packages.application.auth.wechat_oauth_service.requests.get") as mock_get: -- 2.54.0 From 5b212bb571ed46f9ff1a3768b0638220c5cfd36b Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Wed, 22 Jul 2026 22:31:52 +0800 Subject: [PATCH 08/11] style: black + isort format fix --- scripts/ci/pr_auto_scan.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/scripts/ci/pr_auto_scan.py b/scripts/ci/pr_auto_scan.py index 9dd7efc93..9ddfc936f 100644 --- a/scripts/ci/pr_auto_scan.py +++ b/scripts/ci/pr_auto_scan.py @@ -6,11 +6,11 @@ PR自动扫描器:扫描所有open PR,对CI全绿的进行自动审批/合 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): -- 2.54.0 From 5241ec7d427198e1a1970aa4e99e8540e2bdc7e1 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Thu, 23 Jul 2026 00:56:25 +0800 Subject: [PATCH 09/11] =?UTF-8?q?style:=20=E4=BF=AE=E5=A4=8Dpr=5Fauto=5Fsc?= =?UTF-8?q?an.py=E7=9A=84ruff=20lint=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- scripts/ci/pr_auto_scan.py | 24 +++++++++++------------- 1 file changed, 11 insertions(+), 13 deletions(-) diff --git a/scripts/ci/pr_auto_scan.py b/scripts/ci/pr_auto_scan.py index 9ddfc936f..c4c5a3efb 100644 --- a/scripts/ci/pr_auto_scan.py +++ b/scripts/ci/pr_auto_scan.py @@ -6,8 +6,6 @@ PR自动扫描器:扫描所有open PR,对CI全绿的进行自动审批/合 import argparse import json -import os -import sys import time import urllib.error import urllib.request @@ -272,12 +270,12 @@ def main(): # === 自动审批 === if args.approve and all_ok and not failed: if has_approval(args.token, args.repo, pr_num): - print(f" ✅ 已有审批,跳过") + print(" ✅ 已有审批,跳过") else: if dry_run: - print(f" 🎯 [DRY-RUN] 将自动审批") + print(" 🎯 [DRY-RUN] 将自动审批") else: - print(f" 🎯 执行自动审批...") + print(" 🎯 执行自动审批...") ok, msg = approve_pr(args.token, args.repo, pr_num) if ok: print(f" ✅ 审批成功: {msg}") @@ -285,9 +283,9 @@ def main(): else: print(f" ❌ 审批失败: {msg}") elif failed: - print(f" ❌ CI有失败项,跳过审批") + print(" ❌ CI有失败项,跳过审批") elif pending: - print(f" ⏳ CI仍在运行,跳过") + print(" ⏳ CI仍在运行,跳过") # === 自动合并 === if args.merge: @@ -301,9 +299,9 @@ def main(): if merge_ok and approved and not merge_failed: if dry_run: - print(f" 🎯 [DRY-RUN] 将自动合并") + print(" 🎯 [DRY-RUN] 将自动合并") else: - print(f" 🎯 执行自动合并...") + print(" 🎯 执行自动合并...") ok, msg = merge_pr(args.token, args.repo, pr_num) if ok: print(f" ✅ 合并成功: {msg}") @@ -311,13 +309,13 @@ def main(): else: print(f" ⚠️ 合并失败: {msg}") elif merge_pending: - print(f" ⏳ 合并条件未满足: CI运行中") + print(" ⏳ 合并条件未满足: CI运行中") elif merge_failed: - print(f" ❌ 合并条件未满足: CI有失败") + print(" ❌ 合并条件未满足: CI有失败") elif not approved: - print(f" ⏳ 合并条件未满足: 无审批") + print(" ⏳ 合并条件未满足: 无审批") - print(f"\n=== 扫描结果 ===") + print("\n=== 扫描结果 ===") print(f" 处理PR数: {min(len(prs), args.max_prs)}") print(f" 自动审批: {approved_count} 个") print(f" 自动合并: {merged_count} 个") -- 2.54.0 From 40dccdb0b2c37dec1eaa82cead7ed681fe69020b Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Thu, 23 Jul 2026 07:02:44 +0800 Subject: [PATCH 10/11] =?UTF-8?q?fix(bandit):=20B017=20-=20=E7=94=A8Invali?= =?UTF-8?q?dTokenError=E6=9B=BF=E6=8D=A2=E5=AE=BD=E6=B3=9BException?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tests/unit/test_auth_handlers.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tests/unit/test_auth_handlers.py b/tests/unit/test_auth_handlers.py index 0a31fd0a1..e486bbebf 100755 --- a/tests/unit/test_auth_handlers.py +++ b/tests/unit/test_auth_handlers.py @@ -9,6 +9,7 @@ JWT + Password 委托层单元测试(第二十一波) from unittest.mock import patch import pytest +from jwt.exceptions import InvalidTokenError from packages.application.auth.jwt_handler import ( JWTHandler, @@ -76,7 +77,7 @@ class TestJWTHandler: def test_verify_invalid_token_raises(self): """无效 token 验证失败""" handler = JWTHandler(secret_key=SECRET_KEY) - with pytest.raises(Exception): + with pytest.raises(InvalidTokenError): handler.verify_access_token("invalid-token") def test_verify_wrong_secret(self): @@ -85,7 +86,7 @@ class TestJWTHandler: handler2 = JWTHandler(secret_key="key-b") token = handler1.create_access_token(user_id="user-1") - with pytest.raises(Exception): + with pytest.raises(InvalidTokenError): handler2.verify_access_token(token) def test_custom_algorithm(self): -- 2.54.0 From f6d0705eeb77b016cb2a8b66efcf36c10d9659d9 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Thu, 23 Jul 2026 09:50:55 +0800 Subject: [PATCH 11/11] =?UTF-8?q?fix:=20=E7=94=A8develop=E7=89=88=E6=9C=AC?= =?UTF-8?q?=E7=9A=84pr=5Fauto=5Fscan.py=E8=A7=A3=E5=86=B3=E5=86=B2?= =?UTF-8?q?=E7=AA=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- scripts/ci/pr_auto_scan.py | 95 +++++++++++++++++++++++++++++++++++--- 1 file changed, 88 insertions(+), 7 deletions(-) diff --git a/scripts/ci/pr_auto_scan.py b/scripts/ci/pr_auto_scan.py index c4c5a3efb..f17753d27 100644 --- a/scripts/ci/pr_auto_scan.py +++ b/scripts/ci/pr_auto_scan.py @@ -2,10 +2,15 @@ """ PR自动扫描器:扫描所有open PR,对CI全绿的进行自动审批/合并 作为短作业模式的兜底机制,每5分钟运行一次 + +新增:AI审查联动 - AI代码审查发现严重问题时,不自动审批 """ import argparse import json +import os +import re +import sys import time import urllib.error import urllib.request @@ -108,7 +113,55 @@ def has_approval(token, repo, pr_number): return any(r.get("state") == "APPROVED" for r in reviews if isinstance(r, dict)) -def approve_pr(token, repo, pr_number): +def get_ai_review_result(token, repo, pr_number): + """ + 检查AI代码审查结果,返回 (has_critical, review_body) + has_critical: 是否有严重问题(需修改的问题 > 0) + review_body: 最新的AI审查评论文本 + """ + # AI审查评论标记 + AI_REVIEW_MARKER = "AI_CODE_REVIEW_AUTO_COMMENT" + + comments, code = api_request(token, repo, f"issues/{pr_number}/comments") + if code != 200: + return False, None + + # 找最新的AI审查评论 + ai_comments = [c for c in comments if isinstance(c, dict) and AI_REVIEW_MARKER in c.get("body", "")] + + if not ai_comments: + return False, None + + # 按时间排序,取最新的 + latest = max(ai_comments, key=lambda c: c.get("created_at", "")) + body = latest.get("body", "") + + # 解析严重问题数量 + # 匹配 "严重问题数量:X 个" 或 "需修改的问题(严重)" 下的列表 + critical_count = 0 + + # 方式1:直接匹配数字 + match = re.search(r"严重问题数量[::]\s*(\d+)\s*个", body) + if match: + critical_count = int(match.group(1)) + else: + # 方式2:数 "需修改的问题" 章节下的条目数 + critical_section = re.search( + r"###\s*[❌⚠️].*?(?:需修改|问题).*?\n(.*?)(?=\n###|\Z)", + body, + re.DOTALL, + ) + if critical_section: + section_text = critical_section.group(1) + # 数编号条目 1. 2. 3. + items = re.findall(r"^\d+\.\s+\*\*", section_text, re.MULTILINE) + critical_count = len(items) + + has_critical = critical_count > 0 + return has_critical, body + + +def approve_pr(token, repo, pr_number, reason="CI全绿,自动审批通过。"): """审批PR""" # 创建review data, code = api_request( @@ -116,7 +169,7 @@ def approve_pr(token, repo, pr_number): repo, f"pulls/{pr_number}/reviews", method="POST", - data={"event": "PENDING", "body": "CI全绿,自动审批通过。"}, + data={"event": "PENDING", "body": reason}, ) if code not in (200, 201): @@ -135,7 +188,7 @@ def approve_pr(token, repo, pr_number): repo, f"pulls/{pr_number}/reviews/{review_id}/events", method="POST", - data={"event": "APPROVED", "body": "CI全绿,自动审批通过。"}, + data={"event": "APPROVED", "body": reason}, ) if code2 in (200, 201): @@ -147,13 +200,25 @@ def approve_pr(token, repo, pr_number): repo, f"pulls/{pr_number}/reviews/{review_id}", method="POST", - data={"event": "APPROVED", "body": "CI全绿,自动审批通过。"}, + data={"event": "APPROVED", "body": reason}, ) if code3 in (200, 201): return True, "审批提交成功(备用端点)" return False, f"审批提交失败: HTTP {code2}/{code3}" +def add_pr_label(token, repo, pr_number, label): + """给PR添加标签""" + data, code = api_request( + token, + repo, + f"issues/{pr_number}/labels", + method="POST", + data={"labels": [label]}, + ) + return code in (200, 201) + + def merge_pr(token, repo, pr_number): """合并PR(squash merge)""" # 等待几秒让状态同步 @@ -198,6 +263,7 @@ 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数") + parser.add_argument("--skip-ai-review", action="store_true", help="跳过AI审查检查(强制审批)") args = parser.parse_args() @@ -231,12 +297,13 @@ def main(): approved_count = 0 merged_count = 0 skipped_count = 0 + ai_blocked_count = 0 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", "") + base_ref = pr.get("base", {}).get("re", "") # 跳过draft if pr.get("draft"): @@ -267,8 +334,19 @@ def main(): # 检查审批用的CI状态 all_ok, pending, failed, _ = check_required_contexts(args.token, args.repo, head_sha, approve_contexts) + # === AI审查检查 === + ai_has_critical = False + if not args.skip_ai_review and all_ok and not failed and args.approve: + ai_has_critical, ai_body = get_ai_review_result(args.token, args.repo, pr_num) + if ai_has_critical: + print(" ⚠️ AI审查发现严重问题,阻止自动审批") + ai_blocked_count += 1 + # 给PR打标签便于人工识别 + if not dry_run: + add_pr_label(args.token, args.repo, pr_num, "ai-review/需修改") + # === 自动审批 === - if args.approve and all_ok and not failed: + if args.approve and all_ok and not failed and not ai_has_critical: if has_approval(args.token, args.repo, pr_num): print(" ✅ 已有审批,跳过") else: @@ -282,6 +360,8 @@ def main(): approved_count += 1 else: print(f" ❌ 审批失败: {msg}") + elif ai_has_critical: + print(" 🚫 AI审查阻止审批(人工可手动审批覆盖)") elif failed: print(" ❌ CI有失败项,跳过审批") elif pending: @@ -319,8 +399,9 @@ def main(): print(f" 处理PR数: {min(len(prs), args.max_prs)}") print(f" 自动审批: {approved_count} 个") print(f" 自动合并: {merged_count} 个") + print(f" AI审查阻止: {ai_blocked_count} 个") print(f" 跳过: {skipped_count} 个") - print(f" 模式: {'DRY-RUN' if dry_run else '正式执行'}") + print(" 模式: {'DRY-RUN' if dry_run else '正式执行'}") if __name__ == "__main__": -- 2.54.0