Skip to content

How to set up ETL with Apache Airflow

DAG

Parametrize DAG

@dag(
    params={
        'environment': Param('dev', type='string', enum=['dev', 'stage', 'prod']),
    },
)

these params are injected by dag into task automatically

eg.

@task
def transform(data, params):
    pass

How to dynamically create dags for all environments

from functools import wraps
from typing import Any, Callable, Dict, Optional
from airflow.decorators import dag


def multi_env_dag(
    dag_id: str,
    **envs: Optional[Dict[str, Any]]
) -> Callable:
    """
    Decorator that creates multiple environment-specific DAGs from a single DAG definition.
    Automatically adds environment-specific tags and env param to each DAG.

    Args:
        dag_id: Base DAG ID (suffixes will be appended)
        **envs: Environment configurations as keyword arguments (e.g., prod={...}, stage={...}, integration={...})
                Each environment accepts same parameters as @dag decorator except dag_id

    Example:
        @multi_env_dag(
            dag_id="my_pipeline",
            prod={"schedule": "0 0 * * *", "tags": ["critical"]},
            stage={"schedule": "0 6 * * *"},
            integration={"schedule": None}
        )
        def my_dag_function():
            # Your DAG tasks here
            # Access env with: context['params']['env']
            pass
    """

    def decorator(func: Callable) -> Callable:
        # Helper function to add environment tag and param
        def add_env_metadata(config: Dict[str, Any], env_name: str) -> Dict[str, Any]:
            config_copy = config.copy()

            # Add environment tag
            existing_tags = config_copy.get('tags', [])

            # Ensure tags is a list
            if not isinstance(existing_tags, list):
                existing_tags = [existing_tags]

            # Add environment tag if not already present
            if env_name not in existing_tags:
                config_copy['tags'] = existing_tags + [env_name]
            else:
                config_copy['tags'] = existing_tags

            # Add environment to params
            existing_params = config_copy.get('params', {})

            # Ensure params is a dict
            if not isinstance(existing_params, dict):
                existing_params = {}

            # Add env param (don't override if already exists)
            if 'env' not in existing_params:
                existing_params['env'] = env_name

            config_copy['params'] = existing_params

            return config_copy

        # Store decorated DAG functions
        decorated_dags = {}

        # Create a DAG for each environment
        for env_name, env_config in envs.items():
            if env_config is None:
                env_config = {}

            # Add environment-specific tag and param
            env_config_with_metadata = add_env_metadata(env_config, env_name)

            # Create DAG with environment suffix
            env_dag_id = f"{dag_id}_{env_name}"
            decorated_dags[env_name] = dag(dag_id=env_dag_id, **env_config_with_metadata)(func)

        # Return a wrapper that creates all DAG instances
        @wraps(func)
        def wrapper(*args, **kwargs):
            # Instantiate all DAGs
            dag_instances = {}
            for env_name, decorated_func in decorated_dags.items():
                dag_instances[env_name] = decorated_func(*args, **kwargs)

            return dag_instances

        # Store references to the decorated functions for Airflow to discover
        for env_name, decorated_func in decorated_dags.items():
            setattr(wrapper, env_name, decorated_func)

        return wrapper

    return decorator

Snowflake to Postgres Example

@dag(schedule=None, start_date=datetime(2024, 1, 1))
def snowflake_to_postgres_etl():

    @task()
    def extract_from_snowflake():
        hook = SnowflakeHook(snowflake_conn_id='snowflake_default')

        sql = """
            SELECT 1
        """

        df = hook.get_pandas_df(sql)
        return df.to_dict('records')

    @task()
    def transform_data(data: list):
        """Transform the extracted data"""
        df = pd.DataFrame(data)

        # Do transformations here

        return df.to_dict('records')

    @task()
    def load_to_postgres(data: list):
        """Load transformed data to Postgres"""
        if not data:
            print("No data to load")
            return

        df = pd.DataFrame(data)
        hook = PostgresHook(postgres_conn_id='postgres_default')
        engine = hook.get_sqlalchemy_engine()

        # Load data to Postgres
        df.to_sql(
            name='target_table',
            con=engine,
            schema='public',
            if_exists='append',  # or 'replace' depending on your needs
            index=False,
            method='multi',
            chunksize=1000
        )

        print(f"Successfully loaded {len(df)} rows to Postgres")

    # Define task dependencies
    data = extract_from_snowflake()
    transformed_data = transform_data(data)
    load_to_postgres(transformed_data)


# Instantiate the DAG
dag = snowflake_to_postgres_etl()

Tasks

Cleanup task

@task(trigger_rule='always')
def cleanup():
    pass

Standard Python task

@task
def my_python_task():
    pass

Docker task

@task.docker(image="python:3.9")
def my_docker_task():
    pass

Virtualenv task

@task.virtualenv(requirements=["pandas==1.5.0"])
def my_venv_task():
    pass

Kubernetes task

@task.kubernetes(
    image="python:3.9-slim",
    namespace="default",
    name="my-k8s-pod"
)
def my_k8s_task():
    pass

Branching task

@task.branch
def branch_task():
    if condition:
        return "task_a"
    return "task_b"

Short circuiting task

@task.short_circuit
def short_circuit_task():
    if condition:
        return False
    return True

Task Dependencies

SyntaxMeaningExample
task(input)Implicit dependency, TaskFlow APItransform(extract())
a >> ba runs before bextract >> transform
b << ab runs after atransform << extract
a >> [b, c]a runs before b and c (parallel)extract >> [validate, load]
[a, b] >> cc waits for both a and b[check1, check2] >> process
a >> b >> cSequential chainextract >> transform >> load

Examples

Prepare staging data

