import logging from functools import partial from typing import Any, Dict, List, Optional from starlette import status from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint from starlette.requests import Request from starlette.responses import Response from starlette.types import ASGIApp from audience_common.logger.utils import get_status_code_log_level from audience_common.utils.dictutil import exclude from audience_common.utils.urlpath import is_path_match class RequestLoggerMiddleware(BaseHTTPMiddleware): LOG_MESSAGE = "{status} - {verb} {resource}" def __init__( self, app: ASGIApp, exclude_paths: Optional[List[str]] = None, extra_headers: bool = False, logger: logging.Logger = logging.getLogger(__name__), ): super().__init__(app) self.exclude_paths = exclude_paths self.extra_headers = extra_headers self.logger = logger async def dispatch( self, request: Request, call_next: RequestResponseEndpoint ) -> Response: if is_path_match(request.url.path, match=self.exclude_paths): return await call_next(request) response = Response(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR) try: response = await call_next(request) finally: self.log_response(request, response) return response def log_response(self, request: Request, response: Response) -> None: extra: Dict[str, Any] = {} if self.extra_headers: filter_headers = partial(exclude, keys=["authorization"]) extra.update({"request.headers": filter_headers(dict(request.headers))}) if request.headers: extra.update( {"response.headers": filter_headers(dict(response.headers))} ) level = get_status_code_log_level(response.status_code) self.logger.log( level, self.LOG_MESSAGE.format( status=response.status_code, verb=request.method, resource=request.url.path, ), extra=extra, )