test: P3-1 第二十二波 TTS音频合并+流式服务单元测试 14个 #728
+163
-73
@@ -2,32 +2,34 @@
|
||||
"""
|
||||
PR自动扫描器:扫描所有open PR,对CI全绿的进行自动审批/合并
|
||||
作为短作业模式的兜底机制,每5分钟运行一次
|
||||
|
||||
新增:AI审查联动 - AI代码审查发现严重问题时,不自动审批
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
import urllib.request
|
||||
import urllib.error
|
||||
import time
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
import time
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
|
||||
|
||||
def api_request(token, repo, endpoint, method='GET', data=None):
|
||||
def api_request(token, repo, endpoint, method="GET", data=None):
|
||||
"""Gitea API请求"""
|
||||
url = f"https://git.xiaoxiajianji.com/api/v1/repos/{repo}/{endpoint}"
|
||||
headers = {
|
||||
"Authorization": f"token {token}",
|
||||
"Content-Type": "application/json"
|
||||
}
|
||||
headers = {"Authorization": f"token {token}", "Content-Type": "application/json"}
|
||||
body = json.dumps(data).encode() if data else None
|
||||
req = urllib.request.Request(url, data=body, headers=headers, method=method)
|
||||
|
||||
|
||||
# 跳过SSL验证
|
||||
import ssl
|
||||
|
||||
ctx = ssl.create_default_context()
|
||||
ctx.check_hostname = False
|
||||
ctx.verify_mode = ssl.CERT_NONE
|
||||
|
||||
|
||||
try:
|
||||
resp = urllib.request.urlopen(req, context=ctx)
|
||||
return json.loads(resp.read().decode()), resp.status
|
||||
@@ -35,13 +37,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 +64,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 +77,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 +86,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)
|
||||
@@ -113,65 +113,139 @@ 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(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": reason},
|
||||
)
|
||||
|
||||
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": reason},
|
||||
)
|
||||
|
||||
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": 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)"""
|
||||
# 等待几秒让状态同步
|
||||
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 +263,12 @@ 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()
|
||||
|
||||
|
||||
dry_run = args.dry_run.lower() == "true"
|
||||
|
||||
|
||||
# required contexts(与分支保护一致)
|
||||
REQUIRED_CONTEXTS_FULL = [
|
||||
"CI/CD Pipeline / Validate - Code Quality (pull_request)",
|
||||
@@ -213,39 +288,40 @@ 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]:
|
||||
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"):
|
||||
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,45 +330,58 @@ 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)
|
||||
|
||||
# === 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(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}")
|
||||
approved_count += 1
|
||||
else:
|
||||
print(f" ❌ 审批失败: {msg}")
|
||||
elif ai_has_critical:
|
||||
print(" 🚫 AI审查阻止审批(人工可手动审批覆盖)")
|
||||
elif failed:
|
||||
print(f" ❌ CI有失败项,跳过审批")
|
||||
print(" ❌ CI有失败项,跳过审批")
|
||||
elif pending:
|
||||
print(f" ⏳ CI仍在运行,跳过")
|
||||
|
||||
print(" ⏳ 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] 将自动合并")
|
||||
print(" 🎯 [DRY-RUN] 将自动合并")
|
||||
else:
|
||||
print(f" 🎯 执行自动合并...")
|
||||
print(" 🎯 执行自动合并...")
|
||||
ok, msg = merge_pr(args.token, args.repo, pr_num)
|
||||
if ok:
|
||||
print(f" ✅ 合并成功: {msg}")
|
||||
@@ -300,18 +389,19 @@ 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(f"\n=== 扫描结果 ===")
|
||||
print(" ⏳ 合并条件未满足: 无审批")
|
||||
|
||||
print("\n=== 扫描结果 ===")
|
||||
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__":
|
||||
|
||||
Executable
+244
@@ -0,0 +1,244 @@
|
||||
"""
|
||||
JWT + Password 委托层单元测试(第二十一波)
|
||||
|
||||
覆盖:
|
||||
- JWTHandler (create/verify/configure/get)
|
||||
- PasswordHandler (hash/verify/needs_rehash/validate_strength/configure/get)
|
||||
"""
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from jwt.exceptions import InvalidTokenError
|
||||
|
||||
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(InvalidTokenError):
|
||||
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(InvalidTokenError):
|
||||
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
|
||||
Executable
+320
@@ -0,0 +1,320 @@
|
||||
"""
|
||||
生成任务应用层用例单元测试(第十九波)
|
||||
|
||||
覆盖:
|
||||
- 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
|
||||
Executable
+576
@@ -0,0 +1,576 @@
|
||||
"""
|
||||
Job 应用层用例单元测试(第十八波)
|
||||
|
||||
覆盖:
|
||||
- CreateJobUseCase
|
||||
- SubmitJobUseCase
|
||||
- UpdateJobProgressUseCase
|
||||
- CompleteJobUseCase
|
||||
- FailJobUseCase
|
||||
- RetryJobUseCase
|
||||
- CancelJobUseCase
|
||||
- GetJobUseCase
|
||||
- ListJobsUseCase
|
||||
- GetJobStatisticsUseCase
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.jobs import (
|
||||
CancelJobUseCase,
|
||||
CompleteJobCommand,
|
||||
CompleteJobUseCase,
|
||||
CreateJobCommand,
|
||||
CreateJobUseCase,
|
||||
FailJobCommand,
|
||||
FailJobUseCase,
|
||||
GetJobStatisticsUseCase,
|
||||
GetJobUseCase,
|
||||
ListJobsUseCase,
|
||||
RetryJobUseCase,
|
||||
SubmitJobUseCase,
|
||||
UpdateJobProgressCommand,
|
||||
UpdateJobProgressUseCase,
|
||||
)
|
||||
from packages.domain.job import Job, JobStatus, JobType
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
def make_job(
|
||||
status=JobStatus.PENDING,
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
project_id="proj-1",
|
||||
**kwargs,
|
||||
):
|
||||
job = Job.create(
|
||||
project_id=project_id,
|
||||
job_type=job_type,
|
||||
**kwargs,
|
||||
)
|
||||
# 绕过状态机直接设置状态(测试构造用)
|
||||
if status != JobStatus.PENDING:
|
||||
object.__setattr__(job, "status", status)
|
||||
return job
|
||||
|
||||
|
||||
# ============================================================
|
||||
# CreateJobUseCase
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestCreateJobUseCase:
|
||||
"""CreateJobUseCase 创建任务"""
|
||||
|
||||
def test_create_success(self, mock_repo):
|
||||
"""正常创建任务"""
|
||||
mock_repo.create.side_effect = lambda j: j
|
||||
|
||||
cmd = CreateJobCommand(
|
||||
project_id="proj-1",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
payload={"key": "val"},
|
||||
source_id="src-1",
|
||||
created_by_user_id="user-1",
|
||||
max_retries=5,
|
||||
)
|
||||
uc = CreateJobUseCase(mock_repo)
|
||||
job = uc.execute(cmd)
|
||||
|
||||
assert job.project_id == "proj-1"
|
||||
assert job.job_type == JobType.VIDEO_COMPOSE
|
||||
assert job.payload == {"key": "val"}
|
||||
assert job.source_id == "src-1"
|
||||
assert job.created_by_user_id == "user-1"
|
||||
assert job.max_retries == 5
|
||||
assert job.status == JobStatus.PENDING
|
||||
assert job.progress == 0.0
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_create_default_values(self, mock_repo):
|
||||
"""默认参数"""
|
||||
mock_repo.create.side_effect = lambda j: j
|
||||
|
||||
cmd = CreateJobCommand(project_id="proj-1", job_type="video_compose")
|
||||
uc = CreateJobUseCase(mock_repo)
|
||||
job = uc.execute(cmd)
|
||||
|
||||
assert job.payload == {}
|
||||
assert job.source_id == ""
|
||||
assert job.created_by_user_id == ""
|
||||
assert job.max_retries == 3
|
||||
|
||||
def test_create_string_job_type(self, mock_repo):
|
||||
"""字符串类型的 job_type 也支持"""
|
||||
mock_repo.create.side_effect = lambda j: j
|
||||
|
||||
cmd = CreateJobCommand(project_id="proj-1", job_type="asset_ingest")
|
||||
uc = CreateJobUseCase(mock_repo)
|
||||
job = uc.execute(cmd)
|
||||
|
||||
assert job.job_type == JobType.ASSET_INGEST
|
||||
|
||||
|
||||
# ============================================================
|
||||
# SubmitJobUseCase
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestSubmitJobUseCase:
|
||||
"""SubmitJobUseCase 提交任务"""
|
||||
|
||||
def test_submit_success(self, mock_repo):
|
||||
"""正常提交 pending 任务"""
|
||||
job = make_job(status=JobStatus.PENDING)
|
||||
mock_repo.get.return_value = job
|
||||
mock_repo.update.side_effect = lambda j: j
|
||||
|
||||
uc = SubmitJobUseCase(mock_repo)
|
||||
result = uc.execute(job.id, celery_task_id="celery-123")
|
||||
|
||||
assert result.status == JobStatus.RUNNING
|
||||
assert result.celery_task_id == "celery-123"
|
||||
assert result.current_stage == "已提交,等待执行"
|
||||
assert result.started_at is not None
|
||||
mock_repo.update.assert_called_once()
|
||||
|
||||
def test_submit_job_not_found(self, mock_repo):
|
||||
"""任务不存在"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
uc = SubmitJobUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="任务不存在"):
|
||||
uc.execute("nonexistent")
|
||||
|
||||
def test_submit_already_running(self, mock_repo):
|
||||
"""已经是 running 状态不能再提交"""
|
||||
job = make_job(status=JobStatus.RUNNING)
|
||||
mock_repo.get.return_value = job
|
||||
|
||||
uc = SubmitJobUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="只有 pending 状态"):
|
||||
uc.execute(job.id)
|
||||
|
||||
def test_submit_without_celery_id(self, mock_repo):
|
||||
"""不传 celery_task_id 也可以"""
|
||||
job = make_job(status=JobStatus.PENDING)
|
||||
mock_repo.get.return_value = job
|
||||
mock_repo.update.side_effect = lambda j: j
|
||||
|
||||
uc = SubmitJobUseCase(mock_repo)
|
||||
result = uc.execute(job.id)
|
||||
|
||||
assert result.status == JobStatus.RUNNING
|
||||
assert result.celery_task_id == ""
|
||||
|
||||
|
||||
# ============================================================
|
||||
# UpdateJobProgressUseCase
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestUpdateJobProgressUseCase:
|
||||
"""UpdateJobProgressUseCase 更新进度"""
|
||||
|
||||
def test_update_progress_success(self, mock_repo):
|
||||
"""正常更新进度"""
|
||||
job = make_job(status=JobStatus.RUNNING)
|
||||
mock_repo.get.return_value = job
|
||||
mock_repo.update.side_effect = lambda j: j
|
||||
|
||||
cmd = UpdateJobProgressCommand(job_id=job.id, progress=50.0, current_stage="处理中")
|
||||
uc = UpdateJobProgressUseCase(mock_repo)
|
||||
result = uc.execute(cmd)
|
||||
|
||||
assert result.progress == 50.0
|
||||
assert result.current_stage == "处理中"
|
||||
mock_repo.update.assert_called_once()
|
||||
|
||||
def test_update_progress_job_not_found(self, mock_repo):
|
||||
"""任务不存在"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
cmd = UpdateJobProgressCommand(job_id="nope", progress=50.0)
|
||||
uc = UpdateJobProgressUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="任务不存在"):
|
||||
uc.execute(cmd)
|
||||
|
||||
def test_update_progress_not_running(self, mock_repo):
|
||||
"""非 running 状态不能更新进度"""
|
||||
job = make_job(status=JobStatus.PENDING)
|
||||
mock_repo.get.return_value = job
|
||||
|
||||
cmd = UpdateJobProgressCommand(job_id=job.id, progress=50.0)
|
||||
uc = UpdateJobProgressUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="只有 running 状态"):
|
||||
uc.execute(cmd)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# CompleteJobUseCase
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestCompleteJobUseCase:
|
||||
"""CompleteJobUseCase 完成任务"""
|
||||
|
||||
def test_complete_from_running(self, mock_repo):
|
||||
"""从 running 状态完成"""
|
||||
job = make_job(status=JobStatus.RUNNING)
|
||||
mock_repo.get.return_value = job
|
||||
mock_repo.update.side_effect = lambda j: j
|
||||
|
||||
cmd = CompleteJobCommand(job_id=job.id, result={"output": "ok"})
|
||||
uc = CompleteJobUseCase(mock_repo)
|
||||
result = uc.execute(cmd)
|
||||
|
||||
assert result.status == JobStatus.SUCCESS
|
||||
assert result.progress == 100.0
|
||||
assert result.result == {"output": "ok"}
|
||||
assert result.completed_at is not None
|
||||
mock_repo.update.assert_called_once()
|
||||
|
||||
def test_complete_from_pending(self, mock_repo):
|
||||
"""从 pending 状态也可以直接完成"""
|
||||
job = make_job(status=JobStatus.PENDING)
|
||||
mock_repo.get.return_value = job
|
||||
mock_repo.update.side_effect = lambda j: j
|
||||
|
||||
cmd = CompleteJobCommand(job_id=job.id)
|
||||
uc = CompleteJobUseCase(mock_repo)
|
||||
result = uc.execute(cmd)
|
||||
|
||||
assert result.status == JobStatus.SUCCESS
|
||||
assert result.progress == 100.0
|
||||
|
||||
def test_complete_job_not_found(self, mock_repo):
|
||||
"""任务不存在"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
cmd = CompleteJobCommand(job_id="nope")
|
||||
uc = CompleteJobUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="任务不存在"):
|
||||
uc.execute(cmd)
|
||||
|
||||
def test_complete_already_failed(self, mock_repo):
|
||||
"""已失败的任务不能直接标记完成"""
|
||||
job = make_job(status=JobStatus.FAILED)
|
||||
job.error_message = "some error"
|
||||
mock_repo.get.return_value = job
|
||||
|
||||
cmd = CompleteJobCommand(job_id=job.id)
|
||||
uc = CompleteJobUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="只有 running/pending"):
|
||||
uc.execute(cmd)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# FailJobUseCase
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestFailJobUseCase:
|
||||
"""FailJobUseCase 失败任务"""
|
||||
|
||||
def test_fail_from_running(self, mock_repo):
|
||||
"""从 running 状态失败"""
|
||||
job = make_job(status=JobStatus.RUNNING)
|
||||
mock_repo.get.return_value = job
|
||||
mock_repo.update.side_effect = lambda j: j
|
||||
|
||||
cmd = FailJobCommand(job_id=job.id, error_message="网络超时")
|
||||
uc = FailJobUseCase(mock_repo)
|
||||
result = uc.execute(cmd)
|
||||
|
||||
assert result.status == JobStatus.FAILED
|
||||
assert result.error_message == "网络超时"
|
||||
assert result.current_stage == "失败"
|
||||
assert result.completed_at is not None
|
||||
mock_repo.update.assert_called_once()
|
||||
|
||||
def test_fail_pending_rejected_by_domain(self, mock_repo):
|
||||
"""pending 状态不能直接失败(领域状态机约束)"""
|
||||
job = make_job(status=JobStatus.PENDING)
|
||||
mock_repo.get.return_value = job
|
||||
|
||||
cmd = FailJobCommand(job_id=job.id, error_message="资源不足")
|
||||
uc = FailJobUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="非法状态转换"):
|
||||
uc.execute(cmd)
|
||||
|
||||
def test_fail_job_not_found(self, mock_repo):
|
||||
"""任务不存在"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
cmd = FailJobCommand(job_id="nope", error_message="err")
|
||||
uc = FailJobUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="任务不存在"):
|
||||
uc.execute(cmd)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# RetryJobUseCase
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestRetryJobUseCase:
|
||||
"""RetryJobUseCase 重试任务"""
|
||||
|
||||
def test_retry_success(self, mock_repo):
|
||||
"""失败任务重试成功"""
|
||||
job = make_job(status=JobStatus.FAILED, max_retries=3)
|
||||
job.retry_count = 0
|
||||
job.error_message = "timeout"
|
||||
mock_repo.get.return_value = job
|
||||
mock_repo.update.side_effect = lambda j: j
|
||||
|
||||
uc = RetryJobUseCase(mock_repo)
|
||||
result = uc.execute(job.id)
|
||||
|
||||
assert result.status == JobStatus.PENDING
|
||||
assert result.retry_count == 1
|
||||
assert result.progress == 0.0
|
||||
assert result.error_message == ""
|
||||
assert result.started_at is None
|
||||
assert result.completed_at is None
|
||||
assert result.celery_task_id == ""
|
||||
mock_repo.update.assert_called_once()
|
||||
|
||||
def test_retry_job_not_found(self, mock_repo):
|
||||
"""任务不存在"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
uc = RetryJobUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="任务不存在"):
|
||||
uc.execute("nope")
|
||||
|
||||
def test_retry_exceeds_max_retries(self, mock_repo):
|
||||
"""超过最大重试次数不可重试"""
|
||||
job = make_job(status=JobStatus.FAILED, max_retries=3)
|
||||
job.retry_count = 3
|
||||
mock_repo.get.return_value = job
|
||||
|
||||
uc = RetryJobUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="不可重试"):
|
||||
uc.execute(job.id)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# CancelJobUseCase
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestCancelJobUseCase:
|
||||
"""CancelJobUseCase 取消任务"""
|
||||
|
||||
def test_cancel_pending(self, mock_repo):
|
||||
"""取消 pending 任务"""
|
||||
job = make_job(status=JobStatus.PENDING)
|
||||
mock_repo.get.return_value = job
|
||||
mock_repo.update.side_effect = lambda j: j
|
||||
|
||||
uc = CancelJobUseCase(mock_repo)
|
||||
result = uc.execute(job.id)
|
||||
|
||||
assert result.status == JobStatus.CANCELLED
|
||||
assert result.current_stage == "已取消"
|
||||
mock_repo.update.assert_called_once()
|
||||
|
||||
def test_cancel_running(self, mock_repo):
|
||||
"""取消 running 任务"""
|
||||
job = make_job(status=JobStatus.RUNNING)
|
||||
mock_repo.get.return_value = job
|
||||
mock_repo.update.side_effect = lambda j: j
|
||||
|
||||
uc = CancelJobUseCase(mock_repo)
|
||||
result = uc.execute(job.id)
|
||||
|
||||
assert result.status == JobStatus.CANCELLED
|
||||
|
||||
def test_cancel_job_not_found(self, mock_repo):
|
||||
"""任务不存在"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
uc = CancelJobUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="任务不存在"):
|
||||
uc.execute("nope")
|
||||
|
||||
def test_cancel_already_success(self, mock_repo):
|
||||
"""已成功的任务不能取消"""
|
||||
job = make_job(status=JobStatus.SUCCESS)
|
||||
mock_repo.get.return_value = job
|
||||
|
||||
uc = CancelJobUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="终态"):
|
||||
uc.execute(job.id)
|
||||
|
||||
def test_cancel_already_failed(self, mock_repo):
|
||||
"""已失败的任务不能取消(走重试)"""
|
||||
job = make_job(status=JobStatus.FAILED)
|
||||
job.error_message = "err"
|
||||
mock_repo.get.return_value = job
|
||||
|
||||
uc = CancelJobUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="终态"):
|
||||
uc.execute(job.id)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# GetJobUseCase
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestGetJobUseCase:
|
||||
"""GetJobUseCase 获取任务"""
|
||||
|
||||
def test_get_existing(self, mock_repo):
|
||||
"""获取存在的任务"""
|
||||
job = make_job()
|
||||
mock_repo.get.return_value = job
|
||||
|
||||
uc = GetJobUseCase(mock_repo)
|
||||
result = uc.execute(job.id)
|
||||
|
||||
assert result is job
|
||||
mock_repo.get.assert_called_once_with(job.id)
|
||||
|
||||
def test_get_not_found(self, mock_repo):
|
||||
"""获取不存在的任务返回 None"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
uc = GetJobUseCase(mock_repo)
|
||||
result = uc.execute("nope")
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
# ============================================================
|
||||
# ListJobsUseCase
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestListJobsUseCase:
|
||||
"""ListJobsUseCase 列出任务"""
|
||||
|
||||
def test_list_by_project(self, mock_repo):
|
||||
"""按项目列出"""
|
||||
jobs = [make_job(), make_job()]
|
||||
mock_repo.list_by_project.return_value = jobs
|
||||
|
||||
uc = ListJobsUseCase(mock_repo)
|
||||
result = uc.execute(project_id="proj-1")
|
||||
|
||||
assert len(result) == 2
|
||||
mock_repo.list_by_project.assert_called_once()
|
||||
|
||||
def test_list_by_project_with_filters(self, mock_repo):
|
||||
"""按项目 + 类型 + 状态过滤"""
|
||||
mock_repo.list_by_project.return_value = []
|
||||
|
||||
uc = ListJobsUseCase(mock_repo)
|
||||
uc.execute(
|
||||
project_id="proj-1",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
status=JobStatus.RUNNING,
|
||||
limit=20,
|
||||
offset=10,
|
||||
)
|
||||
|
||||
mock_repo.list_by_project.assert_called_once_with(
|
||||
"proj-1",
|
||||
job_type=JobType.VIDEO_COMPOSE,
|
||||
status=JobStatus.RUNNING,
|
||||
limit=20,
|
||||
offset=10,
|
||||
)
|
||||
|
||||
def test_list_by_user(self, mock_repo):
|
||||
"""按用户列出"""
|
||||
jobs = [make_job()]
|
||||
mock_repo.list_by_user.return_value = jobs
|
||||
|
||||
uc = ListJobsUseCase(mock_repo)
|
||||
result = uc.execute(user_id="user-1")
|
||||
|
||||
assert len(result) == 1
|
||||
mock_repo.list_by_user.assert_called_once()
|
||||
|
||||
def test_list_no_filter_raises(self, mock_repo):
|
||||
"""不指定 project_id 或 user_id 报错"""
|
||||
uc = ListJobsUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="必须指定"):
|
||||
uc.execute()
|
||||
|
||||
def test_list_project_takes_precedence(self, mock_repo):
|
||||
"""同时传 project_id 和 user_id,优先按项目查"""
|
||||
mock_repo.list_by_project.return_value = []
|
||||
|
||||
uc = ListJobsUseCase(mock_repo)
|
||||
uc.execute(project_id="proj-1", user_id="user-1")
|
||||
|
||||
mock_repo.list_by_project.assert_called_once()
|
||||
mock_repo.list_by_user.assert_not_called()
|
||||
|
||||
|
||||
# ============================================================
|
||||
# GetJobStatisticsUseCase
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestGetJobStatisticsUseCase:
|
||||
"""GetJobStatisticsUseCase 任务统计"""
|
||||
|
||||
def test_stats_counts(self, mock_repo):
|
||||
"""统计各状态数量"""
|
||||
mock_repo.count_by_project.side_effect = lambda pid, status=None: {
|
||||
None: 10, # total
|
||||
JobStatus.PENDING: 2,
|
||||
JobStatus.RUNNING: 3,
|
||||
JobStatus.SUCCESS: 4,
|
||||
JobStatus.FAILED: 1,
|
||||
}[status]
|
||||
|
||||
uc = GetJobStatisticsUseCase(mock_repo)
|
||||
stats = uc.execute("proj-1")
|
||||
|
||||
assert stats["project_id"] == "proj-1"
|
||||
assert stats["total"] == 10
|
||||
assert stats["pending"] == 2
|
||||
assert stats["running"] == 3
|
||||
assert stats["success"] == 4
|
||||
assert stats["failed"] == 1
|
||||
# 总共调用 5 次 count_by_project
|
||||
assert mock_repo.count_by_project.call_count == 5
|
||||
|
||||
|
||||
# ============================================================
|
||||
# Command 对象
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestCommandObjects:
|
||||
"""命令对象基本属性"""
|
||||
|
||||
def test_create_job_command_defaults(self):
|
||||
cmd = CreateJobCommand(project_id="p1", job_type=JobType.VIDEO_COMPOSE)
|
||||
assert cmd.payload == {}
|
||||
assert cmd.source_id == ""
|
||||
assert cmd.created_by_user_id == ""
|
||||
assert cmd.max_retries == 3
|
||||
|
||||
def test_update_progress_command_defaults(self):
|
||||
cmd = UpdateJobProgressCommand(job_id="j1", progress=50.0)
|
||||
assert cmd.current_stage == ""
|
||||
|
||||
def test_complete_job_command_defaults(self):
|
||||
cmd = CompleteJobCommand(job_id="j1")
|
||||
assert cmd.result == {}
|
||||
|
||||
def test_fail_job_command(self):
|
||||
cmd = FailJobCommand(job_id="j1", error_message="err")
|
||||
assert cmd.error_message == "err"
|
||||
Executable
+313
@@ -0,0 +1,313 @@
|
||||
"""
|
||||
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()
|
||||
Executable
+449
@@ -0,0 +1,449 @@
|
||||
"""
|
||||
验证码服务单元测试(第十七波)
|
||||
|
||||
覆盖:
|
||||
- VerificationCodeService.generate
|
||||
- VerificationCodeService.verify
|
||||
- 频控逻辑(冷却 + 每日上限)
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.auth.verification_code_service import (
|
||||
CODE_TYPE_EMAIL_BIND,
|
||||
CODE_TYPE_EMAIL_LOGIN,
|
||||
CODE_TYPE_PHONE_BIND,
|
||||
DAILY_LIMIT,
|
||||
DEFAULT_TTL_SECONDS,
|
||||
MAX_ATTEMPTS,
|
||||
RESEND_COOLDOWN_SECONDS,
|
||||
VerificationCodeService,
|
||||
)
|
||||
from packages.domain.verification_code import VerificationCode
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
"""mock 验证码仓储"""
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def service(mock_repo):
|
||||
"""验证码服务实例"""
|
||||
return VerificationCodeService(repo=mock_repo)
|
||||
|
||||
|
||||
def make_code(
|
||||
recipient="test@example.com",
|
||||
code_type=CODE_TYPE_EMAIL_BIND,
|
||||
code="123456",
|
||||
ttl=300,
|
||||
used=False,
|
||||
attempts=0,
|
||||
created_at=None,
|
||||
):
|
||||
"""构造一个验证码实体"""
|
||||
now = created_at or datetime.now(timezone.utc)
|
||||
return VerificationCode(
|
||||
id="test-code-id",
|
||||
recipient=recipient,
|
||||
code=code,
|
||||
code_type=code_type,
|
||||
expires_at=now + timedelta(seconds=ttl),
|
||||
used_at=now if used else None,
|
||||
attempts=attempts,
|
||||
created_at=now,
|
||||
)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# generate - 参数校验
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestGenerateParamValidation:
|
||||
"""generate 参数校验"""
|
||||
|
||||
def test_empty_recipient(self, service):
|
||||
"""空接收方"""
|
||||
code, err = service.generate("", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "不能为空" in err
|
||||
|
||||
def test_whitespace_recipient_stripped(self, service, mock_repo):
|
||||
"""前后空格会被 strip 掉,正常生成"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
code, err = service.generate(" test@example.com ", CODE_TYPE_EMAIL_BIND)
|
||||
assert err is None
|
||||
assert code is not None
|
||||
assert code.recipient == "test@example.com"
|
||||
|
||||
def test_invalid_code_type(self, service):
|
||||
"""无效验证码类型"""
|
||||
code, err = service.generate("test@example.com", "invalid_type")
|
||||
assert code is None
|
||||
assert "无效的验证码类型" in err
|
||||
|
||||
|
||||
# ============================================================
|
||||
# generate - 正常生成
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestGenerateNormal:
|
||||
"""generate 正常生成场景"""
|
||||
|
||||
def test_generate_success(self, service, mock_repo):
|
||||
"""正常生成验证码"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert err is None
|
||||
assert code is not None
|
||||
assert code.recipient == "test@example.com"
|
||||
assert code.code_type == CODE_TYPE_EMAIL_BIND
|
||||
assert len(code.code) == 6
|
||||
assert code.code.isdigit()
|
||||
assert not code.is_used
|
||||
assert not code.is_expired
|
||||
mock_repo.save.assert_called_once()
|
||||
|
||||
def test_custom_code(self, service, mock_repo):
|
||||
"""自定义验证码"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND, custom_code="888888")
|
||||
|
||||
assert err is None
|
||||
assert code.code == "888888"
|
||||
|
||||
def test_custom_ttl(self, service, mock_repo):
|
||||
"""自定义有效期"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND, ttl_seconds=60)
|
||||
|
||||
assert err is None
|
||||
# 过期时间 - 创建时间 ≈ 60 秒
|
||||
delta = (code.expires_at - code.created_at).total_seconds()
|
||||
assert delta == 60
|
||||
|
||||
def test_default_ttl_used_when_not_specified(self, service, mock_repo):
|
||||
"""未指定 ttl 时使用默认值"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert err is None
|
||||
delta = (code.expires_at - code.created_at).total_seconds()
|
||||
assert delta == DEFAULT_TTL_SECONDS
|
||||
|
||||
def test_phone_bind_type(self, service, mock_repo):
|
||||
"""手机号绑定类型也支持"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
|
||||
code, err = service.generate("13800138000", CODE_TYPE_PHONE_BIND)
|
||||
|
||||
assert err is None
|
||||
assert code.code_type == CODE_TYPE_PHONE_BIND
|
||||
|
||||
|
||||
# ============================================================
|
||||
# generate - 频控
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestGenerateRateLimit:
|
||||
"""generate 频控逻辑"""
|
||||
|
||||
def test_resend_cooldown_blocked(self, service, mock_repo):
|
||||
"""冷却期内发送被拒绝"""
|
||||
# 10 秒前刚发过一条
|
||||
recent = make_code(created_at=datetime.now(timezone.utc) - timedelta(seconds=10))
|
||||
mock_repo.find_latest.return_value = recent
|
||||
mock_repo.count_today.return_value = 1
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert code is None
|
||||
assert "发送太频繁" in err
|
||||
assert "秒后再试" in err
|
||||
# 等待时间应接近 50 秒(60-10)
|
||||
# 提取数字验证范围
|
||||
import re
|
||||
|
||||
match = re.search(r"(\d+)\s*秒", err)
|
||||
assert match
|
||||
wait = int(match.group(1))
|
||||
assert 45 <= wait <= 55
|
||||
|
||||
def test_resend_after_cooldown_ok(self, service, mock_repo):
|
||||
"""超过冷却期可以重发"""
|
||||
# 2 分钟前发的,已过冷却
|
||||
old = make_code(created_at=datetime.now(timezone.utc) - timedelta(seconds=120))
|
||||
mock_repo.find_latest.return_value = old
|
||||
mock_repo.count_today.return_value = 1
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert err is None
|
||||
assert code is not None
|
||||
|
||||
def test_daily_limit_reached(self, service, mock_repo):
|
||||
"""达到每日上限"""
|
||||
# 没有最近的(过了冷却),但今日已达上限
|
||||
old = make_code(created_at=datetime.now(timezone.utc) - timedelta(hours=2))
|
||||
mock_repo.find_latest.return_value = old
|
||||
mock_repo.count_today.return_value = DAILY_LIMIT
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert code is None
|
||||
assert "今日发送次数已达上限" in err
|
||||
|
||||
def test_daily_limit_not_reached(self, service, mock_repo):
|
||||
"""未达每日上限可以发"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = DAILY_LIMIT - 1
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert err is None
|
||||
assert code is not None
|
||||
|
||||
def test_no_history_first_time_ok(self, service, mock_repo):
|
||||
"""首次发送,无历史记录"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert err is None
|
||||
assert code is not None
|
||||
mock_repo.save.assert_called_once()
|
||||
|
||||
|
||||
# ============================================================
|
||||
# generate - 自定义频控参数
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestGenerateCustomRateLimitParams:
|
||||
"""自定义频控参数"""
|
||||
|
||||
def test_custom_cooldown(self, mock_repo):
|
||||
"""自定义冷却时间"""
|
||||
svc = VerificationCodeService(repo=mock_repo, resend_cooldown=300, daily_limit=5)
|
||||
# 60 秒前发的,默认冷却 60 秒就够了,但这里设了 300 秒
|
||||
recent = make_code(created_at=datetime.now(timezone.utc) - timedelta(seconds=60))
|
||||
mock_repo.find_latest.return_value = recent
|
||||
mock_repo.count_today.return_value = 1
|
||||
|
||||
code, err = svc.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert code is None
|
||||
assert "发送太频繁" in err
|
||||
|
||||
def test_custom_daily_limit(self, mock_repo):
|
||||
"""自定义每日上限"""
|
||||
svc = VerificationCodeService(repo=mock_repo, resend_cooldown=60, daily_limit=3)
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 3
|
||||
|
||||
code, err = svc.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert code is None
|
||||
assert "今日发送次数已达上限" in err
|
||||
|
||||
|
||||
# ============================================================
|
||||
# verify - 参数校验
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestVerifyParamValidation:
|
||||
"""verify 参数校验"""
|
||||
|
||||
def test_empty_recipient(self, service):
|
||||
"""空接收方"""
|
||||
ok, err = service.verify("", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
assert not ok
|
||||
assert "参数不完整" in err
|
||||
|
||||
def test_empty_code(self, service):
|
||||
"""空验证码"""
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "")
|
||||
assert not ok
|
||||
assert "参数不完整" in err
|
||||
|
||||
def test_whitespace_stripped(self, service, mock_repo):
|
||||
"""前后空格会被 strip"""
|
||||
code = make_code(code="123456")
|
||||
mock_repo.find_latest.return_value = code
|
||||
mock_repo.count_today.return_value = 0
|
||||
|
||||
ok, err = service.verify(" test@example.com ", CODE_TYPE_EMAIL_BIND, " 123456 ")
|
||||
|
||||
assert ok
|
||||
assert err is None
|
||||
|
||||
|
||||
# ============================================================
|
||||
# verify - 正常验证
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestVerifyNormal:
|
||||
"""verify 正常验证场景"""
|
||||
|
||||
def test_verify_success_consume(self, service, mock_repo):
|
||||
"""验证成功并消耗"""
|
||||
code = make_code(code="123456")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456", consume=True)
|
||||
|
||||
assert ok
|
||||
assert err is None
|
||||
assert code.is_used # 被标记为已使用
|
||||
# save 被调用了两次:一次 increment_attempts 后,一次 mark_used 后
|
||||
assert mock_repo.save.call_count >= 2
|
||||
|
||||
def test_verify_success_no_consume(self, service, mock_repo):
|
||||
"""验证成功但不消耗"""
|
||||
code = make_code(code="123456")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456", consume=False)
|
||||
|
||||
assert ok
|
||||
assert err is None
|
||||
assert not code.is_used # 未被标记
|
||||
|
||||
def test_verify_code_not_found(self, service, mock_repo):
|
||||
"""验证码不存在"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
|
||||
assert not ok
|
||||
assert "不存在或已过期" in err
|
||||
|
||||
def test_verify_wrong_code(self, service, mock_repo):
|
||||
"""验证码错误"""
|
||||
code = make_code(code="123456")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "999999")
|
||||
|
||||
assert not ok
|
||||
assert "验证码错误" in err
|
||||
# 尝试次数增加了
|
||||
assert code.attempts == 1
|
||||
|
||||
def test_verify_already_used(self, service, mock_repo):
|
||||
"""验证码已使用"""
|
||||
code = make_code(code="123456", used=True)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
|
||||
assert not ok
|
||||
assert "已使用" in err
|
||||
|
||||
def test_verify_expired(self, service, mock_repo):
|
||||
"""验证码已过期"""
|
||||
code = make_code(code="123456", ttl=-60) # 已过期 60 秒
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
|
||||
assert not ok
|
||||
assert "已过期" in err
|
||||
|
||||
def test_verify_attempts_exceeded(self, service, mock_repo):
|
||||
"""超过最大尝试次数"""
|
||||
code = make_code(code="123456", attempts=MAX_ATTEMPTS)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
|
||||
assert not ok
|
||||
assert "验证次数过多" in err
|
||||
# verify 里先 increment_attempts 再判断,所以这里 attempts 应该是 MAX_ATTEMPTS + 1
|
||||
assert code.attempts == MAX_ATTEMPTS + 1
|
||||
|
||||
def test_attempts_increment_on_wrong_code(self, service, mock_repo):
|
||||
"""错误验证码会增加尝试次数"""
|
||||
code = make_code(code="123456", attempts=0)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "000000")
|
||||
assert code.attempts == 1
|
||||
|
||||
service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "000001")
|
||||
assert code.attempts == 2
|
||||
|
||||
|
||||
# ============================================================
|
||||
# verify - 不同 code_type 互不干扰
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestVerifyCodeTypeIsolation:
|
||||
"""不同验证码类型互不干扰"""
|
||||
|
||||
def test_email_bind_vs_email_login(self, service, mock_repo):
|
||||
"""用 email_login 类型的验证码去验证 email_bind 应该失败"""
|
||||
code = make_code(code_type=CODE_TYPE_EMAIL_LOGIN, code="123456")
|
||||
mock_repo.find_latest.return_value = None # 按 email_bind 查不到
|
||||
|
||||
# find_latest 按 code_type 查询,传 email_bind 返回 None
|
||||
def side_effect(recipient, ct):
|
||||
if ct == CODE_TYPE_EMAIL_LOGIN:
|
||||
return code
|
||||
return None
|
||||
|
||||
mock_repo.find_latest.side_effect = side_effect
|
||||
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
assert not ok
|
||||
assert "不存在或已过期" in err
|
||||
|
||||
|
||||
# ============================================================
|
||||
# 常量值检查
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestConstants:
|
||||
"""常量默认值校验"""
|
||||
|
||||
def test_default_cooldown_60(self):
|
||||
assert RESEND_COOLDOWN_SECONDS == 60
|
||||
|
||||
def test_default_daily_limit_10(self):
|
||||
assert DAILY_LIMIT == 10
|
||||
|
||||
def test_default_max_attempts_5(self):
|
||||
assert MAX_ATTEMPTS == 5
|
||||
|
||||
def test_default_ttl_300(self):
|
||||
assert DEFAULT_TTL_SECONDS == 300
|
||||
|
||||
def test_valid_code_types_count(self):
|
||||
"""5 种验证码类型"""
|
||||
from packages.application.auth.verification_code_service import VALID_CODE_TYPES
|
||||
|
||||
assert len(VALID_CODE_TYPES) == 5
|
||||
Executable
+354
@@ -0,0 +1,354 @@
|
||||
"""
|
||||
微信 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)
|
||||
Reference in New Issue
Block a user