middleware.py 3.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899
  1. """请求日志中间件 — 记录最近 MAX_ENTRIES 条请求到内存"""
  2. import time
  3. import asyncio
  4. from collections import deque
  5. from typing import Optional
  6. from fastapi import Request, Response
  7. import logging
  8. logger = logging.getLogger(__name__)
  9. MAX_ENTRIES = 500
  10. _request_log: deque = deque(maxlen=MAX_ENTRIES)
  11. def get_request_log() -> list:
  12. """返回请求日志副本(最近在前)"""
  13. return list(_request_log)[::-1]
  14. def clear_request_log() -> int:
  15. """清空日志,返回清空条数"""
  16. n = len(_request_log)
  17. _request_log.clear()
  18. return n
  19. async def request_log_middleware(request: Request, call_next) -> Response:
  20. start_time = time.perf_counter()
  21. client_ip = request.client.host if request.client else "-"
  22. # 读取请求体(不重复消费)
  23. body = None
  24. if request.method in ("POST", "PUT", "PATCH"):
  25. try:
  26. raw = await request.body()
  27. if raw:
  28. try:
  29. body = raw.decode("utf-8")
  30. if len(body) > 2000:
  31. body = body[:2000] + "...(truncated)"
  32. except Exception:
  33. body = f"<binary {len(raw)} bytes>"
  34. except Exception:
  35. pass
  36. try:
  37. response = await call_next(request)
  38. except Exception as exc:
  39. elapsed = time.perf_counter() - start_time
  40. status_code = 500
  41. response_body = ""
  42. logger.error("REQUEST EXCEPTION: %s %s %.2fs error=%r",
  43. request.method, request.url.path, elapsed, exc)
  44. _request_log.appendleft({
  45. "ts": time.strftime("%H:%M:%S"),
  46. "method": request.method,
  47. "path": request.url.path,
  48. "status": 500,
  49. "duration_ms": round(elapsed * 1000),
  50. "client_ip": client_ip,
  51. "body": body,
  52. "error": str(exc)[:500],
  53. })
  54. # 构造错误响应
  55. from fastapi.responses import JSONResponse
  56. return JSONResponse(status_code=500, content={"detail": str(exc)})
  57. elapsed = time.perf_counter() - start_time
  58. status_code = response.status_code
  59. # 读取响应体(不可逆,只取内容长度)
  60. response_body = ""
  61. if hasattr(response, "body") and response.body:
  62. try:
  63. response_body = response.body.decode("utf-8")[:1000]
  64. except Exception:
  65. response_body = f"<binary {len(response.body)} bytes>"
  66. # 写入日志
  67. entry = {
  68. "ts": time.strftime("%H:%M:%S"),
  69. "method": request.method,
  70. "path": request.url.path,
  71. "status": status_code,
  72. "duration_ms": round(elapsed * 1000),
  73. "client_ip": client_ip,
  74. "body": body,
  75. "response": response_body,
  76. }
  77. _request_log.appendleft(entry)
  78. # 超5秒标记 warning
  79. if elapsed > 5:
  80. logger.warning("SLOW_REQUEST: %s %s %.2fs", request.method, request.url.path, elapsed)
  81. else:
  82. logger.debug("REQUEST: %s %s %.2fs", request.method, request.url.path, elapsed)
  83. response.headers["X-Response-Time"] = f"{elapsed:.3f}s"
  84. return response