test: P3-1 第39波单元测试(audio_merger/jwt_service/verification_code/register_user)
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 12s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m10s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m6s
CI/CD Pipeline / Validate - Code Quality (pull_request) Failing after 1m48s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 41s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 30s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 59s
Preview Deploy / Deploy Preview Environment (pull_request) Failing after 25s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m14s
AI Code Review / AI Code Review (pull_request) Successful in 4m37s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 5m17s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 11m24s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 3m7s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m34s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 25s
CI/CD Pipeline / Production Browser E2E (pull_request) Failing after 1298h38m45s
CI/CD Pipeline / ACR Image Cleanup (pull_request) Failing after 1298h44m58s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 1298h45m0s
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 1298h45m2s
CI/CD Pipeline / Deploy Production (pull_request) Failing after 1298h45m4s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Failing after 1298h56m13s
CI/CD Pipeline / Build Production Web Image (pull_request) Failing after 1298h56m31s
CI/CD Pipeline / Build Production API Image (pull_request) Failing after 1298h56m33s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Failing after 1298h56m35s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Failing after 1299h10m14s
CI/CD Pipeline / Build Staging Web Image (pull_request) Failing after 1299h10m16s
CI/CD Pipeline / Build Staging API Image (pull_request) Failing after 1299h10m18s
CI/CD Pipeline / Build Production Worker Image (pull_request) Failing after 1299h28m56s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 12s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m10s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m6s
CI/CD Pipeline / Validate - Code Quality (pull_request) Failing after 1m48s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 41s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 30s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 59s
Preview Deploy / Deploy Preview Environment (pull_request) Failing after 25s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m14s
AI Code Review / AI Code Review (pull_request) Successful in 4m37s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 5m17s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 11m24s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 3m7s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m34s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 25s
CI/CD Pipeline / Production Browser E2E (pull_request) Failing after 1298h38m45s
CI/CD Pipeline / ACR Image Cleanup (pull_request) Failing after 1298h44m58s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 1298h45m0s
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 1298h45m2s
CI/CD Pipeline / Deploy Production (pull_request) Failing after 1298h45m4s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Failing after 1298h56m13s
CI/CD Pipeline / Build Production Web Image (pull_request) Failing after 1298h56m31s
CI/CD Pipeline / Build Production API Image (pull_request) Failing after 1298h56m33s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Failing after 1298h56m35s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Failing after 1299h10m14s
CI/CD Pipeline / Build Staging Web Image (pull_request) Failing after 1299h10m16s
CI/CD Pipeline / Build Staging API Image (pull_request) Failing after 1299h10m18s
CI/CD Pipeline / Build Production Worker Image (pull_request) Failing after 1299h28m56s
- test_audio_merger: 13个(FFmpeg音频合并器) - test_jwt_service: 30个(JWT配置+服务+类型校验) - test_verification_code_service: 31个(验证码生成/验证/频控+手机号邮箱校验) - test_register_user_use_case: 20个(用户注册+邮箱验证) - 合计+94个测试,全量4577 passed
This commit is contained in:
Executable
+207
@@ -0,0 +1,207 @@
|
||||
"""音频合并器单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.tts_job.audio_merger import AudioMergeError, AudioMerger
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_audio_dir():
|
||||
"""创建临时目录,放几个模拟音频文件"""
|
||||
tmpdir = tempfile.mkdtemp()
|
||||
files = []
|
||||
for i in range(3):
|
||||
fpath = os.path.join(tmpdir, f"part{i}.mp3")
|
||||
with open(fpath, "wb") as f:
|
||||
f.write(f"audio_data_{i}".encode() * 100)
|
||||
files.append(fpath)
|
||||
yield files
|
||||
import shutil
|
||||
shutil.rmtree(tmpdir, ignore_errors=True)
|
||||
|
||||
|
||||
class TestAudioMerger:
|
||||
"""AudioMerger 测试"""
|
||||
|
||||
def test_empty_list_raises_error(self):
|
||||
"""空列表抛出 AudioMergeError"""
|
||||
merger = AudioMerger()
|
||||
with pytest.raises(AudioMergeError, match="没有可合并的音频文件"):
|
||||
merger.merge([])
|
||||
|
||||
def test_single_file_returns_content(self, sample_audio_dir):
|
||||
"""单文件直接返回文件内容"""
|
||||
merger = AudioMerger()
|
||||
result = merger.merge([sample_audio_dir[0]])
|
||||
|
||||
with open(sample_audio_dir[0], "rb") as f:
|
||||
expected = f.read()
|
||||
|
||||
assert result == expected
|
||||
|
||||
def test_single_file_no_ffmpeg_needed(self, sample_audio_dir):
|
||||
"""单文件不需要调用 FFmpeg"""
|
||||
with patch("packages.application.tts_job.audio_merger.run_ffmpeg") as mock_ffmpeg:
|
||||
merger = AudioMerger()
|
||||
merger.merge([sample_audio_dir[0]])
|
||||
mock_ffmpeg.assert_not_called()
|
||||
|
||||
def test_merge_multiple_files(self, sample_audio_dir):
|
||||
"""多文件合并调用 FFmpeg"""
|
||||
with patch("packages.application.tts_job.audio_merger.run_ffmpeg") as mock_ffmpeg, \
|
||||
patch("packages.application.tts_job.audio_merger.FFMPEG_BIN", "ffmpeg"):
|
||||
# 模拟 FFmpeg 成功:在 output_path 写点数据
|
||||
def fake_run_ffmpeg(cmd, timeout=120):
|
||||
output_idx = cmd.index("-c") + 2 # -c copy 后面是 output_path
|
||||
output_path = cmd[-1]
|
||||
with open(output_path, "wb") as f:
|
||||
f.write(b"merged_audio_data")
|
||||
return MagicMock(stdout=b"", stderr=b"")
|
||||
|
||||
mock_ffmpeg.side_effect = fake_run_ffmpeg
|
||||
|
||||
merger = AudioMerger()
|
||||
result = merger.merge(sample_audio_dir)
|
||||
|
||||
assert result == b"merged_audio_data"
|
||||
mock_ffmpeg.assert_called_once()
|
||||
|
||||
def test_merge_concat_list_generated(self, sample_audio_dir):
|
||||
"""生成正确的 concat demuxer 列表文件"""
|
||||
import subprocess
|
||||
|
||||
with patch("packages.application.tts_job.audio_merger.run_ffmpeg") as mock_ffmpeg, \
|
||||
patch("packages.application.tts_job.audio_merger.FFMPEG_BIN", "ffmpeg"):
|
||||
|
||||
captured_list_content = []
|
||||
|
||||
def fake_run_ffmpeg(cmd, timeout=120):
|
||||
# 找到 -i 参数后面的文件路径
|
||||
# 命令结构: ffmpeg -y -f concat -safe 0 -i LIST_PATH -c copy OUTPUT
|
||||
for i, arg in enumerate(cmd):
|
||||
if arg == "-i" and i + 1 < len(cmd):
|
||||
list_path = cmd[i + 1]
|
||||
if list_path.endswith(".txt"):
|
||||
with open(list_path, "r") as f:
|
||||
captured_list_content.append(f.read())
|
||||
break
|
||||
# 写输出文件
|
||||
output_path = cmd[-1]
|
||||
with open(output_path, "wb") as f:
|
||||
f.write(b"fake")
|
||||
return MagicMock(stdout=b"", stderr=b"")
|
||||
|
||||
mock_ffmpeg.side_effect = fake_run_ffmpeg
|
||||
|
||||
merger = AudioMerger()
|
||||
merger.merge(sample_audio_dir, output_format="mp3")
|
||||
|
||||
# 检查列表文件包含所有输入文件
|
||||
assert len(captured_list_content) == 1
|
||||
list_content = captured_list_content[0]
|
||||
for fpath in sample_audio_dir:
|
||||
assert fpath in list_content.replace("'\\''", "'")
|
||||
|
||||
def test_merge_ffmpeg_failure_raises(self, sample_audio_dir):
|
||||
"""FFmpeg 失败抛出 AudioMergeError"""
|
||||
from subprocess import CalledProcessError
|
||||
|
||||
with patch("packages.application.tts_job.audio_merger.run_ffmpeg") as mock_ffmpeg, \
|
||||
patch("packages.application.tts_job.audio_merger.FFMPEG_BIN", "ffmpeg"):
|
||||
mock_ffmpeg.side_effect = CalledProcessError(
|
||||
returncode=1, cmd=["ffmpeg"], stderr=b"error message"
|
||||
)
|
||||
|
||||
merger = AudioMerger()
|
||||
with pytest.raises(AudioMergeError, match="FFmpeg 合并失败"):
|
||||
merger.merge(sample_audio_dir)
|
||||
|
||||
def test_merge_timeout_raises(self, sample_audio_dir):
|
||||
"""合并超时抛出 AudioMergeError"""
|
||||
from subprocess import TimeoutExpired
|
||||
|
||||
with patch("packages.application.tts_job.audio_merger.run_ffmpeg") as mock_ffmpeg, \
|
||||
patch("packages.application.tts_job.audio_merger.FFMPEG_BIN", "ffmpeg"):
|
||||
mock_ffmpeg.side_effect = TimeoutExpired(cmd=["ffmpeg"], timeout=120)
|
||||
|
||||
merger = AudioMerger()
|
||||
with pytest.raises(AudioMergeError, match="超时"):
|
||||
merger.merge(sample_audio_dir)
|
||||
|
||||
def test_merge_cleanup_temp_dir(self, sample_audio_dir):
|
||||
"""合并完成后清理临时目录"""
|
||||
with patch("packages.application.tts_job.audio_merger.run_ffmpeg") as mock_ffmpeg, \
|
||||
patch("packages.application.tts_job.audio_merger.FFMPEG_BIN", "ffmpeg"), \
|
||||
patch("packages.application.tts_job.audio_merger.shutil.rmtree") as mock_rmtree:
|
||||
|
||||
def fake_run_ffmpeg(cmd, timeout=120):
|
||||
output_path = cmd[-1]
|
||||
with open(output_path, "wb") as f:
|
||||
f.write(b"data")
|
||||
return MagicMock()
|
||||
|
||||
mock_ffmpeg.side_effect = fake_run_ffmpeg
|
||||
|
||||
merger = AudioMerger()
|
||||
merger.merge(sample_audio_dir)
|
||||
|
||||
mock_rmtree.assert_called_once()
|
||||
# 第一个参数是临时目录路径
|
||||
temp_dir_path = mock_rmtree.call_args[0][0]
|
||||
assert "tts_merge_" in temp_dir_path
|
||||
|
||||
def test_merge_cleanup_on_error(self, sample_audio_dir):
|
||||
"""合并失败也清理临时目录"""
|
||||
with patch("packages.application.tts_job.audio_merger.run_ffmpeg") as mock_ffmpeg, \
|
||||
patch("packages.application.tts_job.audio_merger.FFMPEG_BIN", "ffmpeg"), \
|
||||
patch("packages.application.tts_job.audio_merger.shutil.rmtree") as mock_rmtree:
|
||||
from subprocess import CalledProcessError
|
||||
mock_ffmpeg.side_effect = CalledProcessError(1, ["ffmpeg"])
|
||||
|
||||
merger = AudioMerger()
|
||||
try:
|
||||
merger.merge(sample_audio_dir)
|
||||
except AudioMergeError:
|
||||
pass
|
||||
|
||||
mock_rmtree.assert_called_once()
|
||||
|
||||
def test_merge_custom_output_format(self, sample_audio_dir):
|
||||
"""自定义输出格式"""
|
||||
with patch("packages.application.tts_job.audio_merger.run_ffmpeg") as mock_ffmpeg, \
|
||||
patch("packages.application.tts_job.audio_merger.FFMPEG_BIN", "ffmpeg"):
|
||||
|
||||
def fake_run_ffmpeg(cmd, timeout=120):
|
||||
output_path = cmd[-1]
|
||||
assert output_path.endswith(".wav")
|
||||
with open(output_path, "wb") as f:
|
||||
f.write(b"data")
|
||||
return MagicMock()
|
||||
|
||||
mock_ffmpeg.side_effect = fake_run_ffmpeg
|
||||
|
||||
merger = AudioMerger()
|
||||
merger.merge(sample_audio_dir, output_format="wav")
|
||||
|
||||
def test_merge_two_files(self, sample_audio_dir):
|
||||
"""两个文件合并"""
|
||||
with patch("packages.application.tts_job.audio_merger.run_ffmpeg") as mock_ffmpeg, \
|
||||
patch("packages.application.tts_job.audio_merger.FFMPEG_BIN", "ffmpeg"):
|
||||
|
||||
def fake_run_ffmpeg(cmd, timeout=120):
|
||||
output_path = cmd[-1]
|
||||
with open(output_path, "wb") as f:
|
||||
f.write(b"two_files_merged")
|
||||
return MagicMock()
|
||||
|
||||
mock_ffmpeg.side_effect = fake_run_ffmpeg
|
||||
|
||||
merger = AudioMerger()
|
||||
result = merger.merge(sample_audio_dir[:2])
|
||||
assert result == b"two_files_merged"
|
||||
+192
-290
@@ -1,362 +1,264 @@
|
||||
"""
|
||||
JWT Service 单元测试
|
||||
"""
|
||||
"""JWT 服务单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
import jwt
|
||||
import pytest
|
||||
from jwt.exceptions import ExpiredSignatureError, InvalidTokenError
|
||||
|
||||
from packages.application.auth.jwt_service import (
|
||||
JWTConfig,
|
||||
JWTService,
|
||||
TokenType,
|
||||
)
|
||||
from packages.application.auth.jwt_service import JWTConfig, JWTService, TokenType
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def jwt_config():
|
||||
return JWTConfig(
|
||||
secret_key="test-secret-key-strong-enough-123456",
|
||||
algorithm="HS256",
|
||||
access_token_expire_minutes=30,
|
||||
refresh_token_expire_days=7,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def jwt_service(jwt_config):
|
||||
return JWTService(jwt_config)
|
||||
|
||||
|
||||
class TestJWTConfig:
|
||||
"""JWT 配置测试"""
|
||||
"""JWTConfig 测试"""
|
||||
|
||||
def test_config_init_success(self):
|
||||
"""测试正常初始化"""
|
||||
config = JWTConfig(secret_key="a-very-strong-secret-key-for-testing")
|
||||
assert config.SECRET_KEY == "a-very-strong-secret-key-for-testing"
|
||||
assert config.ALGORITHM == "HS256"
|
||||
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 15
|
||||
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 7
|
||||
|
||||
def test_config_custom_values(self):
|
||||
"""测试自定义配置值"""
|
||||
config = JWTConfig(
|
||||
secret_key="test-secret",
|
||||
algorithm="HS512",
|
||||
access_token_expire_minutes=30,
|
||||
refresh_token_expire_days=14,
|
||||
)
|
||||
assert config.ALGORITHM == "HS512"
|
||||
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 30
|
||||
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 14
|
||||
|
||||
def test_config_empty_secret_raises(self):
|
||||
"""测试空密钥报错"""
|
||||
with pytest.raises(ValueError, match="secret_key must be provided"):
|
||||
def test_empty_secret_raises(self):
|
||||
"""空 secret_key 抛出 ValueError"""
|
||||
with pytest.raises(ValueError, match="must be provided"):
|
||||
JWTConfig(secret_key="")
|
||||
|
||||
def test_config_whitespace_secret_raises(self):
|
||||
"""测试全空格密钥报错"""
|
||||
with pytest.raises(ValueError, match="secret_key must be provided"):
|
||||
def test_whitespace_secret_raises(self):
|
||||
"""纯空白 secret_key 抛出 ValueError"""
|
||||
with pytest.raises(ValueError, match="must be provided"):
|
||||
JWTConfig(secret_key=" ")
|
||||
|
||||
def test_config_insecure_default_secret_raises(self):
|
||||
"""测试不安全的默认密钥报错"""
|
||||
insecure_keys = [
|
||||
def test_insecure_default_secret_raises(self):
|
||||
"""不安全的默认 secret 抛出 ValueError"""
|
||||
insecure_secrets = [
|
||||
"your-secret-key-change-in-production",
|
||||
"your-secret-key",
|
||||
"secret",
|
||||
"changeme",
|
||||
"password",
|
||||
"YOUR-SECRET-KEY",
|
||||
"Secret",
|
||||
"SECRET",
|
||||
]
|
||||
for key in insecure_keys:
|
||||
for secret in insecure_secrets:
|
||||
with pytest.raises(ValueError, match="insecure"):
|
||||
JWTConfig(secret_key=key)
|
||||
JWTConfig(secret_key=secret)
|
||||
|
||||
def test_strong_secret_accepted(self):
|
||||
"""强 secret 可以正常创建"""
|
||||
config = JWTConfig(secret_key="my-strong-secret-key-1234567890")
|
||||
assert config.SECRET_KEY == "my-strong-secret-key-1234567890"
|
||||
|
||||
class TestJWTService:
|
||||
"""JWT 服务测试"""
|
||||
def test_default_values(self):
|
||||
"""默认配置值正确"""
|
||||
config = JWTConfig(secret_key="test-secret-12345")
|
||||
assert config.ALGORITHM == "HS256"
|
||||
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 15
|
||||
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 7
|
||||
|
||||
@pytest.fixture
|
||||
def config(self):
|
||||
return JWTConfig(
|
||||
secret_key="test-secret-key-for-jwt-unit-tests-12345",
|
||||
algorithm="HS256",
|
||||
access_token_expire_minutes=30,
|
||||
refresh_token_expire_days=7,
|
||||
def test_custom_expiry_values(self):
|
||||
"""自定义过期时间"""
|
||||
config = JWTConfig(
|
||||
secret_key="test-secret-12345",
|
||||
access_token_expire_minutes=60,
|
||||
refresh_token_expire_days=30,
|
||||
)
|
||||
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 60
|
||||
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 30
|
||||
|
||||
@pytest.fixture
|
||||
def service(self, config):
|
||||
return JWTService(config=config)
|
||||
|
||||
def test_service_init_without_config_raises(self):
|
||||
"""测试无 config 初始化报错"""
|
||||
class TestTokenType:
|
||||
"""TokenType 测试"""
|
||||
|
||||
def test_access_token_type(self):
|
||||
"""access token 类型值"""
|
||||
assert TokenType.ACCESS == "access"
|
||||
|
||||
def test_refresh_token_type(self):
|
||||
"""refresh token 类型值"""
|
||||
assert TokenType.REFRESH == "refresh"
|
||||
|
||||
|
||||
class TestJWTServiceInit:
|
||||
"""JWTService 初始化测试"""
|
||||
|
||||
def test_none_config_raises(self):
|
||||
"""不传 config 抛出 ValueError"""
|
||||
with pytest.raises(ValueError, match="requires a JWTConfig"):
|
||||
JWTService(config=None)
|
||||
JWTService(None)
|
||||
|
||||
# --- create_access_token ---
|
||||
def test_with_config_creates_service(self, jwt_config):
|
||||
"""传入 config 正常创建"""
|
||||
service = JWTService(jwt_config)
|
||||
assert service.config is jwt_config
|
||||
|
||||
def test_create_access_token_success(self, service):
|
||||
"""测试创建 access token 成功"""
|
||||
token = service.create_access_token(user_id="user-123")
|
||||
|
||||
class TestCreateAccessToken:
|
||||
"""create_access_token 测试"""
|
||||
|
||||
def test_returns_string(self, jwt_service):
|
||||
"""返回非空字符串"""
|
||||
token = jwt_service.create_access_token(user_id="user_001")
|
||||
assert isinstance(token, str)
|
||||
assert len(token) > 0
|
||||
|
||||
def test_create_access_token_contains_user_id(self, service, config):
|
||||
"""测试 access token 包含正确的 user_id"""
|
||||
token = service.create_access_token(user_id="user-456")
|
||||
payload = jwt.decode(token, config.SECRET_KEY, algorithms=[config.ALGORITHM])
|
||||
assert payload["sub"] == "user-456"
|
||||
def test_contains_user_id(self, jwt_service):
|
||||
"""payload 包含正确的 user_id(sub字段)"""
|
||||
token = jwt_service.create_access_token(user_id="user_123")
|
||||
payload = jwt_service.verify_token(token)
|
||||
assert payload["sub"] == "user_123"
|
||||
|
||||
def test_create_access_token_has_correct_type(self, service, config):
|
||||
"""测试 access token 类型正确"""
|
||||
token = service.create_access_token(user_id="user-123")
|
||||
payload = jwt.decode(token, config.SECRET_KEY, algorithms=[config.ALGORITHM])
|
||||
assert payload["type"] == TokenType.ACCESS
|
||||
|
||||
def test_create_access_token_contains_role(self, service, config):
|
||||
"""测试 access token 包含角色"""
|
||||
token = service.create_access_token(user_id="user-123", role="admin")
|
||||
payload = jwt.decode(token, config.SECRET_KEY, algorithms=[config.ALGORITHM])
|
||||
def test_contains_role(self, jwt_service):
|
||||
"""payload 包含 role"""
|
||||
token = jwt_service.create_access_token(user_id="user_001", role="admin")
|
||||
payload = jwt_service.verify_token(token)
|
||||
assert payload["role"] == "admin"
|
||||
|
||||
def test_create_access_token_additional_claims(self, service, config):
|
||||
"""测试 access token 包含额外声明"""
|
||||
token = service.create_access_token(
|
||||
user_id="user-123",
|
||||
additional_claims={"custom_field": "custom_value", "sid": "session-abc"},
|
||||
def test_default_role_empty(self, jwt_service):
|
||||
"""不传 role 默认为空字符串"""
|
||||
token = jwt_service.create_access_token(user_id="user_001")
|
||||
payload = jwt_service.verify_token(token)
|
||||
assert payload["role"] == ""
|
||||
|
||||
def test_token_type_is_access(self, jwt_service):
|
||||
"""access token 的 type 为 access"""
|
||||
token = jwt_service.create_access_token(user_id="user_001")
|
||||
payload = jwt_service.verify_token(token)
|
||||
assert payload["type"] == TokenType.ACCESS
|
||||
|
||||
def test_additional_claims(self, jwt_service):
|
||||
"""额外声明被包含在 payload 中"""
|
||||
token = jwt_service.create_access_token(
|
||||
user_id="user_001",
|
||||
additional_claims={"email": "test@example.com", "tenant": "t1"},
|
||||
)
|
||||
payload = jwt.decode(token, config.SECRET_KEY, algorithms=[config.ALGORITHM])
|
||||
assert payload["custom_field"] == "custom_value"
|
||||
assert payload["sid"] == "session-abc"
|
||||
payload = jwt_service.verify_token(token)
|
||||
assert payload["email"] == "test@example.com"
|
||||
assert payload["tenant"] == "t1"
|
||||
|
||||
def test_create_access_token_has_iat_and_exp(self, service, config):
|
||||
"""测试 access token 包含 iat 和 exp"""
|
||||
before = datetime.utcnow() - timedelta(seconds=1)
|
||||
token = service.create_access_token(user_id="user-123")
|
||||
after = datetime.utcnow() + timedelta(seconds=1)
|
||||
|
||||
payload = jwt.decode(token, config.SECRET_KEY, algorithms=[config.ALGORITHM])
|
||||
def test_has_iat_and_exp(self, jwt_service):
|
||||
"""payload 包含 iat 和 exp"""
|
||||
token = jwt_service.create_access_token(user_id="user_001")
|
||||
payload = jwt_service.verify_token(token)
|
||||
assert "iat" in payload
|
||||
assert "exp" in payload
|
||||
assert payload["exp"] > payload["iat"]
|
||||
|
||||
iat = datetime.utcfromtimestamp(payload["iat"])
|
||||
exp = datetime.utcfromtimestamp(payload["exp"])
|
||||
def test_expiry_correct_duration(self, jwt_service):
|
||||
"""过期时间设置正确"""
|
||||
token = jwt_service.create_access_token(user_id="user_001")
|
||||
payload = jwt_service.verify_token(token)
|
||||
# 30分钟 = 1800秒
|
||||
duration = payload["exp"] - payload["iat"]
|
||||
assert 1790 <= duration <= 1810 # 允许10秒误差
|
||||
|
||||
assert before <= iat <= after
|
||||
assert exp > iat
|
||||
# 过期时间约等于配置的分钟数
|
||||
expected_expiry = timedelta(minutes=config.ACCESS_TOKEN_EXPIRE_MINUTES)
|
||||
actual_expiry = exp - iat
|
||||
assert abs((actual_expiry - expected_expiry).total_seconds()) < 5
|
||||
|
||||
# --- create_refresh_token ---
|
||||
class TestCreateRefreshToken:
|
||||
"""create_refresh_token 测试"""
|
||||
|
||||
def test_create_refresh_token_success(self, service):
|
||||
"""测试创建 refresh token 成功"""
|
||||
token = service.create_refresh_token(user_id="user-123", session_id="sess-abc")
|
||||
def test_returns_string(self, jwt_service):
|
||||
"""返回非空字符串"""
|
||||
token = jwt_service.create_refresh_token(user_id="user_001", session_id="sess_001")
|
||||
assert isinstance(token, str)
|
||||
assert len(token) > 0
|
||||
|
||||
def test_create_refresh_token_contains_correct_data(self, service, config):
|
||||
"""测试 refresh token 包含正确数据"""
|
||||
token = service.create_refresh_token(user_id="user-789", session_id="sess-xyz")
|
||||
payload = jwt.decode(token, config.SECRET_KEY, algorithms=[config.ALGORITHM])
|
||||
assert payload["sub"] == "user-789"
|
||||
assert payload["session_id"] == "sess-xyz"
|
||||
def test_contains_user_and_session(self, jwt_service):
|
||||
"""包含 user_id 和 session_id"""
|
||||
token = jwt_service.create_refresh_token(user_id="user_123", session_id="sess_456")
|
||||
payload = jwt_service.verify_token(token)
|
||||
assert payload["sub"] == "user_123"
|
||||
assert payload["session_id"] == "sess_456"
|
||||
|
||||
def test_token_type_is_refresh(self, jwt_service):
|
||||
"""refresh token 的 type 为 refresh"""
|
||||
token = jwt_service.create_refresh_token(user_id="user_001", session_id="s1")
|
||||
payload = jwt_service.verify_token(token)
|
||||
assert payload["type"] == TokenType.REFRESH
|
||||
|
||||
def test_create_refresh_token_expiry(self, service, config):
|
||||
"""测试 refresh token 过期时间正确"""
|
||||
token = service.create_refresh_token(user_id="user-123", session_id="sess-abc")
|
||||
payload = jwt.decode(token, config.SECRET_KEY, algorithms=[config.ALGORITHM])
|
||||
|
||||
iat = datetime.utcfromtimestamp(payload["iat"])
|
||||
exp = datetime.utcfromtimestamp(payload["exp"])
|
||||
expected_expiry = timedelta(days=config.REFRESH_TOKEN_EXPIRE_DAYS)
|
||||
actual_expiry = exp - iat
|
||||
assert abs((actual_expiry - expected_expiry).total_seconds()) < 5
|
||||
class TestVerifyToken:
|
||||
"""verify_token 测试"""
|
||||
|
||||
# --- verify_token ---
|
||||
def test_valid_token(self, jwt_service):
|
||||
"""有效 token 验证通过"""
|
||||
token = jwt_service.create_access_token(user_id="u1")
|
||||
payload = jwt_service.verify_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
|
||||
def test_verify_valid_token(self, service):
|
||||
"""测试验证有效 token"""
|
||||
token = service.create_access_token(user_id="user-123")
|
||||
payload = service.verify_token(token)
|
||||
assert payload["sub"] == "user-123"
|
||||
|
||||
def test_verify_expired_token_raises(self, service, config):
|
||||
"""测试验证过期 token 报错"""
|
||||
# 创建一个已经过期的 token
|
||||
payload = {
|
||||
"sub": "user-123",
|
||||
"type": TokenType.ACCESS,
|
||||
"iat": datetime.utcnow() - timedelta(hours=1),
|
||||
"exp": datetime.utcnow() - timedelta(minutes=30),
|
||||
}
|
||||
expired_token = jwt.encode(payload, config.SECRET_KEY, algorithm=config.ALGORITHM)
|
||||
|
||||
with pytest.raises(ExpiredSignatureError, match="expired"):
|
||||
service.verify_token(expired_token)
|
||||
|
||||
def test_verify_invalid_token_raises(self, service):
|
||||
"""测试验证无效 token 报错"""
|
||||
def test_invalid_token_raises(self, jwt_service):
|
||||
"""无效 token 抛出 InvalidTokenError"""
|
||||
with pytest.raises(InvalidTokenError):
|
||||
service.verify_token("this-is-not-a-valid-jwt-token")
|
||||
|
||||
def test_verify_token_with_wrong_secret_raises(self, service, config):
|
||||
"""测试用错误密钥签发的 token 验证失败"""
|
||||
wrong_config = JWTConfig(secret_key="different-secret-key")
|
||||
wrong_service = JWTService(config=wrong_config)
|
||||
token = wrong_service.create_access_token(user_id="user-123")
|
||||
jwt_service.verify_token("not.a.valid.token")
|
||||
|
||||
def test_empty_token_raises(self, jwt_service):
|
||||
"""空字符串 token 抛出异常"""
|
||||
with pytest.raises(InvalidTokenError):
|
||||
service.verify_token(token)
|
||||
jwt_service.verify_token("")
|
||||
|
||||
# --- verify_access_token ---
|
||||
def test_wrong_secret_fails(self, jwt_config):
|
||||
"""不同密钥的 token 无法验证"""
|
||||
service1 = JWTService(JWTConfig(secret_key="secret-one-123456"))
|
||||
service2 = JWTService(JWTConfig(secret_key="secret-two-1234567"))
|
||||
|
||||
def test_verify_access_token_success(self, service):
|
||||
"""测试验证有效的 access token"""
|
||||
token = service.create_access_token(user_id="user-123", role="user")
|
||||
payload = service.verify_access_token(token)
|
||||
assert payload["sub"] == "user-123"
|
||||
assert payload["type"] == TokenType.ACCESS
|
||||
token = service1.create_access_token(user_id="u1")
|
||||
with pytest.raises(InvalidTokenError):
|
||||
service2.verify_token(token)
|
||||
|
||||
def test_verify_access_token_with_refresh_token_raises(self, service):
|
||||
"""测试用 refresh token 调用 verify_access_token 报错"""
|
||||
refresh_token = service.create_refresh_token(user_id="user-123", session_id="sess-abc")
|
||||
|
||||
class TestVerifyAccessToken:
|
||||
"""verify_access_token 测试"""
|
||||
|
||||
def test_valid_access_token(self, jwt_service):
|
||||
"""有效 access token 验证通过"""
|
||||
token = jwt_service.create_access_token(user_id="u1")
|
||||
payload = jwt_service.verify_access_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
|
||||
def test_refresh_token_fails(self, jwt_service):
|
||||
"""refresh token 不能当 access token 用"""
|
||||
token = jwt_service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
with pytest.raises(ValueError, match="Token type must be 'access'"):
|
||||
service.verify_access_token(refresh_token)
|
||||
jwt_service.verify_access_token(token)
|
||||
|
||||
# --- verify_refresh_token ---
|
||||
|
||||
def test_verify_refresh_token_success(self, service):
|
||||
"""测试验证有效的 refresh token"""
|
||||
token = service.create_refresh_token(user_id="user-123", session_id="sess-abc")
|
||||
payload = service.verify_refresh_token(token)
|
||||
assert payload["sub"] == "user-123"
|
||||
assert payload["session_id"] == "sess-abc"
|
||||
assert payload["type"] == TokenType.REFRESH
|
||||
class TestVerifyRefreshToken:
|
||||
"""verify_refresh_token 测试"""
|
||||
|
||||
def test_verify_refresh_token_with_access_token_raises(self, service):
|
||||
"""测试用 access token 调用 verify_refresh_token 报错"""
|
||||
access_token = service.create_access_token(user_id="user-123")
|
||||
def test_valid_refresh_token(self, jwt_service):
|
||||
"""有效 refresh token 验证通过"""
|
||||
token = jwt_service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
payload = jwt_service.verify_refresh_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
assert payload["session_id"] == "s1"
|
||||
|
||||
def test_access_token_fails(self, jwt_service):
|
||||
"""access token 不能当 refresh token 用"""
|
||||
token = jwt_service.create_access_token(user_id="u1")
|
||||
with pytest.raises(ValueError, match="Token type must be 'refresh'"):
|
||||
service.verify_refresh_token(access_token)
|
||||
|
||||
def test_access_and_refresh_tokens_are_different(self, service):
|
||||
"""测试 access token 和 refresh token 不相同"""
|
||||
access = service.create_access_token(user_id="user-123")
|
||||
refresh = service.create_refresh_token(user_id="user-123", session_id="sess-abc")
|
||||
assert access != refresh
|
||||
jwt_service.verify_refresh_token(token)
|
||||
|
||||
|
||||
class TestJWTHandler:
|
||||
"""JWT Handler 委托层测试"""
|
||||
class TestExpiredToken:
|
||||
"""过期 token 测试"""
|
||||
|
||||
def test_create_access_token(self):
|
||||
"""测试创建 access token"""
|
||||
from packages.application.auth.jwt_handler import JWTHandler
|
||||
|
||||
handler = JWTHandler(secret_key="test-secret-key")
|
||||
token = handler.create_access_token(user_id="user-123", role="admin")
|
||||
assert token is not None
|
||||
assert len(token) > 20
|
||||
|
||||
# 验证token内容
|
||||
payload = jwt.decode(token, "test-secret-key", algorithms=["HS256"])
|
||||
assert payload["sub"] == "user-123"
|
||||
assert payload["role"] == "admin"
|
||||
assert payload["type"] == "access"
|
||||
|
||||
def test_create_access_token_with_additional_claims(self):
|
||||
"""测试带额外声明创建 token"""
|
||||
from packages.application.auth.jwt_handler import JWTHandler
|
||||
|
||||
handler = JWTHandler(secret_key="test-secret-key")
|
||||
token = handler.create_access_token(
|
||||
user_id="user-456",
|
||||
additional_claims={"custom_field": "custom_value"},
|
||||
def test_expired_access_token_raises(self):
|
||||
"""过期 token 验证抛出 ExpiredSignatureError"""
|
||||
config = JWTConfig(
|
||||
secret_key="test-secret-12345",
|
||||
access_token_expire_minutes=-1, # 立即过期
|
||||
)
|
||||
service = JWTService(config)
|
||||
token = service.create_access_token(user_id="u1")
|
||||
|
||||
payload = jwt.decode(token, "test-secret-key", algorithms=["HS256"])
|
||||
assert payload["sub"] == "user-456"
|
||||
assert payload["custom_field"] == "custom_value"
|
||||
|
||||
def test_verify_access_token(self):
|
||||
"""测试验证 access token"""
|
||||
from packages.application.auth.jwt_handler import JWTHandler
|
||||
|
||||
handler = JWTHandler(secret_key="test-secret-key")
|
||||
token = handler.create_access_token(user_id="user-123", role="user")
|
||||
payload = handler.verify_access_token(token)
|
||||
|
||||
assert payload["sub"] == "user-123"
|
||||
assert payload["role"] == "user"
|
||||
assert payload["type"] == "access"
|
||||
|
||||
def test_verify_access_token_expired(self):
|
||||
"""测试验证过期的 access token"""
|
||||
from packages.application.auth.jwt_handler import JWTHandler
|
||||
|
||||
handler = JWTHandler(secret_key="test-secret-key", access_token_expire_minutes=0)
|
||||
token = handler.create_access_token(user_id="user-123")
|
||||
|
||||
time.sleep(1) # 确保过期
|
||||
time.sleep(0.1)
|
||||
|
||||
with pytest.raises(ExpiredSignatureError):
|
||||
handler.verify_access_token(token)
|
||||
|
||||
def test_verify_token(self):
|
||||
"""测试验证任意类型 token"""
|
||||
from packages.application.auth.jwt_handler import JWTHandler
|
||||
|
||||
handler = JWTHandler(secret_key="test-secret-key")
|
||||
token = handler.create_access_token(user_id="user-123")
|
||||
payload = handler.verify_token(token)
|
||||
|
||||
assert payload["sub"] == "user-123"
|
||||
|
||||
def test_verify_invalid_token(self):
|
||||
"""测试验证无效 token"""
|
||||
from packages.application.auth.jwt_handler import JWTHandler
|
||||
|
||||
handler = JWTHandler(secret_key="test-secret-key")
|
||||
with pytest.raises(InvalidTokenError):
|
||||
handler.verify_token("invalid.token.here")
|
||||
|
||||
def test_configure_and_get_default_handler(self):
|
||||
"""测试配置和获取全局默认 handler"""
|
||||
from packages.application.auth import jwt_handler as handler_module
|
||||
from packages.application.auth.jwt_handler import (
|
||||
configure_jwt_handler,
|
||||
get_jwt_handler,
|
||||
)
|
||||
|
||||
# 重置全局状态
|
||||
handler_module._default_handler = None
|
||||
|
||||
# 配置
|
||||
handler = configure_jwt_handler(
|
||||
secret_key="global-secret",
|
||||
algorithm="HS256",
|
||||
access_token_expire_minutes=60,
|
||||
)
|
||||
assert handler is not None
|
||||
|
||||
# 获取
|
||||
same_handler = get_jwt_handler()
|
||||
assert same_handler is handler
|
||||
|
||||
# 验证能正常工作
|
||||
token = same_handler.create_access_token(user_id="global-user")
|
||||
payload = jwt.decode(token, "global-secret", algorithms=["HS256"])
|
||||
assert payload["sub"] == "global-user"
|
||||
|
||||
# 重置全局状态,避免影响其他测试
|
||||
handler_module._default_handler = None
|
||||
|
||||
def test_get_jwt_handler_not_configured(self):
|
||||
"""测试未配置时获取 handler 抛出异常"""
|
||||
from packages.application.auth import jwt_handler as handler_module
|
||||
from packages.application.auth.jwt_handler import get_jwt_handler
|
||||
|
||||
# 确保未配置
|
||||
handler_module._default_handler = None
|
||||
|
||||
with pytest.raises(RuntimeError, match="JWT handler not configured"):
|
||||
get_jwt_handler()
|
||||
service.verify_access_token(token)
|
||||
|
||||
Regular → Executable
+350
-147
@@ -1,12 +1,12 @@
|
||||
"""
|
||||
用户注册 Use Case 测试
|
||||
"""
|
||||
"""用户注册 UseCase 单元测试."""
|
||||
|
||||
from unittest.mock import Mock
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.auth import (
|
||||
from packages.application.auth.register_user_use_case import (
|
||||
RegisterUserRequest,
|
||||
RegisterUserUseCase,
|
||||
VerifyEmailRequest,
|
||||
@@ -15,213 +15,416 @@ from packages.application.auth import (
|
||||
from packages.domain.entities import User
|
||||
|
||||
|
||||
class TestRegisterUserUseCase:
|
||||
"""注册用例测试"""
|
||||
@pytest.fixture
|
||||
def mock_user_repo():
|
||||
return MagicMock()
|
||||
|
||||
@pytest.fixture
|
||||
def mock_user_repo(self):
|
||||
"""Mock 用户仓储"""
|
||||
repo = Mock()
|
||||
repo.find_by_email = Mock(return_value=None)
|
||||
repo.find_by_username = Mock(return_value=None)
|
||||
repo.find_by_verification_token = Mock(return_value=None)
|
||||
repo.save = Mock()
|
||||
return repo
|
||||
|
||||
@pytest.fixture
|
||||
def use_case(self, mock_user_repo):
|
||||
"""创建注册用例"""
|
||||
email_service = Mock()
|
||||
email_service.send_verification_email.return_value = (True, None)
|
||||
return RegisterUserUseCase(
|
||||
user_repository=mock_user_repo,
|
||||
base_url="https://test.com",
|
||||
email_service=email_service,
|
||||
)
|
||||
@pytest.fixture
|
||||
def mock_email_service():
|
||||
svc = MagicMock()
|
||||
svc.send_verification_email.return_value = (True, None)
|
||||
return svc
|
||||
|
||||
def test_register_user_success(self, use_case, mock_user_repo):
|
||||
"""测试注册成功"""
|
||||
|
||||
request = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="SecurePass123",
|
||||
@pytest.fixture
|
||||
def sample_user():
|
||||
user = User(
|
||||
id="user_001",
|
||||
email="test@example.com",
|
||||
username="testuser",
|
||||
display_name="测试用户",
|
||||
password_hash="hashed_pw",
|
||||
)
|
||||
user.email_verified = False
|
||||
user.email_verification_token = "some_token"
|
||||
return user
|
||||
|
||||
|
||||
class TestRegisterUserRequest:
|
||||
"""RegisterUserRequest 测试"""
|
||||
|
||||
def test_email_lowercased_stripped(self):
|
||||
"""邮箱转小写并去空格"""
|
||||
req = RegisterUserRequest(
|
||||
email=" Test@Example.COM ",
|
||||
password="TestPass1!",
|
||||
username="testuser",
|
||||
display_name="Test User",
|
||||
display_name="测试用户",
|
||||
)
|
||||
assert req.email == "test@example.com"
|
||||
|
||||
def test_username_stripped(self):
|
||||
"""用户名去空格"""
|
||||
req = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="TestPass1!",
|
||||
username=" testuser ",
|
||||
display_name="测试用户",
|
||||
)
|
||||
assert req.username == "testuser"
|
||||
|
||||
def test_display_name_stripped(self):
|
||||
"""显示名去空格"""
|
||||
req = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="TestPass1!",
|
||||
username="testuser",
|
||||
display_name=" 测试用户 ",
|
||||
)
|
||||
assert req.display_name == "测试用户"
|
||||
|
||||
|
||||
class TestRegisterUserUseCase:
|
||||
"""RegisterUserUseCase 测试"""
|
||||
|
||||
def test_register_success(self, mock_user_repo, mock_email_service):
|
||||
"""注册成功"""
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = None
|
||||
mock_user_repo.save.return_value = None
|
||||
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RegisterUserRequest(
|
||||
email="newuser@example.com",
|
||||
password="StrongPass1!",
|
||||
username="newuser",
|
||||
display_name="新用户",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.email == "test@example.com"
|
||||
assert response.username == "testuser"
|
||||
assert response.display_name == "Test User"
|
||||
assert response.email == "newuser@example.com"
|
||||
assert response.username == "newuser"
|
||||
assert response.display_name == "新用户"
|
||||
assert response.email_verification_sent is True
|
||||
|
||||
# 验证保存了用户
|
||||
assert response.user_id is not None
|
||||
mock_user_repo.save.assert_called_once()
|
||||
saved_user = mock_user_repo.save.call_args[0][0]
|
||||
assert saved_user.email == "test@example.com"
|
||||
assert saved_user.password_hash != ""
|
||||
assert saved_user.email_verified is False
|
||||
assert saved_user.email_verification_token is not None
|
||||
mock_email_service.send_verification_email.assert_called_once()
|
||||
|
||||
def test_register_user_weak_password(self, use_case):
|
||||
"""测试弱密码"""
|
||||
def test_register_empty_email(self, mock_user_repo, mock_email_service):
|
||||
"""空邮箱返回错误"""
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RegisterUserRequest(
|
||||
email="",
|
||||
password="TestPass1!",
|
||||
username="testuser",
|
||||
display_name="测试",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "Email is required" in error
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_register_empty_username(self, mock_user_repo, mock_email_service):
|
||||
"""空用户名返回错误"""
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="TestPass1!",
|
||||
username="",
|
||||
display_name="测试",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "Username is required" in error
|
||||
|
||||
def test_register_empty_display_name(self, mock_user_repo, mock_email_service):
|
||||
"""空显示名返回错误"""
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="TestPass1!",
|
||||
username="testuser",
|
||||
display_name="",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "Display name is required" in error
|
||||
|
||||
def test_register_weak_password(self, mock_user_repo, mock_email_service):
|
||||
"""弱密码返回错误"""
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = None
|
||||
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="weak",
|
||||
username="testuser",
|
||||
display_name="Test User",
|
||||
display_name="测试",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert error is not None
|
||||
assert "at least 8 characters" in error
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_register_user_email_exists(self, use_case, mock_user_repo):
|
||||
"""测试邮箱已存在"""
|
||||
# Mock 返回已存在的用户
|
||||
existing_user = User(
|
||||
id="existing-id",
|
||||
email="test@example.com",
|
||||
username="existing",
|
||||
display_name="Existing",
|
||||
def test_register_email_already_exists(self, mock_user_repo, mock_email_service, sample_user):
|
||||
"""邮箱已被注册"""
|
||||
mock_user_repo.find_by_email.return_value = sample_user
|
||||
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
mock_user_repo.find_by_email.return_value = existing_user
|
||||
|
||||
request = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="SecurePass123",
|
||||
password="TestPass1!",
|
||||
username="testuser",
|
||||
display_name="Test User",
|
||||
display_name="测试",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert error == "Email already registered"
|
||||
assert "Email already registered" in error
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_register_user_username_taken(self, use_case, mock_user_repo):
|
||||
"""测试用户名已被占用"""
|
||||
existing_user = User(
|
||||
id="existing-id",
|
||||
email="other@example.com",
|
||||
username="testuser",
|
||||
display_name="Other",
|
||||
def test_register_username_already_taken(self, mock_user_repo, mock_email_service, sample_user):
|
||||
"""用户名已被占用"""
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = sample_user
|
||||
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
mock_user_repo.find_by_username.return_value = existing_user
|
||||
|
||||
request = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="SecurePass123",
|
||||
username="testuser",
|
||||
display_name="Test User",
|
||||
email="new@example.com",
|
||||
password="TestPass1!",
|
||||
username="existinguser",
|
||||
display_name="测试",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert error == "Username already taken"
|
||||
assert "Username already taken" in error
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_register_user_missing_email(self, use_case):
|
||||
"""测试缺少邮箱"""
|
||||
request = RegisterUserRequest(
|
||||
email="",
|
||||
password="SecurePass123",
|
||||
username="testuser",
|
||||
display_name="Test User",
|
||||
def test_register_password_is_hashed(self, mock_user_repo, mock_email_service):
|
||||
"""用户密码被哈希存储,不是明文"""
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = None
|
||||
saved_user = None
|
||||
|
||||
def capture_save(user):
|
||||
nonlocal saved_user
|
||||
saved_user = user
|
||||
|
||||
mock_user_repo.save.side_effect = capture_save
|
||||
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert error == "Email is required"
|
||||
|
||||
def test_register_user_email_send_failure(self, use_case, mock_user_repo):
|
||||
"""测试邮件发送失败(用户仍然创建)"""
|
||||
use_case.email_service.send_verification_email.return_value = (
|
||||
False,
|
||||
"SMTP error",
|
||||
)
|
||||
|
||||
request = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="SecurePass123",
|
||||
password="MySecretPass1!",
|
||||
username="testuser",
|
||||
display_name="Test User",
|
||||
display_name="测试",
|
||||
)
|
||||
use_case.execute(request)
|
||||
|
||||
assert saved_user is not None
|
||||
assert saved_user.password_hash != "MySecretPass1!"
|
||||
assert len(saved_user.password_hash) > 0
|
||||
|
||||
def test_register_verification_token_generated(self, mock_user_repo, mock_email_service):
|
||||
"""生成邮箱验证令牌"""
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = None
|
||||
saved_user = None
|
||||
|
||||
def capture_save(user):
|
||||
nonlocal saved_user
|
||||
saved_user = user
|
||||
|
||||
mock_user_repo.save.side_effect = capture_save
|
||||
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="TestPass1!",
|
||||
username="testuser",
|
||||
display_name="测试",
|
||||
)
|
||||
use_case.execute(request)
|
||||
|
||||
assert saved_user.email_verification_token is not None
|
||||
assert len(saved_user.email_verification_token) > 0
|
||||
assert saved_user.email_verified is False
|
||||
|
||||
def test_register_verification_email_contains_url(self, mock_user_repo, mock_email_service):
|
||||
"""验证邮件包含正确的验证链接"""
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = None
|
||||
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://app.example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="TestPass1!",
|
||||
username="testuser",
|
||||
display_name="测试",
|
||||
)
|
||||
use_case.execute(request)
|
||||
|
||||
call_args = mock_email_service.send_verification_email.call_args
|
||||
verif_url = call_args[1].get("verification_url", "") or ""
|
||||
assert "https://app.example.com/verify-email?token=" in verif_url
|
||||
|
||||
def test_register_email_failure_still_creates_user(self, mock_user_repo, mock_email_service):
|
||||
"""邮件发送失败但用户仍被创建"""
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = None
|
||||
mock_email_service.send_verification_email.return_value = (False, "SMTP error")
|
||||
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="TestPass1!",
|
||||
username="testuser",
|
||||
display_name="测试",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None # 用户创建成功
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.email_verification_sent is False # 但邮件发送失败
|
||||
assert response.email_verification_sent is False
|
||||
mock_user_repo.save.assert_called_once()
|
||||
|
||||
def test_register_generates_user_id(self, mock_user_repo, mock_email_service):
|
||||
"""新用户有 id"""
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = None
|
||||
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RegisterUserRequest(
|
||||
email="test@example.com",
|
||||
password="TestPass1!",
|
||||
username="testuser",
|
||||
display_name="测试",
|
||||
)
|
||||
response, _ = use_case.execute(request)
|
||||
|
||||
assert response.user_id is not None
|
||||
assert len(response.user_id) > 0
|
||||
|
||||
def test_register_two_users_different_ids(self, mock_user_repo, mock_email_service):
|
||||
"""两个用户的 id 不同"""
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = None
|
||||
|
||||
use_case = RegisterUserUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
r1 = RegisterUserRequest(
|
||||
email="user1@example.com", password="TestPass1!",
|
||||
username="user1", display_name="用户1",
|
||||
)
|
||||
r2 = RegisterUserRequest(
|
||||
email="user2@example.com", password="TestPass1!",
|
||||
username="user2", display_name="用户2",
|
||||
)
|
||||
|
||||
resp1, _ = use_case.execute(r1)
|
||||
resp2, _ = use_case.execute(r2)
|
||||
|
||||
assert resp1.user_id != resp2.user_id
|
||||
|
||||
|
||||
class TestVerifyEmailUseCase:
|
||||
"""邮箱验证用例测试"""
|
||||
"""VerifyEmailUseCase 测试"""
|
||||
|
||||
@pytest.fixture
|
||||
def mock_user_repo(self):
|
||||
repo = Mock()
|
||||
repo.find_by_verification_token = Mock(return_value=None)
|
||||
repo.save = Mock()
|
||||
return repo
|
||||
def test_verify_success(self, mock_user_repo, sample_user):
|
||||
"""邮箱验证成功"""
|
||||
mock_user_repo.find_by_verification_token.return_value = sample_user
|
||||
mock_user_repo.save.return_value = None
|
||||
|
||||
@pytest.fixture
|
||||
def use_case(self, mock_user_repo):
|
||||
return VerifyEmailUseCase(user_repository=mock_user_repo)
|
||||
|
||||
def test_verify_email_success(self, use_case, mock_user_repo):
|
||||
"""测试验证成功"""
|
||||
user = User(
|
||||
id="user-123",
|
||||
email="test@example.com",
|
||||
username="testuser",
|
||||
display_name="Test User",
|
||||
email_verified=False,
|
||||
email_verification_token="valid-token",
|
||||
)
|
||||
mock_user_repo.find_by_verification_token.return_value = user
|
||||
|
||||
request = VerifyEmailRequest(token="valid-token")
|
||||
use_case = VerifyEmailUseCase(mock_user_repo)
|
||||
request = VerifyEmailRequest(token="some_token")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is True
|
||||
assert error is None
|
||||
|
||||
# 验证用户状态已更新
|
||||
assert user.email_verified is True
|
||||
assert user.email_verification_token is None
|
||||
assert sample_user.email_verified is True
|
||||
assert sample_user.email_verification_token is None
|
||||
mock_user_repo.save.assert_called_once()
|
||||
|
||||
def test_verify_email_invalid_token(self, use_case, mock_user_repo):
|
||||
"""测试无效令牌"""
|
||||
mock_user_repo.find_by_verification_token.return_value = None
|
||||
|
||||
request = VerifyEmailRequest(token="invalid-token")
|
||||
def test_verify_empty_token(self, mock_user_repo):
|
||||
"""空 token 返回错误"""
|
||||
use_case = VerifyEmailUseCase(mock_user_repo)
|
||||
request = VerifyEmailRequest(token="")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is False
|
||||
assert error == "Invalid or expired verification token"
|
||||
assert "Verification token is required" in error
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_verify_email_already_verified(self, use_case, mock_user_repo):
|
||||
"""测试已验证的邮箱"""
|
||||
user = User(
|
||||
id="user-123",
|
||||
email="test@example.com",
|
||||
username="testuser",
|
||||
display_name="Test User",
|
||||
email_verified=True,
|
||||
email_verification_token="old-token",
|
||||
)
|
||||
mock_user_repo.find_by_verification_token.return_value = user
|
||||
def test_verify_invalid_token(self, mock_user_repo):
|
||||
"""无效 token 返回错误"""
|
||||
mock_user_repo.find_by_verification_token.return_value = None
|
||||
|
||||
request = VerifyEmailRequest(token="old-token")
|
||||
use_case = VerifyEmailUseCase(mock_user_repo)
|
||||
request = VerifyEmailRequest(token="invalid_token")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is True # 已验证也返回成功
|
||||
assert success is False
|
||||
assert "Invalid or expired" in error
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_verify_already_verified(self, mock_user_repo, sample_user):
|
||||
"""已验证的用户再次验证也返回成功"""
|
||||
sample_user.email_verified = True
|
||||
mock_user_repo.find_by_verification_token.return_value = sample_user
|
||||
|
||||
use_case = VerifyEmailUseCase(mock_user_repo)
|
||||
request = VerifyEmailRequest(token="some_token")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is True
|
||||
assert error is None
|
||||
|
||||
@@ -1,11 +1,6 @@
|
||||
"""
|
||||
验证码服务单元测试(第十七波)
|
||||
"""验证码服务单元测试."""
|
||||
|
||||
覆盖:
|
||||
- VerificationCodeService.generate
|
||||
- VerificationCodeService.verify
|
||||
- 频控逻辑(冷却 + 每日上限)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import MagicMock
|
||||
@@ -14,436 +9,328 @@ 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,
|
||||
DAILY_LIMIT,
|
||||
MAX_ATTEMPTS,
|
||||
RESEND_COOLDOWN_SECONDS,
|
||||
VerificationCodeService,
|
||||
normalize_phone,
|
||||
validate_email,
|
||||
validate_phone,
|
||||
)
|
||||
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 code_service(mock_repo):
|
||||
return VerificationCodeService(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,
|
||||
@pytest.fixture
|
||||
def sample_code():
|
||||
code = VerificationCode.create(
|
||||
recipient="test@example.com",
|
||||
code_type=CODE_TYPE_EMAIL_BIND,
|
||||
ttl_seconds=300,
|
||||
)
|
||||
return code
|
||||
|
||||
|
||||
# ============================================================
|
||||
# generate - 参数校验
|
||||
# ============================================================
|
||||
class TestVerificationCodeServiceGenerate:
|
||||
"""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 掉,正常生成"""
|
||||
def test_generate_success(self, code_service, mock_repo, sample_code):
|
||||
"""生成验证码成功"""
|
||||
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"
|
||||
mock_repo.save.return_value = None
|
||||
|
||||
def test_invalid_code_type(self, service):
|
||||
"""无效验证码类型"""
|
||||
code, err = service.generate("test@example.com", "invalid_type")
|
||||
assert code is None
|
||||
assert "无效的验证码类型" in err
|
||||
code, error = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
|
||||
# ============================================================
|
||||
# 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 error 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
|
||||
def test_generate_empty_recipient(self, code_service):
|
||||
"""空接收方返回错误"""
|
||||
code, error = code_service.generate("", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "接收方不能为空" in error
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND, custom_code="888888")
|
||||
def test_generate_invalid_type(self, code_service):
|
||||
"""无效验证码类型返回错误"""
|
||||
code, error = code_service.generate("test@example.com", "invalid_type")
|
||||
assert code is None
|
||||
assert "无效的验证码类型" in error
|
||||
|
||||
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
|
||||
def test_generate_cooldown(self, code_service, mock_repo, sample_code):
|
||||
"""冷却期内返回频控错误"""
|
||||
# 最新的验证码刚创建10秒前
|
||||
sample_code.created_at = datetime.now(timezone.utc) - timedelta(seconds=10)
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
mock_repo.count_today.return_value = 1
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
code, error = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert code is None
|
||||
assert "发送太频繁" in err
|
||||
assert "秒后再试" in err
|
||||
# 等待时间应接近 50 秒(60-10)
|
||||
# 提取数字验证范围
|
||||
import re
|
||||
assert "发送太频繁" in error
|
||||
assert "秒后再试" in error
|
||||
|
||||
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
|
||||
def test_generate_daily_limit_exceeded(self, code_service, mock_repo):
|
||||
"""超过每日上限返回错误"""
|
||||
mock_repo.find_latest.return_value = None # 没有冷却期问题
|
||||
mock_repo.count_today.return_value = DAILY_LIMIT
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
code, error = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert code is None
|
||||
assert "今日发送次数已达上限" in err
|
||||
assert "今日发送次数已达上限" in error
|
||||
|
||||
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):
|
||||
"""首次发送,无历史记录"""
|
||||
def test_generate_recipient_stripped(self, code_service, mock_repo, sample_code):
|
||||
"""recipient 会被 strip"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
mock_repo.save.return_value = None
|
||||
|
||||
code, err = service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
code_service.generate(" test@example.com ", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert err is None
|
||||
assert code is not None
|
||||
mock_repo.save.assert_called_once()
|
||||
# 传给 repo 的应该是 strip 后的值
|
||||
save_call = mock_repo.save.call_args[0][0]
|
||||
assert save_call.recipient == "test@example.com"
|
||||
|
||||
|
||||
# ============================================================
|
||||
# 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)
|
||||
def test_generate_custom_code(self, code_service, mock_repo):
|
||||
"""使用自定义验证码"""
|
||||
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
|
||||
mock_repo.save.return_value = None
|
||||
|
||||
ok, err = service.verify(" test@example.com ", CODE_TYPE_EMAIL_BIND, " 123456 ")
|
||||
code, _ = code_service.generate(
|
||||
"test@example.com", CODE_TYPE_EMAIL_BIND, custom_code="123456"
|
||||
)
|
||||
assert code.code == "123456"
|
||||
|
||||
assert ok
|
||||
assert err is None
|
||||
def test_generate_custom_ttl(self, code_service, mock_repo):
|
||||
"""自定义 TTL"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
mock_repo.save.return_value = None
|
||||
|
||||
code, _ = code_service.generate(
|
||||
"test@example.com", CODE_TYPE_EMAIL_BIND, ttl_seconds=600
|
||||
)
|
||||
assert code is not None
|
||||
|
||||
|
||||
# ============================================================
|
||||
# verify - 正常验证
|
||||
# ============================================================
|
||||
class TestVerificationCodeServiceVerify:
|
||||
"""verify 方法测试"""
|
||||
|
||||
def test_verify_success(self, code_service, mock_repo, sample_code):
|
||||
"""验证成功"""
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
|
||||
class TestVerifyNormal:
|
||||
"""verify 正常验证场景"""
|
||||
success, error = code_service.verify(
|
||||
"test@example.com", CODE_TYPE_EMAIL_BIND, sample_code.code
|
||||
)
|
||||
|
||||
def test_verify_success_consume(self, service, mock_repo):
|
||||
"""验证成功并消耗"""
|
||||
code = make_code(code="123456")
|
||||
mock_repo.find_latest.return_value = code
|
||||
assert success is True
|
||||
assert error is None
|
||||
assert sample_code.is_used is True
|
||||
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456", consume=True)
|
||||
def test_verify_wrong_code(self, code_service, mock_repo, sample_code):
|
||||
"""验证码错误"""
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
|
||||
assert ok
|
||||
assert err is None
|
||||
assert code.is_used # 被标记为已使用
|
||||
# save 被调用了两次:一次 increment_attempts 后,一次 mark_used 后
|
||||
assert mock_repo.save.call_count >= 2
|
||||
success, error = code_service.verify(
|
||||
"test@example.com", CODE_TYPE_EMAIL_BIND, "wrongcode"
|
||||
)
|
||||
|
||||
def test_verify_success_no_consume(self, service, mock_repo):
|
||||
"""验证成功但不消耗"""
|
||||
code = make_code(code="123456")
|
||||
mock_repo.find_latest.return_value = code
|
||||
assert success is False
|
||||
assert "验证码错误" in error
|
||||
|
||||
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):
|
||||
def test_verify_not_found(self, code_service, mock_repo):
|
||||
"""验证码不存在"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
success, error = code_service.verify(
|
||||
"test@example.com", CODE_TYPE_EMAIL_BIND, "123456"
|
||||
)
|
||||
|
||||
assert not ok
|
||||
assert "不存在或已过期" in err
|
||||
assert success is False
|
||||
assert "不存在或已过期" in error
|
||||
|
||||
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):
|
||||
def test_verify_expired(self, code_service, mock_repo):
|
||||
"""验证码已过期"""
|
||||
code = make_code(code="123456", ttl=-60) # 已过期 60 秒
|
||||
mock_repo.find_latest.return_value = code
|
||||
expired_code = VerificationCode.create(
|
||||
recipient="test@example.com",
|
||||
code_type=CODE_TYPE_EMAIL_BIND,
|
||||
ttl_seconds=1, # 1秒过期
|
||||
)
|
||||
# 手动设置过期时间
|
||||
expired_code.expires_at = datetime.now(timezone.utc) - timedelta(seconds=10)
|
||||
mock_repo.find_latest.return_value = expired_code
|
||||
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
success, error = code_service.verify(
|
||||
"test@example.com", CODE_TYPE_EMAIL_BIND, expired_code.code
|
||||
)
|
||||
|
||||
assert not ok
|
||||
assert "已过期" in err
|
||||
assert success is False
|
||||
assert "已过期" in error
|
||||
|
||||
def test_verify_attempts_exceeded(self, service, mock_repo):
|
||||
"""超过最大尝试次数"""
|
||||
code = make_code(code="123456", attempts=MAX_ATTEMPTS)
|
||||
mock_repo.find_latest.return_value = code
|
||||
def test_verify_already_used(self, code_service, mock_repo, sample_code):
|
||||
"""验证码已使用"""
|
||||
sample_code.mark_used()
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
|
||||
ok, err = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
success, error = code_service.verify(
|
||||
"test@example.com", CODE_TYPE_EMAIL_BIND, sample_code.code
|
||||
)
|
||||
|
||||
assert not ok
|
||||
assert "验证次数过多" in err
|
||||
# verify 里先 increment_attempts 再判断,所以这里 attempts 应该是 MAX_ATTEMPTS + 1
|
||||
assert code.attempts == MAX_ATTEMPTS + 1
|
||||
assert success is False
|
||||
assert "已使用" in error
|
||||
|
||||
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
|
||||
def test_verify_max_attempts_exceeded(self, code_service, mock_repo, sample_code):
|
||||
"""尝试次数过多"""
|
||||
# 先把尝试次数加到超过上限
|
||||
for _ in range(MAX_ATTEMPTS + 1):
|
||||
sample_code.increment_attempts()
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
|
||||
service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "000000")
|
||||
assert code.attempts == 1
|
||||
success, error = code_service.verify(
|
||||
"test@example.com", CODE_TYPE_EMAIL_BIND, sample_code.code
|
||||
)
|
||||
|
||||
service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "000001")
|
||||
assert code.attempts == 2
|
||||
assert success is False
|
||||
assert "验证次数过多" in error
|
||||
|
||||
def test_verify_empty_params(self, code_service):
|
||||
"""空参数返回错误"""
|
||||
success, error = code_service.verify("", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
assert success is False
|
||||
assert "参数不完整" in error
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "")
|
||||
assert success is False
|
||||
assert "参数不完整" in error
|
||||
|
||||
def test_verify_increments_attempts(self, code_service, mock_repo, sample_code):
|
||||
"""验证会增加尝试次数"""
|
||||
initial_attempts = sample_code.attempts
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
|
||||
code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "wrong")
|
||||
|
||||
assert sample_code.attempts == initial_attempts + 1
|
||||
|
||||
def test_verify_no_consume(self, code_service, mock_repo, sample_code):
|
||||
"""consume=False 时不标记为已使用"""
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
|
||||
success, _ = code_service.verify(
|
||||
"test@example.com", CODE_TYPE_EMAIL_BIND, sample_code.code, consume=False
|
||||
)
|
||||
|
||||
assert success is True
|
||||
assert sample_code.is_used is False
|
||||
|
||||
|
||||
# ============================================================
|
||||
# verify - 不同 code_type 互不干扰
|
||||
# ============================================================
|
||||
class TestVerifyPhone:
|
||||
"""validate_phone 函数测试"""
|
||||
|
||||
def test_valid_phone(self):
|
||||
"""有效手机号"""
|
||||
ok, err = validate_phone("13800000001")
|
||||
assert ok is True
|
||||
assert err == ""
|
||||
|
||||
def test_valid_phone_with_plus86(self):
|
||||
"""带 +86 前缀的手机号"""
|
||||
ok, err = validate_phone("+8613800000001")
|
||||
assert ok is True
|
||||
|
||||
def test_invalid_phone_short(self):
|
||||
"""太短的手机号"""
|
||||
ok, err = validate_phone("123")
|
||||
assert ok is False
|
||||
assert "格式不正确" in err
|
||||
|
||||
def test_invalid_phone_wrong_prefix(self):
|
||||
"""号段不对的手机号"""
|
||||
ok, err = validate_phone("11000000000")
|
||||
assert ok is False
|
||||
|
||||
def test_empty_phone(self):
|
||||
"""空手机号"""
|
||||
ok, err = validate_phone("")
|
||||
assert ok is False
|
||||
assert "不能为空" in err
|
||||
|
||||
def test_phone_with_spaces(self):
|
||||
"""带空格的手机号会被 strip"""
|
||||
ok, _ = validate_phone(" 13800000001 ")
|
||||
assert ok is True
|
||||
|
||||
|
||||
class TestVerifyCodeTypeIsolation:
|
||||
"""不同验证码类型互不干扰"""
|
||||
class TestNormalizePhone:
|
||||
"""normalize_phone 函数测试"""
|
||||
|
||||
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 查不到
|
||||
def test_removes_plus86(self):
|
||||
"""去掉 +86 前缀"""
|
||||
assert normalize_phone("+8613800000001") == "13800000001"
|
||||
|
||||
# find_latest 按 code_type 查询,传 email_bind 返回 None
|
||||
def side_effect(recipient, ct):
|
||||
if ct == CODE_TYPE_EMAIL_LOGIN:
|
||||
return code
|
||||
return None
|
||||
def test_no_prefix_stays_same(self):
|
||||
"""没有前缀保持不变"""
|
||||
assert normalize_phone("13800000001") == "13800000001"
|
||||
|
||||
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
|
||||
def test_strips_whitespace(self):
|
||||
"""去掉两端空白"""
|
||||
assert normalize_phone(" 13800000001 ") == "13800000001"
|
||||
|
||||
|
||||
# ============================================================
|
||||
# 常量值检查
|
||||
# ============================================================
|
||||
class TestValidateEmail:
|
||||
"""validate_email 函数测试"""
|
||||
|
||||
def test_valid_email(self):
|
||||
"""有效邮箱"""
|
||||
ok, err = validate_email("test@example.com")
|
||||
assert ok is True
|
||||
assert err == ""
|
||||
|
||||
class TestConstants:
|
||||
"""常量默认值校验"""
|
||||
def test_valid_email_with_subdomain(self):
|
||||
"""带子域名的邮箱"""
|
||||
ok, _ = validate_email("user@mail.example.com")
|
||||
assert ok is True
|
||||
|
||||
def test_default_cooldown_60(self):
|
||||
assert RESEND_COOLDOWN_SECONDS == 60
|
||||
def test_valid_email_with_plus(self):
|
||||
"""带 + 号的邮箱"""
|
||||
ok, _ = validate_email("user+tag@example.com")
|
||||
assert ok is True
|
||||
|
||||
def test_default_daily_limit_10(self):
|
||||
assert DAILY_LIMIT == 10
|
||||
def test_invalid_email_no_at(self):
|
||||
"""没有 @ 的邮箱"""
|
||||
ok, err = validate_email("notanemail")
|
||||
assert ok is False
|
||||
assert "格式不正确" in err
|
||||
|
||||
def test_default_max_attempts_5(self):
|
||||
assert MAX_ATTEMPTS == 5
|
||||
def test_invalid_email_no_domain(self):
|
||||
"""没有域名的邮箱"""
|
||||
ok, err = validate_email("user@")
|
||||
assert ok is False
|
||||
|
||||
def test_default_ttl_300(self):
|
||||
assert DEFAULT_TTL_SECONDS == 300
|
||||
def test_empty_email(self):
|
||||
"""空邮箱"""
|
||||
ok, err = validate_email("")
|
||||
assert ok is False
|
||||
assert "不能为空" in err
|
||||
|
||||
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
|
||||
def test_email_with_spaces(self):
|
||||
"""带空格的邮箱会被 strip"""
|
||||
ok, _ = validate_email(" test@example.com ")
|
||||
assert ok is True
|
||||
|
||||
Reference in New Issue
Block a user