spark-data-engineering-pipeline
Production PySpark ETL pipeline with AWS S3, PostgreSQL, schema validation, data quality checks, and automated testing
来源信息
- 仓库
- reason-machines/data-skills
- 最近来源活动
- 2026年8月5日 00:47
- 检测到的 SKILL.md 语言
- 英语
- 星标
- 5
- 分支
- 1
安装方式
默认使用会先检查来源的 Prompt;你也可以切换为直接命令,或下载本地副本。
检查来源文件
决定是否安装前,请先阅读 SKILL.md,以及 SkillsMP 当前展示的配套文件。
正在显示 SKILL.md
SKILL.md
来源说明 · 只读预览- 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](https://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
```bash
# Clone the repository
git clone https://github.com/Giri-25/spark-data-engineering-pipeline.git
cd spark-data-engineering-pipeline
# Create and activate virtual environment
python3 -m venv .venv
source .venv/bin/activate # On Windows: .venv\Scripts\activate
# Install dependencies
pip install -r requirements.txt
```
## Configuration
Create a `.env` file in the project root:
```bash
# AWS Configuration
AWS_ACCESS_KEY_ID=your_access_key
AWS_SECRET_ACCESS_KEY=your_secret_key
AWS_REGION=us-east-1
S3_BUCKET=your-bucket-name
# PostgreSQL Configuration
POSTGRES_HOST=localhost
POSTGRES_PORT=5432
POSTGRES_DB=data_warehouse
POSTGRES_USER=your_user
POSTGRES_PASSWORD=your_password
```
## Running the Pipeline
```bash
# Run the complete pipeline
python main.py
# Run with custom configuration
python main.py --config configs/pipeline_config.yaml
# Run unit tests
pytest tests/
# Run tests with coverage
pytest --cov=src tests/
```
## Core Components
### 1. Data Extraction from S3
```python
# src/extract/s3_extractor.py
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
"""
# Configure Spark for S3 access
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"
)
# Read JSON data
df = spark.read.json(s3_path)
return df
# Usage
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
```python
# src/validation/schema_validator.py
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())
# Check for missing fields
missing_fields = expected_fields - df_fields
if missing_fields:
logger.error(f"Missing required fields: {missing_fields}")
return False
# Check for extra fields
extra_fields = df_fields - expected_fields
if extra_fields:
logger.warning(f"Extra fields found: {extra_fields}")
# Validate data types
for field in expected_schema.fields:
if field.name in df.schema.fieldNames():
df_field = df.schema[field.name]
if df_field.dataType != field.dataType:
logger.error(
f"Type mismatch for {field.name}: "
f"expected {field.dataType}, got {df_field.dataType}"
)
return False
logger.info("Schema validation passed")
return True
# Usage
expected_schema = define_expected_schema()
if validate_schema(raw_df, expected_schema):
validated_df = raw_df
else:
raise ValueError("Schema validation failed")
```
### 3. Data Quality Checks
```python
# src/validation/quality_checker.py
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
if duplicates > 0:
logger.warning(f"Found {duplicates} duplicate records")
self.quality_report['duplicates'] = duplicates
return duplicates
def check_value_ranges(self, column: str, min_val, max_val) -> int:
"""Check if values are within expected range"""
out_of_range = self.df.filter(
(F.col(column) < min_val) | (F.col(column) > max_val)
).count()
if out_of_range > 0:
logger.warning(
f"Column '{column}' has {out_of_range} values outside range "
f"[{min_val}, {max_val}]"
)
self.quality_report[f'{column}_out_of_range'] = out_of_range
return out_of_range
def check_categorical_values(self, column: str, allowed_values: list) -> int:
"""Check if categorical values are in allowed list"""
invalid_count = self.df.filter(
~F.col(column).isin(allowed_values)
).count()
if invalid_count > 0:
logger.warning(
f"Column '{column}' has {invalid_count} invalid values"
)
self.quality_report[f'{column}_invalid'] = invalid_count
return invalid_count
def get_report(self) -> dict:
"""Return complete quality report"""
return self.quality_report
# Usage
qc = DataQualityChecker(validated_df)
# Check for nulls in required fields
qc.check_null_values(['user_id', 'transaction_id', 'amount'])
# Check for duplicates
qc.check_duplicates(['transaction_id'])
# Check value ranges
qc.check_value_ranges('amount', min_val=0, max_val=1000000)
# Check categorical values
qc.check_categorical_values('status', ['pending', 'completed', 'failed'])
# Get quality report
quality_report = qc.get_report()
logger.info(f"Quality report: {quality_report}")
```
### 4. Data Transformation
```python
# src/transform/transformer.py
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"""
# Remove duplicates
df_cleaned = df.dropDuplicates(['transaction_id'])
# Fill nulls in optional fields
df_cleaned = df_cleaned.fillna({
'status': 'unknown',
'merchant_id': 'unknown'
})
# Remove rows with nulls in required fields
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"""
# Extract date components
df_transformed = df.withColumn('date', F.to_date('timestamp'))
df_transformed = df_transformed.withColumn('year', F.year('timestamp'))
df_transformed = df_transformed.withColumn('month', F.month('timestamp'))
df_transformed = df_transformed.withColumn('day', F.dayofmonth('timestamp'))
df_transformed = df_transformed.withColumn('hour', F.hour('timestamp'))
# Add day of week
df_transformed = df_transformed.withColumn(
'day_of_week',
F.dayofweek('timestamp')
)
# Convert amount to USD (assuming conversion logic)
df_transformed = df_transformed.withColumn(
'amount_usd',
F.when(F.col('currency') == 'EUR', F.col('amount') * 1.1)
.when(F.col('currency') == 'GBP', F.col('amount') * 1.3)
.otherwise(F.col('amount'))
)
return df_transformed
@staticmethod
def add_aggregations(df: DataFrame) -> DataFrame:
"""Add user-level aggregations"""
# Window specification for user-level aggregations
user_window = Window.partitionBy('user_id')
df_agg = df.withColumn(
'user_total_transactions',
F.count('transaction_id').over(user_window)
)
df_agg = df_agg.withColumn(
'user_total_amount',
F.sum('amount_usd').over(user_window)
)
df_agg = df_agg.withColumn(
'user_avg_amount',
F.avg('amount_usd').over(user_window)
)
return df_agg
# Usage
transformer = DataTransformer()
# Clean data
cleaned_df = transformer.clean_data(validated_df)
# Add derived columns
enriched_df = transformer.add_derived_columns(cleaned_df)
# Add aggregations
final_df = transformer.add_aggregations(enriched_df)
logger.info(f"Transformation complete: {final_df.count()} rows")
```
### 5. Write to Parquet
```python
# src/load/parquet_writer.py
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)
# Write with compression
writer.option('compression', 'snappy').save(output_path)
logger.info(f"Data written to Parquet at {output_path}")
# Usage
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
```python
在 GitHub 查看这个 SKILL.md 很大,SkillsMP 这里只预览前一段内容。 在 GitHub 查看