from collections.abc import Sequence from typing import Any, cast import sqlalchemy as sa from anydi import singleton from fansifter_common.utils import timezone from pydantic import TypeAdapter from app.adapters.db import Repository from app.enums import CampaignStatus, MessageSendQueueStatus from app.models import AudienceFan, Campaign, MessageSendQueue from app.types import AudienceFanData @singleton class CampaignRepository(Repository[Campaign]): def save_as_in_progress(self, campaign: Campaign) -> None: """Set the given campaign in progress.""" campaign.status = CampaignStatus.IN_PROGRESS self.save(campaign) def save_as_sent(self, campaign: Campaign) -> None: """Set the given campaign as sent.""" campaign.status = CampaignStatus.SENT self.save(campaign) @singleton class MessageSendQueueRepository(Repository[MessageSendQueue]): def get_for_update(self, id: str) -> MessageSendQueue: query = ( sa.select(MessageSendQueue) .where(MessageSendQueue.id == id) .with_for_update() ) result = self.db.session.execute(query) return result.scalar_one() def find_ready(self) -> Sequence[MessageSendQueue]: query = sa.select(MessageSendQueue) query = self._apply_ready_query(query) result = self.db.session.execute(query) return result.scalars().all() def find_ready_by_campaign_ids( self, campaign_ids: set[str] ) -> Sequence[MessageSendQueue]: query = sa.select(MessageSendQueue).where( MessageSendQueue.campaign_id.in_(campaign_ids), MessageSendQueue.total > MessageSendQueue.offset, MessageSendQueue.status.in_( [ MessageSendQueueStatus.QUEUED, MessageSendQueueStatus.PROCESSING, ] ), ) result = self.db.session.execute(query) return result.scalars().all() @staticmethod def _apply_ready_query(query: sa.Select[Any]) -> sa.Select[Any]: return query.where( MessageSendQueue.total > MessageSendQueue.offset, MessageSendQueue.status.in_( [ MessageSendQueueStatus.QUEUED, MessageSendQueueStatus.PROCESSING, ] ), MessageSendQueue.send_at <= timezone.now(), ) AudienceFanDataList = TypeAdapter(list[AudienceFanData]) @singleton class AudienceFanRepository(Repository[AudienceFan]): def count_by_snapshot_id(self, snapshot_id: str) -> int: """Count fans by snapshot id.""" query = self.db.query_from_template( "get-audience-fans-total.sql", context={ "snapshot_id": snapshot_id, "limit": None, "offset": None, }, ) result = self.db.session.execute(query) return cast(int, result.scalar_one()) def find_by_snapshot_id( self, snapshot_id: str, limit: int, offset: int, ) -> list[AudienceFanData]: """Find fan data by snapshot id.""" query = self.db.query_from_template( "get-audience-fans.sql", context={ "snapshot_id": snapshot_id, "limit": limit, "offset": offset, }, ) result = self.db.session.execute(query) return AudienceFanDataList.validate_python(result.mappings())