32 lines
1.1 KiB
Python
32 lines
1.1 KiB
Python
"""每请求 trace_id:contextvar 存链路号,中间件生成 / 透传并回写响应头。"""
|
||||
|
|
import contextvars
|
|||
|
|
import uuid
|
|||
|
|
|
|||
|
|
from starlette.middleware.base import BaseHTTPMiddleware
|
|||
|
|
from starlette.requests import Request
|
|||
|
|
from starlette.responses import Response
|
|||
|
|
|
|||
|
|
_request_id: contextvars.ContextVar[str | None] = contextvars.ContextVar("request_id", default=None)
|
|||
|
|
|
|||
|
|
|
|||
|
|
def get_request_id() -> str | None:
|
|||
|
|
"""获取当前请求的 trace_id(请求上下文外返回 None)。"""
|
|||
|
|
return _request_id.get()
|
|||
|
|
|
|||
|
|
|
|||
|
|
def new_request_id() -> str:
|
|||
|
|
return uuid.uuid4().hex[:16]
|
|||
|
|
|
|||
|
|
|
|||
|
|
class RequestIdMiddleware(BaseHTTPMiddleware):
|
|||
|
|
"""透传客户端 X-Trace-Id(无则生成),供日志/响应头/审计链路关联。"""
|
|||
|
|
|
|||
|
|
async def dispatch(self, request: Request, call_next):
|
|||
|
|
rid = request.headers.get("X-Trace-Id") or new_request_id()
|
|||
|
|
token = _request_id.set(rid)
|
|||
|
|
try:
|
|||
|
|
response: Response = await call_next(request)
|
|||
|
|
response.headers["X-Trace-Id"] = rid
|
|||
|
|
return response
|
|||
|
|
finally:
|
|||
|
|
_request_id.reset(token)
|