74 lines
2.9 KiB
Python
74 lines
2.9 KiB
Python
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),
|
|
}
|
|
},
|
|
)
|