68 lines
3.2 KiB
Python
68 lines
3.2 KiB
Python
"""统一 API 响应外壳(Wave 0 ok/fail + 风控模块 error_body/register_error_handlers)。"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from typing import Any
|
|
|
|
from fastapi import FastAPI, Request
|
|
from fastapi.exceptions import RequestValidationError
|
|
from fastapi.responses import JSONResponse
|
|
from starlette.exceptions import HTTPException as StarletteHTTPException
|
|
|
|
from app.model.schemas import ApiResponse
|
|
from app.utils.exceptions import ApiError, PermissionDenied
|
|
from app.utils.trace import current_request_id, current_trace, new_request_id, new_trace
|
|
|
|
_HTTP_ERROR_CODES = {404: "NOT_FOUND", 405: "METHOD_NOT_ALLOWED"}
|
|
|
|
|
|
def ok(data: Any, trace_id: str, message: str = "ok") -> ApiResponse:
|
|
return ApiResponse(code=0, message=message, data=data, trace_id=trace_id)
|
|
|
|
|
|
def fail(code: int, message: str, trace_id: str, *, data: Any = None) -> ApiResponse:
|
|
return ApiResponse(code=code, message=message, data=data, trace_id=trace_id)
|
|
|
|
|
|
def error_body(error_code: str, message: str) -> dict[str, Any]:
|
|
"""手册 §10 错误体(trace/request 缺失时兜底生成,保证响应可归因)。"""
|
|
tid = current_trace() or new_trace()
|
|
rid = current_request_id() or new_request_id()
|
|
return {"error_code": error_code, "message": message, "trace_id": tid, "request_id": rid}
|
|
|
|
|
|
def register_error_handlers(app: FastAPI) -> None:
|
|
@app.exception_handler(ApiError)
|
|
async def _api_error_handler(request: Request, exc: ApiError) -> JSONResponse:
|
|
body = error_body(exc.error_code, exc.message)
|
|
# T-9:合入结构性附加字段(架构 §8.3 契约)。
|
|
# 仅 convert 的错误子类(`ConvertError`)带 `extra` 属性,既有 ApiError
|
|
# 无此属性 → `getattr` 取 None → 既有错误体**逐字节不变**。
|
|
extra = getattr(exc, "extra", None)
|
|
if extra:
|
|
body.update(extra)
|
|
return JSONResponse(status_code=exc.status_code, content=body)
|
|
|
|
@app.exception_handler(PermissionDenied)
|
|
async def _permission_denied_handler(request: Request, exc: PermissionDenied) -> JSONResponse:
|
|
return JSONResponse(status_code=403, content=error_body(exc.code, str(exc)))
|
|
|
|
@app.exception_handler(StarletteHTTPException)
|
|
async def _http_exception_handler(request: Request, exc: StarletteHTTPException) -> JSONResponse:
|
|
code = _HTTP_ERROR_CODES.get(exc.status_code, "HTTP_ERROR")
|
|
return JSONResponse(status_code=exc.status_code, content=error_body(code, str(exc.detail)))
|
|
|
|
@app.exception_handler(RequestValidationError)
|
|
async def _validation_error_handler(request: Request, exc: RequestValidationError) -> JSONResponse:
|
|
return JSONResponse(
|
|
status_code=422,
|
|
content=error_body("REQUEST_VALIDATION_FAILED", "request validation failed"),
|
|
)
|
|
|
|
|
|
def trace_headers(request: Request) -> dict[str, str]:
|
|
"""从 request.state 取 trace/request 头(异常路径 contextvar 可能已复位)。"""
|
|
trace_id = getattr(request.state, "trace_id", None) or current_trace() or new_trace()
|
|
request_id = getattr(request.state, "request_id", None) or current_request_id() or new_request_id()
|
|
return {"X-Trace-Id": trace_id, "X-Request-Id": request_id}
|