from collections.abc import Sequence import sqlalchemy as sa from sqlalchemy.orm import selectinload from dmp.adapters.db import Repository from dmp.tiktok.models import TikTokUserConnection class TikTokUserConnectionRepository(Repository[TikTokUserConnection]): default_options = [selectinload(TikTokUserConnection.ad_accounts)] def get_by_identity_id(self, identity_id: str) -> Sequence[TikTokUserConnection]: query = ( sa.select(TikTokUserConnection) .where(TikTokUserConnection.identity_id == identity_id) .options(*self.default_options) ) result = self.db.session.execute(query) return result.scalars().all() def invalidate(self, user_connection: TikTokUserConnection) -> None: user_connection.is_valid = False def get_by_identity_id_and_user_id( self, identity_id: str, user_id: str ) -> TikTokUserConnection | None: query = ( sa.select(TikTokUserConnection) .where( TikTokUserConnection.identity_id == identity_id, TikTokUserConnection.user_id == user_id, ) .options(*self.default_options) ) result = self.db.session.execute(query) return result.scalar_one_or_none() def find_active_by_audience_and_ad_account_id( self, *, audience_id: str, ad_account_id: str, identity_id: str, vendor_ids: list[int] | None = None, subaccount_ids: list[int] | None = None, ) -> Sequence[TikTokUserConnection]: query = self.db.query_from_template( "tiktok-user-connection/find-active-by-audience-id-and-ad-account-id.sql", context={ "audience_id": audience_id, "ad_account_id": ad_account_id, "identity_id": identity_id, "vendor_ids": vendor_ids, "subaccount_ids": subaccount_ids, }, ) result = self.db.session.execute( sa.select(TikTokUserConnection) .from_statement(query) .options(*self.default_options) ) return result.scalars().all()