| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117 |
- 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
|