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