import datetime import json import logging from pathlib import Path from tempfile import TemporaryDirectory from airflow.exceptions import AirflowFailException from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook from airflow.providers.amazon.aws.hooks.s3 import S3Hook from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook from flows.awal_salesforce import api, config from flows.awal_salesforce.config import AWS_CONN_ID THIS_DIR = Path(__file__).parent logger = logging.getLogger(__name__) def bootstrap(dag_run, ds, **kwargs): context = dag_run.conf logger.info(f'Context {context}') context_date = dag_run.logical_date today_started = datetime.datetime.now(tz=datetime.timezone.utc).replace(minute=0, hour=0, second=0, microsecond=0) logger.info(f'Today {today_started}, context: {context_date}') if context_date < today_started: raise AirflowFailException(f'Cannot process historical date {context_date}. This DAG doesn''t support backfill') s3_bucket = config.ARCHIVE_S3_BUCKET s3_key = config.ARCHIVE_S3_KEY_PREFIX_TEMPLATE.format(date=ds) s3_path = f's3://{s3_bucket}/{s3_key}' return dict( archive_s3_bucket=s3_bucket, archive_s3_key=s3_key, archive_s3_path=s3_path, ) def fetch(salesforce_entity, s3_bucket, s3_key, **kwargs): sf = api.connect() object_type = salesforce_entity.lower() desc = api.describe_salesforce_object(sf, salesforce_entity) with TemporaryDirectory() as temp_dir: create_ddl_filename = f'{object_type}_describe.json' temp_path = Path(temp_dir) with open(temp_path / create_ddl_filename, 'w') as file: json.dump(desc, fp=file, indent=4) csv_filename = f'{object_type}_data.csv' with open(temp_path / csv_filename, 'w') as file: api.save_to_csv_file(sf, desc, file) s3_hook = S3Hook(aws_conn_id=AWS_CONN_ID) for file in create_ddl_filename, csv_filename: s3_kwargs = dict( bucket_name=s3_bucket, key=f'{s3_key}{file}', ) logger.info(f'Uploading file {file} to S3 {s3_kwargs}') s3_hook.load_file( filename=temp_path / file, replace=True, **s3_kwargs ) logger.info(f'Upload done') # create_ddl_filename, csv_filename = 'opportunity_create.sql', 'opportunity_data.csv' return {'structure_file': create_ddl_filename, 'data_file': csv_filename} def create_table(s3_bucket, s3_key, structure_file, table_name, **kwargs): # download file from s3 with TemporaryDirectory() as tmp_dir: s3_hook = S3Hook(aws_conn_id=AWS_CONN_ID) s3_kwargs = dict( bucket_name=s3_bucket, key=f'{s3_key}{structure_file}', ) logger.info(f'Download file from S3 {s3_kwargs} to {tmp_dir}') downloaded_file = s3_hook.download_file( local_path=tmp_dir, **s3_kwargs ) with open(downloaded_file) as file: desc = json.load(file) # generate DDL ddl_stmt = api.gen_create_table_dds(desc, table_name) # execute DDL snowflake_hook = SnowflakeHook( snowflake_conn_id='snowflake_default', ) parameters = {} execution_info = snowflake_hook.run(sql=ddl_stmt, parameters=parameters) return execution_info def copy_data(data_file, table_name, s3_path, **kwargs): snowflake_hook = SnowflakeHook(snowflake_conn_id='snowflake_default') query_file = 'queries/copy_into_table.sql' sql = Path(THIS_DIR / query_file).read_text() aws_hook = AwsBaseHook(aws_conn_id=AWS_CONN_ID) credentials = aws_hook.get_session().get_credentials() sql = sql.replace('%(table_name)i', table_name) parameters = dict( table_name=table_name, aws_key_id=credentials.access_key, aws_secret_key=credentials.secret_key, s3_path=s3_path, files=data_file ) execution_info = snowflake_hook.run(sql=sql, parameters=parameters) return execution_info