From f0d0edf449c6f58f4d16c4474754ee458ccab06f Mon Sep 17 00:00:00 2001 From: Xiaoxia AI Date: Wed, 17 Jun 2026 01:24:45 +0800 Subject: [PATCH] feat(auth): add login, logout and refresh token use cases - Implement LoginUseCase with password verification and JWT token generation - Generate user_auth token (without workspace) for initial login - Create session with refresh_token in Redis - Track last_login_at and last_login_ip - Implement LogoutUseCase for single device or all devices - Add RefreshTokenUseCase placeholder (to be implemented) - Add 9 comprehensive unit tests (all passed) Phase 4 Task 7/68 completed --- packages/application/auth/__init__.py | 16 ++ packages/application/auth/login_use_case.py | 232 ++++++++++++++++++++ tests/unit/test_login_use_case.py | 198 +++++++++++++++++ 3 files changed, 446 insertions(+) create mode 100644 packages/application/auth/login_use_case.py create mode 100644 tests/unit/test_login_use_case.py diff --git a/packages/application/auth/__init__.py b/packages/application/auth/__init__.py index e73afd1bc..6debd1bdc 100644 --- a/packages/application/auth/__init__.py +++ b/packages/application/auth/__init__.py @@ -6,6 +6,15 @@ from packages.application.auth.register_user_use_case import ( VerifyEmailUseCase, VerifyEmailRequest, ) +from packages.application.auth.login_use_case import ( + LoginUseCase, + LoginRequest, + LoginResponse, + RefreshTokenUseCase, + RefreshTokenRequest, + LogoutUseCase, + LogoutRequest, +) __all__ = [ "RegisterUserUseCase", @@ -13,4 +22,11 @@ __all__ = [ "RegisterUserResponse", "VerifyEmailUseCase", "VerifyEmailRequest", + "LoginUseCase", + "LoginRequest", + "LoginResponse", + "RefreshTokenUseCase", + "RefreshTokenRequest", + "LogoutUseCase", + "LogoutRequest", ] diff --git a/packages/application/auth/login_use_case.py b/packages/application/auth/login_use_case.py new file mode 100644 index 000000000..3cdf600a6 --- /dev/null +++ b/packages/application/auth/login_use_case.py @@ -0,0 +1,232 @@ +""" +用户登录 Use Case +""" +import secrets +from datetime import datetime, timezone, timedelta +from typing import Optional + +from packages.domain.auth import ( + password_hasher, + jwt_service, + session_store, +) + + +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): + self.user_repository = user_repository + + 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. 生成基础 JWT token(不包含 workspace,登录后需要选择工作空间) + # 这里使用一个特殊的 "user_token",不包含 workspace 和 role + # 用户选择工作空间后,会换取包含 workspace 的 access_token + import jwt as pyjwt + now = datetime.now(timezone.utc) + access_token_payload = { + "sub": user.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 + ) + + # 生成 refresh_token + refresh_token = secrets.token_urlsafe(32) + + # 6. 创建 session 并保存 refresh_token + session_id = secrets.token_urlsafe(16) + 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 天 + ) + + # 7. 更新最后登录信息 + user.last_login_at = datetime.now(timezone.utc) + user.last_login_ip = request.ip_address + self.user_repository.save(user) + + # 8. 返回响应 + 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): + pass + + def execute(self, request: LogoutRequest) -> tuple[bool, Optional[str]]: + """ + 执行登出 + + Args: + request: 登出请求 + + Returns: + (是否成功, 错误信息) + """ + try: + if request.logout_all_devices: + # 删除所有设备的 session + count = 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) + if success: + return True, None + else: + return False, "Session not found" + + except Exception as e: + return False, f"Logout failed: {str(e)}" diff --git a/tests/unit/test_login_use_case.py b/tests/unit/test_login_use_case.py new file mode 100644 index 000000000..21b7e3969 --- /dev/null +++ b/tests/unit/test_login_use_case.py @@ -0,0 +1,198 @@ +""" +用户登录 Use Case 测试 +""" +import pytest +from unittest.mock import Mock, patch +from datetime import datetime, timezone +from packages.application.auth import ( + LoginUseCase, + LoginRequest, + LogoutUseCase, + LogoutRequest, +) +from packages.domain.entities import User +from packages.domain.auth import password_hasher + + +class TestLoginUseCase: + """登录用例测试""" + + @pytest.fixture + def mock_user_repo(self): + repo = Mock() + repo.find_by_email = Mock(return_value=None) + repo.save = Mock() + return repo + + @pytest.fixture + def use_case(self, mock_user_repo): + return LoginUseCase(user_repository=mock_user_repo) + + @pytest.fixture + def test_user(self): + """创建测试用户""" + password_hash = password_hasher.hash_password("SecurePass123") + return User( + id="user-123", + email="test@example.com", + username="testuser", + display_name="Test User", + password_hash=password_hash, + 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): + """测试登录成功""" + mock_user_repo.find_by_email.return_value = test_user + mock_session_store.save_session.return_value = True + + request = LoginRequest( + email="test@example.com", + password="SecurePass123", + device_info="Chrome/Windows", + ip_address="192.168.1.1", + ) + + response, error = use_case.execute(request) + + assert error is None + assert response is not None + assert response.user_id == "user-123" + assert response.email == "test@example.com" + assert response.username == "testuser" + assert response.access_token != "" + assert response.refresh_token != "" + assert response.expires_in > 0 + + # 验证保存了 session + mock_session_store.save_session.assert_called_once() + + # 验证更新了最后登录信息 + mock_user_repo.save.assert_called_once() + + def test_login_invalid_email(self, use_case, mock_user_repo): + """测试邮箱不存在""" + mock_user_repo.find_by_email.return_value = None + + request = LoginRequest( + email="nonexistent@example.com", + password="SecurePass123", + ) + + response, error = use_case.execute(request) + + assert response is None + assert error == "Invalid email or password" + + def test_login_wrong_password(self, use_case, mock_user_repo, test_user): + """测试密码错误""" + mock_user_repo.find_by_email.return_value = test_user + + request = LoginRequest( + email="test@example.com", + password="WrongPassword123", + ) + + response, error = use_case.execute(request) + + assert response is None + assert error == "Invalid email or password" + + def test_login_missing_email(self, use_case): + """测试缺少邮箱""" + request = LoginRequest( + email="", + password="SecurePass123", + ) + + response, error = use_case.execute(request) + + assert response is None + assert error == "Email is required" + + def test_login_missing_password(self, use_case, mock_user_repo, test_user): + """测试缺少密码""" + mock_user_repo.find_by_email.return_value = test_user + + request = LoginRequest( + email="test@example.com", + password="", + ) + + response, error = use_case.execute(request) + + assert response is None + assert error == "Password is required" + + +class TestLogoutUseCase: + """登出用例测试""" + + @pytest.fixture + def use_case(self): + return LogoutUseCase() + + @patch('packages.application.auth.login_use_case.session_store') + def test_logout_current_device(self, mock_session_store, use_case): + """测试登出当前设备""" + mock_session_store.delete_session.return_value = True + + request = LogoutRequest( + user_id="user-123", + session_id="session-abc", + logout_all_devices=False, + ) + + success, error = use_case.execute(request) + + assert success is True + assert error is None + + mock_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): + """测试登出所有设备""" + mock_session_store.delete_all_user_sessions.return_value = 3 + + request = LogoutRequest( + user_id="user-123", + logout_all_devices=True, + ) + + success, error = use_case.execute(request) + + assert success is True + assert error is None + + mock_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): + """测试 session 不存在""" + mock_session_store.delete_session.return_value = False + + request = LogoutRequest( + user_id="user-123", + session_id="nonexistent", + logout_all_devices=False, + ) + + success, error = use_case.execute(request) + + assert success is False + assert error == "Session not found" + + def test_logout_missing_session_id(self, use_case): + """测试缺少 session_id""" + request = LogoutRequest( + user_id="user-123", + session_id=None, + logout_all_devices=False, + ) + + success, error = use_case.execute(request) + + assert success is False + assert error == "Session ID is required"