From 93271f7ddac73e50bc3d30c04459ce7fa4483dc2 Mon Sep 17 00:00:00 2001 From: Xiaoxia AI Date: Sun, 21 Jun 2026 08:22:49 +0800 Subject: [PATCH] refactor(auth): route simple auth through use cases --- apps/api/app/api/routes/auth_simple.py | 183 +++++++----------- apps/api/app/dependencies.py | 8 + .../sqlalchemy_impl/user_repository.py | 79 ++++++++ packages/application/auth/login_use_case.py | 35 +++- tests/unit/test_architecture_boundaries.py | 1 - tests/unit/test_auth_simple.py | 157 ++++++++++----- 6 files changed, 295 insertions(+), 168 deletions(-) create mode 100644 packages/adapters/sqlalchemy_impl/user_repository.py diff --git a/apps/api/app/api/routes/auth_simple.py b/apps/api/app/api/routes/auth_simple.py index e779e9ca5..41351f4bb 100644 --- a/apps/api/app/api/routes/auth_simple.py +++ b/apps/api/app/api/routes/auth_simple.py @@ -1,20 +1,20 @@ """ -认证 API(SQLAlchemy ORM) +认证 API compatibility routes. + +The route layer is intentionally thin: repository construction lives in +app.dependencies and authentication behavior lives in application use cases. """ -import hashlib -import secrets -from datetime import datetime, timedelta, timezone - -import jwt from app.config import settings -from app.dependencies import get_db_session +from app.dependencies import get_user_repository from fastapi import APIRouter, Depends, HTTPException, status from pydantic import BaseModel, EmailStr -from sqlalchemy.orm import Session -from packages.adapters.sqlalchemy_impl.models import UserModel -from packages.domain.auth import password_hasher, password_validator +from packages.application.auth.login_use_case import LoginRequest as LoginUseCaseRequest +from packages.application.auth.login_use_case import LoginUseCase +from packages.application.auth.register_user_use_case import RegisterUserRequest as RegisterUseCaseRequest +from packages.application.auth.register_user_use_case import RegisterUserUseCase +from packages.ports.user_repository import UserRepository router = APIRouter(prefix="/auth", tags=["认证"]) @@ -49,122 +49,57 @@ class LoginResponse(BaseModel): expires_in: int -ACCESS_TOKEN_EXPIRE_MINUTES = 30 -JWT_ALGORITHM = "HS256" -LEGACY_SHA256_HEX_LENGTH = 64 - - -def _normalize_email(email: str) -> str: - return email.strip().lower() - - -def _normalize_username(username: str) -> str: - return username.strip() - - -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() - - -def _verify_password_with_legacy_upgrade(password: str, user: UserModel, db: Session) -> bool: - stored_hash = user.password_hash or "" - if password_hasher.verify_password(password, stored_hash): - return True - - if _is_legacy_sha256_hash(stored_hash) and secrets.compare_digest(_legacy_sha256(password), stored_hash): - user.password_hash = password_hasher.hash_password(password) - db.add(user) - db.commit() - db.refresh(user) - return True - - return False - - -def _create_access_token(user: UserModel) -> tuple[str, int]: - expires_delta = timedelta(minutes=ACCESS_TOKEN_EXPIRE_MINUTES) - now = datetime.now(timezone.utc) - payload = { - "sub": user.id, - "email": user.email, - "type": "user_auth", - "iat": now, - "exp": now + expires_delta, - } - token = jwt.encode(payload, settings.JWT_SECRET_KEY, algorithm=JWT_ALGORITHM) - return token, int(expires_delta.total_seconds()) - - @router.post("/register", response_model=RegisterResponse, status_code=status.HTTP_201_CREATED) -async def register(request: RegisterRequest, db: Session = Depends(get_db_session)): - email = _normalize_email(request.email) - username = _normalize_username(request.username) - display_name = request.display_name.strip() - - if not username: - raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="用户名不能为空") - - if not display_name: - raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="显示名称不能为空") - - password_valid, password_error = password_validator.validate(request.password) - if not password_valid: - raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=password_error) - - existing_user = db.query(UserModel).filter(UserModel.email == email).first() - if existing_user: - raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="邮箱已被注册") - - existing_username = db.query(UserModel).filter(UserModel.username == username).first() - if existing_username: - raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="用户名已被使用") - - new_user = UserModel( - id=f"user_{secrets.token_hex(8)}", - email=email, - username=username, - display_name=display_name, - password_hash=password_hasher.hash_password(request.password), - email_verified=False, - created_at=datetime.now(timezone.utc), +async def register( + request: RegisterRequest, + user_repository: UserRepository = Depends(get_user_repository), +): + use_case = RegisterUserUseCase( + user_repository=user_repository, + base_url="http://localhost:3000", + email_service=_NoopEmailService(), ) - - db.add(new_user) - db.commit() - db.refresh(new_user) + response, error = use_case.execute( + RegisterUseCaseRequest( + email=request.email, + password=request.password, + username=request.username, + display_name=request.display_name, + ) + ) + if error or response is None: + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=_translate_auth_error(error)) return RegisterResponse( - user_id=new_user.id, - email=new_user.email, - username=new_user.username or "", - display_name=new_user.display_name, + user_id=response.user_id, + email=response.email, + username=response.username, + display_name=response.display_name, message="注册成功!", ) @router.post("/login", response_model=LoginResponse) -async def login(request: LoginRequest, db: Session = Depends(get_db_session)): - email = _normalize_email(request.email) - user = db.query(UserModel).filter(UserModel.email == email).first() - - if not user or not _verify_password_with_legacy_upgrade(request.password, user, db): +async def login( + request: LoginRequest, + user_repository: UserRepository = Depends(get_user_repository), +): + use_case = LoginUseCase( + user_repository=user_repository, + session_store=_NoopSessionStore(), + jwt_secret_key=settings.JWT_SECRET_KEY, + ) + response, error = use_case.execute(LoginUseCaseRequest(email=request.email, password=request.password)) + if error or response is None: raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="邮箱或密码错误") - access_token, expires_in = _create_access_token(user) - return LoginResponse( - access_token=access_token, - user_id=user.id, - email=user.email, - username=user.username or "", - display_name=user.display_name, - expires_in=expires_in, + access_token=response.access_token, + user_id=response.user_id, + email=response.email, + username=response.username, + display_name=response.display_name, + expires_in=response.expires_in, ) @@ -174,3 +109,23 @@ async def get_current_user_info(): status_code=status.HTTP_501_NOT_IMPLEMENTED, detail="/auth/me requires bearer-token dependency integration", ) + + +class _NoopSessionStore: + def save_session(self, **kwargs): + return None + + +class _NoopEmailService: + def send_verification_email(self, **kwargs): + return False, "Email delivery is disabled for compatibility auth routes" + + +def _translate_auth_error(error: str | None) -> str: + translations = { + "Email already registered": "邮箱已被注册", + "Username already taken": "用户名已被使用", + "Username is required": "用户名不能为空", + "Display name is required": "显示名称不能为空", + } + return translations.get(error or "", error or "注册失败") diff --git a/apps/api/app/dependencies.py b/apps/api/app/dependencies.py index c0cf958d2..74a2ad7f9 100644 --- a/apps/api/app/dependencies.py +++ b/apps/api/app/dependencies.py @@ -22,6 +22,8 @@ from packages.adapters.sqlalchemy_impl.project_repository import ( SQLAlchemyProjectRepository, ) from packages.adapters.sqlalchemy_impl.session import build_session_factory +from packages.adapters.sqlalchemy_impl.user_repository import SQLAlchemyUserRepository +from packages.ports.user_repository import UserRepository _engine, _SessionLocal = build_session_factory(settings.DATABASE_URL) @@ -74,3 +76,9 @@ def get_project_repository( session: Session = Depends(get_db_session), ) -> SQLAlchemyProjectRepository: return SQLAlchemyProjectRepository(session) + + +def get_user_repository( + session: Session = Depends(get_db_session), +) -> UserRepository: + return SQLAlchemyUserRepository(session) diff --git a/packages/adapters/sqlalchemy_impl/user_repository.py b/packages/adapters/sqlalchemy_impl/user_repository.py new file mode 100644 index 000000000..7b7f269a7 --- /dev/null +++ b/packages/adapters/sqlalchemy_impl/user_repository.py @@ -0,0 +1,79 @@ +from __future__ import annotations + +from sqlalchemy.orm import Session + +from packages.adapters.sqlalchemy_impl.models import UserModel +from packages.domain.entities import User +from packages.ports.user_repository import UserRepository + + +class SQLAlchemyUserRepository(UserRepository): + def __init__(self, session: Session): + self.session = session + + def save(self, user: User) -> None: + model = self.session.get(UserModel, user.id) + if model is None: + model = UserModel(id=user.id) + self.session.add(model) + + model.email = user.email + model.username = user.username + model.display_name = user.display_name + model.password_hash = user.password_hash + model.email_verified = user.email_verified + model.email_verification_token = user.email_verification_token + model.password_reset_token = user.password_reset_token + model.password_reset_expires_at = user.password_reset_expires_at + model.last_login_at = user.last_login_at + model.last_login_ip = user.last_login_ip + model.created_at = user.created_at + + self.session.commit() + self.session.refresh(model) + + def find_by_id(self, user_id: str) -> User | None: + return self._to_entity(self.session.get(UserModel, user_id)) + + def find_by_email(self, email: str) -> User | None: + model = self.session.query(UserModel).filter(UserModel.email == email.strip().lower()).first() + return self._to_entity(model) + + def find_by_username(self, username: str) -> User | None: + model = self.session.query(UserModel).filter(UserModel.username == username.strip()).first() + return self._to_entity(model) + + def find_by_verification_token(self, token: str) -> User | None: + model = self.session.query(UserModel).filter(UserModel.email_verification_token == token).first() + return self._to_entity(model) + + def find_by_password_reset_token(self, token: str) -> User | None: + model = self.session.query(UserModel).filter(UserModel.password_reset_token == token).first() + return self._to_entity(model) + + def delete(self, user_id: str) -> bool: + model = self.session.get(UserModel, user_id) + if model is None: + return False + self.session.delete(model) + self.session.commit() + return True + + @staticmethod + def _to_entity(model: UserModel | None) -> User | None: + if model is None: + return None + return User( + id=model.id, + email=model.email, + username=model.username or "", + display_name=model.display_name, + password_hash=model.password_hash, + email_verified=model.email_verified, + email_verification_token=model.email_verification_token, + password_reset_token=model.password_reset_token, + password_reset_expires_at=model.password_reset_expires_at, + last_login_at=model.last_login_at, + last_login_ip=model.last_login_ip, + created_at=model.created_at, + ) diff --git a/packages/application/auth/login_use_case.py b/packages/application/auth/login_use_case.py index febd925b4..41e8102a6 100644 --- a/packages/application/auth/login_use_case.py +++ b/packages/application/auth/login_use_case.py @@ -2,13 +2,28 @@ 用户登录 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.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: """登录请求""" @@ -51,9 +66,10 @@ class LoginResponse: class LoginUseCase: """用户登录用例""" - def __init__(self, user_repository, session_store=None): + def __init__(self, user_repository, session_store=None, jwt_secret_key: str | None = None): self.user_repository = user_repository self.session_store = 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]]: """ @@ -79,7 +95,14 @@ class LoginUseCase: return None, "Invalid email or password" # 3. 验证密码 - if not password_hasher.verify_password(request.password, user.password_hash): + 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. 检查邮箱是否已验证(可选,根据需求决定是否强制) @@ -93,19 +116,17 @@ class LoginUseCase: # 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(未绑定工作空间) + "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, - jwt_service.config.SECRET_KEY, + self.jwt_secret_key, algorithm=jwt_service.config.ALGORITHM, ) self.session_store.save_session( diff --git a/tests/unit/test_architecture_boundaries.py b/tests/unit/test_architecture_boundaries.py index 61be16f14..253c33e11 100644 --- a/tests/unit/test_architecture_boundaries.py +++ b/tests/unit/test_architecture_boundaries.py @@ -3,7 +3,6 @@ from pathlib import Path ALLOWED_API_ADAPTER_IMPORTS = { Path("apps/api/app/dependencies.py"), Path("apps/api/app/db.py"), - Path("apps/api/app/api/routes/auth_simple.py"), } diff --git a/tests/unit/test_auth_simple.py b/tests/unit/test_auth_simple.py index 30a5e9b0c..dd3d47e08 100644 --- a/tests/unit/test_auth_simple.py +++ b/tests/unit/test_auth_simple.py @@ -8,70 +8,135 @@ API_ROOT = ROOT / "apps" / "api" if str(API_ROOT) not in sys.path: sys.path.insert(0, str(API_ROOT)) -import importlib.util - -spec = importlib.util.spec_from_file_location("auth_simple", API_ROOT / "app" / "api" / "routes" / "auth_simple.py") -auth_simple = importlib.util.module_from_spec(spec) -assert spec.loader is not None -spec.loader.exec_module(auth_simple) - -_create_access_token = auth_simple._create_access_token -_verify_password_with_legacy_upgrade = auth_simple._verify_password_with_legacy_upgrade from app.config import settings -from packages.adapters.sqlalchemy_impl.models import UserModel +from packages.application.auth.login_use_case import LoginRequest, LoginUseCase +from packages.application.auth.register_user_use_case import RegisterUserRequest, RegisterUserUseCase from packages.domain.auth import password_hasher +from packages.domain.entities import User -class DummySession: +class InMemoryUserRepository: def __init__(self): - self.committed = False - self.refreshed = False - self.added = [] + self.users = {} - def add(self, item): - self.added.append(item) + def save(self, user): + self.users[user.id] = user - def commit(self): - self.committed = True + def find_by_id(self, user_id): + return self.users.get(user_id) - def refresh(self, item): - self.refreshed = True + def find_by_email(self, email): + return next((user for user in self.users.values() if user.email == email), None) + + def find_by_username(self, username): + return next((user for user in self.users.values() if user.username == username), None) + + def find_by_verification_token(self, token): + return None + + def find_by_password_reset_token(self, token): + return None + + def delete(self, user_id): + return self.users.pop(user_id, None) is not None -def test_create_access_token_returns_verifiable_jwt(): - user = UserModel(id="user-1", email="user@example.com", username="user", display_name="User") +class DummySessionStore: + def __init__(self): + self.saved = [] - token, expires_in = _create_access_token(user) - payload = jwt.decode(token, settings.JWT_SECRET_KEY, algorithms=["HS256"]) - - assert expires_in == 1800 - assert payload["sub"] == "user-1" - assert payload["email"] == "user@example.com" - assert payload["type"] == "user_auth" + def save_session(self, **kwargs): + self.saved.append(kwargs) -def test_verify_password_accepts_bcrypt_hash(): - db = DummySession() - user = UserModel(password_hash=password_hasher.hash_password("Password1")) - - assert _verify_password_with_legacy_upgrade("Password1", user, db) is True - assert db.committed is False +class DummyEmailService: + def send_verification_email(self, **kwargs): + return False, "disabled" -def test_verify_password_upgrades_legacy_sha256_hash(): - db = DummySession() - user = UserModel(password_hash="19513fdc9da4fb72a4a05eb66917548d3c90ff94d5419e1f2363eea89dfee1dd") +def test_register_use_case_hashes_password_and_normalizes_email(): + repo = InMemoryUserRepository() + use_case = RegisterUserUseCase(repo, email_service=DummyEmailService()) - assert _verify_password_with_legacy_upgrade("Password1", user, db) is True + response, error = use_case.execute( + RegisterUserRequest( + email="USER@EXAMPLE.COM", + password="Password1", + username="user", + display_name="User", + ) + ) + + assert error is None + assert response is not None + user = repo.find_by_email("user@example.com") + assert user is not None assert user.password_hash.startswith("$2") - assert db.committed is True - assert db.refreshed is True + assert password_hasher.verify_password("Password1", user.password_hash) -def test_verify_password_rejects_wrong_password(): - db = DummySession() - user = UserModel(password_hash=password_hasher.hash_password("Password1")) +def test_login_use_case_returns_verifiable_jwt(): + repo = InMemoryUserRepository() + session_store = DummySessionStore() + user = User( + id="user-1", + email="user@example.com", + username="user", + display_name="User", + password_hash=password_hasher.hash_password("Password1"), + ) + repo.save(user) - assert _verify_password_with_legacy_upgrade("WrongPassword1", user, db) is False - assert db.committed is False + response, error = LoginUseCase(repo, session_store=session_store, jwt_secret_key=settings.JWT_SECRET_KEY).execute( + LoginRequest("user@example.com", "Password1") + ) + + assert error is None + assert response is not None + payload = jwt.decode(response.access_token, settings.JWT_SECRET_KEY, algorithms=["HS256"]) + assert response.expires_in == 1800 + assert payload["sub"] == "user-1" + assert payload["type"] == "user_auth" + assert session_store.saved + + +def test_login_use_case_upgrades_legacy_sha256_hash(): + repo = InMemoryUserRepository() + user = User( + id="user-1", + email="user@example.com", + username="user", + display_name="User", + password_hash="19513fdc9da4fb72a4a05eb66917548d3c90ff94d5419e1f2363eea89dfee1dd", + ) + repo.save(user) + + response, error = LoginUseCase( + repo, session_store=DummySessionStore(), jwt_secret_key=settings.JWT_SECRET_KEY + ).execute(LoginRequest("user@example.com", "Password1")) + + assert error is None + assert response is not None + assert user.password_hash.startswith("$2") + assert password_hasher.verify_password("Password1", user.password_hash) + + +def test_login_use_case_rejects_wrong_password(): + repo = InMemoryUserRepository() + repo.save( + User( + id="user-1", + email="user@example.com", + username="user", + display_name="User", + password_hash=password_hasher.hash_password("Password1"), + ) + ) + + response, error = LoginUseCase(repo, session_store=DummySessionStore()).execute( + LoginRequest("user@example.com", "WrongPassword1") + ) + + assert response is None + assert error == "Invalid email or password"