| name | spark-data-engineering-pipeline |
| description | Production PySpark ETL pipeline with AWS S3, PostgreSQL, schema validation, data quality checks, and automated testing |
| triggers | ["build a pyspark data pipeline","create an etl pipeline with spark and s3","set up spark data engineering workflow","validate and transform data using pyspark","load parquet data from s3 to postgres","implement data quality checks in spark","create production spark pipeline","set up spark with aws s3 and postgresql"] |
Spark Data Engineering Pipeline
Skill by ara.so — Data Skills collection.
Overview
A production-grade ETL pipeline built with PySpark that extracts JSON data from AWS S3, performs schema validation and data quality checks, transforms the data, writes it as Parquet, and loads it into PostgreSQL. This project demonstrates end-to-end data engineering best practices including logging, testing, and CI/CD.
Project Structure
spark-data-engineering-pipeline/
├── configs/ # Configuration files
├── data/ # Sample data
├── drivers/ # JDBC drivers
├── logs/ # Application logs
├── notebooks/ # Jupyter notebooks
├── sql/ # SQL scripts
├── src/
│ ├── extract/ # Data extraction modules
│ ├── validation/ # Schema and quality validation
│ ├── transform/ # Data transformation logic
│ ├── load/ # Data loading to PostgreSQL
│ └── utils/ # Utility functions
├── tests/ # Unit tests
├── main.py # Pipeline entry point
└── requirements.txt
Installation
git clone https://github.com/Giri-25/spark-data-engineering-pipeline.git
cd spark-data-engineering-pipeline
python3 -m venv .venv
source .venv/bin/activate
pip install -r requirements.txt
Configuration
Create a .env file in the project root:
AWS_ACCESS_KEY_ID=your_access_key
AWS_SECRET_ACCESS_KEY=your_secret_key
AWS_REGION=us-east-1
S3_BUCKET=your-bucket-name
POSTGRES_HOST=localhost
POSTGRES_PORT=5432
POSTGRES_DB=data_warehouse
POSTGRES_USER=your_user
POSTGRES_PASSWORD=your_password
Running the Pipeline
python main.py
python main.py --config configs/pipeline_config.yaml
pytest tests/
pytest --cov=src tests/
Core Components
1. Data Extraction from S3
from pyspark.sql import SparkSession
import os
def extract_from_s3(spark, s3_path):
"""
Extract JSON data from S3
Args:
spark: SparkSession instance
s3_path: S3 path to JSON files (s3a://bucket/path/)
Returns:
DataFrame with raw data
"""
spark._jsc.hadoopConfiguration().set(
"fs.s3a.access.key",
os.getenv("AWS_ACCESS_KEY_ID")
)
spark._jsc.hadoopConfiguration().set(
"fs.s3a.secret.key",
os.getenv("AWS_SECRET_ACCESS_KEY")
)
spark._jsc.hadoopConfiguration().set(
"fs.s3a.endpoint",
f"s3.{os.getenv('AWS_REGION')}.amazonaws.com"
)
df = spark.read.json(s3_path)
return df
spark = SparkSession.builder \
.appName("DataPipeline") \
.config("spark.jars.packages", "org.apache.hadoop:hadoop-aws:3.3.1") \
.getOrCreate()
s3_bucket = os.getenv("S3_BUCKET")
raw_df = extract_from_s3(spark, f"s3a://{s3_bucket}/raw/data/*.json")
2. Schema Validation
from pyspark.sql.types import StructType, StructField, StringType, IntegerType, TimestampType
from pyspark.sql import DataFrame
import logging
logger = logging.getLogger(__name__)
def define_expected_schema():
"""Define the expected schema for incoming data"""
return StructType([
StructField("user_id", StringType(), nullable=False),
StructField("transaction_id", StringType(), nullable=False),
StructField("amount", IntegerType(), nullable=False),
StructField("currency", StringType(), nullable=False),
StructField("timestamp", TimestampType(), nullable=False),
StructField("status", StringType(), nullable=True),
StructField("merchant_id", StringType(), nullable=True)
])
def validate_schema(df: DataFrame, expected_schema: StructType) -> bool:
"""
Validate DataFrame schema against expected schema
Args:
df: Input DataFrame
expected_schema: Expected StructType schema
Returns:
True if schema matches, False otherwise
"""
df_fields = set(df.schema.fieldNames())
expected_fields = set(expected_schema.fieldNames())
missing_fields = expected_fields - df_fields
if missing_fields:
logger.error(f"Missing required fields: {missing_fields}")
return False
extra_fields = df_fields - expected_fields
extra_fields:
logger.warning()
field expected_schema.fields:
field.name df.schema.fieldNames():
df_field = df.schema[field.name]
df_field.dataType != field.dataType:
logger.error(
)
logger.info()
expected_schema = define_expected_schema()
validate_schema(raw_df, expected_schema):
validated_df = raw_df
:
ValueError()
3. Data Quality Checks
from pyspark.sql import DataFrame
from pyspark.sql import functions as F
import logging
logger = logging.getLogger(__name__)
class DataQualityChecker:
"""Perform data quality checks on DataFrames"""
def __init__(self, df: DataFrame):
self.df = df
self.quality_report = {}
def check_null_values(self, columns: list) -> dict:
"""Check for null values in specified columns"""
null_counts = {}
for col in columns:
null_count = self.df.filter(F.col(col).isNull()).count()
null_counts[col] = null_count
if null_count > 0:
logger.warning(f"Column '{col}' has {null_count} null values")
self.quality_report['null_counts'] = null_counts
return null_counts
def check_duplicates(self, key_columns: list) -> int:
"""Check for duplicate records based on key columns"""
total_count = self.df.count()
distinct_count = self.df.dropDuplicates(key_columns).count()
duplicates = total_count - distinct_count
duplicates > :
logger.warning()
.quality_report[] = duplicates
duplicates
() -> :
out_of_range = .df.(
(F.col(column) < min_val) | (F.col(column) > max_val)
).count()
out_of_range > :
logger.warning(
)
.quality_report[] = out_of_range
out_of_range
() -> :
invalid_count = .df.(
~F.col(column).isin(allowed_values)
).count()
invalid_count > :
logger.warning(
)
.quality_report[] = invalid_count
invalid_count
() -> :
.quality_report
qc = DataQualityChecker(validated_df)
qc.check_null_values([, , ])
qc.check_duplicates([])
qc.check_value_ranges(, min_val=, max_val=)
qc.check_categorical_values(, [, , ])
quality_report = qc.get_report()
logger.info()
4. Data Transformation
from pyspark.sql import DataFrame
from pyspark.sql import functions as F
from pyspark.sql.window import Window
import logging
logger = logging.getLogger(__name__)
class DataTransformer:
"""Transform and clean data"""
@staticmethod
def clean_data(df: DataFrame) -> DataFrame:
"""Clean data by removing nulls and duplicates"""
df_cleaned = df.dropDuplicates(['transaction_id'])
df_cleaned = df_cleaned.fillna({
'status': 'unknown',
'merchant_id': 'unknown'
})
df_cleaned = df_cleaned.dropna(
subset=['user_id', 'transaction_id', 'amount']
)
logger.info(f"Cleaned data: {df_cleaned.count()} rows")
return df_cleaned
@staticmethod
def add_derived_columns(df: DataFrame) -> DataFrame:
"""Add derived columns for analytics"""
df_transformed = df.withColumn('date', F.to_date('timestamp'))
df_transformed = df_transformed.withColumn('year', F.year('timestamp'))
df_transformed = df_transformed.withColumn(, F.month())
df_transformed = df_transformed.withColumn(, F.dayofmonth())
df_transformed = df_transformed.withColumn(, F.hour())
df_transformed = df_transformed.withColumn(
,
F.dayofweek()
)
df_transformed = df_transformed.withColumn(
,
F.when(F.col() == , F.col() * )
.when(F.col() == , F.col() * )
.otherwise(F.col())
)
df_transformed
() -> DataFrame:
user_window = Window.partitionBy()
df_agg = df.withColumn(
,
F.count().over(user_window)
)
df_agg = df_agg.withColumn(
,
F.().over(user_window)
)
df_agg = df_agg.withColumn(
,
F.avg().over(user_window)
)
df_agg
transformer = DataTransformer()
cleaned_df = transformer.clean_data(validated_df)
enriched_df = transformer.add_derived_columns(cleaned_df)
final_df = transformer.add_aggregations(enriched_df)
logger.info()
5. Write to Parquet
from pyspark.sql import DataFrame
import os
import logging
logger = logging.getLogger(__name__)
def write_to_parquet(df: DataFrame, output_path: str, partition_cols: list = None):
"""
Write DataFrame to Parquet format
Args:
df: DataFrame to write
output_path: S3 or local path for Parquet files
partition_cols: Columns to partition by (optional)
"""
writer = df.write.mode('overwrite').format('parquet')
if partition_cols:
writer = writer.partitionBy(*partition_cols)
writer.option('compression', 'snappy').save(output_path)
logger.info(f"Data written to Parquet at {output_path}")
s3_bucket = os.getenv("S3_BUCKET")
output_path = f"s3a://{s3_bucket}/processed/transactions/"
write_to_parquet(
final_df,
output_path,
partition_cols=['year', 'month', 'day']
)
6. Load to PostgreSQL
from pyspark.sql import DataFrame
import os
import logging
logger = logging.getLogger(__name__)
def load_to_postgres(df: DataFrame, table_name: str, mode: str = 'overwrite'):
"""
Load DataFrame to PostgreSQL
Args:
df: DataFrame to load
table_name: Target table name
mode: Write mode ('overwrite', 'append')
"""
postgres_url = (
f"jdbc:postgresql://{os.getenv('POSTGRES_HOST')}:"
f"{os.getenv('POSTGRES_PORT')}/{os.getenv('POSTGRES_DB')}"
)
connection_properties = {
"user": os.getenv("POSTGRES_USER"),
"password": os.getenv("POSTGRES_PASSWORD"),
"driver": "org.postgresql.Driver"
}
df.write.jdbc(
url=postgres_url,
table=table_name,
mode=mode,
properties=connection_properties
)
logger.info(f"Data loaded to PostgreSQL table '{table_name}'")
load_to_postgres(
final_df,
table_name='transactions_processed',
mode='overwrite'
)
7. Complete Pipeline
from pyspark.sql import SparkSession
from src.extract.s3_extractor import extract_from_s3
from src.validation.schema_validator import define_expected_schema, validate_schema
from src.validation.quality_checker import DataQualityChecker
from src.transform.transformer import DataTransformer
from src.load.parquet_writer import write_to_parquet
from src.load.postgres_loader import load_to_postgres
import logging
import os
from dotenv import load_dotenv
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
handlers=[
logging.FileHandler('logs/pipeline.log'),
logging.StreamHandler()
]
)
logger = logging.getLogger(__name__)
def main():
"""Main pipeline execution"""
load_dotenv()
spark = SparkSession.builder \
.appName("DataEngineeringPipeline") \
.config("spark.jars.packages",
"org.apache.hadoop:hadoop-aws:3.3.1,"
"org.postgresql:postgresql:42.5.0") \
.config("spark.sql.adaptive.enabled", "true") \
.config("spark.sql.adaptive.coalescePartitions.enabled", "true") \
.getOrCreate()
try:
logger.info("Step 1: Extracting data from S3")
s3_bucket = os.getenv()
raw_df = extract_from_s3(spark, )
logger.info()
expected_schema = define_expected_schema()
validate_schema(raw_df, expected_schema):
ValueError()
logger.info()
qc = DataQualityChecker(raw_df)
qc.check_null_values([, , ])
qc.check_duplicates([])
qc.check_value_ranges(, , )
qc.check_categorical_values(, [, , ])
logger.info()
transformer = DataTransformer()
cleaned_df = transformer.clean_data(raw_df)
enriched_df = transformer.add_derived_columns(cleaned_df)
final_df = transformer.add_aggregations(enriched_df)
logger.info()
parquet_path =
write_to_parquet(final_df, parquet_path, [, ])
logger.info()
load_to_postgres(final_df, , mode=)
logger.info()
Exception e:
logger.error(, exc_info=)
:
spark.stop()
__name__ == :
main()
Testing
import pytest
from pyspark.sql import SparkSession
from pyspark.sql.types import StructType, StructField, StringType, IntegerType, TimestampType
from src.transform.transformer import DataTransformer
from datetime import datetime
@pytest.fixture(scope="session")
def spark():
"""Create Spark session for testing"""
return SparkSession.builder \
.appName("TestPipeline") \
.master("local[2]") \
.getOrCreate()
@pytest.fixture
def sample_data(spark):
"""Create sample test data"""
schema = StructType([
StructField("user_id", StringType()),
StructField("transaction_id", StringType()),
StructField("amount", IntegerType()),
StructField("currency", StringType()),
StructField("timestamp", TimestampType())
])
data = [
("user1", "txn1", 100, "USD", datetime(2024, 1, 15, 10, 30)),
("user1", "txn2", 200, "EUR", datetime(2024, 1, 16, 14, )),
(, , , , datetime(, , , , ))
]
spark.createDataFrame(data, schema)
():
transformer = DataTransformer()
result = transformer.add_derived_columns(sample_data)
result.columns
result.columns
result.columns
result.columns
first_row = result.(result.transaction_id == ).first()
first_row[] ==
first_row[] ==
first_row[] ==
():
data = [
(, , , , datetime(, , )),
(, , , , datetime(, , )),
(, , , , datetime(, , ))
]
schema = StructType([
StructField(, StringType()),
StructField(, StringType()),
StructField(, IntegerType()),
StructField(, StringType()),
StructField(, TimestampType())
])
df = spark.createDataFrame(data, schema)
transformer = DataTransformer()
result = transformer.clean_data(df)
result.count() ==
Common Patterns
Pattern 1: Incremental Data Processing
from datetime import datetime, timedelta
last_processed_date = datetime(2024, 1, 1)
incremental_df = raw_df.filter(
F.col('timestamp') > F.lit(last_processed_date)
)
load_to_postgres(incremental_df, 'transactions_processed', mode='append')
Pattern 2: Error Handling and Data Quarantine
valid_df = raw_df.filter(F.col('amount') > 0)
invalid_df = raw_df.filter(F.col('amount') <= 0)
invalid_df.write.mode('append').parquet(
f"s3a://{s3_bucket}/quarantine/invalid_amounts/"
)
process(valid_df)
Pattern 3: Data Versioning
from datetime import datetime
versioned_df = final_df.withColumn(
'processing_date', F.lit(datetime.now())
).withColumn(
'version', F.lit('v1.0')
)
write_to_parquet(
versioned_df,
f"s3a://{s3_bucket}/processed/transactions/",
partition_cols=['year', 'month', 'version']
)
Troubleshooting
Issue: S3 Connection Errors
spark._jsc.hadoopConfiguration().set("fs.s3a.aws.credentials.provider",
"org.apache.hadoop.fs.s3a.SimpleAWSCredentialsProvider")
spark.sparkContext.setLogLevel("DEBUG")
Issue: PostgreSQL Connection Timeout
connection_properties = {
"user": os.getenv("POSTGRES_USER"),
"password": os.getenv("POSTGRES_PASSWORD"),
"driver": "org.postgresql.Driver",
"connectTimeout": "60",
"socketTimeout": "60"
}
Issue: Out of Memory Errors
df_repartitioned = large_df.repartition(100)
df_coalesced = df_repartitioned.coalesce(10)
spark = SparkSession.builder \
.config("spark.executor.memory", "4g") \
.config("spark.driver.memory", "2g") \
.getOrCreate()
Issue: Schema Evolution
from pyspark.sql.utils import AnalysisException
try:
df = spark.read.parquet(path)
except AnalysisException:
df = spark.read.option("mergeSchema", "true").parquet(path)
Performance Optimization
Optimize Parquet Writes
df.write \
.option("maxRecordsPerFile", 100000) \
.option("compression", "snappy") \
.partitionBy("year", "month") \
.parquet(output_path)
Cache Intermediate Results
cleaned_df = transformer.clean_data(raw_df)
cleaned_df.cache()
enriched_df = transformer.add_derived_columns(cleaned_df)
aggregated_df = transformer.add_aggregations(cleaned_df)
cleaned_df.unpersist()
Broadcast Small Lookup Tables
from pyspark.sql.functions import broadcast
merchants_df = spark.read.parquet("s3a://bucket/merchants/")
result = transactions_df.join(
broadcast(merchants_df),
on='merchant_id',
how='left'
)