"""Transaction Persister. Handles CRUD operations on the transaction table. """ from datetime import date, datetime, timezone from decimal import Decimal from typing import List, Optional, Tuple from sqlalchemy import Row, case, func, insert, select, tuple_ from sqlalchemy.orm.session import Session from sqlalchemy.sql.functions import coalesce from collaborator.connectors import mysql from collaborator.constants import error from collaborator.constants.statement_period import StatementPeriodStatus from collaborator.constants.transaction import TYPE_CREDIT, TYPE_REVENUE from collaborator.models.rds.collaborator import Collaborator from collaborator.models.rds.report import Report 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.models.rds.transferwise_transaction import TransferwiseTransaction from collaborator.utils import logging from collaborator.utils.error import OwsError class TransactionPersister: """Handles high level operations for transactions.""" @classmethod @mysql.db_session def create_transaction( cls, collaborator_id: int, transaction_type: str, transaction_date: Optional[date], description: Optional[str], original_amount: Decimal, collaborator_share: Optional[float], chargeable_amount: Decimal, statement_period_id: int, transferwise_transaction_id: Optional[int], report_id: Optional[int], voided_transaction_id: Optional[int], currency: Optional[str], session: Session, credited_payment_id: Optional[int] = None, ) -> dict: """Create a transaction. Args: collaborator_id (int): ID of the collaborator transaction_type (str): Type of transaction transaction_date (date): Date the transaction was made description (str): Description of the transaction original_amount (Decimal): Original amount collaborator_share (float): Collaborator share chargeable_amount (Decimal): Chargeable amount transferwise_transaction_id (int): ID of a TW transaction report_id (int): ID of a report voided_transaction_id (int): Transaction to void currency (str): Currency represented with the ISO 4217 code credited_payment_id (int): ID of the original transaction being reversed Returns: dict: the newly created transaction """ transaction = Transaction( collaborator_id=collaborator_id, date=transaction_date, transaction_type=transaction_type, description=description, original_amount=original_amount, collaborator_share=collaborator_share, chargeable_amount=chargeable_amount, statement_period_id=statement_period_id, transferwise_transaction_id=transferwise_transaction_id, report_id=report_id, voided_transaction_id=voided_transaction_id, currency=currency, credited_payment_id=credited_payment_id, ) session.add(transaction) session.flush() return transaction.to_dict() @classmethod @mysql.db_session def create_transactions( cls, transactions_data: list, creation_batch_uuid: str, session: Session ) -> list: """Create multiple transactions. Args: transactions_data (list): List of transactions to create created (datetime): The datetime of txn creation Returns: list: the newly created transactions """ session.execute(insert(Transaction.__table__), transactions_data) session.flush() rows = ( session.query(Transaction) .filter( Transaction.creation_batch_uuid == creation_batch_uuid, ) .all() ) result = [row.to_dict() for row in rows] return result @classmethod @mysql.db_session def get_transactions( cls, session: Session, collaborator_id: Optional[int] = None, statement_period_id: Optional[int] = None, transaction_ids: Optional[List[int]] = None, limit: int = 0, offset: int = 0, ) -> Tuple[list, int]: """Get all the transactions for a collaborator. Args: collaborator_id (int): The collaborator's unique identifier limit (int): how many transactions to retrieve. offset (int): the offset (for pagination). Returns: [dict]: with the all the collaborator's transactions. """ filters = [ Transaction.deleted_date == None, # noqa E711 ] if collaborator_id: filters.append(Transaction.collaborator_id == collaborator_id) if statement_period_id: filters.append(Transaction.statement_period_id == statement_period_id) if transaction_ids: filters.append(Transaction.transaction_id.in_(transaction_ids)) query = ( session.query( Transaction, TransferwiseTransaction.transferwise_batch_id, TransferwiseTransaction.status, ) .outerjoin( TransferwiseTransaction, TransferwiseTransaction.transaction_id == Transaction.transferwise_transaction_id, ) .filter(*filters) .order_by( Transaction.created_date.desc(), Transaction.date.desc(), ) ) total_records = query.count() if limit != 0: limited_query = query.limit(limit).offset(offset) else: limited_query = query.offset(offset) results = limited_query.all() transactions = [ { **t.Transaction.to_dict(), "transferwise_transaction_status": t.status, "transferwise_batch_id": t.transferwise_batch_id, } for t in results ] return transactions, total_records @classmethod @mysql.db_session def latest_transaction( cls, collaborator_id: int, session: Session ) -> Optional[dict]: """Get the latest transaction for a collaborator. Args: collaborator_id (int): ID of the collaborator. session (Session): Database session. Returns: dict: The transaction. """ result = ( session.query(Transaction) .filter_by(collaborator_id=collaborator_id, deleted_date=None) .order_by(Transaction.transaction_id.desc()) .first() ) return result.to_dict() if result else None @classmethod @mysql.db_session def get_transaction_by_transferwise_transaction_id( cls, transferwise_transaction_id: int, session: Session ): """Get Transaction by transaction_id id. Args: transferwise_transaction_id (int): TW Transaction ID session (sqlalchemy.orm.session.Session): database session. Returns: dict: Transaction from Transaction Table """ query = session.query(Transaction) query = query.filter_by( transferwise_transaction_id=transferwise_transaction_id, deleted_date=None ) results = query.first() return results.to_dict() if results else None @classmethod @mysql.db_session def get_by_id( cls, transaction_id: int, session: Session, include_deleted: bool = False, ): """Get a single transaction based on the ID. Args: transaction_id (int): ID of the transaction to get. session (sqlalchemy.orm.session.Session): database session. Returns: Response: a transaction """ filters = [Transaction.transaction_id == transaction_id] if not include_deleted: filters.append(Transaction.deleted_date == None) # noqa E711 result = session.query(Transaction).filter(*filters).first() return result.to_dict() if result else None @classmethod @mysql.db_session def get_collaborator_by_transaction_id(cls, transaction_id: int, session: Session): """Get a single transaction based on the ID. Args: transaction_id (int): ID of the transaction to get. session (sqlalchemy.orm.session.Session): database session. Returns: Response: """ result = ( session.query(Collaborator) .join(Transaction) .filter_by(transaction_id=transaction_id, deleted_date=None) .first() ) return result.to_dict() if result else None @classmethod @mysql.db_session def check_void_transaction(cls, transaction_id: int, session: Session): """Check if a transaction is voided. Args: transaction_id (int): Transaction ID session (sqlalchemy.orm.session.Session): database session. Returns: bool: If transaction voided """ query = session.query(Transaction) query = query.filter_by(voided_transaction_id=transaction_id) result = query.first() if not result: return False return True @classmethod @mysql.db_session def delete_by_ids(cls, transaction_ids: List[int], session: Session): """Soft-delete a single transaction based on the ID. Args: transaction_id (int): ID of the transaction to get. session (sqlalchemy.orm.session.Session): database session. Returns: Response: a transaction """ transactions_base_query = session.query(Transaction).filter( Transaction.transaction_id.in_(transaction_ids), Transaction.deleted_date == None, # noqa E711 ) transactions = transactions_base_query.all() if not transactions or len(transactions) != len(transaction_ids): raise OwsError.not_found( code=error.ERROR_CODE_TRANSACTION_NOT_FOUND, message=error.ERROR_MESSAGE_TRANSACTION_NOT_FOUND, ) if any( transaction.transaction_type == TYPE_CREDIT for transaction in transactions ): raise OwsError.forbidden( code=error.ERROR_CODE_TRANSACTION_FORBIDDEN, message=error.ERROR_MESSAGE_TRANSACTION_FORBIDDEN, ) statement_period_ids = list( {transaction.statement_period_id for transaction in transactions} ) statement_periods = ( session.query(StatementPeriod) .filter( StatementPeriod.statement_period_id.in_(statement_period_ids), StatementPeriod.status == StatementPeriodStatus.OPEN, ) .all() ) if not statement_periods or len(statement_periods) != len(statement_period_ids): raise OwsError.not_found( code=error.ERROR_CODE_CANNOT_DELETE_CLOSED_TRANSACTION, message=error.ERROR_MESSAGE_CANNOT_DELETE_CLOSED_TRANSACTION, ) original_transaction_dicts = { transaction.transaction_id: transaction.to_dict() for transaction in transactions } transactions_base_query.update({"deleted_date": datetime.now(timezone.utc)}) session.commit() updated_transaction_dicts = { transaction.transaction_id: transaction.to_dict() for transaction in transactions } logging.bulk_log_events( event_type=logging.LOG_EVENT_DELETE, entity_type="transaction", data=[ { "id": transaction_id, "original": original_transaction_dicts[transaction_id], "updated": updated_transaction_dicts[transaction_id], } for transaction_id in transaction_ids ], user=None, ) @classmethod @mysql.db_session def is_transaction_credited(cls, payment_id: int, session: Session): """Check if a given (failed) payment is already credited. Args: payment_id (int): ID of the transaction to get. session (sqlalchemy.orm.session.Session): database session. Returns: Response: True if the txn is already credited. """ query = session.query(Transaction).filter_by( transferwise_transaction_id=payment_id, transaction_type=TYPE_CREDIT, ) result = query.first() return False if not result else True @classmethod @mysql.db_session def is_credited_by_credited_payment_id( cls, transaction_id: int, session: Session ) -> bool: """Check if a transaction has already been reversed via a credit. Used for idempotency when creating TYPE_CREDIT transactions for failed/cancelled Payoneer payments. Args: transaction_id (int): ID of the original transaction to check. session (sqlalchemy.orm.session.Session): database session. Returns: bool: True if a credit already exists for this transaction. """ result = ( session.query(Transaction) .filter_by( credited_payment_id=transaction_id, transaction_type=TYPE_CREDIT, ) .first() ) return result is not None @classmethod @mysql.db_session def get_transactions_for_statement_period( cls, statement_period_id: int, limit: int, offset: int, session: Session ) -> Tuple[list, int]: """Get all the transactions for a collaborator. Args: collaborator_id (int): The collaborator's unique identifier limit (int): how many transactions to retrieve. offset (int): the offset (for pagination). Returns: [dict]: with the all the collaborator's transactions. """ filters = [ Transaction.statement_period_id == statement_period_id, Transaction.deleted_date == None, # noqa E711 ] query = ( session.query(Transaction) .filter(*filters) .order_by( Transaction.created_date.desc(), Transaction.date.desc(), ) ) total_records = query.count() if limit != 0: limited_query = query.limit(limit).offset(offset) else: limited_query = query.offset(offset) results = limited_query.all() transactions = [t.to_dict() for t in results] return transactions, total_records @classmethod @mysql.db_session def get_transactions_for_open_period( cls, vendor_id: int, session: Session ) -> Tuple[list, int]: """Get the transactions for the open period. Args: vendor_id (int): The vendor ID. Returns: [list, int]: The transactions and txn count for the open period. """ query = ( session.query(Transaction) .join( StatementPeriod, Transaction.statement_period_id == StatementPeriod.statement_period_id, ) .join( Collaborator, Transaction.collaborator_id == Collaborator.collaborator_id, ) .filter( StatementPeriod.status == StatementPeriodStatus.OPEN, Collaborator.vendor_id == vendor_id, Transaction.deleted_date == None, # noqa E711 ) ) total_records = query.count() results = query.all() transactions = [t.to_dict() for t in results] return transactions, total_records @classmethod @mysql.db_session def get_transactions_by_report_ids(cls, report_ids: list, session: Session) -> list: """Get transactions by report identifiers. Args: report_ids (list): list with report_ids session (sqlalchemy.orm.session.Session): database session. Returns: list: with transactions """ filters = [ Transaction.report_id.in_(report_ids), Transaction.deleted_date == None, # noqa E711 ] query = session.query(Transaction).filter(*filters) result = query.all() return [item.to_dict() for item in result] @classmethod @mysql.db_session def get_vendor_ids_for_transactions( cls, transaction_ids: list[int], session: Session ) -> list[int]: """Get vendor ids for transactions. Args: transaction_ids (list[int]): list of transaction ids session (sqlalchemy.orm.session.Session): database session. Returns: list[int]: with vendor ids """ query = ( select(Collaborator.vendor_id) .select_from(Transaction) .filter( Transaction.transaction_id.in_(transaction_ids), Transaction.deleted_date == None, # noqa E711 ) .join( Collaborator, Collaborator.collaborator_id == Transaction.collaborator_id, ) ) vendor_ids = session.execute(query).scalars().all() return list(vendor_ids) @classmethod @mysql.db_session def get_aggregations_for_report_run( cls, report_run_id: int, session: Session, collaborator_dp_enabled: bool, ) -> Row[Tuple[int, int, Decimal, datetime]]: """Get transaction aggregations for a report run. Args: report_run_id (int): The report run ID. session (sqlalchemy.orm.session.Session): database session. collaborator_dp_enabled (bool): Filter for report.collaborator_dp_enabled. Returns: aggregations result. """ filters = [ Transaction.transaction_type == TYPE_REVENUE, ReportRun.report_run_id == report_run_id, ] if collaborator_dp_enabled: filters.append( Report.collaborator_dp_enabled == True, # noqa E711 ) query = ( select( func.count().label("total_count"), func.count(case((Transaction.chargeable_amount != 0, 1))).label( "non_zero_count" ), coalesce(func.sum(Transaction.chargeable_amount), 0).label( "currency_agnostic_total" ), func.min(Transaction.created_date).label("completed_date"), ) .join( Report, Report.report_id == Transaction.report_id, ) .join( ReportRun, ReportRun.report_run_id == Report.report_run_id, ) .filter(*filters) ) result = session.execute(query).one() return result @classmethod @mysql.db_session def get_transactions_for_participations( cls, participations: list, session: Session, ) -> dict: """Fetch transactions for multiple (collaborator_id, statement_period_id) pairs. Args: participations (list): List of (collaborator_id, statement_period_id) tuples. session (Session): Database session. Returns: dict: Mapping of (collaborator_id, statement_period_id) to {transactions, total_count}. """ results = ( session.query( Transaction, TransferwiseTransaction.transferwise_batch_id, TransferwiseTransaction.status, ) .outerjoin( TransferwiseTransaction, TransferwiseTransaction.transaction_id == Transaction.transferwise_transaction_id, ) .filter( Transaction.deleted_date == None, # noqa E711 tuple_( Transaction.collaborator_id, Transaction.statement_period_id ).in_(participations), ) .order_by( Transaction.created_date.desc(), Transaction.date.desc(), ) .all() ) grouped: dict = {p: [] for p in participations} for row in results: key = ( row.Transaction.collaborator_id, row.Transaction.statement_period_id, ) grouped[key].append( { **row.Transaction.to_dict(), "transferwise_transaction_status": row.status, "transferwise_batch_id": row.transferwise_batch_id, } ) return { key: {"transactions": txns, "total_count": len(txns)} for key, txns in grouped.items() }