"""ASGI request logging middleware.""" import logging from typing import Optional from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint from starlette.requests import Request from starlette.responses import Response from starlette.types import ASGIApp from owslogger.constants import AUTOLOG_MESSAGE class RequestLoggingMiddleware(BaseHTTPMiddleware): """Middleware to log request / response information.""" def __init__( self, app: ASGIApp, exclude_paths: Optional[list[str]] = None, logger: Optional[logging.Logger] = None, ) -> None: super().__init__(app) self.exclude_paths = exclude_paths or [] self.logger = logger or logging.getLogger(__name__) async def dispatch( self, request: Request, call_next: RequestResponseEndpoint ) -> Response: if request.url.path in self.exclude_paths: return await call_next(request) try: response = await call_next(request) except Exception: self._log_request(request, status_code=500) raise else: self._log_request(request, response.status_code) return response def _log_request(self, request: Request, status_code: int) -> None: """Log request details with appropriate log level.""" self.logger.log( level=self._get_log_level(status_code), msg=AUTOLOG_MESSAGE.format( status=status_code, verb=request.method, resource=request.url.path, ), ) @staticmethod def _get_log_level(status_code: Optional[int]) -> int: if status_code is None: return logging.INFO if 299 < status_code < 499: return logging.WARNING elif 499 < status_code: return logging.ERROR return logging.INFO