videos.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318
  1. import math
  2. from datetime import datetime, timedelta, timezone
  3. from typing import Any, List, Literal, Optional, Tuple
  4. from fastapi import APIRouter, Request
  5. from pydantic import Field, PrivateAttr, field_validator, model_validator
  6. from api.base import ApiParams, BaseApi
  7. from api.errors import BusinessValidationError
  8. from config import settings
  9. router = APIRouter(prefix='/api/v1/crawler/videos', tags=['垂直视频'])
  10. CHINA_TIMEZONE = timezone(timedelta(hours=8))
  11. FilterField = Literal[
  12. 'play_cnt',
  13. 'like_cnt',
  14. 'share_cnt',
  15. 'collection_cnt',
  16. 'comment_cnt',
  17. 'duration',
  18. 'publish_time',
  19. 'create_time',
  20. ]
  21. FilterOperator = Literal['>', '>=', '=', '<', '<=', 'between', 'in', 'not_in']
  22. SqlFragment = Tuple[str, List[Any]]
  23. SUPPORTED_PLATFORMS = frozenset({'xiaoniangao', 'xiaoniangaotuijianliu'})
  24. MAX_FILTER_SET_VALUES = 100
  25. def normalize_datetime(value: Any, field_name: str) -> datetime:
  26. """把毫秒时间戳或日期字符串统一转换为东八区无时区时间。"""
  27. if isinstance(value, bool):
  28. raise ValueError(f'{field_name}时间格式错误: {value}')
  29. if isinstance(value, datetime):
  30. parsed = value
  31. elif isinstance(value, (int, float)) or (isinstance(value, str) and value.strip().isdigit()):
  32. try:
  33. timestamp = float(value)
  34. if not math.isfinite(timestamp):
  35. raise ValueError
  36. if abs(timestamp) >= 10_000_000_000:
  37. timestamp /= 1000
  38. parsed = datetime.fromtimestamp(timestamp, tz=CHINA_TIMEZONE)
  39. except (OverflowError, OSError, ValueError) as exc:
  40. raise ValueError(f'{field_name}时间格式错误: {value}') from exc
  41. else:
  42. try:
  43. parsed = datetime.fromisoformat(str(value).strip())
  44. except ValueError as exc:
  45. raise ValueError(f'{field_name}时间格式错误: {value}') from exc
  46. if parsed.tzinfo is not None:
  47. return parsed.astimezone(CHINA_TIMEZONE).replace(tzinfo=None)
  48. return parsed
  49. # ==================== 请求参数 ====================
  50. class FilterCondition(ApiParams):
  51. """单个结构化筛选条件,字段和操作符均通过白名单限制。"""
  52. field: FilterField
  53. operator: FilterOperator
  54. value: Any
  55. @model_validator(mode='after')
  56. def validate_collection_size(self):
  57. if isinstance(self.value, list) and len(self.value) > MAX_FILTER_SET_VALUES:
  58. raise ValueError(f'筛选值最多允许{MAX_FILTER_SET_VALUES}项')
  59. return self
  60. class PageCursor(ApiParams):
  61. """稳定翻页游标,并固定默认近3天查询的时间锚点。"""
  62. id: int = Field(gt=0)
  63. query_time: Optional[datetime] = None
  64. @field_validator('query_time', mode='before')
  65. @classmethod
  66. def normalize_query_time(cls, value):
  67. return normalize_datetime(value, 'cursor.query_time') if value is not None else None
  68. class VideoQueryParams(ApiParams):
  69. """查询垂直视频的请求参数。"""
  70. _default_query_end: Optional[datetime] = PrivateAttr(default=None)
  71. platforms: List[str] = Field(
  72. default_factory=lambda: ['xiaoniangao', 'xiaoniangaotuijianliu'],
  73. min_length=1,
  74. max_length=20,
  75. )
  76. keywords: List[str] = Field(default_factory=list, max_length=20)
  77. filter_match_mode: Literal[1, 2] = 2 # 1=OR,2=AND
  78. filters: List[FilterCondition] = Field(default_factory=list, max_length=50)
  79. limit: int = Field(default=500, ge=1, le=settings.API_MAX_LIMIT)
  80. cursor: Optional[PageCursor] = None
  81. @field_validator('platforms')
  82. @classmethod
  83. def validate_platforms(cls, values: List[str]) -> List[str]:
  84. normalized = [value.strip() for value in values if value and value.strip()]
  85. if not normalized:
  86. raise ValueError('platforms不能为空')
  87. if any(len(value) > 50 for value in normalized):
  88. raise ValueError('platform长度不能超过50')
  89. normalized = list(dict.fromkeys(normalized))
  90. unsupported = sorted(set(normalized) - SUPPORTED_PLATFORMS)
  91. if unsupported:
  92. raise ValueError(f'不支持的平台: {", ".join(unsupported)}')
  93. return normalized
  94. @field_validator('keywords')
  95. @classmethod
  96. def validate_keywords(cls, values: List[str]) -> List[str]:
  97. normalized = [value.strip() for value in values if value and value.strip()]
  98. if values and not normalized:
  99. raise ValueError('keywords不能全部为空')
  100. if any(len(value) > 200 for value in normalized):
  101. raise ValueError('keyword长度不能超过200')
  102. return normalized
  103. @model_validator(mode='after')
  104. def set_default_query_time(self):
  105. """没有创建时间筛选时使用近3天;翻页时沿用首屏时间锚点。"""
  106. if not any(condition.field == 'create_time' for condition in self.filters):
  107. cursor_time = self.cursor.query_time if self.cursor else None
  108. self._default_query_end = cursor_time or datetime.now(CHINA_TIMEZONE).replace(tzinfo=None)
  109. return self
  110. # ==================== 功能实现 ====================
  111. SELECT_COLUMNS = (
  112. 'video_id', 'user_id', 'out_user_id', 'platform', 'strategy',
  113. 'out_video_id', 'video_title', 'cover_url', 'video_url', 'duration',
  114. 'publish_time', 'play_cnt', 'like_cnt', 'share_cnt', 'collection_cnt',
  115. 'comment_cnt', 'width', 'height', 'id', 'create_time',
  116. )
  117. COMPARISON_OPERATORS = frozenset({'>', '>=', '=', '<', '<='})
  118. def normalize_filter_value(field: FilterField, value: Any) -> Any:
  119. """将请求值转换为数据库可比较的数字或日期字符串。"""
  120. if field in ('publish_time', 'create_time'):
  121. try:
  122. return normalize_datetime(value, field).strftime('%Y-%m-%d %H:%M:%S')
  123. except ValueError as exc:
  124. raise BusinessValidationError(str(exc)) from exc
  125. if isinstance(value, bool):
  126. raise BusinessValidationError(f'{field}筛选值必须是数字: {value}')
  127. try:
  128. number = float(value)
  129. except (TypeError, ValueError) as exc:
  130. raise BusinessValidationError(f'{field}筛选值必须是数字: {value}') from exc
  131. if not math.isfinite(number):
  132. raise BusinessValidationError(f'{field}筛选值必须是有限数字: {value}')
  133. return int(number) if number.is_integer() else number
  134. def qualified_column(field: str, table_alias: Optional[str] = None) -> str:
  135. return f'`{table_alias}`.`{field}`' if table_alias else f'`{field}`'
  136. def compile_filter_condition(
  137. condition: FilterCondition,
  138. table_alias: Optional[str] = None,
  139. ) -> SqlFragment:
  140. """将结构化筛选条件转换成参数化SQL。"""
  141. field, operator, value = condition.field, condition.operator, condition.value
  142. column = qualified_column(field, table_alias)
  143. if operator in COMPARISON_OPERATORS:
  144. if value in (None, '') or isinstance(value, list):
  145. raise BusinessValidationError(f'{field}的{operator}操作需要单个有效值')
  146. return f'{column} {operator} %s', [normalize_filter_value(field, value)]
  147. if not isinstance(value, list):
  148. raise BusinessValidationError(f'{field}的{operator}操作需要数组值')
  149. if operator == 'between':
  150. if len(value) != 2:
  151. raise BusinessValidationError(f'{field}的between操作必须传两个值')
  152. normalized_values = [normalize_filter_value(field, item) for item in value]
  153. if normalized_values[0] > normalized_values[1]:
  154. raise BusinessValidationError(f'{field}的between起始值不能大于结束值')
  155. return f'{column} BETWEEN %s AND %s', normalized_values
  156. if operator in ('in', 'not_in'):
  157. if not value:
  158. raise BusinessValidationError(f'{field}的{operator}操作不能为空')
  159. placeholders = ', '.join(['%s'] * len(value))
  160. sql_operator = 'IN' if operator == 'in' else 'NOT IN'
  161. return f'{column} {sql_operator} ({placeholders})', [
  162. normalize_filter_value(field, item) for item in value
  163. ]
  164. raise BusinessValidationError(f'不支持的筛选操作符: {operator}')
  165. def build_filter_scope(
  166. params: VideoQueryParams,
  167. table_alias: Optional[str] = None,
  168. ) -> SqlFragment:
  169. """生成可复用的筛选范围,供主查询和去重子查询保持一致。"""
  170. column = lambda field: qualified_column(field, table_alias)
  171. platform_placeholders = ', '.join(['%s'] * len(params.platforms))
  172. clauses = [
  173. f'{column("platform")} IN ({platform_placeholders})',
  174. f"{column('video_url')} <> ''",
  175. ]
  176. sql_params: List[Any] = [*params.platforms]
  177. if params._default_query_end is not None:
  178. clauses.append(f'{column("create_time")} >= %s')
  179. clauses.append(f'{column("create_time")} < %s')
  180. sql_params.extend([
  181. params._default_query_end - timedelta(days=3),
  182. params._default_query_end,
  183. ])
  184. if params.keywords:
  185. keyword_clause = f'{column("video_title")} LIKE %s'
  186. clauses.append(f"({' OR '.join([keyword_clause] * len(params.keywords))})")
  187. sql_params.extend([f'%{keyword}%' for keyword in params.keywords])
  188. grouped_filters = {}
  189. for condition in params.filters:
  190. clause, values = compile_filter_condition(condition, table_alias)
  191. field_clauses, field_params = grouped_filters.setdefault(condition.field, ([], []))
  192. field_clauses.append(clause)
  193. field_params.extend(values)
  194. # 创建时间是查询范围,始终与其他条件使用AND。
  195. create_time_group = grouped_filters.pop('create_time', None)
  196. if create_time_group:
  197. field_clauses, field_params = create_time_group
  198. clauses.append(f"({' AND '.join(field_clauses)})")
  199. sql_params.extend(field_params)
  200. # 同一字段的上下界必须使用AND;不同字段之间才应用计划配置的AND/OR模式。
  201. if grouped_filters:
  202. filter_groups, filter_params = [], []
  203. for field_clauses, field_params in grouped_filters.values():
  204. filter_groups.append(f"({' AND '.join(field_clauses)})")
  205. filter_params.extend(field_params)
  206. joiner = ' OR ' if params.filter_match_mode == 1 else ' AND '
  207. clauses.append(f"({joiner.join(filter_groups)})")
  208. sql_params.extend(filter_params)
  209. return ' AND '.join(clauses), sql_params
  210. def build_query(params: VideoQueryParams) -> SqlFragment:
  211. """先按out_video_id取最大id,再按主键回表读取当前页完整数据。"""
  212. filter_scope, sql_params = build_filter_scope(params, 'source')
  213. columns = ', '.join(f'`cv`.`{column}`' for column in SELECT_COLUMNS)
  214. page_size = min(params.limit, settings.API_MAX_LIMIT)
  215. having_clause = ''
  216. if params.cursor:
  217. # 游标必须作用在分组结果上,不能放入WHERE,否则旧重复记录会在下一页再次出现。
  218. having_clause = 'HAVING MAX(`source`.`id`) < %s'
  219. sql_params.append(params.cursor.id)
  220. sql = f'''SELECT {columns}
  221. FROM `crawler_video` AS `cv`
  222. INNER JOIN (
  223. SELECT MAX(`source`.`id`) AS `selected_id`
  224. FROM `crawler_video` AS `source`
  225. WHERE {filter_scope}
  226. GROUP BY
  227. CASE WHEN `source`.`out_video_id` = '' THEN `source`.`id` ELSE 0 END,
  228. `source`.`out_video_id`
  229. {having_clause}
  230. ORDER BY `selected_id` DESC
  231. LIMIT %s
  232. ) AS `dedup` ON `dedup`.`selected_id` = `cv`.`id`
  233. ORDER BY `cv`.`id` DESC'''
  234. # 多取1条判断下一页,避免为每次请求额外执行COUNT(*)。
  235. sql_params.append(page_size + 1)
  236. return sql, sql_params
  237. # ==================== 接口实现 ====================
  238. def timestamp_ms(value: Optional[datetime]) -> Optional[int]:
  239. if value is None:
  240. return None
  241. if value.tzinfo is None:
  242. value = value.replace(tzinfo=CHINA_TIMEZONE)
  243. else:
  244. value = value.astimezone(CHINA_TIMEZONE)
  245. return int(value.timestamp() * 1000)
  246. class VideoQueryApi(BaseApi):
  247. """管理视频查询的入参、SQL、分页和业务返回数据。"""
  248. path = '/query'
  249. async def logic(self, params: VideoQueryParams, request: Request):
  250. sql, sql_params = build_query(params)
  251. rows = await self.fetch_all(request, sql, sql_params)
  252. page_size = min(params.limit, settings.API_MAX_LIMIT)
  253. has_more = len(rows) > page_size
  254. rows = rows[:page_size]
  255. next_cursor = {'id': rows[-1]['id']} if has_more and rows else None
  256. if next_cursor is not None and params._default_query_end is not None:
  257. next_cursor['query_time'] = timestamp_ms(params._default_query_end)
  258. return {
  259. 'data': rows,
  260. 'count': len(rows),
  261. 'has_more': has_more,
  262. 'next_cursor': next_cursor,
  263. }
  264. VideoQueryApi().register(router)