"""Snowflake executor class for the Chartmetric Tracks tasks.""" from snowflake_connector.etl_connector import SnowflakeSQLExecutor from snowflake_connector.etl_connector import SQLLoader from feed_ingestion.common.chartmetric_share_retry import ( ChartmetricShareRetryMixin) sql_loader = SQLLoader(__file__) class SnowflakeExecutor(ChartmetricShareRetryMixin, SnowflakeSQLExecutor): """Helper class to abstract Snowflake operations. This class inherits from SnowflakeSQLExecutor class, which provides basic set of methods. This class extends SnowflakeSQLExecutor with some specific methods, which are useful to encapsulate some flow specific operations. """ def grant_select_to_facts_db_prod_schema_read(self, table_name): """Grant SELECT privilege to the FACTS_DB_PROD_SCHEMA_READ role. Args: table_name (str): Name of a table. """ params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], table_name=table_name) sql_template = """ GRANT SELECT ON %(db)i.%(schema)i.%(table_name)i TO FACTS_DB_PROD_SCHEMA_READ; """ sql_template, _ = self.validator.format_identifiers( sql_template, params) self.execute(sql_template) def execute_query(self, query_name, **kwargs): """Execute a query. Args: query_name (str): The name of the query. kwargs (dict): Additional parameters to pass to the query. """ sql_template = sql_loader.load_query(query_name) params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], **kwargs) sql_template, non_identifier_params = self.validator.\ format_identifiers(sql_template, params) self.execute(sql_template, params=non_identifier_params) def fetchone_query(self, query_name, **kwargs): """Load query, resolve params and fetchone. Args: query_name (str): name of a query to load. kwargs (dict): Additional parameters to pass to the query. Returns: tuple: A first row produced by executing an SQL statement. """ params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], **kwargs) sql_template = sql_loader.load_query(query_name) sql, non_identifier_params = self.validator.format_identifiers( sql_template, params) return self.fetchone(sql, params=non_identifier_params) def fetchall_dict_query(self, query_name, **kwargs): """Load query, resolve params and fetch all as a dict. Args: query_name (str): name of a query to load. kwargs (dict): Additional parameters to pass to the query. Returns: tuple: A first row produced by executing an SQL statement. """ params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], **kwargs) sql_template = sql_loader.load_query(query_name) sql, non_identifier_params = self.validator.format_identifiers( sql_template, params) return self.fetchall( sql, params=non_identifier_params, dict_cursor=True)