Files
track/backend/app/core/middleware.py

88 lines
3.4 KiB
Python
Raw Normal View History

"""请求上下文中间件 — request_id 生成/透传 + 结构化访问日志"""
from __future__ import annotations
import logging
import time
import uuid
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.requests import Request
from app.core.logging import request_id_var, user_var
access_log = logging.getLogger("track.access")
# 探针被高频轮询,降级为 DEBUG 避免把有价值的信息淹掉
_QUIET_PATHS = frozenset({"/health", "/health/live", "/health/ready"})
class RequestContextMiddleware(BaseHTTPMiddleware):
"""为每个请求建立可追踪上下文。
- request_id优先沿用上游网关传来的 X-Request-ID实现全链路追踪
没有就生成一个响应头回写该 ID前端报错时可直接带上
运维拿 ID 就能在日志里精确定位到这一次请求
- 访问日志method / path / status / duration_ms / client / user
"""
async def dispatch(self, request: Request, call_next):
request_id = request.headers.get("X-Request-ID") or uuid.uuid4().hex
# 同时写入 request.state它由 ASGI scope 承载,作用域比 contextvar 更长。
# FastAPI 把 Exception 处理器交给 ServerErrorMiddleware位于本中间件外层
# 异常传播到那里时 contextvar 已在 finally 中被重置,只有 state 还留着 ID。
request.state.request_id = request_id
rid_token = request_id_var.set(request_id)
user_token = user_var.set(None)
started = time.perf_counter()
logged = False
status_code = 500
try:
response = await call_next(request)
status_code = response.status_code
response.headers["X-Request-ID"] = request_id
self._log_access(request, status_code, started)
logged = True
return response
finally:
# 异常路径也要留下访问记录,否则接口 500 时日志里反而没有痕迹
if not logged:
self._log_access(request, status_code, started)
request_id_var.reset(rid_token)
user_var.reset(user_token)
def _log_access(self, request: Request, status_code: int, started: float) -> None:
path = request.url.path
duration_ms = round((time.perf_counter() - started) * 1000, 1)
# user 必须从 request.state 取:本中间件在独立 task 中执行,路由内
# 写入的 contextvar 不会回流到这里(详见 get_current_user 的说明)。
user = getattr(request.state, "audit_user", None) or user_var.get()
if status_code >= 500:
level = logging.ERROR
elif status_code >= 400:
level = logging.WARNING
elif path in _QUIET_PATHS:
level = logging.DEBUG
else:
level = logging.INFO
access_log.log(
level,
"%s %s -> %s (%.1fms)",
request.method,
path,
status_code,
duration_ms,
extra={
"extra_fields": {
"method": request.method,
"path": path,
"status": status_code,
"duration_ms": duration_ms,
"client": request.client.host if request.client else None,
"user": user,
}
},
)