from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook from jinja2 import Template from lib.config import SF_SCHEMA QUERY_TEMPLATE = Template(""" SELECT COUNT(*) AS "count" FROM ROYALTY_ACCOUNTING.{{schema}}.STMT_DB_SALES_DISTRO_STAGING WHERE BATCH_ID = '{{batch_id}}' """) def check_ingested_sales(**kwargs): batch_id = kwargs.get('params').get('batch_id') query = QUERY_TEMPLATE.render(schema=SF_SCHEMA, batch_id=batch_id) sf_hook = SnowflakeHook() db_engine = sf_hook.get_sqlalchemy_engine() with db_engine.begin() as conn: row = conn.execute(query).fetchone() row_count = dict(row).get('count') if row_count > 0: raise Exception(f'Batch {batch_id} has already been ingested')