"""Payment Group Payment model.""" from abacus_common_logic.models.base import BaseModel, db from sqlalchemy import and_, func, literal_column, or_, select, table from payment.constants.constants import ( PAYMENT_GROUP_PAYMENT_ACTION_STATUSES, PAYMENT_GROUP_PAYMENT_ACTIONS, ) from payment.models.payment_group_payment_account import PaymentGroupPaymentAccount class PaymentGroupPayment(BaseModel): """Payment Group Payment model.""" __tablename__ = 'payment_group_payment' payment_group_payment_id = db.Column(db.Integer, primary_key=True) payment_group_id = db.Column( db.Integer, db.ForeignKey('payment_group.payment_group_id'), nullable=False ) statement_period_id = db.Column(db.Integer, nullable=False) payment_name = db.Column(db.String(180), nullable=False) deleted_at = db.Column(db.DateTime, nullable=True) deleted_by = db.Column(db.String(255), nullable=True) reports = db.relationship( 'ReportPaymentGroupPayment', backref='payment_group_payment', cascade='all, delete-orphan', ) payment_accounts = db.relationship( 'PaymentGroupPaymentAccount', backref='payment_group_payment', cascade='all, delete-orphan', ) payment_batches = db.relationship( 'PaymentGroupPaymentBatch', backref='payment_group_payment', cascade='all, delete-orphan', ) @property def account_count(self): """Return the number of accounts associated with the payment group payment.""" return len(self.payment_accounts) @classmethod def base_list_query(cls): """Query to get list of payment_group_payments that have not been deleted.""" return ( select( cls.created_at, cls.payment_group_id, cls.payment_group_payment_id, cls.payment_name, func.coalesce( func.count(PaymentGroupPaymentAccount.payment_group_payment_id), 0 ).label('account_count'), ) .outerjoin( PaymentGroupPaymentAccount, and_( PaymentGroupPaymentAccount.payment_group_payment_id == cls.payment_group_payment_id, PaymentGroupPaymentAccount.deleted_by.is_(None), PaymentGroupPaymentAccount.prior_payment_group_payment_id.is_(None), ), ) .where(cls.deleted_at.is_(None)) .group_by( cls.created_at, cls.payment_group_id, cls.payment_group_payment_id, cls.payment_name, ) .order_by(cls.created_at.desc()) ) @classmethod def get_paginated(cls, limit: int, offset: int): """Return (items, total_count) for the payment_group_payment list.""" stmt = cls.base_list_query() items = db.session.execute(stmt.offset(offset).limit(limit)).mappings().all() count = db.session.execute( select(func.count()).select_from(stmt.subquery()) ).scalar_one() return items, count @classmethod def find_by_name(cls, payment_name): """Override base model's find_by_name.""" return ( db.session.execute(select(cls).where(cls.payment_name == payment_name)) .scalars() .first() ) def is_sent(self): """Check if a payment_group_payment's send_payments action is complete.""" query = ( select(literal_column('1')) .where( and_( literal_column('parent_table_name') == self.__tablename__, literal_column('parent_table_id') == self.payment_group_payment_id, literal_column('action_name') == PAYMENT_GROUP_PAYMENT_ACTIONS.SEND_PAYMENTS, or_( literal_column('action_status') == PAYMENT_GROUP_PAYMENT_ACTION_STATUSES.COMPLETE, literal_column('action_status') == PAYMENT_GROUP_PAYMENT_ACTION_STATUSES.RUNNING, ), ) ) .select_from(table('abacus_state')) ) posted_status = db.session.execute(query).scalars().all() return bool(posted_status)