Compare commits
10 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 93db1d69e7 | |||
| 206f13bc16 | |||
| f0cf3f8504 | |||
| 9581398d56 | |||
| 0d91b461c7 | |||
| 6a89be1034 | |||
| d7453696b9 | |||
| 540df908f5 | |||
| c2ec722644 | |||
| 92d5b3f26c |
@@ -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]
|
||||
|
||||
Executable
+501
@@ -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"
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user