from __future__ import annotations import os from datetime import datetime import logging from airflow import DAG from airflow.providers.snowflake.operators.snowflake import SnowflakeOperator from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook from airflow.operators.python import PythonOperator from airflow.utils.trigger_rule import TriggerRule from airflow.operators.empty import EmptyOperator from airflow.models import Variable from airflow.exceptions import AirflowFailException logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) SNOWFLAKE_CONN_ID = "doczy_dev_snowflake" DAG_ID = "load_request_and_contract_submissions" DATABASE="DOCZY_DEV" # bucket = "airflow-data-ingestion" TAGS=["dev","contract_config_interface","dataload"] # Trigger rules ALL_SUCCESS = 'all_success' ALL_FAILED = 'all_failed' ALL_DONE = 'all_done' ONE_SUCCESS = 'one_success' ONE_FAILED = 'one_failed' # Passing empty params for now, this will be overridden by the payload from the trigger # These params can also be set from the Airflow UI while manually triggering the DAG default_params = {"request_submission_file_name": "", "contract_config_file_name":""} # This will be replaced with the payload from the event after API connection is setup # training_results_file_name = "training_results_sample.csv" # attempt_logs_file_name = "attempt_logs_sample.csv" def call_stored_proc(proc_name,file_type, params): dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) with dwh_hook.get_conn() as conn: # dwh_hook.set_autocommit(conn,autocommit=False) cur = conn.cursor() # Added new parameter file_type to determine the file name to be passed to the stored procedure # The bucket name is set by default to "doczy-dev-infra-raw-data-ingestion" and the files should be ALWAYS save under training_interface/ path for now if file_type == 'request_submission': file_name = params['request_submission_file_name'] elif file_type == 'contract_config': file_name = params['contract_config_file_name'] cur.execute(f"CALL {DATABASE}.STG.{proc_name}('{file_name}');") result = cur.fetchone() if result[0] == 'Setup, Load, and Audit Complete': logger.info('PROCEDURE EXECUTED SUCCESSFULLY') else: raise AirflowFailException("Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", result) logger.info(f"QUERY EXECUTION RESULT: {str(result)}") dag = DAG( DAG_ID, start_date=datetime(2024, 1, 1), default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries":0}, tags=TAGS, catchup=False, schedule=None, params = default_params ) begin_job = EmptyOperator(task_id='Begin') load_training_results = PythonOperator( task_id="load_request_submissions", python_callable=call_stored_proc, dag=dag, op_kwargs={'proc_name':'LOAD_REQUEST_SUBMISSION', 'file_type':'request_submission'} ) load_attempt_logs_sp = PythonOperator( task_id="load_contract_config", python_callable=call_stored_proc, dag=dag, op_kwargs={'proc_name':'LOAD_CONTRACT_CONFIG', 'file_type':'contract_config'} ) end_job = EmptyOperator(task_id='End') begin_job >> load_training_results >> load_attempt_logs_sp >> end_job