refactor(auth): route simple auth through use cases

This commit is contained in:
Xiaoxia AI
2026-06-21 08:22:49 +08:00
parent 0f7cc8f12a
commit 93271f7dda
6 changed files with 295 additions and 168 deletions
+69 -114
View File
@@ -1,20 +1,20 @@
"""
认证 APISQLAlchemy 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 "注册失败")
+8
View File
@@ -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)
@@ -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,
)
+28 -7
View File
@@ -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(
@@ -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"),
}
+111 -46
View File
@@ -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"