test: 新增认证模块和封面服务单元测试,覆盖率提升至96%+ #654

Merged
auto-approve-bot merged 2 commits from test/unit-test-coverage-improvement into develop 2026-07-20 20:43:22 +08:00
3 changed files with 1069 additions and 0 deletions
+400
View File
@@ -0,0 +1,400 @@
"""
封面管理服务单元测试
"""
import subprocess
from pathlib import Path
from unittest.mock import Mock, patch
import pytest
from apps.api.app.services.cover_service import (
COVER_STORAGE_PREFIX,
DEFAULT_COVER_HEIGHT,
DEFAULT_COVER_QUALITY,
DEFAULT_COVER_WIDTH,
CoverService,
)
class TestGetCoverConfig:
"""get_cover_config 静态方法测试"""
def test_get_cover_config_default(self):
"""测试默认封面配置"""
config = {}
result = CoverService.get_cover_config(config)
assert result["type"] == "ai_frame"
assert result["image_url"] == ""
assert result["frame_time"] is None
def test_get_cover_config_with_custom_values(self):
"""测试自定义封面配置"""
config = {
"cover": {
"type": "manual",
"image_url": "https://example.com/cover.jpg",
"frame_time": 5.5,
}
}
result = CoverService.get_cover_config(config)
assert result["type"] == "manual"
assert result["image_url"] == "https://example.com/cover.jpg"
assert result["frame_time"] == 5.5
def test_get_cover_config_cover_not_dict(self):
"""测试 cover 不是 dict 时返回默认值"""
config = {"cover": "not-a-dict"}
result = CoverService.get_cover_config(config)
assert result["type"] == "ai_frame"
assert result["image_url"] == ""
assert result["frame_time"] is None
def test_get_cover_config_partial_fields(self):
"""测试部分字段存在时,其余字段用默认值"""
config = {"cover": {"type": "custom"}}
result = CoverService.get_cover_config(config)
assert result["type"] == "custom"
assert result["image_url"] == ""
assert result["frame_time"] is None
def test_get_cover_config_empty_cover_dict(self):
"""测试空的 cover dict"""
config = {"cover": {}}
result = CoverService.get_cover_config(config)
assert result["type"] == "ai_frame"
assert result["image_url"] == ""
class TestExtractCoverFromClip:
"""extract_cover_from_clip 测试"""
@pytest.fixture
def mock_storage(self):
storage = Mock()
storage.download_file = Mock()
storage.upload_file = Mock()
storage.get_url = Mock(return_value="https://oss.example.com/covers/plan1/cover_1000.jpg")
return storage
@pytest.fixture
def mock_asset_repo(self):
repo = Mock()
repo.get = Mock(return_value=None)
return repo
@pytest.fixture
def video_asset(self):
asset = Mock()
asset.storage_key = "videos/test-video.mp4"
asset.mime_type = "video/mp4"
return asset
@pytest.fixture
def service(self, mock_storage, mock_asset_repo):
return CoverService(storage_service=mock_storage, asset_repository=mock_asset_repo)
def test_extract_cover_asset_not_found(self, service, mock_asset_repo):
"""测试素材不存在时报错"""
mock_asset_repo.get.return_value = None
with pytest.raises(ValueError, match="素材不存在"):
service.extract_cover_from_clip(plan_id="plan-1", asset_id="nonexistent")
def test_extract_cover_asset_no_storage_key(self, service, mock_asset_repo):
"""测试素材没有文件时报错"""
asset = Mock()
asset.storage_key = ""
asset.mime_type = "video/mp4"
mock_asset_repo.get.return_value = asset
with pytest.raises(ValueError, match="素材没有文件"):
service.extract_cover_from_clip(plan_id="plan-1", asset_id="asset-no-file")
def test_extract_cover_asset_not_video(self, service, mock_asset_repo):
"""测试非视频素材报错"""
asset = Mock()
asset.storage_key = "images/photo.jpg"
asset.mime_type = "image/jpeg"
mock_asset_repo.get.return_value = asset
with pytest.raises(ValueError, match="素材不是视频类型"):
service.extract_cover_from_clip(plan_id="plan-1", asset_id="asset-img")
def test_extract_cover_download_failure(self, service, mock_asset_repo, mock_storage, video_asset):
"""测试下载素材失败"""
mock_asset_repo.get.return_value = video_asset
mock_storage.download_file.side_effect = Exception("网络错误")
with pytest.raises(RuntimeError, match="下载素材失败"):
service.extract_cover_from_clip(plan_id="plan-1", asset_id="asset-1")
def test_extract_cover_upload_failure(self, service, mock_asset_repo, mock_storage, video_asset):
"""测试上传封面失败"""
mock_asset_repo.get.return_value = video_asset
def fake_download(storage_key, local_path):
# 创建一个假的视频文件
Path(local_path).parent.mkdir(parents=True, exist_ok=True)
with open(local_path, "wb") as f:
f.write(b"fake video data")
mock_storage.download_file.side_effect = fake_download
mock_storage.upload_file.side_effect = Exception("上传失败")
# mock _extract_frame 避免真的调 ffmpeg
with patch.object(CoverService, "_extract_frame") as mock_extract:
def fake_extract(video_path, output_path, **kwargs):
# 创建假的封面文件
with open(output_path, "wb") as f:
f.write(b"\xff\xd8\xff\xe0fake jpeg data")
mock_extract.side_effect = fake_extract
with pytest.raises(RuntimeError, match="上传封面失败"):
service.extract_cover_from_clip(plan_id="plan-1", asset_id="asset-1")
def test_extract_cover_get_url_falls_back_to_key(self, service, mock_asset_repo, mock_storage, video_asset):
"""测试获取 URL 失败时降级为 storage_key"""
mock_asset_repo.get.return_value = video_asset
def fake_download(storage_key, local_path):
Path(local_path).parent.mkdir(parents=True, exist_ok=True)
with open(local_path, "wb") as f:
f.write(b"fake video data")
mock_storage.download_file.side_effect = fake_download
mock_storage.get_url.side_effect = Exception("URL服务不可用")
with patch.object(CoverService, "_extract_frame") as mock_extract:
def fake_extract(video_path, output_path, **kwargs):
with open(output_path, "wb") as f:
f.write(b"\xff\xd8\xff\xe0fake jpeg")
mock_extract.side_effect = fake_extract
result = service.extract_cover_from_clip(plan_id="plan-abc", asset_id="asset-xyz", frame_time=2.5)
assert result["type"] == "manual"
assert result["frame_time"] == 2.5
# URL 失败时返回 storage_key
assert COVER_STORAGE_PREFIX in result["image_url"]
assert "plan-abc" in result["image_url"]
def test_extract_cover_success(self, service, mock_asset_repo, mock_storage, video_asset):
"""测试抽帧成功完整流程"""
mock_asset_repo.get.return_value = video_asset
def fake_download(storage_key, local_path):
Path(local_path).parent.mkdir(parents=True, exist_ok=True)
with open(local_path, "wb") as f:
f.write(b"fake video data for testing")
mock_storage.download_file.side_effect = fake_download
with patch.object(CoverService, "_extract_frame") as mock_extract:
def fake_extract(video_path, output_path, **kwargs):
with open(output_path, "wb") as f:
f.write(b"\xff\xd8\xff\xe0fake jpeg image data")
mock_extract.side_effect = fake_extract
result = service.extract_cover_from_clip(
plan_id="plan-123",
asset_id="asset-456",
frame_time=3.0,
width=720,
height=1280,
quality=3,
)
assert result["type"] == "manual"
assert result["image_url"] == "https://oss.example.com/covers/plan1/cover_1000.jpg"
assert result["frame_time"] == 3.0
# 验证上传被调用
mock_storage.upload_file.assert_called_once()
upload_args = mock_storage.upload_file.call_args[1]
assert upload_args["content_type"] == "image/jpeg"
assert "plan-123" in upload_args["storage_key"]
assert "3000" in upload_args["storage_key"] # frame_time * 1000
# 验证 _extract_frame 被调用且参数正确
mock_extract.assert_called_once()
extract_kwargs = mock_extract.call_args[1]
assert extract_kwargs["time_sec"] == 3.0
assert extract_kwargs["width"] == 720
assert extract_kwargs["height"] == 1280
assert extract_kwargs["quality"] == 3
class TestGenerateSmartCover:
"""generate_smart_cover 测试"""
@pytest.fixture
def service(self):
return CoverService(storage_service=Mock(), asset_repository=Mock())
def test_generate_smart_cover_calls_extract_with_default_time(self, service):
"""测试智能封面调用 extract_cover_from_clip 并设置 type 为 ai_frame"""
fake_result = {"type": "manual", "image_url": "test.jpg", "frame_time": 3.0}
with patch.object(service, "extract_cover_from_clip", return_value=fake_result) as mock_extract:
result = service.generate_smart_cover(plan_id="plan-1", asset_id="asset-1")
mock_extract.assert_called_once()
call_kwargs = mock_extract.call_args[1]
assert call_kwargs["plan_id"] == "plan-1"
assert call_kwargs["asset_id"] == "asset-1"
assert call_kwargs["frame_time"] == 3.0 # 默认第3秒
assert result["type"] == "ai_frame"
assert result["image_url"] == "test.jpg"
def test_generate_smart_cover_passes_dimensions(self, service):
"""测试智能封面传递尺寸和质量参数"""
fake_result = {"type": "manual", "image_url": "test.jpg", "frame_time": 3.0}
with patch.object(service, "extract_cover_from_clip", return_value=fake_result) as mock_extract:
service.generate_smart_cover(
plan_id="plan-1",
asset_id="asset-1",
width=1080,
height=1920,
quality=5,
)
call_kwargs = mock_extract.call_args[1]
assert call_kwargs["width"] == 1080
assert call_kwargs["height"] == 1920
assert call_kwargs["quality"] == 5
class TestExtractFrame:
"""_extract_frame 静态方法测试(mock subprocess"""
@pytest.fixture
def video_path(self, tmp_path):
path = tmp_path / "test_video.mp4"
path.write_bytes(b"fake video")
return path
@pytest.fixture
def output_path(self, tmp_path):
return tmp_path / "cover.jpg"
def test_extract_frame_success(self, video_path, output_path):
"""测试 FFmpeg 抽帧成功"""
fake_result = Mock()
fake_result.returncode = 0
with patch("subprocess.run", return_value=fake_result) as mock_run:
CoverService._extract_frame(
video_path=video_path,
output_path=output_path,
time_sec=2.5,
width=1080,
height=1920,
quality=5,
)
assert mock_run.call_count == 1
cmd = mock_run.call_args[0][0]
assert cmd[0] == "ffmpeg"
assert "-ss" in cmd
assert "2.500" in cmd
assert "-vframes" in cmd
# 验证 scale+crop 滤镜存在
vf_index = cmd.index("-vf") + 1
assert "scale=" in cmd[vf_index]
assert "crop=" in cmd[vf_index]
def test_extract_frame_fallback_to_simple_command(self, video_path, output_path):
"""测试主命令失败时回退到简化命令"""
fail_result = Mock()
fail_result.returncode = 1
fail_result.stderr = "Filter graph error"
success_result = Mock()
success_result.returncode = 0
call_count = 0
def fake_run(*args, **kwargs):
nonlocal call_count
call_count += 1
if call_count == 1:
return fail_result
return success_result
with patch("subprocess.run", side_effect=fake_run) as mock_run:
CoverService._extract_frame(
video_path=video_path,
output_path=output_path,
time_sec=1.0,
width=1080,
height=1920,
quality=5,
)
assert mock_run.call_count == 2
# 第二次是简化命令(没有 -vf 参数)
second_cmd = mock_run.call_args_list[1][0][0]
assert "-vf" not in second_cmd
def test_extract_frame_both_commands_fail(self, video_path, output_path):
"""测试两个命令都失败时报错"""
fail_result = Mock()
fail_result.returncode = 1
fail_result.stderr = "Invalid data found when processing input"
with patch("subprocess.run", return_value=fail_result):
with pytest.raises(RuntimeError, match="FFmpeg 抽帧失败"):
CoverService._extract_frame(
video_path=video_path,
output_path=output_path,
time_sec=1.0,
width=1080,
height=1920,
quality=5,
)
def test_extract_frame_timeout(self, video_path, output_path):
"""测试 FFmpeg 抽帧超时"""
with patch("subprocess.run", side_effect=subprocess.TimeoutExpired(cmd="ffmpeg", timeout=60)):
with pytest.raises(RuntimeError, match="FFmpeg 抽帧超时"):
CoverService._extract_frame(
video_path=video_path,
output_path=output_path,
time_sec=1.0,
width=1080,
height=1920,
quality=5,
)
def test_extract_frame_ffmpeg_not_found(self, video_path, output_path):
"""测试 FFmpeg 不可用"""
with patch("subprocess.run", side_effect=FileNotFoundError("ffmpeg not found")):
with pytest.raises(RuntimeError, match="FFmpeg 不可用"):
CoverService._extract_frame(
video_path=video_path,
output_path=output_path,
time_sec=1.0,
width=1080,
height=1920,
quality=5,
)
class TestDefaults:
"""默认常量测试"""
def test_default_dimensions(self):
"""测试默认尺寸常量"""
assert DEFAULT_COVER_WIDTH == 1080
assert DEFAULT_COVER_HEIGHT == 1920
assert DEFAULT_COVER_QUALITY == 5
assert COVER_STORAGE_PREFIX == "covers"
+245
View File
@@ -0,0 +1,245 @@
"""
JWT Service 单元测试
"""
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,
)
class TestJWTConfig:
"""JWT 配置测试"""
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"):
JWTConfig(secret_key="")
def test_config_whitespace_secret_raises(self):
"""测试全空格密钥报错"""
with pytest.raises(ValueError, match="secret_key must be provided"):
JWTConfig(secret_key=" ")
def test_config_insecure_default_secret_raises(self):
"""测试不安全的默认密钥报错"""
insecure_keys = [
"your-secret-key-change-in-production",
"your-secret-key",
"secret",
"changeme",
"password",
"YOUR-SECRET-KEY",
"Secret",
]
for key in insecure_keys:
with pytest.raises(ValueError, match="insecure"):
JWTConfig(secret_key=key)
class TestJWTService:
"""JWT 服务测试"""
@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,
)
@pytest.fixture
def service(self, config):
return JWTService(config=config)
def test_service_init_without_config_raises(self):
"""测试无 config 初始化报错"""
with pytest.raises(ValueError, match="requires a JWTConfig"):
JWTService(config=None)
# --- create_access_token ---
def test_create_access_token_success(self, service):
"""测试创建 access token 成功"""
token = service.create_access_token(user_id="user-123")
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_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])
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"},
)
payload = jwt.decode(token, config.SECRET_KEY, algorithms=[config.ALGORITHM])
assert payload["custom_field"] == "custom_value"
assert payload["sid"] == "session-abc"
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])
assert "iat" in payload
assert "exp" in payload
iat = datetime.utcfromtimestamp(payload["iat"])
exp = datetime.utcfromtimestamp(payload["exp"])
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 ---
def test_create_refresh_token_success(self, service):
"""测试创建 refresh token 成功"""
token = service.create_refresh_token(user_id="user-123", session_id="sess-abc")
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"
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
# --- verify_token ---
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 报错"""
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")
with pytest.raises(InvalidTokenError):
service.verify_token(token)
# --- verify_access_token ---
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
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")
with pytest.raises(ValueError, match="Token type must be 'access'"):
service.verify_access_token(refresh_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
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")
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
+424
View File
@@ -0,0 +1,424 @@
"""
登录/登出/刷新令牌 Use Case 测试
"""
from unittest.mock import Mock, patch
import pytest
from packages.application.auth.login_use_case import (
LoginRequest,
LoginUseCase,
LogoutRequest,
LogoutUseCase,
RefreshTokenRequest,
RefreshTokenUseCase,
_is_legacy_sha256_hash,
_legacy_sha256,
)
from packages.domain.entities import User
class TestLegacyHashHelpers:
"""旧版密码哈希工具函数测试"""
def test_is_legacy_sha256_hash_valid(self):
"""测试识别有效的 SHA256 哈希"""
valid_hash = "a" * 64 # 64个十六进制字符
assert _is_legacy_sha256_hash(valid_hash) is True
def test_is_legacy_sha256_hash_wrong_length(self):
"""测试长度不对的不是 SHA256"""
assert _is_legacy_sha256_hash("abc123") is False
assert _is_legacy_sha256_hash("a" * 63) is False
assert _is_legacy_sha256_hash("a" * 65) is False
def test_is_legacy_sha256_hash_non_hex(self):
"""测试包含非十六进制字符的不是 SHA256"""
non_hex = "g" * 64
assert _is_legacy_sha256_hash(non_hex) is False
def test_legacy_sha256_produces_correct_hash(self):
"""测试 SHA256 哈希生成正确"""
result = _legacy_sha256("password123")
assert len(result) == 64
assert all(c in "0123456789abcdef" for c in result)
# 相同输入产生相同输出
assert _legacy_sha256("password123") == result
class TestLoginUseCase:
"""登录用例测试"""
@pytest.fixture
def mock_user_repo(self):
"""Mock 用户仓储"""
repo = Mock()
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()
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)
store.delete_all_user_sessions = Mock()
return store
@pytest.fixture
def test_user(self):
"""测试用户"""
user = User(
id="user-123",
email="test@example.com",
username="testuser",
display_name="Test User",
password_hash="hashed_password",
)
return user
@pytest.fixture
def use_case(self, mock_user_repo, mock_session_store):
"""创建登录用例(使用测试用JWT密钥)"""
return LoginUseCase(
user_repository=mock_user_repo,
session_store=mock_session_store,
jwt_secret_key="test-secret-key-for-unit-tests",
)
def test_login_success(self, use_case, mock_user_repo, mock_session_store, test_user):
"""测试登录成功"""
mock_user_repo.find_by_email.return_value = test_user
with patch("packages.application.auth.login_use_case.password_hasher") as mock_hasher:
mock_hasher.verify_password.return_value = True
request = LoginRequest(
email="test@example.com",
password="CorrectPass123",
device_info="Test Device",
ip_address="192.168.1.1",
)
response, error = use_case.execute(request)
assert error is None
assert response is not None
assert response.user_id == "user-123"
assert response.email == "test@example.com"
assert response.username == "testuser"
assert response.display_name == "Test User"
assert response.access_token != ""
assert response.refresh_token != ""
assert response.expires_in > 0
# 验证 session 已保存
mock_session_store.save_session.assert_called_once()
save_args = mock_session_store.save_session.call_args[1]
assert save_args["user_id"] == "user-123"
assert save_args["device_info"] == "Test Device"
assert save_args["ip_address"] == "192.168.1.1"
# 验证最后登录信息已更新
mock_user_repo.save.assert_called()
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 == "192.168.1.1"
def test_login_email_empty(self, use_case):
"""测试邮箱为空"""
request = LoginRequest(email="", password="password123")
response, error = use_case.execute(request)
assert response is None
assert error == "Email is required"
def test_login_password_empty(self, use_case, mock_user_repo):
"""测试密码为空"""
mock_user_repo.find_by_email.return_value = Mock() # 即使有用户也应该在密码检查前失败
request = LoginRequest(email="test@example.com", password="")
response, error = use_case.execute(request)
assert response is None
assert error == "Password is required"
def test_login_user_not_found(self, use_case, mock_user_repo):
"""测试用户不存在"""
mock_user_repo.find_by_email.return_value = None
request = LoginRequest(email="nonexistent@example.com", password="password123")
response, error = use_case.execute(request)
assert response is None
assert error == "Invalid email or password"
def test_login_wrong_password(self, use_case, mock_user_repo, test_user):
"""测试密码错误"""
mock_user_repo.find_by_email.return_value = test_user
with patch("packages.application.auth.login_use_case.password_hasher") as mock_hasher:
mock_hasher.verify_password.return_value = False
request = LoginRequest(email="test@example.com", password="WrongPass")
response, error = use_case.execute(request)
assert response is None
assert error == "Invalid email or password"
def test_login_legacy_sha256_password_success_and_upgrade(self, use_case, mock_user_repo, mock_session_store):
"""测试旧版 SHA256 密码登录成功并自动升级哈希"""
legacy_hash = _legacy_sha256("OldPassword123")
legacy_user = User(
id="user-legacy",
email="legacy@example.com",
username="legacyuser",
display_name="Legacy User",
password_hash=legacy_hash,
)
mock_user_repo.find_by_email.return_value = legacy_user
with patch("packages.application.auth.login_use_case.password_hasher") as mock_hasher:
mock_hasher.verify_password.return_value = False # 现代哈希验证失败
mock_hasher.hash_password.return_value = "new_bcrypt_hash"
request = LoginRequest(email="legacy@example.com", password="OldPassword123")
response, error = use_case.execute(request)
assert error is None
assert response is not None
assert response.user_id == "user-legacy"
# 验证密码哈希已升级
mock_user_repo.save.assert_called()
saved_user = mock_user_repo.save.call_args[0][0]
assert saved_user.password_hash == "new_bcrypt_hash"
def test_login_legacy_sha256_password_wrong(self, use_case, mock_user_repo):
"""测试旧版 SHA256 密码错误"""
legacy_hash = _legacy_sha256("CorrectPassword")
legacy_user = User(
id="user-legacy",
email="legacy@example.com",
username="legacyuser",
display_name="Legacy User",
password_hash=legacy_hash,
)
mock_user_repo.find_by_email.return_value = legacy_user
with patch("packages.application.auth.login_use_case.password_hasher") as mock_hasher:
mock_hasher.verify_password.return_value = False
request = LoginRequest(email="legacy@example.com", password="WrongPassword")
response, error = use_case.execute(request)
assert response is None
assert error == "Invalid email or password"
def test_login_email_normalized_to_lowercase(self, use_case, mock_user_repo, test_user):
"""测试邮箱自动转小写并去空格"""
mock_user_repo.find_by_email.return_value = test_user
with patch("packages.application.auth.login_use_case.password_hasher") as mock_hasher:
mock_hasher.verify_password.return_value = True
request = LoginRequest(email=" TEST@Example.COM ", password="pass123")
response, error = use_case.execute(request)
assert error is None
assert response is not None
# find_by_email 应该收到小写去空格后的邮箱
mock_user_repo.find_by_email.assert_called_with("test@example.com")
def test_login_default_device_and_ip(self, use_case, mock_user_repo, mock_session_store, test_user):
"""测试设备信息和IP的默认值"""
mock_user_repo.find_by_email.return_value = test_user
with patch("packages.application.auth.login_use_case.password_hasher") as mock_hasher:
mock_hasher.verify_password.return_value = True
request = LoginRequest(email="test@example.com", password="pass123")
response, error = use_case.execute(request)
assert error is None
save_args = mock_session_store.save_session.call_args[1]
assert save_args["device_info"] == "Unknown"
assert save_args["ip_address"] == "unknown"
class TestRefreshTokenUseCase:
"""刷新令牌用例测试"""
@pytest.fixture
def mock_user_repo(self):
repo = Mock()
repo.get = Mock(return_value=None)
return repo
@pytest.fixture
def mock_session_store(self):
store = Mock()
store.get_session_by_refresh_token = Mock(return_value=None)
store.get_refresh_token = Mock(return_value=None)
return store
@pytest.fixture
def test_user(self):
return User(
id="user-123",
email="test@example.com",
username="testuser",
display_name="Test User",
password_hash="hashed",
)
@pytest.fixture
def use_case(self, mock_user_repo, mock_session_store):
# 用 patch 替换 jwt_service.config
with patch("packages.application.auth.login_use_case.jwt_service") as mock_jwt:
mock_jwt.config.SECRET_KEY = "test-secret-key"
mock_jwt.config.ALGORITHM = "HS256"
mock_jwt.config.ACCESS_TOKEN_EXPIRE_MINUTES = 30
uc = RefreshTokenUseCase(
user_repository=mock_user_repo,
session_store=mock_session_store,
)
uc._jwt_secret_key = "test-secret-key"
uc.jwt_service.config = mock_jwt.config
yield uc
def test_refresh_success(self, use_case, mock_user_repo, mock_session_store, test_user):
"""测试刷新令牌成功"""
mock_session_store.get_session_by_refresh_token.return_value = {
"session_id": "sess-abc",
"user_id": "user-123",
}
mock_session_store.get_refresh_token.return_value = "valid-refresh-token"
mock_user_repo.get.return_value = test_user
request = RefreshTokenRequest(refresh_token="valid-refresh-token")
response, error = use_case.execute(request)
assert error is None
assert response is not None
assert response.user_id == "user-123"
assert response.email == "test@example.com"
assert response.access_token != ""
assert response.refresh_token == "valid-refresh-token"
def test_refresh_token_empty(self, use_case):
"""测试 refresh_token 为空"""
request = RefreshTokenRequest(refresh_token="")
response, error = use_case.execute(request)
assert response is None
assert error == "Refresh token is required"
def test_refresh_invalid_token(self, use_case, mock_session_store):
"""测试无效的 refresh_token"""
mock_session_store.get_session_by_refresh_token.return_value = None
request = RefreshTokenRequest(refresh_token="invalid-token")
response, error = use_case.execute(request)
assert response is None
assert error == "Invalid or expired refresh token"
def test_refresh_session_data_invalid(self, use_case, mock_session_store):
"""测试 session 数据不完整"""
mock_session_store.get_session_by_refresh_token.return_value = {
"session_id": "sess-abc",
# 缺少 user_id
}
request = RefreshTokenRequest(refresh_token="some-token")
response, error = use_case.execute(request)
assert response is None
assert error == "Invalid session data"
def test_refresh_token_mismatch(self, use_case, mock_user_repo, mock_session_store, test_user):
"""测试 refresh_token 不匹配"""
mock_session_store.get_session_by_refresh_token.return_value = {
"session_id": "sess-abc",
"user_id": "user-123",
}
mock_session_store.get_refresh_token.return_value = "different-token"
mock_user_repo.get.return_value = test_user
request = RefreshTokenRequest(refresh_token="user-provided-token")
response, error = use_case.execute(request)
assert response is None
assert error == "Refresh token mismatch"
def test_refresh_user_not_found(self, use_case, mock_user_repo, mock_session_store):
"""测试用户不存在"""
mock_session_store.get_session_by_refresh_token.return_value = {
"session_id": "sess-abc",
"user_id": "user-nonexistent",
}
mock_session_store.get_refresh_token.return_value = "valid-token"
mock_user_repo.get.return_value = None
request = RefreshTokenRequest(refresh_token="valid-token")
response, error = use_case.execute(request)
assert response is None
assert error == "User not found"
class TestLogoutUseCase:
"""登出用例测试"""
@pytest.fixture
def mock_session_store(self):
store = Mock()
store.delete_session = Mock(return_value=True)
store.delete_all_user_sessions = Mock()
return store
@pytest.fixture
def use_case(self, mock_session_store):
return LogoutUseCase(session_store=mock_session_store)
def test_logout_single_device_success(self, use_case, mock_session_store):
"""测试单设备登出成功"""
request = LogoutRequest(user_id="user-123", session_id="sess-abc")
success, error = use_case.execute(request)
assert success is True
assert error is None
mock_session_store.delete_session.assert_called_once_with("sess-abc")
mock_session_store.delete_all_user_sessions.assert_not_called()
def test_logout_all_devices(self, use_case, mock_session_store):
"""测试所有设备登出"""
request = LogoutRequest(user_id="user-123", logout_all_devices=True)
success, error = use_case.execute(request)
assert success is True
assert error is None
mock_session_store.delete_all_user_sessions.assert_called_once_with("user-123")
mock_session_store.delete_session.assert_not_called()
def test_logout_missing_session_id(self, use_case):
"""测试缺少 session_id"""
request = LogoutRequest(user_id="user-123", session_id=None)
success, error = use_case.execute(request)
assert success is False
assert error == "Session ID is required"
def test_logout_session_not_found(self, use_case, mock_session_store):
"""测试 session 不存在"""
mock_session_store.delete_session.return_value = False
request = LogoutRequest(user_id="user-123", session_id="nonexistent-sess")
success, error = use_case.execute(request)
assert success is False
assert error == "Session not found"