diff --git a/docs/全面代码审计报告-2026-06-21.md b/docs/全面代码审计报告-2026-06-21.md index fd89a9158..c456f6b57 100644 --- a/docs/全面代码审计报告-2026-06-21.md +++ b/docs/全面代码审计报告-2026-06-21.md @@ -137,11 +137,35 @@ - 测试步骤改为 `python -m pytest tests -q`。 - Gitea 与 GitHub workflow 保持一致,避免 CI 双轨漂移。 +### 8. P1:Domain 层直接依赖 Redis/SMTP + +**涉及文件**:`packages/domain/auth/session_store.py`、`packages/domain/auth/email_service.py`、`packages/adapters/redis/session_store.py`、`packages/adapters/smtp/email_service.py` + +**问题**: +- Domain 层直接 import Redis、SMTP、email MIME 等基础设施实现。 +- `packages.domain.auth.__init__` 导入时创建 Redis/SMTP 全局单例。 +- Application use case 通过模块全局变量调用外部服务,测试只能 patch 全局变量。 + +**根因**: +- 基础设施实现被放进了 domain 包,破坏 Clean Architecture 依赖方向。 +- 外部服务没有通过构造注入进入 use case。 + +**修复**: +- Redis session 实现迁移到 `packages/adapters/redis/session_store.py`。 +- SMTP email 实现迁移到 `packages/adapters/smtp/email_service.py`。 +- Domain 原路径保留兼容 shim,但不再导出全局外部服务单例。 +- 注册、登录、登出、密码重置、邀请成员 use case 改为构造注入,默认使用 adapter 懒加载工厂。 +- 相关测试从 patch 模块全局变量改为注入 mock 服务。 + ## 三、已补充测试 ### 新增 - `tests/unit/test_auth_simple.py` +- `tests/unit/test_login_use_case.py` +- `tests/unit/test_register_user_use_case.py` +- `tests/unit/test_password_reset_use_case.py` +- `tests/unit/test_invite_member_use_case.py` 覆盖: - 登录 token 是可验证 JWT。 @@ -153,9 +177,10 @@ ```bash python -m pytest tests/unit/test_auth_simple.py tests/unit/test_password_hasher.py tests/integration/test_generation_pipeline.py tests/integration/test_projects.py -q +python -m pytest tests/unit/test_login_use_case.py tests/unit/test_register_user_use_case.py tests/unit/test_password_reset_use_case.py tests/unit/test_invite_member_use_case.py tests/unit/test_session_store.py tests/unit/test_email_service.py -q ``` -结果:`30 passed` +结果:`30 passed`,`52 passed` ## 四、仍需继续治理的问题 diff --git a/packages/adapters/redis/__init__.py b/packages/adapters/redis/__init__.py new file mode 100644 index 000000000..1d7a1be95 --- /dev/null +++ b/packages/adapters/redis/__init__.py @@ -0,0 +1,3 @@ +from packages.adapters.redis.session_store import RedisConfig, SessionStore, get_session_store + +__all__ = ["RedisConfig", "SessionStore", "get_session_store"] diff --git a/packages/adapters/redis/session_store.py b/packages/adapters/redis/session_store.py new file mode 100644 index 000000000..574537107 --- /dev/null +++ b/packages/adapters/redis/session_store.py @@ -0,0 +1,300 @@ +""" +Redis Session 存储 +用于存储 refresh_token 和 Session 信息 +""" +from typing import Optional +from datetime import datetime, timedelta, timezone +import json +import redis +from redis import Redis + + +class RedisConfig: + """Redis 配置""" + HOST: str = "localhost" + PORT: int = 6379 + DB: int = 0 + PASSWORD: Optional[str] = None + DECODE_RESPONSES: bool = True + + +class SessionStore: + """Session 存储服务""" + + def __init__(self, redis_client: Optional[Redis] = None, config: Optional[RedisConfig] = None): + """ + 初始化 Session 存储 + + Args: + redis_client: Redis 客户端(可选,用于注入) + config: Redis 配置(可选) + """ + if redis_client: + self.redis = redis_client + else: + cfg = config or RedisConfig() + self.redis = redis.Redis( + host=cfg.HOST, + port=cfg.PORT, + db=cfg.DB, + password=cfg.PASSWORD, + decode_responses=cfg.DECODE_RESPONSES, + ) + + def _session_key(self, session_id: str) -> str: + """生成 Session key""" + return f"session:{session_id}" + + def _refresh_token_key(self, session_id: str) -> str: + """生成 refresh_token key""" + return f"refresh_token:{session_id}" + + def _user_sessions_key(self, user_id: str) -> str: + """生成用户所有 Session 的 key""" + return f"user_sessions:{user_id}" + + def save_session( + self, + session_id: str, + user_id: str, + refresh_token: str, + device_info: str, + ip_address: str, + expires_in_seconds: int = 30 * 24 * 60 * 60, # 30 天 + ) -> bool: + """ + 保存 Session + + Args: + session_id: Session ID + user_id: 用户 ID + refresh_token: 刷新令牌 + device_info: 设备信息 + ip_address: IP 地址 + expires_in_seconds: 过期时间(秒) + + Returns: + 是否保存成功 + """ + try: + now = datetime.now(timezone.utc) + expires_at = now + timedelta(seconds=expires_in_seconds) + + session_data = { + "session_id": session_id, + "user_id": user_id, + "device_info": device_info, + "ip_address": ip_address, + "created_at": now.isoformat(), + "last_active_at": now.isoformat(), + "expires_at": expires_at.isoformat(), + } + + # 保存 Session 数据 + session_key = self._session_key(session_id) + self.redis.setex( + session_key, + expires_in_seconds, + json.dumps(session_data) + ) + + # 保存 refresh_token 映射 + refresh_token_key = self._refresh_token_key(session_id) + self.redis.setex( + refresh_token_key, + expires_in_seconds, + refresh_token + ) + + # 添加到用户的 Session 集合 + user_sessions_key = self._user_sessions_key(user_id) + self.redis.sadd(user_sessions_key, session_id) + self.redis.expire(user_sessions_key, expires_in_seconds) + + return True + except Exception as e: + print(f"Failed to save session: {e}") + return False + + def get_session(self, session_id: str) -> Optional[dict]: + """ + 获取 Session + + Args: + session_id: Session ID + + Returns: + Session 数据,如果不存在返回 None + """ + try: + session_key = self._session_key(session_id) + data = self.redis.get(session_key) + + if data: + return json.loads(data) + return None + except Exception as e: + print(f"Failed to get session: {e}") + return None + + def get_refresh_token(self, session_id: str) -> Optional[str]: + """ + 获取 refresh_token + + Args: + session_id: Session ID + + Returns: + refresh_token,如果不存在返回 None + """ + try: + refresh_token_key = self._refresh_token_key(session_id) + return self.redis.get(refresh_token_key) + except Exception as e: + print(f"Failed to get refresh_token: {e}") + return None + + def update_last_active(self, session_id: str) -> bool: + """ + 更新 Session 最后活跃时间 + + Args: + session_id: Session ID + + Returns: + 是否更新成功 + """ + try: + session = self.get_session(session_id) + if not session: + return False + + session["last_active_at"] = datetime.now(timezone.utc).isoformat() + + session_key = self._session_key(session_id) + ttl = self.redis.ttl(session_key) + + if ttl > 0: + self.redis.setex( + session_key, + ttl, + json.dumps(session) + ) + return True + + return False + except Exception as e: + print(f"Failed to update last active: {e}") + return False + + def delete_session(self, session_id: str) -> bool: + """ + 删除 Session(登出) + + Args: + session_id: Session ID + + Returns: + 是否删除成功 + """ + try: + session = self.get_session(session_id) + if not session: + return False + + user_id = session["user_id"] + + # 删除 Session 数据 + session_key = self._session_key(session_id) + self.redis.delete(session_key) + + # 删除 refresh_token + refresh_token_key = self._refresh_token_key(session_id) + self.redis.delete(refresh_token_key) + + # 从用户 Session 集合中移除 + user_sessions_key = self._user_sessions_key(user_id) + self.redis.srem(user_sessions_key, session_id) + + return True + except Exception as e: + print(f"Failed to delete session: {e}") + return False + + def get_user_sessions(self, user_id: str) -> list[dict]: + """ + 获取用户的所有活跃 Session + + Args: + user_id: 用户 ID + + Returns: + Session 列表 + """ + try: + user_sessions_key = self._user_sessions_key(user_id) + session_ids = self.redis.smembers(user_sessions_key) + + sessions = [] + for session_id in session_ids: + session = self.get_session(session_id) + if session: + sessions.append(session) + + return sessions + except Exception as e: + print(f"Failed to get user sessions: {e}") + return [] + + def delete_all_user_sessions(self, user_id: str) -> int: + """ + 删除用户的所有 Session(强制登出所有设备) + + Args: + user_id: 用户 ID + + Returns: + 删除的 Session 数量 + """ + try: + sessions = self.get_user_sessions(user_id) + count = 0 + + for session in sessions: + if self.delete_session(session["session_id"]): + count += 1 + + # 清空用户 Session 集合 + user_sessions_key = self._user_sessions_key(user_id) + self.redis.delete(user_sessions_key) + + return count + except Exception as e: + print(f"Failed to delete all user sessions: {e}") + return 0 + + def session_exists(self, session_id: str) -> bool: + """ + 检查 Session 是否存在 + + Args: + session_id: Session ID + + Returns: + 是否存在 + """ + try: + session_key = self._session_key(session_id) + return self.redis.exists(session_key) > 0 + except Exception: + return False + + +_session_store = None + + +def get_session_store() -> SessionStore: + global _session_store + if _session_store is None: + _session_store = SessionStore() + return _session_store diff --git a/packages/adapters/smtp/__init__.py b/packages/adapters/smtp/__init__.py new file mode 100644 index 000000000..63833654e --- /dev/null +++ b/packages/adapters/smtp/__init__.py @@ -0,0 +1,3 @@ +from packages.adapters.smtp.email_service import EmailConfig, EmailService, get_email_service + +__all__ = ["EmailConfig", "EmailService", "get_email_service"] diff --git a/packages/adapters/smtp/email_service.py b/packages/adapters/smtp/email_service.py new file mode 100644 index 000000000..c64d75cbc --- /dev/null +++ b/packages/adapters/smtp/email_service.py @@ -0,0 +1,335 @@ +""" +邮件服务 +支持 SMTP 发送邮件(验证/重置密码/邀请等) +""" +import smtplib +from email.mime.text import MIMEText +from email.mime.multipart import MIMEMultipart +from typing import Optional, List +from dataclasses import dataclass + + +@dataclass +class EmailConfig: + """邮件配置""" + smtp_host: str = "smtp.gmail.com" + smtp_port: int = 587 + smtp_user: str = "" + smtp_password: str = "" + from_email: str = "" + from_name: str = "小虾 SaaS" + use_tls: bool = True + + +class EmailService: + """邮件服务类""" + + def __init__(self, config: Optional[EmailConfig] = None): + """ + 初始化邮件服务 + + Args: + config: 邮件配置 + """ + self.config = config or EmailConfig() + + def send_email( + self, + to_email: str, + subject: str, + html_body: str, + text_body: Optional[str] = None, + cc: Optional[List[str]] = None, + bcc: Optional[List[str]] = None, + ) -> tuple[bool, Optional[str]]: + """ + 发送邮件 + + Args: + to_email: 收件人邮箱 + subject: 邮件主题 + html_body: HTML 正文 + text_body: 纯文本正文(可选,作为 HTML 的备用) + cc: 抄送列表 + bcc: 密送列表 + + Returns: + (是否成功, 错误信息) + """ + try: + # 创建邮件 + msg = MIMEMultipart("alternative") + msg["Subject"] = subject + msg["From"] = f"{self.config.from_name} <{self.config.from_email}>" + msg["To"] = to_email + + if cc: + msg["Cc"] = ", ".join(cc) + + # 添加纯文本正文 + if text_body: + part1 = MIMEText(text_body, "plain", "utf-8") + msg.attach(part1) + + # 添加 HTML 正文 + part2 = MIMEText(html_body, "html", "utf-8") + msg.attach(part2) + + # 连接 SMTP 服务器 + with smtplib.SMTP(self.config.smtp_host, self.config.smtp_port) as server: + if self.config.use_tls: + server.starttls() + + # 登录 + if self.config.smtp_user and self.config.smtp_password: + server.login(self.config.smtp_user, self.config.smtp_password) + + # 发送 + recipients = [to_email] + if cc: + recipients.extend(cc) + if bcc: + recipients.extend(bcc) + + server.sendmail( + self.config.from_email, + recipients, + msg.as_string() + ) + + return True, None + + except Exception as e: + return False, str(e) + + def send_verification_email( + self, + to_email: str, + username: str, + verification_url: str, + ) -> tuple[bool, Optional[str]]: + """ + 发送邮箱验证邮件 + + Args: + to_email: 收件人邮箱 + username: 用户名 + verification_url: 验证链接 + + Returns: + (是否成功, 错误信息) + """ + subject = "验证您的邮箱 - 小虾 SaaS" + + html_body = f""" + + + + + + +
+

