"""Snowflake ETL connector base class, and SQL loader class.""" from contextlib import contextmanager from datetime import datetime import os from snowflake import connector from . import validator class SnowflakeSQLExecutor: """Helper base class to abstract Snowflake operations within ETLs. Operates with the connection to Snowflake, provides the methods for retrieving information about Snowflake tables, etc. Could be used as is, if provided methods are enough, and also could be extended with more specific method to incapsulate db logic. """ def __init__( self, sf_config, autocommit=True, statement_timeout_in_seconds=None): """Open a connection to database, set up the db operations. Args: sf_config (dict): Dict with all the required credentials. autocommit (bool): If True, use autocommit mode. statement_timeout_in_seconds (int): A global timeout for all the statements issued by this connection. A timed out query throws an exception, and closes the connection. """ self.sf_config = sf_config self.autocommit = autocommit self.statement_timeout_in_seconds = statement_timeout_in_seconds self.validator = validator.BaseValidator() self.snowflake_conn = self.get_connection() def __enter__(self): """Implement the context manager protocol. Returns: (SnowflakeSQLExecutor obj): An instance of SnowflakeSQLExecutor. """ if self.statement_timeout_in_seconds: self.execute( 'ALTER SESSION SET STATEMENT_TIMEOUT_IN_SECONDS={};'.format( self.statement_timeout_in_seconds)) return self def __exit__(self, ext_type, exc_value, traceback): """Implement the context manager protocol. Calls methods to free resources. """ self.snowflake_conn.close() def get_connection(self): """Create connection to the Snowflake. If autocommit == False, then automatically rollback all the statements of the current connection in case if db exception occurs within any of them. Returns: Connection: Connection that supports DB API v2 interface. """ return connector.connect( user=self.sf_config.get('user', None), password=self.sf_config.get('password', None), account=self.sf_config['account'], warehouse=self.sf_config['warehouse'], database=self.sf_config['db'], schema=self.sf_config['schema'], role=self.sf_config['role'], private_key=self.sf_config.get('private_key', None), ocsp_fail_open=self.sf_config.get('ocsp_fail_open', True), autocommit=self.autocommit) @contextmanager def get_cursor(self, dict_cursor=None): """Convenience context manager to provide a cursor. Usage example: with self.get_cursor() as cursor: cursor.execute('DROP TABLE super_important_stuff') Args: dict_cursor (bool): If True, use DictCursor. Yields: Cursor: initialized cursor object as per DB API v2. """ if dict_cursor: cursor = self.snowflake_conn.cursor(connector.DictCursor) else: cursor = self.snowflake_conn.cursor() if self.autocommit: try: yield cursor cursor.close() except Exception: cursor.close() self.snowflake_conn.close() raise else: try: cursor.execute('BEGIN') yield cursor self.snowflake_conn.commit() cursor.close() except Exception: self.snowflake_conn.rollback() cursor.close() self.snowflake_conn.close() raise def execute(self, sql_template, params=None): """Execute an SQL statement. Args: sql_template (str): A template to be parametrized. All the identifiers (db name, schema name, table name, column names) should be already present. params (dict): A dict of params to bind to sql_template. """ with self.get_cursor() as cursor: return cursor.execute(sql_template, params) def executemany(self, sql_template, params_list): """Execute an SQL statement many times with different params to bind. Args: sql_template (str): A template to be parametrized. All the identifiers (db name, schema name, table name, column names) should be already present. params_list (list): A list of params dicts to bind to sql_template. """ with self.get_cursor() as cursor: cursor.executemany(sql_template, params_list) def fetchone(self, sql_template, params=None, dict_cursor=False): """Execute SQL and fetch first tuple. Useful for SQL statements which are supposed to return just one row, like COUNT, etc.. Args: sql_template (str): A template to be parametrized. All the identifiers (db name, schema name, table name, column names) should be already present. params (dict): A dict of params to bind to sql_template. dict_cursor (bool): If True, use DictCursor. Returns: tuple: A first row produced by executing an SQL statement. """ with self.get_cursor(dict_cursor) as cursor: cursor.execute(sql_template, params) return cursor.fetchone() def fetchall(self, sql_template, params=None, dict_cursor=False): """Execute query and fetch all results. Args: sql_template (str): A template to be parametrized. All the identifiers (db name, schema name, table name, column names) should be already present. params (dict): A dict of params to bind to sql_template. dict_cursor (bool): If True, use DictCursor. Returns: list(tuple): List of rows with query result. """ with self.get_cursor(dict_cursor) as cursor: cursor.execute(sql_template, params) return cursor.fetchall() def fetchmany(self, sql_template, size, params=None, dict_cursor=False): """Generator fetches set of n rows of a query result (n=size). Stops when no more results available. Args: sql_template (str): A template to be parametrized. All the identifiers (db name, schema name, table name, column names) should be already present. size (int): Number of rows to fetch within one batch. params (dict): A dict of params to bind to sql_template. dict_cursor (bool): If True, use DictCursor. Yields: list(tuple): List of maximum n rows with query result (n = size). """ with self.get_cursor(dict_cursor) as cursor: cursor.execute(sql_template, params) while True: batch = cursor.fetchmany(size) if not batch: break yield batch def table_exists(self, table): """Check if table exists in the Snowflake. Schema name will be taken from sf_config. Args: table (str): Table name in Snowflake. Returns: bool: True if table exists, False othewise. """ sql_template = ( 'SELECT COUNT(*) FROM INFORMATION_SCHEMA.TABLES ' 'WHERE TABLE_NAME = %(table)s ' 'AND TABLE_SCHEMA = %(schema)s ' 'AND TABLE_CATALOG = %(db)s;') params = { 'table': table.upper(), 'schema': self.sf_config['schema'].upper(), 'db': self.sf_config['db'].upper()} with self.get_cursor() as cursor: cursor.execute(sql_template, params) return bool(cursor.fetchone()[0]) def drop_table(self, table, db=None, schema=None): """Drop table if it exists in the Snowflake. Default db and schema names will be taken from sf_config. Args: table (str): Table name in Snowflake. db (str): Optional db name. schema (str): Optional schema name. """ sql_template = 'DROP TABLE IF EXISTS %(db)i.%(schema)i.%(table)i;' params = dict( db=db or self.sf_config['db'], schema=schema or self.sf_config['schema'], table=table) sql, _ = self.validator.format_identifiers(sql_template, params) self.execute(sql) def create_table(self, table, transient=False, db=None, schema=None): """Create table in the Snowflake. This is a simple helper, if you want more parameters, please use custom SQL statement and execute() method. Default db and schema names will be taken from sf_config. Args: table (str): Table name in Snowflake. transient (bool): True if you want to create a transient table. db (str): Optional db name. schema (str): Optional schema name. """ sql_template = ( 'CREATE %(transient)i TABLE IF NOT EXISTS ' '%(db)i.%(schema)i.%(table)i;') params = dict( transient='' if not transient else 'TRANSIENT', db=self.sf_config['db'] or db, schema=self.sf_config['schema'] or schema, table=table) sql, _ = self.validator.format_identifiers(sql_template, params) self.execute(sql) def create_table_like( self, table, source_table, db=None, schema=None, source_db=None, source_schema=None, transient=False): """Create table LIKE other table. Args: table (str): Destination table name in Snowflake. source_table (str): Source table name in Snowflake. db (str): Optional destination db name. schema (str): Optional destination schema name. source_db (str): Optional source db name. source_schema (str): Optional source schema name. transient (bool): True if you want to create a transient table. """ sql_template = ( 'CREATE %(transient)i TABLE IF NOT EXISTS ' '%(db)i.%(schema)i.%(table)i LIKE ' '%(source_db)i.%(source_schema)i.%(source_table)i;') params = dict( transient='' if not transient else 'TRANSIENT', table=table, source_table=source_table, db=db or self.sf_config['db'], schema=schema or self.sf_config['schema'], source_db=source_db or self.sf_config['db'], source_schema=source_schema or self.sf_config['schema']) sql, _ = self.validator.format_identifiers(sql_template, params) self.execute(sql) def truncate_table(self, table, db=None, schema=None): """Create table in the Snowflake. Default db and schema names will be taken from sf_config. Args: table (str): Table name in Snowflake. db (str): Optional db name. schema (str): Optional schema name. """ sql_template = 'TRUNCATE %(db)i.%(schema)i.%(table)i;' params = dict( db=db or self.sf_config['db'], schema=schema or self.sf_config['schema'], table=table) sql, _ = self.validator.format_identifiers(sql_template, params) self.execute(sql) def get_column_names(self, table, db=None, schema=None): """Get a list of column names of a Snowflake table. Args: table (str): A table name to extract column names for. db (str): Optional db name. schema (str): Optional schema name. Returns: list: A list of column names. """ sql_template = 'DESC TABLE %(db)i.%(schema)i.%(table)i;' params = dict( db=db or self.sf_config['db'], schema=schema or self.sf_config['schema'], table=table) sql, _ = self.validator.format_identifiers(sql_template, params) result = self.fetchall(sql) column_names = [] for table_desc in result: column_names.append(table_desc[0]) return column_names def swap_tables(self, table1, table2): """Swap two tables. Args: table1 (str): A first table name. table2 (str): A second table name. """ sql_template = ( 'ALTER TABLE %(db)i.%(schema)i.%(table1)i SWAP WITH ' '%(db)i.%(schema)i.%(table2)i;') params = dict( db=self.sf_config['db'], schema=self.sf_config['schema'], table1=table1, table2=table2) sql, _ = self.validator.format_identifiers(sql_template, params) self.execute(sql) def execute_query(self, sql_loader, query_name, params): """Load query, resolve params and execute. Args: sql_loader (SQLLoader): loader instance. query_name (str): name of a query to load. params (dict): sql parameters. """ sql_template = sql_loader.load_query(query_name) sql, non_identifier_params = self.validator.format_identifiers( sql_template, params) self.execute(sql, params=non_identifier_params) def fetchone_query(self, sql_loader, query_name, params): """Load query, resolve params and fetchone. Args: sql_loader (SQLLoader): loader instance. query_name (str): name of a query to load. params (dict): sql parameters. Returns: tuple: A first row produced by executing an SQL statement. """ 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, sql_loader, query_name, params): """Load query, resolve params and fetchall. Args: sql_loader (SQLLoader): loader instance. query_name (str): name of a query to load. params (dict): sql parameters. Returns: list(tuple): List of rows with query result. """ 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) class SQLLoader(object): """Loads SQL queries/templates from text files. Do not forget to include these files to setup.py of the repo: setup( package_data={ # include all the .sql files 'snowflake_etl.flows.rs2sf': ['queries/*.sql']}) """ def __init__(self, path_to_file, folder='/queries', date=None): """Constructor for SQL loader. To add possibility loading different versions it is needed to specify date and put queries in sub folders with version release date as name. Usage examples: import os sql = SQLLoader(__file__) sql = SQLLoader(__file__, date='2019-01-01') Args: path_to_file (str): absolute path to the file which imports SQLLoader. date (str): YYYY-MM-DD date is required when some queries are placed in sub folder with date name. """ self.sql_files_root = os.path.realpath( os.path.dirname(path_to_file)) + folder self.query_cash = {} self.date = date self.folder_version = self._get_sub_folder(date) if date else None def _get_sub_folder(self, date): """Return sub folder name in format YYYY-MM-DD. Args: date (str): YYYY-MM-DD date. Returns: str or None: Sub folder name in format YYYY-MM-DD if it exists and None otherwise. """ current_date = datetime.strptime(date, '%Y-%m-%d') for item in sorted(os.listdir(self.sql_files_root), reverse=True): path = '{root}/{folder_version}'.format( root=self.sql_files_root, folder_version=item) if os.path.isdir(path): try: if current_date >= datetime.strptime(item, '%Y-%m-%d'): return item except ValueError: # skip folders which are not in YYYY-MM-DD format pass def _get_query_path(self, query_name): """Return a query path depending on sub_folder. Args: query_name (str): name of the query (filename without extension). Returns: str: The full query path. """ if self.folder_version: path = '{root}/{folder_version}/{query_name}.sql'.format( root=self.sql_files_root, folder_version=self.folder_version, query_name=query_name) if os.path.isfile(path): return path path = '{root}/{query_name}.sql'.format( root=self.sql_files_root, query_name=query_name) return path def _load_query(self, query_name): """Load query from the disk. Args: query_name (str): name of the query (filename without extension). """ path = self._get_query_path(query_name) with open(path, 'r') as q: query = q.read() return query def load_query(self, item): """Access queries by name. This method uses cache to reduce I/O operations. Usage example: sql_queries = SQLLoader('query_path') sql_queries.get_query('query_name') Args: item (str): name of the query (filename without extension). Returns: str: SQL query template. """ if item not in self.query_cash: self.query_cash[item] = self._load_query(item) return self.query_cash[item]