"""MR Snowflake executor.""" from snowflake_connector.etl_connector import SnowflakeSQLExecutor from snowflake_connector.etl_connector import SQLLoader from flows import config # Load SQL templates sql_loader = SQLLoader(__file__) class SnowflakeSQLExecutorMR(SnowflakeSQLExecutor): """MR Snowflake executor.""" def __init__(self): """Snowflake connector init.""" super().__init__( config.SF_CONFIG, autocommit=True) self.tmp_table_name = config.SF_TEMP_TABLE_NAME def drop_tmp_table(self): """Drop temporary import table.""" self.drop_table(self.tmp_table_name) def drop_audit_tmp_table(self): """Drop temporary import audit table.""" self.drop_table(config.SF_AUDIT_TEMP_TABLE_NAME) def create_tmp_table(self): """Create temporary import table.""" sql_template = sql_loader.load_query('create_tmp_table') sql_params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], tmp_table_name=self.tmp_table_name) sql, _ = self.validator.format_identifiers(sql_template, sql_params) self.execute(sql) def create_import_stage(self, s3_path, import_stage_name): """Create a stage for importing from S3 bucket into a temp table. We need to create a stage because we are going to use data transformations during the import process. Args: s3_path (str): S3 path to import the data from import_stage_name (str): import stage name """ sql_template = sql_loader.load_query('create_import_stage') sql_params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], s3_location=s3_path, aws_secret_key=config.AWS_SECRET_ACCESS_KEY, aws_key_id=config.AWS_ACCESS_KEY_ID, import_stage_name=import_stage_name ) sql, non_identifier_params = self.validator.format_identifiers( sql_template, params=sql_params) self.execute(sql, non_identifier_params) def import_from_s3(self): """Import data from S3 bucket into a temp table.""" sql_template = sql_loader.load_query('import_from_s3') sql_params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], tmp_table_name=self.tmp_table_name, import_stage_name=config.SF_IMPORT_STAGE_NAME ) sql, _ = self.validator.format_identifiers(sql_template, sql_params) self.execute(sql) def merge_into_mr_table(self): """Merge log of changes into MR table.""" sql_template = sql_loader.load_query('merge_into_mr_table') sql_params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], tmp_table_name=self.tmp_table_name, mr_table_name=config.SF_MR_TABLE_NAME, grouped_view_name=config.SF_GROUPED_VIEW_NAME ) sql, _ = self.validator.format_identifiers(sql_template, sql_params) self.execute(sql) def create_flattened_mr_table(self): """Create a table with flattened data from MR table.""" sql_template = sql_loader.load_query('create_flattened_mr_table') sql_params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], flattened_mr_table_name=config.SF_MR_FLATTENED_TABLE_NAME ) sql, _ = self.validator.format_identifiers(sql_template, sql_params) self.execute(sql) def fill_flattened_mr_table(self): """Fill flattened MR table with data from MR table.""" sql_template = sql_loader.load_query('fill_flattened_mr_table') sql_params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], flattened_mr_table_name=config.SF_MR_FLATTENED_TABLE_NAME, mr_table_name=config.SF_MR_TABLE_NAME ) sql, _ = self.validator.format_identifiers(sql_template, sql_params) self.execute(sql) def fill_flattened_mr_locked_table(self): """Fill flattened MR table with locked territories.""" sql_params = { 'db': self.sf_config['db'], 'schema': self.sf_config['schema'], 'flattened_mr_locked_table_name': config.SF_MR_LOCKED_FLATTENED_TABLE_NAME, 'mr_table_name': config.SF_MR_TABLE_NAME } self.execute_query( sql_loader, 'fill_flattened_mr_locked_table', sql_params) def import_audit_from_s3(self): """Import audit table data from S3 bucket into a temp table.""" sql_params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], tmp_table_name=config.SF_MR_AUDIT_TABLE_NAME, import_stage_name=config.SF_AUDIT_IMPORT_STAGE_NAME ) self.execute_query(sql_loader, 'import_audit_from_s3', sql_params) def fill_flattened_mr_audit_table(self): """Fill flattened MR_AUDIT table.""" sql_params = { 'db': self.sf_config['db'], 'schema': self.sf_config['schema'], 'flattened_mr_audit_table_name': config.SF_MR_AUDIT_FLATTENED_TABLE_NAME, 'mr_audit_table_name': config.SF_MR_AUDIT_TABLE_NAME } self.execute_query( sql_loader, 'fill_flattened_mr_audit_table', sql_params)