"""Snowflake executor class for the Chartmetric Charts 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 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_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.fetchall(sql, params=non_identifier_params) def fetchall_dict_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.fetchall( sql, params=non_identifier_params, dict_cursor=True) def fetchmany_dict_query(self, query_name, size=None, **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. size (int): batch size for result set. 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.fetchmany( sql, size=size, params=non_identifier_params, dict_cursor=True)