Quellcode durchsuchen

优化框架增加全局异常处理和响应处理

liujintao vor 1 Monat
Ursprung
Commit
7d9f82365f
6 geänderte Dateien mit 379 neuen und 46 gelöschten Zeilen
  1. 119 0
      app/exception_handlers.py
  2. 40 17
      app/main.py
  3. 176 0
      app/response.py
  4. 27 17
      app/routers/search.py
  5. 15 6
      app/schemas.py
  6. 2 6
      app/services/zendesk_client.py

+ 119 - 0
app/exception_handlers.py

@@ -0,0 +1,119 @@
+"""全局异常处理器,统一所有错误响应格式为 {code, data, msg}。"""
+from __future__ import annotations
+
+import logging
+from uuid import uuid4
+
+from fastapi import FastAPI, Request, status
+from fastapi.encoders import jsonable_encoder
+from fastapi.exceptions import RequestValidationError
+from starlette.exceptions import HTTPException as StarletteHTTPException
+from starlette.responses import JSONResponse
+
+from app.response import BusinessError, ResponseCode, error_response
+
+logger = logging.getLogger(__name__)
+
+
+def _get_request_info(request: Request) -> str:
+    """提取请求信息用于日志记录(不记录 query string 避免敏感信息泄露)。"""
+    x_forwarded_for = request.headers.get("x-forwarded-for", "")
+    real_ip = x_forwarded_for.split(",")[0].strip() if x_forwarded_for else ""
+    client_host = real_ip or (request.client.host if request.client else "unknown")
+    return f"{request.method} {request.url.path} (client: {client_host})"
+
+
+# HTTP 状态码 -> 业务错误码 的映射
+_HTTP_TO_CODE: dict[int, ResponseCode] = {
+    400: ResponseCode.REQUEST_INVALID,
+    404: ResponseCode.NOT_FOUND,
+    405: ResponseCode.METHOD_NOT_ALLOWED,
+    422: ResponseCode.PARAM_INVALID,
+    502: ResponseCode.UPSTREAM_ERROR,
+    504: ResponseCode.UPSTREAM_TIMEOUT,
+}
+
+
+async def validation_exception_handler(
+    request: Request, exc: RequestValidationError
+) -> JSONResponse:
+    """处理 Pydantic 请求验证错误(如 page=-1、缺失必填参数等)。
+
+    返回示例:
+        {
+            "code": 1001,
+            "data": [{"loc": [...], "msg": "...", "type": "..."}],
+            "msg": "参数验证失败"
+        }
+    """
+    request_info = _get_request_info(request)
+    logger.warning("参数验证失败 [%s]: %s", request_info, exc.errors())
+
+    return error_response(
+        code=ResponseCode.PARAM_INVALID,
+        data=jsonable_encoder(exc.errors()),
+        http_status=status.HTTP_422_UNPROCESSABLE_ENTITY,
+    )
+
+
+async def http_exception_handler(
+    request: Request, exc: StarletteHTTPException
+) -> JSONResponse:
+    """处理 Starlette/FastAPI HTTPException(404/405/502 等)。"""
+    request_info = _get_request_info(request)
+    status_code = exc.status_code
+
+    if status_code >= 500:
+        logger.error("服务器错误 [%s]: %s", request_info, exc.detail)
+    elif status_code >= 400:
+        logger.warning("客户端错误 [%s]: %s", request_info, exc.detail)
+
+    code = _HTTP_TO_CODE.get(status_code, ResponseCode.INTERNAL_ERROR)
+    return error_response(
+        code=code,
+        msg=str(exc.detail) if exc.detail else None,
+        http_status=status_code,
+    )
+
+
+async def business_exception_handler(
+    request: Request, exc: BusinessError
+) -> JSONResponse:
+    """处理业务异常(BusinessError)。"""
+    request_info = _get_request_info(request)
+    logger.warning(
+        "业务异常 [%s]: code=%s msg=%s", request_info, exc.code, exc.msg
+    )
+    return error_response(
+        code=exc.response_code,
+        msg=exc.msg,
+        data=exc.data,
+        http_status=exc.http_status,
+    )
+
+
+async def general_exception_handler(
+    request: Request, exc: Exception
+) -> JSONResponse:
+    """兜底异常处理器,捕获所有未处理的异常,避免堆栈信息泄露。
+
+    注意:仅捕获 Exception,不捕获 BaseException(如 asyncio.CancelledError),
+    以确保 FastAPI lifespan 的取消逻辑正常工作。严禁改为捕获 BaseException!
+    """
+    request_info = _get_request_info(request)
+    request_id = str(uuid4())
+    logger.exception("未捕获异常 [%s] [request_id=%s]", request_info, request_id)
+
+    return error_response(
+        code=ResponseCode.INTERNAL_ERROR,
+        data={"request_id": request_id},
+        http_status=status.HTTP_500_INTERNAL_SERVER_ERROR,
+    )
+
+
+def register_exception_handlers(app: FastAPI) -> None:
+    """向 FastAPI 应用注册所有全局异常处理器。"""
+    app.add_exception_handler(RequestValidationError, validation_exception_handler)
+    app.add_exception_handler(StarletteHTTPException, http_exception_handler)
+    app.add_exception_handler(BusinessError, business_exception_handler)  # type: ignore[arg-type]
+    app.add_exception_handler(Exception, general_exception_handler)  # type: ignore[arg-type]

