| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318 |
- import math
- from datetime import datetime, timedelta, timezone
- from typing import Any, List, Literal, Optional, Tuple
- from fastapi import APIRouter, Request
- from pydantic import Field, PrivateAttr, field_validator, model_validator
- from api.base import ApiParams, BaseApi
- from api.errors import BusinessValidationError
- from config import settings
- router = APIRouter(prefix='/api/v1/crawler/videos', tags=['垂直视频'])
- CHINA_TIMEZONE = timezone(timedelta(hours=8))
- FilterField = Literal[
- 'play_cnt',
- 'like_cnt',
- 'share_cnt',
- 'collection_cnt',
- 'comment_cnt',
- 'duration',
- 'publish_time',
- 'create_time',
- ]
- FilterOperator = Literal['>', '>=', '=', '<', '<=', 'between', 'in', 'not_in']
- SqlFragment = Tuple[str, List[Any]]
- SUPPORTED_PLATFORMS = frozenset({'xiaoniangao', 'xiaoniangaotuijianliu'})
- MAX_FILTER_SET_VALUES = 100
- def normalize_datetime(value: Any, field_name: str) -> datetime:
- """把毫秒时间戳或日期字符串统一转换为东八区无时区时间。"""
- if isinstance(value, bool):
- raise ValueError(f'{field_name}时间格式错误: {value}')
- if isinstance(value, datetime):
- parsed = value
- elif isinstance(value, (int, float)) or (isinstance(value, str) and value.strip().isdigit()):
- try:
- timestamp = float(value)
- if not math.isfinite(timestamp):
- raise ValueError
- if abs(timestamp) >= 10_000_000_000:
- timestamp /= 1000
- parsed = datetime.fromtimestamp(timestamp, tz=CHINA_TIMEZONE)
- except (OverflowError, OSError, ValueError) as exc:
- raise ValueError(f'{field_name}时间格式错误: {value}') from exc
- else:
- try:
- parsed = datetime.fromisoformat(str(value).strip())
- except ValueError as exc:
- raise ValueError(f'{field_name}时间格式错误: {value}') from exc
- if parsed.tzinfo is not None:
- return parsed.astimezone(CHINA_TIMEZONE).replace(tzinfo=None)
- return parsed
- # ==================== 请求参数 ====================
- class FilterCondition(ApiParams):
- """单个结构化筛选条件,字段和操作符均通过白名单限制。"""
- field: FilterField
- operator: FilterOperator
- value: Any
- @model_validator(mode='after')
- def validate_collection_size(self):
- if isinstance(self.value, list) and len(self.value) > MAX_FILTER_SET_VALUES:
- raise ValueError(f'筛选值最多允许{MAX_FILTER_SET_VALUES}项')
- return self
- class PageCursor(ApiParams):
- """稳定翻页游标,并固定默认近3天查询的时间锚点。"""
- id: int = Field(gt=0)
- query_time: Optional[datetime] = None
- @field_validator('query_time', mode='before')
- @classmethod
- def normalize_query_time(cls, value):
- return normalize_datetime(value, 'cursor.query_time') if value is not None else None
- class VideoQueryParams(ApiParams):
- """查询垂直视频的请求参数。"""
- _default_query_end: Optional[datetime] = PrivateAttr(default=None)
- platforms: List[str] = Field(
- default_factory=lambda: ['xiaoniangao', 'xiaoniangaotuijianliu'],
- min_length=1,
- max_length=20,
- )
- keywords: List[str] = Field(default_factory=list, max_length=20)
- filter_match_mode: Literal[1, 2] = 2 # 1=OR,2=AND
- filters: List[FilterCondition] = Field(default_factory=list, max_length=50)
- limit: int = Field(default=500, ge=1, le=settings.API_MAX_LIMIT)
- cursor: Optional[PageCursor] = None
- @field_validator('platforms')
- @classmethod
- def validate_platforms(cls, values: List[str]) -> List[str]:
- normalized = [value.strip() for value in values if value and value.strip()]
- if not normalized:
- raise ValueError('platforms不能为空')
- if any(len(value) > 50 for value in normalized):
- raise ValueError('platform长度不能超过50')
- normalized = list(dict.fromkeys(normalized))
- unsupported = sorted(set(normalized) - SUPPORTED_PLATFORMS)
- if unsupported:
- raise ValueError(f'不支持的平台: {", ".join(unsupported)}')
- return normalized
- @field_validator('keywords')
- @classmethod
- def validate_keywords(cls, values: List[str]) -> List[str]:
- normalized = [value.strip() for value in values if value and value.strip()]
- if values and not normalized:
- raise ValueError('keywords不能全部为空')
- if any(len(value) > 200 for value in normalized):
- raise ValueError('keyword长度不能超过200')
- return normalized
- @model_validator(mode='after')
- def set_default_query_time(self):
- """没有创建时间筛选时使用近3天;翻页时沿用首屏时间锚点。"""
- if not any(condition.field == 'create_time' for condition in self.filters):
- cursor_time = self.cursor.query_time if self.cursor else None
- self._default_query_end = cursor_time or datetime.now(CHINA_TIMEZONE).replace(tzinfo=None)
- return self
- # ==================== 功能实现 ====================
- SELECT_COLUMNS = (
- 'video_id', 'user_id', 'out_user_id', 'platform', 'strategy',
- 'out_video_id', 'video_title', 'cover_url', 'video_url', 'duration',
- 'publish_time', 'play_cnt', 'like_cnt', 'share_cnt', 'collection_cnt',
- 'comment_cnt', 'width', 'height', 'id', 'create_time',
- )
- COMPARISON_OPERATORS = frozenset({'>', '>=', '=', '<', '<='})
- def normalize_filter_value(field: FilterField, value: Any) -> Any:
- """将请求值转换为数据库可比较的数字或日期字符串。"""
- if field in ('publish_time', 'create_time'):
- try:
- return normalize_datetime(value, field).strftime('%Y-%m-%d %H:%M:%S')
- except ValueError as exc:
- raise BusinessValidationError(str(exc)) from exc
- if isinstance(value, bool):
- raise BusinessValidationError(f'{field}筛选值必须是数字: {value}')
- try:
- number = float(value)
- except (TypeError, ValueError) as exc:
- raise BusinessValidationError(f'{field}筛选值必须是数字: {value}') from exc
- if not math.isfinite(number):
- raise BusinessValidationError(f'{field}筛选值必须是有限数字: {value}')
- return int(number) if number.is_integer() else number
- def qualified_column(field: str, table_alias: Optional[str] = None) -> str:
- return f'`{table_alias}`.`{field}`' if table_alias else f'`{field}`'
- def compile_filter_condition(
- condition: FilterCondition,
- table_alias: Optional[str] = None,
- ) -> SqlFragment:
- """将结构化筛选条件转换成参数化SQL。"""
- field, operator, value = condition.field, condition.operator, condition.value
- column = qualified_column(field, table_alias)
- if operator in COMPARISON_OPERATORS:
- if value in (None, '') or isinstance(value, list):
- raise BusinessValidationError(f'{field}的{operator}操作需要单个有效值')
- return f'{column} {operator} %s', [normalize_filter_value(field, value)]
- if not isinstance(value, list):
- raise BusinessValidationError(f'{field}的{operator}操作需要数组值')
- if operator == 'between':
- if len(value) != 2:
- raise BusinessValidationError(f'{field}的between操作必须传两个值')
- normalized_values = [normalize_filter_value(field, item) for item in value]
- if normalized_values[0] > normalized_values[1]:
- raise BusinessValidationError(f'{field}的between起始值不能大于结束值')
- return f'{column} BETWEEN %s AND %s', normalized_values
- if operator in ('in', 'not_in'):
- if not value:
- raise BusinessValidationError(f'{field}的{operator}操作不能为空')
- placeholders = ', '.join(['%s'] * len(value))
- sql_operator = 'IN' if operator == 'in' else 'NOT IN'
- return f'{column} {sql_operator} ({placeholders})', [
- normalize_filter_value(field, item) for item in value
- ]
- raise BusinessValidationError(f'不支持的筛选操作符: {operator}')
- def build_filter_scope(
- params: VideoQueryParams,
- table_alias: Optional[str] = None,
- ) -> SqlFragment:
- """生成可复用的筛选范围,供主查询和去重子查询保持一致。"""
- column = lambda field: qualified_column(field, table_alias)
- platform_placeholders = ', '.join(['%s'] * len(params.platforms))
- clauses = [
- f'{column("platform")} IN ({platform_placeholders})',
- f"{column('video_url')} <> ''",
- ]
- sql_params: List[Any] = [*params.platforms]
- if params._default_query_end is not None:
- clauses.append(f'{column("create_time")} >= %s')
- clauses.append(f'{column("create_time")} < %s')
- sql_params.extend([
- params._default_query_end - timedelta(days=3),
- params._default_query_end,
- ])
- if params.keywords:
- keyword_clause = f'{column("video_title")} LIKE %s'
- clauses.append(f"({' OR '.join([keyword_clause] * len(params.keywords))})")
- sql_params.extend([f'%{keyword}%' for keyword in params.keywords])
- grouped_filters = {}
- for condition in params.filters:
- clause, values = compile_filter_condition(condition, table_alias)
- field_clauses, field_params = grouped_filters.setdefault(condition.field, ([], []))
- field_clauses.append(clause)
- field_params.extend(values)
- # 创建时间是查询范围,始终与其他条件使用AND。
- create_time_group = grouped_filters.pop('create_time', None)
- if create_time_group:
- field_clauses, field_params = create_time_group
- clauses.append(f"({' AND '.join(field_clauses)})")
- sql_params.extend(field_params)
- # 同一字段的上下界必须使用AND;不同字段之间才应用计划配置的AND/OR模式。
- if grouped_filters:
- filter_groups, filter_params = [], []
- for field_clauses, field_params in grouped_filters.values():
- filter_groups.append(f"({' AND '.join(field_clauses)})")
- filter_params.extend(field_params)
- joiner = ' OR ' if params.filter_match_mode == 1 else ' AND '
- clauses.append(f"({joiner.join(filter_groups)})")
- sql_params.extend(filter_params)
- return ' AND '.join(clauses), sql_params
- def build_query(params: VideoQueryParams) -> SqlFragment:
- """先按out_video_id取最大id,再按主键回表读取当前页完整数据。"""
- filter_scope, sql_params = build_filter_scope(params, 'source')
- columns = ', '.join(f'`cv`.`{column}`' for column in SELECT_COLUMNS)
- page_size = min(params.limit, settings.API_MAX_LIMIT)
- having_clause = ''
- if params.cursor:
- # 游标必须作用在分组结果上,不能放入WHERE,否则旧重复记录会在下一页再次出现。
- having_clause = 'HAVING MAX(`source`.`id`) < %s'
- sql_params.append(params.cursor.id)
- sql = f'''SELECT {columns}
- FROM `crawler_video` AS `cv`
- INNER JOIN (
- SELECT MAX(`source`.`id`) AS `selected_id`
- FROM `crawler_video` AS `source`
- WHERE {filter_scope}
- GROUP BY
- CASE WHEN `source`.`out_video_id` = '' THEN `source`.`id` ELSE 0 END,
- `source`.`out_video_id`
- {having_clause}
- ORDER BY `selected_id` DESC
- LIMIT %s
- ) AS `dedup` ON `dedup`.`selected_id` = `cv`.`id`
- ORDER BY `cv`.`id` DESC'''
- # 多取1条判断下一页,避免为每次请求额外执行COUNT(*)。
- sql_params.append(page_size + 1)
- return sql, sql_params
- # ==================== 接口实现 ====================
- def timestamp_ms(value: Optional[datetime]) -> Optional[int]:
- if value is None:
- return None
- if value.tzinfo is None:
- value = value.replace(tzinfo=CHINA_TIMEZONE)
- else:
- value = value.astimezone(CHINA_TIMEZONE)
- return int(value.timestamp() * 1000)
- class VideoQueryApi(BaseApi):
- """管理视频查询的入参、SQL、分页和业务返回数据。"""
- path = '/query'
- async def logic(self, params: VideoQueryParams, request: Request):
- sql, sql_params = build_query(params)
- rows = await self.fetch_all(request, sql, sql_params)
- page_size = min(params.limit, settings.API_MAX_LIMIT)
- has_more = len(rows) > page_size
- rows = rows[:page_size]
- next_cursor = {'id': rows[-1]['id']} if has_more and rows else None
- if next_cursor is not None and params._default_query_end is not None:
- next_cursor['query_time'] = timestamp_ms(params._default_query_end)
- return {
- 'data': rows,
- 'count': len(rows),
- 'has_more': has_more,
- 'next_cursor': next_cursor,
- }
- VideoQueryApi().register(router)
|