| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148 |
- import asyncio
- import time
- from abc import ABC, abstractmethod
- from functools import wraps
- from typing import Any
- from fastapi import APIRouter, Request
- from fastapi.exceptions import RequestValidationError
- from pydantic import BaseModel, ConfigDict
- from starlette.exceptions import HTTPException as StarletteHTTPException
- from starlette.responses import Response
- from api.common import api_response
- from api.errors import BusinessValidationError, DatabaseQueryError, ServiceBusyError
- from api.reporting import format_validation_errors, report_api_request, set_failure
- from config import settings
- class ApiParams(BaseModel):
- """所有API请求参数的公共基类,默认拒绝未声明字段。"""
- model_config = ConfigDict(extra='forbid')
- def unified_endpoint(api):
- """保留logic签名供FastAPI解析,并由基类统一完成执行和出口处理。"""
- @wraps(api.logic)
- async def wrapper(*args, **kwargs):
- return await api.dispatch(*args, **kwargs)
- return wrapper
- class BaseApi(ABC):
- """API模板:统一注册、执行、响应、异常和日志,子类只实现业务。"""
- path: str
- methods: tuple[str, ...] = ('POST',)
- def register(self, router: APIRouter) -> None:
- router.add_api_route(
- self.path,
- unified_endpoint(self),
- methods=list(self.methods),
- name=self.__class__.__name__,
- summary=(self.__doc__ or '').strip().splitlines()[0] or None,
- )
- async def dispatch(self, *args, **kwargs) -> Response:
- request = self._find_request(args, kwargs)
- try:
- data = await self.logic(*args, **kwargs)
- response = data if isinstance(data, Response) else api_response(data)
- except Exception as exc:
- response = self.exception_response(request, exc)
- if request is not None:
- response.headers['X-Request-ID'] = request.state.request_id
- await self.report_once(request, response)
- return response
- async def fetch_all(self, request: Request, sql: str, params: list[Any]) -> list[dict]:
- """统一执行列表查询,并记录并发、超时、耗时和结果数量。"""
- semaphore = request.app.state.query_semaphore
- try:
- await asyncio.wait_for(semaphore.acquire(), timeout=settings.API_QUEUE_TIMEOUT)
- except asyncio.TimeoutError as exc:
- raise ServiceBusyError from exc
- query_started_at = time.perf_counter()
- try:
- async with asyncio.timeout(settings.API_QUERY_TIMEOUT):
- rows = await request.app.state.mysql.fetch_all(sql, params)
- request.state.query_metrics = {
- 'result_count': len(rows),
- 'duration_ms': round((time.perf_counter() - query_started_at) * 1000, 2),
- 'success': True,
- }
- return rows
- except asyncio.TimeoutError:
- request.state.query_metrics = {
- 'result_count': None,
- 'duration_ms': round((time.perf_counter() - query_started_at) * 1000, 2),
- 'success': False,
- }
- raise
- except Exception as exc:
- request.state.query_metrics = {
- 'result_count': None,
- 'duration_ms': round((time.perf_counter() - query_started_at) * 1000, 2),
- 'success': False,
- }
- raise DatabaseQueryError('database query failed') from exc
- finally:
- semaphore.release()
- @staticmethod
- def _find_request(args, kwargs) -> Request | None:
- candidates = (*args, *kwargs.values())
- return next((item for item in candidates if isinstance(item, Request)), None)
- @staticmethod
- def exception_response(request: Request | None, exc: Exception) -> Response:
- """所有API共用的异常到响应映射。"""
- if request is None:
- raise exc
- if isinstance(exc, RequestValidationError):
- errors = exc.errors()
- message = format_validation_errors(errors)
- set_failure(request, 'request_validation', exc, message)
- return api_response(code=400, msg=message, data=errors)
- if isinstance(exc, BusinessValidationError):
- set_failure(request, 'business_validation', exc)
- return api_response(code=422, msg=str(exc))
- if isinstance(exc, ServiceBusyError):
- set_failure(request, 'concurrency_limit', exc, str(exc) or 'server busy')
- return api_response(code=503, msg='server busy')
- if isinstance(exc, asyncio.TimeoutError):
- set_failure(request, 'database_query', exc, str(exc) or 'query timeout')
- return api_response(code=504, msg='query timeout')
- if isinstance(exc, DatabaseQueryError):
- set_failure(request, 'database_query', exc)
- request.app.state.logger.exception(str(exc))
- return api_response(code=500, msg='database query failed')
- if isinstance(exc, StarletteHTTPException):
- set_failure(request, 'routing', exc, str(exc.detail))
- return api_response(code=exc.status_code, msg=str(exc.detail))
- set_failure(request, 'internal', exc, f'{type(exc).__name__}: {exc}')
- request.app.state.logger.exception(f'API request failed: {exc}')
- return api_response(code=500, msg='internal server error')
- @staticmethod
- async def report_once(request: Request, response: Response) -> None:
- """保证一个请求最多记录和上报一次访问日志。"""
- if getattr(request.state, 'api_reported', False):
- return
- try:
- await report_api_request(request, response)
- except Exception:
- request.app.state.logger.exception('API访问日志记录失败')
- finally:
- request.state.api_reported = True
- @abstractmethod
- async def logic(self, *args, **kwargs):
- """只实现参数对应的业务逻辑,并直接返回业务data。"""
- raise NotImplementedError
|