欢迎加入小虾 SaaS!

+

你好 {username}

+

感谢您注册小虾 SaaS!请点击下面的按钮验证您的邮箱地址:

+
+ + 验证邮箱 + +
+

+ 如果按钮无法点击,请复制以下链接到浏览器:
+ {verification_url} +

+

+ 此链接将在 24 小时后过期。 +

+
+

+ 如果您没有注册小虾 SaaS,请忽略此邮件。 +

+
+ + + """ + + text_body = f""" + 欢迎加入小虾 SaaS! + + 你好 {username}, + + 感谢您注册小虾 SaaS!请访问以下链接验证您的邮箱地址: + + {verification_url} + + 此链接将在 24 小时后过期。 + + 如果您没有注册小虾 SaaS,请忽略此邮件。 + """ + + return self.send_email(to_email, subject, html_body, text_body) + + def send_password_reset_email( + self, + to_email: str, + username: str, + reset_url: str, + ) -> tuple[bool, Optional[str]]: + """ + 发送密码重置邮件 + + Args: + to_email: 收件人邮箱 + username: 用户名 + reset_url: 重置链接 + + Returns: + (是否成功, 错误信息) + """ + subject = "重置您的密码 - 小虾 SaaS" + + html_body = f""" + + + + + + +
+

重置密码请求

