from collections.abc import Sequence import sqlalchemy as sa from sqlalchemy.orm import selectinload from dmp.adapters.db import Repository from dmp.app_connections.enums import AppConnectionStatus from dmp.tiktok.models import TikTokAdReportingConnection class TikTokAdReportingConnectionRepository(Repository[TikTokAdReportingConnection]): default_options = (selectinload(TikTokAdReportingConnection.ad_accounts),) def get_connection_by_identity_id_and_user_id( self, identity_id: str, user_id: str ) -> TikTokAdReportingConnection | None: query = ( sa.select(TikTokAdReportingConnection) .where( TikTokAdReportingConnection.identity_id == identity_id, TikTokAdReportingConnection.user_id == user_id, ) .options(*self.default_options) ) result = self.db.session.execute(query) return result.scalar_one_or_none() def get_by_identity_id( self, identity_id: str ) -> Sequence[TikTokAdReportingConnection]: query = ( sa.select(TikTokAdReportingConnection) .where(TikTokAdReportingConnection.identity_id == identity_id) .options(*self.default_options) ) result = self.db.session.execute(query) return result.scalars().all() def exists_connection_by_identity_id_and_user_id( self, identity_id: str, user_id: str ) -> bool: query = sa.select( sa.select(1) .exists() .where( sa.and_( TikTokAdReportingConnection.identity_id == identity_id, TikTokAdReportingConnection.user_id == user_id, ), ) ) result = self.db.session.execute(query) return bool(result.scalar_one()) def get_by_fivetran_connection_id( self, fivetran_connection_id: str ) -> TikTokAdReportingConnection | None: query = ( sa.select(TikTokAdReportingConnection) .where( TikTokAdReportingConnection.fivetran_connector_id == fivetran_connection_id ) .options(*self.default_options) ) result = self.db.session.execute(query) return result.scalar_one_or_none() def find_not_notified_connections( self, ) -> list[TikTokAdReportingConnection]: query = ( sa.select(TikTokAdReportingConnection) .where( sa.and_( TikTokAdReportingConnection.status.in_( ( AppConnectionStatus.WAITING_FOR_PROCESSING, AppConnectionStatus.CONNECTED, ) ), TikTokAdReportingConnection.initial_email_sent == sa.false(), ) ) .options(*self.default_options) ) result = self.db.session.execute(query) return list(result.scalars().all()) def get_fivetran_connections_ids(self) -> list[str]: query = sa.select(TikTokAdReportingConnection.fivetran_connector_id).where( TikTokAdReportingConnection.status.in_( AppConnectionStatus.schema_update_statuses() ) ) return list(self.db.session.execute(query).scalars().all())