app.py 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475
  1. import asyncio
  2. import hmac
  3. import json
  4. import re
  5. import time
  6. import traceback
  7. import uuid
  8. from typing import Any
  9. from aiohttp import web
  10. from pydantic import ValidationError
  11. from api.chui_zhi.videos import (
  12. MYSQL_KEY,
  13. QUERY_METRICS_KEY,
  14. QUERY_SEMAPHORE_KEY,
  15. BusinessValidationError,
  16. DatabaseQueryError,
  17. ServiceBusyError,
  18. api_response,
  19. query_videos,
  20. )
  21. from config import settings
  22. from core.base.async_mysql_client import AsyncMySQLClient
  23. from core.utils.log.logger_manager import LoggerManager
  24. LOGGER_KEY = web.AppKey('logger', object)
  25. ALIYUN_LOGGER_KEY = web.AppKey('aliyun_logger', object)
  26. CLOUD_LOG_QUEUE_KEY = web.AppKey('cloud_log_queue', asyncio.Queue)
  27. CLOUD_LOG_WORKER_KEY = web.AppKey('cloud_log_worker', asyncio.Task)
  28. ERROR_REASON_KEY = web.AppKey('error_reason', str)
  29. ERROR_TRACEBACK_KEY = web.AppKey('error_traceback', str)
  30. ERROR_TYPE_KEY = web.AppKey('error_type', str)
  31. FAILURE_STAGE_KEY = web.AppKey('failure_stage', str)
  32. REQUEST_ID_KEY = web.AppKey('request_id', str)
  33. REQUEST_PARAMS_KEY = web.AppKey('request_params', dict)
  34. REQUEST_STARTED_AT_KEY = web.AppKey('request_started_at', float)
  35. API_PATH = '/api/v1/crawler/videos/query'
  36. HEALTH_PATH = '/health'
  37. READY_PATH = '/ready'
  38. AUTH_EXEMPT_PATHS = frozenset({HEALTH_PATH, READY_PATH})
  39. MAX_CLOUD_LOG_VALUE_LENGTH = 64 * 1024
  40. MAX_LOG_BATCH_SIZE = 50
  41. MAX_LOG_RETRIES = 3
  42. SENSITIVE_FIELD_PARTS = ('authorization', 'api_key', 'apikey', 'token', 'password', 'secret')
  43. def print_api_routes(app: web.Application) -> None:
  44. print('\n已注册 API:', flush=True)
  45. for route in app.router.routes():
  46. path = route.resource.canonical
  47. handler_name = getattr(route.handler, '__name__', route.handler.__class__.__name__)
  48. print(f' {route.method:<6} {path:<36} -> {handler_name}', flush=True)
  49. print('', flush=True)
  50. def parse_json_text(text: str):
  51. """尽量将日志内容还原为JSON;非JSON内容保留原字符串。"""
  52. if not text:
  53. return None
  54. try:
  55. return json.loads(text)
  56. except (TypeError, ValueError):
  57. return text
  58. def format_validation_error(exc: ValidationError) -> str:
  59. """把Pydantic错误转换为调用方可直接定位的简洁参数信息。"""
  60. messages = []
  61. for error in exc.errors(include_url=False):
  62. location = '.'.join(str(item) for item in error.get('loc', ())) or 'body'
  63. error_type = error.get('type', '')
  64. input_value = str(error.get('input', ''))[:100]
  65. context = error.get('ctx') or {}
  66. if error_type == 'extra_forbidden':
  67. message = f'不支持的参数: {location}'
  68. elif error_type == 'literal_error':
  69. message = f'{location}不支持值 {input_value},允许值为{context.get("expected", "白名单值")}'
  70. elif error_type == 'missing':
  71. message = f'{location}不能为空'
  72. elif error_type == 'list_too_long':
  73. message = f'{location}最多允许{context.get("max_length")}项'
  74. elif error_type == 'list_too_short':
  75. message = f'{location}至少需要{context.get("min_length")}项'
  76. elif error_type == 'less_than_equal':
  77. message = f'{location}不能大于{context.get("le")}'
  78. elif error_type == 'greater_than_equal':
  79. message = f'{location}不能小于{context.get("ge")}'
  80. elif error_type == 'value_error':
  81. detail = error.get('msg', '').removeprefix('Value error, ')
  82. message = f'{location}: {detail}'
  83. else:
  84. message = f'{location}: {error.get("msg", "参数格式错误")}'
  85. messages.append(message)
  86. return '参数校验失败: ' + '; '.join(messages)
  87. def redact_value(value: Any):
  88. """递归脱敏凭证字段,避免未来扩展请求参数时意外泄露。"""
  89. if isinstance(value, dict):
  90. return {
  91. key: '***' if any(part in str(key).lower() for part in SENSITIVE_FIELD_PARTS)
  92. else redact_value(item)
  93. for key, item in value.items()
  94. }
  95. if isinstance(value, list):
  96. return [redact_value(item) for item in value]
  97. return value
  98. def redact_text(value: str) -> str:
  99. """保留异常堆栈,同时移除配置中已知的密钥内容。"""
  100. result = value
  101. for secret in (
  102. settings.CHUI_ZHI_API_TOKEN,
  103. settings.DB_PASSWORD,
  104. settings.ALIYUN_ACCESS_KEY_ID,
  105. settings.ALIYUN_ACCESS_KEY_SECRET,
  106. ):
  107. if secret:
  108. result = result.replace(str(secret), '***')
  109. return result
  110. def limit_cloud_log_value(value):
  111. """限制单个日志字段大小,避免请求参数超过SLS单条日志限制。"""
  112. value = redact_value(value)
  113. text = json.dumps(value, ensure_ascii=False, default=str)
  114. if len(text) <= MAX_CLOUD_LOG_VALUE_LENGTH:
  115. return value
  116. return {
  117. 'truncated': True,
  118. 'original_length': len(text),
  119. 'content': text[:MAX_CLOUD_LOG_VALUE_LENGTH],
  120. }
  121. def get_request_url(request: web.Request) -> str:
  122. scheme = request.headers.get('X-Forwarded-Proto', request.scheme).split(',', 1)[0].strip()
  123. host = request.headers.get('X-Forwarded-Host', request.host).split(',', 1)[0].strip()
  124. return f'{scheme}://{host}{request.rel_url}'
  125. async def send_cloud_log_batch(app: web.Application, events: list[dict]) -> None:
  126. await asyncio.wait_for(
  127. asyncio.to_thread(app[ALIYUN_LOGGER_KEY].logging_batch, events),
  128. timeout=settings.API_LOG_FLUSH_TIMEOUT,
  129. )
  130. async def cloud_log_worker(app: web.Application) -> None:
  131. """后台批量上报SLS,日志服务异常不阻塞业务请求。"""
  132. queue = app[CLOUD_LOG_QUEUE_KEY]
  133. logger = app[LOGGER_KEY]
  134. while True:
  135. event = await queue.get()
  136. if event is None:
  137. queue.task_done()
  138. return
  139. events = [event]
  140. while len(events) < MAX_LOG_BATCH_SIZE:
  141. try:
  142. next_event = queue.get_nowait()
  143. except asyncio.QueueEmpty:
  144. break
  145. if next_event is None:
  146. queue.task_done()
  147. break
  148. events.append(next_event)
  149. try:
  150. for attempt in range(MAX_LOG_RETRIES):
  151. try:
  152. await send_cloud_log_batch(app, events)
  153. break
  154. except Exception:
  155. if attempt + 1 >= MAX_LOG_RETRIES:
  156. logger.exception(f'阿里云API日志批量上报失败: count={len(events)}')
  157. else:
  158. await asyncio.sleep(0.5 * (2 ** attempt))
  159. finally:
  160. for _ in events:
  161. queue.task_done()
  162. async def enqueue_cloud_log(app: web.Application, event: dict) -> None:
  163. queue = app.get(CLOUD_LOG_QUEUE_KEY)
  164. if queue is None:
  165. # 单元测试或嵌入运行未启动cleanup context时,仍可验证日志行为。
  166. try:
  167. await send_cloud_log_batch(app, [event])
  168. except Exception:
  169. app[LOGGER_KEY].exception('阿里云API日志上报失败')
  170. return
  171. try:
  172. queue.put_nowait(event)
  173. except asyncio.QueueFull:
  174. app[LOGGER_KEY].error(
  175. f'阿里云API日志队列已满,丢弃日志: request_id={event.get("trace_id", "")}'
  176. )
  177. async def report_api_request(
  178. request: web.Request,
  179. response: web.StreamResponse | None,
  180. exception: Exception | None = None,
  181. ) -> None:
  182. """本地日志同步落地,SLS日志仅入队,不增加接口响应延迟。"""
  183. logger = request.app[LOGGER_KEY]
  184. status_code = response.status if response is not None else getattr(exception, 'status', 500)
  185. request_duration_ms = round(
  186. (time.perf_counter() - request.get(REQUEST_STARTED_AT_KEY, time.perf_counter())) * 1000,
  187. 2,
  188. )
  189. request_params = request.get(REQUEST_PARAMS_KEY, {'query': dict(request.query), 'body': None})
  190. safe_request_params = limit_cloud_log_value(request_params)
  191. request_id = request[REQUEST_ID_KEY]
  192. error_message = request.get(ERROR_TRACEBACK_KEY, '') or request.get(ERROR_REASON_KEY, '')
  193. if not error_message and exception is not None:
  194. error_message = f'{type(exception).__name__}: {exception}'
  195. error_message = redact_text(error_message)
  196. error_type = request.get(ERROR_TYPE_KEY, '')
  197. if not error_type and exception is not None:
  198. error_type = type(exception).__name__
  199. failure_stage = request.get(FAILURE_STAGE_KEY, '')
  200. request_url = get_request_url(request)
  201. local_message = (
  202. f'API请求 request_id={request_id} method={request.method} url={request_url} '
  203. f'params={json.dumps(safe_request_params, ensure_ascii=False, default=str)} '
  204. f'status={status_code} duration_ms={request_duration_ms}'
  205. )
  206. if error_message:
  207. logger.error(f'{local_message} error={error_message}')
  208. else:
  209. logger.info(local_message)
  210. query_metrics = request.get(QUERY_METRICS_KEY, {})
  211. request_body = request_params.get('body') if isinstance(request_params, dict) else None
  212. request_body = request_body if isinstance(request_body, dict) else {}
  213. cloud_data = {
  214. 'url': request_url,
  215. 'path': request.path,
  216. 'method': request.method,
  217. 'request_id': request_id,
  218. 'request_params': safe_request_params,
  219. 'platforms': request_body.get('platforms'),
  220. 'status_code': status_code,
  221. 'success': status_code < 400,
  222. 'request_duration_ms': request_duration_ms,
  223. 'failure_stage': failure_stage,
  224. 'error_type': error_type,
  225. 'message': error_message,
  226. 'query_result_count': query_metrics.get('result_count'),
  227. 'query_duration_ms': query_metrics.get('duration_ms'),
  228. 'query_success': query_metrics.get('success'),
  229. }
  230. await enqueue_cloud_log(request.app, {
  231. 'code': '2000' if status_code < 400 else '9000',
  232. 'message': 'API请求成功' if status_code < 400 else error_message,
  233. 'data': cloud_data,
  234. 'trace_id': request_id,
  235. })
  236. @web.middleware
  237. async def request_context_middleware(request: web.Request, handler):
  238. """建立请求上下文,确保成功、失败和未鉴权请求都有访问日志。"""
  239. incoming_request_id = request.headers.get('X-Request-ID', '').strip()
  240. if not re.fullmatch(r'[A-Za-z0-9._-]{1,128}', incoming_request_id):
  241. incoming_request_id = uuid.uuid4().hex
  242. request[REQUEST_ID_KEY] = incoming_request_id
  243. request[REQUEST_STARTED_AT_KEY] = time.perf_counter()
  244. request[REQUEST_PARAMS_KEY] = {
  245. 'query': dict(request.query),
  246. 'body': None,
  247. }
  248. response = None
  249. exception = None
  250. try:
  251. response = await handler(request)
  252. response.headers['X-Request-ID'] = request[REQUEST_ID_KEY]
  253. return response
  254. except Exception as exc:
  255. exception = exc
  256. request[ERROR_TRACEBACK_KEY] = traceback.format_exc()
  257. raise
  258. finally:
  259. try:
  260. await report_api_request(request, response, exception)
  261. except Exception:
  262. # 日志异常不能覆盖API原始响应或异常。
  263. request.app[LOGGER_KEY].exception('API访问日志记录失败')
  264. @web.middleware
  265. async def request_body_middleware(request: web.Request, handler):
  266. """鉴权通过后再读取请求体,避免未授权请求消耗JSON解析资源。"""
  267. raw_body = await request.read()
  268. request[REQUEST_PARAMS_KEY]['body'] = parse_json_text(
  269. raw_body.decode('utf-8', errors='replace')
  270. )
  271. return await handler(request)
  272. @web.middleware
  273. async def error_middleware(request: web.Request, handler):
  274. try:
  275. return await handler(request)
  276. except json.JSONDecodeError as exc:
  277. request[ERROR_REASON_KEY] = str(exc)
  278. request[ERROR_TRACEBACK_KEY] = traceback.format_exc()
  279. request[ERROR_TYPE_KEY] = type(exc).__name__
  280. request[FAILURE_STAGE_KEY] = 'request_validation'
  281. return api_response(code=400, msg=f'请求JSON格式错误: {exc.msg}')
  282. except ValidationError as exc:
  283. message = format_validation_error(exc)
  284. request[ERROR_REASON_KEY] = message
  285. request[ERROR_TRACEBACK_KEY] = traceback.format_exc()
  286. request[ERROR_TYPE_KEY] = type(exc).__name__
  287. request[FAILURE_STAGE_KEY] = 'request_validation'
  288. return api_response(code=400, msg=message, data=exc.errors(include_url=False))
  289. except BusinessValidationError as exc:
  290. request[ERROR_REASON_KEY] = str(exc)
  291. request[ERROR_TRACEBACK_KEY] = traceback.format_exc()
  292. request[ERROR_TYPE_KEY] = type(exc).__name__
  293. request[FAILURE_STAGE_KEY] = 'business_validation'
  294. return api_response(code=422, msg=str(exc))
  295. except ServiceBusyError as exc:
  296. request[ERROR_REASON_KEY] = str(exc) or 'server busy'
  297. request[ERROR_TRACEBACK_KEY] = traceback.format_exc()
  298. request[ERROR_TYPE_KEY] = type(exc).__name__
  299. request[FAILURE_STAGE_KEY] = 'concurrency_limit'
  300. return api_response(code=503, msg='server busy')
  301. except asyncio.TimeoutError as exc:
  302. request[ERROR_REASON_KEY] = str(exc) or 'query timeout'
  303. request[ERROR_TRACEBACK_KEY] = traceback.format_exc()
  304. request[ERROR_TYPE_KEY] = type(exc).__name__
  305. request[FAILURE_STAGE_KEY] = 'database_query'
  306. return api_response(code=504, msg='query timeout')
  307. except DatabaseQueryError as exc:
  308. request[ERROR_REASON_KEY] = str(exc)
  309. request[ERROR_TRACEBACK_KEY] = traceback.format_exc()
  310. request[ERROR_TYPE_KEY] = type(exc).__name__
  311. request[FAILURE_STAGE_KEY] = 'database_query'
  312. request.app[LOGGER_KEY].exception(str(exc))
  313. return api_response(code=500, msg='database query failed')
  314. except web.HTTPException:
  315. raise
  316. except Exception as exc:
  317. request[ERROR_REASON_KEY] = f'{type(exc).__name__}: {exc}'
  318. request[ERROR_TRACEBACK_KEY] = traceback.format_exc()
  319. request[ERROR_TYPE_KEY] = type(exc).__name__
  320. request[FAILURE_STAGE_KEY] = 'internal'
  321. request.app[LOGGER_KEY].exception(f'API request failed: {exc}')
  322. return api_response(code=500, msg='internal server error')
  323. @web.middleware
  324. async def auth_middleware(request: web.Request, handler):
  325. if request.path in AUTH_EXEMPT_PATHS:
  326. return await handler(request)
  327. supplied_token = request.headers.get('X-API-Key', '')
  328. if not hmac.compare_digest(supplied_token, settings.CHUI_ZHI_API_TOKEN):
  329. request[ERROR_REASON_KEY] = 'unauthorized'
  330. request[ERROR_TYPE_KEY] = 'AuthenticationError'
  331. request[FAILURE_STAGE_KEY] = 'authentication'
  332. return api_response(code=401, msg='unauthorized')
  333. return await handler(request)
  334. async def health(_request: web.Request) -> web.Response:
  335. return api_response({'status': 'ok'})
  336. async def ready(request: web.Request) -> web.Response:
  337. try:
  338. async with asyncio.timeout(min(3, settings.API_QUERY_TIMEOUT)):
  339. row = await request.app[MYSQL_KEY].fetch_one('SELECT 1 AS ok')
  340. if not row or row.get('ok') != 1:
  341. raise DatabaseQueryError('database readiness check failed')
  342. return api_response({'status': 'ready'})
  343. except Exception as exc:
  344. request[ERROR_REASON_KEY] = f'{type(exc).__name__}: {exc}'
  345. request[ERROR_TRACEBACK_KEY] = traceback.format_exc()
  346. request[ERROR_TYPE_KEY] = type(exc).__name__
  347. request[FAILURE_STAGE_KEY] = 'readiness'
  348. return api_response(code=503, msg='not ready')
  349. async def app_context(app: web.Application):
  350. if not settings.CHUI_ZHI_API_TOKEN:
  351. raise RuntimeError('CHUI_ZHI_API_TOKEN未配置,拒绝启动垂直spider API')
  352. logger = LoggerManager.get_logger(platform='chui_zhi', mode='api')
  353. aliyun_logger = LoggerManager.get_aliyun_logger(platform='chui_zhi', mode='api')
  354. mysql = AsyncMySQLClient(
  355. host=settings.DB_HOST,
  356. port=settings.DB_PORT,
  357. user=settings.DB_USER,
  358. password=settings.DB_PASSWORD,
  359. db=settings.DB_NAME,
  360. charset=settings.DB_CHARSET,
  361. minsize=min(5, settings.DB_POOL_SIZE),
  362. maxsize=settings.DB_POOL_SIZE,
  363. pool_recycle=settings.DB_POOL_RECYCLE,
  364. logger=logger,
  365. # API中间件统一上报异常,避免数据库客户端同步调用SLS。
  366. aliyun_logr=None,
  367. )
  368. await mysql.init_pool()
  369. query_concurrency = min(settings.API_MAX_CONCURRENT_REQUESTS, settings.DB_POOL_SIZE)
  370. app[MYSQL_KEY] = mysql
  371. app[QUERY_SEMAPHORE_KEY] = asyncio.Semaphore(query_concurrency)
  372. app[LOGGER_KEY] = logger
  373. app[ALIYUN_LOGGER_KEY] = aliyun_logger
  374. app[CLOUD_LOG_QUEUE_KEY] = asyncio.Queue(maxsize=settings.API_LOG_QUEUE_SIZE)
  375. app[CLOUD_LOG_WORKER_KEY] = asyncio.create_task(cloud_log_worker(app))
  376. logger.info(
  377. f'垂直spider API启动: db_pool={settings.DB_POOL_SIZE}, '
  378. f'query_concurrency={query_concurrency}'
  379. )
  380. print_api_routes(app)
  381. yield
  382. queue = app[CLOUD_LOG_QUEUE_KEY]
  383. try:
  384. await asyncio.wait_for(queue.join(), timeout=settings.API_LOG_FLUSH_TIMEOUT)
  385. except asyncio.TimeoutError:
  386. logger.error(f'API关闭时日志队列未完全清空: remaining={queue.qsize()}')
  387. app[CLOUD_LOG_WORKER_KEY].cancel()
  388. try:
  389. await app[CLOUD_LOG_WORKER_KEY]
  390. except asyncio.CancelledError:
  391. pass
  392. else:
  393. await queue.put(None)
  394. await app[CLOUD_LOG_WORKER_KEY]
  395. await mysql.close()
  396. logger.info('垂直spider API已关闭')
  397. def create_app() -> web.Application:
  398. app = web.Application(
  399. middlewares=[
  400. request_context_middleware,
  401. error_middleware,
  402. auth_middleware,
  403. request_body_middleware,
  404. ],
  405. client_max_size=1024 * 1024,
  406. )
  407. app.cleanup_ctx.append(app_context)
  408. app.add_routes([
  409. web.post(API_PATH, query_videos, name='chui_zhi_videos'),
  410. web.get(HEALTH_PATH, health, name='health'),
  411. web.get(READY_PATH, ready, name='ready'),
  412. ])
  413. return app
  414. def main():
  415. web.run_app(
  416. create_app(),
  417. host=settings.API_HOST,
  418. port=settings.API_PORT,
  419. access_log=None,
  420. )
  421. if __name__ == '__main__':
  422. main()