38 lines
1.2 KiB
Python
38 lines
1.2 KiB
Python
|
|
"""TraceId 中间件:为每个请求生成/注入 trace_id。
|
|||
|
|
|
|||
|
|
优先从 X-Request-Id header 读取(支持上游链路传递),
|
|||
|
|
否则生成 uuid4。写入 ContextVar 供下游使用。
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
from __future__ import annotations
|
|||
|
|
|
|||
|
|
from starlette.middleware.base import BaseHTTPMiddleware
|
|||
|
|
from starlette.requests import Request
|
|||
|
|
from starlette.responses import Response
|
|||
|
|
from yuxi.utils.trace_context import (
|
|||
|
|
generate_trace_id,
|
|||
|
|
reset_trace_id,
|
|||
|
|
set_trace_id,
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
_TRACE_ID_HEADER = "X-Request-Id"
|
|||
|
|
|
|||
|
|
|
|||
|
|
class TraceIdMiddleware(BaseHTTPMiddleware):
|
|||
|
|
"""为每个请求注入 trace_id 的中间件。
|
|||
|
|
|
|||
|
|
注册时应最后注册(LIFO 顺序中最先执行),确保所有后续中间件、
|
|||
|
|
异常处理器都能从 ContextVar 读取 trace_id。
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
async def dispatch(self, request: Request, call_next):
|
|||
|
|
trace_id = request.headers.get(_TRACE_ID_HEADER) or generate_trace_id()
|
|||
|
|
set_trace_id(trace_id)
|
|||
|
|
try:
|
|||
|
|
response: Response = await call_next(request)
|
|||
|
|
# 回写 trace_id 到响应头,前端可关联
|
|||
|
|
response.headers[_TRACE_ID_HEADER] = trace_id
|
|||
|
|
return response
|
|||
|
|
finally:
|
|||
|
|
reset_trace_id()
|