+ 40 - 17
app/main.py

@@ -8,8 +8,10 @@ from typing import AsyncIterator
 
 from fastapi import FastAPI
 
+from app.exception_handlers import register_exception_handlers
+from app.response import ApiResponse
 from app.routers import search
-from app.schemas import CacheStatus
+from app.schemas import CacheRefreshResult, CacheStatus
 from app.services.cache import faq_cache, scheduler_loop
 
 logging.basicConfig(
@@ -46,32 +48,53 @@ app = FastAPI(
     lifespan=lifespan,
 )
 
+# 注册全局异常处理器(统一错误响应格式 {code, data, msg})
+register_exception_handlers(app)
+
 # 路由
 app.include_router(search.router)
 
 
-@app.get("/health", tags=["meta"])
-async def health() -> dict[str, str]:
-    return {"status": "ok"}
+@app.get(
+    "/health",
+    response_model=ApiResponse[dict[str, str]],
+    tags=["meta"],
+)
+async def health() -> ApiResponse[dict[str, str]]:
+    """健康检查。"""
+    return ApiResponse.ok(data={"status": "ok"})
 
 
-@app.get("/health/cache", response_model=CacheStatus, tags=["meta"])
-async def cache_status() -> CacheStatus:
+@app.get(
+    "/health/cache",
+    response_model=ApiResponse[CacheStatus],
+    tags=["meta"],
+)
+async def cache_status() -> ApiResponse[CacheStatus]:
+    """缓存状态查看。"""
     snap = faq_cache.snapshot()
-    return CacheStatus(
-        sec_ids_count=len(snap["sec_ids"]),
-        last_updated_at=snap["last_updated_at"],
-        next_refresh_at=snap["next_refresh_at"],
+    return ApiResponse.ok(
+        data=CacheStatus(
+            sec_ids_count=len(snap["sec_ids"]),
+            last_updated_at=snap["last_updated_at"],
+            next_refresh_at=snap["next_refresh_at"],
+        )
     )
 
 
-@app.post("/admin/cache/refresh", tags=["meta"])
-async def refresh_cache_now() -> dict[str, object]:
+@app.post(
+    "/admin/cache/refresh",
+    response_model=ApiResponse[CacheRefreshResult],
+    tags=["meta"],
+)
+async def refresh_cache_now() -> ApiResponse[CacheRefreshResult]:
     """手动触发刷新(运维 / 测试用途)。"""
     await faq_cache.refresh()
     snap = faq_cache.snapshot()
-    return {
-        "success": True,
-        "sec_ids_count": len(snap["sec_ids"]),
-        "last_updated_at": snap["last_updated_at"],
-    }
+    return ApiResponse.ok(
+        data=CacheRefreshResult(
+            sec_ids_count=len(snap["sec_ids"]),
+            last_updated_at=snap["last_updated_at"],
+        ),
+        msg="缓存刷新成功",
+    )

+ 176 - 0
app/response.py

@@ -0,0 +1,176 @@
+"""统一响应类与业务错误码定义。
+
+所有接口响应统一使用 `{code, data, msg}` 信封格式:
+    - code: 0 表示成功;其他值代表不同的失败类型
+    - data: 业务数据(成功时)或 None(失败时)
+    - msg:  描述信息
+
+错误码使用 `(code, default_msg)` 的元组定义,code 和默认提示一一对应。
+"""
+from __future__ import annotations
+
+from enum import Enum
+from typing import Any, Generic, TypeVar
+
+from fastapi.responses import JSONResponse
+from pydantic import BaseModel, Field
+
+T = TypeVar("T")
+
+
+class ResponseCode(Enum):
+    """业务错误码枚举:每个成员都是 (code, default_msg) 元组。
+
+    分段规则:
+        0          成功
+        1000-1999  通用错误(参数、请求格式等)
+        2000-2999  业务逻辑错误(缓存未就绪、资源不存在等)
+        3000-3999  外部依赖错误(Zendesk、上游服务等)
+        9000-9999  服务器内部错误
+
+    用法:
+        ResponseCode.SUCCESS.code   # -> 0
+        ResponseCode.SUCCESS.msg    # -> "成功"
+    """
+
+    # 成功
+    SUCCESS = (0, "成功")
+
+    # 通用错误 1xxx
+    PARAM_INVALID = (1001, "参数验证失败")
+    REQUEST_INVALID = (1002, "请求格式不正确")
+    METHOD_NOT_ALLOWED = (1003, "请求方法不允许")
+    NOT_FOUND = (1004, "资源不存在")
+
+    # 业务错误 2xxx
+    CACHE_NOT_READY = (2001, "FAQ 缓存未就绪,请稍后重试")
+    RESOURCE_NOT_FOUND = (2002, "业务资源不存在")
+
+    # 外部依赖错误 3xxx
+    UPSTREAM_ERROR = (3001, "上游服务调用失败")
+    UPSTREAM_TIMEOUT = (3002, "上游服务响应超时")
+
+    # 服务器错误 9xxx
+    INTERNAL_ERROR = (9999, "服务器内部错误,请稍后重试")
+
+    def __init__(self, code: int, msg: str) -> None:
+        self.code = code
+        self.msg = msg
+
+    def __repr__(self) -> str:  # pragma: no cover
+        return f"<{self.__class__.__name__}.{self.name}: code={self.code} msg={self.msg!r}>"
+
+
+class ApiResponse(BaseModel, Generic[T]):
+    """统一 API 响应模型。
+
+    用于在路由 `response_model` 中声明,让 OpenAPI 文档展示标准结构。
+
+    用法示例:
+        @router.get("/search", response_model=ApiResponse[SearchData])
+        async def search(...) -> ApiResponse[SearchData]:
+            return ApiResponse.ok(data=SearchData(...))
+
+        # 失败响应(自动使用错误码对应的默认 msg)
+        return ApiResponse.fail(ResponseCode.PARAM_INVALID)
+
+        # 失败响应(覆盖默认 msg)
+        return ApiResponse.fail(ResponseCode.UPSTREAM_ERROR, msg="Zendesk 限流")
+    """
+
+    code: int = Field(default=0, description="业务码,0=成功,非0=失败")
+    data: T | None = Field(default=None, description="业务数据,失败时为 null")
+    msg: str = Field(default="成功", description="描述信息")
+
+    @classmethod
+    def ok(
+        cls,
+        data: T | None = None,
+        msg: str | None = None,
+    ) -> "ApiResponse[T]":
+        """构造成功响应。msg 为空时使用 SUCCESS 默认提示。"""
+        return cls(
+            code=ResponseCode.SUCCESS.code,
+            data=data,
+            msg=msg if msg is not None else ResponseCode.SUCCESS.msg,
+        )
+
+    @classmethod
+    def fail(
+        cls,
+        code: ResponseCode,
+        msg: str | None = None,
+        data: T | None = None,
+    ) -> "ApiResponse[T]":
+        """构造失败响应。msg 为空时使用错误码自带的默认提示。"""
+        return cls(
+            code=code.code,
+            data=data,
+            msg=msg if msg is not None else code.msg,
+        )
+
+
+def success_response(
+    data: Any = None,
+    msg: str | None = None,
+    http_status: int = 200,
+) -> JSONResponse:
+    """构造成功的 JSONResponse(适用于异常处理器、中间件等场景)。"""
+    return JSONResponse(
+        status_code=http_status,
+        content={
+            "code": ResponseCode.SUCCESS.code,
+            "data": data,
+            "msg": msg if msg is not None else ResponseCode.SUCCESS.msg,
+        },
+    )
+
+
+def error_response(
+    code: ResponseCode,
+    msg: str | None = None,
+    data: Any = None,
+    http_status: int = 200,
+) -> JSONResponse:
+    """构造失败的 JSONResponse(适用于异常处理器、中间件等场景)。
+
+    Args:
+        code: 业务错误码枚举
+        msg:  错误提示,为空时使用错误码自带的默认提示
+        data: 附加数据(如参数验证错误的详细字段列表)
+        http_status: HTTP 状态码,默认 200(业务错误也走 200,由 code 区分成败)
+
+    返回示例:
+        {"code": 1001, "data": [...], "msg": "参数验证失败"}
+    """
+    return JSONResponse(
+        status_code=http_status,
+        content={
+            "code": code.code,
+            "data": data,
+            "msg": msg if msg is not None else code.msg,
+        },
+    )
+
+
+class BusinessError(Exception):
+    """业务异常基类,用于在业务层抛出后由全局处理器统一转为 ApiResponse。
+
+    用法:
+        raise BusinessError(ResponseCode.CACHE_NOT_READY)
+        raise BusinessError(ResponseCode.UPSTREAM_ERROR, msg="Zendesk 限流")
+    """
+
+    def __init__(
+        self,
+        code: ResponseCode,
+        msg: str | None = None,
+        data: Any = None,
+        http_status: int = 200,
+    ) -> None:
+        self.response_code = code
+        self.code = code.code
+        self.msg = msg if msg is not None else code.msg
+        self.data = data
+        self.http_status = http_status
+        super().__init__(self.msg)

+ 27 - 17
app/routers/search.py

@@ -3,9 +3,10 @@ from __future__ import annotations
 
 import logging
 
-from fastapi import APIRouter, HTTPException, Query
+from fastapi import APIRouter, Query
 
-from app.schemas import SearchRequest, SearchResponse
+from app.response import ApiResponse, BusinessError, ResponseCode
+from app.schemas import SearchData, SearchRequest
 from app.services.cache import faq_cache
 from app.services.zendesk_client import ZendeskError, search_articles
 
@@ -14,7 +15,7 @@ logger = logging.getLogger(__name__)
 router = APIRouter(tags=["search"])
 
 
-async def _do_search(req: SearchRequest) -> SearchResponse:
+async def _do_search(req: SearchRequest) -> ApiResponse[SearchData]:
     """共享的搜索逻辑:使用内存中的 FAQ sec_ids 调用 Zendesk。"""
     snapshot = faq_cache.snapshot()
     sec_ids: list[int] = snapshot["sec_ids"]
@@ -31,27 +32,32 @@ async def _do_search(req: SearchRequest) -> SearchResponse:
             per_page=req.per_page,
         )
     except ZendeskError as exc:
-        raise HTTPException(status_code=502, detail=str(exc)) from exc
+        raise BusinessError(ResponseCode.UPSTREAM_ERROR, msg=str(exc)) from exc
 
-    return SearchResponse(
-        success=True,
-        query=req.query,
-        count=int(data.get("count", 0)),
-        page=req.page,
-        per_page=req.per_page,
-        next_page=data.get("next_page"),
-        sec_ids_used=sec_ids,
-        results=data.get("results", []),
+    return ApiResponse.ok(
+        data=SearchData(
+            query=req.query,
+            count=int(data.get("count", 0)),
+            page=req.page,
+            per_page=req.per_page,
+            next_page=data.get("next_page"),
+            sec_ids_used=sec_ids,
+            results=data.get("results", []),
+        )
     )
 
 
-@router.get("/search", response_model=SearchResponse, summary="按 QUERY 搜索 FAQ 文章")
+@router.get(
+    "/search",
+    response_model=ApiResponse[SearchData],
+    summary="按 QUERY 搜索 FAQ 文章",
+)
 async def search_get(
     query: str = Query(..., min_length=1, description="搜索关键词"),
     locale: str | None = Query(default=None),
     page: int = Query(default=1, ge=1),
     per_page: int = Query(default=25, ge=1, le=100),
-) -> SearchResponse:
+) -> ApiResponse[SearchData]:
     """GET 版本,便于浏览器直接测试。"""
     return await _do_search(
         SearchRequest(
@@ -60,7 +66,11 @@ async def search_get(
     )
 
 
-@router.post("/search", response_model=SearchResponse, summary="按 QUERY 搜索 FAQ 文章 (POST)")
-async def search_post(req: SearchRequest) -> SearchResponse:
+@router.post(
+    "/search",
+    response_model=ApiResponse[SearchData],
+    summary="按 QUERY 搜索 FAQ 文章 (POST)",
+)
+async def search_post(req: SearchRequest) -> ApiResponse[SearchData]:
     """POST 版本,请求体:{"query": "...", "locale": "...", ...}。"""
     return await _do_search(req)

+ 15 - 6
app/schemas.py

@@ -1,4 +1,8 @@
-"""请求 / 响应 Pydantic 模型。"""
+"""请求 / 响应 Pydantic 模型。
+
+注意:所有响应都使用 app/response.py 中的 ApiResponse 信封包裹,
+本文件仅定义 data 字段内的业务数据结构。
+"""
 from typing import Any
 
 from pydantic import BaseModel, Field
@@ -13,10 +17,9 @@ class SearchRequest(BaseModel):
     per_page: int = Field(default=25, ge=1, le=100)
 
 
-class SearchResponse(BaseModel):
-    """统一响应信封。"""
+class SearchData(BaseModel):
+    """搜索响应数据(放入 ApiResponse.data 字段)。"""
 
-    success: bool
     query: str
     count: int
     page: int
@@ -24,12 +27,18 @@ class SearchResponse(BaseModel):
     next_page: str | None = None
     sec_ids_used: list[int]
     results: list[dict[str, Any]]
-    error: str | None = None
 
 
 class CacheStatus(BaseModel):
-    """缓存状态。"""
+    """缓存状态(放入 ApiResponse.data 字段)。"""
 
     sec_ids_count: int
     last_updated_at: str | None
     next_refresh_at: str | None
+
+
+class CacheRefreshResult(BaseModel):
+    """手动刷新缓存的返回数据。"""
+
+    sec_ids_count: int
+    last_updated_at: str | None

+ 2 - 6
app/services/zendesk_client.py

@@ -63,11 +63,7 @@ async def list_all_faq_sections(
             logger.info("扫描 Locale: %s", locale)
 
             while url:
-                params: dict[str, Any] = {"per_page": 100}
-                # next_page 已包含 page 参数;首次请求才需要显式传
-                if "page=" not in url:
-                    params["page"] = page
-
+                params: dict[str, Any] = {"per_page": 100, "page": page} if "page" in url else {"per_page": 100}
                 try:
                     resp = await client.get(url, params=params)
                 except httpx.HTTPError as exc:
@@ -84,7 +80,7 @@ async def list_all_faq_sections(
 
                 data = resp.json()
                 sections = data.get("sections", [])
-                # logger.info("第 %s 页,共 %s 个 Section", page, len(sections))
+                logger.info("第 %s 页,共 %s 个 Section", page, len(sections))
 
                 for sec in sections:
                     name = (sec.get("name") or "").strip()