diff --git a/packages/adapters/in_memory/user_repository.py b/packages/adapters/in_memory/user_repository.py index 68db1479c..72cc8118d 100755 --- a/packages/adapters/in_memory/user_repository.py +++ b/packages/adapters/in_memory/user_repository.py @@ -2,6 +2,7 @@ 用户仓储 In-Memory 实现 """ +import copy from typing import Dict, Optional from packages.domain.entities import User @@ -22,28 +23,59 @@ class InMemoryUserRepository(UserRepository): self._phone_index: Dict[str, str] = {} # phone -> user_id def save(self, user: User) -> None: - """保存用户""" - # 如果是更新,先清理旧索引 - old = self._users.get(user.id) - if old: - self._email_index.pop(old.email.lower(), None) - if old.username: - self._username_index.pop(old.username.lower(), None) - if old.email_verification_token: - self._verification_token_index.pop(old.email_verification_token, None) - if old.password_reset_token: - self._reset_token_index.pop(old.password_reset_token, None) - if old.wechat_openid: - self._wechat_openid_index.pop(old.wechat_openid, None) - if old.wechat_unionid: - self._wechat_unionid_index.pop(old.wechat_unionid, None) - if old.phone: - self._phone_index.pop(old.phone, None) + """保存用户(存储独立副本,避免外部修改影响内部状态)""" + # 第一步:唯一性约束检查(O(1),全部检查通过再动数据) + # email 是必填字段,做空值防御 + if not user.email: + raise ValueError("User email cannot be empty") + new_email = user.email.lower() + existing_id = self._email_index.get(new_email) + if existing_id and existing_id != user.id: + raise ValueError(f"Email already in use: {user.email}") - self._users[user.id] = user - self._email_index[user.email.lower()] = user.id - if user.username: - self._username_index[user.username.lower()] = user.id + new_username = user.username.lower() if user.username else None + if new_username: + existing_id = self._username_index.get(new_username) + if existing_id and existing_id != user.id: + raise ValueError(f"Username already in use: {user.username}") + + if user.phone: + existing_id = self._phone_index.get(user.phone) + if existing_id and existing_id != user.id: + raise ValueError(f"Phone already in use: {user.phone}") + + if user.wechat_openid: + existing_id = self._wechat_openid_index.get(user.wechat_openid) + if existing_id and existing_id != user.id: + raise ValueError(f"WeChat openid already in use: {user.wechat_openid}") + + if user.wechat_unionid: + existing_id = self._wechat_unionid_index.get(user.wechat_unionid) + if existing_id and existing_id != user.id: + raise ValueError(f"WeChat unionid already in use: {user.wechat_unionid}") + + # 验证令牌和重置令牌也做唯一性防御(防止生成器异常导致重复) + if user.email_verification_token: + existing_id = self._verification_token_index.get(user.email_verification_token) + if existing_id and existing_id != user.id: + raise ValueError("Email verification token already in use") + + if user.password_reset_token: + existing_id = self._reset_token_index.get(user.password_reset_token) + if existing_id and existing_id != user.id: + raise ValueError("Password reset token already in use") + + # 第二步:如果是更新,用旧对象属性清理旧索引(O(1),因存储的是独立副本) + old_user = self._users.get(user.id) + if old_user is not None: + self._remove_indexes_of_user(old_user) + + # 第三步:存储独立副本 + 写入新索引 + stored_user = copy.copy(user) + self._users[user.id] = stored_user + self._email_index[new_email] = user.id + if new_username: + self._username_index[new_username] = user.id if user.email_verification_token: self._verification_token_index[user.email_verification_token] = user.id if user.password_reset_token: @@ -55,12 +87,31 @@ class InMemoryUserRepository(UserRepository): if user.phone: self._phone_index[user.phone] = user.id + def _remove_indexes_of_user(self, user: User) -> None: + """利用已知用户对象属性清理所有索引,时间复杂度 O(1).""" + if user.email: + self._email_index.pop(user.email.lower(), None) + if user.username: + self._username_index.pop(user.username.lower(), None) + if user.email_verification_token: + self._verification_token_index.pop(user.email_verification_token, None) + if user.password_reset_token: + self._reset_token_index.pop(user.password_reset_token, None) + if user.wechat_openid: + self._wechat_openid_index.pop(user.wechat_openid, None) + if user.wechat_unionid: + self._wechat_unionid_index.pop(user.wechat_unionid, None) + if user.phone: + self._phone_index.pop(user.phone, None) + def find_by_id(self, user_id: str) -> Optional[User]: """根据 ID 查找用户""" return self._users.get(user_id) def find_by_email(self, email: str) -> Optional[User]: """根据邮箱查找用户""" + if not email: + return None user_id = self._email_index.get(email.lower()) if user_id: return self._users.get(user_id) @@ -68,6 +119,8 @@ class InMemoryUserRepository(UserRepository): def find_by_username(self, username: str) -> Optional[User]: """根据用户名查找用户""" + if not username: + return None user_id = self._username_index.get(username.lower()) if user_id: return self._users.get(user_id) @@ -115,17 +168,11 @@ class InMemoryUserRepository(UserRepository): def delete(self, user_id: str) -> bool: """删除用户""" user = self._users.get(user_id) - if not user: + if user is None: return False - # 清理索引 - self._email_index.pop(user.email.lower(), None) - if user.username: - self._username_index.pop(user.username.lower(), None) - if user.email_verification_token: - self._verification_token_index.pop(user.email_verification_token, None) - if user.password_reset_token: - self._reset_token_index.pop(user.password_reset_token, None) + # O(1) 清理所有索引 + self._remove_indexes_of_user(user) # 删除用户 del self._users[user_id] diff --git a/tests/unit/test_inmemory_small_repos.py b/tests/unit/test_inmemory_small_repos.py new file mode 100755 index 000000000..87ed5f667 --- /dev/null +++ b/tests/unit/test_inmemory_small_repos.py @@ -0,0 +1,501 @@ +"""InMemory 小型仓储模块单元测试(asset_library/tag/project/ingest_job/classification_job).""" + +from datetime import datetime, timezone + +import pytest + +from packages.adapters.in_memory.asset_library_repository import InMemoryAssetLibraryRepository +from packages.adapters.in_memory.classification_job_repository import InMemoryClassificationJobRepository +from packages.adapters.in_memory.ingest_job_repository import InMemoryIngestJobRepository +from packages.adapters.in_memory.project_repository import InMemoryProjectRepository +from packages.adapters.in_memory.tag_repository import InMemoryTagRepository +from packages.domain.classification import ( + AssetLibraryKind, + ClassificationJob, + ClassificationJobStatus, + IngestJobStatus, +) +from packages.domain.entities import AssetLibrary, IngestJob, Project, User +from packages.domain.tag import Tag + +# ==================== AssetLibrary ==================== + + +class TestInMemoryAssetLibraryRepository: + @pytest.fixture + def repo(self): + return InMemoryAssetLibraryRepository() + + @pytest.fixture + def lib_video(self): + return AssetLibrary.create(project_id="p1", name="视频库", kind=AssetLibraryKind.VIDEO) + + @pytest.fixture + def lib_image(self): + return AssetLibrary.create(project_id="p1", name="图片库", kind=AssetLibraryKind.IMAGE) + + @pytest.fixture + def lib_other_project(self): + return AssetLibrary.create(project_id="p2", name="其他项目库", kind=AssetLibraryKind.VIDEO) + + def test_create_and_get(self, repo, lib_video): + result = repo.create(lib_video) + assert result.id == lib_video.id + assert result.name == "视频库" + + fetched = repo.get(lib_video.id) + assert fetched is not None + assert fetched.id == lib_video.id + + def test_get_nonexistent(self, repo): + assert repo.get("nonexistent") is None + + def test_find_by_id_alias(self, repo, lib_video): + repo.create(lib_video) + assert repo.find_by_id(lib_video.id).id == repo.get(lib_video.id).id + + def test_find_by_project(self, repo, lib_video, lib_image, lib_other_project): + repo.create(lib_video) + repo.create(lib_image) + repo.create(lib_other_project) + + p1_libs = repo.find_by_project("p1") + assert len(p1_libs) == 2 + + p2_libs = repo.find_by_project("p2") + assert len(p2_libs) == 1 + assert p2_libs[0].id == lib_other_project.id + + def test_find_by_project_with_kind_filter(self, repo, lib_video, lib_image): + repo.create(lib_video) + repo.create(lib_image) + + video_libs = repo.find_by_project("p1", kind=AssetLibraryKind.VIDEO) + assert len(video_libs) == 1 + assert video_libs[0].kind == AssetLibraryKind.VIDEO + + image_libs = repo.find_by_project("p1", kind=AssetLibraryKind.IMAGE) + assert len(image_libs) == 1 + + def test_find_by_project_empty(self, repo): + assert repo.find_by_project("nonexistent") == [] + + def test_update(self, repo, lib_video): + repo.create(lib_video) + lib_video.name = "新名称" + result = repo.update(lib_video) + assert result.name == "新名称" + assert repo.get(lib_video.id).name == "新名称" + + def test_delete(self, repo, lib_video): + repo.create(lib_video) + assert repo.delete(lib_video.id) is True + assert repo.get(lib_video.id) is None + + def test_delete_nonexistent(self, repo): + assert repo.delete("nonexistent") is False + + def test_increment_asset_count(self, repo, lib_video): + repo.create(lib_video) + repo.increment_asset_count(lib_video.id, 1024) + + lib = repo.get(lib_video.id) + assert lib.asset_count == 1 + assert lib.total_size == 1024 + + repo.increment_asset_count(lib_video.id, 512) + lib = repo.get(lib_video.id) + assert lib.asset_count == 2 + assert lib.total_size == 1536 + + def test_increment_asset_count_nonexistent(self, repo): + # 不报错,静默忽略 + repo.increment_asset_count("nonexistent", 100) + + def test_decrement_asset_count(self, repo, lib_video): + repo.create(lib_video) + repo.increment_asset_count(lib_video.id, 1024) + repo.increment_asset_count(lib_video.id, 512) + + repo.decrement_asset_count(lib_video.id, 512) + lib = repo.get(lib_video.id) + assert lib.asset_count == 1 + assert lib.total_size == 1024 + + def test_decrement_asset_count_not_below_zero(self, repo, lib_video): + repo.create(lib_video) + repo.decrement_asset_count(lib_video.id, 9999) + lib = repo.get(lib_video.id) + assert lib.asset_count == 0 + assert lib.total_size == 0 + + def test_decrement_asset_count_nonexistent(self, repo): + repo.decrement_asset_count("nonexistent", 100) + + +# ==================== Tag ==================== + + +class TestInMemoryTagRepository: + @pytest.fixture + def repo(self): + return InMemoryTagRepository() + + @pytest.fixture + def tag1(self): + return Tag.create(user_id="u1", name="风景") + + @pytest.fixture + def tag2(self): + return Tag.create(user_id="u1", name="人物") + + @pytest.fixture + def tag_other_user(self): + return Tag.create(user_id="u2", name="风景") + + def test_create_and_get(self, repo, tag1): + result = repo.create(tag1) + assert result.id == tag1.id + assert result.name == "风景" + + fetched = repo.get(tag1.id) + assert fetched is not None + assert fetched.id == tag1.id + + def test_get_nonexistent(self, repo): + assert repo.get("nonexistent") is None + + def test_find_by_name(self, repo, tag1, tag_other_user): + repo.create(tag1) + repo.create(tag_other_user) + + # 同用户同名 + found = repo.find_by_name("u1", "风景") + assert found is not None + assert found.id == tag1.id + + # 不同用户同名不冲突 + found2 = repo.find_by_name("u2", "风景") + assert found2 is not None + assert found2.id == tag_other_user.id + + def test_find_by_name_not_found(self, repo, tag1): + repo.create(tag1) + assert repo.find_by_name("u1", "不存在") is None + assert repo.find_by_name("u2", "风景") is None + + def test_list_by_user(self, repo, tag1, tag2, tag_other_user): + repo.create(tag1) + repo.create(tag2) + repo.create(tag_other_user) + + u1_tags = repo.list_by_user("u1") + assert len(u1_tags) == 2 + + u2_tags = repo.list_by_user("u2") + assert len(u2_tags) == 1 + assert u2_tags[0].id == tag_other_user.id + + def test_list_by_user_pagination(self, repo): + for i in range(5): + repo.create(Tag.create(user_id="u1", name=f"tag-{i}")) + + page1 = repo.list_by_user("u1", limit=2) + assert len(page1) == 2 + + page2 = repo.list_by_user("u1", skip=2, limit=2) + assert len(page2) == 2 + + def test_list_by_user_sorted_by_created_at_desc(self, repo): + t1 = Tag.create(user_id="u1", name="old") + t2 = Tag.create(user_id="u1", name="new") + repo.create(t1) + repo.create(t2) + + tags = repo.list_by_user("u1") + # 新创建的排前面 + assert tags[0].id == t2.id + assert tags[1].id == t1.id + + def test_count_by_user(self, repo, tag1, tag2, tag_other_user): + repo.create(tag1) + repo.create(tag2) + repo.create(tag_other_user) + + assert repo.count_by_user("u1") == 2 + assert repo.count_by_user("u2") == 1 + assert repo.count_by_user("u3") == 0 + + def test_delete(self, repo, tag1): + repo.create(tag1) + assert repo.delete(tag1.id) is True + assert repo.get(tag1.id) is None + + def test_delete_nonexistent(self, repo): + assert repo.delete("nonexistent") is False + + +# ==================== Project ==================== + + +class TestInMemoryProjectRepository: + @pytest.fixture + def repo(self): + return InMemoryProjectRepository() + + @pytest.fixture + def project1(self): + return Project(id="proj-1", owner_user_id="u1", name="项目一", shared_users=[]) + + @pytest.fixture + def project2(self): + return Project(id="proj-2", owner_user_id="u1", name="项目二", shared_users=["u2"]) + + @pytest.fixture + def project_other(self): + return Project(id="proj-3", owner_user_id="u3", name="他人项目", shared_users=["u2"]) + + def test_save_and_find_by_id(self, repo, project1): + result = repo.save(project1) + assert result.id == "proj-1" + + found = repo.find_by_id("proj-1") + assert found is not None + assert found.name == "项目一" + + def test_find_by_id_not_found(self, repo): + assert repo.find_by_id("nonexistent") is None + + def test_find_by_owner_user_id(self, repo, project1, project2, project_other): + repo.save(project1) + repo.save(project2) + repo.save(project_other) + + u1_projects = repo.find_by_owner_user_id("u1") + assert len(u1_projects) == 2 + + u3_projects = repo.find_by_owner_user_id("u3") + assert len(u3_projects) == 1 + + def test_find_accessible_projects_owner(self, repo, project1, project_other): + repo.save(project1) + repo.save(project_other) + + # u1 可以访问自己的项目 + accessible = repo.find_accessible_projects("u1") + assert len(accessible) == 1 + assert accessible[0].id == "proj-1" + + def test_find_accessible_projects_shared(self, repo, project2, project_other): + repo.save(project2) + repo.save(project_other) + + # u2 被两个项目共享 + accessible = repo.find_accessible_projects("u2") + assert len(accessible) == 2 + ids = {p.id for p in accessible} + assert ids == {"proj-2", "proj-3"} + + def test_find_accessible_projects_none(self, repo, project1): + repo.save(project1) + assert repo.find_accessible_projects("nobody") == [] + + def test_count_by_owner(self, repo, project1, project2, project_other): + repo.save(project1) + repo.save(project2) + repo.save(project_other) + + assert repo.count_by_owner("u1") == 2 + assert repo.count_by_owner("u3") == 1 + assert repo.count_by_owner("nobody") == 0 + + def test_delete(self, repo, project1): + repo.save(project1) + assert repo.delete("proj-1") is True + assert repo.find_by_id("proj-1") is None + + def test_delete_nonexistent(self, repo): + assert repo.delete("nonexistent") is False + + +# ==================== IngestJob ==================== + + +class TestInMemoryIngestJobRepository: + @pytest.fixture + def repo(self): + return InMemoryIngestJobRepository() + + @pytest.fixture + def job(self): + return IngestJob.create(project_id="p1", library_id="l1", storage_key="key1", file_hash="hash1") + + def test_create_and_get(self, repo, job): + result = repo.create(job) + assert result.id == job.id + assert result.status == IngestJobStatus.PENDING + + fetched = repo.get(job.id) + assert fetched is not None + assert fetched.storage_key == "key1" + + def test_get_nonexistent(self, repo): + assert repo.get("nonexistent") is None + + def test_update(self, repo, job): + repo.create(job) + job.status = IngestJobStatus.PROCESSING + job.error_message = "" + result = repo.update(job) + assert result.status == IngestJobStatus.PROCESSING + + fetched = repo.get(job.id) + assert fetched.status == IngestJobStatus.PROCESSING + + def test_update_with_result(self, repo, job): + repo.create(job) + job.status = IngestJobStatus.COMPLETED + job.result_asset_id = "asset-123" + repo.update(job) + + fetched = repo.get(job.id) + assert fetched.status == IngestJobStatus.COMPLETED + assert fetched.result_asset_id == "asset-123" + + +# ==================== ClassificationJob ==================== + + +class TestInMemoryClassificationJobRepository: + @pytest.fixture + def repo(self): + return InMemoryClassificationJobRepository() + + @pytest.fixture + def job(self): + return ClassificationJob.create(project_id="p1", asset_id="a1") + + def test_create_and_get(self, repo, job): + result = repo.create(job) + assert result.id == job.id + assert result.status == ClassificationJobStatus.PENDING + assert result.confidence == 0.0 + + fetched = repo.get(job.id) + assert fetched is not None + assert fetched.asset_id == "a1" + + def test_get_nonexistent(self, repo): + assert repo.get("nonexistent") is None + + def test_update_status_and_result(self, repo, job): + repo.create(job) + job.status = ClassificationJobStatus.COMPLETED + job.classification = "video" + job.confidence = 0.95 + result = repo.update(job) + + assert result.status == ClassificationJobStatus.COMPLETED + assert result.classification == "video" + assert result.confidence == 0.95 + + def test_update_failed(self, repo, job): + repo.create(job) + job.status = ClassificationJobStatus.FAILED + job.error_message = "something went wrong" + repo.update(job) + + fetched = repo.get(job.id) + assert fetched.status == ClassificationJobStatus.FAILED + assert fetched.error_message == "something went wrong" + + +class TestUserRepositoryUniqueness: + """唯一性约束测试 - 模拟数据库唯一索引冲突.""" + + @pytest.fixture + def repo(self): + from packages.adapters.in_memory.user_repository import InMemoryUserRepository + + return InMemoryUserRepository() + + @pytest.fixture + def user1(self): + return User( + id="user-1", + email="user1@example.com", + display_name="User One", + username="user1", + phone="13800000001", + wechat_openid="wx-openid-1", + wechat_unionid="wx-unionid-1", + created_at=datetime.now(timezone.utc), + ) + + @pytest.fixture + def user2(self): + return User( + id="user-2", + email="user2@example.com", + display_name="User Two", + username="user2", + phone="13800000002", + wechat_openid="wx-openid-2", + wechat_unionid="wx-unionid-2", + created_at=datetime.now(timezone.utc), + ) + + def test_duplicate_email_raises(self, repo, user1, user2): + repo.save(user1) + user2.email = "User1@example.com" # 大小写不同,应视为冲突 + with pytest.raises(ValueError, match="Email already in use"): + repo.save(user2) + + def test_duplicate_username_raises(self, repo, user1, user2): + repo.save(user1) + user2.username = "USER1" # 大小写不同,应视为冲突 + with pytest.raises(ValueError, match="Username already in use"): + repo.save(user2) + + def test_duplicate_phone_raises(self, repo, user1, user2): + repo.save(user1) + user2.phone = "13800000001" + with pytest.raises(ValueError, match="Phone already in use"): + repo.save(user2) + + def test_duplicate_wechat_openid_raises(self, repo, user1, user2): + repo.save(user1) + user2.wechat_openid = "wx-openid-1" + with pytest.raises(ValueError, match="openid already in use"): + repo.save(user2) + + def test_duplicate_wechat_unionid_raises(self, repo, user1, user2): + repo.save(user1) + user2.wechat_unionid = "wx-unionid-1" + with pytest.raises(ValueError, match="unionid already in use"): + repo.save(user2) + + def test_same_user_update_email_ok(self, repo, user1): + """同一用户更新自己的邮箱不视为冲突.""" + repo.save(user1) + user1.email = "newemail@example.com" + repo.save(user1) # 不应抛异常 + + found = repo.find_by_email("newemail@example.com") + assert found is not None + assert found.id == "user-1" + assert repo.find_by_email("user1@example.com") is None + + def test_duplicate_email_fails_cleanly(self, repo, user1, user2): + """唯一性冲突时,用户数据不应被部分写入.""" + repo.save(user1) + user2.email = "user1@example.com" + + with pytest.raises(ValueError): + repo.save(user2) + + # user2 不应该被保存 + assert repo.find_by_id("user-2") is None + # user1 仍然完好 + assert repo.find_by_id("user-1") is not None + assert repo.find_by_email("user1@example.com").id == "user-1" diff --git a/tests/unit/test_inmemory_user_repository.py b/tests/unit/test_inmemory_user_repository.py index f7298c489..6d8eb4153 100755 --- a/tests/unit/test_inmemory_user_repository.py +++ b/tests/unit/test_inmemory_user_repository.py @@ -1,5 +1,6 @@ """InMemoryUserRepository 单元测试.""" +import copy from datetime import datetime, timezone import pytest @@ -173,19 +174,140 @@ class TestDelete: class TestIndexUpdates: - def test_save_new_user_with_same_email_overwrites_index(self, repo, sample_user): - """不同用户同邮箱,后者覆盖索引.""" + def test_save_new_user_with_same_email_raises_uniqueness_error(self, repo, sample_user): + """不同用户同邮箱应触发唯一约束异常.""" repo.save(sample_user) user2 = User( id="user-2", - email="test@example.com", # 同邮箱不同大小写 + email="test@example.com", # 同邮箱 + display_name="User 2", + username="user2", + ) + with pytest.raises(ValueError, match="Email already in use"): + repo.save(user2) + + def test_same_user_update_email_allowed(self, repo, sample_user): + """同一用户更新邮箱不触发唯一约束.""" + repo.save(sample_user) + sample_user.email = "new@example.com" + repo.save(sample_user) + + found = repo.find_by_email("new@example.com") + assert found is not None + assert found.id == "user-1" + assert repo.find_by_email("test@example.com") is None + + def test_update_email_conflict_does_not_corrupt_indexes(self, repo, sample_user): + """更新邮箱与其他用户冲突时,索引必须保持一致,旧邮箱索引不丢失.""" + repo.save(sample_user) + # 第二个用户 + user2 = User( + id="user-2", + email="other@example.com", display_name="User 2", username="user2", ) repo.save(user2) - # 邮箱索引指向最后保存的用户 - found = repo.find_by_email("test@example.com") - assert found.id == "user-2" - # 原用户仍然可通过ID找到 - assert repo.find_by_id("user-1") is not None + # 尝试把 user2 的邮箱改成 sample_user 的邮箱(冲突) + user2_new = copy.deepcopy(user2) + user2_new.email = "test@example.com" + with pytest.raises(ValueError, match="Email already in use"): + repo.save(user2_new) + + # 索引必须保持一致(user-1仍占test@example.com,user-2仍占other@example.com) + assert repo.find_by_email("test@example.com").id == "user-1" + assert repo.find_by_email("other@example.com").id == "user-2" + assert repo.find_by_username("testuser").id == "user-1" + assert repo.find_by_username("user2").id == "user-2" + + def test_update_username_conflict_does_not_corrupt_indexes(self, repo, sample_user): + """更新用户名与他人冲突时,各索引保持一致不丢失.""" + repo.save(sample_user) + user2 = User( + id="user-2", + email="other@example.com", + display_name="User 2", + username="user2", + phone="13900000002", + ) + repo.save(user2) + + # 尝试把 user2 用户名改成 testuser(冲突) + user2_new = copy.deepcopy(user2) + user2_new.username = "testuser" + with pytest.raises(ValueError, match="Username already in use"): + repo.save(user2_new) + + # 各索引必须保持一致 + assert repo.find_by_username("testuser").id == "user-1" + assert repo.find_by_username("user2").id == "user-2" + assert repo.find_by_email("other@example.com").id == "user-2" + assert repo.find_by_phone("13900000002").id == "user-2" + + def test_save_stores_independent_copy(self, repo, sample_user): + """save存储独立副本,外部修改对象不影响仓储内部状态.""" + repo.save(sample_user) + + # 外部修改对象属性 + original_email = sample_user.email + sample_user.email = "hacked@example.com" + sample_user.display_name = "Hacked" + + # 仓储中数据不应受影响 + stored = repo.find_by_id(sample_user.id) + assert stored.email == original_email + assert repo.find_by_email(original_email) is not None + assert repo.find_by_email("hacked@example.com") is None + + def test_save_empty_email_raises(self, repo): + """空邮箱应抛出异常.""" + user = User( + id="user-empty", + email="", + display_name="Empty Email", + username="emptyuser", + ) + with pytest.raises(ValueError, match="email cannot be empty"): + repo.save(user) + + def test_find_by_empty_email_returns_none(self, repo, sample_user): + """空邮箱查询返回None.""" + repo.save(sample_user) + assert repo.find_by_email("") is None + assert repo.find_by_email(None) is None + + def test_find_by_empty_username_returns_none(self, repo, sample_user): + """空用户名查询返回None.""" + repo.save(sample_user) + assert repo.find_by_username("") is None + + def test_duplicate_verification_token_raises(self, repo, sample_user): + """相同邮箱验证令牌应触发唯一约束.""" + sample_user.email_verification_token = "verify-token-abc" + repo.save(sample_user) + + user2 = User( + id="user-2", + email="user2@example.com", + display_name="User 2", + username="user2", + email_verification_token="verify-token-abc", # 重复 + ) + with pytest.raises(ValueError, match="verification token already in use"): + repo.save(user2) + + def test_duplicate_reset_token_raises(self, repo, sample_user): + """相同密码重置令牌应触发唯一约束.""" + sample_user.password_reset_token = "reset-token-xyz" + repo.save(sample_user) + + user2 = User( + id="user-2", + email="user2@example.com", + display_name="User 2", + username="user2", + password_reset_token="reset-token-xyz", # 重复 + ) + with pytest.raises(ValueError, match="reset token already in use"): + repo.save(user2)