fastapi: 给每个请求增加请求id

一,代码

1,middleware

"""request_id 中间件。

优先沿用调用方传入的 ``X-Request-Id``,否则生成一个;把它绑进 contextvars
供日志使用,并在响应头回写。
"""

from collections.abc import Awaitable, Callable
from typing import Any
from uuid import uuid4

from starlette.datastructures import Headers, MutableHeaders
from starlette.types import ASGIApp, Message, Receive, Scope, Send

from app.core.context import bind_context


class RequestContextMiddleware:
    def __init__(self, app: ASGIApp) -> None:
        self.app = app

    async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
        if scope["type"] != "http":
            await self.app(scope, receive, send)
            return

        request_id = Headers(scope=scope).get("x-request-id") or uuid4().hex
        # 额外落到 scope 上:最外层 ServerErrorMiddleware 的 500 处理器运行在
        # 本中间件之外,此时 contextvars 已被 reset,只能靠 scope 把 request_id 传出去。
        scope.setdefault("state", {})["request_id"] = request_id

        async def send_with_request_id(message: Message) -> None:
            if message["type"] == "http.response.start":
                MutableHeaders(scope=message)["X-Request-Id"] = request_id
            await send(message)

        with bind_context(request_id=request_id):
            await self.app(scope, receive, send_with_request_id)


__all__: list[str] = ["RequestContextMiddleware"]

# 供类型检查器识别 Callable/Awaitable 的导入用途
_type_hints = (Callable, Awaitable, Any)

2, 用到的类

"""请求上下文:request_id / user_id / channel 的 contextvars 存取。

这些值由中间件与认证依赖写入,日志 patcher 读取后附加到每条日志上,
从而把一次请求涉及的所有日志串成一条链路。
"""

from collections.abc import Iterator
from contextlib import contextmanager
from contextvars import ContextVar, Token
from typing import Any

_request_id: ContextVar[str] = ContextVar("request_id", default="-")
_user_id: ContextVar[int | None] = ContextVar("user_id", default=None)
_channel: ContextVar[str] = ContextVar("channel", default="-")


def current_request_id() -> str:

    return _request_id.get()


def current_user_id() -> int | None:
    return _user_id.get()


def current_channel() -> str:
    return _channel.get()


@contextmanager
def bind_context(
    request_id: str | None = None,
    user_id: int | None = None,
    channel: str | None = None,
) -> Iterator[None]:
    """在 with 块内绑定上下文,退出时自动还原。"""
    tokens: list[tuple[ContextVar[Any], Token[Any]]] = []
    if request_id is not None:
        tokens.append((_request_id, _request_id.set(request_id)))
    if user_id is not None:
        tokens.append((_user_id, _user_id.set(user_id)))
    if channel is not None:
        tokens.append((_channel, _channel.set(channel)))
    try:
        yield
    finally:
        for var, token in reversed(tokens):
            var.reset(token)


def set_current_user(user_id: int | None, channel: str | None = None) -> None:
    """认证依赖拿到用户后回填,供后续日志与审计使用。"""
    _user_id.set(user_id)
    if channel is not None:
        _channel.set(channel)


def reset_context() -> None:
    """测试与清理用:把上下文恢复到默认值。"""
    print("context清理:reset_context:")
    _request_id.set("-")
    _user_id.set(None)
    _channel.set("-")

 

二,测试效果:

image

posted @ 2026-09-28 20:00  刘宏缔的架构森林  阅读(2)  评论(0)    收藏  举报