from collections.abc import Sequence from typing import cast from pydantic import TypeAdapter from dmp.ad_reporting.dtos import ( AdReportingAdSet, AdReportingAdSetCriteria, AdReportingAdSetIdName, AdReportingSummary, ) from dmp.ad_reporting.models import AdReportingAdSetCountryDbt, AdReportingAdSetDbt from dmp.ad_reporting.types import ( AdReportingAdSetCIFields, AdReportingAdSetOrderBy, ) from dmp.adapters.db import ReportingRepository AdReportingAdSetList = TypeAdapter(list[AdReportingAdSet]) AdReportingAdSetIdNameList = TypeAdapter(list[AdReportingAdSetIdName]) class AdReportingAdSetRepository(ReportingRepository[AdReportingAdSetDbt]): def find_by_criteria( self, criteria: AdReportingAdSetCriteria, *, order_by: list[AdReportingAdSetOrderBy] | None = None, limit: int | None = None, offset: int | None = None, ) -> list[AdReportingAdSet]: if criteria.countries: template_name = "ad-reporting/ad-set/find-by-criteria-country.sql" else: template_name = "ad-reporting/ad-set/find-by-criteria.sql" query = self.db.query_from_template( template_name, context={ "criteria": criteria, "order_by": order_by, "ci_fields": AdReportingAdSetCIFields, "limit": limit, "offset": offset, }, ) result = self.db.session.execute(query) return AdReportingAdSetList.validate_python(result.mappings()) def summary_by_criteria( self, criteria: AdReportingAdSetCriteria ) -> AdReportingSummary: if criteria.countries: data_table_name = AdReportingAdSetCountryDbt.__tablename__ else: data_table_name = AdReportingAdSetDbt.__tablename__ query = self.db.query_from_template( "ad-reporting/summary.sql", context={ "criteria": criteria, "id_column": AdReportingAdSetDbt.id.name, "data_id_column": AdReportingAdSetDbt.id.name, "table_name": AdReportingAdSetDbt.__tablename__, "data_table_name": data_table_name, "campaign_id_column": AdReportingAdSetDbt.campaign_id.name, "data_campaign_id_column": AdReportingAdSetDbt.campaign_id.name, "ad_set_id_column": AdReportingAdSetDbt.id.name, "global_participant_id_column": AdReportingAdSetDbt.campaign_global_participant_id.name, "objective_column": AdReportingAdSetDbt.campaign_objective.name, "budget_type_for": "ad_set", }, ) result = self.db.session.execute(query) return AdReportingSummary.model_validate(result.mappings().one()) def count_by_criteria(self, criteria: AdReportingAdSetCriteria) -> int: query = self.db.query_from_template( "ad-reporting/ad-set/count-by-criteria.sql", context={ "criteria": criteria, }, ) result = self.db.session.execute(query) return cast(int, result.scalar_one()) def find_id_names_by_criteria( self, criteria: AdReportingAdSetCriteria, *, limit: int | None = None, offset: int | None = None, ) -> Sequence[AdReportingAdSetIdName]: query = self.db.query_from_template( "ad-reporting/ad-set/find-id-names-by-criteria.sql", context={"criteria": criteria, "limit": limit, "offset": offset}, ) result = self.db.session.execute(query) return AdReportingAdSetIdNameList.validate_python(result.mappings())