+

你好 {username}

+

我们收到了重置您账号密码的请求。请点击下面的按钮重置密码:

+
+ + 重置密码 + +
+

+ 如果按钮无法点击,请复制以下链接到浏览器:
+ {reset_url} +

+

+ 此链接将在 1 小时后过期。 +

+
+

+ 如果您没有请求重置密码,请忽略此邮件,您的密码不会被更改。 +

+
+ + + """ + + text_body = f""" + 重置密码请求 + + 你好 {username}, + + 我们收到了重置您账号密码的请求。请访问以下链接重置密码: + + {reset_url} + + 此链接将在 1 小时后过期。 + + 如果您没有请求重置密码,请忽略此邮件,您的密码不会被更改。 + """ + + return self.send_email(to_email, subject, html_body, text_body) + + def send_workspace_invitation_email( + self, + to_email: str, + inviter_name: str, + workspace_name: str, + role: str, + invitation_url: str, + ) -> tuple[bool, Optional[str]]: + """ + 发送 Workspace 邀请邮件 + + Args: + to_email: 收件人邮箱 + inviter_name: 邀请人姓名 + workspace_name: 工作空间名称 + role: 角色(Admin/Member/Viewer) + invitation_url: 邀请链接 + + Returns: + (是否成功, 错误信息) + """ + subject = f"{inviter_name} 邀请您加入 {workspace_name} - 小虾 SaaS" + + role_names = { + "owner": "所有者", + "admin": "管理员", + "member": "成员", + "viewer": "查看者", + } + role_display = role_names.get(role.lower(), role) + + html_body = f""" + + + + + + +
+

工作空间邀请

+

{inviter_name} 邀请您以 {role_display} 身份加入工作空间:

+
+

{workspace_name}

+

角色:{role_display}

+
+
+ + 接受邀请 + +
+

+ 如果按钮无法点击,请复制以下链接到浏览器:
+ {invitation_url} +

+

+ 此邀请将在 7 天后过期。 +

+
+

+ 如果您不认识邀请人或不想加入此工作空间,请忽略此邮件。 +

