refactor(domain): move auth infrastructure to adapters
This commit is contained in:
@@ -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`
|
||||
|
||||
## 四、仍需继续治理的问题
|
||||
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
from packages.adapters.redis.session_store import RedisConfig, SessionStore, get_session_store
|
||||
|
||||
__all__ = ["RedisConfig", "SessionStore", "get_session_store"]
|
||||
@@ -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
|
||||
@@ -0,0 +1,3 @@
|
||||
from packages.adapters.smtp.email_service import EmailConfig, EmailService, get_email_service
|
||||
|
||||
__all__ = ["EmailConfig", "EmailService", "get_email_service"]
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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):
|
||||
"""测试缺少邮箱"""
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user