multi_demand_video_point_repo.py 4.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117
  1. from __future__ import annotations
  2. from collections.abc import Iterable
  3. from typing import Any
  4. from sqlalchemy import delete, select
  5. from supply_infra.db.models.multi_demand_video_point import MultiDemandVideoPoint
  6. from supply_infra.db.repositories.base import BaseRepository
  7. from supply_infra.video_points import json_fields_from_point_rows
  8. _BATCH_SIZE = 1000
  9. class MultiDemandVideoPointRepository(BaseRepository[MultiDemandVideoPoint]):
  10. """需求池视频点位表 — 按 video_id 批量替换与查询。"""
  11. model = MultiDemandVideoPoint
  12. def list_video_ids_with_points(self, video_ids: Iterable[str]) -> set[str]:
  13. """返回 video_ids 中在点位表已有记录的视频 id。"""
  14. vid_list = [v for v in video_ids if v]
  15. if not vid_list:
  16. return set()
  17. existing: set[str] = set()
  18. for i in range(0, len(vid_list), _BATCH_SIZE):
  19. batch = vid_list[i : i + _BATCH_SIZE]
  20. stmt = (
  21. select(MultiDemandVideoPoint.video_id)
  22. .where(MultiDemandVideoPoint.video_id.in_(batch))
  23. .distinct()
  24. )
  25. existing.update(
  26. str(v) for v in self.session.scalars(stmt).all() if v
  27. )
  28. return existing
  29. def list_by_video_ids(
  30. self, video_ids: Iterable[str]
  31. ) -> dict[str, list[dict[str, Any]]]:
  32. """按 video_id 批量查询点位行,返回 video_id → 行列表。"""
  33. vid_list = [v for v in video_ids if v]
  34. if not vid_list:
  35. return {}
  36. result: dict[str, list[dict[str, Any]]] = {}
  37. for i in range(0, len(vid_list), _BATCH_SIZE):
  38. batch = vid_list[i : i + _BATCH_SIZE]
  39. stmt = (
  40. select(MultiDemandVideoPoint)
  41. .where(MultiDemandVideoPoint.video_id.in_(batch))
  42. .order_by(
  43. MultiDemandVideoPoint.video_id,
  44. MultiDemandVideoPoint.point_type,
  45. MultiDemandVideoPoint.id,
  46. )
  47. )
  48. for row in self.session.scalars(stmt).all():
  49. result.setdefault(str(row.video_id), []).append(
  50. {
  51. "point_type": row.point_type,
  52. "point_data": row.point_data,
  53. "point_desc": row.point_desc,
  54. }
  55. )
  56. return result
  57. def json_fields_by_video_ids(
  58. self, video_ids: Iterable[str]
  59. ) -> dict[str, dict[str, str | None]]:
  60. """按 video_id 返回三个 JSON 列(API 兼容)。"""
  61. rows_by_vid = self.list_by_video_ids(video_ids)
  62. return {
  63. vid: json_fields_from_point_rows(rows)
  64. for vid, rows in rows_by_vid.items()
  65. }
  66. def replace_for_video_ids(
  67. self, points_by_video_id: dict[str, list[dict[str, Any]]]
  68. ) -> int:
  69. """按 video_id 全量替换点位:先删后插。"""
  70. if not points_by_video_id:
  71. return 0
  72. video_ids = [v for v in points_by_video_id if v]
  73. if not video_ids:
  74. return 0
  75. for i in range(0, len(video_ids), _BATCH_SIZE):
  76. batch = video_ids[i : i + _BATCH_SIZE]
  77. self.session.execute(
  78. delete(MultiDemandVideoPoint).where(
  79. MultiDemandVideoPoint.video_id.in_(batch)
  80. )
  81. )
  82. insert_rows: list[dict[str, Any]] = []
  83. for video_id, rows in points_by_video_id.items():
  84. if not video_id or not rows:
  85. continue
  86. for row in rows:
  87. insert_rows.append(
  88. {
  89. "video_id": video_id,
  90. "point_type": row["point_type"],
  91. "point_data": row.get("point_data"),
  92. "point_desc": row.get("point_desc"),
  93. }
  94. )
  95. inserted = 0
  96. for i in range(0, len(insert_rows), _BATCH_SIZE):
  97. batch = insert_rows[i : i + _BATCH_SIZE]
  98. self.session.bulk_insert_mappings(MultiDemandVideoPoint, batch)
  99. inserted += len(batch)
  100. return inserted