"""Abacus Statement Period Persister.""" from typing import List import uuid from sqlalchemy import and_, func, insert, literal, select, union, update from sqlalchemy.orm.session import Session from collaborator.connectors import mysql from collaborator.constants.reports import TriggerType from collaborator.constants.statement_period import StatementPeriodStatus from collaborator.constants.transaction import TransactionType from collaborator.models.rds.collaborator import Collaborator from collaborator.models.rds.dp_payment import DpPayment from collaborator.models.rds.report_run import ReportRun from collaborator.models.rds.statement_period import StatementPeriod from collaborator.models.rds.transaction import Transaction from collaborator.utils.typing import User def _get_transaction_aggregation_subquery( transaction_type: TransactionType, ): """Get aggregation subquery by a given transaction type.""" prefix = transaction_type.value.lower() return ( select( func.count(Transaction.transaction_id).label( f"{prefix}_transactions_count" ), func.sum(Transaction.chargeable_amount).label(f"{prefix}_transactions_sum"), StatementPeriod.abacus_statement_period_id, ) .select_from(Transaction) .join( StatementPeriod, Transaction.statement_period_id == StatementPeriod.statement_period_id, ) .where( and_( Transaction.transaction_type == transaction_type, Transaction.deleted_date == None, # noqa: E711 ) ) .group_by(StatementPeriod.abacus_statement_period_id) .subquery() ) def _get_dp_payment_approval_aggregation_subquery(): """Get payment approvals aggregation subquery from dp_payment table.""" return ( select( func.count(DpPayment.dp_payment_id).label("payment_approvals_count"), func.sum(DpPayment.amount).label("payment_approvals_sum"), DpPayment.abacus_statement_period_id, ) .select_from(DpPayment) .group_by(DpPayment.abacus_statement_period_id) .subquery() ) def _get_dp_payments_aggregation_subquery(select_clause): """Get DP payment aggregation subquery from dp_payment table.""" return ( select_clause.select_from(DpPayment) .join(Collaborator, Collaborator.collaborator_id == DpPayment.collaborator_id) .join( Transaction, Transaction.transaction_id == DpPayment.collaborator_transaction_id, ) .where(DpPayment.collaborator_transaction_id != None) # noqa: E711 .group_by(DpPayment.abacus_statement_period_id) .subquery() ) class AbacusStatementPeriodPersister: """Abacus Statement Period Persister.""" @classmethod @mysql.db_session def get_abacus_statement_periods( cls, abacus_statement_period_ids: List[int], session: Session, ): """Get Abacus statement periods by IDs.""" abacus_statement_periods_subquery = union( *( select(literal(id_).label("abacus_statement_period_id")) for id_ in abacus_statement_period_ids ) ).subquery() report_runs_subquery = ( select( ReportRun.report_run_id, ReportRun.period_ids.label("abacus_statement_period_id"), ) .where( ReportRun.trigger_type == TriggerType.AUTO, ) .group_by(ReportRun.period_ids) # We need to exclude cases where multiple report runs exist for the same period .having(func.count(ReportRun.report_run_id) == 1) .subquery() ) dp_payments_subquery = _get_dp_payments_aggregation_subquery( select( func.count(Transaction.transaction_id).label( "direct_payment_transactions_count" ), func.sum(Transaction.chargeable_amount).label( "direct_payment_transactions_sum" ), DpPayment.abacus_statement_period_id, ) ) dp_payments_counts_subquery = _get_dp_payments_aggregation_subquery( select( func.count(func.distinct(DpPayment.collaborator_id)).label( "direct_payment_transactions_collaborators_count" ), func.count(func.distinct(Collaborator.vendor_id)).label( "direct_payment_transactions_vendors_count" ), DpPayment.abacus_statement_period_id, ) ) payment_approvals_subquery = _get_dp_payment_approval_aggregation_subquery() payment_fees_subquery = _get_transaction_aggregation_subquery( TransactionType.PAYMENT_FEES, ) wht_allocation_subquery = _get_transaction_aggregation_subquery( TransactionType.WHT_ALLOCATION, ) query = ( select( abacus_statement_periods_subquery.c.abacus_statement_period_id, report_runs_subquery.c.report_run_id, payment_fees_subquery.c.payment_fees_transactions_count, payment_fees_subquery.c.payment_fees_transactions_sum, wht_allocation_subquery.c.wht_allocation_transactions_count, wht_allocation_subquery.c.wht_allocation_transactions_sum, dp_payments_subquery.c.direct_payment_transactions_count, dp_payments_subquery.c.direct_payment_transactions_sum, dp_payments_counts_subquery.c.direct_payment_transactions_collaborators_count, dp_payments_counts_subquery.c.direct_payment_transactions_vendors_count, payment_approvals_subquery.c.payment_approvals_count, payment_approvals_subquery.c.payment_approvals_sum, ) .outerjoin( report_runs_subquery, abacus_statement_periods_subquery.c.abacus_statement_period_id == report_runs_subquery.c.abacus_statement_period_id, ) .outerjoin( payment_fees_subquery, abacus_statement_periods_subquery.c.abacus_statement_period_id == payment_fees_subquery.c.abacus_statement_period_id, ) .outerjoin( wht_allocation_subquery, abacus_statement_periods_subquery.c.abacus_statement_period_id == wht_allocation_subquery.c.abacus_statement_period_id, ) .outerjoin( dp_payments_subquery, abacus_statement_periods_subquery.c.abacus_statement_period_id == dp_payments_subquery.c.abacus_statement_period_id, ) .outerjoin( dp_payments_counts_subquery, abacus_statement_periods_subquery.c.abacus_statement_period_id == dp_payments_counts_subquery.c.abacus_statement_period_id, ) .outerjoin( payment_approvals_subquery, abacus_statement_periods_subquery.c.abacus_statement_period_id == payment_approvals_subquery.c.abacus_statement_period_id, ) ) return session.execute(query).all() @classmethod @mysql.db_session def create_dp_payment_transactions( cls, abacus_statement_period_id: int, user: User, session: Session, ): """Create transactions from DP payment entries for an abacus statement.""" creation_batch_uuid = str(uuid.uuid4()) # Create transactions insert_statement = insert(Transaction).from_select( [ Transaction.collaborator_id, Transaction.date, Transaction.transaction_type, Transaction.description, Transaction.original_amount, Transaction.chargeable_amount, Transaction.created_date, Transaction.currency, Transaction.statement_period_id, Transaction.creation_batch_uuid, ], select( DpPayment.collaborator_id, func.current_date(), literal(TransactionType.DIRECT_PAYMENT.value), DpPayment.abacus_statement_period_name + literal(" Balance Payment"), DpPayment.amount * -1, DpPayment.amount * -1, func.now(), DpPayment.currency, StatementPeriod.statement_period_id, literal(creation_batch_uuid), ) .select_from(DpPayment) .join( Collaborator, Collaborator.collaborator_id == DpPayment.collaborator_id, ) .join( StatementPeriod, and_( StatementPeriod.vendor_id == Collaborator.vendor_id, StatementPeriod.status == StatementPeriodStatus.OPEN, ), ) .where( DpPayment.abacus_statement_period_id == abacus_statement_period_id, DpPayment.collaborator_transaction_id.is_(None), ), ) session.execute(insert_statement) # Update DP payments with created transaction IDs update_statement = ( update(DpPayment) .where( and_( DpPayment.abacus_statement_period_id == abacus_statement_period_id, DpPayment.collaborator_id.in_( select(Transaction.collaborator_id).where( Transaction.creation_batch_uuid == creation_batch_uuid ) ), ) ) .values( collaborator_transaction_id=( select(Transaction.transaction_id) .where( and_( Transaction.collaborator_id == DpPayment.collaborator_id, Transaction.creation_batch_uuid == creation_batch_uuid, ) ) .scalar_subquery() ), updated_by=user.id, ) .execution_options(synchronize_session=False) ) session.execute(update_statement) session.commit()