diff --git a/docs/全面代码审计报告-2026-06-21.md b/docs/全面代码审计报告-2026-06-21.md
index fd89a9158..c456f6b57 100644
--- a/docs/全面代码审计报告-2026-06-21.md
+++ b/docs/全面代码审计报告-2026-06-21.md
@@ -137,11 +137,35 @@
- 测试步骤改为 `python -m pytest tests -q`。
- Gitea 与 GitHub workflow 保持一致,避免 CI 双轨漂移。
+### 8. P1:Domain 层直接依赖 Redis/SMTP
+
+**涉及文件**:`packages/domain/auth/session_store.py`、`packages/domain/auth/email_service.py`、`packages/adapters/redis/session_store.py`、`packages/adapters/smtp/email_service.py`
+
+**问题**:
+- Domain 层直接 import Redis、SMTP、email MIME 等基础设施实现。
+- `packages.domain.auth.__init__` 导入时创建 Redis/SMTP 全局单例。
+- Application use case 通过模块全局变量调用外部服务,测试只能 patch 全局变量。
+
+**根因**:
+- 基础设施实现被放进了 domain 包,破坏 Clean Architecture 依赖方向。
+- 外部服务没有通过构造注入进入 use case。
+
+**修复**:
+- Redis session 实现迁移到 `packages/adapters/redis/session_store.py`。
+- SMTP email 实现迁移到 `packages/adapters/smtp/email_service.py`。
+- Domain 原路径保留兼容 shim,但不再导出全局外部服务单例。
+- 注册、登录、登出、密码重置、邀请成员 use case 改为构造注入,默认使用 adapter 懒加载工厂。
+- 相关测试从 patch 模块全局变量改为注入 mock 服务。
+
## 三、已补充测试
### 新增
- `tests/unit/test_auth_simple.py`
+- `tests/unit/test_login_use_case.py`
+- `tests/unit/test_register_user_use_case.py`
+- `tests/unit/test_password_reset_use_case.py`
+- `tests/unit/test_invite_member_use_case.py`
覆盖:
- 登录 token 是可验证 JWT。
@@ -153,9 +177,10 @@
```bash
python -m pytest tests/unit/test_auth_simple.py tests/unit/test_password_hasher.py tests/integration/test_generation_pipeline.py tests/integration/test_projects.py -q
+python -m pytest tests/unit/test_login_use_case.py tests/unit/test_register_user_use_case.py tests/unit/test_password_reset_use_case.py tests/unit/test_invite_member_use_case.py tests/unit/test_session_store.py tests/unit/test_email_service.py -q
```
-结果:`30 passed`
+结果:`30 passed`,`52 passed`
## 四、仍需继续治理的问题
diff --git a/packages/adapters/redis/__init__.py b/packages/adapters/redis/__init__.py
new file mode 100644
index 000000000..1d7a1be95
--- /dev/null
+++ b/packages/adapters/redis/__init__.py
@@ -0,0 +1,3 @@
+from packages.adapters.redis.session_store import RedisConfig, SessionStore, get_session_store
+
+__all__ = ["RedisConfig", "SessionStore", "get_session_store"]
diff --git a/packages/adapters/redis/session_store.py b/packages/adapters/redis/session_store.py
new file mode 100644
index 000000000..574537107
--- /dev/null
+++ b/packages/adapters/redis/session_store.py
@@ -0,0 +1,300 @@
+"""
+Redis Session 存储
+用于存储 refresh_token 和 Session 信息
+"""
+from typing import Optional
+from datetime import datetime, timedelta, timezone
+import json
+import redis
+from redis import Redis
+
+
+class RedisConfig:
+ """Redis 配置"""
+ HOST: str = "localhost"
+ PORT: int = 6379
+ DB: int = 0
+ PASSWORD: Optional[str] = None
+ DECODE_RESPONSES: bool = True
+
+
+class SessionStore:
+ """Session 存储服务"""
+
+ def __init__(self, redis_client: Optional[Redis] = None, config: Optional[RedisConfig] = None):
+ """
+ 初始化 Session 存储
+
+ Args:
+ redis_client: Redis 客户端(可选,用于注入)
+ config: Redis 配置(可选)
+ """
+ if redis_client:
+ self.redis = redis_client
+ else:
+ cfg = config or RedisConfig()
+ self.redis = redis.Redis(
+ host=cfg.HOST,
+ port=cfg.PORT,
+ db=cfg.DB,
+ password=cfg.PASSWORD,
+ decode_responses=cfg.DECODE_RESPONSES,
+ )
+
+ def _session_key(self, session_id: str) -> str:
+ """生成 Session key"""
+ return f"session:{session_id}"
+
+ def _refresh_token_key(self, session_id: str) -> str:
+ """生成 refresh_token key"""
+ return f"refresh_token:{session_id}"
+
+ def _user_sessions_key(self, user_id: str) -> str:
+ """生成用户所有 Session 的 key"""
+ return f"user_sessions:{user_id}"
+
+ def save_session(
+ self,
+ session_id: str,
+ user_id: str,
+ refresh_token: str,
+ device_info: str,
+ ip_address: str,
+ expires_in_seconds: int = 30 * 24 * 60 * 60, # 30 天
+ ) -> bool:
+ """
+ 保存 Session
+
+ Args:
+ session_id: Session ID
+ user_id: 用户 ID
+ refresh_token: 刷新令牌
+ device_info: 设备信息
+ ip_address: IP 地址
+ expires_in_seconds: 过期时间(秒)
+
+ Returns:
+ 是否保存成功
+ """
+ try:
+ now = datetime.now(timezone.utc)
+ expires_at = now + timedelta(seconds=expires_in_seconds)
+
+ session_data = {
+ "session_id": session_id,
+ "user_id": user_id,
+ "device_info": device_info,
+ "ip_address": ip_address,
+ "created_at": now.isoformat(),
+ "last_active_at": now.isoformat(),
+ "expires_at": expires_at.isoformat(),
+ }
+
+ # 保存 Session 数据
+ session_key = self._session_key(session_id)
+ self.redis.setex(
+ session_key,
+ expires_in_seconds,
+ json.dumps(session_data)
+ )
+
+ # 保存 refresh_token 映射
+ refresh_token_key = self._refresh_token_key(session_id)
+ self.redis.setex(
+ refresh_token_key,
+ expires_in_seconds,
+ refresh_token
+ )
+
+ # 添加到用户的 Session 集合
+ user_sessions_key = self._user_sessions_key(user_id)
+ self.redis.sadd(user_sessions_key, session_id)
+ self.redis.expire(user_sessions_key, expires_in_seconds)
+
+ return True
+ except Exception as e:
+ print(f"Failed to save session: {e}")
+ return False
+
+ def get_session(self, session_id: str) -> Optional[dict]:
+ """
+ 获取 Session
+
+ Args:
+ session_id: Session ID
+
+ Returns:
+ Session 数据,如果不存在返回 None
+ """
+ try:
+ session_key = self._session_key(session_id)
+ data = self.redis.get(session_key)
+
+ if data:
+ return json.loads(data)
+ return None
+ except Exception as e:
+ print(f"Failed to get session: {e}")
+ return None
+
+ def get_refresh_token(self, session_id: str) -> Optional[str]:
+ """
+ 获取 refresh_token
+
+ Args:
+ session_id: Session ID
+
+ Returns:
+ refresh_token,如果不存在返回 None
+ """
+ try:
+ refresh_token_key = self._refresh_token_key(session_id)
+ return self.redis.get(refresh_token_key)
+ except Exception as e:
+ print(f"Failed to get refresh_token: {e}")
+ return None
+
+ def update_last_active(self, session_id: str) -> bool:
+ """
+ 更新 Session 最后活跃时间
+
+ Args:
+ session_id: Session ID
+
+ Returns:
+ 是否更新成功
+ """
+ try:
+ session = self.get_session(session_id)
+ if not session:
+ return False
+
+ session["last_active_at"] = datetime.now(timezone.utc).isoformat()
+
+ session_key = self._session_key(session_id)
+ ttl = self.redis.ttl(session_key)
+
+ if ttl > 0:
+ self.redis.setex(
+ session_key,
+ ttl,
+ json.dumps(session)
+ )
+ return True
+
+ return False
+ except Exception as e:
+ print(f"Failed to update last active: {e}")
+ return False
+
+ def delete_session(self, session_id: str) -> bool:
+ """
+ 删除 Session(登出)
+
+ Args:
+ session_id: Session ID
+
+ Returns:
+ 是否删除成功
+ """
+ try:
+ session = self.get_session(session_id)
+ if not session:
+ return False
+
+ user_id = session["user_id"]
+
+ # 删除 Session 数据
+ session_key = self._session_key(session_id)
+ self.redis.delete(session_key)
+
+ # 删除 refresh_token
+ refresh_token_key = self._refresh_token_key(session_id)
+ self.redis.delete(refresh_token_key)
+
+ # 从用户 Session 集合中移除
+ user_sessions_key = self._user_sessions_key(user_id)
+ self.redis.srem(user_sessions_key, session_id)
+
+ return True
+ except Exception as e:
+ print(f"Failed to delete session: {e}")
+ return False
+
+ def get_user_sessions(self, user_id: str) -> list[dict]:
+ """
+ 获取用户的所有活跃 Session
+
+ Args:
+ user_id: 用户 ID
+
+ Returns:
+ Session 列表
+ """
+ try:
+ user_sessions_key = self._user_sessions_key(user_id)
+ session_ids = self.redis.smembers(user_sessions_key)
+
+ sessions = []
+ for session_id in session_ids:
+ session = self.get_session(session_id)
+ if session:
+ sessions.append(session)
+
+ return sessions
+ except Exception as e:
+ print(f"Failed to get user sessions: {e}")
+ return []
+
+ def delete_all_user_sessions(self, user_id: str) -> int:
+ """
+ 删除用户的所有 Session(强制登出所有设备)
+
+ Args:
+ user_id: 用户 ID
+
+ Returns:
+ 删除的 Session 数量
+ """
+ try:
+ sessions = self.get_user_sessions(user_id)
+ count = 0
+
+ for session in sessions:
+ if self.delete_session(session["session_id"]):
+ count += 1
+
+ # 清空用户 Session 集合
+ user_sessions_key = self._user_sessions_key(user_id)
+ self.redis.delete(user_sessions_key)
+
+ return count
+ except Exception as e:
+ print(f"Failed to delete all user sessions: {e}")
+ return 0
+
+ def session_exists(self, session_id: str) -> bool:
+ """
+ 检查 Session 是否存在
+
+ Args:
+ session_id: Session ID
+
+ Returns:
+ 是否存在
+ """
+ try:
+ session_key = self._session_key(session_id)
+ return self.redis.exists(session_key) > 0
+ except Exception:
+ return False
+
+
+_session_store = None
+
+
+def get_session_store() -> SessionStore:
+ global _session_store
+ if _session_store is None:
+ _session_store = SessionStore()
+ return _session_store
diff --git a/packages/adapters/smtp/__init__.py b/packages/adapters/smtp/__init__.py
new file mode 100644
index 000000000..63833654e
--- /dev/null
+++ b/packages/adapters/smtp/__init__.py
@@ -0,0 +1,3 @@
+from packages.adapters.smtp.email_service import EmailConfig, EmailService, get_email_service
+
+__all__ = ["EmailConfig", "EmailService", "get_email_service"]
diff --git a/packages/adapters/smtp/email_service.py b/packages/adapters/smtp/email_service.py
new file mode 100644
index 000000000..c64d75cbc
--- /dev/null
+++ b/packages/adapters/smtp/email_service.py
@@ -0,0 +1,335 @@
+"""
+邮件服务
+支持 SMTP 发送邮件(验证/重置密码/邀请等)
+"""
+import smtplib
+from email.mime.text import MIMEText
+from email.mime.multipart import MIMEMultipart
+from typing import Optional, List
+from dataclasses import dataclass
+
+
+@dataclass
+class EmailConfig:
+ """邮件配置"""
+ smtp_host: str = "smtp.gmail.com"
+ smtp_port: int = 587
+ smtp_user: str = ""
+ smtp_password: str = ""
+ from_email: str = ""
+ from_name: str = "小虾 SaaS"
+ use_tls: bool = True
+
+
+class EmailService:
+ """邮件服务类"""
+
+ def __init__(self, config: Optional[EmailConfig] = None):
+ """
+ 初始化邮件服务
+
+ Args:
+ config: 邮件配置
+ """
+ self.config = config or EmailConfig()
+
+ def send_email(
+ self,
+ to_email: str,
+ subject: str,
+ html_body: str,
+ text_body: Optional[str] = None,
+ cc: Optional[List[str]] = None,
+ bcc: Optional[List[str]] = None,
+ ) -> tuple[bool, Optional[str]]:
+ """
+ 发送邮件
+
+ Args:
+ to_email: 收件人邮箱
+ subject: 邮件主题
+ html_body: HTML 正文
+ text_body: 纯文本正文(可选,作为 HTML 的备用)
+ cc: 抄送列表
+ bcc: 密送列表
+
+ Returns:
+ (是否成功, 错误信息)
+ """
+ try:
+ # 创建邮件
+ msg = MIMEMultipart("alternative")
+ msg["Subject"] = subject
+ msg["From"] = f"{self.config.from_name} <{self.config.from_email}>"
+ msg["To"] = to_email
+
+ if cc:
+ msg["Cc"] = ", ".join(cc)
+
+ # 添加纯文本正文
+ if text_body:
+ part1 = MIMEText(text_body, "plain", "utf-8")
+ msg.attach(part1)
+
+ # 添加 HTML 正文
+ part2 = MIMEText(html_body, "html", "utf-8")
+ msg.attach(part2)
+
+ # 连接 SMTP 服务器
+ with smtplib.SMTP(self.config.smtp_host, self.config.smtp_port) as server:
+ if self.config.use_tls:
+ server.starttls()
+
+ # 登录
+ if self.config.smtp_user and self.config.smtp_password:
+ server.login(self.config.smtp_user, self.config.smtp_password)
+
+ # 发送
+ recipients = [to_email]
+ if cc:
+ recipients.extend(cc)
+ if bcc:
+ recipients.extend(bcc)
+
+ server.sendmail(
+ self.config.from_email,
+ recipients,
+ msg.as_string()
+ )
+
+ return True, None
+
+ except Exception as e:
+ return False, str(e)
+
+ def send_verification_email(
+ self,
+ to_email: str,
+ username: str,
+ verification_url: str,
+ ) -> tuple[bool, Optional[str]]:
+ """
+ 发送邮箱验证邮件
+
+ Args:
+ to_email: 收件人邮箱
+ username: 用户名
+ verification_url: 验证链接
+
+ Returns:
+ (是否成功, 错误信息)
+ """
+ subject = "验证您的邮箱 - 小虾 SaaS"
+
+ html_body = f"""
+
+
+
+
+
+
+
+
欢迎加入小虾 SaaS!
+
你好 {username},
+
感谢您注册小虾 SaaS!请点击下面的按钮验证您的邮箱地址:
+
+
+ 如果按钮无法点击,请复制以下链接到浏览器:
+ {verification_url}
+
+
+ 此链接将在 24 小时后过期。
+
+
+
+ 如果您没有注册小虾 SaaS,请忽略此邮件。
+
+
+
+
+ """
+
+ text_body = f"""
+ 欢迎加入小虾 SaaS!
+
+ 你好 {username},
+
+ 感谢您注册小虾 SaaS!请访问以下链接验证您的邮箱地址:
+
+ {verification_url}
+
+ 此链接将在 24 小时后过期。
+
+ 如果您没有注册小虾 SaaS,请忽略此邮件。
+ """
+
+ return self.send_email(to_email, subject, html_body, text_body)
+
+ def send_password_reset_email(
+ self,
+ to_email: str,
+ username: str,
+ reset_url: str,
+ ) -> tuple[bool, Optional[str]]:
+ """
+ 发送密码重置邮件
+
+ Args:
+ to_email: 收件人邮箱
+ username: 用户名
+ reset_url: 重置链接
+
+ Returns:
+ (是否成功, 错误信息)
+ """
+ subject = "重置您的密码 - 小虾 SaaS"
+
+ html_body = f"""
+
+
+
+
+
+
+
+
重置密码请求
+
你好 {username},
+
我们收到了重置您账号密码的请求。请点击下面的按钮重置密码:
+
+
+ 如果按钮无法点击,请复制以下链接到浏览器:
+ {reset_url}
+
+
+ 此链接将在 1 小时后过期。
+
+
+
+ 如果您没有请求重置密码,请忽略此邮件,您的密码不会被更改。
+
+
+
+
+ """
+
+ text_body = f"""
+ 重置密码请求
+
+ 你好 {username},
+
+ 我们收到了重置您账号密码的请求。请访问以下链接重置密码:
+
+ {reset_url}
+
+ 此链接将在 1 小时后过期。
+
+ 如果您没有请求重置密码,请忽略此邮件,您的密码不会被更改。
+ """
+
+ return self.send_email(to_email, subject, html_body, text_body)
+
+ def send_workspace_invitation_email(
+ self,
+ to_email: str,
+ inviter_name: str,
+ workspace_name: str,
+ role: str,
+ invitation_url: str,
+ ) -> tuple[bool, Optional[str]]:
+ """
+ 发送 Workspace 邀请邮件
+
+ Args:
+ to_email: 收件人邮箱
+ inviter_name: 邀请人姓名
+ workspace_name: 工作空间名称
+ role: 角色(Admin/Member/Viewer)
+ invitation_url: 邀请链接
+
+ Returns:
+ (是否成功, 错误信息)
+ """
+ subject = f"{inviter_name} 邀请您加入 {workspace_name} - 小虾 SaaS"
+
+ role_names = {
+ "owner": "所有者",
+ "admin": "管理员",
+ "member": "成员",
+ "viewer": "查看者",
+ }
+ role_display = role_names.get(role.lower(), role)
+
+ html_body = f"""
+
+
+
+
+
+
+
+
工作空间邀请
+
{inviter_name} 邀请您以 {role_display} 身份加入工作空间:
+
+
{workspace_name}
+
角色:{role_display}
+
+
+
+ 如果按钮无法点击,请复制以下链接到浏览器:
+ {invitation_url}
+
+
+ 此邀请将在 7 天后过期。
+
+
+
+ 如果您不认识邀请人或不想加入此工作空间,请忽略此邮件。
+
+
+
+
+ """
+
+ text_body = f"""
+ 工作空间邀请
+
+ {inviter_name} 邀请您以 {role_display} 身份加入工作空间:{workspace_name}
+
+ 请访问以下链接接受邀请:
+
+ {invitation_url}
+
+ 此邀请将在 7 天后过期。
+
+ 如果您不认识邀请人或不想加入此工作空间,请忽略此邮件。
+ """
+
+ return self.send_email(to_email, subject, html_body, text_body)
+
+
+_email_service = None
+
+
+def get_email_service() -> EmailService:
+ global _email_service
+ if _email_service is None:
+ _email_service = EmailService()
+ return _email_service
diff --git a/packages/application/auth/login_use_case.py b/packages/application/auth/login_use_case.py
index 188d7a8c8..484c4f979 100644
--- a/packages/application/auth/login_use_case.py
+++ b/packages/application/auth/login_use_case.py
@@ -5,10 +5,10 @@ import secrets
from datetime import datetime, timezone, timedelta
from typing import Optional
+from packages.adapters.redis import get_session_store
from packages.domain.auth import (
password_hasher,
jwt_service,
- session_store,
)
@@ -53,8 +53,9 @@ class LoginResponse:
class LoginUseCase:
"""用户登录用例"""
- def __init__(self, user_repository):
+ def __init__(self, user_repository, session_store=None):
self.user_repository = user_repository
+ self.session_store = session_store or get_session_store()
def execute(self, request: LoginRequest) -> tuple[Optional[LoginResponse], Optional[str]]:
"""
@@ -108,7 +109,7 @@ class LoginUseCase:
jwt_service.config.SECRET_KEY,
algorithm=jwt_service.config.ALGORITHM
)
- session_store.save_session(
+ self.session_store.save_session(
session_id=session_id,
user_id=user.id,
refresh_token=refresh_token,
@@ -198,8 +199,8 @@ class LogoutRequest:
class LogoutUseCase:
"""用户登出用例"""
- def __init__(self):
- pass
+ def __init__(self, session_store=None):
+ self.session_store = session_store or get_session_store()
def execute(self, request: LogoutRequest) -> tuple[bool, Optional[str]]:
"""
@@ -214,14 +215,14 @@ class LogoutUseCase:
try:
if request.logout_all_devices:
# 删除所有设备的 session
- count = session_store.delete_all_user_sessions(request.user_id)
+ count = self.session_store.delete_all_user_sessions(request.user_id)
return True, None
else:
# 删除当前 session
if not request.session_id:
return False, "Session ID is required"
- success = session_store.delete_session(request.session_id)
+ success = self.session_store.delete_session(request.session_id)
if success:
return True, None
else:
diff --git a/packages/application/auth/password_reset_use_case.py b/packages/application/auth/password_reset_use_case.py
index 31942d71d..b482e4c8c 100644
--- a/packages/application/auth/password_reset_use_case.py
+++ b/packages/application/auth/password_reset_use_case.py
@@ -5,7 +5,8 @@ import secrets
from datetime import datetime, timedelta, timezone
from typing import Optional
-from packages.domain.auth import password_hasher, password_validator, email_service
+from packages.adapters.smtp import get_email_service
+from packages.domain.auth import password_hasher, password_validator
class RequestPasswordResetRequest:
@@ -23,10 +24,12 @@ class RequestPasswordResetUseCase:
user_repository,
base_url: str = "http://localhost:3000",
token_expire_hours: int = 1,
+ email_service=None,
):
self.user_repository = user_repository
self.base_url = base_url
self.token_expire_hours = token_expire_hours
+ self.email_service = email_service or get_email_service()
def execute(self, request: RequestPasswordResetRequest) -> tuple[bool, Optional[str]]:
"""
@@ -64,7 +67,7 @@ class RequestPasswordResetUseCase:
# 发送重置邮件
try:
- success, error = email_service.send_password_reset_email(
+ success, error = self.email_service.send_password_reset_email(
to_email=user.email,
username=user.username or user.display_name,
reset_url=reset_url,
diff --git a/packages/application/auth/register_user_use_case.py b/packages/application/auth/register_user_use_case.py
index 74a41c640..be43cd276 100644
--- a/packages/application/auth/register_user_use_case.py
+++ b/packages/application/auth/register_user_use_case.py
@@ -6,8 +6,9 @@ from datetime import datetime, timedelta, timezone
from typing import Optional
from uuid import uuid4
+from packages.adapters.smtp import get_email_service
from packages.domain.entities import User
-from packages.domain.auth import password_hasher, password_validator, email_service
+from packages.domain.auth import password_hasher, password_validator
class RegisterUserRequest:
@@ -51,6 +52,7 @@ class RegisterUserUseCase:
self,
user_repository,
base_url: str = "http://localhost:3000",
+ email_service=None,
):
"""
初始化注册用例
@@ -61,6 +63,7 @@ class RegisterUserUseCase:
"""
self.user_repository = user_repository
self.base_url = base_url
+ self.email_service = email_service or get_email_service()
def execute(self, request: RegisterUserRequest) -> tuple[Optional[RegisterUserResponse], Optional[str]]:
"""
@@ -124,7 +127,7 @@ class RegisterUserUseCase:
email_sent = False
try:
- success, error = email_service.send_verification_email(
+ success, error = self.email_service.send_verification_email(
to_email=user.email,
username=user.username,
verification_url=verification_url,
diff --git a/packages/application/workspace/invite_member_use_case.py b/packages/application/workspace/invite_member_use_case.py
index ac64dbafe..c5b9ab160 100644
--- a/packages/application/workspace/invite_member_use_case.py
+++ b/packages/application/workspace/invite_member_use_case.py
@@ -6,12 +6,12 @@ from datetime import datetime, timedelta, timezone
from typing import Optional
from uuid import uuid4
+from packages.adapters.smtp import get_email_service
from packages.domain.entities import (
WorkspaceInvitation,
WorkspaceMemberRole,
InvitationStatus,
)
-from packages.domain.auth import email_service
class InviteMemberRequest:
@@ -63,6 +63,7 @@ class InviteMemberUseCase:
user_repository,
base_url: str = "http://localhost:3000",
invitation_expire_days: int = 7,
+ email_service=None,
):
self.workspace_repository = workspace_repository
self.workspace_member_repository = workspace_member_repository
@@ -70,6 +71,7 @@ class InviteMemberUseCase:
self.user_repository = user_repository
self.base_url = base_url
self.invitation_expire_days = invitation_expire_days
+ self.email_service = email_service or get_email_service()
def execute(self, request: InviteMemberRequest) -> tuple[Optional[InviteMemberResponse], Optional[str]]:
"""
@@ -160,7 +162,7 @@ class InviteMemberUseCase:
inviter = self.user_repository.find_by_id(request.inviter_user_id)
inviter_name = inviter.display_name if inviter else "Someone"
- success, error = email_service.send_workspace_invitation_email(
+ success, error = self.email_service.send_workspace_invitation_email(
to_email=request.invitee_email,
inviter_name=inviter_name,
workspace_name=workspace.name,
diff --git a/packages/domain/auth/__init__.py b/packages/domain/auth/__init__.py
index 0e594885e..12129fe3d 100644
--- a/packages/domain/auth/__init__.py
+++ b/packages/domain/auth/__init__.py
@@ -1,13 +1,19 @@
-"""认证模块"""
-from packages.domain.auth.jwt_service import JWTService, JWTConfig, TokenType, jwt_service
+"""Authentication domain services.
+
+Only pure domain authentication helpers are exported here. Infrastructure-backed
+services such as Redis session storage and SMTP email delivery live under
+`packages.adapters` and should be injected into use cases.
+"""
+
+from packages.domain.auth.jwt_service import JWTConfig, JWTService, TokenType, jwt_service
from packages.domain.auth.password_hasher import (
PasswordHasher,
PasswordValidator,
password_hasher,
password_validator,
)
-from packages.domain.auth.session_store import SessionStore, RedisConfig, session_store
-from packages.domain.auth.email_service import EmailService, EmailConfig, email_service
+from packages.domain.auth.session_store import RedisConfig, SessionStore
+from packages.domain.auth.email_service import EmailConfig, EmailService
__all__ = [
"JWTService",
@@ -20,8 +26,6 @@ __all__ = [
"password_validator",
"SessionStore",
"RedisConfig",
- "session_store",
"EmailService",
"EmailConfig",
- "email_service",
]
diff --git a/packages/domain/auth/email_service.py b/packages/domain/auth/email_service.py
index ae1f5824a..b2e9ad0ae 100644
--- a/packages/domain/auth/email_service.py
+++ b/packages/domain/auth/email_service.py
@@ -1,329 +1,10 @@
+"""Compatibility import for SMTP email delivery.
+
+Infrastructure implementations live under `packages.adapters`. New code should
+import `packages.adapters.smtp.email_service` directly or inject an email-sender
+port into the use case.
"""
-邮件服务
-支持 SMTP 发送邮件(验证/重置密码/邀请等)
-"""
-import smtplib
-from email.mime.text import MIMEText
-from email.mime.multipart import MIMEMultipart
-from typing import Optional, List
-from dataclasses import dataclass
+from packages.adapters.smtp.email_service import EmailConfig, EmailService
-@dataclass
-class EmailConfig:
- """邮件配置"""
- smtp_host: str = "smtp.gmail.com"
- smtp_port: int = 587
- smtp_user: str = ""
- smtp_password: str = ""
- from_email: str = ""
- from_name: str = "小虾 SaaS"
- use_tls: bool = True
-
-
-class EmailService:
- """邮件服务类"""
-
- def __init__(self, config: Optional[EmailConfig] = None):
- """
- 初始化邮件服务
-
- Args:
- config: 邮件配置
- """
- self.config = config or EmailConfig()
-
- def send_email(
- self,
- to_email: str,
- subject: str,
- html_body: str,
- text_body: Optional[str] = None,
- cc: Optional[List[str]] = None,
- bcc: Optional[List[str]] = None,
- ) -> tuple[bool, Optional[str]]:
- """
- 发送邮件
-
- Args:
- to_email: 收件人邮箱
- subject: 邮件主题
- html_body: HTML 正文
- text_body: 纯文本正文(可选,作为 HTML 的备用)
- cc: 抄送列表
- bcc: 密送列表
-
- Returns:
- (是否成功, 错误信息)
- """
- try:
- # 创建邮件
- msg = MIMEMultipart("alternative")
- msg["Subject"] = subject
- msg["From"] = f"{self.config.from_name} <{self.config.from_email}>"
- msg["To"] = to_email
-
- if cc:
- msg["Cc"] = ", ".join(cc)
-
- # 添加纯文本正文
- if text_body:
- part1 = MIMEText(text_body, "plain", "utf-8")
- msg.attach(part1)
-
- # 添加 HTML 正文
- part2 = MIMEText(html_body, "html", "utf-8")
- msg.attach(part2)
-
- # 连接 SMTP 服务器
- with smtplib.SMTP(self.config.smtp_host, self.config.smtp_port) as server:
- if self.config.use_tls:
- server.starttls()
-
- # 登录
- if self.config.smtp_user and self.config.smtp_password:
- server.login(self.config.smtp_user, self.config.smtp_password)
-
- # 发送
- recipients = [to_email]
- if cc:
- recipients.extend(cc)
- if bcc:
- recipients.extend(bcc)
-
- server.sendmail(
- self.config.from_email,
- recipients,
- msg.as_string()
- )
-
- return True, None
-
- except Exception as e:
- return False, str(e)
-
- def send_verification_email(
- self,
- to_email: str,
- username: str,
- verification_url: str,
- ) -> tuple[bool, Optional[str]]:
- """
- 发送邮箱验证邮件
-
- Args:
- to_email: 收件人邮箱
- username: 用户名
- verification_url: 验证链接
-
- Returns:
- (是否成功, 错误信息)
- """
- subject = "验证您的邮箱 - 小虾 SaaS"
-
- html_body = f"""
-
-
-
-
-
-
-
-
欢迎加入小虾 SaaS!
-
你好 {username},
-
感谢您注册小虾 SaaS!请点击下面的按钮验证您的邮箱地址:
-
-
- 如果按钮无法点击,请复制以下链接到浏览器:
- {verification_url}
-
-
- 此链接将在 24 小时后过期。
-
-
-
- 如果您没有注册小虾 SaaS,请忽略此邮件。
-
-
-
-
- """
-
- text_body = f"""
- 欢迎加入小虾 SaaS!
-
- 你好 {username},
-
- 感谢您注册小虾 SaaS!请访问以下链接验证您的邮箱地址:
-
- {verification_url}
-
- 此链接将在 24 小时后过期。
-
- 如果您没有注册小虾 SaaS,请忽略此邮件。
- """
-
- return self.send_email(to_email, subject, html_body, text_body)
-
- def send_password_reset_email(
- self,
- to_email: str,
- username: str,
- reset_url: str,
- ) -> tuple[bool, Optional[str]]:
- """
- 发送密码重置邮件
-
- Args:
- to_email: 收件人邮箱
- username: 用户名
- reset_url: 重置链接
-
- Returns:
- (是否成功, 错误信息)
- """
- subject = "重置您的密码 - 小虾 SaaS"
-
- html_body = f"""
-
-
-
-
-
-
-
-
重置密码请求
-
你好 {username},
-
我们收到了重置您账号密码的请求。请点击下面的按钮重置密码:
-
-
- 如果按钮无法点击,请复制以下链接到浏览器:
- {reset_url}
-
-
- 此链接将在 1 小时后过期。
-
-
-
- 如果您没有请求重置密码,请忽略此邮件,您的密码不会被更改。
-
-
-
-
- """
-
- text_body = f"""
- 重置密码请求
-
- 你好 {username},
-
- 我们收到了重置您账号密码的请求。请访问以下链接重置密码:
-
- {reset_url}
-
- 此链接将在 1 小时后过期。
-
- 如果您没有请求重置密码,请忽略此邮件,您的密码不会被更改。
- """
-
- return self.send_email(to_email, subject, html_body, text_body)
-
- def send_workspace_invitation_email(
- self,
- to_email: str,
- inviter_name: str,
- workspace_name: str,
- role: str,
- invitation_url: str,
- ) -> tuple[bool, Optional[str]]:
- """
- 发送 Workspace 邀请邮件
-
- Args:
- to_email: 收件人邮箱
- inviter_name: 邀请人姓名
- workspace_name: 工作空间名称
- role: 角色(Admin/Member/Viewer)
- invitation_url: 邀请链接
-
- Returns:
- (是否成功, 错误信息)
- """
- subject = f"{inviter_name} 邀请您加入 {workspace_name} - 小虾 SaaS"
-
- role_names = {
- "owner": "所有者",
- "admin": "管理员",
- "member": "成员",
- "viewer": "查看者",
- }
- role_display = role_names.get(role.lower(), role)
-
- html_body = f"""
-
-
-
-
-
-
-
-
工作空间邀请
-
{inviter_name} 邀请您以 {role_display} 身份加入工作空间:
-
-
{workspace_name}
-
角色:{role_display}
-
-
-
- 如果按钮无法点击,请复制以下链接到浏览器:
- {invitation_url}
-
-
- 此邀请将在 7 天后过期。
-
-
-
- 如果您不认识邀请人或不想加入此工作空间,请忽略此邮件。
-
-
-
-
- """
-
- text_body = f"""
- 工作空间邀请
-
- {inviter_name} 邀请您以 {role_display} 身份加入工作空间:{workspace_name}
-
- 请访问以下链接接受邀请:
-
- {invitation_url}
-
- 此邀请将在 7 天后过期。
-
- 如果您不认识邀请人或不想加入此工作空间,请忽略此邮件。
- """
-
- return self.send_email(to_email, subject, html_body, text_body)
-
-
-# 全局实例(生产环境应该从配置读取)
-email_service = EmailService()
+__all__ = ["EmailConfig", "EmailService"]
diff --git a/packages/domain/auth/session_store.py b/packages/domain/auth/session_store.py
index 0f2bc2b53..bb65a83f3 100644
--- a/packages/domain/auth/session_store.py
+++ b/packages/domain/auth/session_store.py
@@ -1,294 +1,10 @@
+"""Compatibility import for Redis-backed session storage.
+
+Infrastructure implementations live under `packages.adapters`. New code should
+import `packages.adapters.redis.session_store` directly or inject a session-store
+port into the use case.
"""
-Redis Session 存储
-用于存储 refresh_token 和 Session 信息
-"""
-from typing import Optional
-from datetime import datetime, timedelta
-import json
-import redis
-from redis import Redis
+from packages.adapters.redis.session_store import RedisConfig, SessionStore
-class RedisConfig:
- """Redis 配置"""
- HOST: str = "localhost"
- PORT: int = 6379
- DB: int = 0
- PASSWORD: Optional[str] = None
- DECODE_RESPONSES: bool = True
-
-
-class SessionStore:
- """Session 存储服务"""
-
- def __init__(self, redis_client: Optional[Redis] = None, config: Optional[RedisConfig] = None):
- """
- 初始化 Session 存储
-
- Args:
- redis_client: Redis 客户端(可选,用于注入)
- config: Redis 配置(可选)
- """
- if redis_client:
- self.redis = redis_client
- else:
- cfg = config or RedisConfig()
- self.redis = redis.Redis(
- host=cfg.HOST,
- port=cfg.PORT,
- db=cfg.DB,
- password=cfg.PASSWORD,
- decode_responses=cfg.DECODE_RESPONSES,
- )
-
- def _session_key(self, session_id: str) -> str:
- """生成 Session key"""
- return f"session:{session_id}"
-
- def _refresh_token_key(self, session_id: str) -> str:
- """生成 refresh_token key"""
- return f"refresh_token:{session_id}"
-
- def _user_sessions_key(self, user_id: str) -> str:
- """生成用户所有 Session 的 key"""
- return f"user_sessions:{user_id}"
-
- def save_session(
- self,
- session_id: str,
- user_id: str,
- refresh_token: str,
- device_info: str,
- ip_address: str,
- expires_in_seconds: int = 30 * 24 * 60 * 60, # 30 天
- ) -> bool:
- """
- 保存 Session
-
- Args:
- session_id: Session ID
- user_id: 用户 ID
- refresh_token: 刷新令牌
- device_info: 设备信息
- ip_address: IP 地址
- expires_in_seconds: 过期时间(秒)
-
- Returns:
- 是否保存成功
- """
- try:
- now = datetime.utcnow()
- expires_at = now + timedelta(seconds=expires_in_seconds)
-
- session_data = {
- "session_id": session_id,
- "user_id": user_id,
- "device_info": device_info,
- "ip_address": ip_address,
- "created_at": now.isoformat(),
- "last_active_at": now.isoformat(),
- "expires_at": expires_at.isoformat(),
- }
-
- # 保存 Session 数据
- session_key = self._session_key(session_id)
- self.redis.setex(
- session_key,
- expires_in_seconds,
- json.dumps(session_data)
- )
-
- # 保存 refresh_token 映射
- refresh_token_key = self._refresh_token_key(session_id)
- self.redis.setex(
- refresh_token_key,
- expires_in_seconds,
- refresh_token
- )
-
- # 添加到用户的 Session 集合
- user_sessions_key = self._user_sessions_key(user_id)
- self.redis.sadd(user_sessions_key, session_id)
- self.redis.expire(user_sessions_key, expires_in_seconds)
-
- return True
- except Exception as e:
- print(f"Failed to save session: {e}")
- return False
-
- def get_session(self, session_id: str) -> Optional[dict]:
- """
- 获取 Session
-
- Args:
- session_id: Session ID
-
- Returns:
- Session 数据,如果不存在返回 None
- """
- try:
- session_key = self._session_key(session_id)
- data = self.redis.get(session_key)
-
- if data:
- return json.loads(data)
- return None
- except Exception as e:
- print(f"Failed to get session: {e}")
- return None
-
- def get_refresh_token(self, session_id: str) -> Optional[str]:
- """
- 获取 refresh_token
-
- Args:
- session_id: Session ID
-
- Returns:
- refresh_token,如果不存在返回 None
- """
- try:
- refresh_token_key = self._refresh_token_key(session_id)
- return self.redis.get(refresh_token_key)
- except Exception as e:
- print(f"Failed to get refresh_token: {e}")
- return None
-
- def update_last_active(self, session_id: str) -> bool:
- """
- 更新 Session 最后活跃时间
-
- Args:
- session_id: Session ID
-
- Returns:
- 是否更新成功
- """
- try:
- session = self.get_session(session_id)
- if not session:
- return False
-
- session["last_active_at"] = datetime.utcnow().isoformat()
-
- session_key = self._session_key(session_id)
- ttl = self.redis.ttl(session_key)
-
- if ttl > 0:
- self.redis.setex(
- session_key,
- ttl,
- json.dumps(session)
- )
- return True
-
- return False
- except Exception as e:
- print(f"Failed to update last active: {e}")
- return False
-
- def delete_session(self, session_id: str) -> bool:
- """
- 删除 Session(登出)
-
- Args:
- session_id: Session ID
-
- Returns:
- 是否删除成功
- """
- try:
- session = self.get_session(session_id)
- if not session:
- return False
-
- user_id = session["user_id"]
-
- # 删除 Session 数据
- session_key = self._session_key(session_id)
- self.redis.delete(session_key)
-
- # 删除 refresh_token
- refresh_token_key = self._refresh_token_key(session_id)
- self.redis.delete(refresh_token_key)
-
- # 从用户 Session 集合中移除
- user_sessions_key = self._user_sessions_key(user_id)
- self.redis.srem(user_sessions_key, session_id)
-
- return True
- except Exception as e:
- print(f"Failed to delete session: {e}")
- return False
-
- def get_user_sessions(self, user_id: str) -> list[dict]:
- """
- 获取用户的所有活跃 Session
-
- Args:
- user_id: 用户 ID
-
- Returns:
- Session 列表
- """
- try:
- user_sessions_key = self._user_sessions_key(user_id)
- session_ids = self.redis.smembers(user_sessions_key)
-
- sessions = []
- for session_id in session_ids:
- session = self.get_session(session_id)
- if session:
- sessions.append(session)
-
- return sessions
- except Exception as e:
- print(f"Failed to get user sessions: {e}")
- return []
-
- def delete_all_user_sessions(self, user_id: str) -> int:
- """
- 删除用户的所有 Session(强制登出所有设备)
-
- Args:
- user_id: 用户 ID
-
- Returns:
- 删除的 Session 数量
- """
- try:
- sessions = self.get_user_sessions(user_id)
- count = 0
-
- for session in sessions:
- if self.delete_session(session["session_id"]):
- count += 1
-
- # 清空用户 Session 集合
- user_sessions_key = self._user_sessions_key(user_id)
- self.redis.delete(user_sessions_key)
-
- return count
- except Exception as e:
- print(f"Failed to delete all user sessions: {e}")
- return 0
-
- def session_exists(self, session_id: str) -> bool:
- """
- 检查 Session 是否存在
-
- Args:
- session_id: Session ID
-
- Returns:
- 是否存在
- """
- try:
- session_key = self._session_key(session_id)
- return self.redis.exists(session_key) > 0
- except Exception:
- return False
-
-
-# 全局实例(生产环境应该从配置读取)
-session_store = SessionStore()
+__all__ = ["RedisConfig", "SessionStore"]
diff --git a/tests/unit/test_invite_member_use_case.py b/tests/unit/test_invite_member_use_case.py
index cca83daef..73ec5ee19 100644
--- a/tests/unit/test_invite_member_use_case.py
+++ b/tests/unit/test_invite_member_use_case.py
@@ -2,7 +2,7 @@
邀请成员 Use Case 测试
"""
import pytest
-from unittest.mock import Mock, patch
+from unittest.mock import Mock
from datetime import datetime, timezone
from packages.application.workspace.invite_member_use_case import (
InviteMemberUseCase,
@@ -47,6 +47,8 @@ class TestInviteMemberUseCase:
@pytest.fixture
def use_case(self, mock_workspace_repo, mock_member_repo, mock_invitation_repo, mock_user_repo):
+ email_service = Mock()
+ email_service.send_workspace_invitation_email.return_value = (True, None)
return InviteMemberUseCase(
workspace_repository=mock_workspace_repo,
workspace_member_repository=mock_member_repo,
@@ -54,6 +56,7 @@ class TestInviteMemberUseCase:
user_repository=mock_user_repo,
base_url="https://test.com",
invitation_expire_days=7,
+ email_service=email_service,
)
@pytest.fixture
@@ -91,10 +94,8 @@ class TestInviteMemberUseCase:
role=WorkspaceMemberRole.ADMIN,
)
- @patch('packages.application.workspace.invite_member_use_case.email_service')
def test_invite_member_success_by_owner(
self,
- mock_email_service,
use_case,
mock_workspace_repo,
mock_member_repo,
@@ -108,7 +109,6 @@ class TestInviteMemberUseCase:
mock_workspace_repo.find_by_id.return_value = test_workspace
mock_member_repo.find_by_workspace_and_user.return_value = owner_member
mock_user_repo.find_by_id.return_value = test_inviter
- mock_email_service.send_workspace_invitation_email.return_value = (True, None)
request = InviteMemberRequest(
workspace_id="workspace-123",
@@ -132,12 +132,10 @@ class TestInviteMemberUseCase:
assert invitation.status == "pending"
# 验证发送了邮件
- mock_email_service.send_workspace_invitation_email.assert_called_once()
+ use_case.email_service.send_workspace_invitation_email.assert_called_once()
- @patch('packages.application.workspace.invite_member_use_case.email_service')
def test_invite_member_success_by_admin(
self,
- mock_email_service,
use_case,
mock_workspace_repo,
mock_member_repo,
@@ -151,7 +149,6 @@ class TestInviteMemberUseCase:
mock_workspace_repo.find_by_id.return_value = test_workspace
mock_member_repo.find_by_workspace_and_user.return_value = admin_member
mock_user_repo.find_by_id.return_value = test_inviter
- mock_email_service.send_workspace_invitation_email.return_value = (True, None)
request = InviteMemberRequest(
workspace_id="workspace-123",
diff --git a/tests/unit/test_login_use_case.py b/tests/unit/test_login_use_case.py
index 21b7e3969..642149ee8 100644
--- a/tests/unit/test_login_use_case.py
+++ b/tests/unit/test_login_use_case.py
@@ -2,7 +2,7 @@
用户登录 Use Case 测试
"""
import pytest
-from unittest.mock import Mock, patch
+from unittest.mock import Mock
from datetime import datetime, timezone
from packages.application.auth import (
LoginUseCase,
@@ -26,7 +26,9 @@ class TestLoginUseCase:
@pytest.fixture
def use_case(self, mock_user_repo):
- return LoginUseCase(user_repository=mock_user_repo)
+ session_store = Mock()
+ session_store.save_session.return_value = True
+ return LoginUseCase(user_repository=mock_user_repo, session_store=session_store)
@pytest.fixture
def test_user(self):
@@ -41,11 +43,9 @@ class TestLoginUseCase:
email_verified=True,
)
- @patch('packages.application.auth.login_use_case.session_store')
- def test_login_success(self, mock_session_store, use_case, mock_user_repo, test_user):
+ def test_login_success(self, use_case, mock_user_repo, test_user):
"""测试登录成功"""
mock_user_repo.find_by_email.return_value = test_user
- mock_session_store.save_session.return_value = True
request = LoginRequest(
email="test@example.com",
@@ -66,7 +66,7 @@ class TestLoginUseCase:
assert response.expires_in > 0
# 验证保存了 session
- mock_session_store.save_session.assert_called_once()
+ use_case.session_store.save_session.assert_called_once()
# 验证更新了最后登录信息
mock_user_repo.save.assert_called_once()
@@ -131,12 +131,12 @@ class TestLogoutUseCase:
@pytest.fixture
def use_case(self):
- return LogoutUseCase()
+ session_store = Mock()
+ return LogoutUseCase(session_store=session_store)
- @patch('packages.application.auth.login_use_case.session_store')
- def test_logout_current_device(self, mock_session_store, use_case):
+ def test_logout_current_device(self, use_case):
"""测试登出当前设备"""
- mock_session_store.delete_session.return_value = True
+ use_case.session_store.delete_session.return_value = True
request = LogoutRequest(
user_id="user-123",
@@ -149,12 +149,11 @@ class TestLogoutUseCase:
assert success is True
assert error is None
- mock_session_store.delete_session.assert_called_once_with("session-abc")
+ use_case.session_store.delete_session.assert_called_once_with("session-abc")
- @patch('packages.application.auth.login_use_case.session_store')
- def test_logout_all_devices(self, mock_session_store, use_case):
+ def test_logout_all_devices(self, use_case):
"""测试登出所有设备"""
- mock_session_store.delete_all_user_sessions.return_value = 3
+ use_case.session_store.delete_all_user_sessions.return_value = 3
request = LogoutRequest(
user_id="user-123",
@@ -166,12 +165,11 @@ class TestLogoutUseCase:
assert success is True
assert error is None
- mock_session_store.delete_all_user_sessions.assert_called_once_with("user-123")
+ use_case.session_store.delete_all_user_sessions.assert_called_once_with("user-123")
- @patch('packages.application.auth.login_use_case.session_store')
- def test_logout_session_not_found(self, mock_session_store, use_case):
+ def test_logout_session_not_found(self, use_case):
"""测试 session 不存在"""
- mock_session_store.delete_session.return_value = False
+ use_case.session_store.delete_session.return_value = False
request = LogoutRequest(
user_id="user-123",
diff --git a/tests/unit/test_password_reset_use_case.py b/tests/unit/test_password_reset_use_case.py
index 426900b43..762ed9413 100644
--- a/tests/unit/test_password_reset_use_case.py
+++ b/tests/unit/test_password_reset_use_case.py
@@ -2,7 +2,7 @@
密码重置 Use Case 测试
"""
import pytest
-from unittest.mock import Mock, patch
+from unittest.mock import Mock
from datetime import datetime, timedelta, timezone
from packages.application.auth.password_reset_use_case import (
RequestPasswordResetUseCase,
@@ -25,10 +25,13 @@ class TestRequestPasswordResetUseCase:
@pytest.fixture
def use_case(self, mock_user_repo):
+ email_service = Mock()
+ email_service.send_password_reset_email.return_value = (True, None)
return RequestPasswordResetUseCase(
user_repository=mock_user_repo,
base_url="https://test.com",
token_expire_hours=1,
+ email_service=email_service,
)
@pytest.fixture
@@ -41,11 +44,9 @@ class TestRequestPasswordResetUseCase:
password_hash="hash",
)
- @patch('packages.application.auth.password_reset_use_case.email_service')
- def test_request_reset_success(self, mock_email_service, use_case, mock_user_repo, test_user):
+ def test_request_reset_success(self, use_case, mock_user_repo, test_user):
"""测试请求重置成功"""
mock_user_repo.find_by_email.return_value = test_user
- mock_email_service.send_password_reset_email.return_value = (True, None)
request = RequestPasswordResetRequest(email="test@example.com")
success, error = use_case.execute(request)
@@ -60,10 +61,9 @@ class TestRequestPasswordResetUseCase:
assert saved_user.password_reset_expires_at is not None
# 验证发送了邮件
- mock_email_service.send_password_reset_email.assert_called_once()
+ use_case.email_service.send_password_reset_email.assert_called_once()
- @patch('packages.application.auth.password_reset_use_case.email_service')
- def test_request_reset_user_not_exists(self, mock_email_service, use_case, mock_user_repo):
+ def test_request_reset_user_not_exists(self, use_case, mock_user_repo):
"""测试用户不存在(仍返回成功,避免暴露)"""
mock_user_repo.find_by_email.return_value = None
@@ -74,7 +74,7 @@ class TestRequestPasswordResetUseCase:
assert error is None
# 不发送邮件
- mock_email_service.send_password_reset_email.assert_not_called()
+ use_case.email_service.send_password_reset_email.assert_not_called()
def test_request_reset_missing_email(self, use_case):
"""测试缺少邮箱"""
diff --git a/tests/unit/test_register_user_use_case.py b/tests/unit/test_register_user_use_case.py
index 049f0ff8c..4cfa8405b 100644
--- a/tests/unit/test_register_user_use_case.py
+++ b/tests/unit/test_register_user_use_case.py
@@ -2,7 +2,7 @@
用户注册 Use Case 测试
"""
import pytest
-from unittest.mock import Mock, patch
+from unittest.mock import Mock
from packages.application.auth import (
RegisterUserUseCase,
RegisterUserRequest,
@@ -28,15 +28,16 @@ class TestRegisterUserUseCase:
@pytest.fixture
def use_case(self, mock_user_repo):
"""创建注册用例"""
+ email_service = Mock()
+ email_service.send_verification_email.return_value = (True, None)
return RegisterUserUseCase(
user_repository=mock_user_repo,
- base_url="https://test.com"
+ base_url="https://test.com",
+ email_service=email_service,
)
- @patch('packages.application.auth.register_user_use_case.email_service')
- def test_register_user_success(self, mock_email_service, use_case, mock_user_repo):
+ def test_register_user_success(self, use_case, mock_user_repo):
"""测试注册成功"""
- mock_email_service.send_verification_email.return_value = (True, None)
request = RegisterUserRequest(
email="test@example.com",
@@ -136,10 +137,9 @@ class TestRegisterUserUseCase:
assert response is None
assert error == "Email is required"
- @patch('packages.application.auth.register_user_use_case.email_service')
- def test_register_user_email_send_failure(self, mock_email_service, use_case, mock_user_repo):
+ def test_register_user_email_send_failure(self, use_case, mock_user_repo):
"""测试邮件发送失败(用户仍然创建)"""
- mock_email_service.send_verification_email.return_value = (False, "SMTP error")
+ use_case.email_service.send_verification_email.return_value = (False, "SMTP error")
request = RegisterUserRequest(
email="test@example.com",