refactor(domain): move auth infrastructure to adapters

This commit is contained in:
Xiaoxia AI
2026-06-21 00:46:56 +08:00
parent 3e2b579e03
commit 0e93d2e71f
16 changed files with 750 additions and 679 deletions
+26 -1
View File
@@ -137,11 +137,35 @@
- 测试步骤改为 `python -m pytest tests -q`
- Gitea 与 GitHub workflow 保持一致,避免 CI 双轨漂移。
### 8. P1Domain 层直接依赖 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`
## 四、仍需继续治理的问题
+3
View File
@@ -0,0 +1,3 @@
from packages.adapters.redis.session_store import RedisConfig, SessionStore, get_session_store
__all__ = ["RedisConfig", "SessionStore", "get_session_store"]
+300
View File
@@ -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
+3
View File
@@ -0,0 +1,3 @@
from packages.adapters.smtp.email_service import EmailConfig, EmailService, get_email_service
__all__ = ["EmailConfig", "EmailService", "get_email_service"]
+335
View File
@@ -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"""
<!DOCTYPE html>
<html>
<head>
<meta charset="utf-8">
</head>
<body style="font-family: Arial, sans-serif; line-height: 1.6; color: #333;">
<div style="max-width: 600px; margin: 0 auto; padding: 20px;">
<h2 style="color: #2563eb;">欢迎加入小虾 SaaS</h2>
<p>你好 <strong>{username}</strong></p>
<p>感谢您注册小虾 SaaS!请点击下面的按钮验证您的邮箱地址:</p>
<div style="text-align: center; margin: 30px 0;">
<a href="{verification_url}"
style="background-color: #2563eb; color: white; padding: 12px 30px;
text-decoration: none; border-radius: 5px; display: inline-block;">
验证邮箱
</a>
</div>
<p style="color: #666; font-size: 14px;">
如果按钮无法点击,请复制以下链接到浏览器:<br>
<a href="{verification_url}">{verification_url}</a>
</p>
<p style="color: #666; font-size: 14px;">
此链接将在 24 小时后过期。
</p>
<hr style="border: none; border-top: 1px solid #eee; margin: 30px 0;">
<p style="color: #999; font-size: 12px;">
如果您没有注册小虾 SaaS,请忽略此邮件。
</p>
</div>
</body>
</html>
"""
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"""
<!DOCTYPE html>
<html>
<head>
<meta charset="utf-8">
</head>
<body style="font-family: Arial, sans-serif; line-height: 1.6; color: #333;">
<div style="max-width: 600px; margin: 0 auto; padding: 20px;">
<h2 style="color: #dc2626;">重置密码请求</h2>
<p>你好 <strong>{username}</strong></p>
<p>我们收到了重置您账号密码的请求。请点击下面的按钮重置密码:</p>
<div style="text-align: center; margin: 30px 0;">
<a href="{reset_url}"
style="background-color: #dc2626; color: white; padding: 12px 30px;
text-decoration: none; border-radius: 5px; display: inline-block;">
重置密码
</a>
</div>
<p style="color: #666; font-size: 14px;">
如果按钮无法点击,请复制以下链接到浏览器:<br>
<a href="{reset_url}">{reset_url}</a>
</p>
<p style="color: #666; font-size: 14px;">
此链接将在 1 小时后过期。
</p>
<hr style="border: none; border-top: 1px solid #eee; margin: 30px 0;">
<p style="color: #999; font-size: 12px;">
如果您没有请求重置密码,请忽略此邮件,您的密码不会被更改。
</p>
</div>
</body>
</html>
"""
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"""
<!DOCTYPE html>
<html>
<head>
<meta charset="utf-8">
</head>
<body style="font-family: Arial, sans-serif; line-height: 1.6; color: #333;">
<div style="max-width: 600px; margin: 0 auto; padding: 20px;">
<h2 style="color: #2563eb;">工作空间邀请</h2>
<p><strong>{inviter_name}</strong> 邀请您以 <strong>{role_display}</strong> 身份加入工作空间:</p>
<div style="background: #f3f4f6; padding: 15px; border-radius: 5px; margin: 20px 0;">
<h3 style="margin: 0 0 10px 0; color: #1f2937;">{workspace_name}</h3>
<p style="margin: 0; color: #6b7280;">角色:{role_display}</p>
</div>
<div style="text-align: center; margin: 30px 0;">
<a href="{invitation_url}"
style="background-color: #2563eb; color: white; padding: 12px 30px;
text-decoration: none; border-radius: 5px; display: inline-block;">
接受邀请
</a>
</div>
<p style="color: #666; font-size: 14px;">
如果按钮无法点击,请复制以下链接到浏览器:<br>
<a href="{invitation_url}">{invitation_url}</a>
</p>
<p style="color: #666; font-size: 14px;">
此邀请将在 7 天后过期。
</p>
<hr style="border: none; border-top: 1px solid #eee; margin: 30px 0;">
<p style="color: #999; font-size: 12px;">
如果您不认识邀请人或不想加入此工作空间,请忽略此邮件。
</p>
</div>
</body>
</html>
"""
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
+8 -7
View File
@@ -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:
@@ -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,
@@ -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,
@@ -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,
+10 -6
View File
@@ -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",
]
+7 -326
View File
@@ -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"""
<!DOCTYPE html>
<html>
<head>
<meta charset="utf-8">
</head>
<body style="font-family: Arial, sans-serif; line-height: 1.6; color: #333;">
<div style="max-width: 600px; margin: 0 auto; padding: 20px;">
<h2 style="color: #2563eb;">欢迎加入小虾 SaaS</h2>
<p>你好 <strong>{username}</strong></p>
<p>感谢您注册小虾 SaaS!请点击下面的按钮验证您的邮箱地址:</p>
<div style="text-align: center; margin: 30px 0;">
<a href="{verification_url}"
style="background-color: #2563eb; color: white; padding: 12px 30px;
text-decoration: none; border-radius: 5px; display: inline-block;">
验证邮箱
</a>
</div>
<p style="color: #666; font-size: 14px;">
如果按钮无法点击,请复制以下链接到浏览器:<br>
<a href="{verification_url}">{verification_url}</a>
</p>
<p style="color: #666; font-size: 14px;">
此链接将在 24 小时后过期。
</p>
<hr style="border: none; border-top: 1px solid #eee; margin: 30px 0;">
<p style="color: #999; font-size: 12px;">
如果您没有注册小虾 SaaS,请忽略此邮件。
</p>
</div>
</body>
</html>
"""
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"""
<!DOCTYPE html>
<html>
<head>
<meta charset="utf-8">
</head>
<body style="font-family: Arial, sans-serif; line-height: 1.6; color: #333;">
<div style="max-width: 600px; margin: 0 auto; padding: 20px;">
<h2 style="color: #dc2626;">重置密码请求</h2>
<p>你好 <strong>{username}</strong></p>
<p>我们收到了重置您账号密码的请求。请点击下面的按钮重置密码:</p>
<div style="text-align: center; margin: 30px 0;">
<a href="{reset_url}"
style="background-color: #dc2626; color: white; padding: 12px 30px;
text-decoration: none; border-radius: 5px; display: inline-block;">
重置密码
</a>
</div>
<p style="color: #666; font-size: 14px;">
如果按钮无法点击,请复制以下链接到浏览器:<br>
<a href="{reset_url}">{reset_url}</a>
</p>
<p style="color: #666; font-size: 14px;">
此链接将在 1 小时后过期。
</p>
<hr style="border: none; border-top: 1px solid #eee; margin: 30px 0;">
<p style="color: #999; font-size: 12px;">
如果您没有请求重置密码,请忽略此邮件,您的密码不会被更改。
</p>
</div>
</body>
</html>
"""
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"""
<!DOCTYPE html>
<html>
<head>
<meta charset="utf-8">
</head>
<body style="font-family: Arial, sans-serif; line-height: 1.6; color: #333;">
<div style="max-width: 600px; margin: 0 auto; padding: 20px;">
<h2 style="color: #2563eb;">工作空间邀请</h2>
<p><strong>{inviter_name}</strong> 邀请您以 <strong>{role_display}</strong> 身份加入工作空间:</p>
<div style="background: #f3f4f6; padding: 15px; border-radius: 5px; margin: 20px 0;">
<h3 style="margin: 0 0 10px 0; color: #1f2937;">{workspace_name}</h3>
<p style="margin: 0; color: #6b7280;">角色:{role_display}</p>
</div>
<div style="text-align: center; margin: 30px 0;">
<a href="{invitation_url}"
style="background-color: #2563eb; color: white; padding: 12px 30px;
text-decoration: none; border-radius: 5px; display: inline-block;">
接受邀请
</a>
</div>
<p style="color: #666; font-size: 14px;">
如果按钮无法点击,请复制以下链接到浏览器:<br>
<a href="{invitation_url}">{invitation_url}</a>
</p>
<p style="color: #666; font-size: 14px;">
此邀请将在 7 天后过期。
</p>
<hr style="border: none; border-top: 1px solid #eee; margin: 30px 0;">
<p style="color: #999; font-size: 12px;">
如果您不认识邀请人或不想加入此工作空间,请忽略此邮件。
</p>
</div>
</body>
</html>
"""
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"]
+7 -291
View File
@@ -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"]
+5 -8
View File
@@ -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",
+16 -18
View File
@@ -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",
+8 -8
View File
@@ -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):
"""测试缺少邮箱"""
+8 -8
View File
@@ -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",