Skip to content
Good engineers know 5 min read · Testing

Testing PySpark code

Unit test transformations with a local SparkSession and assertDataFrameEqual, and the cases worth covering.

Read first: select · groupBy + agg (skip if you know them)

Why it comes up

Interviews for data engineering roles increasingly ask how you test pipelines, not only how you write them. The honest answer for most PySpark code is: pull the logic into functions that take and return DataFrames, and test those functions on tiny hand-written data with a local SparkSession.

Write testable transformations

Keep reading and writing at the edges. Everything between is a function from DataFrames to a DataFrame, with no I/O inside:

PySpark · Logic as pure functions, chained with transform
from pyspark.sql import DataFrame, functions as F

def paid_only(df: DataFrame) -> DataFrame:
    return df.filter(F.col("status") == "PAID")

def revenue_by_country(df: DataFrame) -> DataFrame:
    return df.groupBy("country").agg(F.sum("amount").alias("revenue"))

# The job: I/O at the edges, logic in the middle.
orders = spark.read.table("raw.orders")
result = orders.transform(paid_only).transform(revenue_by_country)
result.write.mode("overwrite").saveAsTable("gold.revenue_by_country")

DataFrame.transform applies a function and returns its result, so steps chain like built-in methods. Functions that need parameters can take them first and return the inner function, or be called with a lambda.

A fast local SparkSession

PySpark · conftest.py: one session for the whole test run
import pytest
from pyspark.sql import SparkSession

@pytest.fixture(scope="session")
def spark():
    session = (SparkSession.builder
        .master("local[2]")
        .appName("tests")
        .config("spark.sql.shuffle.partitions", "2")   # not 200: tiny data
        .config("spark.ui.enabled", "false")
        .getOrCreate())
    yield session
    session.stop()
  • Session scope: starting a JVM takes seconds; do it once per test run, not per test.
  • Few shuffleshuffle: Moving rows between machines so that all rows with the same key end up together. Needed by joins, groupBy and sorting, and usually the most expensive step of a job. Learn more → partitionspartition: A chunk of a DataFrame's rows. Spark processes each partition as one task, so partitions decide how much work runs in parallel. Learn more →: with 200, every groupBy on five rows launches 200 taskstask: The work for one partition in one stage, run on one CPU core. Learn more →. Two keeps tests fast.
  • Tiny data: build DataFrames inline with spark.createDataFrame and an explicit schema, a handful of rows chosen to hit each case.

Comparing DataFrames

PySpark 3.5 added pyspark.testing. assertDataFrameEqual compares schemas and rows and prints a readable diff when they differ:

PySpark · A test with assertDataFrameEqual
from pyspark.testing import assertDataFrameEqual, assertSchemaEqual

def test_revenue_by_country(spark):
    orders = spark.createDataFrame(
        [("IN", "PAID", 100), ("IN", "PAID", 50), ("US", "PAID", 70), ("US", "REFUNDED", 30)],
        "country string, status string, amount int")
    expected = spark.createDataFrame(
        [("IN", 150), ("US", 70)], "country string, revenue bigint")

    actual = orders.transform(paid_only).transform(revenue_by_country)

    assertDataFrameEqual(actual, expected)          # row order ignored by default
    assertSchemaEqual(actual.schema, expected.schema)
OptionUse
checkRowOrder=False (default)Rows compared as a bag. Set True only when order is part of the contract, after an orderBy.
rtol, atolTolerances for floating-point columns, so averages and ratios do not fail on the last digit.
Spark 4.0 additionsOptions to ignore column order, names or types, and to show only the differing rows, for wider DataFrames.

Note the expected type: sum of an int column is a bigint. Type mismatches like that are exactly what schema checks catch. On Spark versions before 3.5, the chispa library offers similar assertions.

What to test

  • Nulls in keys, in values, in filter columns. Most real bugs live here.
  • Duplicates: duplicate keys in a join or a dedup step, ties in a ranking.
  • Empty input: an empty DataFrame should give an empty result with the right schema, not an error.
  • Boundaries: dates at month and year ends, time zones, the first and last row of a window.
  • Schema: the output columns and types a downstream consumer relies on.
Tip: unit tests on tiny data check logic. They do not check performance or data quality in production. Pair them with data quality checks in the pipeline (constraints, row counts, null rates) and a run on a production-sized sample before release.

Common mistakes

  • Testing whole jobs end to end only — Slow, and a failure does not say which step broke. Test the transformation functions.
  • Reading files inside transformation functions — Then tests need files. Keep I/O at the edges and pass DataFrames in.
  • Comparing collect() output with == — Row order is not guaranteed and floats differ slightly. Use assertDataFrameEqual with tolerances.
  • A new SparkSession per test — Each one costs seconds; use a session-scoped fixture.

Interview prep

THE QUESTION

"How do you unit test a PySpark transformation?"

Structure the logic as functions that take and return DataFrames, with reading and writing outside them. In pytest, use a session-scoped local SparkSession with few shuffle partitions, build tiny input and expected DataFrames inline with explicit schemas, and compare with assertDataFrameEqual (PySpark 3.5+), which ignores row order by default and supports float tolerances. Cover nulls, duplicates, empty input and boundaries, and check the output schema.

Avoid saying: "we test in production with data checks" or "we run the job on a sample and eyeball it". Neither tests the logic in isolation.

What the interviewer asks next. Answer out loud first, then open the strong answer.

Follow-up"How do you test a function that needs the current date?"
Pass the date in as a parameter (or a column) instead of calling current_date() inside, so the test can fix it. The same goes for anything non-deterministic: random salts, run ids, environment config.
Scenario"Your test suite takes 15 minutes. How do you speed it up?"
Use one session-scoped SparkSession, set shuffle partitions to 1 or 2, disable the UI, keep test data to a handful of rows, and avoid writing to disk in unit tests. Move slow end-to-end tests to a separate, less frequent stage.
Trap"assertDataFrameEqual passed, so the pipeline is correct."
It proves the function is right on the cases you wrote. Coverage of nulls, duplicates and boundaries decides how much that is worth, and unit tests say nothing about real data volumes or quality. Pair them with data checks in the pipeline.

What you learned

  • How to structure PySpark code so it can be tested
  • How to set up a fast local SparkSession for tests
  • How to compare DataFrames with assertDataFrameEqual and assertSchemaEqual
  • Which cases a good test suite covers

Key takeaways

  • Put logic in DataFrame-to-DataFrame functions; keep I/O at the edges.
  • Use one local SparkSession per test run, with few shuffle partitions.
  • assertDataFrameEqual compares schema and rows, ignoring order by default.
  • Test nulls, duplicates, empty input and boundaries.

Check yourself

3 questions

Why set spark.sql.shuffle.partitions low in tests?

Show the answer

Tiny data with 200 partitions launches 200 near-empty tasks per shuffle. Task overhead dominates on tiny data.

By default, does assertDataFrameEqual care about row order?

Show the answer

No; set checkRowOrder=True if it should. Rows are compared as a bag unless checkRowOrder is True.

What does DataFrame.transform(f) do?

Show the answer

Calls f(df) and returns its result, so functions chain. It is a convenience for chaining DataFrame functions.

Keep going

Up next · lesson 28 of 37 · 5 min read
rangeBetween
Window frames by value instead of row count: the last 7 calendar days, not the last 7 rows.

Related lessons

Previous: UDFs and pandas UDFs

Primary sources: pyspark.testing · Testing PySpark