235 lines
7.3 KiB
Python
235 lines
7.3 KiB
Python
"""
|
||
用户登录 Use Case
|
||
"""
|
||
|
||
import secrets
|
||
from datetime import datetime, timedelta, timezone
|
||
from typing import Optional
|
||
|
||
from packages.adapters.redis import get_session_store
|
||
from packages.domain.auth import jwt_service, password_hasher
|
||
|
||
|
||
class LoginRequest:
|
||
"""登录请求"""
|
||
|
||
def __init__(
|
||
self,
|
||
email: str,
|
||
password: str,
|
||
device_info: Optional[str] = None,
|
||
ip_address: Optional[str] = None,
|
||
):
|
||
self.email = email.strip().lower()
|
||
self.password = password
|
||
self.device_info = device_info or "Unknown"
|
||
self.ip_address = ip_address or "0.0.0.0"
|
||
|
||
|
||
class LoginResponse:
|
||
"""登录响应"""
|
||
|
||
def __init__(
|
||
self,
|
||
access_token: str,
|
||
refresh_token: str,
|
||
user_id: str,
|
||
email: str,
|
||
username: str,
|
||
display_name: str,
|
||
expires_in: int,
|
||
):
|
||
self.access_token = access_token
|
||
self.refresh_token = refresh_token
|
||
self.user_id = user_id
|
||
self.email = email
|
||
self.username = username
|
||
self.display_name = display_name
|
||
self.expires_in = expires_in
|
||
|
||
|
||
class LoginUseCase:
|
||
"""用户登录用例"""
|
||
|
||
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]]:
|
||
"""
|
||
执行登录
|
||
|
||
Args:
|
||
request: 登录请求
|
||
|
||
Returns:
|
||
(登录响应, 错误信息)
|
||
"""
|
||
try:
|
||
# 1. 验证输入
|
||
if not request.email:
|
||
return None, "Email is required"
|
||
|
||
if not request.password:
|
||
return None, "Password is required"
|
||
|
||
# 2. 查找用户
|
||
user = self.user_repository.find_by_email(request.email)
|
||
if not user:
|
||
return None, "Invalid email or password"
|
||
|
||
# 3. 验证密码
|
||
if not password_hasher.verify_password(request.password, user.password_hash):
|
||
return None, "Invalid email or password"
|
||
|
||
# 4. 检查邮箱是否已验证(可选,根据需求决定是否强制)
|
||
# if not user.email_verified:
|
||
# return None, "Please verify your email first"
|
||
|
||
# 5. 创建 session 并生成 refresh_token
|
||
session_id = secrets.token_urlsafe(16)
|
||
refresh_token = secrets.token_urlsafe(32)
|
||
|
||
# 6. 生成基础 JWT token(包含 session_id,不包含 workspace)
|
||
# 这里使用一个特殊的 "user_token",不包含 workspace 和 role
|
||
# 用户选择工作空间后,会换取包含 workspace 的 access_token
|
||
import jwt as pyjwt
|
||
|
||
now = datetime.now(timezone.utc)
|
||
access_token_payload = {
|
||
"sub": user.id,
|
||
"sid": session_id, # 添加 session_id
|
||
"type": "user_auth", # 标记为用户认证 token(未绑定工作空间)
|
||
"iat": now,
|
||
"exp": now + timedelta(minutes=jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES),
|
||
}
|
||
access_token = pyjwt.encode(
|
||
access_token_payload,
|
||
jwt_service.config.SECRET_KEY,
|
||
algorithm=jwt_service.config.ALGORITHM,
|
||
)
|
||
self.session_store.save_session(
|
||
session_id=session_id,
|
||
user_id=user.id,
|
||
refresh_token=refresh_token,
|
||
device_info=request.device_info,
|
||
ip_address=request.ip_address,
|
||
expires_in_seconds=30 * 24 * 3600, # 30 天
|
||
)
|
||
|
||
# 8. 更新最后登录信息
|
||
user.last_login_at = datetime.now(timezone.utc)
|
||
user.last_login_ip = request.ip_address
|
||
self.user_repository.save(user)
|
||
|
||
# 9. 返回响应
|
||
return (
|
||
LoginResponse(
|
||
access_token=access_token,
|
||
refresh_token=refresh_token,
|
||
user_id=user.id,
|
||
email=user.email,
|
||
username=user.username,
|
||
display_name=user.display_name,
|
||
expires_in=jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES * 60,
|
||
),
|
||
None,
|
||
)
|
||
|
||
except Exception as e:
|
||
return None, f"Login failed: {str(e)}"
|
||
|
||
|
||
class RefreshTokenRequest:
|
||
"""刷新令牌请求"""
|
||
|
||
def __init__(self, refresh_token: str):
|
||
self.refresh_token = refresh_token
|
||
|
||
|
||
class RefreshTokenUseCase:
|
||
"""刷新令牌用例"""
|
||
|
||
def __init__(self, user_repository):
|
||
self.user_repository = user_repository
|
||
|
||
def execute(self, request: RefreshTokenRequest) -> tuple[Optional[LoginResponse], Optional[str]]:
|
||
"""
|
||
执行令牌刷新
|
||
|
||
Args:
|
||
request: 刷新请求
|
||
|
||
Returns:
|
||
(登录响应, 错误信息)
|
||
"""
|
||
try:
|
||
if not request.refresh_token:
|
||
return None, "Refresh token is required"
|
||
|
||
# 1. 查找 session(通过遍历所有 session)
|
||
# 注意:这里为了简化,先用遍历实现,生产环境应该用 refresh_token -> session_id 的索引
|
||
session = None
|
||
session_id = None
|
||
|
||
# 这是一个简化实现,实际应该在 SessionStore 中添加 find_by_refresh_token 方法
|
||
# 这里我们假设 refresh_token 就是 session_id(简化处理)
|
||
# 生产环境需要更复杂的映射
|
||
|
||
# 临时方案:从 Redis 获取(需要在 session_store 中添加方法)
|
||
# 现在先返回错误,提示需要实现
|
||
return None, "Refresh token implementation pending (需要完善 session_store)"
|
||
|
||
except Exception as e:
|
||
return None, f"Token refresh failed: {str(e)}"
|
||
|
||
|
||
class LogoutRequest:
|
||
"""登出请求"""
|
||
|
||
def __init__(
|
||
self,
|
||
user_id: str,
|
||
session_id: Optional[str] = None,
|
||
logout_all_devices: bool = False,
|
||
):
|
||
self.user_id = user_id
|
||
self.session_id = session_id
|
||
self.logout_all_devices = logout_all_devices
|
||
|
||
|
||
class LogoutUseCase:
|
||
"""用户登出用例"""
|
||
|
||
def __init__(self, session_store=None):
|
||
self.session_store = session_store or get_session_store()
|
||
|
||
def execute(self, request: LogoutRequest) -> tuple[bool, Optional[str]]:
|
||
"""
|
||
执行登出
|
||
|
||
Args:
|
||
request: 登出请求
|
||
|
||
Returns:
|
||
(是否成功, 错误信息)
|
||
"""
|
||
try:
|
||
if request.logout_all_devices:
|
||
# 删除所有设备的 session
|
||
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 = self.session_store.delete_session(request.session_id)
|
||
if success:
|
||
return True, None
|
||
else:
|
||
return False, "Session not found"
|
||
|
||
except Exception as e:
|
||
return False, f"Logout failed: {str(e)}"
|