import re from pathlib import Path from typing import Optional, Any, List from airflow.contrib.hooks.aws_hook import AwsHook from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook from airflow.providers.snowflake.operators.snowflake import SnowflakeOperator from jinjasql import JinjaSql COMMON_QUERIES_DIR = Path(__file__).parent / 'queries' jinja_sql = JinjaSql( param_style='pyformat' ) class SQLTemplateOperator(SnowflakeOperator): template_fields = ('parameters',) def __init__( self, *, template: str, parameters: dict, template_dir: Optional[Path] = None, aws_conn_id: Optional[str] = None, sql: str = None, **kwargs, ) -> None: if sql: raise ValueError('sql argument not supported. Use template instead') super().__init__(sql='NOT_USED_PLACEHOLDER', **kwargs) self.template = template self.parameters = parameters self.template_dir = template_dir self.aws_conn_id = aws_conn_id def execute(self, context: Any) -> None: """Run query on snowflake""" self.log.info(f'Executing SQL template: {self.template}') hook = self.get_db_hook() template_dirs = [self.template_dir] if self.template_dir else [] template_dirs.append(COMMON_QUERIES_DIR) db_conf = hook._get_conn_params() auto_parameters = dict( db=db_conf['database'], schema=db_conf['schema'], ) if self.aws_conn_id: aws_hook = AwsHook(aws_conn_id=self.aws_conn_id) creds = aws_hook.get_credentials() auto_parameters['aws_creds'] = creds query_template = load_query_template(filename=self.template, template_dirs=template_dirs) params = {**auto_parameters, **self.parameters} query, parameters = prepare_query(query_template=query_template, params=params) execution_info = hook.run(sql=query, autocommit=self.autocommit, parameters=parameters) self.query_ids = hook.query_ids if self.do_xcom_push: return execution_info def load_query_template(filename: str, template_dirs: Optional[List[Path]] = None): for template_dir in template_dirs: path = template_dir / filename if path.exists(): return path.read_text() raise FileNotFoundError(f'Cannot find file {filename} across dirs {template_dirs}') def run_template(template: str, parameters: dict, snowflake_conn_id: str): snowflake_hook = SnowflakeHook( snowflake_conn_id=snowflake_conn_id, ) sql = prepare_query(template, parameters) execution_info = snowflake_hook.run(sql=sql, parameters=parameters) return execution_info def prepare_query(query_template: str, params: Optional[dict]): params = params or {} query_template = re.sub(r'%\(([\w_]+)\)i', r'{{ \1|sqlsafe }}', query_template) query_template = re.sub(r'%\(([\w_]+)\)s', r'{{ \1 }}', query_template) query, query_params = jinja_sql.prepare_query(source=query_template, data=params) return query, query_params # return query, {**compatible_params, **query_params}