test: P3-1 第40波单元测试(wechat_oauth + wechat_sync) #821
Executable
+197
@@ -0,0 +1,197 @@
|
||||
"""Assets UseCase 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.assets import (
|
||||
CreateAssetCommand,
|
||||
CreateAssetUseCase,
|
||||
ListAssetsUseCase,
|
||||
)
|
||||
from packages.domain import Asset, AssetStatus, ClassificationStatus
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_asset_repo():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_asset():
|
||||
asset = Asset.create(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
name="test_video.mp4",
|
||||
storage_key="videos/test.mp4",
|
||||
mime_type="video/mp4",
|
||||
file_size=1024000,
|
||||
duration=15.5,
|
||||
width=1920,
|
||||
height=1080,
|
||||
)
|
||||
asset.id = "asset_001"
|
||||
return asset
|
||||
|
||||
|
||||
class TestListAssetsUseCase:
|
||||
"""ListAssetsUseCase 测试"""
|
||||
|
||||
def test_list_returns_repo_results(self, mock_asset_repo, sample_asset):
|
||||
"""正常返回 repository 的查询结果"""
|
||||
mock_asset_repo.find_by_library.return_value = [sample_asset]
|
||||
use_case = ListAssetsUseCase(mock_asset_repo)
|
||||
|
||||
result = use_case.execute("lib_001")
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0].id == "asset_001"
|
||||
mock_asset_repo.find_by_library.assert_called_once_with("lib_001")
|
||||
|
||||
def test_empty_library_id_raises_value_error(self, mock_asset_repo):
|
||||
"""空 library_id 抛出 ValueError"""
|
||||
use_case = ListAssetsUseCase(mock_asset_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="library_id 不能为空"):
|
||||
use_case.execute("")
|
||||
|
||||
mock_asset_repo.find_by_library.assert_not_called()
|
||||
|
||||
def test_whitespace_library_id_raises_value_error(self, mock_asset_repo):
|
||||
"""纯空格 library_id 抛出 ValueError"""
|
||||
use_case = ListAssetsUseCase(mock_asset_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="library_id 不能为空"):
|
||||
use_case.execute(" ")
|
||||
|
||||
mock_asset_repo.find_by_library.assert_not_called()
|
||||
|
||||
def test_library_id_stripped_before_query(self, mock_asset_repo, sample_asset):
|
||||
"""library_id 会被 strip 后再查询"""
|
||||
mock_asset_repo.find_by_library.return_value = [sample_asset]
|
||||
use_case = ListAssetsUseCase(mock_asset_repo)
|
||||
|
||||
use_case.execute(" lib_001 ")
|
||||
|
||||
mock_asset_repo.find_by_library.assert_called_once_with("lib_001")
|
||||
|
||||
def test_empty_list(self, mock_asset_repo):
|
||||
"""素材库为空时返回空列表"""
|
||||
mock_asset_repo.find_by_library.return_value = []
|
||||
use_case = ListAssetsUseCase(mock_asset_repo)
|
||||
|
||||
result = use_case.execute("lib_001")
|
||||
|
||||
assert result == []
|
||||
mock_asset_repo.find_by_library.assert_called_once_with("lib_001")
|
||||
|
||||
|
||||
class TestCreateAssetUseCase:
|
||||
"""CreateAssetUseCase 测试"""
|
||||
|
||||
def test_create_asset_success(self, mock_asset_repo):
|
||||
"""正常创建素材"""
|
||||
mock_asset_repo.create.side_effect = lambda a: a
|
||||
use_case = CreateAssetUseCase(mock_asset_repo)
|
||||
|
||||
command = CreateAssetCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
name="test.png",
|
||||
storage_key="images/test.png",
|
||||
mime_type="image/png",
|
||||
file_size=512000,
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.name == "test.png"
|
||||
assert result.library_id == "lib_001"
|
||||
assert result.mime_type == "image/png"
|
||||
assert result.status == AssetStatus.UPLOADING
|
||||
assert result.classification_status == ClassificationStatus.PENDING
|
||||
mock_asset_repo.create.assert_called_once()
|
||||
|
||||
def test_create_asset_with_metadata(self, mock_asset_repo):
|
||||
"""创建带 metadata 的素材"""
|
||||
mock_asset_repo.create.side_effect = lambda a: a
|
||||
use_case = CreateAssetUseCase(mock_asset_repo)
|
||||
|
||||
command = CreateAssetCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
name="test.mp3",
|
||||
storage_key="audio/test.mp3",
|
||||
mime_type="audio/mpeg",
|
||||
metadata={"bitrate": 320, "sample_rate": 44100},
|
||||
duration=180.0,
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.metadata["bitrate"] == 320
|
||||
assert result.duration == 180.0
|
||||
|
||||
def test_create_asset_with_quality_score(self, mock_asset_repo):
|
||||
"""创建带质量分的素材"""
|
||||
mock_asset_repo.create.side_effect = lambda a: a
|
||||
use_case = CreateAssetUseCase(mock_asset_repo)
|
||||
|
||||
command = CreateAssetCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
name="high_quality.mp4",
|
||||
storage_key="videos/hq.mp4",
|
||||
mime_type="video/mp4",
|
||||
quality_score=95.5,
|
||||
uploaded_by_user_id="user_001",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.quality_score == 95.5
|
||||
assert result.uploaded_by_user_id == "user_001"
|
||||
|
||||
def test_create_asset_custom_status(self, mock_asset_repo):
|
||||
"""创建时指定自定义状态"""
|
||||
mock_asset_repo.create.side_effect = lambda a: a
|
||||
use_case = CreateAssetUseCase(mock_asset_repo)
|
||||
|
||||
command = CreateAssetCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
name="ready.mp4",
|
||||
storage_key="videos/ready.mp4",
|
||||
mime_type="video/mp4",
|
||||
status=AssetStatus.READY,
|
||||
classification_status=ClassificationStatus.COMPLETED,
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.status == AssetStatus.READY
|
||||
assert result.classification_status == ClassificationStatus.COMPLETED
|
||||
|
||||
def test_create_asset_with_video_info(self, mock_asset_repo):
|
||||
"""创建带视频参数的素材"""
|
||||
mock_asset_repo.create.side_effect = lambda a: a
|
||||
use_case = CreateAssetUseCase(mock_asset_repo)
|
||||
|
||||
command = CreateAssetCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
name="video.mp4",
|
||||
storage_key="videos/v.mp4",
|
||||
mime_type="video/mp4",
|
||||
width=1920,
|
||||
height=1080,
|
||||
fps=30.0,
|
||||
codec="h264",
|
||||
duration=60.0,
|
||||
thumbnail_url="https://cdn.example.com/thumb.jpg",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.width == 1920
|
||||
assert result.height == 1080
|
||||
assert result.fps == 30.0
|
||||
assert result.codec == "h264"
|
||||
assert result.thumbnail_url == "https://cdn.example.com/thumb.jpg"
|
||||
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"
|
||||
Executable
+484
@@ -0,0 +1,484 @@
|
||||
"""绑定联系方式 UseCase 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.auth.bind_contact_use_case import (
|
||||
BindContactRequest,
|
||||
BindContactUseCase,
|
||||
SendVerificationCodeRequest,
|
||||
SendVerificationCodeUseCase,
|
||||
)
|
||||
from packages.domain.entities import User
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_user_repo():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_verification_service():
|
||||
svc = MagicMock()
|
||||
svc.verify.return_value = (True, None)
|
||||
return svc
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_user():
|
||||
user = User(
|
||||
id="user_001",
|
||||
email="",
|
||||
display_name="测试用户",
|
||||
phone_verified=False,
|
||||
email_verified=False,
|
||||
)
|
||||
user.phone = None
|
||||
return user
|
||||
|
||||
|
||||
class TestBindContactRequest:
|
||||
"""BindContactRequest 测试"""
|
||||
|
||||
def test_phone_strips_plus86(self):
|
||||
"""手机号 +86 前缀会被去掉"""
|
||||
req = BindContactRequest(
|
||||
user_id="u1", phone="+8613800000001", phone_code="1234"
|
||||
)
|
||||
assert req.phone == "13800000001"
|
||||
|
||||
def test_email_lowercased(self):
|
||||
"""邮箱会被转小写"""
|
||||
req = BindContactRequest(
|
||||
user_id="u1", email="Test@Example.COM", email_code="1234"
|
||||
)
|
||||
assert req.email == "test@example.com"
|
||||
|
||||
def test_code_stripped(self):
|
||||
"""验证码会被 strip"""
|
||||
req = BindContactRequest(
|
||||
user_id="u1", phone="13800000001", phone_code=" 1234 "
|
||||
)
|
||||
assert req.phone_code == "1234"
|
||||
|
||||
def test_empty_fields(self):
|
||||
"""空字段处理"""
|
||||
req = BindContactRequest(user_id="u1")
|
||||
assert req.phone == ""
|
||||
assert req.email == ""
|
||||
assert req.phone_code == ""
|
||||
assert req.email_code == ""
|
||||
|
||||
|
||||
class TestBindContactUseCase:
|
||||
"""BindContactUseCase 测试"""
|
||||
|
||||
def test_bind_phone_success(self, mock_user_repo, mock_verification_service, sample_user):
|
||||
"""绑定手机号成功"""
|
||||
mock_user_repo.find_by_id.return_value = sample_user
|
||||
mock_user_repo.find_by_phone.return_value = None
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="user_001",
|
||||
phone="13800000001",
|
||||
phone_code="123456",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.user.phone == "13800000001"
|
||||
assert response.user.phone_verified is True
|
||||
mock_user_repo.save.assert_called_once()
|
||||
|
||||
def test_bind_email_success(self, mock_user_repo, mock_verification_service, sample_user):
|
||||
"""绑定邮箱成功"""
|
||||
mock_user_repo.find_by_id.return_value = sample_user
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="user_001",
|
||||
email="test@example.com",
|
||||
email_code="123456",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.user.email == "test@example.com"
|
||||
assert response.user.email_verified is True
|
||||
|
||||
def test_bind_phone_and_email(self, mock_user_repo, mock_verification_service, sample_user):
|
||||
"""同时绑定手机和邮箱"""
|
||||
mock_user_repo.find_by_id.return_value = sample_user
|
||||
mock_user_repo.find_by_phone.return_value = None
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="user_001",
|
||||
phone="13800000001",
|
||||
phone_code="123456",
|
||||
email="test@example.com",
|
||||
email_code="123456",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response.user.phone == "13800000001"
|
||||
assert response.user.phone_verified is True
|
||||
assert response.user.email == "test@example.com"
|
||||
assert response.user.email_verified is True
|
||||
# 两个都绑定完成,binding_completed_at 应该被设置
|
||||
assert response.user.binding_completed_at is not None
|
||||
|
||||
def test_no_contact_info_returns_error(self, mock_user_repo, mock_verification_service):
|
||||
"""既没填手机也没填邮箱返回错误"""
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(user_id="user_001")
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "至少填写" in error
|
||||
mock_user_repo.find_by_id.assert_not_called()
|
||||
|
||||
def test_user_not_found(self, mock_user_repo, mock_verification_service):
|
||||
"""用户不存在返回错误"""
|
||||
mock_user_repo.find_by_id.return_value = None
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="nonexistent",
|
||||
phone="13800000001",
|
||||
phone_code="123456",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "用户不存在" in error
|
||||
|
||||
def test_phone_already_bound_by_other(self, mock_user_repo, mock_verification_service, sample_user):
|
||||
"""手机号已被其他账号绑定"""
|
||||
other_user = MagicMock()
|
||||
other_user.id = "user_other"
|
||||
mock_user_repo.find_by_id.return_value = sample_user
|
||||
mock_user_repo.find_by_phone.return_value = other_user
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="user_001",
|
||||
phone="13800000001",
|
||||
phone_code="123456",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "已被其他账号绑定" in error
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_phone_bound_by_self_ok(self, mock_user_repo, mock_verification_service, sample_user):
|
||||
"""手机号已被自己绑定,允许"""
|
||||
sample_user.phone = "13800000001"
|
||||
mock_user_repo.find_by_id.return_value = sample_user
|
||||
mock_user_repo.find_by_phone.return_value = sample_user
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="user_001",
|
||||
phone="13800000001",
|
||||
phone_code="123456",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
assert error is None
|
||||
assert response is not None
|
||||
|
||||
def test_wrong_phone_code(self, mock_user_repo, mock_verification_service, sample_user):
|
||||
"""手机验证码错误"""
|
||||
mock_user_repo.find_by_id.return_value = sample_user
|
||||
mock_user_repo.find_by_phone.return_value = None
|
||||
mock_verification_service.verify.return_value = (False, "验证码过期")
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="user_001",
|
||||
phone="13800000001",
|
||||
phone_code="000000",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "手机验证码错误" in error
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_missing_phone_code(self, mock_user_repo, mock_verification_service, sample_user):
|
||||
"""缺少手机验证码"""
|
||||
mock_user_repo.find_by_id.return_value = sample_user
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="user_001",
|
||||
phone="13800000001",
|
||||
phone_code="",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "请输入手机验证码" in error
|
||||
|
||||
def test_invalid_phone_format(self, mock_user_repo, mock_verification_service, sample_user):
|
||||
"""手机号格式不正确"""
|
||||
mock_user_repo.find_by_id.return_value = sample_user
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="user_001",
|
||||
phone="123", # 太短
|
||||
phone_code="123456",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert error is not None
|
||||
|
||||
def test_email_already_bound_by_other(self, mock_user_repo, mock_verification_service, sample_user):
|
||||
"""邮箱已被其他账号绑定"""
|
||||
other_user = MagicMock()
|
||||
other_user.id = "user_other"
|
||||
mock_user_repo.find_by_id.return_value = sample_user
|
||||
mock_user_repo.find_by_email.return_value = other_user
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="user_001",
|
||||
email="test@example.com",
|
||||
email_code="123456",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "已被其他账号绑定" in error
|
||||
|
||||
def test_missing_email_code(self, mock_user_repo, mock_verification_service, sample_user):
|
||||
"""缺少邮箱验证码"""
|
||||
mock_user_repo.find_by_id.return_value = sample_user
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="user_001",
|
||||
email="test@example.com",
|
||||
email_code="",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "请输入邮箱验证码" in error
|
||||
|
||||
def test_invalid_email_format(self, mock_user_repo, mock_verification_service, sample_user):
|
||||
"""邮箱格式不正确"""
|
||||
mock_user_repo.find_by_id.return_value = sample_user
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="user_001",
|
||||
email="not_an_email",
|
||||
email_code="123456",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert error is not None
|
||||
|
||||
def test_response_to_dict(self, mock_user_repo, mock_verification_service, sample_user):
|
||||
"""BindContactResponse.to_dict 返回正确格式"""
|
||||
mock_user_repo.find_by_id.return_value = sample_user
|
||||
mock_user_repo.find_by_phone.return_value = None
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
use_case = BindContactUseCase(mock_user_repo, mock_verification_service)
|
||||
request = BindContactRequest(
|
||||
user_id="user_001",
|
||||
phone="13800000001",
|
||||
phone_code="123456",
|
||||
email="test@example.com",
|
||||
email_code="123456",
|
||||
)
|
||||
response, _ = use_case.execute(request)
|
||||
data = response.to_dict()
|
||||
|
||||
assert "user" in data
|
||||
assert data["user"]["id"] == "user_001"
|
||||
assert "email" in data["user"]
|
||||
assert "phone" in data["user"]
|
||||
assert "phone_verified" in data["user"]
|
||||
assert "display_name" in data["user"]
|
||||
assert "binding_complete" in data["user"]
|
||||
|
||||
|
||||
class TestSendVerificationCodeRequest:
|
||||
"""SendVerificationCodeRequest 测试"""
|
||||
|
||||
def test_value_stripped(self):
|
||||
"""value 会被 strip"""
|
||||
req = SendVerificationCodeRequest(target="phone", value=" 13800000001 ", purpose="bind")
|
||||
assert req.value == "13800000001"
|
||||
|
||||
|
||||
class TestSendVerificationCodeUseCase:
|
||||
"""SendVerificationCodeUseCase 测试"""
|
||||
|
||||
def test_send_phone_code_success(self, mock_verification_service):
|
||||
"""发送手机验证码成功"""
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
code_obj = MagicMock()
|
||||
code_obj.code = "123456"
|
||||
code_obj.created_at = datetime.now(timezone.utc)
|
||||
code_obj.expires_at = datetime.now(timezone.utc) + timedelta(minutes=5)
|
||||
mock_verification_service.generate.return_value = (code_obj, None)
|
||||
|
||||
mock_sms = MagicMock()
|
||||
use_case = SendVerificationCodeUseCase(
|
||||
mock_verification_service,
|
||||
sms_service=mock_sms,
|
||||
)
|
||||
request = SendVerificationCodeRequest(
|
||||
target="phone",
|
||||
value="13800000001",
|
||||
purpose="bind",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.expires_in > 0
|
||||
assert response.resend_after == 60
|
||||
mock_sms.send_verification_code.assert_called_once()
|
||||
|
||||
def test_send_email_code_success(self, mock_verification_service):
|
||||
"""发送邮箱验证码成功"""
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
code_obj = MagicMock()
|
||||
code_obj.code = "654321"
|
||||
code_obj.created_at = datetime.now(timezone.utc)
|
||||
code_obj.expires_at = datetime.now(timezone.utc) + timedelta(minutes=5)
|
||||
mock_verification_service.generate.return_value = (code_obj, None)
|
||||
|
||||
mock_email = MagicMock()
|
||||
use_case = SendVerificationCodeUseCase(
|
||||
mock_verification_service,
|
||||
email_service=mock_email,
|
||||
)
|
||||
request = SendVerificationCodeRequest(
|
||||
target="email",
|
||||
value="test@example.com",
|
||||
purpose="bind",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
mock_email.send_email.assert_called_once()
|
||||
|
||||
def test_invalid_target_returns_error(self, mock_verification_service):
|
||||
"""不支持的目标类型返回错误"""
|
||||
use_case = SendVerificationCodeUseCase(mock_verification_service)
|
||||
request = SendVerificationCodeRequest(
|
||||
target="wechat",
|
||||
value="some_value",
|
||||
purpose="bind",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "不支持的目标类型" in error
|
||||
|
||||
def test_invalid_phone_format(self, mock_verification_service):
|
||||
"""手机号格式错误返回错误"""
|
||||
use_case = SendVerificationCodeUseCase(mock_verification_service)
|
||||
request = SendVerificationCodeRequest(
|
||||
target="phone",
|
||||
value="123",
|
||||
purpose="bind",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert error is not None
|
||||
mock_verification_service.generate.assert_not_called()
|
||||
|
||||
def test_invalid_email_format(self, mock_verification_service):
|
||||
"""邮箱格式错误返回错误"""
|
||||
use_case = SendVerificationCodeUseCase(mock_verification_service)
|
||||
request = SendVerificationCodeRequest(
|
||||
target="email",
|
||||
value="not_email",
|
||||
purpose="bind",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert error is not None
|
||||
mock_verification_service.generate.assert_not_called()
|
||||
|
||||
def test_generate_failure_returns_error(self, mock_verification_service):
|
||||
"""生成验证码失败返回错误"""
|
||||
mock_verification_service.generate.return_value = (None, "发送太频繁")
|
||||
|
||||
use_case = SendVerificationCodeUseCase(mock_verification_service)
|
||||
request = SendVerificationCodeRequest(
|
||||
target="phone",
|
||||
value="13800000001",
|
||||
purpose="bind",
|
||||
)
|
||||
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "发送太频繁" in error
|
||||
|
||||
def test_response_to_dict(self, mock_verification_service):
|
||||
"""SendVerificationCodeResponse.to_dict 格式正确"""
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
code_obj = MagicMock()
|
||||
code_obj.code = "123456"
|
||||
code_obj.created_at = datetime.now(timezone.utc)
|
||||
code_obj.expires_at = datetime.now(timezone.utc) + timedelta(seconds=300)
|
||||
mock_verification_service.generate.return_value = (code_obj, None)
|
||||
|
||||
use_case = SendVerificationCodeUseCase(mock_verification_service)
|
||||
request = SendVerificationCodeRequest(
|
||||
target="phone",
|
||||
value="13800000001",
|
||||
purpose="bind",
|
||||
)
|
||||
response, _ = use_case.execute(request)
|
||||
data = response.to_dict()
|
||||
|
||||
assert "expires_in" in data
|
||||
assert "resend_after" in data
|
||||
Executable
+96
@@ -0,0 +1,96 @@
|
||||
"""AI分类任务 UseCase 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.classification_jobs import (
|
||||
SubmitClassificationJobCommand,
|
||||
SubmitClassificationJobUseCase,
|
||||
)
|
||||
from packages.domain import ClassificationJob
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
class TestSubmitClassificationJobUseCase:
|
||||
"""SubmitClassificationJobUseCase 测试"""
|
||||
|
||||
def test_submit_job_success(self, mock_repo):
|
||||
"""正常提交分类任务"""
|
||||
mock_repo.create.side_effect = lambda j: j
|
||||
use_case = SubmitClassificationJobUseCase(mock_repo)
|
||||
|
||||
command = SubmitClassificationJobCommand(
|
||||
project_id="proj_001",
|
||||
asset_id="asset_001",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert isinstance(result, ClassificationJob)
|
||||
assert result.project_id == "proj_001"
|
||||
assert result.asset_id == "asset_001"
|
||||
assert result.status == "pending"
|
||||
assert result.confidence == 0.0
|
||||
assert result.error_message == ""
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_submit_job_generates_id(self, mock_repo):
|
||||
"""提交任务时生成 id"""
|
||||
mock_repo.create.side_effect = lambda j: j
|
||||
use_case = SubmitClassificationJobUseCase(mock_repo)
|
||||
|
||||
command = SubmitClassificationJobCommand(
|
||||
project_id="proj_001",
|
||||
asset_id="asset_001",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.id is not None
|
||||
assert len(result.id) > 0
|
||||
|
||||
def test_submit_job_two_different_ids(self, mock_repo):
|
||||
"""两次提交生成不同的 id"""
|
||||
mock_repo.create.side_effect = lambda j: j
|
||||
use_case = SubmitClassificationJobUseCase(mock_repo)
|
||||
|
||||
command = SubmitClassificationJobCommand(
|
||||
project_id="proj_001",
|
||||
asset_id="asset_001",
|
||||
)
|
||||
r1 = use_case.execute(command)
|
||||
r2 = use_case.execute(command)
|
||||
|
||||
assert r1.id != r2.id
|
||||
|
||||
def test_submit_job_initial_classification_empty(self, mock_repo):
|
||||
"""初始 classification 为空"""
|
||||
mock_repo.create.side_effect = lambda j: j
|
||||
use_case = SubmitClassificationJobUseCase(mock_repo)
|
||||
|
||||
command = SubmitClassificationJobCommand(
|
||||
project_id="proj_001",
|
||||
asset_id="asset_001",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.classification == ""
|
||||
|
||||
def test_submit_job_returns_repo_result(self, mock_repo):
|
||||
"""返回 repository.create 的结果"""
|
||||
expected = MagicMock(spec=ClassificationJob)
|
||||
mock_repo.create.return_value = expected
|
||||
|
||||
use_case = SubmitClassificationJobUseCase(mock_repo)
|
||||
command = SubmitClassificationJobCommand(
|
||||
project_id="proj_001",
|
||||
asset_id="asset_001",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result is expected
|
||||
+167
-227
@@ -1,13 +1,6 @@
|
||||
"""
|
||||
生成任务应用层用例单元测试(第十九波)
|
||||
"""生成任务 UseCase 单元测试."""
|
||||
|
||||
覆盖:
|
||||
- CreateGenerationTaskUseCase
|
||||
- GetGenerationTaskUseCase
|
||||
- ListUserTasksFilteredUseCase
|
||||
- RetryGenerationTaskUseCase
|
||||
- Command / Filter / Result 对象
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
@@ -22,7 +15,7 @@ from packages.application.generation_tasks import (
|
||||
ListUserTasksFilteredUseCase,
|
||||
RetryGenerationTaskUseCase,
|
||||
)
|
||||
from packages.domain.generation_task import GenerationTask, GenerationTaskStatus
|
||||
from packages.domain import GenerationTask
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -30,291 +23,238 @@ 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)
|
||||
@pytest.fixture
|
||||
def sample_task():
|
||||
task = MagicMock(spec=GenerationTask)
|
||||
task.id = "task_001"
|
||||
task.project_id = "proj_001"
|
||||
task.status = "pending"
|
||||
return task
|
||||
|
||||
|
||||
# ============================================================
|
||||
# CreateGenerationTaskUseCase
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestCreateGenerationTaskUseCase:
|
||||
"""CreateGenerationTaskUseCase 创建生成任务"""
|
||||
"""CreateGenerationTaskUseCase 测试"""
|
||||
|
||||
def test_create_success(self, mock_repo):
|
||||
"""正常创建任务"""
|
||||
def test_create_task_success(self, mock_repo):
|
||||
"""正常创建生成任务"""
|
||||
mock_repo.create.side_effect = lambda t: t
|
||||
use_case = CreateGenerationTaskUseCase(mock_repo)
|
||||
|
||||
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="我的视频",
|
||||
command = CreateGenerationTaskCommand(
|
||||
project_id="proj_001",
|
||||
template_id="tpl_001",
|
||||
asset_library_id="lib_001",
|
||||
voice_library_id="voice_lib_001",
|
||||
created_by_user_id="user_001",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert isinstance(result, GenerationTask)
|
||||
assert result.project_id == "proj_001"
|
||||
assert result.template_id == "tpl_001"
|
||||
assert result.status == "pending"
|
||||
assert result.progress == 0.0
|
||||
assert result.result_count == 0
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_create_task_generates_id(self, mock_repo):
|
||||
"""创建任务时生成 id"""
|
||||
mock_repo.create.side_effect = lambda t: t
|
||||
use_case = CreateGenerationTaskUseCase(mock_repo)
|
||||
|
||||
command = CreateGenerationTaskCommand(project_id="proj_001")
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.id is not None
|
||||
assert len(result.id) > 0
|
||||
|
||||
def test_create_task_with_asset_ids(self, mock_repo):
|
||||
"""创建带 asset_ids 的任务"""
|
||||
mock_repo.create.side_effect = lambda t: t
|
||||
use_case = CreateGenerationTaskUseCase(mock_repo)
|
||||
|
||||
command = CreateGenerationTaskCommand(
|
||||
project_id="proj_001",
|
||||
asset_ids=["asset_1", "asset_2", "asset_3"],
|
||||
title_ids=["title_1", "title_2"],
|
||||
voice_ids=["voice_1"],
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert len(result.asset_ids) == 3
|
||||
assert len(result.title_ids) == 2
|
||||
assert len(result.voice_ids) == 1
|
||||
|
||||
def test_create_task_with_auto_retry(self, mock_repo):
|
||||
"""创建带自动重试配置的任务"""
|
||||
mock_repo.create.side_effect = lambda t: t
|
||||
use_case = CreateGenerationTaskUseCase(mock_repo)
|
||||
|
||||
command = CreateGenerationTaskCommand(
|
||||
project_id="proj_001",
|
||||
auto_retry_enabled=True,
|
||||
auto_retry_max=3,
|
||||
)
|
||||
uc = CreateGenerationTaskUseCase(mock_repo)
|
||||
task = uc.execute(cmd)
|
||||
result = use_case.execute(command)
|
||||
|
||||
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()
|
||||
assert result.auto_retry_enabled is True
|
||||
assert result.auto_retry_max == 3
|
||||
|
||||
def test_create_default_values(self, mock_repo):
|
||||
"""默认参数值"""
|
||||
def test_create_task_with_bgm_config(self, mock_repo):
|
||||
"""创建带 BGM 配置的任务"""
|
||||
mock_repo.create.side_effect = lambda t: t
|
||||
use_case = CreateGenerationTaskUseCase(mock_repo)
|
||||
|
||||
cmd = CreateGenerationTaskCommand(
|
||||
project_id="proj-1",
|
||||
asset_library_id="lib-1",
|
||||
bgm = {"enabled": True, "volume": 0.5, "library_id": "bgm_lib"}
|
||||
command = CreateGenerationTaskCommand(
|
||||
project_id="proj_001",
|
||||
bgm_config=bgm,
|
||||
resolution="1080p",
|
||||
video_title="测试视频",
|
||||
)
|
||||
uc = CreateGenerationTaskUseCase(mock_repo)
|
||||
task = uc.execute(cmd)
|
||||
result = use_case.execute(command)
|
||||
|
||||
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
|
||||
assert result.bgm_config == bgm
|
||||
assert result.resolution == "1080p"
|
||||
assert result.video_title == "测试视频"
|
||||
|
||||
def test_create_id_is_generated(self, mock_repo):
|
||||
"""ID 会自动生成"""
|
||||
def test_create_task_defaults(self, mock_repo):
|
||||
"""默认参数的任务"""
|
||||
mock_repo.create.side_effect = lambda t: t
|
||||
use_case = CreateGenerationTaskUseCase(mock_repo)
|
||||
|
||||
cmd = CreateGenerationTaskCommand(project_id="proj-1", asset_library_id="lib-1")
|
||||
uc = CreateGenerationTaskUseCase(mock_repo)
|
||||
task = uc.execute(cmd)
|
||||
command = CreateGenerationTaskCommand()
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert task.id
|
||||
assert isinstance(task.id, str)
|
||||
assert len(task.id) > 10 # uuid hex
|
||||
|
||||
|
||||
# ============================================================
|
||||
# GetGenerationTaskUseCase
|
||||
# ============================================================
|
||||
assert result.project_id == ""
|
||||
assert result.asset_ids == []
|
||||
assert result.auto_retry_enabled is False
|
||||
assert result.auto_retry_max == 0
|
||||
|
||||
|
||||
class TestGetGenerationTaskUseCase:
|
||||
"""GetGenerationTaskUseCase 获取任务"""
|
||||
"""GetGenerationTaskUseCase 测试"""
|
||||
|
||||
def test_get_existing(self, mock_repo):
|
||||
"""获取存在的任务"""
|
||||
task = make_task()
|
||||
mock_repo.get.return_value = task
|
||||
def test_get_task_success(self, mock_repo, sample_task):
|
||||
"""获取任务成功"""
|
||||
mock_repo.get.return_value = sample_task
|
||||
|
||||
uc = GetGenerationTaskUseCase(mock_repo)
|
||||
result = uc.execute("task-1")
|
||||
use_case = GetGenerationTaskUseCase(mock_repo)
|
||||
result = use_case.execute("task_001")
|
||||
|
||||
assert result is task
|
||||
mock_repo.get.assert_called_once_with("task-1")
|
||||
assert result is sample_task
|
||||
mock_repo.get.assert_called_once_with("task_001")
|
||||
|
||||
def test_get_not_found(self, mock_repo):
|
||||
"""获取不存在的任务返回 None"""
|
||||
def test_get_task_not_found(self, mock_repo):
|
||||
"""任务不存在返回 None"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
uc = GetGenerationTaskUseCase(mock_repo)
|
||||
result = uc.execute("nonexistent")
|
||||
use_case = GetGenerationTaskUseCase(mock_repo)
|
||||
result = use_case.execute("nonexistent")
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
# ============================================================
|
||||
# ListUserTasksFilteredUseCase
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestListUserTasksFilteredUseCase:
|
||||
"""ListUserTasksFilteredUseCase 按用户筛选任务"""
|
||||
"""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
|
||||
def test_list_without_filter(self, mock_repo, sample_task):
|
||||
"""不带筛选条件查询"""
|
||||
mock_repo.list_by_user_filtered.return_value = [sample_task]
|
||||
mock_repo.count_by_user_filtered.return_value = 1
|
||||
|
||||
uc = ListUserTasksFilteredUseCase(mock_repo)
|
||||
result = uc.execute("user-1")
|
||||
use_case = ListUserTasksFilteredUseCase(mock_repo)
|
||||
result = use_case.execute("user_001")
|
||||
|
||||
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)
|
||||
assert len(result.items) == 1
|
||||
assert result.total == 1
|
||||
mock_repo.list_by_user_filtered.assert_called_once_with(
|
||||
"user_001", status=None, limit=None, offset=0
|
||||
)
|
||||
|
||||
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")
|
||||
use_case = ListUserTasksFilteredUseCase(mock_repo)
|
||||
result = use_case.execute("user_001", status="completed")
|
||||
|
||||
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")
|
||||
assert result.total == 0
|
||||
mock_repo.list_by_user_filtered.assert_called_once_with(
|
||||
"user_001", status="completed", limit=None, offset=0
|
||||
)
|
||||
|
||||
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
|
||||
mock_repo.count_by_user_filtered.return_value = 50
|
||||
|
||||
uc = ListUserTasksFilteredUseCase(mock_repo)
|
||||
result = uc.execute("user-1", limit=10, offset=20)
|
||||
use_case = ListUserTasksFilteredUseCase(mock_repo)
|
||||
result = use_case.execute("user_001", 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)
|
||||
assert result.total == 50
|
||||
mock_repo.list_by_user_filtered.assert_called_once_with(
|
||||
"user_001", status=None, limit=10, offset=20
|
||||
)
|
||||
|
||||
def test_list_empty_result(self, mock_repo):
|
||||
"""空结果"""
|
||||
def test_list_with_all_params(self, mock_repo):
|
||||
"""带所有筛选和分页参数"""
|
||||
mock_repo.list_by_user_filtered.return_value = []
|
||||
mock_repo.count_by_user_filtered.return_value = 0
|
||||
mock_repo.count_by_user_filtered.return_value = 5
|
||||
|
||||
uc = ListUserTasksFilteredUseCase(mock_repo)
|
||||
result = uc.execute("user-1", status="failed")
|
||||
use_case = ListUserTasksFilteredUseCase(mock_repo)
|
||||
use_case.execute("user_001", status="failed", limit=20, offset=0)
|
||||
|
||||
assert result.items == []
|
||||
assert result.total == 0
|
||||
|
||||
|
||||
# ============================================================
|
||||
# RetryGenerationTaskUseCase
|
||||
# ============================================================
|
||||
mock_repo.list_by_user_filtered.assert_called_once_with(
|
||||
"user_001", status="failed", limit=20, offset=0
|
||||
)
|
||||
mock_repo.count_by_user_filtered.assert_called_once_with(
|
||||
"user_001", status="failed"
|
||||
)
|
||||
|
||||
|
||||
class TestRetryGenerationTaskUseCase:
|
||||
"""RetryGenerationTaskUseCase 重试失败任务"""
|
||||
"""RetryGenerationTaskUseCase 测试"""
|
||||
|
||||
def test_retry_success(self, mock_repo):
|
||||
"""失败任务重试成功"""
|
||||
task = make_task(
|
||||
status=GenerationTaskStatus.FAILED,
|
||||
error_message="网络超时",
|
||||
retry_count=0,
|
||||
)
|
||||
def test_retry_failed_task(self, mock_repo):
|
||||
"""重试失败的任务"""
|
||||
task = MagicMock(spec=GenerationTask)
|
||||
task.is_failed = True
|
||||
mock_repo.get.return_value = task
|
||||
mock_repo.update.side_effect = lambda t: t
|
||||
mock_repo.update.return_value = task
|
||||
|
||||
uc = RetryGenerationTaskUseCase(mock_repo)
|
||||
result = uc.execute("task-1")
|
||||
use_case = RetryGenerationTaskUseCase(mock_repo)
|
||||
result = use_case.execute("task_001")
|
||||
|
||||
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()
|
||||
task.mark_pending_from_failed.assert_called_once()
|
||||
mock_repo.update.assert_called_once_with(task)
|
||||
assert result is task
|
||||
|
||||
def test_retry_not_found(self, mock_repo):
|
||||
"""任务不存在"""
|
||||
"""任务不存在抛出 ValueError"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
uc = RetryGenerationTaskUseCase(mock_repo)
|
||||
use_case = RetryGenerationTaskUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="任务不存在"):
|
||||
uc.execute("nonexistent")
|
||||
use_case.execute("nonexistent")
|
||||
|
||||
def test_retry_not_failed(self, mock_repo):
|
||||
"""非失败状态不能重试"""
|
||||
task = make_task(status=GenerationTaskStatus.RUNNING)
|
||||
mock_repo.update.assert_not_called()
|
||||
|
||||
def test_retry_non_failed_task(self, mock_repo):
|
||||
"""非失败状态的任务不能重试"""
|
||||
task = MagicMock(spec=GenerationTask)
|
||||
task.is_failed = False
|
||||
task.status = MagicMock()
|
||||
task.status.value = "running"
|
||||
mock_repo.get.return_value = task
|
||||
|
||||
uc = RetryGenerationTaskUseCase(mock_repo)
|
||||
with pytest.raises(ValueError, match="只有失败状态"):
|
||||
uc.execute("task-1")
|
||||
use_case = RetryGenerationTaskUseCase(mock_repo)
|
||||
|
||||
def test_retry_pending_not_allowed(self, mock_repo):
|
||||
"""pending 状态不能重试"""
|
||||
task = make_task(status=GenerationTaskStatus.PENDING)
|
||||
mock_repo.get.return_value = task
|
||||
with pytest.raises(ValueError, match="只有失败状态的任务才能重试"):
|
||||
use_case.execute("task_001")
|
||||
|
||||
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
|
||||
mock_repo.update.assert_not_called()
|
||||
task.mark_pending_from_failed.assert_not_called()
|
||||
|
||||
Executable
+72
@@ -0,0 +1,72 @@
|
||||
"""素材入库任务 UseCase 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.ingest_jobs import (
|
||||
SubmitIngestJobCommand,
|
||||
SubmitIngestJobUseCase,
|
||||
)
|
||||
from packages.domain import IngestJob
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
class TestSubmitIngestJobUseCase:
|
||||
"""SubmitIngestJobUseCase 测试"""
|
||||
|
||||
def test_submit_job_success(self, mock_repo):
|
||||
"""正常提交入库任务"""
|
||||
mock_repo.create.side_effect = lambda j: j
|
||||
use_case = SubmitIngestJobUseCase(mock_repo)
|
||||
|
||||
command = SubmitIngestJobCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
storage_key="videos/test.mp4",
|
||||
file_hash="abc123def",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert isinstance(result, IngestJob)
|
||||
assert result.project_id == "proj_001"
|
||||
assert result.library_id == "lib_001"
|
||||
assert result.storage_key == "videos/test.mp4"
|
||||
assert result.file_hash == "abc123def"
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_submit_job_without_hash(self, mock_repo):
|
||||
"""不传 file_hash 时默认为空"""
|
||||
mock_repo.create.side_effect = lambda j: j
|
||||
use_case = SubmitIngestJobUseCase(mock_repo)
|
||||
|
||||
command = SubmitIngestJobCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
storage_key="images/test.png",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.file_hash == ""
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_submit_job_returns_repo_result(self, mock_repo):
|
||||
"""返回 repository.create 的结果"""
|
||||
expected_job = MagicMock(spec=IngestJob)
|
||||
mock_repo.create.return_value = expected_job
|
||||
|
||||
use_case = SubmitIngestJobUseCase(mock_repo)
|
||||
command = SubmitIngestJobCommand(
|
||||
project_id="proj_001",
|
||||
library_id="lib_001",
|
||||
storage_key="test.mp4",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result is expected_job
|
||||
Executable
+169
@@ -0,0 +1,169 @@
|
||||
"""JWT Handler 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.auth.jwt_handler import (
|
||||
JWTHandler,
|
||||
configure_jwt_handler,
|
||||
get_jwt_handler,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def jwt_handler():
|
||||
return JWTHandler(
|
||||
secret_key="test-secret-key-12345",
|
||||
algorithm="HS256",
|
||||
access_token_expire_minutes=30,
|
||||
)
|
||||
|
||||
|
||||
class TestJWTHandler:
|
||||
"""JWTHandler 测试"""
|
||||
|
||||
def test_create_access_token_returns_string(self, jwt_handler):
|
||||
"""创建 access_token 返回非空字符串"""
|
||||
token = jwt_handler.create_access_token(user_id="user_001")
|
||||
|
||||
assert isinstance(token, str)
|
||||
assert len(token) > 0
|
||||
|
||||
def test_create_access_token_with_role(self, jwt_handler):
|
||||
"""创建带 role 的 access_token"""
|
||||
token = jwt_handler.create_access_token(user_id="user_001", role="admin")
|
||||
payload = jwt_handler.verify_access_token(token)
|
||||
|
||||
assert payload["sub"] == "user_001"
|
||||
assert payload["role"] == "admin"
|
||||
|
||||
def test_create_access_token_with_additional_claims(self, jwt_handler):
|
||||
"""创建带额外声明的 access_token"""
|
||||
token = jwt_handler.create_access_token(
|
||||
user_id="user_001",
|
||||
additional_claims={"email": "test@example.com", "tenant": "t1"},
|
||||
)
|
||||
payload = jwt_handler.verify_access_token(token)
|
||||
|
||||
assert payload["sub"] == "user_001"
|
||||
assert payload["email"] == "test@example.com"
|
||||
assert payload["tenant"] == "t1"
|
||||
|
||||
def test_verify_access_token_success(self, jwt_handler):
|
||||
"""验证有效 access_token"""
|
||||
token = jwt_handler.create_access_token(user_id="user_001")
|
||||
payload = jwt_handler.verify_access_token(token)
|
||||
|
||||
assert payload["sub"] == "user_001"
|
||||
assert "exp" in payload
|
||||
assert "iat" in payload
|
||||
|
||||
def test_verify_access_token_type_check(self, jwt_handler):
|
||||
"""verify_access_token 验证 token 类型为 access"""
|
||||
token = jwt_handler.create_access_token(user_id="user_001")
|
||||
payload = jwt_handler.verify_access_token(token)
|
||||
|
||||
assert payload.get("type") == "access" or "type" in payload
|
||||
|
||||
def test_verify_token_no_type_restriction(self, jwt_handler):
|
||||
"""verify_token 不限制 token 类型"""
|
||||
token = jwt_handler.create_access_token(user_id="user_001")
|
||||
payload = jwt_handler.verify_token(token)
|
||||
|
||||
assert payload["sub"] == "user_001"
|
||||
|
||||
def test_expired_token_raises_error(self):
|
||||
"""过期 token 验证失败"""
|
||||
handler = JWTHandler(
|
||||
secret_key="test-secret",
|
||||
access_token_expire_minutes=-1, # 立即过期
|
||||
)
|
||||
token = handler.create_access_token(user_id="user_001")
|
||||
|
||||
# 等待一小段时间确保过期
|
||||
time.sleep(0.1)
|
||||
|
||||
with pytest.raises(Exception):
|
||||
handler.verify_access_token(token)
|
||||
|
||||
def test_invalid_token_raises_error(self, jwt_handler):
|
||||
"""无效 token 验证失败"""
|
||||
with pytest.raises(Exception):
|
||||
jwt_handler.verify_access_token("invalid.token.here")
|
||||
|
||||
def test_empty_token_raises_error(self, jwt_handler):
|
||||
"""空字符串 token 验证失败"""
|
||||
with pytest.raises(Exception):
|
||||
jwt_handler.verify_access_token("")
|
||||
|
||||
def test_different_secret_fails_verification(self):
|
||||
"""不同密钥生成的 token 无法互相验证"""
|
||||
handler1 = JWTHandler(secret_key="secret-one")
|
||||
handler2 = JWTHandler(secret_key="secret-two")
|
||||
|
||||
token = handler1.create_access_token(user_id="user_001")
|
||||
|
||||
with pytest.raises(Exception):
|
||||
handler2.verify_access_token(token)
|
||||
|
||||
def test_custom_algorithm(self):
|
||||
"""支持自定义算法"""
|
||||
handler = JWTHandler(
|
||||
secret_key="test-secret",
|
||||
algorithm="HS256",
|
||||
)
|
||||
token = handler.create_access_token(user_id="user_001")
|
||||
payload = handler.verify_access_token(token)
|
||||
|
||||
assert payload["sub"] == "user_001"
|
||||
|
||||
def test_default_role_is_empty_string(self, jwt_handler):
|
||||
"""不传 role 时默认为空字符串"""
|
||||
token = jwt_handler.create_access_token(user_id="user_001")
|
||||
payload = jwt_handler.verify_access_token(token)
|
||||
|
||||
assert payload.get("role", "") == ""
|
||||
|
||||
|
||||
class TestGlobalJWTHandler:
|
||||
"""全局 JWT handler 配置测试"""
|
||||
|
||||
def test_configure_creates_handler(self):
|
||||
"""configure_jwt_handler 创建并返回 handler"""
|
||||
import packages.application.auth.jwt_handler as jwt_module
|
||||
|
||||
# 重置全局状态
|
||||
jwt_module._default_handler = None
|
||||
|
||||
handler = configure_jwt_handler(
|
||||
secret_key="global-secret",
|
||||
access_token_expire_minutes=60,
|
||||
)
|
||||
|
||||
assert isinstance(handler, JWTHandler)
|
||||
assert get_jwt_handler() is handler
|
||||
|
||||
def test_get_jwt_handler_without_config_raises(self):
|
||||
"""未配置时调用 get_jwt_handler 抛出 RuntimeError"""
|
||||
import packages.application.auth.jwt_handler as jwt_module
|
||||
|
||||
# 重置全局状态
|
||||
jwt_module._default_handler = None
|
||||
|
||||
with pytest.raises(RuntimeError, match="JWT handler not configured"):
|
||||
get_jwt_handler()
|
||||
|
||||
def test_configure_overwrites_existing(self):
|
||||
"""重新配置会覆盖之前的 handler"""
|
||||
import packages.application.auth.jwt_handler as jwt_module
|
||||
|
||||
jwt_module._default_handler = None
|
||||
|
||||
handler1 = configure_jwt_handler(secret_key="first-secret")
|
||||
handler2 = configure_jwt_handler(secret_key="second-secret")
|
||||
|
||||
assert handler1 is not handler2
|
||||
assert get_jwt_handler() is handler2
|
||||
+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)
|
||||
|
||||
+155
-248
@@ -1,12 +1,6 @@
|
||||
"""
|
||||
pagination 通用分页器单元测试
|
||||
"""通用分页器单元测试."""
|
||||
|
||||
覆盖:
|
||||
- PaginationParams: 默认值/边界/校验/offset/limit
|
||||
- PaginationMeta: from_params 各种边界场景
|
||||
- PaginatedResponse: create 工厂方法
|
||||
- paginate: 内存分页函数
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
@@ -18,320 +12,233 @@ from packages.application.common.pagination import (
|
||||
paginate,
|
||||
)
|
||||
|
||||
# ============================================================
|
||||
# PaginationParams
|
||||
# ============================================================
|
||||
|
||||
class TestPaginationParams:
|
||||
"""PaginationParams 测试"""
|
||||
|
||||
class TestPaginationParamsDefaults:
|
||||
"""默认值测试"""
|
||||
|
||||
def test_default_page_is_1(self):
|
||||
def test_default_values(self):
|
||||
"""默认值正确"""
|
||||
params = PaginationParams()
|
||||
assert params.page == 1
|
||||
|
||||
def test_default_page_size_is_20(self):
|
||||
params = PaginationParams()
|
||||
assert params.page_size == 20
|
||||
|
||||
def test_default_offset_is_0(self):
|
||||
params = PaginationParams()
|
||||
def test_offset_first_page(self):
|
||||
"""第一页 offset 为 0"""
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
assert params.offset == 0
|
||||
|
||||
def test_default_limit_is_20(self):
|
||||
params = PaginationParams()
|
||||
assert params.limit == 20
|
||||
def test_offset_second_page(self):
|
||||
"""第二页 offset 计算正确"""
|
||||
params = PaginationParams(page=2, page_size=20)
|
||||
assert params.offset == 20
|
||||
|
||||
def test_offset_custom_page_size(self):
|
||||
"""自定义 page_size 的 offset"""
|
||||
params = PaginationParams(page=3, page_size=10)
|
||||
assert params.offset == 20
|
||||
|
||||
class TestPaginationParamsValidation:
|
||||
"""参数校验"""
|
||||
def test_limit_equals_page_size(self):
|
||||
"""limit 等于 page_size"""
|
||||
params = PaginationParams(page_size=50)
|
||||
assert params.limit == 50
|
||||
|
||||
@pytest.mark.parametrize("page", [1, 2, 100, 9999])
|
||||
def test_valid_page_values(self, page):
|
||||
params = PaginationParams(page=page)
|
||||
assert params.page == page
|
||||
|
||||
def test_page_zero_raises(self):
|
||||
def test_page_must_be_at_least_1(self):
|
||||
"""page 不能小于 1"""
|
||||
with pytest.raises(ValidationError):
|
||||
PaginationParams(page=0)
|
||||
|
||||
def test_page_negative_raises(self):
|
||||
"""page 不能为负数"""
|
||||
with pytest.raises(ValidationError):
|
||||
PaginationParams(page=-1)
|
||||
|
||||
@pytest.mark.parametrize("page_size", [1, 20, 50, 100])
|
||||
def test_valid_page_size_values(self, page_size):
|
||||
params = PaginationParams(page_size=page_size)
|
||||
assert params.page_size == page_size
|
||||
|
||||
def test_page_size_zero_raises(self):
|
||||
def test_page_size_must_be_at_least_1(self):
|
||||
"""page_size 不能小于 1"""
|
||||
with pytest.raises(ValidationError):
|
||||
PaginationParams(page_size=0)
|
||||
|
||||
def test_page_size_negative_raises(self):
|
||||
with pytest.raises(ValidationError):
|
||||
PaginationParams(page_size=-5)
|
||||
|
||||
def test_page_size_over_100_raises(self):
|
||||
def test_page_size_max_100(self):
|
||||
"""page_size 最大 100"""
|
||||
with pytest.raises(ValidationError):
|
||||
PaginationParams(page_size=101)
|
||||
|
||||
def test_invalid_page_type_raises(self):
|
||||
with pytest.raises(ValidationError):
|
||||
PaginationParams(page="abc")
|
||||
|
||||
def test_invalid_page_size_type_raises(self):
|
||||
with pytest.raises(ValidationError):
|
||||
PaginationParams(page_size="abc")
|
||||
|
||||
|
||||
class TestPaginationParamsOffset:
|
||||
"""offset 属性计算"""
|
||||
|
||||
def test_page_1_offset_0(self):
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
assert params.offset == 0
|
||||
|
||||
def test_page_2_offset_page_size(self):
|
||||
params = PaginationParams(page=2, page_size=20)
|
||||
assert params.offset == 20
|
||||
|
||||
def test_page_3_offset_2x_page_size(self):
|
||||
params = PaginationParams(page=3, page_size=20)
|
||||
assert params.offset == 40
|
||||
|
||||
def test_page_5_page_size_10_offset_40(self):
|
||||
params = PaginationParams(page=5, page_size=10)
|
||||
assert params.offset == 40
|
||||
|
||||
def test_page_1_page_size_100_offset_0(self):
|
||||
params = PaginationParams(page=1, page_size=100)
|
||||
assert params.offset == 0
|
||||
|
||||
|
||||
class TestPaginationParamsLimit:
|
||||
"""limit 属性"""
|
||||
|
||||
def test_limit_equals_page_size(self):
|
||||
params = PaginationParams(page_size=20)
|
||||
assert params.limit == 20
|
||||
|
||||
def test_limit_1(self):
|
||||
params = PaginationParams(page_size=1)
|
||||
assert params.limit == 1
|
||||
|
||||
def test_limit_100(self):
|
||||
def test_page_size_100_is_valid(self):
|
||||
"""page_size=100 是合法的"""
|
||||
params = PaginationParams(page_size=100)
|
||||
assert params.limit == 100
|
||||
assert params.page_size == 100
|
||||
|
||||
|
||||
# ============================================================
|
||||
# PaginationMeta.from_params
|
||||
# ============================================================
|
||||
class TestPaginationMeta:
|
||||
"""PaginationMeta 测试"""
|
||||
|
||||
def test_from_params_first_page(self):
|
||||
"""第一页元数据"""
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=25)
|
||||
|
||||
class TestPaginationMetaFromParams:
|
||||
"""from_params 工厂方法"""
|
||||
assert meta.page == 1
|
||||
assert meta.page_size == 10
|
||||
assert meta.total == 25
|
||||
assert meta.total_pages == 3
|
||||
assert meta.has_next is True
|
||||
assert meta.has_prev is False
|
||||
|
||||
def test_empty_total_zero(self):
|
||||
def test_from_params_last_page(self):
|
||||
"""最后一页元数据"""
|
||||
params = PaginationParams(page=3, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=25)
|
||||
|
||||
assert meta.page == 3
|
||||
assert meta.total_pages == 3
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is True
|
||||
|
||||
def test_from_params_middle_page(self):
|
||||
"""中间页元数据"""
|
||||
params = PaginationParams(page=2, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=50)
|
||||
|
||||
assert meta.page == 2
|
||||
assert meta.total_pages == 5
|
||||
assert meta.has_next is True
|
||||
assert meta.has_prev is True
|
||||
|
||||
def test_from_params_zero_total(self):
|
||||
"""总数为 0 时"""
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=0)
|
||||
|
||||
assert meta.total == 0
|
||||
assert meta.total_pages == 0
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is False
|
||||
|
||||
def test_exactly_one_page(self):
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=20)
|
||||
def test_from_params_exact_multiple(self):
|
||||
"""总数刚好是 page_size 的整数倍"""
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=30)
|
||||
|
||||
assert meta.total_pages == 3
|
||||
|
||||
def test_from_params_single_page(self):
|
||||
"""单页即可放下所有数据"""
|
||||
params = PaginationParams(page=1, page_size=100)
|
||||
meta = PaginationMeta.from_params(params, total=50)
|
||||
|
||||
assert meta.total_pages == 1
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is False
|
||||
|
||||
def test_less_than_one_page(self):
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=15)
|
||||
assert meta.total_pages == 1
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is False
|
||||
|
||||
def test_multiple_pages_first_page(self):
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=50)
|
||||
assert meta.total_pages == 3
|
||||
assert meta.has_next is True
|
||||
assert meta.has_prev is False
|
||||
class TestPaginatedResponse:
|
||||
"""PaginatedResponse 测试"""
|
||||
|
||||
def test_multiple_pages_middle_page(self):
|
||||
params = PaginationParams(page=2, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=50)
|
||||
assert meta.total_pages == 3
|
||||
assert meta.has_next is True
|
||||
assert meta.has_prev is True
|
||||
|
||||
def test_multiple_pages_last_page(self):
|
||||
params = PaginationParams(page=3, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=50)
|
||||
assert meta.total_pages == 3
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is True
|
||||
|
||||
def test_exact_division(self):
|
||||
params = PaginationParams(page=2, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=40)
|
||||
assert meta.total_pages == 2
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is True
|
||||
|
||||
def test_non_exact_division_ceil(self):
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=41)
|
||||
assert meta.total_pages == 3
|
||||
|
||||
def test_total_1_page_size_20(self):
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=1)
|
||||
assert meta.total_pages == 1
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is False
|
||||
|
||||
def test_page_beyond_total_pages(self):
|
||||
params = PaginationParams(page=10, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=50)
|
||||
assert meta.total_pages == 3
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is True
|
||||
|
||||
def test_preserves_params_values(self):
|
||||
params = PaginationParams(page=3, page_size=15)
|
||||
meta = PaginationMeta.from_params(params, total=100)
|
||||
assert meta.page == 3
|
||||
assert meta.page_size == 15
|
||||
assert meta.total == 100
|
||||
|
||||
|
||||
# ============================================================
|
||||
# PaginatedResponse.create
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestPaginatedResponseCreate:
|
||||
"""create 工厂方法"""
|
||||
|
||||
def test_create_with_data(self):
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
def test_create_success(self):
|
||||
"""创建分页响应"""
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
data = [1, 2, 3]
|
||||
response = PaginatedResponse.create(data, params, total=100)
|
||||
assert response.data == data
|
||||
assert response.pagination.total == 100
|
||||
assert response.pagination.page == 1
|
||||
assert response.pagination.page_size == 20
|
||||
|
||||
def test_create_with_empty_data(self):
|
||||
response = PaginatedResponse.create(data, params, total=25)
|
||||
|
||||
assert response.data == [1, 2, 3]
|
||||
assert response.pagination.page == 1
|
||||
assert response.pagination.total == 25
|
||||
assert response.pagination.total_pages == 3
|
||||
|
||||
def test_create_empty_data(self):
|
||||
"""空数据分页响应"""
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
response = PaginatedResponse.create([], params, total=0)
|
||||
|
||||
assert response.data == []
|
||||
assert response.pagination.total == 0
|
||||
assert response.pagination.total_pages == 0
|
||||
|
||||
def test_create_preserves_list_type(self):
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
data = ["a", "b", "c"]
|
||||
response = PaginatedResponse.create(data, params, total=10)
|
||||
assert response.data == ["a", "b", "c"]
|
||||
assert len(response.data) == 3
|
||||
|
||||
|
||||
# ============================================================
|
||||
# paginate 函数
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestPaginateFunction:
|
||||
"""内存分页函数"""
|
||||
|
||||
def test_empty_list(self):
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
result = paginate([], params)
|
||||
assert result.data == []
|
||||
assert result.pagination.total == 0
|
||||
assert result.pagination.total_pages == 0
|
||||
"""paginate 函数测试(内存分页)"""
|
||||
|
||||
def test_first_page(self):
|
||||
items = list(range(50))
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
"""第一页分页"""
|
||||
items = list(range(30))
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
|
||||
result = paginate(items, params)
|
||||
assert result.data == list(range(20))
|
||||
assert result.pagination.total == 50
|
||||
|
||||
assert result.data == list(range(10))
|
||||
assert result.pagination.total == 30
|
||||
assert result.pagination.total_pages == 3
|
||||
assert result.pagination.has_next is True
|
||||
assert result.pagination.has_prev is False
|
||||
|
||||
def test_middle_page(self):
|
||||
items = list(range(50))
|
||||
params = PaginationParams(page=2, page_size=20)
|
||||
def test_second_page(self):
|
||||
"""第二页分页"""
|
||||
items = list(range(30))
|
||||
params = PaginationParams(page=2, page_size=10)
|
||||
|
||||
result = paginate(items, params)
|
||||
assert result.data == list(range(20, 40))
|
||||
assert result.pagination.has_next is True
|
||||
assert result.pagination.has_prev is True
|
||||
|
||||
assert result.data == list(range(10, 20))
|
||||
assert result.pagination.page == 2
|
||||
|
||||
def test_last_page(self):
|
||||
items = list(range(50))
|
||||
params = PaginationParams(page=3, page_size=20)
|
||||
"""最后一页分页"""
|
||||
items = list(range(25))
|
||||
params = PaginationParams(page=3, page_size=10)
|
||||
|
||||
result = paginate(items, params)
|
||||
assert result.data == list(range(40, 50))
|
||||
assert len(result.data) == 10
|
||||
|
||||
assert result.data == list(range(20, 25))
|
||||
assert len(result.data) == 5
|
||||
assert result.pagination.has_next is False
|
||||
assert result.pagination.has_prev is True
|
||||
|
||||
def test_empty_list(self):
|
||||
"""空列表分页"""
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
result = paginate([], params)
|
||||
|
||||
assert result.data == []
|
||||
assert result.pagination.total == 0
|
||||
assert result.pagination.total_pages == 0
|
||||
|
||||
def test_page_beyond_total(self):
|
||||
items = list(range(25))
|
||||
params = PaginationParams(page=10, page_size=20)
|
||||
result = paginate(items, params)
|
||||
assert result.data == []
|
||||
assert result.pagination.total == 25
|
||||
assert result.pagination.total_pages == 2
|
||||
|
||||
def test_page_size_larger_than_total(self):
|
||||
"""页码超出总数"""
|
||||
items = list(range(5))
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
params = PaginationParams(page=10, page_size=10)
|
||||
|
||||
result = paginate(items, params)
|
||||
assert result.data == items
|
||||
|
||||
assert result.data == []
|
||||
assert result.pagination.total == 5
|
||||
assert result.pagination.total_pages == 1
|
||||
assert result.pagination.has_next is False
|
||||
|
||||
def test_custom_page_size(self):
|
||||
"""自定义每页数量"""
|
||||
items = list(range(100))
|
||||
params = PaginationParams(page=1, page_size=50)
|
||||
|
||||
result = paginate(items, params)
|
||||
|
||||
assert len(result.data) == 50
|
||||
assert result.pagination.total_pages == 2
|
||||
|
||||
def test_single_item(self):
|
||||
items = [42]
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
"""单条数据"""
|
||||
items = ["only_one"]
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
|
||||
result = paginate(items, params)
|
||||
assert result.data == [42]
|
||||
|
||||
assert result.data == ["only_one"]
|
||||
assert result.pagination.total == 1
|
||||
assert result.pagination.total_pages == 1
|
||||
|
||||
def test_generic_type_preserved(self):
|
||||
"""泛型类型数据正确"""
|
||||
items = [{"id": 1, "name": "a"}, {"id": 2, "name": "b"}]
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
|
||||
def test_page_size_1(self):
|
||||
items = list(range(5))
|
||||
params = PaginationParams(page=3, page_size=1)
|
||||
result = paginate(items, params)
|
||||
assert result.data == [2]
|
||||
assert result.pagination.total_pages == 5
|
||||
|
||||
def test_exact_page_size(self):
|
||||
items = list(range(40))
|
||||
params = PaginationParams(page=2, page_size=20)
|
||||
result = paginate(items, params)
|
||||
assert result.data == list(range(20, 40))
|
||||
assert result.pagination.total_pages == 2
|
||||
assert result.pagination.has_next is False
|
||||
|
||||
def test_string_items(self):
|
||||
items = ["a", "b", "c", "d", "e"]
|
||||
params = PaginationParams(page=2, page_size=2)
|
||||
result = paginate(items, params)
|
||||
assert result.data == ["c", "d"]
|
||||
assert result.pagination.total == 5
|
||||
|
||||
def test_does_not_mutate_original_list(self):
|
||||
items = list(range(10))
|
||||
original = items.copy()
|
||||
params = PaginationParams(page=1, page_size=3)
|
||||
paginate(items, params)
|
||||
assert items == original
|
||||
assert len(result.data) == 2
|
||||
assert result.data[0]["id"] == 1
|
||||
|
||||
Executable
+175
@@ -0,0 +1,175 @@
|
||||
"""Password Handler 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.auth.password_handler import (
|
||||
PasswordHandler,
|
||||
configure_password_handler,
|
||||
get_password_handler,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def password_handler():
|
||||
return PasswordHandler(rounds=4) # 用低rounds加速测试
|
||||
|
||||
|
||||
class TestPasswordHandler:
|
||||
"""PasswordHandler 测试"""
|
||||
|
||||
def test_hash_password_returns_string(self, password_handler):
|
||||
"""哈希密码返回非空字符串"""
|
||||
hashed = password_handler.hash_password("MyP@ssw0rd!")
|
||||
|
||||
assert isinstance(hashed, str)
|
||||
assert len(hashed) > 0
|
||||
assert hashed != "MyP@ssw0rd!"
|
||||
|
||||
def test_hash_password_different_each_time(self, password_handler):
|
||||
"""同一密码每次哈希结果不同(加盐)"""
|
||||
h1 = password_handler.hash_password("TestPass123")
|
||||
h2 = password_handler.hash_password("TestPass123")
|
||||
|
||||
assert h1 != h2
|
||||
|
||||
def test_verify_password_correct(self, password_handler):
|
||||
"""正确密码验证通过"""
|
||||
hashed = password_handler.hash_password("CorrectPass1!")
|
||||
assert password_handler.verify_password("CorrectPass1!", hashed) is True
|
||||
|
||||
def test_verify_password_wrong(self, password_handler):
|
||||
"""错误密码验证失败"""
|
||||
hashed = password_handler.hash_password("RightPass1!")
|
||||
assert password_handler.verify_password("WrongPass1!", hashed) is False
|
||||
|
||||
def test_verify_password_empty_string(self, password_handler):
|
||||
"""空字符串密码也能正确验证(不匹配)"""
|
||||
hashed = password_handler.hash_password("SomePass1!")
|
||||
assert password_handler.verify_password("", hashed) is False
|
||||
|
||||
def test_hash_empty_password_raises(self, password_handler):
|
||||
"""空密码哈希抛出 ValueError"""
|
||||
with pytest.raises(ValueError):
|
||||
password_handler.hash_password("")
|
||||
|
||||
def test_needs_rehash_with_different_rounds(self):
|
||||
"""不同 rounds 的哈希需要重新计算"""
|
||||
handler_low = PasswordHandler(rounds=4)
|
||||
handler_high = PasswordHandler(rounds=5)
|
||||
|
||||
hashed = handler_low.hash_password("TestPass1!")
|
||||
assert handler_low.needs_rehash(hashed) is False
|
||||
assert handler_high.needs_rehash(hashed) is True
|
||||
|
||||
def test_validate_strength_strong_password(self, password_handler):
|
||||
"""强密码通过强度验证"""
|
||||
valid, error = password_handler.validate_strength("Str0ngP@ss!")
|
||||
|
||||
assert valid is True
|
||||
assert error is None
|
||||
|
||||
def test_validate_strength_too_short(self, password_handler):
|
||||
"""密码太短不通过"""
|
||||
valid, error = password_handler.validate_strength("Sh0rt!")
|
||||
|
||||
assert valid is False
|
||||
assert error is not None
|
||||
assert "长度" in error or "length" in error.lower() or "8" in error
|
||||
|
||||
def test_validate_strength_no_uppercase(self, password_handler):
|
||||
"""没有大写字母不通过"""
|
||||
valid, error = password_handler.validate_strength("lowercase1!")
|
||||
|
||||
assert valid is False
|
||||
assert error is not None
|
||||
|
||||
def test_validate_strength_no_lowercase(self, password_handler):
|
||||
"""没有小写字母不通过"""
|
||||
valid, error = password_handler.validate_strength("UPPERCASE1!")
|
||||
|
||||
assert valid is False
|
||||
assert error is not None
|
||||
|
||||
def test_validate_strength_no_digit(self, password_handler):
|
||||
"""没有数字不通过"""
|
||||
valid, error = password_handler.validate_strength("NoDigitPass!")
|
||||
|
||||
assert valid is False
|
||||
assert error is not None
|
||||
|
||||
def test_validate_strength_special_not_required(self, password_handler):
|
||||
"""默认不要求特殊字符"""
|
||||
valid, error = password_handler.validate_strength("NoSpecial1")
|
||||
|
||||
# 没有特殊字符也应该通过(require_special=False)
|
||||
assert valid is True
|
||||
assert error is None
|
||||
|
||||
def test_validate_strength_empty_string(self, password_handler):
|
||||
"""空字符串验证失败"""
|
||||
valid, error = password_handler.validate_strength("")
|
||||
|
||||
assert valid is False
|
||||
assert error is not None
|
||||
|
||||
def test_hash_and_verify_roundtrip(self, password_handler):
|
||||
"""哈希-验证完整往返"""
|
||||
passwords = [
|
||||
"Simple12",
|
||||
"C0mpl3x!Pass",
|
||||
"12345678aA",
|
||||
"user@example.com1",
|
||||
]
|
||||
for pwd in passwords:
|
||||
hashed = password_handler.hash_password(pwd)
|
||||
assert password_handler.verify_password(pwd, hashed)
|
||||
assert not password_handler.verify_password(pwd + "x", hashed)
|
||||
|
||||
|
||||
class TestGlobalPasswordHandler:
|
||||
"""全局密码处理器配置测试"""
|
||||
|
||||
def test_get_password_handler_default(self):
|
||||
"""未配置时 get_password_handler 返回默认实例"""
|
||||
import packages.application.auth.password_handler as pw_module
|
||||
|
||||
pw_module._default_handler = None
|
||||
|
||||
handler = get_password_handler()
|
||||
assert isinstance(handler, PasswordHandler)
|
||||
|
||||
def test_configure_creates_handler(self):
|
||||
"""configure_password_handler 创建并返回 handler"""
|
||||
import packages.application.auth.password_handler as pw_module
|
||||
|
||||
pw_module._default_handler = None
|
||||
|
||||
handler = configure_password_handler(rounds=4)
|
||||
|
||||
assert isinstance(handler, PasswordHandler)
|
||||
assert get_password_handler() is handler
|
||||
|
||||
def test_configure_overwrites_existing(self):
|
||||
"""重新配置会覆盖之前的 handler"""
|
||||
import packages.application.auth.password_handler as pw_module
|
||||
|
||||
pw_module._default_handler = None
|
||||
|
||||
handler1 = configure_password_handler(rounds=4)
|
||||
handler2 = configure_password_handler(rounds=5)
|
||||
|
||||
assert handler1 is not handler2
|
||||
assert get_password_handler() is handler2
|
||||
|
||||
def test_get_password_handler_lazy_init(self):
|
||||
"""未配置时首次调用 get_password_handler 会懒初始化"""
|
||||
import packages.application.auth.password_handler as pw_module
|
||||
|
||||
pw_module._default_handler = None
|
||||
|
||||
assert pw_module._default_handler is None
|
||||
handler = get_password_handler()
|
||||
assert pw_module._default_handler is not None
|
||||
assert pw_module._default_handler is handler
|
||||
+196
-215
@@ -1,269 +1,250 @@
|
||||
"""
|
||||
密码哈希工具测试
|
||||
"""
|
||||
"""密码哈希与验证器单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.auth.password_hasher import PasswordHasher, PasswordValidator
|
||||
from packages.application.auth.password_hasher import (
|
||||
PasswordHasher,
|
||||
PasswordValidator,
|
||||
password_hasher,
|
||||
password_validator,
|
||||
)
|
||||
|
||||
|
||||
class TestPasswordHasher:
|
||||
"""密码哈希测试"""
|
||||
"""PasswordHasher 测试"""
|
||||
|
||||
@pytest.fixture
|
||||
def hasher(self):
|
||||
"""创建密码哈希器"""
|
||||
return PasswordHasher(rounds=4) # 测试用低 cost,加快速度
|
||||
|
||||
def test_hash_password(self, hasher):
|
||||
"""测试密码哈希"""
|
||||
password = "MySecurePassword123"
|
||||
hashed = hasher.hash_password(password)
|
||||
def test_hash_password_returns_string(self):
|
||||
"""哈希密码返回非空字符串"""
|
||||
hasher = PasswordHasher(rounds=4)
|
||||
hashed = hasher.hash_password("TestPass1!")
|
||||
|
||||
assert isinstance(hashed, str)
|
||||
assert len(hashed) > 0
|
||||
assert hashed != password # 哈希后不等于原文
|
||||
assert hashed.startswith("$2b$") # bcrypt 格式
|
||||
assert hashed.startswith("$2") # bcrypt hash 格式
|
||||
|
||||
def test_hash_same_password_different_result(self, hasher):
|
||||
"""测试相同密码每次哈希结果不同(因为 salt 不同)"""
|
||||
password = "MySecurePassword123"
|
||||
hash1 = hasher.hash_password(password)
|
||||
hash2 = hasher.hash_password(password)
|
||||
def test_hash_password_different_salts(self):
|
||||
"""相同密码每次哈希结果不同(加盐)"""
|
||||
hasher = PasswordHasher(rounds=4)
|
||||
|
||||
assert hash1 != hash2 # salt 不同,哈希不同
|
||||
h1 = hasher.hash_password("SamePass1!")
|
||||
h2 = hasher.hash_password("SamePass1!")
|
||||
|
||||
def test_verify_correct_password(self, hasher):
|
||||
"""测试验证正确的密码"""
|
||||
password = "MySecurePassword123"
|
||||
hashed = hasher.hash_password(password)
|
||||
assert h1 != h2
|
||||
|
||||
assert hasher.verify_password(password, hashed) is True
|
||||
def test_verify_correct_password(self):
|
||||
"""正确密码验证通过"""
|
||||
hasher = PasswordHasher(rounds=4)
|
||||
hashed = hasher.hash_password("Correct1!")
|
||||
|
||||
def test_verify_incorrect_password(self, hasher):
|
||||
"""测试验证错误的密码"""
|
||||
password = "MySecurePassword123"
|
||||
hashed = hasher.hash_password(password)
|
||||
assert hasher.verify_password("Correct1!", hashed) is True
|
||||
|
||||
assert hasher.verify_password("WrongPassword", hashed) is False
|
||||
def test_verify_wrong_password(self):
|
||||
"""错误密码验证失败"""
|
||||
hasher = PasswordHasher(rounds=4)
|
||||
hashed = hasher.hash_password("Right123!")
|
||||
|
||||
def test_verify_empty_password(self, hasher):
|
||||
"""测试空密码验证"""
|
||||
hashed = hasher.hash_password("test")
|
||||
assert hasher.verify_password("Wrong123!", hashed) is False
|
||||
|
||||
assert hasher.verify_password("", hashed) is False
|
||||
def test_hash_empty_password_raises(self):
|
||||
"""空密码哈希抛出 ValueError"""
|
||||
hasher = PasswordHasher(rounds=4)
|
||||
|
||||
def test_verify_empty_hash(self, hasher):
|
||||
"""测试空哈希验证"""
|
||||
assert hasher.verify_password("test", "") is False
|
||||
|
||||
def test_verify_invalid_hash(self, hasher):
|
||||
"""测试无效的哈希"""
|
||||
assert hasher.verify_password("test", "invalid-hash") is False
|
||||
|
||||
def test_hash_empty_password(self, hasher):
|
||||
"""测试哈希空密码应该失败"""
|
||||
with pytest.raises(ValueError, match="Password cannot be empty"):
|
||||
hasher.hash_password("")
|
||||
|
||||
def test_invalid_rounds(self):
|
||||
"""测试无效的 rounds 参数"""
|
||||
def test_verify_empty_password_returns_false(self):
|
||||
"""空密码验证返回 False"""
|
||||
hasher = PasswordHasher(rounds=4)
|
||||
hashed = hasher.hash_password("TestPass1!")
|
||||
|
||||
assert hasher.verify_password("", hashed) is False
|
||||
|
||||
def test_verify_empty_hash_returns_false(self):
|
||||
"""空哈希验证返回 False"""
|
||||
hasher = PasswordHasher(rounds=4)
|
||||
|
||||
assert hasher.verify_password("TestPass1!", "") is False
|
||||
|
||||
def test_verify_invalid_hash_format(self):
|
||||
"""无效格式的哈希验证返回 False(不抛异常)"""
|
||||
hasher = PasswordHasher(rounds=4)
|
||||
|
||||
assert hasher.verify_password("TestPass1!", "not_a_valid_hash") is False
|
||||
|
||||
def test_needs_rehash_same_rounds(self):
|
||||
"""相同 rounds 不需要重新哈希"""
|
||||
hasher = PasswordHasher(rounds=4)
|
||||
hashed = hasher.hash_password("TestPass1!")
|
||||
|
||||
assert hasher.needs_rehash(hashed) is False
|
||||
|
||||
def test_needs_rehash_different_rounds(self):
|
||||
"""不同 rounds 需要重新哈希"""
|
||||
hasher_low = PasswordHasher(rounds=4)
|
||||
hasher_high = PasswordHasher(rounds=5)
|
||||
|
||||
hashed = hasher_low.hash_password("TestPass1!")
|
||||
assert hasher_high.needs_rehash(hashed) is True
|
||||
|
||||
def test_needs_rehash_invalid_hash(self):
|
||||
"""无效哈希格式返回 False(不抛异常)"""
|
||||
hasher = PasswordHasher(rounds=4)
|
||||
|
||||
assert hasher.needs_rehash("invalid_hash") is False
|
||||
|
||||
def test_rounds_too_low_raises(self):
|
||||
"""rounds 小于 4 抛出 ValueError"""
|
||||
with pytest.raises(ValueError, match="rounds must be between 4 and 31"):
|
||||
PasswordHasher(rounds=2)
|
||||
PasswordHasher(rounds=3)
|
||||
|
||||
def test_rounds_too_high_raises(self):
|
||||
"""rounds 大于 31 抛出 ValueError"""
|
||||
with pytest.raises(ValueError, match="rounds must be between 4 and 31"):
|
||||
PasswordHasher(rounds=50)
|
||||
PasswordHasher(rounds=32)
|
||||
|
||||
def test_unicode_password(self, hasher):
|
||||
"""测试 Unicode 密码"""
|
||||
password = "密码123!@#"
|
||||
hashed = hasher.hash_password(password)
|
||||
def test_rounds_boundary_values(self):
|
||||
"""rounds 边界值 4 和 31 是合法的"""
|
||||
hasher_low = PasswordHasher(rounds=4)
|
||||
hasher_high = PasswordHasher(rounds=31)
|
||||
|
||||
assert hasher.verify_password(password, hashed) is True
|
||||
assert hasher.verify_password("错误密码", hashed) is False
|
||||
assert hasher_low.rounds == 4
|
||||
assert hasher_high.rounds == 31
|
||||
|
||||
def test_hash_and_verify_various_passwords(self):
|
||||
"""多种密码的哈希-验证往返"""
|
||||
hasher = PasswordHasher(rounds=4)
|
||||
passwords = [
|
||||
"Simple12",
|
||||
"C0mpl3x!@#",
|
||||
" spaces ",
|
||||
"中文密码123",
|
||||
"a" * 50, # 50字节,在72字节限制内
|
||||
"12345678",
|
||||
]
|
||||
for pwd in passwords:
|
||||
hashed = hasher.hash_password(pwd)
|
||||
assert hasher.verify_password(pwd, hashed)
|
||||
assert not hasher.verify_password(pwd + "x", hashed)
|
||||
|
||||
|
||||
class TestPasswordValidator:
|
||||
"""密码验证器测试"""
|
||||
"""PasswordValidator 测试"""
|
||||
|
||||
@pytest.fixture
|
||||
def validator(self):
|
||||
"""创建密码验证器"""
|
||||
return PasswordValidator(
|
||||
min_length=8,
|
||||
require_uppercase=True,
|
||||
require_lowercase=True,
|
||||
require_digit=True,
|
||||
require_special=False,
|
||||
)
|
||||
def test_strong_password_passes(self):
|
||||
"""强密码通过验证"""
|
||||
validator = PasswordValidator()
|
||||
valid, error = validator.validate("Str0ngP@ss")
|
||||
|
||||
def test_valid_password(self, validator):
|
||||
"""测试有效密码"""
|
||||
valid, error = validator.validate("MyPassword123")
|
||||
assert valid is True
|
||||
assert error is None
|
||||
|
||||
def test_password_too_short(self, validator):
|
||||
"""测试密码太短"""
|
||||
valid, error = validator.validate("Pass1")
|
||||
assert valid is False
|
||||
assert "at least 8 characters" in error
|
||||
|
||||
def test_password_no_uppercase(self, validator):
|
||||
"""测试没有大写字母"""
|
||||
valid, error = validator.validate("mypassword123")
|
||||
assert valid is False
|
||||
assert "uppercase letter" in error
|
||||
|
||||
def test_password_no_lowercase(self, validator):
|
||||
"""测试没有小写字母"""
|
||||
valid, error = validator.validate("MYPASSWORD123")
|
||||
assert valid is False
|
||||
assert "lowercase letter" in error
|
||||
|
||||
def test_password_no_digit(self, validator):
|
||||
"""测试没有数字"""
|
||||
valid, error = validator.validate("MyPassword")
|
||||
assert valid is False
|
||||
assert "digit" in error
|
||||
|
||||
def test_password_with_special_chars(self):
|
||||
"""测试要求特殊字符"""
|
||||
validator = PasswordValidator(
|
||||
min_length=8,
|
||||
require_uppercase=True,
|
||||
require_lowercase=True,
|
||||
require_digit=True,
|
||||
require_special=True,
|
||||
)
|
||||
|
||||
# 没有特殊字符
|
||||
valid, error = validator.validate("MyPassword123")
|
||||
assert valid is False
|
||||
assert "special character" in error
|
||||
|
||||
# 有特殊字符
|
||||
valid, error = validator.validate("MyPassword123!")
|
||||
assert valid is True
|
||||
assert error is None
|
||||
|
||||
def test_empty_password(self, validator):
|
||||
"""测试空密码"""
|
||||
def test_empty_password_fails(self):
|
||||
"""空密码验证失败"""
|
||||
validator = PasswordValidator()
|
||||
valid, error = validator.validate("")
|
||||
|
||||
assert valid is False
|
||||
assert "cannot be empty" in error
|
||||
assert "empty" in error.lower()
|
||||
|
||||
def test_too_short_fails(self):
|
||||
"""密码太短失败"""
|
||||
validator = PasswordValidator(min_length=8)
|
||||
valid, error = validator.validate("Sh0rt!")
|
||||
|
||||
assert valid is False
|
||||
assert "at least 8" in error
|
||||
|
||||
def test_no_uppercase_fails(self):
|
||||
"""没有大写字母失败"""
|
||||
validator = PasswordValidator(require_uppercase=True)
|
||||
valid, error = validator.validate("lowercase1!")
|
||||
|
||||
assert valid is False
|
||||
assert "uppercase" in error.lower()
|
||||
|
||||
def test_no_lowercase_fails(self):
|
||||
"""没有小写字母失败"""
|
||||
validator = PasswordValidator(require_lowercase=True)
|
||||
valid, error = validator.validate("UPPERCASE1!")
|
||||
|
||||
assert valid is False
|
||||
assert "lowercase" in error.lower()
|
||||
|
||||
def test_no_digit_fails(self):
|
||||
"""没有数字失败"""
|
||||
validator = PasswordValidator(require_digit=True)
|
||||
valid, error = validator.validate("NoDigitsHere!")
|
||||
|
||||
assert valid is False
|
||||
assert "digit" in error.lower()
|
||||
|
||||
def test_no_special_not_required_passes(self):
|
||||
"""不要求特殊字符时,不含特殊字符也通过"""
|
||||
validator = PasswordValidator(require_special=False)
|
||||
valid, error = validator.validate("NoSpecial1")
|
||||
|
||||
assert valid is True
|
||||
|
||||
def test_no_special_required_fails(self):
|
||||
"""要求特殊字符时,不含特殊字符失败"""
|
||||
validator = PasswordValidator(require_special=True)
|
||||
valid, error = validator.validate("NoSpecial1")
|
||||
|
||||
assert valid is False
|
||||
assert "special" in error.lower()
|
||||
|
||||
def test_custom_min_length(self):
|
||||
"""测试自定义最小长度"""
|
||||
"""自定义最小长度"""
|
||||
validator = PasswordValidator(
|
||||
min_length=12,
|
||||
require_uppercase=False,
|
||||
require_lowercase=False,
|
||||
require_digit=False,
|
||||
)
|
||||
valid, _ = validator.validate("123456789012") # 12字符
|
||||
assert valid is True
|
||||
|
||||
valid, _ = validator.validate("12345678901") # 11字符
|
||||
assert valid is False
|
||||
|
||||
def test_all_requirements_disabled(self):
|
||||
"""所有要求都禁用时,任意非空密码都通过"""
|
||||
validator = PasswordValidator(
|
||||
min_length=1,
|
||||
require_uppercase=False,
|
||||
require_lowercase=False,
|
||||
require_digit=False,
|
||||
require_special=False,
|
||||
)
|
||||
valid, error = validator.validate("x")
|
||||
|
||||
valid, error = validator.validate("short")
|
||||
assert valid is False
|
||||
assert "at least 12 characters" in error
|
||||
|
||||
valid, error = validator.validate("longenoughpassword")
|
||||
assert valid is True
|
||||
assert error is None
|
||||
|
||||
def test_special_characters_recognized(self):
|
||||
"""各种特殊字符都被识别"""
|
||||
validator = PasswordValidator(require_special=True, require_uppercase=False, require_lowercase=False)
|
||||
specials = ["!", "@", "#", "$", "%", "^", "&", "*", "(", ")", "-", "_", "=", "+"]
|
||||
for ch in specials:
|
||||
valid, _ = validator.validate(f"abcd1234{ch}")
|
||||
assert valid is True, f"Special char '{ch}' not recognized"
|
||||
|
||||
class TestPasswordHandler:
|
||||
"""Password Handler 委托层测试"""
|
||||
|
||||
def test_hash_and_verify_password(self):
|
||||
"""测试哈希和验证密码"""
|
||||
from packages.application.auth.password_handler import PasswordHandler
|
||||
class TestGlobalInstances:
|
||||
"""全局实例测试"""
|
||||
|
||||
handler = PasswordHandler(rounds=4)
|
||||
hashed = handler.hash_password("MySecurePass123")
|
||||
def test_global_password_hasher_exists(self):
|
||||
"""全局 password_hasher 实例存在"""
|
||||
assert password_hasher is not None
|
||||
assert isinstance(password_hasher, PasswordHasher)
|
||||
assert password_hasher.rounds == 12
|
||||
|
||||
assert hashed != "MySecurePass123"
|
||||
assert len(hashed) > 20
|
||||
assert handler.verify_password("MySecurePass123", hashed) is True
|
||||
assert handler.verify_password("WrongPassword", hashed) is False
|
||||
|
||||
def test_hash_empty_password_raises(self):
|
||||
"""测试空密码抛出异常"""
|
||||
from packages.application.auth.password_handler import PasswordHandler
|
||||
|
||||
handler = PasswordHandler(rounds=4)
|
||||
with pytest.raises(ValueError):
|
||||
handler.hash_password("")
|
||||
|
||||
def test_needs_rehash(self):
|
||||
"""测试检测需要重新哈希"""
|
||||
from packages.application.auth.password_handler import PasswordHandler
|
||||
|
||||
handler = PasswordHandler(rounds=4)
|
||||
hashed = handler.hash_password("TestPass123")
|
||||
|
||||
# 相同 rounds 不需要重新哈希
|
||||
assert handler.needs_rehash(hashed) is False
|
||||
|
||||
# 用更高 rounds 的 handler 检查,应该需要重新哈希
|
||||
# 注意:bcrypt 的 rounds 体现在 hash 中,这里用不同 rounds 测试
|
||||
high_rounds_handler = PasswordHandler(rounds=5)
|
||||
# 低 rounds 的 hash 在高 rounds 配置下应该需要 rehash
|
||||
assert high_rounds_handler.needs_rehash(hashed) is True
|
||||
|
||||
def test_validate_strength(self):
|
||||
"""测试密码强度验证"""
|
||||
from packages.application.auth.password_handler import PasswordHandler
|
||||
|
||||
handler = PasswordHandler(rounds=4)
|
||||
|
||||
# 弱密码
|
||||
valid, error = handler.validate_strength("weak")
|
||||
assert valid is False
|
||||
assert error is not None
|
||||
|
||||
# 强密码
|
||||
valid, error = handler.validate_strength("StrongPass123")
|
||||
assert valid is True
|
||||
assert error is None
|
||||
|
||||
def test_configure_and_get_default_handler(self):
|
||||
"""测试配置和获取全局默认 handler"""
|
||||
from packages.application.auth import password_handler as handler_module
|
||||
from packages.application.auth.password_handler import (
|
||||
configure_password_handler,
|
||||
get_password_handler,
|
||||
)
|
||||
|
||||
# 重置全局状态
|
||||
handler_module._default_handler = None
|
||||
|
||||
# 配置
|
||||
handler = configure_password_handler(rounds=4)
|
||||
assert handler is not None
|
||||
|
||||
# 获取
|
||||
same_handler = get_password_handler()
|
||||
assert same_handler is handler
|
||||
|
||||
# 验证能正常工作
|
||||
hashed = same_handler.hash_password("TestPass123")
|
||||
assert same_handler.verify_password("TestPass123", hashed) is True
|
||||
|
||||
# 重置全局状态,避免影响其他测试
|
||||
handler_module._default_handler = None
|
||||
|
||||
def test_get_password_handler_auto_creates_default(self):
|
||||
"""测试未配置时获取 handler 会自动创建默认实例"""
|
||||
from packages.application.auth import password_handler as handler_module
|
||||
from packages.application.auth.password_handler import get_password_handler
|
||||
|
||||
# 重置全局状态
|
||||
handler_module._default_handler = None
|
||||
|
||||
# 自动创建默认实例
|
||||
handler = get_password_handler()
|
||||
assert handler is not None
|
||||
|
||||
# 重置
|
||||
handler_module._default_handler = None
|
||||
def test_global_password_validator_exists(self):
|
||||
"""全局 password_validator 实例存在"""
|
||||
assert password_validator is not None
|
||||
assert isinstance(password_validator, PasswordValidator)
|
||||
assert password_validator.min_length == 8
|
||||
assert password_validator.require_uppercase is True
|
||||
assert password_validator.require_special is False
|
||||
|
||||
Regular → Executable
+232
-144
@@ -1,9 +1,9 @@
|
||||
"""
|
||||
密码重置 Use Case 测试
|
||||
"""
|
||||
"""密码重置 UseCase 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import Mock
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -16,196 +16,284 @@ from packages.application.auth.password_reset_use_case import (
|
||||
from packages.domain.entities import User
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_user_repo():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_email_service():
|
||||
svc = MagicMock()
|
||||
svc.send_password_reset_email.return_value = (True, None)
|
||||
return svc
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_user():
|
||||
user = User(
|
||||
id="user_001",
|
||||
email="user@example.com",
|
||||
display_name="测试用户",
|
||||
username="testuser",
|
||||
password_hash="old_hash",
|
||||
)
|
||||
user.password_reset_token = None
|
||||
user.password_reset_expires_at = None
|
||||
return user
|
||||
|
||||
|
||||
class TestRequestPasswordResetRequest:
|
||||
"""RequestPasswordResetRequest 测试"""
|
||||
|
||||
def test_email_lowercased_and_stripped(self):
|
||||
"""邮箱转小写并去空格"""
|
||||
req = RequestPasswordResetRequest(" User@Example.COM ")
|
||||
assert req.email == "user@example.com"
|
||||
|
||||
def test_empty_email(self):
|
||||
"""空邮箱"""
|
||||
req = RequestPasswordResetRequest("")
|
||||
assert req.email == ""
|
||||
|
||||
|
||||
class TestRequestPasswordResetUseCase:
|
||||
"""请求密码重置测试"""
|
||||
"""RequestPasswordResetUseCase 测试"""
|
||||
|
||||
@pytest.fixture
|
||||
def mock_user_repo(self):
|
||||
repo = Mock()
|
||||
repo.find_by_email = Mock(return_value=None)
|
||||
repo.save = Mock()
|
||||
return repo
|
||||
def test_request_success(self, mock_user_repo, mock_email_service, sample_user):
|
||||
"""请求重置成功"""
|
||||
mock_user_repo.find_by_email.return_value = sample_user
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
@pytest.fixture
|
||||
def use_case(self, mock_user_repo):
|
||||
email_service = Mock()
|
||||
email_service.send_password_reset_email.return_value = (True, None)
|
||||
return RequestPasswordResetUseCase(
|
||||
user_repository=mock_user_repo,
|
||||
base_url="https://test.com",
|
||||
token_expire_hours=1,
|
||||
email_service=email_service,
|
||||
use_case = RequestPasswordResetUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RequestPasswordResetRequest("user@example.com")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
@pytest.fixture
|
||||
def test_user(self):
|
||||
return User(
|
||||
id="user-123",
|
||||
email="test@example.com",
|
||||
username="testuser",
|
||||
display_name="Test User",
|
||||
password_hash="hash",
|
||||
assert success is True
|
||||
assert error is None
|
||||
assert sample_user.password_reset_token is not None
|
||||
assert len(sample_user.password_reset_token) > 0
|
||||
assert sample_user.password_reset_expires_at is not None
|
||||
mock_user_repo.save.assert_called_once()
|
||||
mock_email_service.send_password_reset_email.assert_called_once()
|
||||
|
||||
def test_request_user_not_found_returns_success(self, mock_user_repo, mock_email_service):
|
||||
"""用户不存在也返回成功(安全考虑,不暴露用户存在性)"""
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
|
||||
use_case = RequestPasswordResetUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RequestPasswordResetRequest("nonexistent@example.com")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
def test_request_reset_success(self, use_case, mock_user_repo, test_user):
|
||||
"""测试请求重置成功"""
|
||||
mock_user_repo.find_by_email.return_value = test_user
|
||||
assert success is True
|
||||
assert error is None
|
||||
mock_user_repo.save.assert_not_called()
|
||||
mock_email_service.send_password_reset_email.assert_not_called()
|
||||
|
||||
request = RequestPasswordResetRequest(email="test@example.com")
|
||||
def test_request_empty_email_returns_error(self, mock_user_repo, mock_email_service):
|
||||
"""空邮箱返回错误"""
|
||||
use_case = RequestPasswordResetUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RequestPasswordResetRequest("")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is False
|
||||
assert "Email is required" in error
|
||||
|
||||
def test_reset_token_expiry_set(self, mock_user_repo, mock_email_service, sample_user):
|
||||
"""重置令牌过期时间正确设置"""
|
||||
mock_user_repo.find_by_email.return_value = sample_user
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
use_case = RequestPasswordResetUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
token_expire_hours=2,
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RequestPasswordResetRequest("user@example.com")
|
||||
use_case.execute(request)
|
||||
|
||||
assert sample_user.password_reset_expires_at is not None
|
||||
# 过期时间应该在约2小时后
|
||||
expected = datetime.now(timezone.utc) + timedelta(hours=2)
|
||||
diff = abs((sample_user.password_reset_expires_at - expected).total_seconds())
|
||||
assert diff < 10 # 允许10秒误差
|
||||
|
||||
def test_email_contains_reset_url(self, mock_user_repo, mock_email_service, sample_user):
|
||||
"""重置邮件包含正确的重置链接"""
|
||||
mock_user_repo.find_by_email.return_value = sample_user
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
use_case = RequestPasswordResetUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://app.example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RequestPasswordResetRequest("user@example.com")
|
||||
use_case.execute(request)
|
||||
|
||||
call_args = mock_email_service.send_password_reset_email.call_args
|
||||
reset_url = call_args[1]["reset_url"] if "reset_url" in call_args[1] else call_args[0][2]
|
||||
assert "https://app.example.com/reset-password?token=" in reset_url
|
||||
|
||||
def test_email_failure_does_not_affect_result(self, mock_user_repo, mock_email_service, sample_user):
|
||||
"""邮件发送失败不影响返回结果(安全考虑)"""
|
||||
mock_user_repo.find_by_email.return_value = sample_user
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
mock_email_service.send_password_reset_email.return_value = (False, "SMTP error")
|
||||
|
||||
use_case = RequestPasswordResetUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RequestPasswordResetRequest("user@example.com")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is True
|
||||
assert error is None
|
||||
|
||||
# 验证保存了用户
|
||||
mock_user_repo.save.assert_called_once()
|
||||
saved_user = mock_user_repo.save.call_args[0][0]
|
||||
assert saved_user.password_reset_token is not None
|
||||
assert saved_user.password_reset_expires_at is not None
|
||||
def test_different_tokens_each_time(self, mock_user_repo, mock_email_service, sample_user):
|
||||
"""每次请求生成不同的 token"""
|
||||
mock_user_repo.find_by_email.return_value = sample_user
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
# 验证发送了邮件
|
||||
use_case.email_service.send_password_reset_email.assert_called_once()
|
||||
use_case = RequestPasswordResetUseCase(
|
||||
mock_user_repo,
|
||||
base_url="https://example.com",
|
||||
email_service=mock_email_service,
|
||||
)
|
||||
request = RequestPasswordResetRequest("user@example.com")
|
||||
|
||||
def test_request_reset_user_not_exists(self, use_case, mock_user_repo):
|
||||
"""测试用户不存在(仍返回成功,避免暴露)"""
|
||||
mock_user_repo.find_by_email.return_value = None
|
||||
use_case.execute(request)
|
||||
token1 = sample_user.password_reset_token
|
||||
|
||||
request = RequestPasswordResetRequest(email="nonexistent@example.com")
|
||||
success, error = use_case.execute(request)
|
||||
use_case.execute(request)
|
||||
token2 = sample_user.password_reset_token
|
||||
|
||||
assert success is True # 安全考虑,仍返回成功
|
||||
assert error is None
|
||||
assert token1 != token2
|
||||
|
||||
# 不发送邮件
|
||||
use_case.email_service.send_password_reset_email.assert_not_called()
|
||||
|
||||
def test_request_reset_missing_email(self, use_case):
|
||||
"""测试缺少邮箱"""
|
||||
request = RequestPasswordResetRequest(email="")
|
||||
success, error = use_case.execute(request)
|
||||
class TestResetPasswordRequest:
|
||||
"""ResetPasswordRequest 测试"""
|
||||
|
||||
assert success is False
|
||||
assert error == "Email is required"
|
||||
def test_stores_token_and_password(self):
|
||||
"""正确存储 token 和新密码"""
|
||||
req = ResetPasswordRequest(token="abc123", new_password="NewPass1!")
|
||||
assert req.token == "abc123"
|
||||
assert req.new_password == "NewPass1!"
|
||||
|
||||
|
||||
class TestResetPasswordUseCase:
|
||||
"""重置密码测试"""
|
||||
"""ResetPasswordUseCase 测试"""
|
||||
|
||||
@pytest.fixture
|
||||
def mock_user_repo(self):
|
||||
repo = Mock()
|
||||
repo.find_by_password_reset_token = Mock(return_value=None)
|
||||
repo.save = Mock()
|
||||
return repo
|
||||
def test_reset_success(self, mock_user_repo, sample_user):
|
||||
"""重置密码成功"""
|
||||
sample_user.password_reset_token = "valid_token"
|
||||
sample_user.password_reset_expires_at = datetime.now(timezone.utc) + timedelta(hours=1)
|
||||
mock_user_repo.find_by_password_reset_token.return_value = sample_user
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
@pytest.fixture
|
||||
def use_case(self, mock_user_repo):
|
||||
return ResetPasswordUseCase(user_repository=mock_user_repo)
|
||||
|
||||
@pytest.fixture
|
||||
def test_user(self):
|
||||
return User(
|
||||
id="user-123",
|
||||
email="test@example.com",
|
||||
username="testuser",
|
||||
display_name="Test User",
|
||||
password_hash="old-hash",
|
||||
password_reset_token="valid-token",
|
||||
password_reset_expires_at=datetime.now(timezone.utc) + timedelta(hours=1),
|
||||
)
|
||||
|
||||
def test_reset_password_success(self, use_case, mock_user_repo, test_user):
|
||||
"""测试重置密码成功"""
|
||||
mock_user_repo.find_by_password_reset_token.return_value = test_user
|
||||
|
||||
request = ResetPasswordRequest(
|
||||
token="valid-token",
|
||||
new_password="NewSecurePass123",
|
||||
)
|
||||
use_case = ResetPasswordUseCase(mock_user_repo)
|
||||
request = ResetPasswordRequest(token="valid_token", new_password="NewSecurePass1!")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is True
|
||||
assert error is None
|
||||
|
||||
# 验证密码已更新
|
||||
assert test_user.password_hash != "old-hash"
|
||||
assert test_user.password_reset_token is None
|
||||
assert test_user.password_reset_expires_at is None
|
||||
|
||||
# 验证保存了用户
|
||||
assert sample_user.password_reset_token is None
|
||||
assert sample_user.password_reset_expires_at is None
|
||||
assert sample_user.password_hash != "old_hash"
|
||||
mock_user_repo.save.assert_called_once()
|
||||
|
||||
def test_reset_password_success_with_naive_database_datetime(self, use_case, mock_user_repo, test_user):
|
||||
"""测试数据库返回 naive datetime 时仍可重置密码"""
|
||||
test_user.password_reset_expires_at = (datetime.now(timezone.utc) + timedelta(hours=1)).replace(tzinfo=None)
|
||||
mock_user_repo.find_by_password_reset_token.return_value = test_user
|
||||
|
||||
success, error = use_case.execute(ResetPasswordRequest(token="valid-token", new_password="NewSecurePass123"))
|
||||
|
||||
assert success is True
|
||||
assert error is None
|
||||
mock_user_repo.save.assert_called_once()
|
||||
|
||||
def test_reset_password_weak_password(self, use_case, mock_user_repo, test_user):
|
||||
"""测试弱密码"""
|
||||
mock_user_repo.find_by_password_reset_token.return_value = test_user
|
||||
|
||||
request = ResetPasswordRequest(
|
||||
token="valid-token",
|
||||
new_password="weak",
|
||||
)
|
||||
def test_reset_empty_token(self, mock_user_repo):
|
||||
"""空 token 返回错误"""
|
||||
use_case = ResetPasswordUseCase(mock_user_repo)
|
||||
request = ResetPasswordRequest(token="", new_password="NewPass1!")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is False
|
||||
assert "at least 8 characters" in error
|
||||
assert "Reset token is required" in error
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_reset_password_invalid_token(self, use_case, mock_user_repo):
|
||||
"""测试无效令牌"""
|
||||
def test_reset_empty_password(self, mock_user_repo):
|
||||
"""空密码返回错误"""
|
||||
use_case = ResetPasswordUseCase(mock_user_repo)
|
||||
request = ResetPasswordRequest(token="sometoken", new_password="")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is False
|
||||
assert "New password is required" in error
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_reset_weak_password(self, mock_user_repo):
|
||||
"""弱密码返回错误"""
|
||||
use_case = ResetPasswordUseCase(mock_user_repo)
|
||||
request = ResetPasswordRequest(token="sometoken", new_password="weak")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is False
|
||||
assert error is not None
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_reset_invalid_token(self, mock_user_repo):
|
||||
"""无效 token 返回错误"""
|
||||
mock_user_repo.find_by_password_reset_token.return_value = None
|
||||
|
||||
request = ResetPasswordRequest(
|
||||
token="invalid-token",
|
||||
new_password="NewSecurePass123",
|
||||
)
|
||||
use_case = ResetPasswordUseCase(mock_user_repo)
|
||||
request = ResetPasswordRequest(token="invalid_token", new_password="NewPass1!")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is False
|
||||
assert error == "Invalid or expired reset token"
|
||||
assert "Invalid or expired" in error
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_reset_password_expired_token(self, use_case, mock_user_repo, test_user):
|
||||
"""测试过期令牌"""
|
||||
test_user.password_reset_expires_at = datetime.now(timezone.utc) - timedelta(hours=1)
|
||||
mock_user_repo.find_by_password_reset_token.return_value = test_user
|
||||
def test_reset_expired_token(self, mock_user_repo, sample_user):
|
||||
"""过期 token 返回错误"""
|
||||
sample_user.password_reset_token = "expired_token"
|
||||
sample_user.password_reset_expires_at = datetime.now(timezone.utc) - timedelta(hours=1)
|
||||
mock_user_repo.find_by_password_reset_token.return_value = sample_user
|
||||
|
||||
request = ResetPasswordRequest(
|
||||
token="valid-token",
|
||||
new_password="NewSecurePass123",
|
||||
)
|
||||
use_case = ResetPasswordUseCase(mock_user_repo)
|
||||
request = ResetPasswordRequest(token="expired_token", new_password="NewPass1!")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is False
|
||||
assert error == "Reset token has expired"
|
||||
assert "expired" in error.lower()
|
||||
mock_user_repo.save.assert_not_called()
|
||||
|
||||
def test_reset_password_missing_token(self, use_case):
|
||||
"""测试缺少令牌"""
|
||||
request = ResetPasswordRequest(
|
||||
token="",
|
||||
new_password="NewSecurePass123",
|
||||
)
|
||||
def test_reset_naive_datetime_treated_as_utc(self, mock_user_repo, sample_user):
|
||||
"""无时区的过期时间按 UTC 处理"""
|
||||
sample_user.password_reset_token = "naive_token"
|
||||
# 用无时区的时间,设置为过去
|
||||
sample_user.password_reset_expires_at = datetime.utcnow() - timedelta(hours=1)
|
||||
mock_user_repo.find_by_password_reset_token.return_value = sample_user
|
||||
|
||||
use_case = ResetPasswordUseCase(mock_user_repo)
|
||||
request = ResetPasswordRequest(token="naive_token", new_password="NewPass1!")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is False
|
||||
assert error == "Reset token is required"
|
||||
assert "expired" in error.lower()
|
||||
|
||||
def test_reset_password_missing_password(self, use_case, mock_user_repo, test_user):
|
||||
"""测试缺少新密码"""
|
||||
mock_user_repo.find_by_password_reset_token.return_value = test_user
|
||||
def test_reset_no_expiry_set(self, mock_user_repo, sample_user):
|
||||
"""没有设置过期时间的 token 可以使用"""
|
||||
sample_user.password_reset_token = "no_expiry_token"
|
||||
sample_user.password_reset_expires_at = None
|
||||
mock_user_repo.find_by_password_reset_token.return_value = sample_user
|
||||
|
||||
request = ResetPasswordRequest(
|
||||
token="valid-token",
|
||||
new_password="",
|
||||
)
|
||||
use_case = ResetPasswordUseCase(mock_user_repo)
|
||||
request = ResetPasswordRequest(token="no_expiry_token", new_password="NewPass1!")
|
||||
success, error = use_case.execute(request)
|
||||
|
||||
assert success is False
|
||||
assert error == "New password is required"
|
||||
assert success is True
|
||||
|
||||
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
|
||||
|
||||
+114
-282
@@ -1,315 +1,147 @@
|
||||
"""
|
||||
text_splitter 长文本分段工具单元测试
|
||||
"""文本分段工具单元测试."""
|
||||
|
||||
覆盖:
|
||||
- 空文本 / 短文本
|
||||
- 句子边界分段(。!?;\n . ! ? ;)
|
||||
- 超长句子硬切
|
||||
- 过短段落合并
|
||||
- max_chars 参数
|
||||
- 中英文混合
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.tts_job.text_splitter import split_text
|
||||
|
||||
# ============================================================
|
||||
# 基础场景
|
||||
# ============================================================
|
||||
|
||||
class TestSplitText:
|
||||
"""split_text 函数测试"""
|
||||
|
||||
class TestBasicCases:
|
||||
"""基础场景"""
|
||||
|
||||
def test_empty_text_returns_empty_list(self):
|
||||
def test_empty_string_returns_empty_list(self):
|
||||
"""空字符串返回空列表"""
|
||||
assert split_text("") == []
|
||||
|
||||
def test_whitespace_only_returns_empty(self):
|
||||
assert split_text(" \n\n ") == []
|
||||
def test_whitespace_only_returns_empty_list(self):
|
||||
"""纯空白字符返回空列表"""
|
||||
assert split_text(" \n \t ") == []
|
||||
|
||||
def test_short_text_single_segment(self):
|
||||
def test_short_text_returns_single_segment(self):
|
||||
"""短文本直接返回单段"""
|
||||
text = "这是一段短文本。"
|
||||
result = split_text(text, max_chars=500)
|
||||
assert result == [text]
|
||||
|
||||
def test_exactly_max_chars_single_segment(self):
|
||||
text = "a" * 500
|
||||
result = split_text(text, max_chars=500)
|
||||
def test_text_length_equals_max_chars(self):
|
||||
"""文本长度恰好等于 max_chars 时返回单段"""
|
||||
text = "a" * 100
|
||||
result = split_text(text, max_chars=100)
|
||||
assert len(result) == 1
|
||||
assert len(result[0]) == 500
|
||||
assert len(result[0]) == 100
|
||||
|
||||
def test_text_stripped(self):
|
||||
text = " 你好世界。 "
|
||||
result = split_text(text, max_chars=500)
|
||||
assert result == ["你好世界。"]
|
||||
def test_splits_on_sentence_boundary(self):
|
||||
"""在句子边界处分段"""
|
||||
# 构造长文本,确保超过 max_chars
|
||||
sentences = ["今天天气真好。我们一起去公园散步吧。", "公园里有很多花。还有很多小朋友在玩耍。"] * 10
|
||||
text = "".join(sentences)
|
||||
|
||||
result = split_text(text, max_chars=200)
|
||||
|
||||
# ============================================================
|
||||
# 句子边界分段
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestSentenceBoundarySplitting:
|
||||
"""句子边界分段"""
|
||||
|
||||
def test_split_by_chinese_period(self):
|
||||
text = "第一句。第二句。第三句。"
|
||||
# 三句都很短,应该合并成一段
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
|
||||
def test_split_by_chinese_period_long_text(self):
|
||||
"""多段长句子,按句号分段"""
|
||||
sentence1 = "我是第一句" + "啊" * 100 + "。"
|
||||
sentence2 = "我是第二句" + "哦" * 100 + "。"
|
||||
sentence3 = "我是第三句" + "嗯" * 100 + "。"
|
||||
text = sentence1 + sentence2 + sentence3
|
||||
|
||||
result = split_text(text, max_chars=150)
|
||||
# 每句106字符,超过150的阈值?不,106<150
|
||||
# 但累计到一定程度会切
|
||||
assert len(result) >= 2
|
||||
# 每段都不超过 max_chars
|
||||
for seg in result:
|
||||
assert len(seg) <= 150
|
||||
|
||||
def test_split_by_question_mark(self):
|
||||
text = "你是谁?你从哪里来?你要到哪里去?"
|
||||
result = split_text(text, max_chars=500)
|
||||
# 三句都很短,合并成一段
|
||||
assert len(result) == 1
|
||||
|
||||
def test_split_by_exclamation_mark(self):
|
||||
text = "太棒了!太厉害了!太牛了!"
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
|
||||
def test_split_by_newline(self):
|
||||
text = "第一段\n第二段\n第三段"
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
|
||||
def test_split_by_semicolon(self):
|
||||
text = "第一部分;第二部分;第三部分。"
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
|
||||
def test_mixed_punctuation(self):
|
||||
"""混合标点符号的句子边界"""
|
||||
parts = []
|
||||
for i in range(20):
|
||||
parts.append(f"第{i}句的内容" + "字" * 30 + "。")
|
||||
text = "".join(parts)
|
||||
|
||||
result = split_text(text, max_chars=200)
|
||||
# 每句约35字符,200字符大约能放5-6句
|
||||
assert len(result) >= 2
|
||||
for seg in result:
|
||||
assert len(seg) <= 200
|
||||
|
||||
def test_english_period_splitting(self):
|
||||
text = "Hello. How are you. I am fine."
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
def test_all_segments_within_max_chars(self):
|
||||
"""所有分段都不超过 max_chars"""
|
||||
text = "这是第一句话。这是第二句话。这是第三句话。这是第四句话。这是第五句话。" * 10
|
||||
|
||||
def test_english_question(self):
|
||||
text = "What? Why? How?"
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
|
||||
|
||||
# ============================================================
|
||||
# 超长硬切
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestLongSentenceHardCut:
|
||||
"""超长句子硬切"""
|
||||
|
||||
def test_single_very_long_sentence_hard_cut(self):
|
||||
"""单个超长句子,没有标点,硬切"""
|
||||
text = "字" * 1000
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 2
|
||||
assert len(result[0]) == 500
|
||||
assert len(result[1]) == 500
|
||||
|
||||
def test_three_times_max_chars(self):
|
||||
text = "字" * 1500
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 3
|
||||
for seg in result:
|
||||
assert len(seg) == 500
|
||||
|
||||
def test_not_exact_multiple(self):
|
||||
text = "字" * 1250
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 3
|
||||
assert len(result[0]) == 500
|
||||
assert len(result[1]) == 500
|
||||
assert len(result[2]) == 250
|
||||
|
||||
def test_all_segments_within_limit(self):
|
||||
"""所有段都不超过 max_chars"""
|
||||
import random
|
||||
|
||||
random.seed(42)
|
||||
# 生成随机长度的文本
|
||||
text = "".join(random.choices("字字字字。!?;\n", k=5000))
|
||||
for max_chars in [100, 200, 500]:
|
||||
result = split_text(text, max_chars=max_chars)
|
||||
for i, seg in enumerate(result):
|
||||
assert len(seg) <= max_chars, f"Segment {i} length {len(seg)} > {max_chars}"
|
||||
|
||||
|
||||
# ============================================================
|
||||
# 过短段落合并
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestShortSegmentMerging:
|
||||
"""过短段落合并"""
|
||||
|
||||
def test_short_final_segment_merged(self):
|
||||
"""最后一段过短,应该合并到前一段"""
|
||||
# 构造:前一段接近上限,后一段很短
|
||||
long_part = "字" * 480 + "。"
|
||||
short_part = "好的。"
|
||||
text = long_part + short_part
|
||||
|
||||
result = split_text(text, max_chars=500)
|
||||
# 两段加起来 481+3=484 < 500,可能合并
|
||||
# 但要看具体实现...
|
||||
# 至少验证所有段不超长
|
||||
for seg in result:
|
||||
assert len(seg) <= 500
|
||||
|
||||
def test_multiple_short_segments(self):
|
||||
"""多个短段落应该合并"""
|
||||
sentences = ["你好。", "我好。", "大家好。", "今天天气不错。", "适合出去玩。"]
|
||||
text = "".join(sentences)
|
||||
result = split_text(text, max_chars=500)
|
||||
# 5个短句子,应该合并成一段
|
||||
assert len(result) == 1
|
||||
|
||||
|
||||
# ============================================================
|
||||
# max_chars 参数
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestMaxCharsParameter:
|
||||
"""max_chars 参数"""
|
||||
|
||||
def test_small_max_chars(self):
|
||||
text = "一二三四五六七八九十一二三四五六七八九十。"
|
||||
result = split_text(text, max_chars=10)
|
||||
# 应该被切成多段
|
||||
assert len(result) >= 2
|
||||
for seg in result:
|
||||
assert len(seg) <= 10
|
||||
|
||||
def test_custom_max_chars_200(self):
|
||||
text = "测试文本" * 100 # 400字符
|
||||
result = split_text(text, max_chars=200)
|
||||
assert len(result) == 2
|
||||
assert len(result[0]) == 200
|
||||
assert len(result[1]) == 200
|
||||
|
||||
def test_very_small_max_chars(self):
|
||||
text = "abcdefghij"
|
||||
result = split_text(text, max_chars=3)
|
||||
assert len(result) >= 3
|
||||
for seg in result:
|
||||
assert len(seg) <= 3
|
||||
|
||||
|
||||
# ============================================================
|
||||
# 中英文混合
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestMixedContent:
|
||||
"""中英文混合内容"""
|
||||
|
||||
def test_chinese_english_mixed(self):
|
||||
text = "今天天气很好,Today is sunny. 我们去公园玩吧!Let's go to the park."
|
||||
result = split_text(text, max_chars=500)
|
||||
assert len(result) == 1
|
||||
assert result[0] == text.strip()
|
||||
|
||||
def test_mixed_long_text(self):
|
||||
parts = []
|
||||
for i in range(50):
|
||||
parts.append(f"第{i}段中文内容" + "字" * 20 + ". English part " + "word " * 10 + "。")
|
||||
text = "".join(parts)
|
||||
|
||||
result = split_text(text, max_chars=300)
|
||||
assert len(result) >= 2
|
||||
for seg in result:
|
||||
assert len(seg) <= 300
|
||||
|
||||
|
||||
# ============================================================
|
||||
# 输出完整性
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestOutputIntegrity:
|
||||
"""输出完整性验证"""
|
||||
|
||||
def test_combined_length_equals_original(self):
|
||||
"""所有段拼接起来(去掉空段)应该等于原文长度"""
|
||||
text = "这是第一段。这是第二段。这是第三段。这是第四段。这是第五段。" * 20
|
||||
result = split_text(text, max_chars=100)
|
||||
combined = "".join(result)
|
||||
# 由于 strip 可能去掉一些空格,原文也 strip 比较
|
||||
assert len(combined) == len(text.strip())
|
||||
|
||||
def test_order_preserved(self):
|
||||
"""分段后再拼接,文本顺序不变"""
|
||||
text = "第一。第二。第三。第四。第五。" * 10
|
||||
result = split_text(text, max_chars=50)
|
||||
combined = "".join(result)
|
||||
assert combined == text.strip()
|
||||
|
||||
def test_no_empty_strings_in_result(self):
|
||||
"""结果中没有空字符串"""
|
||||
text = "句子一。句子二。句子三。"
|
||||
result = split_text(text, max_chars=10)
|
||||
for seg in result:
|
||||
assert seg != ""
|
||||
assert len(seg) > 0
|
||||
assert len(seg) <= 100
|
||||
|
||||
def test_long_single_sentence_hard_cut(self):
|
||||
"""超长单句会被硬切"""
|
||||
text = "a" * 1000 # 没有标点
|
||||
|
||||
# ============================================================
|
||||
# 边界情况
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestEdgeCases:
|
||||
"""边界情况"""
|
||||
|
||||
def test_single_character(self):
|
||||
assert split_text("一", max_chars=500) == ["一"]
|
||||
|
||||
def test_only_punctuation(self):
|
||||
text = "。。。。。"
|
||||
result = split_text(text, max_chars=500)
|
||||
# 都是标点,也算文本
|
||||
assert len(result) == 1
|
||||
|
||||
def test_only_newlines(self):
|
||||
text = "\n\n\n"
|
||||
result = split_text(text, max_chars=500)
|
||||
assert result == []
|
||||
|
||||
def test_long_text_many_sentences(self):
|
||||
"""大量句子的长文本"""
|
||||
sentences = [f"第{i}句的完整内容。" for i in range(100)]
|
||||
text = "".join(sentences)
|
||||
result = split_text(text, max_chars=200)
|
||||
assert len(result) >= 5
|
||||
|
||||
assert len(result) > 1
|
||||
for seg in result:
|
||||
assert len(seg) <= 200
|
||||
|
||||
def test_newline_is_sentence_end(self):
|
||||
"""换行符作为句子结束符"""
|
||||
text = "第一行内容\n第二行内容\n第三行内容" * 10
|
||||
|
||||
result = split_text(text, max_chars=50)
|
||||
|
||||
assert len(result) > 1
|
||||
for seg in result:
|
||||
assert len(seg) <= 50
|
||||
|
||||
def test_chinese_punctuation(self):
|
||||
"""中文标点(。!?;)作为句子结束符"""
|
||||
text = "你好!今天吃什么?我吃米饭;你呢?我也吃米饭。" * 10
|
||||
|
||||
result = split_text(text, max_chars=80)
|
||||
|
||||
for seg in result:
|
||||
assert len(seg) <= 80
|
||||
|
||||
def test_english_punctuation(self):
|
||||
"""英文标点(.!?;)作为句子结束符"""
|
||||
text = "Hello! How are you? I'm fine; thank you. Good bye." * 10
|
||||
|
||||
result = split_text(text, max_chars=80)
|
||||
|
||||
for seg in result:
|
||||
assert len(seg) <= 80
|
||||
|
||||
def test_merged_short_segments(self):
|
||||
"""过短的段落会被合并"""
|
||||
# 构造很多短句
|
||||
text = "你好。再见。谢谢。抱歉。好的。不行。可以。去吧。" * 5 # 每句3-4字
|
||||
|
||||
result = split_text(text, max_chars=100)
|
||||
|
||||
# 合并后段数应该比单纯按句切的少
|
||||
assert len(result) < len(text) // 3 # 粗略估计
|
||||
for seg in result:
|
||||
assert len(seg) <= 100
|
||||
|
||||
def test_preserves_content(self):
|
||||
"""分段后内容总和与原文基本一致(忽略strip的空白)"""
|
||||
text = "这是测试文本。包含多个句子。用来验证分段正确性。" * 5
|
||||
|
||||
result = split_text(text, max_chars=50)
|
||||
|
||||
# 合并所有分段,去掉空白后应该与原文去掉空白后基本一致
|
||||
combined = "".join(result).replace(" ", "")
|
||||
original = text.strip().replace(" ", "")
|
||||
assert combined == original
|
||||
|
||||
def test_custom_max_chars(self):
|
||||
"""支持自定义 max_chars"""
|
||||
text = "测试" * 100 # 200字
|
||||
|
||||
result_50 = split_text(text, max_chars=50)
|
||||
result_100 = split_text(text, max_chars=100)
|
||||
|
||||
# max_chars 越小,段数应该越多
|
||||
assert len(result_50) >= len(result_100)
|
||||
|
||||
def test_single_char_text(self):
|
||||
"""单字符文本"""
|
||||
assert split_text("好", max_chars=10) == ["好"]
|
||||
|
||||
def test_text_with_only_punctuation(self):
|
||||
"""纯标点文本"""
|
||||
text = "。。。。。。。。。。" # 10个句号
|
||||
result = split_text(text, max_chars=5)
|
||||
|
||||
assert len(result) >= 1
|
||||
for seg in result:
|
||||
assert len(seg) <= 5
|
||||
|
||||
def test_mixed_content(self):
|
||||
"""中英文混合内容"""
|
||||
text = "今天的天气是 sunny and warm。我们去了 park 玩。真的很开心!" * 5
|
||||
|
||||
result = split_text(text, max_chars=80)
|
||||
|
||||
for seg in result:
|
||||
assert len(seg) <= 80
|
||||
|
||||
@@ -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
|
||||
|
||||
Executable
+510
@@ -0,0 +1,510 @@
|
||||
"""视频分享 UseCase 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.video_share.commands import (
|
||||
CreateShareCommand,
|
||||
UpdateShareCommand,
|
||||
)
|
||||
from packages.application.video_share.use_cases import (
|
||||
AccessShareUseCase,
|
||||
CreateShareUseCase,
|
||||
GetShareByTokenUseCase,
|
||||
InvalidPasswordError,
|
||||
ListSharesByUserUseCase,
|
||||
ListSharesByVideoUseCase,
|
||||
NotFoundError,
|
||||
PasswordRequiredError,
|
||||
RecordShareDownloadUseCase,
|
||||
RevokeShareUseCase,
|
||||
ShareExpiredError,
|
||||
UpdateShareUseCase,
|
||||
VideoNotFoundError,
|
||||
)
|
||||
from packages.domain.generated_video import GeneratedVideo
|
||||
from packages.domain.video_share import VideoShare
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_share_repo():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_video_repo():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_video():
|
||||
video = MagicMock(spec=GeneratedVideo)
|
||||
video.id = "video_001"
|
||||
video.user_id = "user_001"
|
||||
return video
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_share():
|
||||
share = VideoShare.create(
|
||||
video_id="video_001",
|
||||
user_id="user_001",
|
||||
)
|
||||
return share
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_share_with_password():
|
||||
share = VideoShare.create(
|
||||
video_id="video_001",
|
||||
user_id="user_001",
|
||||
password="secret123",
|
||||
)
|
||||
return share
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_share_expired():
|
||||
# 直接构造已过期的分享(不经过create方法的校验)
|
||||
share = VideoShare(
|
||||
id="share_expired_001",
|
||||
video_id="video_001",
|
||||
user_id="user_001",
|
||||
share_token="expiredtoken123",
|
||||
expires_at=datetime.now(timezone.utc) - timedelta(hours=1),
|
||||
)
|
||||
return share
|
||||
|
||||
|
||||
class TestCreateShareUseCase:
|
||||
"""CreateShareUseCase 测试"""
|
||||
|
||||
def test_create_share_success(self, mock_share_repo, mock_video_repo, sample_video):
|
||||
"""正常创建分享链接"""
|
||||
mock_video_repo.get.return_value = sample_video
|
||||
mock_share_repo.create.side_effect = lambda s: s
|
||||
|
||||
use_case = CreateShareUseCase(mock_share_repo, mock_video_repo)
|
||||
command = CreateShareCommand(video_id="video_001", user_id="user_001")
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.video_id == "video_001"
|
||||
assert result.user_id == "user_001"
|
||||
assert result.share_token is not None
|
||||
assert result.has_password is False
|
||||
mock_share_repo.create.assert_called_once()
|
||||
|
||||
def test_create_share_with_password(self, mock_share_repo, mock_video_repo, sample_video):
|
||||
"""创建带密码的分享"""
|
||||
mock_video_repo.get.return_value = sample_video
|
||||
mock_share_repo.create.side_effect = lambda s: s
|
||||
|
||||
use_case = CreateShareUseCase(mock_share_repo, mock_video_repo)
|
||||
command = CreateShareCommand(
|
||||
video_id="video_001",
|
||||
user_id="user_001",
|
||||
password="mypassword",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.has_password is True
|
||||
assert result.password_hash is not None
|
||||
|
||||
def test_create_share_with_expiry(self, mock_share_repo, mock_video_repo, sample_video):
|
||||
"""创建带有效期的分享"""
|
||||
mock_video_repo.get.return_value = sample_video
|
||||
mock_share_repo.create.side_effect = lambda s: s
|
||||
|
||||
future = datetime.now(timezone.utc) + timedelta(days=7)
|
||||
use_case = CreateShareUseCase(mock_share_repo, mock_video_repo)
|
||||
command = CreateShareCommand(
|
||||
video_id="video_001",
|
||||
user_id="user_001",
|
||||
expires_at=future,
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.expires_at == future
|
||||
|
||||
def test_create_share_video_not_found(self, mock_share_repo, mock_video_repo):
|
||||
"""视频不存在时抛出 VideoNotFoundError"""
|
||||
mock_video_repo.get.return_value = None
|
||||
|
||||
use_case = CreateShareUseCase(mock_share_repo, mock_video_repo)
|
||||
command = CreateShareCommand(video_id="nonexistent", user_id="user_001")
|
||||
|
||||
with pytest.raises(VideoNotFoundError):
|
||||
use_case.execute(command)
|
||||
|
||||
mock_share_repo.create.assert_not_called()
|
||||
|
||||
def test_create_share_wrong_user(self, mock_share_repo, mock_video_repo, sample_video):
|
||||
"""非视频所有者创建分享失败"""
|
||||
sample_video.user_id = "user_other"
|
||||
mock_video_repo.get.return_value = sample_video
|
||||
|
||||
use_case = CreateShareUseCase(mock_share_repo, mock_video_repo)
|
||||
command = CreateShareCommand(video_id="video_001", user_id="user_001")
|
||||
|
||||
with pytest.raises(VideoNotFoundError):
|
||||
use_case.execute(command)
|
||||
|
||||
mock_share_repo.create.assert_not_called()
|
||||
|
||||
|
||||
class TestGetShareByTokenUseCase:
|
||||
"""GetShareByTokenUseCase 测试"""
|
||||
|
||||
def test_get_share_success(self, mock_share_repo, sample_share):
|
||||
"""通过 token 正常获取分享信息"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share
|
||||
|
||||
use_case = GetShareByTokenUseCase(mock_share_repo)
|
||||
result = use_case.execute(sample_share.share_token)
|
||||
|
||||
assert result.id == sample_share.id
|
||||
mock_share_repo.get_by_token.assert_called_once_with(sample_share.share_token)
|
||||
|
||||
def test_get_share_not_found(self, mock_share_repo):
|
||||
"""token 不存在时抛出 NotFoundError"""
|
||||
mock_share_repo.get_by_token.return_value = None
|
||||
|
||||
use_case = GetShareByTokenUseCase(mock_share_repo)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
use_case.execute("invalid_token")
|
||||
|
||||
def test_get_share_expired_raises(self, mock_share_repo, sample_share_expired):
|
||||
"""已过期的分享不可访问"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share_expired
|
||||
|
||||
use_case = GetShareByTokenUseCase(mock_share_repo)
|
||||
|
||||
with pytest.raises(ShareExpiredError):
|
||||
use_case.execute(sample_share_expired.share_token)
|
||||
|
||||
|
||||
class TestAccessShareUseCase:
|
||||
"""AccessShareUseCase 测试"""
|
||||
|
||||
def test_access_without_password(self, mock_share_repo, mock_video_repo, sample_share, sample_video):
|
||||
"""无密码分享直接访问成功"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share
|
||||
mock_video_repo.get.return_value = sample_video
|
||||
|
||||
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
|
||||
result = use_case.execute(sample_share.share_token)
|
||||
|
||||
assert result.share.id == sample_share.id
|
||||
assert result.video.id == "video_001"
|
||||
assert result.password_verified is True
|
||||
mock_share_repo.increment_view.assert_called_once_with(sample_share.id)
|
||||
assert sample_share.view_count == 1
|
||||
|
||||
def test_access_with_correct_password(self, mock_share_repo, mock_video_repo, sample_share_with_password, sample_video):
|
||||
"""带密码分享输入正确密码访问成功"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share_with_password
|
||||
mock_video_repo.get.return_value = sample_video
|
||||
|
||||
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
|
||||
result = use_case.execute(sample_share_with_password.share_token, password="secret123")
|
||||
|
||||
assert result.password_verified is True
|
||||
mock_share_repo.increment_view.assert_called_once()
|
||||
|
||||
def test_access_password_required_but_not_provided(self, mock_share_repo, mock_video_repo, sample_share_with_password):
|
||||
"""带密码分享不输入密码抛出 PasswordRequiredError"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share_with_password
|
||||
|
||||
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
|
||||
|
||||
with pytest.raises(PasswordRequiredError):
|
||||
use_case.execute(sample_share_with_password.share_token)
|
||||
|
||||
mock_share_repo.increment_view.assert_not_called()
|
||||
|
||||
def test_access_wrong_password(self, mock_share_repo, mock_video_repo, sample_share_with_password):
|
||||
"""密码错误抛出 InvalidPasswordError"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share_with_password
|
||||
|
||||
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
|
||||
|
||||
with pytest.raises(InvalidPasswordError):
|
||||
use_case.execute(sample_share_with_password.share_token, password="wrongpass")
|
||||
|
||||
mock_share_repo.increment_view.assert_not_called()
|
||||
|
||||
def test_access_share_not_found(self, mock_share_repo, mock_video_repo):
|
||||
"""分享不存在抛出 NotFoundError"""
|
||||
mock_share_repo.get_by_token.return_value = None
|
||||
|
||||
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
use_case.execute("invalid_token")
|
||||
|
||||
def test_access_expired_share(self, mock_share_repo, mock_video_repo, sample_share_expired):
|
||||
"""已过期分享不可访问"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share_expired
|
||||
|
||||
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
|
||||
|
||||
with pytest.raises(ShareExpiredError):
|
||||
use_case.execute(sample_share_expired.share_token)
|
||||
|
||||
mock_share_repo.increment_view.assert_not_called()
|
||||
|
||||
def test_access_video_not_found(self, mock_share_repo, mock_video_repo, sample_share):
|
||||
"""分享存在但视频不存在"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share
|
||||
mock_video_repo.get.return_value = None
|
||||
|
||||
use_case = AccessShareUseCase(mock_share_repo, mock_video_repo)
|
||||
|
||||
with pytest.raises(VideoNotFoundError):
|
||||
use_case.execute(sample_share.share_token)
|
||||
|
||||
|
||||
class TestListSharesByVideoUseCase:
|
||||
"""ListSharesByVideoUseCase 测试"""
|
||||
|
||||
def test_list_by_video(self, mock_share_repo, sample_share):
|
||||
"""列出某个视频的所有分享"""
|
||||
mock_share_repo.list_by_video.return_value = [sample_share]
|
||||
|
||||
use_case = ListSharesByVideoUseCase(mock_share_repo)
|
||||
result = use_case.execute("video_001", "user_001")
|
||||
|
||||
assert len(result) == 1
|
||||
mock_share_repo.list_by_video.assert_called_once_with("video_001", "user_001")
|
||||
|
||||
def test_list_by_video_empty(self, mock_share_repo):
|
||||
"""视频没有分享记录时返回空列表"""
|
||||
mock_share_repo.list_by_video.return_value = []
|
||||
|
||||
use_case = ListSharesByVideoUseCase(mock_share_repo)
|
||||
result = use_case.execute("video_001", "user_001")
|
||||
|
||||
assert result == []
|
||||
|
||||
|
||||
class TestListSharesByUserUseCase:
|
||||
"""ListSharesByUserUseCase 测试"""
|
||||
|
||||
def test_list_by_user(self, mock_share_repo, sample_share):
|
||||
"""列出用户的所有分享"""
|
||||
mock_share_repo.list_by_user.return_value = [sample_share]
|
||||
mock_share_repo.count_by_user.return_value = 1
|
||||
|
||||
use_case = ListSharesByUserUseCase(mock_share_repo)
|
||||
items, total = use_case.execute("user_001")
|
||||
|
||||
assert len(items) == 1
|
||||
assert total == 1
|
||||
mock_share_repo.list_by_user.assert_called_once_with("user_001", skip=0, limit=20)
|
||||
|
||||
def test_list_by_user_with_pagination(self, mock_share_repo):
|
||||
"""带分页参数查询"""
|
||||
mock_share_repo.list_by_user.return_value = []
|
||||
mock_share_repo.count_by_user.return_value = 50
|
||||
|
||||
use_case = ListSharesByUserUseCase(mock_share_repo)
|
||||
items, total = use_case.execute("user_001", skip=10, limit=5)
|
||||
|
||||
assert total == 50
|
||||
mock_share_repo.list_by_user.assert_called_once_with("user_001", skip=10, limit=5)
|
||||
|
||||
def test_list_by_user_empty(self, mock_share_repo):
|
||||
"""用户没有分享记录"""
|
||||
mock_share_repo.list_by_user.return_value = []
|
||||
mock_share_repo.count_by_user.return_value = 0
|
||||
|
||||
use_case = ListSharesByUserUseCase(mock_share_repo)
|
||||
items, total = use_case.execute("user_001")
|
||||
|
||||
assert items == []
|
||||
assert total == 0
|
||||
|
||||
|
||||
class TestUpdateShareUseCase:
|
||||
"""UpdateShareUseCase 测试"""
|
||||
|
||||
def test_update_password(self, mock_share_repo, sample_share):
|
||||
"""更新分享密码"""
|
||||
mock_share_repo.get_by_id.return_value = sample_share
|
||||
mock_share_repo.update.side_effect = lambda s: s
|
||||
|
||||
use_case = UpdateShareUseCase(mock_share_repo)
|
||||
command = UpdateShareCommand(
|
||||
share_id=sample_share.id,
|
||||
user_id="user_001",
|
||||
password="newpassword",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.has_password is True
|
||||
mock_share_repo.update.assert_called_once()
|
||||
|
||||
def test_clear_password(self, mock_share_repo, sample_share_with_password):
|
||||
"""清除分享密码(空字符串)"""
|
||||
mock_share_repo.get_by_id.return_value = sample_share_with_password
|
||||
mock_share_repo.update.side_effect = lambda s: s
|
||||
|
||||
use_case = UpdateShareUseCase(mock_share_repo)
|
||||
command = UpdateShareCommand(
|
||||
share_id=sample_share_with_password.id,
|
||||
user_id="user_001",
|
||||
password="", # 空字符串表示清除
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.has_password is False
|
||||
assert result.password_hash is None
|
||||
|
||||
def test_update_password_none_no_change(self, mock_share_repo, sample_share_with_password):
|
||||
"""password=None 不修改密码"""
|
||||
original_hash = sample_share_with_password.password_hash
|
||||
mock_share_repo.get_by_id.return_value = sample_share_with_password
|
||||
mock_share_repo.update.side_effect = lambda s: s
|
||||
|
||||
use_case = UpdateShareUseCase(mock_share_repo)
|
||||
command = UpdateShareCommand(
|
||||
share_id=sample_share_with_password.id,
|
||||
user_id="user_001",
|
||||
password=None, # None表示不修改
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.password_hash == original_hash
|
||||
|
||||
def test_update_expires_at(self, mock_share_repo, sample_share):
|
||||
"""更新有效期"""
|
||||
mock_share_repo.get_by_id.return_value = sample_share
|
||||
mock_share_repo.update.side_effect = lambda s: s
|
||||
|
||||
future = datetime.now(timezone.utc) + timedelta(days=3)
|
||||
use_case = UpdateShareUseCase(mock_share_repo)
|
||||
command = UpdateShareCommand(
|
||||
share_id=sample_share.id,
|
||||
user_id="user_001",
|
||||
expires_at=future,
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.expires_at == future
|
||||
|
||||
def test_update_expires_at_past_raises(self, mock_share_repo, sample_share):
|
||||
"""设置过去的有效期抛出 ValueError"""
|
||||
mock_share_repo.get_by_id.return_value = sample_share
|
||||
|
||||
past = datetime.now(timezone.utc) - timedelta(hours=1)
|
||||
use_case = UpdateShareUseCase(mock_share_repo)
|
||||
command = UpdateShareCommand(
|
||||
share_id=sample_share.id,
|
||||
user_id="user_001",
|
||||
expires_at=past,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="expires_at cannot be in the past"):
|
||||
use_case.execute(command)
|
||||
|
||||
mock_share_repo.update.assert_not_called()
|
||||
|
||||
def test_update_share_not_found(self, mock_share_repo):
|
||||
"""分享不存在抛出 NotFoundError"""
|
||||
mock_share_repo.get_by_id.return_value = None
|
||||
|
||||
use_case = UpdateShareUseCase(mock_share_repo)
|
||||
command = UpdateShareCommand(
|
||||
share_id="nonexistent",
|
||||
user_id="user_001",
|
||||
password="newpass",
|
||||
)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
use_case.execute(command)
|
||||
|
||||
mock_share_repo.update.assert_not_called()
|
||||
|
||||
|
||||
class TestRevokeShareUseCase:
|
||||
"""RevokeShareUseCase 测试"""
|
||||
|
||||
def test_revoke_success(self, mock_share_repo, sample_share):
|
||||
"""撤销分享成功"""
|
||||
mock_share_repo.get_by_id.return_value = sample_share
|
||||
mock_share_repo.delete.return_value = True
|
||||
|
||||
use_case = RevokeShareUseCase(mock_share_repo)
|
||||
result = use_case.execute(sample_share.id, "user_001")
|
||||
|
||||
assert result is True
|
||||
mock_share_repo.delete.assert_called_once_with(sample_share.id, "user_001")
|
||||
|
||||
def test_revoke_not_found(self, mock_share_repo):
|
||||
"""分享不存在抛出 NotFoundError"""
|
||||
mock_share_repo.get_by_id.return_value = None
|
||||
|
||||
use_case = RevokeShareUseCase(mock_share_repo)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
use_case.execute("nonexistent", "user_001")
|
||||
|
||||
mock_share_repo.delete.assert_not_called()
|
||||
|
||||
|
||||
class TestRecordShareDownloadUseCase:
|
||||
"""RecordShareDownloadUseCase 测试"""
|
||||
|
||||
def test_record_download_no_password(self, mock_share_repo, sample_share):
|
||||
"""无密码分享记录下载"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share
|
||||
|
||||
use_case = RecordShareDownloadUseCase(mock_share_repo)
|
||||
use_case.execute(sample_share.share_token)
|
||||
|
||||
mock_share_repo.increment_download.assert_called_once_with(sample_share.id)
|
||||
|
||||
def test_record_download_with_password(self, mock_share_repo, sample_share_with_password):
|
||||
"""带密码分享正确密码记录下载"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share_with_password
|
||||
|
||||
use_case = RecordShareDownloadUseCase(mock_share_repo)
|
||||
use_case.execute(sample_share_with_password.share_token, password="secret123")
|
||||
|
||||
mock_share_repo.increment_download.assert_called_once()
|
||||
|
||||
def test_record_download_wrong_password(self, mock_share_repo, sample_share_with_password):
|
||||
"""密码错误不记录下载"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share_with_password
|
||||
|
||||
use_case = RecordShareDownloadUseCase(mock_share_repo)
|
||||
|
||||
with pytest.raises(InvalidPasswordError):
|
||||
use_case.execute(sample_share_with_password.share_token, password="wrong")
|
||||
|
||||
mock_share_repo.increment_download.assert_not_called()
|
||||
|
||||
def test_record_download_not_found(self, mock_share_repo):
|
||||
"""分享不存在抛出 NotFoundError"""
|
||||
mock_share_repo.get_by_token.return_value = None
|
||||
|
||||
use_case = RecordShareDownloadUseCase(mock_share_repo)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
use_case.execute("invalid_token")
|
||||
|
||||
def test_record_download_expired(self, mock_share_repo, sample_share_expired):
|
||||
"""已过期分享不能下载"""
|
||||
mock_share_repo.get_by_token.return_value = sample_share_expired
|
||||
|
||||
use_case = RecordShareDownloadUseCase(mock_share_repo)
|
||||
|
||||
with pytest.raises(ShareExpiredError):
|
||||
use_case.execute(sample_share_expired.share_token)
|
||||
|
||||
mock_share_repo.increment_download.assert_not_called()
|
||||
@@ -1,12 +1,6 @@
|
||||
"""
|
||||
微信 OAuth 服务单元测试(第二十波)
|
||||
"""微信 OAuth 服务单元测试."""
|
||||
|
||||
覆盖:
|
||||
- MemoryStateStore (put / verify_and_consume / 过期清理)
|
||||
- WechatOAuthService.is_configured
|
||||
- WechatOAuthService.generate_auth_url (正常模式 + mock模式)
|
||||
- WechatOAuthService.handle_callback (正常 / 缺code / state无效 / mock模式 / access_token失败 / userinfo失败 / 网络异常)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from unittest.mock import MagicMock, patch
|
||||
@@ -14,341 +8,393 @@ from unittest.mock import MagicMock, patch
|
||||
import pytest
|
||||
|
||||
from packages.application.auth.wechat_oauth_service import (
|
||||
STATE_TTL_SECONDS,
|
||||
MemoryStateStore,
|
||||
STATE_TTL_SECONDS,
|
||||
WechatOAuthService,
|
||||
WechatUserInfo,
|
||||
get_wechat_oauth_service,
|
||||
)
|
||||
|
||||
# ============================================================
|
||||
# MemoryStateStore
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestMemoryStateStore:
|
||||
"""MemoryStateStore 内存 state 存储"""
|
||||
"""MemoryStateStore 测试"""
|
||||
|
||||
def test_put_and_verify(self):
|
||||
"""放入并验证成功"""
|
||||
"""存入 state 后可以验证通过"""
|
||||
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
|
||||
store.put("state_123")
|
||||
assert store.verify_and_consume("state_123") is True
|
||||
|
||||
def test_verify_nonexistent(self):
|
||||
"""验证不存在的 state"""
|
||||
"""不存在的 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_state_consumed_after_verify(self):
|
||||
"""state 验证后被消费,不能重复使用"""
|
||||
store = MemoryStateStore()
|
||||
store.put("state_123")
|
||||
assert store.verify_and_consume("state_123") is True
|
||||
assert store.verify_and_consume("state_123") is False
|
||||
|
||||
def test_put_cleans_expired(self):
|
||||
"""put 时会清理过期的"""
|
||||
def test_multiple_states(self):
|
||||
"""多个 state 独立管理"""
|
||||
store = MemoryStateStore()
|
||||
store.put("state_a")
|
||||
store.put("state_b")
|
||||
assert store.verify_and_consume("state_a") is True
|
||||
assert store.verify_and_consume("state_b") is True
|
||||
|
||||
def test_expired_state_cleaned(self):
|
||||
"""过期 state 会被清理"""
|
||||
store = MemoryStateStore(ttl_seconds=1)
|
||||
store.put("state-1")
|
||||
store.put("expired_state")
|
||||
time.sleep(1.1)
|
||||
store.put("state-2")
|
||||
# state-1 应该被清理掉了
|
||||
assert len(store._states) == 1
|
||||
assert "state-2" in store._states
|
||||
assert store.verify_and_consume("expired_state") is False
|
||||
|
||||
def test_default_ttl(self):
|
||||
"""默认 TTL 是 10 分钟"""
|
||||
store = MemoryStateStore()
|
||||
assert store._ttl == STATE_TTL_SECONDS
|
||||
def test_custom_ttl(self):
|
||||
"""自定义 TTL"""
|
||||
store = MemoryStateStore(ttl_seconds=60)
|
||||
store.put("my_state")
|
||||
# 立即验证应该通过
|
||||
assert store.verify_and_consume("my_state") is True
|
||||
|
||||
def test_clean_expired_on_put(self):
|
||||
"""put 时清理过期 state"""
|
||||
store = MemoryStateStore(ttl_seconds=1)
|
||||
store.put("old_state")
|
||||
time.sleep(1.1)
|
||||
# put 新 state 时会触发清理
|
||||
store.put("new_state")
|
||||
# old_state 已经过期了,验证应该失败
|
||||
assert store.verify_and_consume("old_state") is False
|
||||
# new_state 应该还在
|
||||
assert store.verify_and_consume("new_state") is True
|
||||
|
||||
# ============================================================
|
||||
# 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
|
||||
# ============================================================
|
||||
def test_clean_expired_on_verify(self):
|
||||
"""verify 时清理过期 state"""
|
||||
store = MemoryStateStore(ttl_seconds=1)
|
||||
store.put("old_state")
|
||||
time.sleep(1.1)
|
||||
# 验证不存在的 state 也会触发清理
|
||||
store.verify_and_consume("other_state")
|
||||
# old_state 已过期,验证失败
|
||||
assert store.verify_and_consume("old_state") is False
|
||||
|
||||
|
||||
class TestWechatUserInfo:
|
||||
"""WechatUserInfo 数据类"""
|
||||
"""WechatUserInfo 测试"""
|
||||
|
||||
def test_minimal_fields(self):
|
||||
info = WechatUserInfo(openid="abc")
|
||||
assert info.openid == "abc"
|
||||
def test_create_with_openid(self):
|
||||
"""仅用 openid 创建"""
|
||||
info = WechatUserInfo(openid="openid_123")
|
||||
assert info.openid == "openid_123"
|
||||
assert info.unionid == ""
|
||||
assert info.nickname == ""
|
||||
assert info.avatar_url == ""
|
||||
|
||||
def test_full_fields(self):
|
||||
def test_create_with_all_fields(self):
|
||||
"""所有字段创建"""
|
||||
info = WechatUserInfo(
|
||||
openid="abc",
|
||||
unionid="def",
|
||||
nickname="测试",
|
||||
openid="openid_123",
|
||||
unionid="unionid_456",
|
||||
nickname="测试用户",
|
||||
avatar_url="https://example.com/avatar.jpg",
|
||||
)
|
||||
assert info.openid == "abc"
|
||||
assert info.unionid == "def"
|
||||
assert info.nickname == "测试"
|
||||
assert info.openid == "openid_123"
|
||||
assert info.unionid == "unionid_456"
|
||||
assert info.nickname == "测试用户"
|
||||
assert info.avatar_url == "https://example.com/avatar.jpg"
|
||||
|
||||
|
||||
# ============================================================
|
||||
# get_wechat_oauth_service
|
||||
# ============================================================
|
||||
class TestWechatOAuthServiceInit:
|
||||
"""WechatOAuthService 初始化测试"""
|
||||
|
||||
def test_not_configured_default(self):
|
||||
"""默认参数(无环境变量)时未配置"""
|
||||
with patch.dict("os.environ", {}, clear=False):
|
||||
# 确保环境变量为空
|
||||
service = WechatOAuthService(
|
||||
app_id="", app_secret="", redirect_uri=""
|
||||
)
|
||||
assert service.is_configured() is False
|
||||
|
||||
def test_configured_with_params(self):
|
||||
"""显式传入配置时已配置"""
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123",
|
||||
app_secret="secret456",
|
||||
redirect_uri="https://example.com/callback",
|
||||
)
|
||||
assert service.is_configured() is True
|
||||
|
||||
def test_missing_app_id_not_configured(self):
|
||||
"""缺少 app_id 未配置"""
|
||||
service = WechatOAuthService(
|
||||
app_id="",
|
||||
app_secret="secret456",
|
||||
redirect_uri="https://example.com/callback",
|
||||
)
|
||||
assert service.is_configured() is False
|
||||
|
||||
def test_default_state_store(self):
|
||||
"""默认使用 MemoryStateStore"""
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123", app_secret="s", redirect_uri="https://x.com"
|
||||
)
|
||||
assert isinstance(service._state_store, MemoryStateStore)
|
||||
|
||||
def test_custom_state_store(self):
|
||||
"""可以自定义 state_store"""
|
||||
custom_store = MagicMock()
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123",
|
||||
app_secret="s",
|
||||
redirect_uri="https://x.com",
|
||||
state_store=custom_store,
|
||||
)
|
||||
assert service._state_store is custom_store
|
||||
|
||||
|
||||
class TestGenerateAuthUrl:
|
||||
"""generate_auth_url 测试"""
|
||||
|
||||
def test_returns_url_and_state(self):
|
||||
"""返回 URL 和 state"""
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123",
|
||||
app_secret="secret456",
|
||||
redirect_uri="https://example.com/callback",
|
||||
)
|
||||
url, state = service.generate_auth_url()
|
||||
assert isinstance(url, str)
|
||||
assert isinstance(state, str)
|
||||
assert len(state) > 0
|
||||
assert "weixin.qq.com" in url
|
||||
|
||||
def test_url_contains_params(self):
|
||||
"""URL 包含必要参数"""
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123",
|
||||
app_secret="secret456",
|
||||
redirect_uri="https://example.com/callback",
|
||||
)
|
||||
url, state = service.generate_auth_url(scope="snsapi_login")
|
||||
|
||||
assert "appid=wx123" in url
|
||||
assert "snsapi_login" in url
|
||||
assert state in url
|
||||
assert "response_type=code" in url
|
||||
|
||||
def test_state_saved_to_store(self):
|
||||
"""生成的 state 存入 store"""
|
||||
mock_store = MagicMock()
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123",
|
||||
app_secret="secret456",
|
||||
redirect_uri="https://example.com/callback",
|
||||
state_store=mock_store,
|
||||
)
|
||||
url, state = service.generate_auth_url()
|
||||
mock_store.put.assert_called_once_with(state)
|
||||
|
||||
def test_mock_mode_when_not_configured(self):
|
||||
"""未配置时返回 mock URL"""
|
||||
service = WechatOAuthService(
|
||||
app_id="", app_secret="", redirect_uri="https://example.com/callback"
|
||||
)
|
||||
url, state = service.generate_auth_url()
|
||||
assert "/mock/wechat/auth" in url
|
||||
assert "mock" in url
|
||||
|
||||
def test_different_states_each_time(self):
|
||||
"""每次生成不同的 state"""
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123",
|
||||
app_secret="secret456",
|
||||
redirect_uri="https://example.com/callback",
|
||||
)
|
||||
_, state1 = service.generate_auth_url()
|
||||
_, state2 = service.generate_auth_url()
|
||||
assert state1 != state2
|
||||
|
||||
|
||||
class TestHandleCallback:
|
||||
"""handle_callback 测试"""
|
||||
|
||||
def test_missing_code_returns_error(self):
|
||||
"""缺少 code 返回错误"""
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123", app_secret="s", redirect_uri="https://x.com"
|
||||
)
|
||||
user_info, error = service.handle_callback("", "some_state")
|
||||
assert user_info is None
|
||||
assert "缺少授权码" in error
|
||||
|
||||
def test_invalid_state_returns_error(self):
|
||||
"""state 无效返回错误"""
|
||||
mock_store = MagicMock()
|
||||
mock_store.verify_and_consume.return_value = False
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123",
|
||||
app_secret="s",
|
||||
redirect_uri="https://x.com",
|
||||
state_store=mock_store,
|
||||
)
|
||||
user_info, error = service.handle_callback("code123", "bad_state")
|
||||
assert user_info is None
|
||||
assert "state" in error
|
||||
|
||||
def test_mock_mode_when_not_configured(self):
|
||||
"""未配置时返回 mock 用户信息"""
|
||||
mock_store = MagicMock()
|
||||
mock_store.verify_and_consume.return_value = True
|
||||
service = WechatOAuthService(
|
||||
app_id="",
|
||||
app_secret="",
|
||||
redirect_uri="https://x.com",
|
||||
state_store=mock_store,
|
||||
)
|
||||
user_info, error = service.handle_callback("mock_code_12345", "valid_state")
|
||||
|
||||
assert error is None
|
||||
assert user_info is not None
|
||||
assert user_info.openid.startswith("mock_")
|
||||
assert "微信测试用户" in user_info.nickname
|
||||
|
||||
def test_state_consumed_after_callback(self):
|
||||
"""回调处理后 state 被消费"""
|
||||
mock_store = MagicMock()
|
||||
mock_store.verify_and_consume.return_value = True
|
||||
service = WechatOAuthService(
|
||||
app_id="",
|
||||
app_secret="",
|
||||
redirect_uri="https://x.com",
|
||||
state_store=mock_store,
|
||||
)
|
||||
service.handle_callback("code", "valid_state")
|
||||
mock_store.verify_and_consume.assert_called_once_with("valid_state")
|
||||
|
||||
def test_real_mode_success(self):
|
||||
"""真实模式下成功获取用户信息"""
|
||||
mock_store = MagicMock()
|
||||
mock_store.verify_and_consume.return_value = True
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123",
|
||||
app_secret="secret456",
|
||||
redirect_uri="https://x.com",
|
||||
state_store=mock_store,
|
||||
)
|
||||
|
||||
mock_token_resp = MagicMock()
|
||||
mock_token_resp.json.return_value = {
|
||||
"access_token": "access_token_123",
|
||||
"openid": "real_openid",
|
||||
"unionid": "real_unionid",
|
||||
}
|
||||
mock_user_resp = MagicMock()
|
||||
mock_user_resp.json.return_value = {
|
||||
"nickname": "真实用户",
|
||||
"headimgurl": "https://wx.qlogo.cn/avatar.jpg",
|
||||
}
|
||||
|
||||
with patch("requests.get") as mock_get:
|
||||
mock_get.side_effect = [mock_token_resp, mock_user_resp]
|
||||
user_info, error = service.handle_callback("auth_code", "valid_state")
|
||||
|
||||
assert error is None
|
||||
assert user_info is not None
|
||||
assert user_info.openid == "real_openid"
|
||||
assert user_info.unionid == "real_unionid"
|
||||
assert user_info.nickname == "真实用户"
|
||||
assert user_info.avatar_url == "https://wx.qlogo.cn/avatar.jpg"
|
||||
|
||||
def test_real_mode_token_error(self):
|
||||
"""真实模式下 access_token 接口返回错误"""
|
||||
mock_store = MagicMock()
|
||||
mock_store.verify_and_consume.return_value = True
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123",
|
||||
app_secret="secret456",
|
||||
redirect_uri="https://x.com",
|
||||
state_store=mock_store,
|
||||
)
|
||||
|
||||
mock_resp = MagicMock()
|
||||
mock_resp.json.return_value = {
|
||||
"errcode": 40029,
|
||||
"errmsg": "invalid code",
|
||||
}
|
||||
|
||||
with patch("requests.get", return_value=mock_resp):
|
||||
user_info, error = service.handle_callback("bad_code", "valid_state")
|
||||
|
||||
assert user_info is None
|
||||
assert error is not None
|
||||
assert "微信授权失败" in error
|
||||
|
||||
def test_real_mode_userinfo_error(self):
|
||||
"""真实模式下用户信息接口返回错误"""
|
||||
mock_store = MagicMock()
|
||||
mock_store.verify_and_consume.return_value = True
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123",
|
||||
app_secret="secret456",
|
||||
redirect_uri="https://x.com",
|
||||
state_store=mock_store,
|
||||
)
|
||||
|
||||
mock_token_resp = MagicMock()
|
||||
mock_token_resp.json.return_value = {
|
||||
"access_token": "access_123",
|
||||
"openid": "open_123",
|
||||
}
|
||||
mock_user_resp = MagicMock()
|
||||
mock_user_resp.json.return_value = {
|
||||
"errcode": 40001,
|
||||
"errmsg": "invalid token",
|
||||
}
|
||||
|
||||
with patch("requests.get") as mock_get:
|
||||
mock_get.side_effect = [mock_token_resp, mock_user_resp]
|
||||
user_info, error = service.handle_callback("code", "state")
|
||||
|
||||
assert user_info is None
|
||||
assert "获取用户信息失败" in error
|
||||
|
||||
def test_real_mode_network_error(self):
|
||||
"""网络异常时返回友好错误"""
|
||||
import requests
|
||||
|
||||
mock_store = MagicMock()
|
||||
mock_store.verify_and_consume.return_value = True
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123",
|
||||
app_secret="secret456",
|
||||
redirect_uri="https://x.com",
|
||||
state_store=mock_store,
|
||||
)
|
||||
|
||||
with patch("requests.get", side_effect=requests.ConnectionError()):
|
||||
user_info, error = service.handle_callback("code", "state")
|
||||
|
||||
assert user_info is None
|
||||
assert "暂不可用" in error
|
||||
|
||||
def test_empty_state_returns_error(self):
|
||||
"""空 state 返回错误"""
|
||||
service = WechatOAuthService(
|
||||
app_id="wx123", app_secret="s", redirect_uri="https://x.com"
|
||||
)
|
||||
user_info, error = service.handle_callback("code123", "")
|
||||
assert user_info is None
|
||||
assert "state" in error
|
||||
|
||||
|
||||
class TestGetWechatOAuthService:
|
||||
"""工厂函数"""
|
||||
"""get_wechat_oauth_service 函数测试"""
|
||||
|
||||
def test_returns_service_instance(self):
|
||||
svc = get_wechat_oauth_service()
|
||||
assert isinstance(svc, WechatOAuthService)
|
||||
"""返回 WechatOAuthService 实例"""
|
||||
service = get_wechat_oauth_service()
|
||||
assert isinstance(service, WechatOAuthService)
|
||||
|
||||
@@ -1,306 +1,360 @@
|
||||
"""
|
||||
微信同步登录/注册 Use Case 测试
|
||||
"""
|
||||
"""微信同步登录 UseCase 单元测试."""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import Mock, patch
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.auth.wechat_sync_use_case import (
|
||||
WechatSyncRequest,
|
||||
WechatSyncResponse,
|
||||
WechatSyncUseCase,
|
||||
)
|
||||
from packages.domain.entities import User
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_user_repo():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_session_store():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_user():
|
||||
user = User(
|
||||
id="user_001",
|
||||
email="test@wechat.local",
|
||||
username="wx_test123",
|
||||
display_name="微信用户",
|
||||
password_hash="hashed",
|
||||
email_verified=True,
|
||||
)
|
||||
user.wechat_openid = "openid_123"
|
||||
user.wechat_unionid = "unionid_456"
|
||||
user.last_login_at = None
|
||||
user.last_login_ip = None
|
||||
return user
|
||||
|
||||
|
||||
class TestWechatSyncRequest:
|
||||
"""微信同步请求对象测试"""
|
||||
"""WechatSyncRequest 测试"""
|
||||
|
||||
def test_request_with_basic_fields(self):
|
||||
"""测试基本字段初始化"""
|
||||
request = WechatSyncRequest(openid="openid123")
|
||||
assert request.openid == "openid123"
|
||||
assert request.unionid == ""
|
||||
assert request.nickname == "微信用户"
|
||||
assert request.avatar_url == ""
|
||||
assert request.source == "miniapp"
|
||||
def test_openid_stripped(self):
|
||||
"""openid 被 strip"""
|
||||
req = WechatSyncRequest(openid=" openid_123 ")
|
||||
assert req.openid == "openid_123"
|
||||
|
||||
def test_request_with_all_fields(self):
|
||||
"""测试完整字段初始化"""
|
||||
request = WechatSyncRequest(
|
||||
openid=" openid123 ",
|
||||
unionid=" unionid456 ",
|
||||
def test_unionid_stripped(self):
|
||||
"""unionid 被 strip"""
|
||||
req = WechatSyncRequest(openid="o1", unionid=" unionid_456 ")
|
||||
assert req.unionid == "unionid_456"
|
||||
|
||||
def test_default_nickname(self):
|
||||
"""默认昵称"""
|
||||
req = WechatSyncRequest(openid="o1")
|
||||
assert req.nickname == "微信用户"
|
||||
|
||||
def test_default_source(self):
|
||||
"""默认来源"""
|
||||
req = WechatSyncRequest(openid="o1")
|
||||
assert req.source == "miniapp"
|
||||
|
||||
def test_empty_unionid(self):
|
||||
"""不传 unionid 默认为空字符串"""
|
||||
req = WechatSyncRequest(openid="o1")
|
||||
assert req.unionid == ""
|
||||
|
||||
|
||||
class TestWechatSyncResponse:
|
||||
"""WechatSyncResponse 测试"""
|
||||
|
||||
def test_to_dict_contains_fields(self):
|
||||
"""to_dict 包含所有必要字段"""
|
||||
resp = WechatSyncResponse(
|
||||
access_token="access_123",
|
||||
refresh_token="refresh_456",
|
||||
user_id="user_001",
|
||||
nickname="测试用户",
|
||||
avatar_url="http://example.com/avatar.jpg",
|
||||
source="h5",
|
||||
avatar_url="https://example.com/avatar.jpg",
|
||||
is_new_user=False,
|
||||
expires_in=1800,
|
||||
)
|
||||
assert request.openid == "openid123" # stripped
|
||||
assert request.unionid == "unionid456" # stripped
|
||||
assert request.nickname == "测试用户"
|
||||
assert request.avatar_url == "http://example.com/avatar.jpg"
|
||||
assert request.source == "h5"
|
||||
data = resp.to_dict()
|
||||
|
||||
def test_request_empty_unionid_stays_empty(self):
|
||||
"""测试空 unionid 处理"""
|
||||
request = WechatSyncRequest(openid="openid123", unionid="")
|
||||
assert request.unionid == ""
|
||||
|
||||
def test_request_none_nickname_defaults(self):
|
||||
"""测试空昵称使用默认值"""
|
||||
request = WechatSyncRequest(openid="openid123", nickname="")
|
||||
assert request.nickname == "微信用户"
|
||||
assert data["access_token"] == "access_123"
|
||||
assert data["token"] == "access_123" # 兼容字段
|
||||
assert data["refresh_token"] == "refresh_456"
|
||||
assert data["user_id"] == "user_001"
|
||||
assert data["is_new_user"] is False
|
||||
assert data["expires_in"] == 1800
|
||||
assert "user" in data
|
||||
assert "user_info" in data
|
||||
assert data["user"]["id"] == "user_001"
|
||||
assert data["user"]["nickname"] == "测试用户"
|
||||
assert data["user"]["display_name"] == "测试用户"
|
||||
|
||||
|
||||
class TestWechatSyncUseCase:
|
||||
"""微信同步登录/注册用例测试"""
|
||||
class TestWechatSyncUseCaseLoginExisting:
|
||||
"""已有用户登录测试"""
|
||||
|
||||
@pytest.fixture
|
||||
def mock_user_repo(self):
|
||||
"""Mock 用户仓储"""
|
||||
repo = Mock()
|
||||
repo.find_by_wechat_openid = Mock(return_value=None)
|
||||
repo.find_by_wechat_unionid = Mock(return_value=None)
|
||||
repo.find_by_username = Mock(return_value=None)
|
||||
repo.find_by_email = Mock(return_value=None)
|
||||
repo.save = Mock()
|
||||
repo.get = Mock(return_value=None)
|
||||
return repo
|
||||
|
||||
@pytest.fixture
|
||||
def mock_session_store(self):
|
||||
"""Mock Session 存储"""
|
||||
store = Mock()
|
||||
store.save_session = Mock(return_value=True)
|
||||
store.get_refresh_token = Mock(return_value=None)
|
||||
store.get_session_by_refresh_token = Mock(return_value=None)
|
||||
store.delete_session = Mock(return_value=True)
|
||||
return store
|
||||
|
||||
@pytest.fixture
|
||||
def test_user(self):
|
||||
"""测试用户"""
|
||||
return User(
|
||||
id="user-123",
|
||||
email="test@example.com",
|
||||
username="testuser",
|
||||
display_name="测试用户",
|
||||
password_hash="hashed_password",
|
||||
wechat_openid="openid123",
|
||||
wechat_unionid="unionid456",
|
||||
)
|
||||
|
||||
@pytest.fixture
|
||||
def use_case(self, mock_user_repo, mock_session_store):
|
||||
"""创建微信同步用例"""
|
||||
return WechatSyncUseCase(
|
||||
user_repository=mock_user_repo,
|
||||
session_store=mock_session_store,
|
||||
jwt_secret_key="test-secret-key-for-unit-tests",
|
||||
)
|
||||
|
||||
# ===== 登录场景:openid 找到用户 =====
|
||||
|
||||
def test_login_by_openid_success(self, use_case, mock_user_repo, mock_session_store, test_user):
|
||||
"""测试通过 openid 登录成功"""
|
||||
mock_user_repo.find_by_wechat_openid.return_value = test_user
|
||||
|
||||
request = WechatSyncRequest(openid="openid123")
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.user_id == "user-123"
|
||||
assert response.nickname == "测试用户"
|
||||
assert response.is_new_user is False
|
||||
assert response.access_token != ""
|
||||
assert response.refresh_token != ""
|
||||
assert response.expires_in > 0
|
||||
|
||||
# 验证 session 已保存
|
||||
mock_session_store.save_session.assert_called_once()
|
||||
save_kwargs = mock_session_store.save_session.call_args.kwargs
|
||||
assert save_kwargs["user_id"] == "user-123"
|
||||
assert "wechat_miniapp" in save_kwargs["device_info"]
|
||||
|
||||
# 验证更新了最后登录信息
|
||||
mock_user_repo.save.assert_called_once()
|
||||
saved_user = mock_user_repo.save.call_args[0][0]
|
||||
assert saved_user.last_login_at is not None
|
||||
assert saved_user.last_login_ip == "bff_gateway"
|
||||
|
||||
# 验证 to_dict 包含兼容字段
|
||||
data = response.to_dict()
|
||||
assert data["access_token"] == response.access_token
|
||||
assert data["token"] == response.access_token # 兼容字段
|
||||
assert data["user"]["id"] == "user-123"
|
||||
assert data["user_info"]["id"] == "user-123"
|
||||
|
||||
# ===== 登录场景:openid 没找到,通过 unionid 找到 =====
|
||||
|
||||
def test_login_by_unionid_binds_openid(self, use_case, mock_user_repo, mock_session_store, test_user):
|
||||
"""测试通过 unionid 找到用户并绑定当前 openid"""
|
||||
# openid 没找到
|
||||
mock_user_repo.find_by_wechat_openid.return_value = None
|
||||
# unionid 找到了(但 openid 字段为空)
|
||||
test_user.wechat_openid = None
|
||||
mock_user_repo.find_by_wechat_unionid.return_value = test_user
|
||||
|
||||
request = WechatSyncRequest(
|
||||
openid="new_openid_789",
|
||||
unionid="unionid456",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.is_new_user is False
|
||||
assert response.user_id == "user-123"
|
||||
|
||||
# 验证绑定了新的 openid(save 被调用了两次:一次绑定 openid,一次更新登录信息)
|
||||
assert mock_user_repo.save.call_count == 2
|
||||
# 第一次 save 应该是绑定 openid
|
||||
first_save_user = mock_user_repo.save.call_args_list[0][0][0]
|
||||
assert first_save_user.wechat_openid == "new_openid_789"
|
||||
|
||||
def test_login_by_unionid_no_binding_needed(self, use_case, mock_user_repo, mock_session_store, test_user):
|
||||
"""测试通过 unionid 找到用户且 openid 已存在时(不需要额外绑定)"""
|
||||
mock_user_repo.find_by_wechat_openid.return_value = None
|
||||
mock_user_repo.find_by_wechat_unionid.return_value = test_user
|
||||
|
||||
request = WechatSyncRequest(
|
||||
openid="openid123", # 跟用户已有的一样
|
||||
unionid="unionid456",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.is_new_user is False
|
||||
# 还是会 save(绑定)+ save(更新登录信息)= 2次
|
||||
assert mock_user_repo.save.call_count == 2
|
||||
|
||||
# ===== 注册场景:openid 和 unionid 都没找到,创建新用户 =====
|
||||
|
||||
def test_register_new_user(self, use_case, mock_user_repo, mock_session_store):
|
||||
"""测试创建新微信用户"""
|
||||
mock_user_repo.find_by_wechat_openid.return_value = None
|
||||
def test_login_by_openid(self, mock_user_repo, mock_session_store, sample_user):
|
||||
"""通过 openid 登录已有用户"""
|
||||
mock_user_repo.find_by_wechat_openid.return_value = sample_user
|
||||
mock_user_repo.find_by_wechat_unionid.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = None
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
use_case = WechatSyncUseCase(
|
||||
mock_user_repo,
|
||||
session_store=mock_session_store,
|
||||
jwt_secret_key="test-secret-key-for-jwt-12345",
|
||||
)
|
||||
request = WechatSyncRequest(openid="openid_123", nickname="测试")
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.user_id == "user_001"
|
||||
assert response.is_new_user is False
|
||||
mock_user_repo.find_by_wechat_openid.assert_called_once_with("openid_123")
|
||||
mock_session_store.save_session.assert_called_once()
|
||||
|
||||
def test_login_by_unionid(self, mock_user_repo, mock_session_store, sample_user):
|
||||
"""openid 没找到,通过 unionid 找到并绑定 openid"""
|
||||
sample_user.wechat_openid = None # 没有当前 openid
|
||||
mock_user_repo.find_by_wechat_openid.return_value = None
|
||||
mock_user_repo.find_by_wechat_unionid.return_value = sample_user
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
use_case = WechatSyncUseCase(
|
||||
mock_user_repo,
|
||||
session_store=mock_session_store,
|
||||
jwt_secret_key="test-secret-key-for-jwt-12345",
|
||||
)
|
||||
request = WechatSyncRequest(
|
||||
openid="new_openid",
|
||||
unionid="new_unionid",
|
||||
unionid="unionid_456",
|
||||
nickname="测试",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.is_new_user is False
|
||||
# 应该保存了新的 openid
|
||||
assert sample_user.wechat_openid == "new_openid"
|
||||
mock_user_repo.save.assert_called()
|
||||
|
||||
def test_updates_last_login(self, mock_user_repo, mock_session_store, sample_user):
|
||||
"""登录时更新最后登录信息"""
|
||||
mock_user_repo.find_by_wechat_openid.return_value = sample_user
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
use_case = WechatSyncUseCase(
|
||||
mock_user_repo,
|
||||
session_store=mock_session_store,
|
||||
jwt_secret_key="test-secret-key-for-jwt-12345",
|
||||
)
|
||||
request = WechatSyncRequest(openid="openid_123")
|
||||
use_case.execute(request)
|
||||
|
||||
assert sample_user.last_login_at is not None
|
||||
assert sample_user.last_login_ip == "bff_gateway"
|
||||
|
||||
def test_returns_tokens(self, mock_user_repo, mock_session_store, sample_user):
|
||||
"""返回 access_token 和 refresh_token"""
|
||||
mock_user_repo.find_by_wechat_openid.return_value = sample_user
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
use_case = WechatSyncUseCase(
|
||||
mock_user_repo,
|
||||
session_store=mock_session_store,
|
||||
jwt_secret_key="test-secret-key-for-jwt-12345",
|
||||
)
|
||||
request = WechatSyncRequest(openid="openid_123")
|
||||
response, _ = use_case.execute(request)
|
||||
|
||||
assert response.access_token is not None
|
||||
assert len(response.access_token) > 0
|
||||
assert response.refresh_token is not None
|
||||
assert len(response.refresh_token) > 0
|
||||
assert response.expires_in > 0
|
||||
|
||||
|
||||
class TestWechatSyncUseCaseNewUser:
|
||||
"""新用户注册测试"""
|
||||
|
||||
def test_create_new_user(self, mock_user_repo, mock_session_store):
|
||||
"""openid 和 unionid 都没找到,创建新用户"""
|
||||
mock_user_repo.find_by_wechat_openid.return_value = None
|
||||
mock_user_repo.find_by_wechat_unionid.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = None # username 不重复
|
||||
|
||||
saved_user = None
|
||||
|
||||
def capture_save(user):
|
||||
nonlocal saved_user
|
||||
saved_user = user
|
||||
|
||||
mock_user_repo.save.side_effect = capture_save
|
||||
|
||||
use_case = WechatSyncUseCase(
|
||||
mock_user_repo,
|
||||
session_store=mock_session_store,
|
||||
jwt_secret_key="test-secret-key-for-jwt-12345",
|
||||
)
|
||||
request = WechatSyncRequest(
|
||||
openid="new_openid_789",
|
||||
unionid="new_union_789",
|
||||
nickname="新用户",
|
||||
avatar_url="http://example.com/avatar.jpg",
|
||||
source="miniapp",
|
||||
avatar_url="https://example.com/avatar.jpg",
|
||||
)
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.is_new_user is True
|
||||
assert response.nickname == "新用户"
|
||||
assert response.access_token != ""
|
||||
assert response.refresh_token != ""
|
||||
assert saved_user is not None
|
||||
assert saved_user.wechat_openid == "new_openid_789"
|
||||
assert saved_user.wechat_unionid == "new_union_789"
|
||||
assert saved_user.email.endswith("@wechat.local")
|
||||
assert saved_user.username.startswith("wx_")
|
||||
assert saved_user.email_verified is True
|
||||
|
||||
# 验证用户被创建并保存
|
||||
assert mock_user_repo.save.call_count >= 1
|
||||
# 找到 save 的用户(可能有多次save,找第一次即创建用户的那次)
|
||||
created_user = None
|
||||
for call in mock_user_repo.save.call_args_list:
|
||||
user = call[0][0]
|
||||
if user.wechat_openid == "new_openid":
|
||||
created_user = user
|
||||
break
|
||||
assert created_user is not None
|
||||
assert created_user.wechat_openid == "new_openid"
|
||||
assert created_user.wechat_unionid == "new_unionid"
|
||||
assert created_user.email_verified is True
|
||||
assert created_user.username.startswith("wx_")
|
||||
assert "@wechat.local" in created_user.email
|
||||
|
||||
def test_register_new_user_without_unionid(self, use_case, mock_user_repo, mock_session_store):
|
||||
"""测试创建无 unionid 的新用户"""
|
||||
mock_user_repo.find_by_wechat_openid.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = None
|
||||
|
||||
request = WechatSyncRequest(openid="openid_no_union")
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.is_new_user is True
|
||||
|
||||
created_user = mock_user_repo.save.call_args_list[0][0][0]
|
||||
assert created_user.wechat_unionid is None
|
||||
|
||||
def test_register_username_conflict_adds_suffix(self, use_case, mock_user_repo, mock_session_store):
|
||||
"""测试用户名冲突时自动加后缀"""
|
||||
def test_new_user_email_based_on_openid(self, mock_user_repo, mock_session_store):
|
||||
"""新用户邮箱基于 openid 生成"""
|
||||
mock_user_repo.find_by_wechat_openid.return_value = None
|
||||
mock_user_repo.find_by_wechat_unionid.return_value = None
|
||||
# 第一次 find_by_username 返回存在(冲突),第二次返回 None(生成了带后缀的新名)
|
||||
mock_user_repo.find_by_username.side_effect = [Mock(), None]
|
||||
mock_user_repo.find_by_username.return_value = None
|
||||
|
||||
request = WechatSyncRequest(openid="conflict_openid")
|
||||
saved_user = None
|
||||
|
||||
def capture_save(user):
|
||||
nonlocal saved_user
|
||||
saved_user = user
|
||||
|
||||
mock_user_repo.save.side_effect = capture_save
|
||||
|
||||
use_case = WechatSyncUseCase(
|
||||
mock_user_repo,
|
||||
session_store=mock_session_store,
|
||||
jwt_secret_key="test-secret-key-for-jwt-12345",
|
||||
)
|
||||
request = WechatSyncRequest(openid="abcdef1234567890")
|
||||
use_case.execute(request)
|
||||
|
||||
assert "abcdef1234567890" in saved_user.email or "abcdef1234567890"[:20] in saved_user.email
|
||||
assert saved_user.email.endswith("@wechat.local")
|
||||
|
||||
def test_username_conflict_adds_suffix(self, mock_user_repo, mock_session_store):
|
||||
"""用户名冲突时加后缀"""
|
||||
call_count = [0]
|
||||
|
||||
def mock_find_by_username(username):
|
||||
# 前两次返回存在(模拟冲突),第三次返回 None(可用)
|
||||
call_count[0] += 1
|
||||
if call_count[0] <= 2:
|
||||
return MagicMock()
|
||||
return None
|
||||
|
||||
mock_user_repo.find_by_wechat_openid.return_value = None
|
||||
mock_user_repo.find_by_wechat_unionid.return_value = None
|
||||
mock_user_repo.find_by_username.side_effect = mock_find_by_username
|
||||
|
||||
use_case = WechatSyncUseCase(
|
||||
mock_user_repo,
|
||||
session_store=mock_session_store,
|
||||
jwt_secret_key="test-secret-key-for-jwt-12345",
|
||||
)
|
||||
request = WechatSyncRequest(openid="test_openid")
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.is_new_user is True
|
||||
# find_by_username 被调用了多次(找不冲突的用户名)
|
||||
assert mock_user_repo.find_by_username.call_count >= 2
|
||||
|
||||
# find_by_username 应该被调用了两次
|
||||
assert mock_user_repo.find_by_username.call_count == 2
|
||||
# 第二个用户名应该带后缀 _1
|
||||
second_call_username = mock_user_repo.find_by_username.call_args_list[1][0][0]
|
||||
assert "_1" in second_call_username
|
||||
|
||||
def test_register_default_nickname_when_empty(self, use_case, mock_user_repo, mock_session_store):
|
||||
"""测试新用户空昵称时使用默认值"""
|
||||
def test_new_user_has_password_hash(self, mock_user_repo, mock_session_store):
|
||||
"""新用户有随机密码哈希(不能是空的)"""
|
||||
mock_user_repo.find_by_wechat_openid.return_value = None
|
||||
mock_user_repo.find_by_wechat_unionid.return_value = None
|
||||
mock_user_repo.find_by_username.return_value = None
|
||||
|
||||
request = WechatSyncRequest(openid="openid123", nickname="")
|
||||
response, error = use_case.execute(request)
|
||||
saved_user = None
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.nickname == "微信用户"
|
||||
def capture_save(user):
|
||||
nonlocal saved_user
|
||||
saved_user = user
|
||||
|
||||
# ===== 错误场景 =====
|
||||
mock_user_repo.save.side_effect = capture_save
|
||||
|
||||
def test_missing_openid(self, use_case):
|
||||
"""测试缺少 openid"""
|
||||
use_case = WechatSyncUseCase(
|
||||
mock_user_repo,
|
||||
session_store=mock_session_store,
|
||||
jwt_secret_key="test-secret-key-for-jwt-12345",
|
||||
)
|
||||
request = WechatSyncRequest(openid="new_openid")
|
||||
use_case.execute(request)
|
||||
|
||||
assert saved_user.password_hash is not None
|
||||
assert len(saved_user.password_hash) > 0
|
||||
|
||||
|
||||
class TestWechatSyncUseCaseErrors:
|
||||
"""错误场景测试"""
|
||||
|
||||
def test_empty_openid(self, mock_user_repo, mock_session_store):
|
||||
"""空 openid 返回错误"""
|
||||
use_case = WechatSyncUseCase(
|
||||
mock_user_repo,
|
||||
session_store=mock_session_store,
|
||||
jwt_secret_key="test-secret-key-for-jwt-12345",
|
||||
)
|
||||
request = WechatSyncRequest(openid="")
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert error == "openid is required"
|
||||
assert "openid is required" in error
|
||||
|
||||
def test_exception_handling(self, use_case, mock_user_repo):
|
||||
"""测试异常处理"""
|
||||
def test_exception_returns_error(self, mock_user_repo, mock_session_store):
|
||||
"""异常时返回友好错误"""
|
||||
mock_user_repo.find_by_wechat_openid.side_effect = Exception("DB error")
|
||||
|
||||
request = WechatSyncRequest(openid="openid123")
|
||||
use_case = WechatSyncUseCase(
|
||||
mock_user_repo,
|
||||
session_store=mock_session_store,
|
||||
jwt_secret_key="test-secret-key-for-jwt-12345",
|
||||
)
|
||||
request = WechatSyncRequest(openid="openid_123")
|
||||
response, error = use_case.execute(request)
|
||||
|
||||
assert response is None
|
||||
assert "Internal error" in error
|
||||
assert "DB error" in error
|
||||
|
||||
# ===== Session 保存验证 =====
|
||||
|
||||
def test_session_saved_with_correct_params(self, use_case, mock_user_repo, mock_session_store, test_user):
|
||||
"""测试 session 保存参数正确"""
|
||||
mock_user_repo.find_by_wechat_openid.return_value = test_user
|
||||
class TestWechatSyncSession:
|
||||
"""Session 相关测试"""
|
||||
|
||||
request = WechatSyncRequest(openid="openid123", source="h5")
|
||||
def test_session_saved(self, mock_user_repo, mock_session_store, sample_user):
|
||||
"""登录时保存 session"""
|
||||
mock_user_repo.find_by_wechat_openid.return_value = sample_user
|
||||
mock_user_repo.save.return_value = sample_user
|
||||
|
||||
use_case = WechatSyncUseCase(
|
||||
mock_user_repo,
|
||||
session_store=mock_session_store,
|
||||
jwt_secret_key="test-secret-key-for-jwt-12345",
|
||||
)
|
||||
request = WechatSyncRequest(openid="openid_123", source="miniapp")
|
||||
use_case.execute(request)
|
||||
|
||||
mock_session_store.save_session.assert_called_once()
|
||||
kwargs = mock_session_store.save_session.call_args.kwargs
|
||||
assert kwargs["user_id"] == "user-123"
|
||||
assert kwargs["refresh_token"] != ""
|
||||
assert "wechat_h5" in kwargs["device_info"]
|
||||
assert kwargs["ip_address"] == "bff_gateway"
|
||||
assert kwargs["expires_in_seconds"] == 30 * 24 * 3600 # 30天
|
||||
call_kwargs = mock_session_store.save_session.call_args[1]
|
||||
assert call_kwargs["user_id"] == "user_001"
|
||||
assert "wechat_miniapp" in call_kwargs["device_info"]
|
||||
assert call_kwargs["expires_in_seconds"] == 30 * 24 * 3600
|
||||
|
||||
Reference in New Issue
Block a user