diff --git a/backend/app/admin/service/login_log_service.py b/backend/app/admin/service/login_log_service.py index 45ed5662..811b397e 100644 --- a/backend/app/admin/service/login_log_service.py +++ b/backend/app/admin/service/login_log_service.py @@ -6,6 +6,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from backend.app.admin.crud.crud_login_log import login_log_dao from backend.app.admin.schema.login_log import CreateLoginLogParam, DeleteLoginLogParam +from backend.common.context import ctx from backend.common.log import log from backend.database.db import async_db_session @@ -53,14 +54,14 @@ class LoginLogService: user_uuid=user_uuid, username=username, status=status, - ip=request.state.ip, - country=request.state.country, - region=request.state.region, - city=request.state.city, - user_agent=request.state.user_agent, - browser=request.state.browser, - os=request.state.os, - device=request.state.device, + ip=ctx.ip, + country=ctx.country, + region=ctx.region, + city=ctx.city, + user_agent=ctx.user_agent, + browser=ctx.browser, + os=ctx.os, + device=ctx.device, msg=msg, login_time=login_time, ) diff --git a/backend/app/admin/service/user_service.py b/backend/app/admin/service/user_service.py index 1d1d8128..af02c310 100644 --- a/backend/app/admin/service/user_service.py +++ b/backend/app/admin/service/user_service.py @@ -14,6 +14,7 @@ from backend.app.admin.schema.user import ( ResetPasswordParam, UpdateUserParam, ) +from backend.common.context import ctx from backend.common.enums import UserPermissionType from backend.common.exception import errors from backend.common.response.response_code import CustomErrorCode @@ -249,12 +250,12 @@ class UserService: user = await user_dao.get(db, token_payload.id) if not user: raise errors.NotFoundError(msg='用户不存在') - captcha_code = await redis_client.get(f'{settings.EMAIL_CAPTCHA_REDIS_PREFIX}:{request.state.ip}') + captcha_code = await redis_client.get(f'{settings.EMAIL_CAPTCHA_REDIS_PREFIX}:{ctx.ip}') if not captcha_code: raise errors.RequestError(msg='验证码已失效,请重新获取') if captcha != captcha_code: raise errors.CustomError(error=CustomErrorCode.CAPTCHA_ERROR) - await redis_client.delete(f'{settings.EMAIL_CAPTCHA_REDIS_PREFIX}:{request.state.ip}') + await redis_client.delete(f'{settings.EMAIL_CAPTCHA_REDIS_PREFIX}:{ctx.ip}') count = await user_dao.update_email(db, token_payload.id, email) await redis_client.delete(f'{settings.JWT_USER_REDIS_PREFIX}:{user.id}') return count diff --git a/backend/common/context.py b/backend/common/context.py new file mode 100644 index 00000000..a807a533 --- /dev/null +++ b/backend/common/context.py @@ -0,0 +1,32 @@ +from datetime import datetime +from typing import Any, Protocol + +from starlette_context.ctx import _Context, context + + +class TypedContextProtocol(Protocol): + perf_time: float + start_time: datetime + + ip: str + country: str | None + region: str | None + city: str | None + + user_agent: str + os: str | None + browser: str | None + device: str | None + + permission: str | None + + +class TypedContext(TypedContextProtocol, _Context): + def __getattr__(self, name: str) -> Any: + return context.get(name) + + def __setattr__(self, name: str, value: Any) -> None: + context[name] = value + + +ctx = TypedContext() diff --git a/backend/common/exception/exception_handler.py b/backend/common/exception/exception_handler.py index c1cb62a0..872a8d72 100644 --- a/backend/common/exception/exception_handler.py +++ b/backend/common/exception/exception_handler.py @@ -4,6 +4,7 @@ from pydantic import ValidationError from starlette.exceptions import HTTPException from uvicorn.protocols.http.h11_impl import STATUS_PHRASES +from backend.common.context import ctx from backend.common.exception.errors import BaseExceptionError from backend.common.i18n import i18n, t from backend.common.response.response_code import CustomResponseCode, StandardResponseCode @@ -72,8 +73,8 @@ async def _validation_exception_handler(request: Request, exc: RequestValidation 'msg': msg, 'data': data, } - request.state.__request_validation_exception__ = content # 用于在中间件中获取异常信息 - content.update(trace_id=get_request_trace_id(request)) + ctx.__request_validation_exception__ = content # 用于在中间件中获取异常信息 + content.update(trace_id=get_request_trace_id()) return MsgSpecJSONResponse(status_code=StandardResponseCode.HTTP_422, content=content) @@ -96,8 +97,8 @@ def register_exception(app: FastAPI) -> None: else: res = response_base.fail(res=CustomResponseCode.HTTP_400) content = res.model_dump() - request.state.__request_http_exception__ = content - content.update(trace_id=get_request_trace_id(request)) + ctx.__request_http_exception__ = content + content.update(trace_id=get_request_trace_id()) return MsgSpecJSONResponse( status_code=_get_exception_code(exc.status_code), content=content, @@ -144,8 +145,8 @@ def register_exception(app: FastAPI) -> None: else: res = response_base.fail(res=CustomResponseCode.HTTP_500) content = res.model_dump() - request.state.__request_assertion_error__ = content - content.update(trace_id=get_request_trace_id(request)) + ctx.__request_assertion_error__ = content + content.update(trace_id=get_request_trace_id()) return MsgSpecJSONResponse( status_code=StandardResponseCode.HTTP_500, content=content, @@ -165,8 +166,8 @@ def register_exception(app: FastAPI) -> None: 'msg': str(exc.msg), 'data': exc.data or None, } - request.state.__request_custom_exception__ = content - content.update(trace_id=get_request_trace_id(request)) + ctx.__request_custom_exception__ = content + content.update(trace_id=get_request_trace_id()) return MsgSpecJSONResponse( status_code=_get_exception_code(exc.code), content=content, @@ -191,7 +192,7 @@ def register_exception(app: FastAPI) -> None: else: res = response_base.fail(res=CustomResponseCode.HTTP_500) content = res.model_dump() - content.update(trace_id=get_request_trace_id(request)) + content.update(trace_id=get_request_trace_id()) return MsgSpecJSONResponse( status_code=StandardResponseCode.HTTP_500, content=content, diff --git a/backend/common/log.py b/backend/common/log.py index cffd40f1..6fcfadbb 100644 --- a/backend/common/log.py +++ b/backend/common/log.py @@ -4,12 +4,13 @@ import os import re import sys -from asgi_correlation_id import correlation_id from loguru import logger +from backend.common.context import ctx from backend.core.conf import settings from backend.core.path_conf import LOG_DIR from backend.utils.timezone import timezone +from backend.utils.trace_id import get_request_trace_id class InterceptHandler(logging.Handler): @@ -75,11 +76,14 @@ def setup_logging() -> None: # 移除 loguru 默认处理器 logger.remove() - # correlation_id 过滤器 - # https://github.com/snok/asgi-correlation-id/issues/7 - def correlation_id_filter(record: logging.LogRecord) -> logging.LogRecord: - cid = correlation_id.get(settings.TRACE_ID_LOG_DEFAULT_VALUE) - record['correlation_id'] = cid[: settings.TRACE_ID_LOG_LENGTH] + # request_id 过滤器 + def request_id_filter(record: logging.LogRecord) -> logging.LogRecord: + if ctx.exists(): + rid = get_request_trace_id() + record['request_id'] = rid[: settings.TRACE_ID_LOG_LENGTH] + else: + record['request_id'] = settings.TRACE_ID_LOG_DEFAULT_VALUE + return record # 配置 loguru 处理器 @@ -89,7 +93,7 @@ def setup_logging() -> None: 'sink': sys.stdout, 'level': settings.LOG_STD_LEVEL, 'format': default_formatter, - 'filter': lambda record: correlation_id_filter(record), + 'filter': lambda record: request_id_filter(record), }, ], ) diff --git a/backend/common/security/permission.py b/backend/common/security/permission.py index 853faded..b5739881 100644 --- a/backend/common/security/permission.py +++ b/backend/common/security/permission.py @@ -3,6 +3,7 @@ from sqlalchemy import ColumnElement, and_, or_ from sqlalchemy.ext.asyncio import AsyncSession from backend.app.admin.crud.crud_data_scope import data_scope_dao +from backend.common.context import ctx from backend.common.enums import RoleDataRuleExpressionType, RoleDataRuleOperatorType from backend.common.exception import errors from backend.core.conf import settings @@ -38,7 +39,7 @@ class RequestPermission: if not isinstance(self.value, str): raise errors.ServerError # 附加权限标识到请求状态 - request.state.permission = self.value + ctx.permission = self.value async def filter_data_permission(db: AsyncSession, request: Request) -> ColumnElement[bool]: # noqa: C901 diff --git a/backend/common/security/rbac.py b/backend/common/security/rbac.py index 8d6e867b..ddd7f1fc 100644 --- a/backend/common/security/rbac.py +++ b/backend/common/security/rbac.py @@ -1,5 +1,6 @@ from fastapi import Depends, Request +from backend.common.context import ctx from backend.common.enums import MethodType, StatusType from backend.common.exception import errors from backend.common.log import log @@ -49,7 +50,7 @@ async def rbac_verify(request: Request, _token: str = DependsJwtAuth) -> None: # RBAC 鉴权 if settings.RBAC_ROLE_MENU_MODE: - path_auth_perm = getattr(request.state, 'permission', None) + path_auth_perm = ctx.permission # 没有菜单操作权限标识不校验 if not path_auth_perm: diff --git a/backend/core/conf.py b/backend/core/conf.py index 166a5d4d..8b14a793 100644 --- a/backend/core/conf.py +++ b/backend/core/conf.py @@ -150,7 +150,7 @@ class Settings(BaseSettings): # 日志 LOG_FORMAT: str = ( - '{time:YYYY-MM-DD HH:mm:ss.SSS} | {level: <8} | {correlation_id} | {message}' + '{time:YYYY-MM-DD HH:mm:ss.SSS} | {level: <8} | {request_id} | {message}' ) # 日志(控制台) diff --git a/backend/core/registrar.py b/backend/core/registrar.py index f20c2f92..1e20e500 100644 --- a/backend/core/registrar.py +++ b/backend/core/registrar.py @@ -6,7 +6,6 @@ from contextlib import asynccontextmanager import socketio -from asgi_correlation_id import CorrelationIdMiddleware from fastapi import Depends, FastAPI from fastapi_limiter import FastAPILimiter from fastapi_pagination import add_pagination @@ -14,10 +13,13 @@ from starlette.middleware.authentication import AuthenticationMiddleware from starlette.middleware.cors import CORSMiddleware from starlette.staticfiles import StaticFiles from starlette.types import ASGIApp +from starlette_context.middleware import ContextMiddleware +from starlette_context.plugins import RequestIdPlugin from backend import __version__ from backend.common.exception.exception_handler import register_exception from backend.common.log import set_custom_logfile, setup_logging +from backend.common.response.response_code import StandardResponseCode from backend.core.conf import settings from backend.core.path_conf import STATIC_DIR, UPLOAD_DIR from backend.database.db import create_tables @@ -154,8 +156,15 @@ def register_middleware(app: FastAPI) -> None: # Access log app.add_middleware(AccessMiddleware) - # Trace ID - app.add_middleware(CorrelationIdMiddleware, validator=False) + # ContextVar + app.add_middleware( + ContextMiddleware, + plugins=[RequestIdPlugin(validate=True)], + default_error_response=MsgSpecJSONResponse( + content={'code': StandardResponseCode.HTTP_400, 'msg': 'BAD_REQUEST', 'data': None}, + status_code=StandardResponseCode.HTTP_400, + ), + ) def register_router(app: FastAPI) -> None: diff --git a/backend/middleware/access_middleware.py b/backend/middleware/access_middleware.py index 2138ad79..c7507634 100644 --- a/backend/middleware/access_middleware.py +++ b/backend/middleware/access_middleware.py @@ -3,6 +3,7 @@ import time from fastapi import Request, Response from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint +from backend.common.context import ctx from backend.common.log import log from backend.utils.timezone import timezone @@ -24,21 +25,19 @@ class AccessMiddleware(BaseHTTPMiddleware): log.debug(f'--> 请求开始[{path}]') perf_time = time.perf_counter() - request.state.perf_time = perf_time + ctx.perf_time = perf_time start_time = timezone.now() - request.state.start_time = start_time + ctx.start_time = start_time response = await call_next(request) - elapsed = (time.perf_counter() - perf_time) * 1000 - if request.method != 'OPTIONS': log.debug('<-- 请求结束') log.info( f'{request.client.host: <15} | {request.method: <8} | {response.status_code: <6} | ' - f'{path} | {elapsed:.3f}ms', + f'{path} | {(time.perf_counter() - perf_time) * 1000:.3f}ms', ) return response diff --git a/backend/middleware/opera_log_middleware.py b/backend/middleware/opera_log_middleware.py index a1615ab3..690e9429 100644 --- a/backend/middleware/opera_log_middleware.py +++ b/backend/middleware/opera_log_middleware.py @@ -11,6 +11,7 @@ from starlette.requests import Request from backend.app.admin.schema.opera_log import CreateOperaLogParam from backend.app.admin.service.opera_log_service import opera_log_service +from backend.common.context import ctx from backend.common.enums import OperaLogCipherType, StatusType from backend.common.log import log from backend.common.queue import batch_dequeue @@ -49,21 +50,21 @@ class OperaLogMiddleware(BaseHTTPMiddleware): error = None try: response = await call_next(request) - elapsed = (time.perf_counter() - request.state.perf_time) * 1000 - for state in [ + elapsed = (time.perf_counter() - ctx.perf_time) * 1000 + for e in [ '__request_http_exception__', '__request_validation_exception__', '__request_assertion_error__', '__request_custom_exception__', ]: - exception = getattr(request.state, state, None) + exception = ctx.get(e) if exception: code = exception.get('code') msg = exception.get('msg') log.error(f'请求异常: {msg}') break except Exception as e: - elapsed = (time.perf_counter() - request.state.perf_time) * 1000 + elapsed = (time.perf_counter() - ctx.perf_time) * 1000 code = getattr(e, 'code', StandardResponseCode.HTTP_500) # 兼容 SQLAlchemy 异常用法 msg = getattr(e, 'msg', str(e)) # 不建议使用 traceback 模块获取错误信息,会暴漏代码信息 status = StatusType.disable @@ -82,30 +83,30 @@ class OperaLogMiddleware(BaseHTTPMiddleware): # 日志记录 log.debug(f'接口摘要:[{summary}]') - log.debug(f'请求地址:[{request.state.ip}]') + log.debug(f'请求地址:[{ctx.ip}]') log.debug(f'请求参数:{args}') # 日志创建 opera_log_in = CreateOperaLogParam( - trace_id=get_request_trace_id(request), + trace_id=get_request_trace_id(), username=username, method=method, title=summary, path=path, - ip=request.state.ip, - country=request.state.country, - region=request.state.region, - city=request.state.city, - user_agent=request.state.user_agent, - os=request.state.os, - browser=request.state.browser, - device=request.state.device, + ip=ctx.ip, + country=ctx.country, + region=ctx.region, + city=ctx.city, + user_agent=ctx.user_agent, + os=ctx.os, + browser=ctx.browser, + device=ctx.device, args=args, status=status, code=str(code), msg=msg, cost_time=elapsed, # 可能和日志存在微小差异(可忽略) - opera_time=request.state.start_time, + opera_time=ctx.start_time, ) await self.opera_log_queue.put(opera_log_in) diff --git a/backend/middleware/state_middleware.py b/backend/middleware/state_middleware.py index 5fae3cfe..852a01f8 100644 --- a/backend/middleware/state_middleware.py +++ b/backend/middleware/state_middleware.py @@ -1,6 +1,7 @@ from fastapi import Request, Response from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint +from backend.common.context import ctx from backend.utils.request_parse import parse_ip_info, parse_user_agent_info @@ -16,16 +17,16 @@ class StateMiddleware(BaseHTTPMiddleware): :return: """ ip_info = await parse_ip_info(request) - request.state.ip = ip_info.ip - request.state.country = ip_info.country - request.state.region = ip_info.region - request.state.city = ip_info.city + ctx.ip = ip_info.ip + ctx.country = ip_info.country + ctx.region = ip_info.region + ctx.city = ip_info.city ua_info = parse_user_agent_info(request) - request.state.user_agent = ua_info.user_agent - request.state.os = ua_info.os - request.state.browser = ua_info.browser - request.state.device = ua_info.device + ctx.user_agent = ua_info.user_agent + ctx.os = ua_info.os + ctx.browser = ua_info.browser + ctx.device = ua_info.device response = await call_next(request) diff --git a/backend/plugin/email/api/v1/email.py b/backend/plugin/email/api/v1/email.py index 963e34ab..e4ba2578 100644 --- a/backend/plugin/email/api/v1/email.py +++ b/backend/plugin/email/api/v1/email.py @@ -4,6 +4,7 @@ from typing import Annotated from fastapi import APIRouter, Body, Request +from backend.common.context import ctx from backend.common.response.response_schema import ResponseModel, response_base from backend.common.security.jwt import DependsJwtAuth from backend.core.conf import settings @@ -21,7 +22,7 @@ async def send_email_captcha( recipients: Annotated[str | list[str], Body(embed=True, description='邮件接收者')], ) -> ResponseModel: code = ''.join([str(random.randint(1, 9)) for _ in range(6)]) - ip = request.state.ip + ip = ctx.ip await redis_client.set( f'{settings.EMAIL_CAPTCHA_REDIS_PREFIX}:{ip}', code, diff --git a/backend/utils/trace_id.py b/backend/utils/trace_id.py index 721ad4df..6629c99d 100644 --- a/backend/utils/trace_id.py +++ b/backend/utils/trace_id.py @@ -1,13 +1,7 @@ -from fastapi import Request - +from backend.common.context import ctx from backend.core.conf import settings -def get_request_trace_id(request: Request) -> str: - """ - 从请求头中获取追踪 ID - - :param request: FastAPI 请求对象 - :return: - """ - return request.headers.get(settings.TRACE_ID_REQUEST_HEADER_KEY) or settings.TRACE_ID_LOG_DEFAULT_VALUE +def get_request_trace_id() -> str: + """从请求头中获取追踪 ID""" + return ctx.get(settings.TRACE_ID_REQUEST_HEADER_KEY, settings.TRACE_ID_LOG_DEFAULT_VALUE) diff --git a/pyproject.toml b/pyproject.toml index d97b415e..d9e93a77 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -13,7 +13,6 @@ requires-python = ">=3.10" dynamic = ['version'] dependencies = [ "alembic>=1.16.5", - "asgi-correlation-id>=4.3.4", "asgiref>=3.10.0", "asyncmy>=0.2.10", "asyncpg>=0.30.0", @@ -50,6 +49,7 @@ dependencies = [ "sqlalchemy-crud-plus>=1.12.1", "sqlalchemy[asyncio]>=2.0.44", "sqlparse>=0.5.3", + "starlette-context>=0.4.0", "user-agents>=2.2.0", ] diff --git a/requirements.txt b/requirements.txt index 08d70708..07b8bcac 100644 --- a/requirements.txt +++ b/requirements.txt @@ -12,8 +12,6 @@ anyio==4.11.0 # httpx # starlette # watchfiles -asgi-correlation-id==4.3.4 - # via fastapi-best-architecture asgiref==3.10.0 # via fastapi-best-architecture async-timeout==5.0.1 ; python_full_version < '3.11.3' @@ -164,7 +162,6 @@ nodeenv==1.9.1 # via pre-commit packaging==25.0 # via - # asgi-correlation-id # kombu # pytest pillow==11.3.0 @@ -270,8 +267,10 @@ sqlparse==0.5.3 # via fastapi-best-architecture starlette==0.48.0 # via - # asgi-correlation-id # fastapi + # starlette-context +starlette-context==0.4.0 + # via fastapi-best-architecture termcolor==3.1.0 # via pytest-sugar tomli==2.3.0 ; python_full_version < '3.11' diff --git a/uv.lock b/uv.lock index fde00b99..b2a14723 100644 --- a/uv.lock +++ b/uv.lock @@ -84,19 +84,6 @@ wheels = [ { url = "https://mirrors.aliyun.com/pypi/packages/15/b3/9b1a8074496371342ec1e796a96f99c82c945a339cd81a8e73de28b4cf9e/anyio-4.11.0-py3-none-any.whl", hash = "sha256:0287e96f4d26d4149305414d4e3bc32f0dcd0862365a4bddea19d7a1ec38c4fc" }, ] -[[package]] -name = "asgi-correlation-id" -version = "4.3.4" -source = { registry = "https://mirrors.aliyun.com/pypi/simple" } -dependencies = [ - { name = "packaging" }, - { name = "starlette" }, -] -sdist = { url = "https://mirrors.aliyun.com/pypi/packages/f4/ff/a6538245ac1eaa7733ec6740774e9d5add019e2c63caa29e758c16c0afdd/asgi_correlation_id-4.3.4.tar.gz", hash = "sha256:ea6bc310380373cb9f731dc2e8b2b6fb978a76afe33f7a2384f697b8d6cd811d" } -wheels = [ - { url = "https://mirrors.aliyun.com/pypi/packages/d9/ab/6936e2663c47a926e0659437b9333ad87d1ff49b1375d239026e0a268eba/asgi_correlation_id-4.3.4-py3-none-any.whl", hash = "sha256:36ce69b06c7d96b4acb89c7556a4c4f01a972463d3d49c675026cbbd08e9a0a2" }, -] - [[package]] name = "asgiref" version = "3.10.0" @@ -683,7 +670,6 @@ name = "fastapi-best-architecture" source = { editable = "." } dependencies = [ { name = "alembic" }, - { name = "asgi-correlation-id" }, { name = "asgiref" }, { name = "asyncmy" }, { name = "asyncpg" }, @@ -718,6 +704,7 @@ dependencies = [ { name = "sqlalchemy", extra = ["asyncio"] }, { name = "sqlalchemy-crud-plus" }, { name = "sqlparse" }, + { name = "starlette-context" }, { name = "user-agents" }, ] @@ -737,7 +724,6 @@ server = [ [package.metadata] requires-dist = [ { name = "alembic", specifier = ">=1.16.5" }, - { name = "asgi-correlation-id", specifier = ">=4.3.4" }, { name = "asgiref", specifier = ">=3.10.0" }, { name = "asyncmy", specifier = ">=0.2.10" }, { name = "asyncpg", specifier = ">=0.30.0" }, @@ -772,6 +758,7 @@ requires-dist = [ { name = "sqlalchemy", extras = ["asyncio"], specifier = ">=2.0.44" }, { name = "sqlalchemy-crud-plus", specifier = ">=1.12.1" }, { name = "sqlparse", specifier = ">=0.5.3" }, + { name = "starlette-context", specifier = ">=0.4.0" }, { name = "user-agents", specifier = ">=2.2.0" }, ] @@ -2586,6 +2573,18 @@ wheels = [ { url = "https://mirrors.aliyun.com/pypi/packages/be/72/2db2f49247d0a18b4f1bb9a5a39a0162869acf235f3a96418363947b3d46/starlette-0.48.0-py3-none-any.whl", hash = "sha256:0764ca97b097582558ecb498132ed0c7d942f233f365b86ba37770e026510659" }, ] +[[package]] +name = "starlette-context" +version = "0.4.0" +source = { registry = "https://mirrors.aliyun.com/pypi/simple" } +dependencies = [ + { name = "starlette" }, +] +sdist = { url = "https://mirrors.aliyun.com/pypi/packages/50/a2/42e29208fea21d03e6487d3ae4b6f1a5bd21cf762b8b7eef95041e8c08ff/starlette_context-0.4.0.tar.gz", hash = "sha256:3242417c9354c067a4ac5009aff762dc0b322074216f664825d5d127108553be" } +wheels = [ + { url = "https://mirrors.aliyun.com/pypi/packages/98/ec/90a4479bc3e307f3ff0b56165d70e1b01cf10882872823b27984bda375e2/starlette_context-0.4.0-py3-none-any.whl", hash = "sha256:dbcc11006587f901edd3d0a989a69a628fccf9d00c1ca3c28fab23ab88bd0093" }, +] + [[package]] name = "termcolor" version = "3.1.0"