from __future__ import annotations import json import logging import re import socket import sys import time from contextvars import ContextVar from datetime import date, datetime, timedelta from logging import Handler, LogRecord from pathlib import Path from typing import Any, Callable from uuid import uuid4 from fastapi import FastAPI, Request from app.core.config import Settings, get_settings # ==================== 链路追踪 ContextVar ==================== request_id_var: ContextVar[str] = ContextVar("request_id", default="-") user_id_var: ContextVar[str] = ContextVar("user_id", default="") client_ip_var: ContextVar[str] = ContextVar("client_ip", default="") # ==================== 敏感数据脱敏 ==================== SENSITIVE_KEYS: set[str] = { "password", "token", "access_token", "refresh_token", "secret_key", "authorization", "bearer", "api_key", "private_key", "secret", "cookie", } FULL_MASK_KEYS: set[str] = { "password", "token", "access_token", "refresh_token", "secret_key", "authorization", "bearer", "api_key", "private_key", "secret", "cookie", } def _mask_phone(value: str) -> str: """手机号脱敏:138****5678""" if len(value) >= 11: return value[:3] + "****" + value[-4:] return value def _mask_id_card(value: str) -> str: """身份证号脱敏:110***********1234""" if len(value) >= 18: return value[:3] + "***********" + value[-4:] return value def mask_value(key: str, value: Any) -> Any: """对单个值进行脱敏处理。""" if value is None: return "" key_lower = key.lower() if key_lower in FULL_MASK_KEYS: return "***" str_val = str(value) # 手机号模式(11位数字,1开头) if re.match(r"^1[3-9]\d{9}$", str_val): return _mask_phone(str_val) # 身份证模式(18位) if re.match(r"^\d{17}[\dXx]$", str_val): return _mask_id_card(str_val) return value def mask_sensitive_dict(data: dict) -> dict: """递归脱敏字典中的敏感字段。""" if not data or not isinstance(data, dict): return data result: dict[str, Any] = {} for key, value in data.items(): if isinstance(value, dict): result[key] = mask_sensitive_dict(value) elif isinstance(value, list): result[key] = [ mask_sensitive_dict(item) if isinstance(item, dict) else item for item in value ] else: result[key] = mask_value(key, value) return result def mask_sensitive_string(text: str) -> str: """从文本中脱敏常见敏感信息。""" if not text: return text patterns: list[tuple[str, str]] = [ (r"Bearer\s+[A-Za-z0-9\-._]+", "Bearer ***"), (r"(?i)token\s*[:=]\s*\S+", "token=***"), (r"(?i)password\s*[:=]\s*\S+", "password=***"), (r"(?i)secret[_-]?key\s*[:=]\s*\S+", "secret_key=***"), (r"(?i)api[-_]?key\s*[:=]\s*\S+", "api_key=***"), (r"(?i)private[_-]?key\s*[:=]\s*\S+", "private_key=***"), (r"(?i)authorization\s*[:=]\s*\S+", "authorization=***"), ] for pattern, replacement in patterns: text = re.sub(pattern, replacement, text) # 手机号脱敏 text = re.sub(r"\b1[3-9]\d{9}\b", lambda m: _mask_phone(m.group()), text) return text # ==================== 大对象截断 ==================== MAX_FIELD_SIZE = 1024 # 超过 1KB 的内容自动截断 def truncate_large_value(value: Any, max_size: int = MAX_FIELD_SIZE) -> Any: """超过 max_size 的字符串自动截断(前 500 + 后 500)。""" if isinstance(value, str) and len(value) > max_size: half = max_size // 2 return value[:half] + f"...[truncated {len(value) - max_size} chars]..." + value[-half:] if isinstance(value, dict): return {k: truncate_large_value(v, max_size) for k, v in value.items()} if isinstance(value, list): return [truncate_large_value(v, max_size) for v in value] return value # ==================== TraceId Filter ==================== class TraceIdFilter(logging.Filter): """自动注入 traceId / userId / clientIp 到每条日志记录。""" def filter(self, record: LogRecord) -> bool: record.traceId = request_id_var.get() record.userId = user_id_var.get("") record.clientIp = client_ip_var.get("") record.host = getattr(self, "_host", None) or socket.gethostname() record.app = getattr(self, "_app", "yg-ft-platform") record.env = getattr(self, "_env", "dev") return True def set_context(self, app: str, env: str, host: str) -> None: self._app = app self._env = env self._host = host # ==================== JSON Formatter ==================== class JsonLogFormatter(logging.Formatter): """ 生产级 JSON 日志格式化器,符合方案文档 §3.2 字段规范。 输出示例: { "@timestamp": "2026-08-19T10:30:45.123+08:00", "level": "INFO", "logger": "app.api.v1.endpoints.platform", "traceId": "abc-123-def-456", "userId": "u_admin", "message": "数据集创建成功", "fields": {"datasetId": "ds_001", "costMs": 23}, "file": "platform.py:156", "thread": "MainThread", "host": "pod-7x9k2", "app": "yg-ft-platform", "env": "dev" } """ # 标准 LogRecord 属性名集合,用于区分 extra 字段 _STD_ATTRS: set[str] = set(vars(logging.LogRecord("", 0, "", 0, "", None, None)).keys()) | { "traceId", "userId", "clientIp", "host", "app", "env", "request_id", "user_id", "client_ip", "asctime", "message", "module", "function", "process", "thread", "threadName", "levelname", "levelno", "name", "pathname", "filename", "lineno", "funcName", "created", "msecs", "relativeCreated", "exc_info", "exc_text", "stack_info", "msg", "args", "processName", "process", } def format(self, record: LogRecord) -> str: # 时间戳:ISO8601 带时区 timestamp = datetime.fromtimestamp(record.created).astimezone().isoformat( timespec="milliseconds" ) payload: dict[str, Any] = { "@timestamp": timestamp, "level": record.levelname, "logger": record.name, "traceId": getattr(record, "traceId", "-"), "message": record.getMessage(), "file": f"{Path(record.pathname).name}:{record.lineno}", "thread": record.threadName, "host": getattr(record, "host", ""), "app": getattr(record, "app", ""), "env": getattr(record, "env", ""), } # userId(业务必填,未登录可为空) user_id = getattr(record, "userId", "") or getattr(record, "user_id", "") if user_id: payload["userId"] = user_id # clientIp client_ip = getattr(record, "clientIp", "") or getattr(record, "client_ip", "") if client_ip: payload["clientIp"] = client_ip # 提取结构化业务字段:只收集通过 extra 传入的非标准属性 fields: dict[str, Any] = {} for attr in dir(record): if attr.startswith("_"): continue if attr in self._STD_ATTRS: continue if attr in ("traceId", "userId", "clientIp", "host", "app", "env"): continue val = getattr(record, attr, None) if val is not None and not callable(val): fields[attr] = truncate_large_value(val) if fields: payload["fields"] = mask_sensitive_dict(fields) # ERROR 级别额外字段 if record.levelname == "ERROR" or record.exc_info: error_obj: dict[str, Any] = { "type": type(record.exc_info[1]).__name__ if record.exc_info and record.exc_info[1] else "Error", "message": record.getMessage(), } if record.exc_info: error_obj["stack_trace"] = self.formatException(record.exc_info) if record.stack_info: error_obj["stack_trace"] = self.formatStack(record.stack_info) payload["error"] = error_obj # 兼容旧字段名 exception if record.exc_info and "error" not in payload: payload["exception"] = self.formatException(record.exc_info) return json.dumps(payload, ensure_ascii=False, separators=(",", ":")) # ==================== DateSizeRotatingFileHandler ==================== class DateSizeRotatingFileHandler(Handler): """按日期+大小滚动的文件日志处理器。 - 按天创建文件,文件名包含日期 - 单文件超过 max_bytes 时自动滚动(带序号后缀) - 自动清理超过 retention_days 的旧日志 """ 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) # ==================== StructuredLogger 封装 ==================== class StructuredLogger: """ 结构化日志记录器,提供符合方案文档 §4.2 的 5W1H 日志接口。 使用方式: logger = get_structured_logger('app.api.dataset') logger.info('数据集创建成功', datasetId='ds_001', costMs=23) """ 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, **fields: Any) -> None: self._log(logging.INFO, message, **fields) def warning(self, message: str, **fields: Any) -> None: self._log(logging.WARNING, message, **fields) def error(self, message: str, **fields: Any) -> None: self._log(logging.ERROR, message, **fields) def debug(self, message: str, **fields: Any) -> None: self._log(logging.DEBUG, message, **fields) def _log(self, level: int, message: str, **fields: Any) -> None: """统一日志记录方法,通过 extra 传递结构化字段。""" extra: dict[str, Any] = {} if self.module: extra["module"] = self.module # 脱敏 + 截断 for k, v in fields.items(): extra[k] = truncate_large_value(v) self.logger.log(level, message, extra=extra, stack_info=False) 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 set_user_context(user_id: str = "", client_ip: str = "") -> None: """设置当前请求的用户上下文(在鉴权后调用)。""" if user_id: user_id_var.set(user_id) if client_ip: client_ip_var.set(client_ip) # ==================== 请求日志中间件 ==================== 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] # 入口生成 traceId(优先使用前端传入的 X-Trace-Id) trace_id = request.headers.get("X-Trace-Id") or request.headers.get("X-Request-ID") or str(uuid4()) token = request_id_var.set(trace_id) started_at = time.perf_counter() # 提取客户端 IP client_ip = "-" if request.client: client_ip = request.client.host # 支持反向代理传递的真实 IP forwarded_for = request.headers.get("X-Forwarded-For", "") if forwarded_for: client_ip = forwarded_for.split(",")[0].strip() client_ip_var.set(client_ip) try: response = await call_next(request) elapsed_ms = (time.perf_counter() - started_at) * 1000 # 噪声路径降级为 DEBUG(健康检查等) noisy_paths = ("/health", "/system-info", "/compute/jobs/", "/model-eval/", "/model-compare/") log_method = logger.info if any(request.url.path.endswith(p) or p in request.url.path for p in noisy_paths) and response.status_code < 400: log_method = logger.debug if response.status_code >= 500: log_method = logger.error elif response.status_code >= 400: log_method = logger.warning # 结构化访问日志(中文 message,方便直接阅读) log_method( f"HTTP请求 {request.method} {request.url.path} → {response.status_code}(耗时{round(elapsed_ms, 2)}ms)", extra={ "request_method": request.method, "request_path": request.url.path, "status_code": response.status_code, "duration_ms": round(elapsed_ms, 2), "client_ip": client_ip, "user_agent": request.headers.get("User-Agent", "")[:200], }, ) # 5xx 系统错误自动写入操作日志 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-Trace-Id"] = trace_id response.headers["X-Request-ID"] = trace_id return response except Exception: elapsed_ms = (time.perf_counter() - started_at) * 1000 logger.error( f"HTTP请求异常 {request.method} {request.url.path}(耗时{round(elapsed_ms, 2)}ms)— 服务内部错误", extra={ "request_method": request.method, "request_path": request.url.path, "duration_ms": round(elapsed_ms, 2), "client_ip": client_ip, }, exc_info=True, ) # 未被捕获的异常,写入操作日志 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) # ==================== 配置函数 ==================== def configure_logging(settings: Settings | None = None) -> None: """ 生产级日志配置,符合方案文档 §二(分类分流)和 §六(性能安全)。 日志分类: - 业务日志 (app-biz): INFO+ 业务流程(保留 7 天) - 系统日志 (app-sys): 框架/中间件日志(保留 7 天) - 访问日志 (app-access): HTTP 请求日志(保留 15 天) - 错误日志 (app-error): ERROR 级别(保留 30 天) """ settings = settings or get_settings() root_logger = logging.getLogger() root_logger.handlers.clear() root_logger.setLevel(settings.log_level.upper()) # ---- Formatter ---- console_formatter = logging.Formatter( fmt=( "%(asctime)s | %(levelname)s | pid=%(process)d | %(threadName)s | " "traceId=%(traceId)s | %(name)s | %(pathname)s:%(lineno)d | %(message)s" ), datefmt="%Y-%m-%d %H:%M:%S", ) json_formatter = JsonLogFormatter() # ---- TraceIdFilter(全局注入 traceId/userId/host/app/env)---- trace_filter = TraceIdFilter() trace_filter.set_context( app=settings.app_name, env=settings.app_env, host=socket.gethostname(), ) # ---- 控制台 Handler ---- console_handler = logging.StreamHandler() console_handler.setFormatter(console_formatter) console_handler.addFilter(trace_filter) # ---- 业务日志文件 Handler (app-biz) ---- biz_file_handler = DateSizeRotatingFileHandler( log_dir=settings.log_dir, file_prefix="app-biz", max_bytes=settings.log_max_bytes, retention_days=7, ) biz_file_handler.setFormatter(json_formatter) biz_file_handler.addFilter(trace_filter) # ---- 访问日志文件 Handler (app-access) ---- access_file_handler = DateSizeRotatingFileHandler( log_dir=settings.log_dir, file_prefix="app-access", max_bytes=settings.log_max_bytes, retention_days=15, ) access_file_handler.setFormatter(json_formatter) access_file_handler.addFilter(trace_filter) # ---- 错误日志文件 Handler (app-error) ---- error_file_handler = DateSizeRotatingFileHandler( log_dir=settings.log_dir, file_prefix="app-error", max_bytes=settings.log_max_bytes, retention_days=30, ) error_file_handler.setLevel(logging.ERROR) error_file_handler.setFormatter(json_formatter) error_file_handler.addFilter(trace_filter) # ---- 注册 Handler ---- root_logger.addHandler(console_handler) root_logger.addHandler(biz_file_handler) root_logger.addHandler(access_file_handler) root_logger.addHandler(error_file_handler) # ---- 访问日志 Logger 独立路由到访问日志文件 ---- access_logger = logging.getLogger("app.access") access_logger.propagate = False # 不向 root 传播,避免重复写入业务日志 access_logger.addHandler(console_handler) access_logger.addHandler(access_file_handler) # 访问日志中的 ERROR 也要进错误日志 access_logger.addHandler(error_file_handler) # ---- 框架类 Logger 降级 ---- for logger_name in ("uvicorn", "uvicorn.error", "uvicorn.access"): lg = logging.getLogger(logger_name) lg.handlers.clear() lg.propagate = True # 框架类日志归入系统日志,生产环境设为 WARN logging.getLogger("uvicorn.access").setLevel(logging.WARNING) logging.getLogger("psycopg.pool").setLevel(logging.ERROR) logging.getLogger("httpx").setLevel(logging.WARNING) # ---- 兼容旧文件前缀(向后兼容)---- # 如果配置了旧的 log_file_prefix,也创建一个对应的 handler if settings.log_file_prefix and settings.log_file_prefix != "app-biz": legacy_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, ) legacy_file_handler.setFormatter(json_formatter) legacy_file_handler.addFilter(trace_filter) root_logger.addHandler(legacy_file_handler) # 旧错误日志前缀兼容 if settings.log_error_file_prefix and settings.log_error_file_prefix != "app-error": legacy_error_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, ) legacy_error_handler.setLevel(logging.ERROR) legacy_error_handler.setFormatter(json_formatter) legacy_error_handler.addFilter(trace_filter) root_logger.addHandler(legacy_error_handler)