base.py 6.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148
  1. import asyncio
  2. import time
  3. from abc import ABC, abstractmethod
  4. from functools import wraps
  5. from typing import Any
  6. from fastapi import APIRouter, Request
  7. from fastapi.exceptions import RequestValidationError
  8. from pydantic import BaseModel, ConfigDict
  9. from starlette.exceptions import HTTPException as StarletteHTTPException
  10. from starlette.responses import Response
  11. from api.common import api_response
  12. from api.errors import BusinessValidationError, DatabaseQueryError, ServiceBusyError
  13. from api.reporting import format_validation_errors, report_api_request, set_failure
  14. from config import settings
  15. class ApiParams(BaseModel):
  16. """所有API请求参数的公共基类,默认拒绝未声明字段。"""
  17. model_config = ConfigDict(extra='forbid')
  18. def unified_endpoint(api):
  19. """保留logic签名供FastAPI解析,并由基类统一完成执行和出口处理。"""
  20. @wraps(api.logic)
  21. async def wrapper(*args, **kwargs):
  22. return await api.dispatch(*args, **kwargs)
  23. return wrapper
  24. class BaseApi(ABC):
  25. """API模板:统一注册、执行、响应、异常和日志,子类只实现业务。"""
  26. path: str
  27. methods: tuple[str, ...] = ('POST',)
  28. def register(self, router: APIRouter) -> None:
  29. router.add_api_route(
  30. self.path,
  31. unified_endpoint(self),
  32. methods=list(self.methods),
  33. name=self.__class__.__name__,
  34. summary=(self.__doc__ or '').strip().splitlines()[0] or None,
  35. )
  36. async def dispatch(self, *args, **kwargs) -> Response:
  37. request = self._find_request(args, kwargs)
  38. try:
  39. data = await self.logic(*args, **kwargs)
  40. response = data if isinstance(data, Response) else api_response(data)
  41. except Exception as exc:
  42. response = self.exception_response(request, exc)
  43. if request is not None:
  44. response.headers['X-Request-ID'] = request.state.request_id
  45. await self.report_once(request, response)
  46. return response
  47. async def fetch_all(self, request: Request, sql: str, params: list[Any]) -> list[dict]:
  48. """统一执行列表查询,并记录并发、超时、耗时和结果数量。"""
  49. semaphore = request.app.state.query_semaphore
  50. try:
  51. await asyncio.wait_for(semaphore.acquire(), timeout=settings.API_QUEUE_TIMEOUT)
  52. except asyncio.TimeoutError as exc:
  53. raise ServiceBusyError from exc
  54. query_started_at = time.perf_counter()
  55. try:
  56. async with asyncio.timeout(settings.API_QUERY_TIMEOUT):
  57. rows = await request.app.state.mysql.fetch_all(sql, params)
  58. request.state.query_metrics = {
  59. 'result_count': len(rows),
  60. 'duration_ms': round((time.perf_counter() - query_started_at) * 1000, 2),
  61. 'success': True,
  62. }
  63. return rows
  64. except asyncio.TimeoutError:
  65. request.state.query_metrics = {
  66. 'result_count': None,
  67. 'duration_ms': round((time.perf_counter() - query_started_at) * 1000, 2),
  68. 'success': False,
  69. }
  70. raise
  71. except Exception as exc:
  72. request.state.query_metrics = {
  73. 'result_count': None,
  74. 'duration_ms': round((time.perf_counter() - query_started_at) * 1000, 2),
  75. 'success': False,
  76. }
  77. raise DatabaseQueryError('database query failed') from exc
  78. finally:
  79. semaphore.release()
  80. @staticmethod
  81. def _find_request(args, kwargs) -> Request | None:
  82. candidates = (*args, *kwargs.values())
  83. return next((item for item in candidates if isinstance(item, Request)), None)
  84. @staticmethod
  85. def exception_response(request: Request | None, exc: Exception) -> Response:
  86. """所有API共用的异常到响应映射。"""
  87. if request is None:
  88. raise exc
  89. if isinstance(exc, RequestValidationError):
  90. errors = exc.errors()
  91. message = format_validation_errors(errors)
  92. set_failure(request, 'request_validation', exc, message)
  93. return api_response(code=400, msg=message, data=errors)
  94. if isinstance(exc, BusinessValidationError):
  95. set_failure(request, 'business_validation', exc)
  96. return api_response(code=422, msg=str(exc))
  97. if isinstance(exc, ServiceBusyError):
  98. set_failure(request, 'concurrency_limit', exc, str(exc) or 'server busy')
  99. return api_response(code=503, msg='server busy')
  100. if isinstance(exc, asyncio.TimeoutError):
  101. set_failure(request, 'database_query', exc, str(exc) or 'query timeout')
  102. return api_response(code=504, msg='query timeout')
  103. if isinstance(exc, DatabaseQueryError):
  104. set_failure(request, 'database_query', exc)
  105. request.app.state.logger.exception(str(exc))
  106. return api_response(code=500, msg='database query failed')
  107. if isinstance(exc, StarletteHTTPException):
  108. set_failure(request, 'routing', exc, str(exc.detail))
  109. return api_response(code=exc.status_code, msg=str(exc.detail))
  110. set_failure(request, 'internal', exc, f'{type(exc).__name__}: {exc}')
  111. request.app.state.logger.exception(f'API request failed: {exc}')
  112. return api_response(code=500, msg='internal server error')
  113. @staticmethod
  114. async def report_once(request: Request, response: Response) -> None:
  115. """保证一个请求最多记录和上报一次访问日志。"""
  116. if getattr(request.state, 'api_reported', False):
  117. return
  118. try:
  119. await report_api_request(request, response)
  120. except Exception:
  121. request.app.state.logger.exception('API访问日志记录失败')
  122. finally:
  123. request.state.api_reported = True
  124. @abstractmethod
  125. async def logic(self, *args, **kwargs):
  126. """只实现参数对应的业务逻辑,并直接返回业务data。"""
  127. raise NotImplementedError