from collections.abc import Mapping, Sequence from typing import Any, cast import sqlalchemy as sa from fansifter_common.auth.identity import Identity from snowflake.connector import SnowflakeConnection from dmp.adapters.db import ReportingRepository from dmp.audiences.dtos import AudienceCriteria from dmp.audiences.enums import ( AudienceExportReason, AudienceSharePlatform, AudienceTarget, ) from dmp.audiences.models import AudienceFan, AudienceTextFan class AudienceFanRepository(ReportingRepository[AudienceFan]): def get_size(self, criteria: AudienceCriteria) -> int: template_name = ( "audience-fan/get-audience-text-fan-total.sql" if criteria.target == AudienceTarget.TEXT else "audience-fan/get-audience-fan-total.sql" ) query = self.db.query_from_template( template_name, context=criteria.get_context(), ) result = self.db.session.execute(query) return cast(int, result.scalar_one()) def get_ids(self, criteria: AudienceCriteria) -> Sequence[str]: query = self.db.query_from_template( "audience-fan/get-audience-fan-ids.sql", context=criteria.get_context(), ) result = self.db.session.execute(query) return result.scalars().all() def insert_bulk(self, snapshot_id: str, criteria: AudienceCriteria) -> None: query = self.db.query_from_template( "audience-fan/insert-audience-fans.sql", context={ "snapshot_id": snapshot_id, **criteria.get_context(), }, ) self.db.session.execute(query) return def insert_text_audience_fans( self, audience_id: str, criteria: AudienceCriteria ) -> None: query = self.db.query_from_template( "audience-fan/insert-text-audience-fans.sql", context={ "audience_id": audience_id, **criteria.get_context(), }, ) self.db.session.execute(query) return def delete_text_audience_fans(self, audience_id: str) -> None: query = sa.delete(AudienceTextFan).where( AudienceTextFan.audience_id == audience_id ) self.db.session.execute(query) def export_to_csv( self, snapshot_id: str, key: str, reason: AudienceExportReason, stage: str, identity_id: str, ) -> None: dbapi_conn = self.db.session.connection().connection.dbapi_connection if not isinstance(dbapi_conn, SnowflakeConnection): return None conn = cast(SnowflakeConnection, dbapi_conn) with conn.cursor() as cur: query, bind_params = self.db.jinja2sql.from_file( "audience-fan/export-audience-fans-to-csv.sql", context={ "snapshot_id": snapshot_id, "reason": reason, "key": key, "stage": stage, "identity_id": identity_id, }, param_style="pyformat", ) cur.execute(query, dict(bind_params)) def get_export_fan_data( self, snapshot_id: str, reason: AudienceExportReason, identity: Identity, ) -> list[Mapping[str, Any]]: """Get fan data for export.""" query = self.db.query_from_template( "audience-fan/get-audience-export-fans.sql", context={ "snapshot_id": snapshot_id, "reason": reason, "identity_id": identity.id, }, ) return cast( list[Mapping[str, Any]], self.db.session.execute(query).mappings().all(), ) def export_share_to_csv( self, snapshot_id: str, key: str, platform: AudienceSharePlatform, stage: str, context: dict[str, Any] | None = None, ) -> None: dbapi_conn = self.db.session.connection().connection.dbapi_connection if not isinstance(dbapi_conn, SnowflakeConnection): return None conn = cast(SnowflakeConnection, dbapi_conn) with conn.cursor() as cur: query, bind_params = self.db.jinja2sql.from_file( "audience-fan/export-audience-share-fans-to-csv.sql", context={ "snapshot_id": snapshot_id, "platform": platform, "key": key, "stage": stage, "context": context or {}, }, param_style="pyformat", ) cur.execute(query, dict(bind_params)) def get_share_fan_data( self, snapshot_id: str, platform: AudienceSharePlatform, context: dict[str, Any] | None = None, ) -> list[Mapping[str, Any]]: """Get fan data for sharing.""" query = self.db.query_from_template( "audience-fan/get-audience-share-fans.sql", context={ "snapshot_id": snapshot_id, "platform": platform, "context": context or {}, }, ) return cast( list[Mapping[str, Any]], self.db.session.execute(query).mappings().all(), )