Update request state usage to context variable (#853)

* Update request state usage to context variable

* Update context to custom ctx

* Fix the exception interception in opera log

* restore elapsed
This commit is contained in:
Wu Clan
2025-10-16 10:23:39 +08:00
committed by GitHub
parent dd08775c85
commit 6a065797f5
17 changed files with 135 additions and 91 deletions
@@ -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,
)
+3 -2
View File
@@ -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
+32
View File
@@ -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()
+10 -9
View File
@@ -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,
+11 -7
View File
@@ -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),
},
],
)
+2 -1
View File
@@ -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
+2 -1
View File
@@ -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:
+1 -1
View File
@@ -150,7 +150,7 @@ class Settings(BaseSettings):
# 日志
LOG_FORMAT: str = (
'<green>{time:YYYY-MM-DD HH:mm:ss.SSS}</> | <lvl>{level: <8}</> | <cyan>{correlation_id}</> | <lvl>{message}</>'
'<green>{time:YYYY-MM-DD HH:mm:ss.SSS}</> | <lvl>{level: <8}</> | <cyan>{request_id}</> | <lvl>{message}</>'
)
# 日志(控制台)
+12 -3
View File
@@ -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:
+4 -5
View File
@@ -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
+16 -15
View File
@@ -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)
+9 -8
View File
@@ -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)
+2 -1
View File
@@ -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,
+4 -10
View File
@@ -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)
+1 -1
View File
@@ -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",
]
+3 -4
View File
@@ -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'
Generated
+14 -15
View File
@@ -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"