Skip to main content

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 查看