Files
xiaoxia-saas/packages/application/auth/login_use_case.py
T
Coze Agent 07b5589253 fix(p1): resolve 4 P1 technical debt issues
- P1-1: CORS configuration security - use DEBUG mode to differentiate
  production vs development CORS settings
- P1-2: Implement token refresh logic in RefreshTokenUseCase
  - Add get_session_by_refresh_token to SessionStore
  - Verify session validity and expiry
  - Generate new access token on refresh
- P1-3: Fix database connection leak in worker ingest task
  - Add proper try-except-finally block
  - Ensure db.close() is always called
- P1-4: Implement real media metadata extraction
  - Use ffprobe for video metadata
  - Use Pillow for image metadata
  - Return empty dict on failure (no mock data)
2026-06-26 18:11:07 +08:00

306 lines
10 KiB
Python
Executable File
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
用户登录 Use Case
"""
import hashlib
import secrets
from datetime import datetime, timedelta, timezone
from typing import Optional
import jwt as pyjwt
from packages.adapters.redis import get_session_store
from packages.adapters.redis.session_store import SessionStore, NoopSessionStore
from packages.domain.auth import jwt_service, password_hasher
LEGACY_SHA256_HEX_LENGTH = 64
def _is_legacy_sha256_hash(password_hash: str) -> bool:
return len(password_hash) == LEGACY_SHA256_HEX_LENGTH and all(
char in "0123456789abcdef" for char in password_hash.lower()
)
def _legacy_sha256(password: str) -> str:
return hashlib.sha256(password.encode()).hexdigest()
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 "unknown"
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, jwt_secret_key: str | None = None):
self.user_repository = user_repository
self.session_store: SessionStore | NoopSessionStore = session_store or get_session_store()
self.jwt_secret_key = jwt_secret_key or jwt_service.config.SECRET_KEY
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. 验证密码
password_is_valid = password_hasher.verify_password(request.password, user.password_hash)
if not password_is_valid and _is_legacy_sha256_hash(user.password_hash):
password_is_valid = secrets.compare_digest(_legacy_sha256(request.password), user.password_hash)
if password_is_valid:
user.password_hash = password_hasher.hash_password(request.password)
self.user_repository.save(user)
if not password_is_valid:
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
now = datetime.now(timezone.utc)
access_token_payload = {
"sub": user.id,
"sid": session_id,
"type": "user_auth",
"iat": now,
"exp": now + timedelta(minutes=jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES),
}
access_token = pyjwt.encode(
access_token_payload,
self.jwt_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, session_store=None, jwt_secret_key: str | None = None):
self.user_repository = user_repository
self.session_store: SessionStore | NoopSessionStore = session_store or get_session_store()
self.jwt_secret_key = jwt_secret_key or jwt_service.config.SECRET_KEY
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. 通过 refresh_token 查找 session
session = self.session_store.get_session_by_refresh_token(request.refresh_token)
if not session:
return None, "Invalid or expired refresh token"
# 2. 检查 session 是否过期
session_id = session.get("session_id")
if not session_id:
return None, "Invalid session data"
# 3. 检查 session 是否在有效期内
expires_at_str = session.get("expires_at")
if expires_at_str:
expires_at = datetime.fromisoformat(expires_at_str)
if datetime.now(timezone.utc) > expires_at:
# session 已过期,删除它
self.session_store.delete_session(session_id)
return None, "Session has expired, please login again"
# 4. 获取用户信息
user_id = session.get("user_id")
if not user_id:
return None, "Invalid session: missing user_id"
user = self.user_repository.get(user_id)
if not user:
return None, "User not found"
# 5. 生成新的 access_token
now = datetime.now(timezone.utc)
access_token_payload = {
"sub": user.id,
"sid": session_id,
"type": "user_auth",
"iat": now,
"exp": now + timedelta(minutes=jwt_service.config.ACCESS_TOKEN_EXPIRE_MINUTES),
}
access_token = pyjwt.encode(
access_token_payload,
self.jwt_secret_key,
algorithm=jwt_service.config.ALGORITHM,
)
# 6. 更新 session 的最后活跃时间
self.session_store.update_last_active(session_id)
# 7. 返回新的登录响应(refresh_token 保持不变)
return (
LoginResponse(
access_token=access_token,
refresh_token=request.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"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)}"