diff --git a/tests/unit/test_assets_use_cases.py b/tests/unit/test_assets_use_cases.py new file mode 100755 index 000000000..bbed9d0fc --- /dev/null +++ b/tests/unit/test_assets_use_cases.py @@ -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" diff --git a/tests/unit/test_jwt_handler.py b/tests/unit/test_jwt_handler.py new file mode 100755 index 000000000..1c8859539 --- /dev/null +++ b/tests/unit/test_jwt_handler.py @@ -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 diff --git a/tests/unit/test_password_handler.py b/tests/unit/test_password_handler.py new file mode 100755 index 000000000..66a068a36 --- /dev/null +++ b/tests/unit/test_password_handler.py @@ -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 diff --git a/tests/unit/test_video_share_use_cases.py b/tests/unit/test_video_share_use_cases.py new file mode 100755 index 000000000..750dc8cad --- /dev/null +++ b/tests/unit/test_video_share_use_cases.py @@ -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()