feature:基本框架
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""基础设施工具层:统一响应、全局异常、分级日志、请求链路 trace_id。"""
|
||||
@@ -0,0 +1,89 @@
|
||||
"""业务异常与全局异常处理器:统一转成 {code, message, data, trace_id}。"""
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.exceptions import RequestValidationError
|
||||
from fastapi.responses import JSONResponse
|
||||
from starlette.exceptions import HTTPException as StarletteHTTPException
|
||||
|
||||
from utils.logger import get_logger
|
||||
from utils.response import Code, fail
|
||||
|
||||
logger = get_logger("api")
|
||||
|
||||
|
||||
class ApiError(Exception):
|
||||
"""业务异常基类。code 见 utils.response.Code。"""
|
||||
|
||||
def __init__(self, code: int, message: str, data=None):
|
||||
self.code = code
|
||||
self.message = message
|
||||
self.data = data
|
||||
super().__init__(message)
|
||||
|
||||
|
||||
class ParamError(ApiError):
|
||||
def __init__(self, message: str = "参数错误"):
|
||||
super().__init__(Code.PARAM_ERROR, message)
|
||||
|
||||
|
||||
class AuthError(ApiError):
|
||||
def __init__(self, message: str = "未登录或登录已过期"):
|
||||
super().__init__(Code.UNAUTHORIZED, message)
|
||||
|
||||
|
||||
class ForbiddenError(ApiError):
|
||||
def __init__(self, message: str = "无权限执行该操作"):
|
||||
super().__init__(Code.FORBIDDEN, message)
|
||||
|
||||
|
||||
class NotFoundError(ApiError):
|
||||
def __init__(self, message: str = "资源不存在"):
|
||||
super().__init__(Code.NOT_FOUND, message)
|
||||
|
||||
|
||||
class LLMFailError(ApiError):
|
||||
"""LLM 调用失败(重试耗尽,§1001)。"""
|
||||
|
||||
def __init__(self, message: str = "LLM 调用失败,请稍后再试"):
|
||||
super().__init__(Code.LLM_FAIL, message)
|
||||
|
||||
|
||||
class RiskTriggeredError(ApiError):
|
||||
"""交易被风控拦截(§1004)。"""
|
||||
|
||||
def __init__(self, message: str = "交易被风控拦截,请在订单列表查看处理进度", data=None):
|
||||
super().__init__(Code.RISK_TRIGGERED, message, data)
|
||||
|
||||
|
||||
class NotSuitableError(ApiError):
|
||||
"""适当性不匹配(§1005)。"""
|
||||
|
||||
def __init__(self, message: str = "客户风险等级与产品风险等级不匹配", data=None):
|
||||
super().__init__(Code.NOT_SUITABLE, message, data)
|
||||
|
||||
|
||||
def _http_status(code: int) -> int:
|
||||
"""业务码若为标准 HTTP 状态码则原样返回,自定义业务码(1001+)统一返回 200。"""
|
||||
return code if 400 <= code < 600 else 200
|
||||
|
||||
|
||||
def register_exception_handlers(app: FastAPI) -> None:
|
||||
@app.exception_handler(StarletteHTTPException)
|
||||
async def http_handler(request: Request, exc: StarletteHTTPException):
|
||||
return JSONResponse(status_code=exc.status_code, content=fail(exc.status_code, str(exc.detail)).model_dump())
|
||||
|
||||
@app.exception_handler(RequestValidationError)
|
||||
async def validation_handler(request: Request, exc: RequestValidationError):
|
||||
return JSONResponse(
|
||||
status_code=422,
|
||||
content=fail(Code.PARAM_ERROR, "参数校验失败", exc.errors()).model_dump(),
|
||||
)
|
||||
|
||||
@app.exception_handler(ApiError)
|
||||
async def api_error_handler(request: Request, exc: ApiError):
|
||||
status = _http_status(exc.code)
|
||||
return JSONResponse(status_code=status, content=fail(exc.code, exc.message, exc.data).model_dump())
|
||||
|
||||
@app.exception_handler(Exception)
|
||||
async def unhandled_handler(request: Request, exc: Exception):
|
||||
logger.exception("unhandled error on %s", request.url.path)
|
||||
return JSONResponse(status_code=500, content=fail(Code.SERVER_ERROR, "服务内部错误").model_dump())
|
||||
@@ -0,0 +1,65 @@
|
||||
"""分级日志(info.log / error.log + console),自动携带 trace_id,金融敏感字段脱敏。"""
|
||||
import logging
|
||||
import re
|
||||
import sys
|
||||
from logging.handlers import TimedRotatingFileHandler
|
||||
from pathlib import Path
|
||||
|
||||
from utils.request_id import get_request_id
|
||||
|
||||
_TRACE_FMT = "%(asctime)s %(levelname)-7s [%(trace_id)s] %(name)s: %(message)s"
|
||||
|
||||
# 手机号(11位)/ 身份证(18位,末位可 X)→ 日志脱敏
|
||||
_PHONE = re.compile(r"(?<!\d)1\d{10}(?!\d)")
|
||||
_IDCARD = re.compile(r"(?<!\d)\d{15}(\d\d[\dXx])(?!\d)|(?<!\d)\d{17}[\dXx](?!\d)")
|
||||
|
||||
_logging_configured = False
|
||||
|
||||
|
||||
class TraceIdFilter(logging.Filter):
|
||||
def filter(self, record: logging.LogRecord) -> bool:
|
||||
record.trace_id = get_request_id() or "-"
|
||||
return True
|
||||
|
||||
|
||||
class SensitiveMaskFilter(logging.Filter):
|
||||
"""身份证保留前4后4、手机号保留前3后4,其余用 * 掩码。"""
|
||||
|
||||
def filter(self, record: logging.LogRecord) -> bool:
|
||||
msg = record.getMessage()
|
||||
msg = _IDCARD.sub(lambda m: m.group()[:4] + "**********" + m.group()[-4:], msg)
|
||||
msg = _PHONE.sub(lambda m: m.group()[:3] + "****" + m.group()[-4:], msg)
|
||||
record.msg = msg
|
||||
record.args = ()
|
||||
return True
|
||||
|
||||
|
||||
def get_logger(name: str = "app") -> logging.Logger:
|
||||
return logging.getLogger(name)
|
||||
|
||||
|
||||
def setup_logging(log_dir: str = "./logs") -> None:
|
||||
"""幂等初始化:console + 按天滚动的 info.log / error.log。"""
|
||||
global _logging_configured
|
||||
if _logging_configured:
|
||||
return
|
||||
_logging_configured = True
|
||||
|
||||
root = logging.getLogger()
|
||||
root.setLevel(logging.INFO)
|
||||
formatter = logging.Formatter(_TRACE_FMT)
|
||||
Path(log_dir).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
handlers = [
|
||||
logging.StreamHandler(sys.stderr),
|
||||
TimedRotatingFileHandler(str(Path(log_dir) / "info.log"), when="midnight", backupCount=30, encoding="utf-8"),
|
||||
TimedRotatingFileHandler(str(Path(log_dir) / "error.log"), when="midnight", backupCount=30, encoding="utf-8"),
|
||||
]
|
||||
handlers[0].setLevel(logging.INFO)
|
||||
handlers[1].setLevel(logging.INFO)
|
||||
handlers[2].setLevel(logging.ERROR)
|
||||
for h in handlers:
|
||||
h.setFormatter(formatter)
|
||||
h.addFilter(TraceIdFilter())
|
||||
h.addFilter(SensitiveMaskFilter())
|
||||
root.addHandler(h)
|
||||
@@ -0,0 +1,32 @@
|
||||
"""每请求 trace_id:contextvar 存链路号,中间件生成 / 透传并回写响应头。"""
|
||||
import contextvars
|
||||
import uuid
|
||||
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
|
||||
_request_id: contextvars.ContextVar[str | None] = contextvars.ContextVar("request_id", default=None)
|
||||
|
||||
|
||||
def get_request_id() -> str | None:
|
||||
"""获取当前请求的 trace_id(请求上下文外返回 None)。"""
|
||||
return _request_id.get()
|
||||
|
||||
|
||||
def new_request_id() -> str:
|
||||
return uuid.uuid4().hex[:16]
|
||||
|
||||
|
||||
class RequestIdMiddleware(BaseHTTPMiddleware):
|
||||
"""透传客户端 X-Trace-Id(无则生成),供日志/响应头/审计链路关联。"""
|
||||
|
||||
async def dispatch(self, request: Request, call_next):
|
||||
rid = request.headers.get("X-Trace-Id") or new_request_id()
|
||||
token = _request_id.set(rid)
|
||||
try:
|
||||
response: Response = await call_next(request)
|
||||
response.headers["X-Trace-Id"] = rid
|
||||
return response
|
||||
finally:
|
||||
_request_id.reset(token)
|
||||
@@ -0,0 +1,37 @@
|
||||
"""统一响应格式:{code, message, data, trace_id},错误码与需求文档 §6.1 对齐。"""
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from utils.request_id import get_request_id
|
||||
|
||||
|
||||
class Code:
|
||||
OK = 200
|
||||
PARAM_ERROR = 400
|
||||
UNAUTHORIZED = 401
|
||||
FORBIDDEN = 403
|
||||
NOT_FOUND = 404
|
||||
SERVER_ERROR = 500
|
||||
LLM_FAIL = 1001 # LLM 调用失败
|
||||
KB_NO_RESULT = 1002 # 知识库检索无结果
|
||||
SQL_GEN_FAIL = 1003 # NL2SQL 生成失败
|
||||
RISK_TRIGGERED = 1004 # 风控规则触发 / 交易被拦截
|
||||
NOT_SUITABLE = 1005 # 适当性不匹配
|
||||
|
||||
|
||||
class ApiResponse(BaseModel):
|
||||
code: int = Code.OK
|
||||
message: str = "success"
|
||||
data: Any = None
|
||||
trace_id: str | None = None
|
||||
|
||||
|
||||
def success(data: Any = None, message: str = "success") -> ApiResponse:
|
||||
"""成功响应。"""
|
||||
return ApiResponse(code=Code.OK, message=message, data=data, trace_id=get_request_id())
|
||||
|
||||
|
||||
def fail(code: int, message: str, data: Any = None) -> ApiResponse:
|
||||
"""失败响应(业务错误码见 Code)。"""
|
||||
return ApiResponse(code=code, message=message, data=data, trace_id=get_request_id())
|
||||
Reference in New Issue
Block a user