first commit
This commit is contained in:
@@ -0,0 +1,73 @@
|
||||
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),
|
||||
}
|
||||
},
|
||||
)
|
||||
@@ -0,0 +1,13 @@
|
||||
from uuid import uuid4
|
||||
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.requests import Request
|
||||
|
||||
|
||||
class RequestIdMiddleware(BaseHTTPMiddleware):
|
||||
async def dispatch(self, request: Request, call_next):
|
||||
request_id = request.headers.get("X-Request-ID") or str(uuid4())
|
||||
request.state.request_id = request_id
|
||||
response = await call_next(request)
|
||||
response.headers["X-Request-ID"] = request_id
|
||||
return response
|
||||
Reference in New Issue
Block a user