"""Snowflake connector class for the amazon_unlimited_marketshare tasks.""" from snowflake_connector.etl_connector import SQLLoader from snowflake_connector.etl_connector import SnowflakeSQLExecutor from analytics_aggregation.util import common # Load SQL templates sql_loader = SQLLoader(__file__) class AmazonUnlimitedSF(SnowflakeSQLExecutor): """Helper class to abstract Snowflake operations.""" def _execute_dynamic_sql(self, template_name, sql_params, format_params): """Format loaded template with format params and then execute. Args: template_name (str): Name of the SQL file in queries directory. sql_params (dict): Query params. format_params (dict): Dynamic SQL part of the query. Must be secured code. """ sql_template = sql_loader.load_query(template_name) formated_sql_template = sql_template.format(**format_params) sql, non_identifier_params = self.validator.format_identifiers( formated_sql_template, sql_params) self.execute(sql, params=non_identifier_params) def cleanup_staging_sos(self, date_range, labelids): """Delete from staging_sos Amazon data by passed params. Args: date_range (dict): Dict with 2 keys start_date and end_date. labelids (list[int]): List of label ids. """ sql_template_name = 'delete_from_staging_sos' sql_params = dict( schema=self.sf_config['schema'], start_date=date_range['start_date'], end_date=date_range['end_date'], labelids=labelids) template_format_params = dict( labelids_clause=common.sos_labelid_clause(labelids)) self._execute_dynamic_sql( sql_template_name, sql_params, template_format_params) def populate_staging_sos(self, date_range, labelids): """Populate staging_sos with Amazon data by passed params. Args: date_range (dict): Dict with 2 keys start_date and end_date. labelids (list[int]): List of label ids. """ sql_template_name = 'populate_staging_with_amazon_unlimited_data' sql_params = dict( schema=self.sf_config['schema'], start_date=date_range['start_date'], end_date=date_range['end_date'], labelids=labelids) template_format_params = dict( labelids_clause=common.sos_labelid_clause(labelids, 'dt')) self._execute_dynamic_sql( sql_template_name, sql_params, template_format_params)