from __future__ import annotations from contextvars import ContextVar from datetime import date, datetime, timedelta import json import logging import sys from logging import Handler, LogRecord from pathlib import Path import re import time from typing import Any, Callable, Optional 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="-") client_ip_var: ContextVar[str] = ContextVar("client_ip", default="") def get_client_ip(request: Request | None) -> str: """获取客户端地址,兼容前置反向代理传递的真实地址。""" if request is None: return "" for header in ("X-Real-IP", "X-Forwarded-For"): value = request.headers.get(header, "") if value: return value.split(",", 1)[0].strip() return request.client.host if request.client else "" # ==================== 敏感数据脱敏规则 ==================== SENSITIVE_PATTERNS: dict[str, Callable | str] = { "token": "***", "password": "***", "access_token": "***", "refresh_token": "***", "secret_key": "***", "authorization": "***", "bearer": "***", "api_key": "***", "private_key": "***", } def mask_value(key: str, value: Any) -> str: """对单个值进行脱敏处理""" if value is None: return "" str_val = str(value) handler = SENSITIVE_PATTERNS.get(key) if callable(handler): return handler(str_val) elif isinstance(handler, str): # 支持正则替换模式,如 r"1\d{3}\d{4}" try: return re.sub(handler, "***", str_val) except re.error: return "***" return handler def mask_sensitive_dict(data: dict) -> dict: """递归脱敏字典中的敏感字段""" if not data or not isinstance(data, dict): return data result = {} for key, value in data.items(): result[key] = mask_value(key, value) return result def mask_sensitive_string(text: str) -> str: """从文本中脱敏常见敏感信息""" if not text: return text # Mask the value as well as the key. Replacing only ``api_key=`` would # still leak the credential in audit messages and exception text. assignment_pattern = ( r"(Bearer\s+|(?:api[-_]?key|access[_-]?token|refresh[_-]?token|" r"secret[_-]?key|password|private[_-]?key|token)\s*[:=]\s*)" r"(\"[^\"]*\"|'[^']*'|[^\s,;]+)" ) try: text = re.sub(assignment_pattern, r"\1***", text, flags=re.IGNORECASE) except re.error: pass patterns = [ (r'Bearer\s+[A-Za-z0-9\-._]+', 'Bearer ***'), (r'\d{11}', r'\d{3}\*\d{4}'), # 手机号/身份证 (r'1[3-9]\d{9}', r'1\*{3}\*{4}'), # 手机号 ] for pattern, replacement in patterns: try: text = re.sub(pattern, replacement, text, flags=re.IGNORECASE) except re.error: pass return text # ==================== RequestId Filter ==================== class RequestIdFilter(logging.Filter): def filter(self, record: LogRecord) -> bool: record.request_id = request_id_var.get() return True # ==================== Enhanced JSON Formatter ==================== class JsonLogFormatter(logging.Formatter): """ 增强的 JSON 日志格式化器,支持结构化字段输出。 输出示例: { "@timestamp": "2026-08-17T18:30:00.123Z", "level": "INFO", "logger": "dataset.router", "message": "数据集创建成功", "module": "dataset.router", "function": "create_dataset", "file": "dataset/router.py", "line": 45, "process": 12345, "thread": "MainThread", "request_id": "req-abc123", "user_id": "u_admin", "client_ip": "192.168.1.100", "extra": {...} } """ 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", "-"), } # 从 record 中提取额外字段(通过 extra 参数传入) for attr in ("user_id", "client_ip", "target_type", "target_id", "duration_ms", "status_code", "error"): val = getattr(record, attr, None) if val is not None: payload[attr] = val # 处理异常信息 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=(",", ":")) # ==================== DateSizeRotatingFileHandler ==================== # (保持不变,已有实现) 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) # ==================== Structured Logger 封装 ==================== class StructuredLogger: """ 结构化日志记录器,提供统一的日志接口。 使用方式: logger = get_structured_logger('dataset.router') logger.info('创建数据集', dataset_id='ds_123') """ def __init__(self, name: str, module: str = ""): self.logger = logging.getLogger(name) self.name = name self.module = module @property def trace_id(self) -> str: return request_id_var.get("-") def info(self, message: str, **extra: Any) -> None: self._log("INFO", message, **extra) def warning(self, message: str, **extra: Any) -> None: self._log("WARNING", message, **extra) def error(self, message: str, **extra: Any) -> None: self._log("ERROR", message, **extra) def debug(self, message: str, **extra: Any) -> None: self._log("DEBUG", message, **extra) def _log(self, level: str, message: str, **extra: Any) -> None: """统一日志记录方法""" log_entry: dict[str, Any] = { "timestamp": datetime.utcnow().isoformat(), "level": level, "logger": self.name, "module": self.module, "message": message, "trace_id": self.trace_id, "extra": extra, } self.logger.log(getattr(logging, level, logging.INFO), json.dumps(log_entry, ensure_ascii=False, default=str)) def get_structured_logger(name: str, module: str = "") -> StructuredLogger: """获取结构化日志记录器""" return StructuredLogger(name, module) # ==================== 快捷函数 ==================== def get_logger(name: str) -> logging.Logger: """获取标准 Python logger""" return logging.getLogger(name) def set_request_id(request_id: str) -> None: """设置当前请求的追踪 ID""" request_id_var.set(request_id) def setup_request_logging(app: FastAPI) -> None: """配置 FastAPI 请求日志中间件""" 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) ip_token = client_ip_var.set(get_client_ip(request)) started_at = time.perf_counter() try: response = await call_next(request) elapsed_ms = (time.perf_counter() - started_at) * 1000 noisy_paths = ("/health", "/system-info", "/compute/jobs/", "/model-eval/", "/model-compare/") log_method = logger.debug if request.method == "GET" and response.status_code < 400 else logger.info if any(request.url.path.endswith(path) or path in request.url.path for path in noisy_paths) and response.status_code < 400: log_method = logger.debug if response.status_code >= 400: log_method = logger.warning log_method( "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 "-", ) # 5xx 系统错误自动写入操作日志(未被 @op_log 覆盖的系统级异常) if response.status_code >= 500: try: from app.core.op_log import log_operation, OpModule, OpStatus log_operation( module=OpModule.SYSTEM, action="request", target_type="api", target_name=request.url.path, status=OpStatus.FAILURE, error_message=f"HTTP {response.status_code} - 系统内部错误", error_type="HTTPError", detail=f'{{"method":"{request.method}","path":"{request.url.path}","status":{response.status_code}}}', func_name="request_logging_middleware", request=request, duration_ms=elapsed_ms, ) except Exception: pass # 日志写入失败不影响主流程 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 "-", ) # 未被捕获的异常,写入操作日志 try: import traceback as _tb from app.core.op_log import log_operation, OpModule, OpStatus log_operation( module=OpModule.SYSTEM, action="request", target_type="api", target_name=request.url.path, status=OpStatus.FAILURE, error_message=str(sys.exc_info()[1])[:1000] if sys.exc_info()[1] else "未知异常", error_type=type(sys.exc_info()[1]).__name__ if sys.exc_info()[1] else "UnknownError", error_traceback="".join(_tb.format_exception(*sys.exc_info()))[:5000], func_name="request_logging_middleware", detail=f'{{"method":"{request.method}","path":"{request.url.path}"}}', request=request, duration_ms=elapsed_ms, ) except Exception: pass # 日志写入失败不影响主流程 raise finally: request_id_var.reset(token) client_ip_var.reset(ip_token) # ==================== 配置函数 ==================== 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 logging.getLogger("uvicorn.access").setLevel(logging.WARNING) logging.getLogger("psycopg.pool").setLevel(logging.ERROR)