| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217 |
- 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,
- })
|