import asyncio import json import time import traceback from typing import Any from fastapi import Request from starlette.responses import Response from config import settings MAX_CLOUD_LOG_VALUE_LENGTH = 64 * 1024 MAX_LOG_BATCH_SIZE = 50 MAX_LOG_RETRIES = 3 SENSITIVE_FIELD_PARTS = ('authorization', 'api_key', 'apikey', 'token', 'password', 'secret') def set_failure(request: Request, stage: str, exc: Exception, message: str | None = None) -> None: request.state.failure_stage = stage request.state.error_type = type(exc).__name__ request.state.error_reason = message if message is not None else str(exc) request.state.error_traceback = traceback.format_exc() def format_validation_errors(errors: list[dict]) -> str: """把FastAPI/Pydantic错误转换为稳定且可定位的中文信息。""" messages = [] for error in errors: location_items = list(error.get('loc', ())) if location_items and location_items[0] == 'body': location_items.pop(0) location = '.'.join(str(item) for item in location_items) or 'body' error_type = error.get('type', '') input_value = str(error.get('input', ''))[:100] context = error.get('ctx') or {} if error_type == 'extra_forbidden': message = f'不支持的参数: {location}' elif error_type == 'literal_error': message = f'{location}不支持值 {input_value},允许值为{context.get("expected", "白名单值")}' elif error_type == 'missing': message = f'{location}不能为空' elif error_type == 'list_too_long': message = f'{location}最多允许{context.get("max_length")}项' elif error_type == 'list_too_short': message = f'{location}至少需要{context.get("min_length")}项' elif error_type == 'less_than_equal': message = f'{location}不能大于{context.get("le")}' elif error_type == 'greater_than_equal': message = f'{location}不能小于{context.get("ge")}' elif error_type == 'value_error': detail = error.get('msg', '').removeprefix('Value error, ') message = f'{location}: {detail}' else: message = f'{location}: {error.get("msg", "参数格式错误")}' messages.append(message) return '参数校验失败: ' + '; '.join(messages) def _redact_value(value: Any): if isinstance(value, dict): return { key: '***' if any(part in str(key).lower() for part in SENSITIVE_FIELD_PARTS) else _redact_value(item) for key, item in value.items() } if isinstance(value, list): return [_redact_value(item) for item in value] return value def _redact_text(value: str) -> str: result = value for secret in ( settings.DB_PASSWORD, settings.ALIYUN_ACCESS_KEY_ID, settings.ALIYUN_ACCESS_KEY_SECRET, ): if secret: result = result.replace(str(secret), '***') return result def _limit_log_value(value): value = _redact_value(value) text = json.dumps(value, ensure_ascii=False, default=str) if len(text) <= MAX_CLOUD_LOG_VALUE_LENGTH: return value return { 'truncated': True, 'original_length': len(text), 'content': text[:MAX_CLOUD_LOG_VALUE_LENGTH], } def _request_url(request: Request) -> str: scheme = request.headers.get('X-Forwarded-Proto', request.url.scheme).split(',', 1)[0].strip() host = request.headers.get('X-Forwarded-Host', request.headers.get('host', '')).split(',', 1)[0].strip() return f'{scheme}://{host}{request.url.path}' + (f'?{request.url.query}' if request.url.query else '') async def _send_cloud_log_batch(app, events: list[dict]) -> None: await asyncio.wait_for( asyncio.to_thread(app.state.aliyun_logger.logging_batch, events), timeout=settings.API_LOG_FLUSH_TIMEOUT, ) async def cloud_log_worker(app) -> None: """异步批量上报访问日志,不阻塞业务请求。""" queue = app.state.cloud_log_queue logger = app.state.logger while True: first_event = await queue.get() if first_event is None: queue.task_done() return events = [first_event] while len(events) < MAX_LOG_BATCH_SIZE: try: event = queue.get_nowait() except asyncio.QueueEmpty: break if event is None: queue.task_done() break events.append(event) try: for attempt in range(MAX_LOG_RETRIES): try: await _send_cloud_log_batch(app, events) break except Exception as exc: if attempt + 1 >= MAX_LOG_RETRIES: logger.exception( f'阿里云API日志批量上报失败: count={len(events)}, ' f'destination={settings.API_ALIYUN_LOG_PROJECT}/' f'{settings.API_ALIYUN_LOGSTORE}, ' f'error={type(exc).__name__}: {exc}' ) else: await asyncio.sleep(0.5 * (2 ** attempt)) finally: for _ in events: queue.task_done() async def _enqueue_cloud_log(app, event: dict) -> None: queue = getattr(app.state, 'cloud_log_queue', None) if queue is None: try: await _send_cloud_log_batch(app, [event]) except Exception as exc: app.state.logger.exception( f'阿里云API日志上报失败: ' f'destination={settings.API_ALIYUN_LOG_PROJECT}/' f'{settings.API_ALIYUN_LOGSTORE}, ' f'error={type(exc).__name__}: {exc}' ) return try: queue.put_nowait(event) except asyncio.QueueFull: app.state.logger.error( f'阿里云API日志队列已满,丢弃日志: request_id={event.get("trace_id", "")}' ) async def report_api_request(request: Request, response: Response) -> None: """统一记录本地访问日志,并异步上报阿里云日志。""" logger = request.app.state.logger duration_ms = round((time.perf_counter() - request.state.started_at) * 1000, 2) request_params = getattr(request.state, 'request_params', {'query': {}, 'body': None}) safe_request_params = _limit_log_value(request_params) error_message = _redact_text( getattr(request.state, 'error_traceback', '') or getattr(request.state, 'error_reason', '') ) request_url = _request_url(request) local_message = ( f'API请求 request_id={request.state.request_id} method={request.method} url={request_url} ' f'params={json.dumps(safe_request_params, ensure_ascii=False, default=str)} ' f'status={response.status_code} duration_ms={duration_ms}' ) if error_message: logger.error(f'{local_message} error={error_message}') else: logger.info(local_message) request_body = request_params.get('body') if isinstance(request_params, dict) else None request_body = request_body if isinstance(request_body, dict) else {} query_metrics = getattr(request.state, 'query_metrics', {}) cloud_data = { 'url': request_url, 'path': request.url.path, 'method': request.method, 'request_id': request.state.request_id, 'request_params': _limit_log_value(request_body), 'platforms': request_body.get('platforms'), 'status_code': response.status_code, 'success': response.status_code < 400, 'request_duration_ms': duration_ms, 'failure_stage': getattr(request.state, 'failure_stage', ''), 'error_type': getattr(request.state, 'error_type', ''), 'message': error_message, 'query_result_count': query_metrics.get('result_count'), 'query_duration_ms': query_metrics.get('duration_ms'), 'query_success': query_metrics.get('success'), } await _enqueue_cloud_log(request.app, { 'code': '2000' if response.status_code < 400 else '9000', 'message': 'API请求成功' if response.status_code < 400 else error_message, 'data': cloud_data, 'trace_id': request.state.request_id, })