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