From 1f95a776f0b8a544494fc03c68e1fe920dd1f1ed Mon Sep 17 00:00:00 2001 From: Wu Clan Date: Mon, 9 Sep 2024 15:46:13 +0800 Subject: [PATCH] Add trace ID to exception handlers (#411) --- backend/common/exception/errors.py | 2 +- backend/common/exception/exception_handler.py | 71 +++++++++++-------- backend/middleware/opera_log_middleware.py | 8 +-- backend/utils/request_parse.py | 3 +- backend/utils/trace_id.py | 9 +++ 5 files changed, 58 insertions(+), 35 deletions(-) create mode 100644 backend/utils/trace_id.py diff --git a/backend/common/exception/errors.py b/backend/common/exception/errors.py index c622d606..a763a646 100644 --- a/backend/common/exception/errors.py +++ b/backend/common/exception/errors.py @@ -4,7 +4,7 @@ 全局业务异常类 业务代码执行异常时,可以使用 raise xxxError 触发内部错误,它尽可能实现带有后台任务的异常,但它不适用于**自定义响应状态码** -如果要求使用**自定义响应状态码**,则可以通过 return await response_base.fail(res=CustomResponseCode.xxx) 直接返回 +如果要求使用**自定义响应状态码**,则可以通过 return response_base.fail(res=CustomResponseCode.xxx) 直接返回 """ # noqa: E501 from typing import Any diff --git a/backend/common/exception/exception_handler.py b/backend/common/exception/exception_handler.py index 5749a98f..0db40994 100644 --- a/backend/common/exception/exception_handler.py +++ b/backend/common/exception/exception_handler.py @@ -18,6 +18,7 @@ from backend.common.schema import ( ) from backend.core.conf import settings from backend.utils.serializers import MsgSpecJSONResponse +from backend.utils.trace_id import get_request_trace_id def _get_exception_code(status_code: int): @@ -70,13 +71,14 @@ async def _validation_exception_handler(request: Request, e: RequestValidationEr error_input = error.get('input') field = str(error.get('loc')[-1]) error_msg = error.get('msg') - message = f'{error_msg}{field},输入:{error_input}' if settings.ENVIRONMENT == 'dev' else error_msg + message = f'{field} {error_msg},输入:{error_input}' if settings.ENVIRONMENT == 'dev' else error_msg msg = f'请求参数非法: {message}' data = {'errors': errors} if settings.ENVIRONMENT == 'dev' else None content = { 'code': StandardResponseCode.HTTP_422, 'msg': msg, 'data': data, + 'trace_id': get_request_trace_id(request), } request.state.__request_validation_exception__ = content # 用于在中间件中获取异常信息 return MsgSpecJSONResponse(status_code=422, content=content) @@ -104,7 +106,7 @@ def register_exception(app: FastAPI): request.state.__request_http_exception__ = content # 用于在中间件中获取异常信息 return MsgSpecJSONResponse( status_code=_get_exception_code(exc.status_code), - content=content, + content=content.update(trace_id=get_request_trace_id(request)), headers=exc.headers, ) @@ -145,6 +147,7 @@ def register_exception(app: FastAPI): 'code': StandardResponseCode.HTTP_500, 'msg': CUSTOM_USAGE_ERROR_MESSAGES.get(exc.code), 'data': None, + 'trace_id': get_request_trace_id(request), }, ) @@ -168,7 +171,27 @@ def register_exception(app: FastAPI): content = res.model_dump() return MsgSpecJSONResponse( status_code=StandardResponseCode.HTTP_500, - content=content, + content=content.update(trace_id=get_request_trace_id(request)), + ) + + @app.exception_handler(BaseExceptionMixin) + async def custom_exception_handler(request: Request, exc: BaseExceptionMixin): + """ + 全局异常处理 + + :param request: + :param exc: + :return: + """ + return MsgSpecJSONResponse( + status_code=_get_exception_code(exc.code), + content={ + 'code': exc.code, + 'msg': str(exc.msg), + 'data': exc.data if exc.data else None, + 'trace_id': get_request_trace_id(request), + }, + background=exc.background, ) @app.exception_handler(Exception) @@ -180,31 +203,23 @@ def register_exception(app: FastAPI): :param exc: :return: """ - if isinstance(exc, BaseExceptionMixin): - return MsgSpecJSONResponse( - status_code=_get_exception_code(exc.code), - content={ - 'code': exc.code, - 'msg': str(exc.msg), - 'data': exc.data if exc.data else None, - }, - background=exc.background, - ) - else: - import traceback + import traceback - log.error(f'未知异常: {exc}') - log.error(traceback.format_exc()) - if settings.ENVIRONMENT == 'dev': - content = { - 'code': StandardResponseCode.HTTP_500, - 'msg': str(exc), - 'data': None, - } - else: - res = response_base.fail(res=CustomResponseCode.HTTP_500) - content = res.model_dump() - return MsgSpecJSONResponse(status_code=StandardResponseCode.HTTP_500, content=content) + log.error(f'未知异常: {exc}') + log.error(traceback.format_exc()) + if settings.ENVIRONMENT == 'dev': + content = { + 'code': StandardResponseCode.HTTP_500, + 'msg': str(exc), + 'data': None, + } + else: + res = response_base.fail(res=CustomResponseCode.HTTP_500) + content = res.model_dump() + return MsgSpecJSONResponse( + status_code=StandardResponseCode.HTTP_500, + content=content.update(trace_id=get_request_trace_id(request)), + ) if settings.MIDDLEWARE_CORS: @@ -238,7 +253,7 @@ def register_exception(app: FastAPI): content = res.model_dump() response = MsgSpecJSONResponse( status_code=exc.code if isinstance(exc, BaseExceptionMixin) else StandardResponseCode.HTTP_500, - content=content, + content=content.update(trace_id=get_request_trace_id(request)), background=exc.background if isinstance(exc, BaseExceptionMixin) else None, ) origin = request.headers.get('origin') diff --git a/backend/middleware/opera_log_middleware.py b/backend/middleware/opera_log_middleware.py index 7b314143..4b30d186 100644 --- a/backend/middleware/opera_log_middleware.py +++ b/backend/middleware/opera_log_middleware.py @@ -17,6 +17,7 @@ from backend.core.conf import settings from backend.utils.encrypt import AESCipher, ItsDCipher, Md5Cipher from backend.utils.request_parse import parse_ip_info, parse_user_agent_info from backend.utils.timezone import timezone +from backend.utils.trace_id import get_request_trace_id class OperaLogMiddleware(BaseHTTPMiddleware): @@ -62,7 +63,7 @@ class OperaLogMiddleware(BaseHTTPMiddleware): # 日志创建 opera_log_in = CreateOperaLogParam( - trace_id=request.headers.get(settings.TRACE_ID_REQUEST_HEADER_KEY) or '-', + trace_id=get_request_trace_id(request), username=username, method=method, title=summary, @@ -100,9 +101,9 @@ class OperaLogMiddleware(BaseHTTPMiddleware): response = None try: response = await call_next(request) + code, msg = self.validation_exception_handler(request, code, msg) except Exception as e: log.exception(e) - code, msg = await self.request_exception_handler(request, code, msg) # code 处理包含 SQLAlchemy 和 Pydantic code = getattr(e, 'code', None) or code msg = getattr(e, 'msg', None) or msg @@ -112,8 +113,7 @@ class OperaLogMiddleware(BaseHTTPMiddleware): return RequestCallNext(code=str(code), msg=msg, status=status, err=err, response=response) @staticmethod - @sync_to_async - def request_exception_handler(request: Request, code: int, msg: str) -> tuple[str, str]: + def validation_exception_handler(request: Request, code: int, msg: str) -> tuple[str, str]: """请求异常处理器""" try: http_exception = request.state.__request_http_exception__ diff --git a/backend/utils/request_parse.py b/backend/utils/request_parse.py index 35540ccd..34c8443a 100644 --- a/backend/utils/request_parse.py +++ b/backend/utils/request_parse.py @@ -14,7 +14,6 @@ from backend.core.path_conf import IP2REGION_XDB from backend.database.db_redis import redis_client -@sync_to_async def get_request_ip(request: Request) -> str: """获取请求的 ip 地址""" real = request.headers.get('X-Real-IP') @@ -78,7 +77,7 @@ def get_location_offline(ip: str) -> dict | None: async def parse_ip_info(request: Request) -> IpInfo: country, region, city = None, None, None - ip = await get_request_ip(request) + ip = get_request_ip(request) location = await redis_client.get(f'{settings.IP_LOCATION_REDIS_PREFIX}:{ip}') if location: country, region, city = location.split(' ') diff --git a/backend/utils/trace_id.py b/backend/utils/trace_id.py new file mode 100644 index 00000000..80d2293c --- /dev/null +++ b/backend/utils/trace_id.py @@ -0,0 +1,9 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +from fastapi import Request + +from backend.core.conf import settings + + +def get_request_trace_id(request: Request) -> str: + return request.headers.get(settings.TRACE_ID_REQUEST_HEADER_KEY) or '-'