""" 性能监控中间件 """ import logging import time from typing import Callable from fastapi import Request from starlette.middleware.base import BaseHTTPMiddleware logger = logging.getLogger(__name__) class PerformanceMonitoringMiddleware(BaseHTTPMiddleware): """性能监控中间件""" def __init__(self, app, slow_request_threshold: float = 1.0): super().__init__(app) self.slow_request_threshold = slow_request_threshold # 慢请求阈值(秒) async def dispatch(self, request: Request, call_next: Callable): # 记录请求开始时间 start_time = time.time() # 生成请求 ID request_id = self._generate_request_id() request.state.request_id = request_id # 处理请求 try: response = await call_next(request) # 计算处理时间 process_time = time.time() - start_time # 添加响应头 response.headers["X-Request-ID"] = request_id response.headers["X-Process-Time"] = f"{process_time:.3f}" # 记录慢请求 if process_time > self.slow_request_threshold: logger.warning( f"Slow request detected: {request.method} {request.url.path} " f"took {process_time:.3f}s (threshold: {self.slow_request_threshold}s) " f"[request_id={request_id}]" ) # 记录请求日志 logger.info( f"{request.method} {request.url.path} " f"status={response.status_code} time={process_time:.3f}s " f"[request_id={request_id}]" ) return response except Exception as e: process_time = time.time() - start_time logger.error( f"Request failed: {request.method} {request.url.path} " f"error={str(e)} time={process_time:.3f}s " f"[request_id={request_id}]", exc_info=True, ) raise def _generate_request_id(self) -> str: """生成请求 ID""" import uuid return str(uuid.uuid4()) class DatabaseQueryLogger: """数据库查询日志记录器""" def __init__(self): self.queries = [] self.total_time = 0 def log_query(self, query: str, params: tuple, duration: float): """记录查询""" self.queries.append( { "query": query, "params": params, "duration": duration, } ) self.total_time += duration # 记录慢查询(超过 100ms) if duration > 0.1: logger.warning(f"Slow query detected: {query[:100]}... " f"took {duration:.3f}s with params {params}") def get_stats(self): """获取统计信息""" return { "total_queries": len(self.queries), "total_time": self.total_time, "avg_time": self.total_time / len(self.queries) if self.queries else 0, "slow_queries": len([q for q in self.queries if q["duration"] > 0.1]), }