from __future__ import annotations from contextvars import ContextVar from datetime import date, datetime, timedelta import json import logging from logging import Handler, LogRecord from pathlib import Path import re import time from typing import Any from uuid import uuid4 from fastapi import FastAPI, Request from app.core.config import Settings, get_settings request_id_var: ContextVar[str] = ContextVar("request_id", default="-") class RequestIdFilter(logging.Filter): def filter(self, record: LogRecord) -> bool: record.request_id = request_id_var.get() return True class JsonLogFormatter(logging.Formatter): """Format one JSON object per line for ELK/Filebeat collection.""" def format(self, record: LogRecord) -> str: payload: dict[str, Any] = { "@timestamp": datetime.fromtimestamp(record.created).astimezone().isoformat( timespec="milliseconds" ), "level": record.levelname, "logger": record.name, "message": record.getMessage(), "module": record.module, "function": record.funcName, "file": record.pathname, "line": record.lineno, "process": record.process, "thread": record.thread, "thread_name": record.threadName, "request_id": getattr(record, "request_id", "-"), } if record.exc_info: payload["exception"] = self.formatException(record.exc_info) if record.stack_info: payload["stack"] = self.formatStack(record.stack_info) return json.dumps(payload, ensure_ascii=False, separators=(",", ":")) class DateSizeRotatingFileHandler(Handler): """Rotate log files by date and size while keeping date in every file name.""" def __init__( self, log_dir: str | Path, file_prefix: str, max_bytes: int, retention_days: int, encoding: str = "utf-8", ) -> None: super().__init__() self.log_dir = Path(log_dir) self.file_prefix = file_prefix self.max_bytes = max_bytes self.retention_days = retention_days self.encoding = encoding self._current_date: date | None = None self._stream: Any | None = None self._current_path: Path | None = None self.log_dir.mkdir(parents=True, exist_ok=True) def emit(self, record: LogRecord) -> None: try: message = self.format(record) + self.terminator encoded_size = len(message.encode(self.encoding)) self._ensure_stream() if self._should_rotate(encoded_size): self._rotate_by_size() self._ensure_stream(force=True) self._stream.write(message) self.flush() self._cleanup_expired_files() except Exception: self.handleError(record) @property def terminator(self) -> str: return "\n" def flush(self) -> None: if self._stream and not self._stream.closed: self._stream.flush() def close(self) -> None: try: if self._stream and not self._stream.closed: self._stream.close() finally: self._stream = None super().close() def _dated_path(self, target_date: date) -> Path: return self.log_dir / f"{self.file_prefix}-{target_date.isoformat()}.log" def _ensure_stream(self, force: bool = False) -> None: today = date.today() if not force and self._stream and self._current_date == today: return if self._stream and not self._stream.closed: self._stream.close() self._current_date = today self._current_path = self._dated_path(today) self._stream = self._current_path.open("a", encoding=self.encoding) def _should_rotate(self, incoming_size: int) -> bool: if not self._current_path or self.max_bytes <= 0: return False if not self._current_path.exists(): return False return self._current_path.stat().st_size + incoming_size > self.max_bytes def _rotate_by_size(self) -> None: if not self._current_path or not self._current_path.exists(): return if self._stream and not self._stream.closed: self._stream.close() self._stream = None stem = self._current_path.stem suffix = self._current_path.suffix index = 1 while True: rotated_path = self.log_dir / f"{stem}.{index}{suffix}" if not rotated_path.exists(): self._current_path.rename(rotated_path) return index += 1 def _cleanup_expired_files(self) -> None: if self.retention_days <= 0: return cutoff = date.today() - timedelta(days=self.retention_days - 1) pattern = re.compile( rf"^{re.escape(self.file_prefix)}-(\d{{4}}-\d{{2}}-\d{{2}})(?:\.\d+)?\.log$" ) for path in self.log_dir.glob(f"{self.file_prefix}-*.log"): match = pattern.match(path.name) if not match: continue file_date = datetime.strptime(match.group(1), "%Y-%m-%d").date() if file_date < cutoff: path.unlink(missing_ok=True) def configure_logging(settings: Settings | None = None) -> None: settings = settings or get_settings() root_logger = logging.getLogger() root_logger.handlers.clear() root_logger.setLevel(settings.log_level.upper()) console_formatter = logging.Formatter( fmt=( "%(asctime)s | %(levelname)s | pid=%(process)d | %(threadName)s | " "request_id=%(request_id)s | %(name)s | %(pathname)s:%(lineno)d | %(message)s" ), datefmt="%Y-%m-%d %H:%M:%S", ) json_formatter = JsonLogFormatter() request_filter = RequestIdFilter() console_handler = logging.StreamHandler() console_handler.setFormatter(console_formatter) console_handler.addFilter(request_filter) file_handler = DateSizeRotatingFileHandler( log_dir=settings.log_dir, file_prefix=settings.log_file_prefix, max_bytes=settings.log_max_bytes, retention_days=settings.log_retention_days, ) file_handler.setFormatter(json_formatter) file_handler.addFilter(request_filter) error_file_handler = DateSizeRotatingFileHandler( log_dir=settings.log_dir, file_prefix=settings.log_error_file_prefix, max_bytes=settings.log_max_bytes, retention_days=settings.log_retention_days, ) error_file_handler.setLevel(logging.ERROR) error_file_handler.setFormatter(json_formatter) error_file_handler.addFilter(request_filter) root_logger.addHandler(console_handler) root_logger.addHandler(file_handler) root_logger.addHandler(error_file_handler) for logger_name in ("uvicorn", "uvicorn.error", "uvicorn.access"): logger = logging.getLogger(logger_name) logger.handlers.clear() logger.propagate = True def get_logger(name: str) -> logging.Logger: return logging.getLogger(name) def set_request_id(request_id: str) -> None: request_id_var.set(request_id) def setup_request_logging(app: FastAPI) -> None: logger = get_logger("app.access") @app.middleware("http") async def request_logging_middleware(request: Request, call_next): # type: ignore[no-untyped-def] request_id = request.headers.get("X-Request-ID") or str(uuid4()) token = request_id_var.set(request_id) started_at = time.perf_counter() try: response = await call_next(request) elapsed_ms = (time.perf_counter() - started_at) * 1000 logger.info( "request completed method=%s path=%s status_code=%s duration_ms=%.2f client=%s", request.method, request.url.path, response.status_code, elapsed_ms, request.client.host if request.client else "-", ) response.headers["X-Request-ID"] = request_id return response except Exception: elapsed_ms = (time.perf_counter() - started_at) * 1000 logger.exception( "request failed method=%s path=%s duration_ms=%.2f client=%s", request.method, request.url.path, elapsed_ms, request.client.host if request.client else "-", ) raise finally: request_id_var.reset(token)