"""Snowflake connector class for YouTube Asset Conflict tasks.""" from snowflake_connector.etl_connector import SnowflakeSQLExecutor from snowflake_connector.etl_connector import SQLLoader from yt_conflict_elasticsearch.flows.elasticsearch_export import config # Load SQL templates sql_loader = SQLLoader(__file__) class YouTubeConflictSFExecutor(SnowflakeSQLExecutor): """Class representing the executor.""" def store_csv_to_s3_conflicts(self, s3_url, conflict_status): """Store csv data to s3 for new conflicts (via Snowflake).""" sql_template = sql_loader.load_query( 'store_csv_to_s3_{}_conflicts'.format(conflict_status)) params = dict( s3_url=s3_url, aws_key=config.AWS_ACCESS_KEY_ID, aws_secret_key=config.AWS_SECRET_ACCESS_KEY, db=self.sf_config['db'], schema=self.sf_config['schema'], ar_db=self.sf_config['ar_db'], ar_schema=self.sf_config['ar_schema'], ows_conflict_db=self.sf_config['ows_conflict_db'], ows_conflict_schema=self.sf_config['ows_conflict_schema'], ) sql, non_identifier_params = self.validator.format_identifiers( sql_template, params) self.execute(sql, params=non_identifier_params) def mark_indexed_conflicts(self): """Set es_indexed=True for indexed conflicts.""" sql_template = sql_loader.load_query('mark_indexed_conflicts') sql, non_identifier_params = self.validator.format_identifiers( sql_template, {'schema': self.sf_config['schema']}) self.execute(sql, params=non_identifier_params) def get_responded_conflicts(self): """Get responded conflicts from Snowflake.""" params = dict(schema=self.sf_config['schema']) return self.fetchall_query( sql_loader, 'get_responded_conflicts', params) def get_resolved_conflicts(self): """Get resolved conflicts from Snowflake.""" params = dict(schema=self.sf_config['schema']) return self.fetchall_query( sql_loader, 'get_resolved_conflicts', params) def create_temp_conflict_to_es_id(self): """Create a temp table that ties conflict_ids to es_id.""" params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], conflict_to_es_id_temp_table=self.sf_config[ 'temp_conflict_es_id_table']) sql_template = sql_loader.load_query('create_temp_conflict_to_es_id') sql, non_identifier_params = self.validator.format_identifiers( sql_template, params) self.execute(sql, params=non_identifier_params) def fill_temp_table_with_es_ids( self, conflict_to_es_id_temp_table, records): """Populate temp table with conflict_id to es_id mappings.""" sql_template = sql_loader.load_query( 'fill_temp_table_with_es_ids') params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], conflict_to_es_id_temp_table=conflict_to_es_id_temp_table) sql, _ = self.validator.format_identifiers(sql_template, params) self.executemany(sql, params_list=records) def populate_es_ids_to_conflicts(self): """Populate es_ids for newly indexed conflicts.""" params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], conflict_to_es_id_temp_table=self.sf_config[ 'temp_conflict_es_id_table']) sql_template = sql_loader.load_query('populate_es_ids_to_conflicts') sql, non_identifier_params = self.validator.format_identifiers( sql_template, params) self.execute(sql, params=non_identifier_params) def remove_selected_es_ids(self, es_ids): """Remove selected es_ids.""" es_ids = [es_id for es_id, in es_ids] params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], es_ids=es_ids) sql_template = sql_loader.load_query('remove_selected_es_ids') sql, non_identifier_params = self.validator.format_identifiers( sql_template, params) self.execute(sql, params=non_identifier_params)