test: 新增认证模块和封面服务单元测试,覆盖率提升至96%+ #654
Executable
+400
@@ -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"
|
||||
Executable
+245
@@ -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
|
||||
Executable
+424
@@ -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"
|
||||
Reference in New Issue
Block a user