| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105 |
- from __future__ import annotations
- from sqlalchemy import func, select
- from supply_infra.db.models.demand_feedback import DemandFeedback
- from supply_infra.db.repositories.base import BaseRepository
- class DemandFeedbackRepository(BaseRepository[DemandFeedback]):
- """需求、视频与命中内容反馈 repository。"""
- model = DemandFeedback
- def get_by_client_request_id(self, client_request_id: str) -> DemandFeedback | None:
- stmt = select(DemandFeedback).where(
- DemandFeedback.client_request_id == client_request_id
- )
- return self.session.scalars(stmt).first()
- def count_demands(self, demand_grade_ids: list[int]) -> dict[int, int]:
- if not demand_grade_ids:
- return {}
- stmt = (
- select(
- DemandFeedback.demand_grade_id,
- func.count(DemandFeedback.id),
- )
- .where(
- DemandFeedback.target_type == "demand",
- DemandFeedback.demand_grade_id.in_(demand_grade_ids),
- )
- .group_by(DemandFeedback.demand_grade_id)
- )
- return {
- int(demand_grade_id): int(count or 0)
- for demand_grade_id, count in self.session.execute(stmt).all()
- }
- def count_for_demand(
- self,
- demand_grade_id: int,
- ) -> tuple[int, dict[str, int], dict[int, int]]:
- stmt = (
- select(
- DemandFeedback.target_type,
- DemandFeedback.video_id,
- DemandFeedback.demand_video_expansion_id,
- func.count(DemandFeedback.id),
- )
- .where(DemandFeedback.demand_grade_id == int(demand_grade_id))
- .group_by(
- DemandFeedback.target_type,
- DemandFeedback.video_id,
- DemandFeedback.demand_video_expansion_id,
- )
- )
- demand_count = 0
- video_counts: dict[str, int] = {}
- expansion_counts: dict[int, int] = {}
- for target_type, video_id, expansion_id, count in self.session.execute(stmt).all():
- value = int(count or 0)
- if target_type == "demand":
- demand_count += value
- elif target_type == "video" and video_id:
- video_counts[str(video_id)] = value
- elif target_type == "hit_content" and expansion_id is not None:
- expansion_counts[int(expansion_id)] = value
- return demand_count, video_counts, expansion_counts
- def list_for_target(
- self,
- *,
- target_type: str,
- demand_grade_id: int,
- video_id: str | None,
- demand_video_expansion_id: int | None,
- limit: int,
- offset: int,
- ) -> tuple[list[DemandFeedback], int]:
- filters = [
- DemandFeedback.target_type == target_type,
- DemandFeedback.demand_grade_id == int(demand_grade_id),
- ]
- if target_type in {"video", "hit_content"}:
- filters.append(DemandFeedback.video_id == video_id)
- else:
- filters.append(DemandFeedback.video_id.is_(None))
- if target_type == "hit_content":
- filters.append(
- DemandFeedback.demand_video_expansion_id
- == int(demand_video_expansion_id or 0)
- )
- else:
- filters.append(DemandFeedback.demand_video_expansion_id.is_(None))
- total_stmt = select(func.count(DemandFeedback.id)).where(*filters)
- total = int(self.session.scalar(total_stmt) or 0)
- rows_stmt = (
- select(DemandFeedback)
- .where(*filters)
- .order_by(DemandFeedback.created_at.desc(), DemandFeedback.id.desc())
- .offset(offset)
- .limit(limit)
- )
- return list(self.session.scalars(rows_stmt).all()), total
|