@task()
def create_staging_table():
    """Create a temporary staging table (dropped at end of session)"""

    hook = PostgresHook(postgres_conn_id='postgres_default')

    # TEMPORARY table - automatically dropped when session ends
    create_temp_table_sql = """
    CREATE TEMPORARY TABLE IF NOT EXISTS staging_table (
        id INTEGER,
        name VARCHAR(255),
        amount NUMERIC(10, 2),
        date TIMESTAMP,
        processed BOOLEAN DEFAULT FALSE
    );
    """

    hook.run(create_temp_table_sql)

    print("Temporary staging table created")

Load data from S3

@task
    def extract_from_s3(bucket_name, file_key) -> str:
        s3_hook = S3Hook(aws_conn_id='aws_default')

        file_content = s3_hook.read_key(
            key=file_key,
            bucket_name=bucket_name
        )

        df = pd.read_csv(io.StringIO(file_content))
        # or
        # df = pd.read_excel(io.BytesIO(file_content))

        return df.to_json()

Load data into postgres (Small Load)

@task()
def load_to_postgres(data: list):
    """Load transformed data to Postgres"""
    if not data:
        print("No data to load")
        return

    df = pd.DataFrame(data)
    hook = PostgresHook(postgres_conn_id='postgres_default')
    engine = hook.get_sqlalchemy_engine()

    # Load data to Postgres
    df.to_sql(
        name='target_table',
        con=engine,
        schema='public',
        if_exists='append',  # or 'replace' depending on your needs
        index=False,
        method='multi',
        chunksize=1000
    )

    print(f"Successfully loaded {len(df)} rows to Postgres")

Load data into postgres (Big Load)

@task
def load_data_in_chunks(file_path, chunk_size=100000):
    hook = PostgresHook(postgres_conn_id='postgres_default')

    # Disable indexes during load
    hook.run("DROP INDEX IF EXISTS idx_name;")

    # Increase maintenance_work_mem
    hook.run("SET maintenance_work_mem = '2GB';")

    # Disable autovacuum
    hook.run("ALTER TABLE target_table SET (autovacuum_enabled = false);")

    # Get connection manually and start one transaction
    conn = hook.get_conn()
    cursor = conn.cursor()
    try:
        for chunk in pd.read_csv(file_path, chunksize=chunk_size):
            temp_file = f'/tmp/chunk_{uuid.uuid4()}.csv'
            chunk.to_csv(temp_file, index=False)

            with open(temp_file, 'r') as f:
                cursor.copy_expert("COPY target_table FROM STDIN WITH CSV HEADER", f)
            os.remove(temp_file)

        conn.commit()
    except Exception:
        conn.rollback()
        raise
    finally:
        cursor.close()
        conn.close()

    # Recreate indexes
    hook.run("CREATE INDEX idx_name ON target_table(column);")

    # Analyze table
    hook.run("ANALYZE target_table;")

    # Re-enable autovacuum
    hook.run("ALTER TABLE target_table SET (autovacuum_enabled = true);")

Fan-out pattern (single param)

from airflow.decorators import dag, task
from datetime import datetime

@dag(schedule=None, start_date=datetime(2024, 1, 1))
def parallel_load_dag():

    @task
    def split_files() -> list:
        return ['file1.csv', 'file2.csv', 'file3.csv', ...]

    @task
    def load_file(file_path: str):
        pass

    @task
    def validate_load(loaded_files: list):
        pass

    # define number of parallel tasks
    files = split_files()
    # will call load_file for each file in parallel
    loaded = load_file.expand(file_path=files)
    # will wait until all files are loaded
    validate_load(loaded)

dag = parallel_load_dag()

Fan-out pattern (multiple param)

from airflow.decorators import dag, task
from datetime import datetime

@dag(schedule=None, start_date=datetime(2024, 1, 1))
def process_data_dag():
    @task
    def split_files():
        return [
            {"path": "data/file1.csv", "format": "csv"},
            {"path": "data/file2.parquet", "format": "parquet"},
            {"path": "data/file3.json", "format": "json"},
        ]

    @task
    def load_file(path: str, format: str):
        pass

    @task
    def validate_load(results: list):
        pass


    # define number of parallel tasks
    files = split_files()
    # will call load_file for each file in parallel
    loaded = load_file.expand_kwargs(files)
    # will wait until all files are loaded
    validate_load(loaded)

dag = process_data_dag()

Operators

Sensors

ExternalTaskSensor

This sensor allows one DAG to wait for a task in another DAG to complete before proceeding. It's crucial for creating dependencies between separate DAGs.

Two modes are available: 'poke' (default) and 'reschedule'. 'poke' keeps a worker waiting and reschedules release worker in between checks. Use 'reschedule' for long waits (hours/days).

It should be placed inside @dag class.

Examples:

# When two dags run at same time
from airflow.sensors.external_task import ExternalTaskSensor

wait_for_extraction = ExternalTaskSensor(
    task_id='wait_for_data_ready',           # Name of THIS sensor task
    external_dag_id='data_extraction_dag',   # The OTHER DAG to monitor
    external_task_id='extract_complete',     # The specific task to wait for
    mode='poke',
    timeout=3600                             # Max wait time in seconds (1 hour)
    poke_interval=120,                       # Check every 2 minutes
)
# DAGs runs at different times
from airflow.sensors.external_task import ExternalTaskSensor

wait_for_extraction = ExternalTaskSensor(
    task_id='wait_for_data_ready',           # Name of THIS sensor task
    external_dag_id='data_extraction_dag',   # The OTHER DAG to monitor
    external_task_id='extract_complete',     # The specific task to wait for
    mode='reschedule',
    poke_interval=60,                        # Check every minute
    timeout=3600                             # Max wait time in seconds (1 hour)
    execution_delta=timedelta(hours=-1),     # How far back to look for the task
)

S3KeySensor

SqlSensor