test: P3-1 第37波单元测试(text_splitter/pagination/password_hasher/bind_contact) #818

Merged
xiaoxia merged 2 commits from test/unit-test-wave37 into develop 2026-07-24 16:45:32 +08:00
8 changed files with 2000 additions and 745 deletions
+197
View File
@@ -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"
+484
View File
@@ -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
+169
View File
@@ -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
+155 -248
View File
@@ -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
+175
View File
@@ -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
View File
@@ -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
+114 -282
View File
@@ -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
+510
View File
@@ -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()