import logging from starlette.middleware.base import BaseHTTPMiddleware from starlette.requests import Request from starlette.responses import JSONResponse from app.core.config import get_settings from app.db.repositories.repositories import UserRepository from app.db.session import SessionLocal from app.services.auth_service import token_digest, verify_token_claims logger = logging.getLogger(__name__) class AuthenticationMiddleware(BaseHTTPMiddleware): async def dispatch(self, request: Request, call_next): path = request.url.path settings = get_settings() prefix = settings.app_api_prefix.rstrip("/") if path in { f"{prefix}/health", f"{prefix}/auth/login", f"{prefix}/auth/refresh", } or not path.startswith(f"{prefix}/"): return await call_next(request) authorization = request.headers.get("Authorization", "") if not authorization.startswith("Bearer "): return self._reject(request, "missing_bearer_token") try: claims = verify_token_claims(authorization[7:], settings) except ValueError: logger.error("Authentication configuration invalid") return JSONResponse( status_code=500, content={"error": {"code": "CONFIGURATION_ERROR", "message": "认证令牌配置无效"}}, ) if claims is None: return self._reject(request, "invalid_token") user_id, token_version, session_id = claims factory = getattr(request.app.state, "session_factory", SessionLocal) with factory() as db: repo = UserRepository(db) user = repo.get_active(user_id) revoked = repo.is_token_revoked(token_digest(authorization[7:])) session_revoked = session_id is not None and not repo.is_refresh_family_active( session_id ) if not user or revoked or session_revoked or user.token_version != token_version: reason = "inactive_user" if not user else "revoked_token" if revoked else "stale_token" return self._reject(request, reason) request.state.user_id = user.id request.state.session_id = session_id return await call_next(request) @staticmethod def _reject(request: Request, reason: str) -> JSONResponse: logger.warning( "Authentication rejected reason=%s request_id=%s", reason, getattr(request.state, "request_id", None), ) return JSONResponse( status_code=401, headers={"WWW-Authenticate": "Bearer"}, content={ "error": { "code": "UNAUTHORIZED", "message": "需要有效的访问令牌", "details": {}, "request_id": getattr(request.state, "request_id", None), } }, )