demand_feedback_repo.py 3.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105
  1. from __future__ import annotations
  2. from sqlalchemy import func, select
  3. from supply_infra.db.models.demand_feedback import DemandFeedback
  4. from supply_infra.db.repositories.base import BaseRepository
  5. class DemandFeedbackRepository(BaseRepository[DemandFeedback]):
  6. """需求、视频与命中内容反馈 repository。"""
  7. model = DemandFeedback
  8. def get_by_client_request_id(self, client_request_id: str) -> DemandFeedback | None:
  9. stmt = select(DemandFeedback).where(
  10. DemandFeedback.client_request_id == client_request_id
  11. )
  12. return self.session.scalars(stmt).first()
  13. def count_demands(self, demand_grade_ids: list[int]) -> dict[int, int]:
  14. if not demand_grade_ids:
  15. return {}
  16. stmt = (
  17. select(
  18. DemandFeedback.demand_grade_id,
  19. func.count(DemandFeedback.id),
  20. )
  21. .where(
  22. DemandFeedback.target_type == "demand",
  23. DemandFeedback.demand_grade_id.in_(demand_grade_ids),
  24. )
  25. .group_by(DemandFeedback.demand_grade_id)
  26. )
  27. return {
  28. int(demand_grade_id): int(count or 0)
  29. for demand_grade_id, count in self.session.execute(stmt).all()
  30. }
  31. def count_for_demand(
  32. self,
  33. demand_grade_id: int,
  34. ) -> tuple[int, dict[str, int], dict[int, int]]:
  35. stmt = (
  36. select(
  37. DemandFeedback.target_type,
  38. DemandFeedback.video_id,
  39. DemandFeedback.demand_video_expansion_id,
  40. func.count(DemandFeedback.id),
  41. )
  42. .where(DemandFeedback.demand_grade_id == int(demand_grade_id))
  43. .group_by(
  44. DemandFeedback.target_type,
  45. DemandFeedback.video_id,
  46. DemandFeedback.demand_video_expansion_id,
  47. )
  48. )
  49. demand_count = 0
  50. video_counts: dict[str, int] = {}
  51. expansion_counts: dict[int, int] = {}
  52. for target_type, video_id, expansion_id, count in self.session.execute(stmt).all():
  53. value = int(count or 0)
  54. if target_type == "demand":
  55. demand_count += value
  56. elif target_type == "video" and video_id:
  57. video_counts[str(video_id)] = value
  58. elif target_type == "hit_content" and expansion_id is not None:
  59. expansion_counts[int(expansion_id)] = value
  60. return demand_count, video_counts, expansion_counts
  61. def list_for_target(
  62. self,
  63. *,
  64. target_type: str,
  65. demand_grade_id: int,
  66. video_id: str | None,
  67. demand_video_expansion_id: int | None,
  68. limit: int,
  69. offset: int,
  70. ) -> tuple[list[DemandFeedback], int]:
  71. filters = [
  72. DemandFeedback.target_type == target_type,
  73. DemandFeedback.demand_grade_id == int(demand_grade_id),
  74. ]
  75. if target_type in {"video", "hit_content"}:
  76. filters.append(DemandFeedback.video_id == video_id)
  77. else:
  78. filters.append(DemandFeedback.video_id.is_(None))
  79. if target_type == "hit_content":
  80. filters.append(
  81. DemandFeedback.demand_video_expansion_id
  82. == int(demand_video_expansion_id or 0)
  83. )
  84. else:
  85. filters.append(DemandFeedback.demand_video_expansion_id.is_(None))
  86. total_stmt = select(func.count(DemandFeedback.id)).where(*filters)
  87. total = int(self.session.scalar(total_stmt) or 0)
  88. rows_stmt = (
  89. select(DemandFeedback)
  90. .where(*filters)
  91. .order_by(DemandFeedback.created_at.desc(), DemandFeedback.id.desc())
  92. .offset(offset)
  93. .limit(limit)
  94. )
  95. return list(self.session.scalars(rows_stmt).all()), total