"""Class DBSetup for seeding mock data into temporary schema.""" import os from snowflake_connector.snowflake_conn import execute from snowflake_connector.snowflake_conn import set_default_sessionmaker from tests.consts import snowflake as snowflake_constants CREATE_USER_TEST_SCHEMA = ( 'CREATE TRANSIENT SCHEMA IF NOT EXISTS {schema}'.format( schema=snowflake_constants.UNIQUE_TEST_SCHEMA)) DROP_USER_TEST_SCHEMA = 'DROP SCHEMA IF EXISTS {schema}'.format( schema=snowflake_constants.UNIQUE_TEST_SCHEMA) def execute_sql(*queries): """Run series of queries.""" for query in queries: execute(query) class DBSetup(object): """Class for creating test schema, seeding with data, and tearing down.""" def __enter__(self): """Set up test database.""" self._set_db_session() return self def __exit__(self, type, value, traceback): """Tear down test database.""" self._tear_down() def execute_queries(self, *queries): """Execute SQL scripts. Args: queries (str): SQL statements to run sequentially. """ execute_sql(*queries) def execute_seed_scripts(self, *seed_file_names): """Seed database with data using seed scripts. Args: seed_file_names (str): file names of seeds scripts to run. """ for seed_file_name in seed_file_names: dir_path = os.path.dirname(os.path.realpath(__file__)) file = open( dir_path + '/../seed_scripts/' + seed_file_name, 'r') sql = file.read() execute_sql(sql) def _set_db_session(self): """Create test schema and set to use in default session.""" set_default_sessionmaker({ 'database': snowflake_constants.TEST_DATABASE, 'schema': snowflake_constants.TEST_SCHEMA}) execute_sql(DROP_USER_TEST_SCHEMA, CREATE_USER_TEST_SCHEMA) set_default_sessionmaker({ 'database': snowflake_constants.TEST_DATABASE, 'schema': snowflake_constants.UNIQUE_TEST_SCHEMA}) def _tear_down(self): """Remove test schema (and corresponding tables).""" execute_sql(DROP_USER_TEST_SCHEMA)