"""Snowflake connector class.""" from datetime import datetime from snowflake_connector.etl_connector import SnowflakeSQLExecutor from snowflake_connector.etl_connector import SQLLoader # Load SQL templates sql_loader = SQLLoader(__file__) class SnowflakeSQLExecutorSME(SnowflakeSQLExecutor): """Helper class to abstract Snowflake operations.""" def get_temp_table_name(self): return 'temp_shared_products_{:%Y%m%d%H%M%S}'.format(datetime.now()) def add_to_shared_products(self, upc_list): """Add new upcs to sme share. Args: upc_list (list): List of new upcs to share. """ def create_temp_share_products_table(temp_table_name): """Create temp share products table. Args: temp_table_name (str): Temp table name. """ sql_template = sql_loader.load_query( 'create_temp_share_products_table') params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], temp_table_name=temp_table_name) sql_template, non_identifier_params = ( self.validator.format_identifiers(sql_template, params)) return self.execute(sql_template, params=non_identifier_params) def drop_temp_share_products_table(temp_table_name): """Drop temp share products table. Args: temp_table_name (str): Temp table name. """ sql_template = sql_loader.load_query( 'drop_temp_share_products_table') params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], temp_table_name=temp_table_name) sql_template, non_identifier_params = ( self.validator.format_identifiers(sql_template, params)) return self.execute(sql_template, params=non_identifier_params) def populate_temp_shared_products(temp_table_name, upc_list): """Populate temp shared products table. Args: temp_table_name (str): Temp table name. upc_list (list): List of UPCs to populate. """ sql_template = sql_loader.load_query( 'populate_temp_shared_products') sql_template = sql_template.format( values_list=','.join(['(%s)'] * len(upc_list))) params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], temp_table_name=temp_table_name) sql_template, non_identifier_params = ( self.validator.format_identifiers(sql_template, params)) self.execute(sql_template, upc_list) temp_table_name = self.get_temp_table_name() create_temp_share_products_table(temp_table_name) populate_temp_shared_products(temp_table_name, upc_list) sql_template = sql_loader.load_query('add_to_shared_products') sql_template = sql_template.format( values_list=','.join(['(%s)'] * len(upc_list))) params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], temp_table_name=temp_table_name) sql_template, non_identifier_params = ( self.validator.format_identifiers(sql_template, params)) self.execute(sql_template, non_identifier_params) drop_temp_share_products_table(temp_table_name) def delete_from_shared_products(self, upc_list): """Delete upcs from sme share. Args: upc_list (list): List of upcs to delete from share. """ sql_template = sql_loader.load_query('delete_from_shared_products') params = dict( db=self.sf_config['db'], schema=self.sf_config['schema']) sql_template, non_identifier_params = ( self.validator.format_identifiers(sql_template, params)) self.execute(sql_template, {'upc_list': upc_list})