AWS Training
Modules Listen Certification
0:00 0:00

← PySpark

Starts this lesson and continues through 5 more to the end of the course.

Production PySpark — testing, structure, and AWS

The structural problem

Most PySpark in the wild is one long script: read, twenty transformations, write. It works, and it is untestable — because there is no unit of behaviour smaller than "the whole job", and running the whole job needs a cluster and real data.

The fix is not a framework. It's one habit: separate the transformation from the I/O.

# transforms.py — pure, testable, no SparkSession creation, no paths
def clean_events(df: DataFrame) -> DataFrame:
    return (df
            .filter(col("event_type").isNotNull())
            .withColumn("event_ts", to_timestamp("raw_ts"))
            .dropDuplicates(["event_id"]))

def daily_totals(events: DataFrame) -> DataFrame:
    return events.groupBy("dt", "user_id").agg(sum("amount").alias("total"))

# job.py — the thin shell that knows about the world
def main(spark, input_path, output_path):
    raw = spark.read.parquet(input_path)
    out = daily_totals(clean_events(raw))
    out.write.mode("overwrite").partitionBy("dt").parquet(output_path)

Every function takes a DataFrame and returns a DataFrame. Now you can test the logic with five rows and no cluster.

Testing without a cluster

A local SparkSession runs in a single JVM on your laptop or in CI. It is slow to start and fine for correctness:

import pytest
from pyspark.sql import SparkSession

@pytest.fixture(scope="session")
def spark():
    s = (SparkSession.builder
         .master("local[2]")
         .appName("tests")
         .config("spark.sql.shuffle.partitions", "2")   # 200 is absurd for 5 rows
         .getOrCreate())
    yield s
    s.stop()

def test_clean_events_drops_null_types(spark):
    df = spark.createDataFrame(
        [("1", "click", "2026-01-01 00:00:00"), ("2", None, "2026-01-01 00:00:00")],
        ["event_id", "event_type", "raw_ts"])
    out = clean_events(df)
    assert out.count() == 1
    assert out.collect()[0]["event_id"] == "1"

⚠️ Set spark.sql.shuffle.partitions low in tests. The default is large, and creating hundreds of empty partitions for a five-row test dominates your test runtime. This one line commonly takes a suite from minutes to seconds.

What to test:

What not to test: that Spark works. Don't write tests that assert filter filters.

Assertions in the job itself

Tests cover logic on synthetic data. They cannot cover today's real data. Put the SQL0 correctness checks in the job:

def assert_unique(df: DataFrame, keys: list[str]) -> None:
    dupes = df.groupBy(*keys).count().filter(col("count") > 1).limit(1).count()
    if dupes:
        raise ValueError(f"Duplicate keys in output on {keys}")

def assert_not_empty(df: DataFrame) -> None:
    if not df.take(1):
        raise ValueError("Output is empty — upstream may have failed")

⚠️ Fail the job rather than writing bad data. A failed run gets attention within the hour. Silently wrong output gets attention in a quarter, from someone outside the team. That asymmetry should decide your design.

Note the cost, though: each assertion is an action, and therefore a full execution unless the DataFrame is cached. Assert on the thing you're about to write, after caching it, or accept the recomputation deliberately.

Configuration and idempotency

Never hardcode paths. Take them as arguments so the same code runs in dev, test, and prod, and so tests can point at a temp directory.

Make the job idempotent. Re-running should produce the same result, not double it:

# Idempotent: replaces just this date's output
out.write.mode("overwrite").partitionBy("dt").parquet(path)

# NOT idempotent: a retry doubles the data
out.write.mode("append").parquet(path)

⚠️ mode("overwrite") semantics with partitionBy are version- and configuration-dependent — whether it replaces the whole dataset or only the partitions being written depends on a dynamic-overwrite setting. I have not verified the current key and default; check the docs for your Spark version before relying on partition-level overwrite. Getting this wrong deletes data.

Retries are not hypothetical: orchestrators retry, spot instances get reclaimed, and Glue and Step Functions both retry by default. Assume every job will run twice.

What changes on AWS

Find out which Spark version you actually have

Every version-dependent claim in this module — Arrow defaults, AQE behaviour, overwrite semantics — depends on this, and Glue and EMR pin versions that are typically behind the current release.

print(spark.version)                                   # from inside the job
print(spark.conf.get("spark.sql.execution.arrow.pyspark.enabled"))
print(spark.sparkContext.getConf().getAll())           # everything, for the record

Do this once per environment and write the output into your runbook. Then read the PySpark docs for that version, not the latest. This is the single most common source of "the docs said it was enabled by default but it isn't".

Glue specifics

The driver is smaller than you think

Lesson 1's driver-killers matter more on managed services, where driver memory is set by the instance/worker type you chose rather than by you directly. A collect() that survived on your laptop can kill a small Glue worker.

Logging

print() from the driver reaches CloudWatch. print() from inside a UDF runs on an executor and goes to executor logs — a different place, and frequently the reason people conclude "my UDF isn't running". Module CW0 covers CloudWatch; the practical point is to know which process a log line comes from before hunting for it.

Check yourself

  1. Why is a single long script untestable?
  2. What one config makes local tests dramatically faster?
  3. Why put correctness assertions in the job when you already have tests?
  4. Why is mode("append") dangerous with retries?
  5. Before trusting any Arrow default in this module, what do you check?
Answers
  1. There's no unit of behaviour smaller than the whole job, and running it requires a cluster and real data. Splitting DataFrame-in/DataFrame-out transformations from I/O makes the logic testable with a handful of rows.
  2. spark.sql.shuffle.partitions set to a small number. The default creates hundreds of empty partitions per shuffle, which dominates runtime for tiny test data.
  3. Tests run against synthetic data you invented. Assertions run against today's real data. Only the second catches an upstream change, and failing loudly beats writing wrong data silently.
  4. Orchestrators, spot reclamation, Glue and Step Functions all retry. An append-mode job that runs twice writes the data twice. Assume every job runs at least twice.
  5. spark.version in your actual environment, then the PySpark docs for that version — not the latest. Glue and EMR pin versions behind the current release.

Teaching this section

← PreviousPartitions, skew, and the shuffle you causedFinished →Cheat sheet, lab & quiz