from collections.abc import Sequence import sqlalchemy as sa from email_campaigns.adapters.db import Repository from email_campaigns.campaigns.enums import EmailCampaignStatus from email_campaigns.campaigns.models import CampaignBatch, EmailCampaign from email_campaigns.campaigns.types import CampaignBatchProgress class CampaignBatchRepository(Repository[CampaignBatch]): def get_progress_by_campaign_id( self, campaign_id: str ) -> CampaignBatchProgress | None: query = ( sa.select( CampaignBatch.campaign_id.label("campaign_id"), sa.func.sum(CampaignBatch.batch_size).label("total_batch_size"), sa.func.sum(CampaignBatch.batch_offset).label("total_batch_offset"), ) .where(CampaignBatch.campaign_id == campaign_id) .group_by(CampaignBatch.campaign_id) ) result = self.db.session.execute(query) mapping = result.mappings().one_or_none() if mapping is None: return None return CampaignBatchProgress( campaign_id=mapping["campaign_id"], total_batch_size=mapping["total_batch_size"], total_batch_offset=mapping["total_batch_offset"], ) def find_active_batches(self, domain_id: str) -> Sequence[CampaignBatch]: query = ( sa.select(CampaignBatch) .join(EmailCampaign, CampaignBatch.campaign_id == EmailCampaign.id) .where( CampaignBatch.domain_id == domain_id, EmailCampaign.status == EmailCampaignStatus.IN_PROGRESS, EmailCampaign.cancelled_at.is_(None), EmailCampaign.deleted_at.is_(None), CampaignBatch.completed_at.is_(None), CampaignBatch.cancelled_at.is_(None), ) ) return self.db.session.execute(query).scalars().all() def find_active_batches_for_domains( self, domain_ids: list[str] ) -> Sequence[CampaignBatch]: if not domain_ids: return [] query = ( sa.select(CampaignBatch) .join(EmailCampaign, CampaignBatch.campaign_id == EmailCampaign.id) .where( CampaignBatch.domain_id.in_(domain_ids), EmailCampaign.status == EmailCampaignStatus.IN_PROGRESS, EmailCampaign.cancelled_at.is_(None), EmailCampaign.deleted_at.is_(None), CampaignBatch.completed_at.is_(None), CampaignBatch.cancelled_at.is_(None), ) ) return self.db.session.execute(query).scalars().all()