Testing PySpark code
Unit test transformations with a local SparkSession and assertDataFrameEqual, and the cases worth covering.
On this page
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:
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
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.createDataFrameand 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:
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)
| Option | Use |
|---|---|
checkRowOrder=False (default) | Rows compared as a bag. Set True only when order is part of the contract, after an orderBy. |
rtol, atol | Tolerances for floating-point columns, so averages and ratios do not fail on the last digit. |
| Spark 4.0 additions | Options 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.
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?"
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?"
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?"
Trap"assertDataFrameEqual passed, so the pipeline is correct."
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 questionsWhy 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 readrangeBetween
Window frames by value instead of row count: the last 7 calendar days, not the last 7 rows.
Related lessons
UDFs and pandas UDFsPySpark functions · 4 min read
cast and schemasPySpark functions · 4 min read
spark.sql and temp views
Previous: UDFs and pandas UDFs
Primary sources: pyspark.testing · Testing PySpark