"""Database test utilities.""" from functools import wraps import sys from collaborator import config from collaborator.connectors import mysql, snowflake from collaborator.models.rds.collaborator import Collaborator from collaborator.models.rds.dp_payment import DpPayment from collaborator.models.rds.recipient import Recipient from collaborator.models.rds.report import Report from collaborator.models.rds.report_contract import ReportContract from collaborator.models.rds.report_run import ReportRun from collaborator.models.rds.split import Split from collaborator.models.rds.split_type import SplitType from collaborator.models.rds.statement_period import StatementPeriod from collaborator.models.rds.terms_and_conditions import TermsAndConditions from collaborator.models.rds.terms_and_conditions_agreement import ( TermsAndConditionsAgreement, ) from collaborator.models.rds.transaction import Transaction from collaborator.models.rds.transferwise_batch import TransferwiseBatch from collaborator.models.rds.transferwise_profile import TransferwiseProfile from collaborator.models.rds.transferwise_transaction import TransferwiseTransaction from collaborator.models.rds.vendor_agreement import VendorAgreement from collaborator.models.snowflake.account_payment_term import AccountPaymentTerm from collaborator.models.snowflake.label_participant import LabelParticipant from collaborator.models.snowflake.label_participant_product_particpation import ( LabelParticipantProductParticipation, ) from collaborator.models.snowflake.label_participant_track_particpation import ( LabelParticipantTrackParticipation, ) from collaborator.models.snowflake.product import Product from collaborator.models.snowflake.split import Split as SnowflakeSplit from collaborator.models.snowflake.track import Track from tests.testutils.seed.account_payment_term_seed import ( account_payment_term_seed_data, ) from tests.testutils.seed.collaborator_seed import collaborator_seed_data from tests.testutils.seed.dp_payment_seed import dp_payment_seed_data from tests.testutils.seed.label_participant_product_participation_seed_data import ( label_participant_product_participation_seed_data, ) from tests.testutils.seed.label_participant_seed_data import label_participant_seed_data from tests.testutils.seed.label_participant_track_participation_seed_data import ( label_participant_track_participation_seed_data, ) from tests.testutils.seed.product_seed import product_seed_data from tests.testutils.seed.recipient_seed import recipient_seed_data from tests.testutils.seed.report_contract_seed import report_contract_seed_data from tests.testutils.seed.report_run_seed import report_run_seed_data from tests.testutils.seed.report_seed import report_seed_data_rds from tests.testutils.seed.split_seed import split_seed_data from tests.testutils.seed.split_type_seed import split_type_seed_data from tests.testutils.seed.statement_period_seed import statement_period_seed_data from tests.testutils.seed.terms_and_conditions_agreement_seed import ( terms_and_conditions_agreement_seed_data, ) from tests.testutils.seed.terms_and_conditions_seed import ( terms_and_conditions_seed_data, ) from tests.testutils.seed.track_seed import track_seed_data from tests.testutils.seed.transaction_seed import transaction_seed_data from tests.testutils.seed.transferwise_batch_seed import transferwise_batch_seed_data from tests.testutils.seed.transferwise_profile_seed import tw_profile_seed_data from tests.testutils.seed.transferwise_transaction_seed import tw_transaction_seed_data from tests.testutils.seed.vendor_agreement_seed import vendor_agreement_seed_data def _exit_if_not_test_environment(session): """For safety, only run tests in test environment pointed to sqlite. Exit immediately if not in test environment or not pointed to sqlite. """ if config.ENVIRONMENT != config.TEST_ENVIRONMENT: sys.exit(f"Environment must be set to {config.TEST_ENVIRONMENT}.") if "sqlite" not in session.bind.url.drivername: sys.exit("Tests must point to sqlite database.") @mysql.db_session def _create_tables(session): """Create the collaborator and split tables.""" _exit_if_not_test_environment(session) mysql.BaseModel.metadata.create_all(mysql._db_engine) @snowflake.db_session def _create_sf_tables(session): """Create the collaborator and split tables.""" _exit_if_not_test_environment(session) snowflake.BaseModel.metadata.create_all(snowflake._db_engine) @mysql.db_session def _seed_tables(session): """Seed the tables.""" _exit_if_not_test_environment(session) session.bulk_insert_mappings(Recipient, recipient_seed_data) session.bulk_insert_mappings(Collaborator, collaborator_seed_data) session.bulk_insert_mappings(Report, report_seed_data_rds) session.bulk_insert_mappings(ReportContract, report_contract_seed_data) session.bulk_insert_mappings(ReportRun, report_run_seed_data) session.bulk_insert_mappings(SplitType, split_type_seed_data) session.bulk_insert_mappings(Split, split_seed_data) session.bulk_insert_mappings(Transaction, transaction_seed_data) session.bulk_insert_mappings(TransferwiseBatch, transferwise_batch_seed_data) session.bulk_insert_mappings(TransferwiseProfile, tw_profile_seed_data) session.bulk_insert_mappings(TransferwiseTransaction, tw_transaction_seed_data) session.bulk_insert_mappings(VendorAgreement, vendor_agreement_seed_data) session.bulk_insert_mappings(StatementPeriod, statement_period_seed_data) session.bulk_insert_mappings(TermsAndConditions, terms_and_conditions_seed_data) session.bulk_insert_mappings( TermsAndConditionsAgreement, terms_and_conditions_agreement_seed_data ) session.bulk_insert_mappings(DpPayment, dp_payment_seed_data) @snowflake.db_session def _seed_sf_tables(session): """Seed the tables.""" session.bulk_insert_mappings(AccountPaymentTerm, account_payment_term_seed_data) session.bulk_insert_mappings(Product, product_seed_data) session.bulk_insert_mappings(SnowflakeSplit, split_seed_data) session.bulk_insert_mappings(Track, track_seed_data) session.bulk_insert_mappings(LabelParticipant, label_participant_seed_data) session.bulk_insert_mappings( LabelParticipantProductParticipation, label_participant_product_participation_seed_data, ) session.bulk_insert_mappings( LabelParticipantTrackParticipation, label_participant_track_participation_seed_data, ) @mysql.db_session def _drop_tables(session): """Drop all tables.""" _exit_if_not_test_environment(session) mysql.BaseModel.metadata.drop_all(mysql._db_engine) @snowflake.db_session def _drop_sf_tables(session): """Drop all tables.""" _exit_if_not_test_environment(session) snowflake.BaseModel.metadata.drop_all(snowflake._db_engine) @mysql.db_session def get_collaborator(participant_id, account, session): """Get a collaborator based on participant ID/account combo. Args: participant_id (str): participant ID of collaborator to fetch. account (Account): Account which created the collaborator. session (sqlalchemy.orm.session.Session): database session. schema. Returns: Collaborator: existing collaborator or None """ _exit_if_not_test_environment(session) result = ( session.query(Collaborator) .filter_by(participant_id=participant_id, vendor_id=account.id) .first() ) return result.to_dict() if result else None @mysql.db_session def get_subaccount_collaborator(subaccount_id, account, session): """Get a collaborator based on subaccount ID/account combo. Args: subaccount_id (int): Subaccount ID of collaborator to fetch. account (Account): Account which created the collaborator. session (sqlalchemy.orm.session.Session): database session. schema. Returns: Collaborator: existing collaborator or None """ _exit_if_not_test_environment(session) result = ( session.query(Collaborator) .filter_by( subaccount_id=subaccount_id, vendor_id=account.id, collaborator_type="SUBACCOUNT", ) .first() ) return result.to_dict() if result else None @mysql.db_session def get_recipient(recipient_id, session): """Get a recipient based on its ID. Args: recipient_id (int): Recipient ID. session (sqlalchemy.orm.session.Session): Database session. Returns: Recipient: existing recipient or None """ _exit_if_not_test_environment(session) result = session.query(Recipient).filter_by(recipient_id=recipient_id).first() return result.to_dict() if result else None @mysql.db_session def get_transferwise_profiles(profile_id: int, session) -> list: """Get a TransferWise profiles based on the profile ID. Args: profile_id (str): ID of the profiles to fetch. session (sqlalchemy.orm.session.Session): Database session. schema. Returns: list: profiles """ _exit_if_not_test_environment(session) result = session.query(TransferwiseProfile).filter_by(profile_id=profile_id).all() return [item.to_dict() for item in result] def test_schema_default_seed(function): """Create and tear down the test DB schema around a function call. Args: Function (func): the function to be called after creating the test schema. Returns: Function: The decorated function. """ @wraps(function) def wrapper(*args, **kwargs): _create_tables() _create_sf_tables() _seed_tables() _seed_sf_tables() try: function_return = function(*args, **kwargs) finally: _drop_tables() _drop_sf_tables() return function_return return wrapper