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에서 보기