from __future__ import annotations from collections.abc import Iterable from typing import Any from sqlalchemy import delete, select from supply_infra.db.models.multi_demand_video_point import MultiDemandVideoPoint from supply_infra.db.repositories.base import BaseRepository from supply_infra.video_points import json_fields_from_point_rows _BATCH_SIZE = 1000 class MultiDemandVideoPointRepository(BaseRepository[MultiDemandVideoPoint]): """需求池视频点位表 — 按 video_id 批量替换与查询。""" model = MultiDemandVideoPoint def list_video_ids_with_points(self, video_ids: Iterable[str]) -> set[str]: """返回 video_ids 中在点位表已有记录的视频 id。""" vid_list = [v for v in video_ids if v] if not vid_list: return set() existing: set[str] = set() for i in range(0, len(vid_list), _BATCH_SIZE): batch = vid_list[i : i + _BATCH_SIZE] stmt = ( select(MultiDemandVideoPoint.video_id) .where(MultiDemandVideoPoint.video_id.in_(batch)) .distinct() ) existing.update( str(v) for v in self.session.scalars(stmt).all() if v ) return existing def list_by_video_ids( self, video_ids: Iterable[str] ) -> dict[str, list[dict[str, Any]]]: """按 video_id 批量查询点位行,返回 video_id → 行列表。""" vid_list = [v for v in video_ids if v] if not vid_list: return {} result: dict[str, list[dict[str, Any]]] = {} for i in range(0, len(vid_list), _BATCH_SIZE): batch = vid_list[i : i + _BATCH_SIZE] stmt = ( select(MultiDemandVideoPoint) .where(MultiDemandVideoPoint.video_id.in_(batch)) .order_by( MultiDemandVideoPoint.video_id, MultiDemandVideoPoint.point_type, MultiDemandVideoPoint.id, ) ) for row in self.session.scalars(stmt).all(): result.setdefault(str(row.video_id), []).append( { "point_type": row.point_type, "point_data": row.point_data, "point_desc": row.point_desc, } ) return result def json_fields_by_video_ids( self, video_ids: Iterable[str] ) -> dict[str, dict[str, str | None]]: """按 video_id 返回三个 JSON 列(API 兼容)。""" rows_by_vid = self.list_by_video_ids(video_ids) return { vid: json_fields_from_point_rows(rows) for vid, rows in rows_by_vid.items() } def replace_for_video_ids( self, points_by_video_id: dict[str, list[dict[str, Any]]] ) -> int: """按 video_id 全量替换点位:先删后插。""" if not points_by_video_id: return 0 video_ids = [v for v in points_by_video_id if v] if not video_ids: return 0 for i in range(0, len(video_ids), _BATCH_SIZE): batch = video_ids[i : i + _BATCH_SIZE] self.session.execute( delete(MultiDemandVideoPoint).where( MultiDemandVideoPoint.video_id.in_(batch) ) ) insert_rows: list[dict[str, Any]] = [] for video_id, rows in points_by_video_id.items(): if not video_id or not rows: continue for row in rows: insert_rows.append( { "video_id": video_id, "point_type": row["point_type"], "point_data": row.get("point_data"), "point_desc": row.get("point_desc"), } ) inserted = 0 for i in range(0, len(insert_rows), _BATCH_SIZE): batch = insert_rows[i : i + _BATCH_SIZE] self.session.bulk_insert_mappings(MultiDemandVideoPoint, batch) inserted += len(batch) return inserted