+
+ + + """ + + text_body = f""" + 工作空间邀请 + + {inviter_name} 邀请您以 {role_display} 身份加入工作空间:{workspace_name} + + 请访问以下链接接受邀请: + + {invitation_url} + + 此邀请将在 7 天后过期。 + + 如果您不认识邀请人或不想加入此工作空间,请忽略此邮件。 + """ + + return self.send_email(to_email, subject, html_body, text_body) + + +_email_service = None + + +def get_email_service() -> EmailService: + global _email_service + if _email_service is None: + _email_service = EmailService() + return _email_service diff --git a/packages/application/auth/login_use_case.py b/packages/application/auth/login_use_case.py index 188d7a8c8..484c4f979 100644 --- a/packages/application/auth/login_use_case.py +++ b/packages/application/auth/login_use_case.py @@ -5,10 +5,10 @@ import secrets from datetime import datetime, timezone, timedelta from typing import Optional +from packages.adapters.redis import get_session_store from packages.domain.auth import ( password_hasher, jwt_service, - session_store, ) @@ -53,8 +53,9 @@ class LoginResponse: class LoginUseCase: """用户登录用例""" - def __init__(self, user_repository): + def __init__(self, user_repository, session_store=None): self.user_repository = user_repository + self.session_store = session_store or get_session_store() def execute(self, request: LoginRequest) -> tuple[Optional[LoginResponse], Optional[str]]: """ @@ -108,7 +109,7 @@ class LoginUseCase: jwt_service.config.SECRET_KEY, algorithm=jwt_service.config.ALGORITHM ) - session_store.save_session( + self.session_store.save_session( session_id=session_id, user_id=user.id, refresh_token=refresh_token, @@ -198,8 +199,8 @@ class LogoutRequest: class LogoutUseCase: """用户登出用例""" - def __init__(self): - pass + def __init__(self, session_store=None): + self.session_store = session_store or get_session_store() def execute(self, request: LogoutRequest) -> tuple[bool, Optional[str]]: """ @@ -214,14 +215,14 @@ class LogoutUseCase: try: if request.logout_all_devices: # 删除所有设备的 session - count = session_store.delete_all_user_sessions(request.user_id) + count = self.session_store.delete_all_user_sessions(request.user_id) return True, None else: # 删除当前 session if not request.session_id: return False, "Session ID is required" - success = session_store.delete_session(request.session_id) + success = self.session_store.delete_session(request.session_id) if success: return True, None else: diff --git a/packages/application/auth/password_reset_use_case.py b/packages/application/auth/password_reset_use_case.py index 31942d71d..b482e4c8c 100644 --- a/packages/application/auth/password_reset_use_case.py +++ b/packages/application/auth/password_reset_use_case.py @@ -5,7 +5,8 @@ import secrets from datetime import datetime, timedelta, timezone from typing import Optional -from packages.domain.auth import password_hasher, password_validator, email_service +from packages.adapters.smtp import get_email_service +from packages.domain.auth import password_hasher, password_validator class RequestPasswordResetRequest: @@ -23,10 +24,12 @@ class RequestPasswordResetUseCase: user_repository, base_url: str = "http://localhost:3000", token_expire_hours: int = 1, + email_service=None, ): self.user_repository = user_repository self.base_url = base_url self.token_expire_hours = token_expire_hours + self.email_service = email_service or get_email_service() def execute(self, request: RequestPasswordResetRequest) -> tuple[bool, Optional[str]]: """ @@ -64,7 +67,7 @@ class RequestPasswordResetUseCase: # 发送重置邮件 try: - success, error = email_service.send_password_reset_email( + success, error = self.email_service.send_password_reset_email( to_email=user.email, username=user.username or user.display_name, reset_url=reset_url, diff --git a/packages/application/auth/register_user_use_case.py b/packages/application/auth/register_user_use_case.py index 74a41c640..be43cd276 100644 --- a/packages/application/auth/register_user_use_case.py +++ b/packages/application/auth/register_user_use_case.py @@ -6,8 +6,9 @@ from datetime import datetime, timedelta, timezone from typing import Optional from uuid import uuid4 +from packages.adapters.smtp import get_email_service from packages.domain.entities import User -from packages.domain.auth import password_hasher, password_validator, email_service +from packages.domain.auth import password_hasher, password_validator class RegisterUserRequest: @@ -51,6 +52,7 @@ class RegisterUserUseCase: self, user_repository, base_url: str = "http://localhost:3000", + email_service=None, ): """ 初始化注册用例 @@ -61,6 +63,7 @@ class RegisterUserUseCase: """ self.user_repository = user_repository self.base_url = base_url + self.email_service = email_service or get_email_service() def execute(self, request: RegisterUserRequest) -> tuple[Optional[RegisterUserResponse], Optional[str]]: """ @@ -124,7 +127,7 @@ class RegisterUserUseCase: email_sent = False try: - success, error = email_service.send_verification_email( + success, error = self.email_service.send_verification_email( to_email=user.email, username=user.username, verification_url=verification_url, diff --git a/packages/application/workspace/invite_member_use_case.py b/packages/application/workspace/invite_member_use_case.py index ac64dbafe..c5b9ab160 100644 --- a/packages/application/workspace/invite_member_use_case.py +++ b/packages/application/workspace/invite_member_use_case.py @@ -6,12 +6,12 @@ from datetime import datetime, timedelta, timezone from typing import Optional from uuid import uuid4 +from packages.adapters.smtp import get_email_service from packages.domain.entities import ( WorkspaceInvitation, WorkspaceMemberRole, InvitationStatus, ) -from packages.domain.auth import email_service class InviteMemberRequest: @@ -63,6 +63,7 @@ class InviteMemberUseCase: user_repository, base_url: str = "http://localhost:3000", invitation_expire_days: int = 7, + email_service=None, ): self.workspace_repository = workspace_repository self.workspace_member_repository = workspace_member_repository @@ -70,6 +71,7 @@ class InviteMemberUseCase: self.user_repository = user_repository self.base_url = base_url self.invitation_expire_days = invitation_expire_days + self.email_service = email_service or get_email_service() def execute(self, request: InviteMemberRequest) -> tuple[Optional[InviteMemberResponse], Optional[str]]: """ @@ -160,7 +162,7 @@ class InviteMemberUseCase: inviter = self.user_repository.find_by_id(request.inviter_user_id) inviter_name = inviter.display_name if inviter else "Someone" - success, error = email_service.send_workspace_invitation_email( + success, error = self.email_service.send_workspace_invitation_email( to_email=request.invitee_email, inviter_name=inviter_name, workspace_name=workspace.name, diff --git a/packages/domain/auth/__init__.py b/packages/domain/auth/__init__.py index 0e594885e..12129fe3d 100644 --- a/packages/domain/auth/__init__.py +++ b/packages/domain/auth/__init__.py @@ -1,13 +1,19 @@ -"""认证模块""" -from packages.domain.auth.jwt_service import JWTService, JWTConfig, TokenType, jwt_service +"""Authentication domain services. + +Only pure domain authentication helpers are exported here. Infrastructure-backed +services such as Redis session storage and SMTP email delivery live under +`packages.adapters` and should be injected into use cases. +""" + +from packages.domain.auth.jwt_service import JWTConfig, JWTService, TokenType, jwt_service from packages.domain.auth.password_hasher import ( PasswordHasher, PasswordValidator, password_hasher, password_validator, ) -from packages.domain.auth.session_store import SessionStore, RedisConfig, session_store -from packages.domain.auth.email_service import EmailService, EmailConfig, email_service +from packages.domain.auth.session_store import RedisConfig, SessionStore +from packages.domain.auth.email_service import EmailConfig, EmailService __all__ = [ "JWTService", @@ -20,8 +26,6 @@ __all__ = [ "password_validator", "SessionStore", "RedisConfig", - "session_store", "EmailService", "EmailConfig", - "email_service", ] diff --git a/packages/domain/auth/email_service.py b/packages/domain/auth/email_service.py index ae1f5824a..b2e9ad0ae 100644 --- a/packages/domain/auth/email_service.py +++ b/packages/domain/auth/email_service.py @@ -1,329 +1,10 @@ +"""Compatibility import for SMTP email delivery. + +Infrastructure implementations live under `packages.adapters`. New code should +import `packages.adapters.smtp.email_service` directly or inject an email-sender +port into the use case. """ -邮件服务 -支持 SMTP 发送邮件(验证/重置密码/邀请等) -""" -import smtplib -from email.mime.text import MIMEText -from email.mime.multipart import MIMEMultipart -from typing import Optional, List -from dataclasses import dataclass +from packages.adapters.smtp.email_service import EmailConfig, EmailService -@dataclass -class EmailConfig: - """邮件配置""" - smtp_host: str = "smtp.gmail.com" - smtp_port: int = 587 - smtp_user: str = "" - smtp_password: str = "" - from_email: str = "" - from_name: str = "小虾 SaaS" - use_tls: bool = True - - -class EmailService: - """邮件服务类""" - - def __init__(self, config: Optional[EmailConfig] = None): - """ - 初始化邮件服务 - - Args: - config: 邮件配置 - """ - self.config = config or EmailConfig() - - def send_email( - self, - to_email: str, - subject: str, - html_body: str, - text_body: Optional[str] = None, - cc: Optional[List[str]] = None, - bcc: Optional[List[str]] = None, - ) -> tuple[bool, Optional[str]]: - """ - 发送邮件 - - Args: - to_email: 收件人邮箱 - subject: 邮件主题 - html_body: HTML 正文 - text_body: 纯文本正文(可选,作为 HTML 的备用) - cc: 抄送列表 - bcc: 密送列表 - - Returns: - (是否成功, 错误信息) - """ - try: - # 创建邮件 - msg = MIMEMultipart("alternative") - msg["Subject"] = subject - msg["From"] = f"{self.config.from_name} <{self.config.from_email}>" - msg["To"] = to_email - - if cc: - msg["Cc"] = ", ".join(cc) - - # 添加纯文本正文 - if text_body: - part1 = MIMEText(text_body, "plain", "utf-8") - msg.attach(part1) - - # 添加 HTML 正文 - part2 = MIMEText(html_body, "html", "utf-8") - msg.attach(part2) - - # 连接 SMTP 服务器 - with smtplib.SMTP(self.config.smtp_host, self.config.smtp_port) as server: - if self.config.use_tls: - server.starttls() - - # 登录 - if self.config.smtp_user and self.config.smtp_password: - server.login(self.config.smtp_user, self.config.smtp_password) - - # 发送 - recipients = [to_email] - if cc: - recipients.extend(cc) - if bcc: - recipients.extend(bcc) - - server.sendmail( - self.config.from_email, - recipients, - msg.as_string() - ) - - return True, None - - except Exception as e: - return False, str(e) - - def send_verification_email( - self, - to_email: str, - username: str, - verification_url: str, - ) -> tuple[bool, Optional[str]]: - """ - 发送邮箱验证邮件 - - Args: - to_email: 收件人邮箱 - username: 用户名 - verification_url: 验证链接 - - Returns: - (是否成功, 错误信息) - """ - subject = "验证您的邮箱 - 小虾 SaaS" - - html_body = f""" - - - - - - -
-

欢迎加入小虾 SaaS!

-

你好 {username}

-

感谢您注册小虾 SaaS!请点击下面的按钮验证您的邮箱地址:

-
- - 验证邮箱 - -
-

- 如果按钮无法点击,请复制以下链接到浏览器:
- {verification_url} -

-

- 此链接将在 24 小时后过期。 -

-
-

- 如果您没有注册小虾 SaaS,请忽略此邮件。 -

-
- - - """ - - text_body = f""" - 欢迎加入小虾 SaaS! - - 你好 {username}, - - 感谢您注册小虾 SaaS!请访问以下链接验证您的邮箱地址: - - {verification_url} - - 此链接将在 24 小时后过期。 - - 如果您没有注册小虾 SaaS,请忽略此邮件。 - """ - - return self.send_email(to_email, subject, html_body, text_body) - - def send_password_reset_email( - self, - to_email: str, - username: str, - reset_url: str, - ) -> tuple[bool, Optional[str]]: - """ - 发送密码重置邮件 - - Args: - to_email: 收件人邮箱 - username: 用户名 - reset_url: 重置链接 - - Returns: - (是否成功, 错误信息) - """ - subject = "重置您的密码 - 小虾 SaaS" - - html_body = f""" - - - - - - -
-

重置密码请求

-

你好 {username}

-

我们收到了重置您账号密码的请求。请点击下面的按钮重置密码:

-
- - 重置密码 - -
-

- 如果按钮无法点击,请复制以下链接到浏览器:
- {reset_url} -

-

- 此链接将在 1 小时后过期。 -

-
-

- 如果您没有请求重置密码,请忽略此邮件,您的密码不会被更改。 -

-
- - - """ - - text_body = f""" - 重置密码请求 - - 你好 {username}, - - 我们收到了重置您账号密码的请求。请访问以下链接重置密码: - - {reset_url} - - 此链接将在 1 小时后过期。 - - 如果您没有请求重置密码,请忽略此邮件,您的密码不会被更改。 - """ - - return self.send_email(to_email, subject, html_body, text_body) - - def send_workspace_invitation_email( - self, - to_email: str, - inviter_name: str, - workspace_name: str, - role: str, - invitation_url: str, - ) -> tuple[bool, Optional[str]]: - """ - 发送 Workspace 邀请邮件 - - Args: - to_email: 收件人邮箱 - inviter_name: 邀请人姓名 - workspace_name: 工作空间名称 - role: 角色(Admin/Member/Viewer) - invitation_url: 邀请链接 - - Returns: - (是否成功, 错误信息) - """ - subject = f"{inviter_name} 邀请您加入 {workspace_name} - 小虾 SaaS" - - role_names = { - "owner": "所有者", - "admin": "管理员", - "member": "成员", - "viewer": "查看者", - } - role_display = role_names.get(role.lower(), role) - - html_body = f""" - - - - - - -
-

工作空间邀请

-

{inviter_name} 邀请您以 {role_display} 身份加入工作空间:

-
-

{workspace_name}

-

角色:{role_display}

-
-
- - 接受邀请 - -
-

- 如果按钮无法点击,请复制以下链接到浏览器:
- {invitation_url} -

-

- 此邀请将在 7 天后过期。 -

-
-

- 如果您不认识邀请人或不想加入此工作空间,请忽略此邮件。 -

-
- - - """ - - text_body = f""" - 工作空间邀请 - - {inviter_name} 邀请您以 {role_display} 身份加入工作空间:{workspace_name} - - 请访问以下链接接受邀请: - - {invitation_url} - - 此邀请将在 7 天后过期。 - - 如果您不认识邀请人或不想加入此工作空间,请忽略此邮件。 - """ - - return self.send_email(to_email, subject, html_body, text_body) - - -# 全局实例(生产环境应该从配置读取) -email_service = EmailService() +__all__ = ["EmailConfig", "EmailService"] diff --git a/packages/domain/auth/session_store.py b/packages/domain/auth/session_store.py index 0f2bc2b53..bb65a83f3 100644 --- a/packages/domain/auth/session_store.py +++ b/packages/domain/auth/session_store.py @@ -1,294 +1,10 @@ +"""Compatibility import for Redis-backed session storage. + +Infrastructure implementations live under `packages.adapters`. New code should +import `packages.adapters.redis.session_store` directly or inject a session-store +port into the use case. """ -Redis Session 存储 -用于存储 refresh_token 和 Session 信息 -""" -from typing import Optional -from datetime import datetime, timedelta -import json -import redis -from redis import Redis +from packages.adapters.redis.session_store import RedisConfig, SessionStore -class RedisConfig: - """Redis 配置""" - HOST: str = "localhost" - PORT: int = 6379 - DB: int = 0 - PASSWORD: Optional[str] = None - DECODE_RESPONSES: bool = True - - -class SessionStore: - """Session 存储服务""" - - def __init__(self, redis_client: Optional[Redis] = None, config: Optional[RedisConfig] = None): - """ - 初始化 Session 存储 - - Args: - redis_client: Redis 客户端(可选,用于注入) - config: Redis 配置(可选) - """ - if redis_client: - self.redis = redis_client - else: - cfg = config or RedisConfig() - self.redis = redis.Redis( - host=cfg.HOST, - port=cfg.PORT, - db=cfg.DB, - password=cfg.PASSWORD, - decode_responses=cfg.DECODE_RESPONSES, - ) - - def _session_key(self, session_id: str) -> str: - """生成 Session key""" - return f"session:{session_id}" - - def _refresh_token_key(self, session_id: str) -> str: - """生成 refresh_token key""" - return f"refresh_token:{session_id}" - - def _user_sessions_key(self, user_id: str) -> str: - """生成用户所有 Session 的 key""" - return f"user_sessions:{user_id}" - - def save_session( - self, - session_id: str, - user_id: str, - refresh_token: str, - device_info: str, - ip_address: str, - expires_in_seconds: int = 30 * 24 * 60 * 60, # 30 天 - ) -> bool: - """ - 保存 Session - - Args: - session_id: Session ID - user_id: 用户 ID - refresh_token: 刷新令牌 - device_info: 设备信息 - ip_address: IP 地址 - expires_in_seconds: 过期时间(秒) - - Returns: - 是否保存成功 - """ - try: - now = datetime.utcnow() - expires_at = now + timedelta(seconds=expires_in_seconds) - - session_data = { - "session_id": session_id, - "user_id": user_id, - "device_info": device_info, - "ip_address": ip_address, - "created_at": now.isoformat(), - "last_active_at": now.isoformat(), - "expires_at": expires_at.isoformat(), - } - - # 保存 Session 数据 - session_key = self._session_key(session_id) - self.redis.setex( - session_key, - expires_in_seconds, - json.dumps(session_data) - ) - - # 保存 refresh_token 映射 - refresh_token_key = self._refresh_token_key(session_id) - self.redis.setex( - refresh_token_key, - expires_in_seconds, - refresh_token - ) - - # 添加到用户的 Session 集合 - user_sessions_key = self._user_sessions_key(user_id) - self.redis.sadd(user_sessions_key, session_id) - self.redis.expire(user_sessions_key, expires_in_seconds) - - return True - except Exception as e: - print(f"Failed to save session: {e}") - return False - - def get_session(self, session_id: str) -> Optional[dict]: - """ - 获取 Session - - Args: - session_id: Session ID - - Returns: - Session 数据,如果不存在返回 None - """ - try: - session_key = self._session_key(session_id) - data = self.redis.get(session_key) - - if data: - return json.loads(data) - return None - except Exception as e: - print(f"Failed to get session: {e}") - return None - - def get_refresh_token(self, session_id: str) -> Optional[str]: - """ - 获取 refresh_token - - Args: - session_id: Session ID - - Returns: - refresh_token,如果不存在返回 None - """ - try: - refresh_token_key = self._refresh_token_key(session_id) - return self.redis.get(refresh_token_key) - except Exception as e: - print(f"Failed to get refresh_token: {e}") - return None - - def update_last_active(self, session_id: str) -> bool: - """ - 更新 Session 最后活跃时间 - - Args: - session_id: Session ID - - Returns: - 是否更新成功 - """ - try: - session = self.get_session(session_id) - if not session: - return False - - session["last_active_at"] = datetime.utcnow().isoformat() - - session_key = self._session_key(session_id) - ttl = self.redis.ttl(session_key) - - if ttl > 0: - self.redis.setex( - session_key, - ttl, - json.dumps(session) - ) - return True - - return False - except Exception as e: - print(f"Failed to update last active: {e}") - return False - - def delete_session(self, session_id: str) -> bool: - """ - 删除 Session(登出) - - Args: - session_id: Session ID - - Returns: - 是否删除成功 - """ - try: - session = self.get_session(session_id) - if not session: - return False - - user_id = session["user_id"] - - # 删除 Session 数据 - session_key = self._session_key(session_id) - self.redis.delete(session_key) - - # 删除 refresh_token - refresh_token_key = self._refresh_token_key(session_id) - self.redis.delete(refresh_token_key) - - # 从用户 Session 集合中移除 - user_sessions_key = self._user_sessions_key(user_id) - self.redis.srem(user_sessions_key, session_id) - - return True - except Exception as e: - print(f"Failed to delete session: {e}") - return False - - def get_user_sessions(self, user_id: str) -> list[dict]: - """ - 获取用户的所有活跃 Session - - Args: - user_id: 用户 ID - - Returns: - Session 列表 - """ - try: - user_sessions_key = self._user_sessions_key(user_id) - session_ids = self.redis.smembers(user_sessions_key) - - sessions = [] - for session_id in session_ids: - session = self.get_session(session_id) - if session: - sessions.append(session) - - return sessions - except Exception as e: - print(f"Failed to get user sessions: {e}") - return [] - - def delete_all_user_sessions(self, user_id: str) -> int: - """ - 删除用户的所有 Session(强制登出所有设备) - - Args: - user_id: 用户 ID - - Returns: - 删除的 Session 数量 - """ - try: - sessions = self.get_user_sessions(user_id) - count = 0 - - for session in sessions: - if self.delete_session(session["session_id"]): - count += 1 - - # 清空用户 Session 集合 - user_sessions_key = self._user_sessions_key(user_id) - self.redis.delete(user_sessions_key) - - return count - except Exception as e: - print(f"Failed to delete all user sessions: {e}") - return 0 - - def session_exists(self, session_id: str) -> bool: - """ - 检查 Session 是否存在 - - Args: - session_id: Session ID - - Returns: - 是否存在 - """ - try: - session_key = self._session_key(session_id) - return self.redis.exists(session_key) > 0 - except Exception: - return False - - -# 全局实例(生产环境应该从配置读取) -session_store = SessionStore() +__all__ = ["RedisConfig", "SessionStore"] diff --git a/tests/unit/test_invite_member_use_case.py b/tests/unit/test_invite_member_use_case.py index cca83daef..73ec5ee19 100644 --- a/tests/unit/test_invite_member_use_case.py +++ b/tests/unit/test_invite_member_use_case.py @@ -2,7 +2,7 @@ 邀请成员 Use Case 测试 """ import pytest -from unittest.mock import Mock, patch +from unittest.mock import Mock from datetime import datetime, timezone from packages.application.workspace.invite_member_use_case import ( InviteMemberUseCase, @@ -47,6 +47,8 @@ class TestInviteMemberUseCase: @pytest.fixture def use_case(self, mock_workspace_repo, mock_member_repo, mock_invitation_repo, mock_user_repo): + email_service = Mock() + email_service.send_workspace_invitation_email.return_value = (True, None) return InviteMemberUseCase( workspace_repository=mock_workspace_repo, workspace_member_repository=mock_member_repo, @@ -54,6 +56,7 @@ class TestInviteMemberUseCase: user_repository=mock_user_repo, base_url="https://test.com", invitation_expire_days=7, + email_service=email_service, ) @pytest.fixture @@ -91,10 +94,8 @@ class TestInviteMemberUseCase: role=WorkspaceMemberRole.ADMIN, ) - @patch('packages.application.workspace.invite_member_use_case.email_service') def test_invite_member_success_by_owner( self, - mock_email_service, use_case, mock_workspace_repo, mock_member_repo, @@ -108,7 +109,6 @@ class TestInviteMemberUseCase: mock_workspace_repo.find_by_id.return_value = test_workspace mock_member_repo.find_by_workspace_and_user.return_value = owner_member mock_user_repo.find_by_id.return_value = test_inviter - mock_email_service.send_workspace_invitation_email.return_value = (True, None) request = InviteMemberRequest( workspace_id="workspace-123", @@ -132,12 +132,10 @@ class TestInviteMemberUseCase: assert invitation.status == "pending" # 验证发送了邮件 - mock_email_service.send_workspace_invitation_email.assert_called_once() + use_case.email_service.send_workspace_invitation_email.assert_called_once() - @patch('packages.application.workspace.invite_member_use_case.email_service') def test_invite_member_success_by_admin( self, - mock_email_service, use_case, mock_workspace_repo, mock_member_repo, @@ -151,7 +149,6 @@ class TestInviteMemberUseCase: mock_workspace_repo.find_by_id.return_value = test_workspace mock_member_repo.find_by_workspace_and_user.return_value = admin_member mock_user_repo.find_by_id.return_value = test_inviter - mock_email_service.send_workspace_invitation_email.return_value = (True, None) request = InviteMemberRequest( workspace_id="workspace-123", diff --git a/tests/unit/test_login_use_case.py b/tests/unit/test_login_use_case.py index 21b7e3969..642149ee8 100644 --- a/tests/unit/test_login_use_case.py +++ b/tests/unit/test_login_use_case.py @@ -2,7 +2,7 @@ 用户登录 Use Case 测试 """ import pytest -from unittest.mock import Mock, patch +from unittest.mock import Mock from datetime import datetime, timezone from packages.application.auth import ( LoginUseCase, @@ -26,7 +26,9 @@ class TestLoginUseCase: @pytest.fixture def use_case(self, mock_user_repo): - return LoginUseCase(user_repository=mock_user_repo) + session_store = Mock() + session_store.save_session.return_value = True + return LoginUseCase(user_repository=mock_user_repo, session_store=session_store) @pytest.fixture def test_user(self): @@ -41,11 +43,9 @@ class TestLoginUseCase: email_verified=True, ) - @patch('packages.application.auth.login_use_case.session_store') - def test_login_success(self, mock_session_store, use_case, mock_user_repo, test_user): + def test_login_success(self, use_case, mock_user_repo, test_user): """测试登录成功""" mock_user_repo.find_by_email.return_value = test_user - mock_session_store.save_session.return_value = True request = LoginRequest( email="test@example.com", @@ -66,7 +66,7 @@ class TestLoginUseCase: assert response.expires_in > 0 # 验证保存了 session - mock_session_store.save_session.assert_called_once() + use_case.session_store.save_session.assert_called_once() # 验证更新了最后登录信息 mock_user_repo.save.assert_called_once() @@ -131,12 +131,12 @@ class TestLogoutUseCase: @pytest.fixture def use_case(self): - return LogoutUseCase() + session_store = Mock() + return LogoutUseCase(session_store=session_store) - @patch('packages.application.auth.login_use_case.session_store') - def test_logout_current_device(self, mock_session_store, use_case): + def test_logout_current_device(self, use_case): """测试登出当前设备""" - mock_session_store.delete_session.return_value = True + use_case.session_store.delete_session.return_value = True request = LogoutRequest( user_id="user-123", @@ -149,12 +149,11 @@ class TestLogoutUseCase: assert success is True assert error is None - mock_session_store.delete_session.assert_called_once_with("session-abc") + use_case.session_store.delete_session.assert_called_once_with("session-abc") - @patch('packages.application.auth.login_use_case.session_store') - def test_logout_all_devices(self, mock_session_store, use_case): + def test_logout_all_devices(self, use_case): """测试登出所有设备""" - mock_session_store.delete_all_user_sessions.return_value = 3 + use_case.session_store.delete_all_user_sessions.return_value = 3 request = LogoutRequest( user_id="user-123", @@ -166,12 +165,11 @@ class TestLogoutUseCase: assert success is True assert error is None - mock_session_store.delete_all_user_sessions.assert_called_once_with("user-123") + use_case.session_store.delete_all_user_sessions.assert_called_once_with("user-123") - @patch('packages.application.auth.login_use_case.session_store') - def test_logout_session_not_found(self, mock_session_store, use_case): + def test_logout_session_not_found(self, use_case): """测试 session 不存在""" - mock_session_store.delete_session.return_value = False + use_case.session_store.delete_session.return_value = False request = LogoutRequest( user_id="user-123", diff --git a/tests/unit/test_password_reset_use_case.py b/tests/unit/test_password_reset_use_case.py index 426900b43..762ed9413 100644 --- a/tests/unit/test_password_reset_use_case.py +++ b/tests/unit/test_password_reset_use_case.py @@ -2,7 +2,7 @@ 密码重置 Use Case 测试 """ import pytest -from unittest.mock import Mock, patch +from unittest.mock import Mock from datetime import datetime, timedelta, timezone from packages.application.auth.password_reset_use_case import ( RequestPasswordResetUseCase, @@ -25,10 +25,13 @@ class TestRequestPasswordResetUseCase: @pytest.fixture def use_case(self, mock_user_repo): + email_service = Mock() + email_service.send_password_reset_email.return_value = (True, None) return RequestPasswordResetUseCase( user_repository=mock_user_repo, base_url="https://test.com", token_expire_hours=1, + email_service=email_service, ) @pytest.fixture @@ -41,11 +44,9 @@ class TestRequestPasswordResetUseCase: password_hash="hash", ) - @patch('packages.application.auth.password_reset_use_case.email_service') - def test_request_reset_success(self, mock_email_service, use_case, mock_user_repo, test_user): + def test_request_reset_success(self, use_case, mock_user_repo, test_user): """测试请求重置成功""" mock_user_repo.find_by_email.return_value = test_user - mock_email_service.send_password_reset_email.return_value = (True, None) request = RequestPasswordResetRequest(email="test@example.com") success, error = use_case.execute(request) @@ -60,10 +61,9 @@ class TestRequestPasswordResetUseCase: assert saved_user.password_reset_expires_at is not None # 验证发送了邮件 - mock_email_service.send_password_reset_email.assert_called_once() + use_case.email_service.send_password_reset_email.assert_called_once() - @patch('packages.application.auth.password_reset_use_case.email_service') - def test_request_reset_user_not_exists(self, mock_email_service, use_case, mock_user_repo): + def test_request_reset_user_not_exists(self, use_case, mock_user_repo): """测试用户不存在(仍返回成功,避免暴露)""" mock_user_repo.find_by_email.return_value = None @@ -74,7 +74,7 @@ class TestRequestPasswordResetUseCase: assert error is None # 不发送邮件 - mock_email_service.send_password_reset_email.assert_not_called() + use_case.email_service.send_password_reset_email.assert_not_called() def test_request_reset_missing_email(self, use_case): """测试缺少邮箱""" diff --git a/tests/unit/test_register_user_use_case.py b/tests/unit/test_register_user_use_case.py index 049f0ff8c..4cfa8405b 100644 --- a/tests/unit/test_register_user_use_case.py +++ b/tests/unit/test_register_user_use_case.py @@ -2,7 +2,7 @@ 用户注册 Use Case 测试 """ import pytest -from unittest.mock import Mock, patch +from unittest.mock import Mock from packages.application.auth import ( RegisterUserUseCase, RegisterUserRequest, @@ -28,15 +28,16 @@ class TestRegisterUserUseCase: @pytest.fixture def use_case(self, mock_user_repo): """创建注册用例""" + email_service = Mock() + email_service.send_verification_email.return_value = (True, None) return RegisterUserUseCase( user_repository=mock_user_repo, - base_url="https://test.com" + base_url="https://test.com", + email_service=email_service, ) - @patch('packages.application.auth.register_user_use_case.email_service') - def test_register_user_success(self, mock_email_service, use_case, mock_user_repo): + def test_register_user_success(self, use_case, mock_user_repo): """测试注册成功""" - mock_email_service.send_verification_email.return_value = (True, None) request = RegisterUserRequest( email="test@example.com", @@ -136,10 +137,9 @@ class TestRegisterUserUseCase: assert response is None assert error == "Email is required" - @patch('packages.application.auth.register_user_use_case.email_service') - def test_register_user_email_send_failure(self, mock_email_service, use_case, mock_user_repo): + def test_register_user_email_send_failure(self, use_case, mock_user_repo): """测试邮件发送失败(用户仍然创建)""" - mock_email_service.send_verification_email.return_value = (False, "SMTP error") + use_case.email_service.send_verification_email.return_value = (False, "SMTP error") request = RegisterUserRequest( email="test@example.com",