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
| Syntax | Meaning | Example |
|---|---|---|
task(input) | Implicit dependency, TaskFlow API | transform(extract()) |
a >> b | a runs before b | extract >> transform |
b << a | b runs after a | transform << extract |
a >> [b, c] | a runs before b and c (parallel) | extract >> [validate, load] |
[a, b] >> c | c waits for both a and b | [check1, check2] >> process |
a >> b >> c | Sequential chain | extract >> 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
)