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) |