In the SQL Server world, testing ETL often meant "run it in UAT and compare row counts". tSQLt existed, but few teams used it. With Spark, the same habit often shows up in notebooks: logic that's only ever tested by running the whole notebook against real bordereaux and eyeballing the output.
That's risky in delegated authority. A small mistake in how commission is netted off, how cancellations are signed, or which resubmission wins can misstate written premium for a whole binding authority, and nobody notices until the reconciliation with the broker fails. PySpark transformations are just Python functions that take DataFrames and return DataFrames, so they can be unit tested with pytest, locally, in seconds, without a cluster. Spark 3.5 also added built-in testing helpers that remove most of the boilerplate. Here's a setup that works well.
Step 1: Get the logic out of the notebook
You can't easily unit test a notebook cell that reads a table, transforms it and writes it back. Separate the three:
1: # src/transforms/premium.py 2: from pyspark.sql import DataFrame, functions as F, Window 3: 4: 5: def latest_submission(df: DataFrame) -> DataFrame: 6: """Keep the latest submitted version of each premium transaction.""" 7: keys = ["umr", "section_no", "certificate_ref", "transaction_seq"] 8: w = Window.partitionBy(*keys).orderBy(F.col("submitted_at").desc(), 9: F.col("bdx_version").desc()) 10: return (df.withColumn("_rn", F.row_number().over(w)) 11: .filter("_rn = 1") 12: .drop("_rn")) 13: 14: 15: def derive_premium_amounts(df: DataFrame) -> DataFrame: 16: """Sign cancellations negative and derive commission and net premium.""" 17: sign = F.when(F.col("transaction_type") == "CANCELLATION", F.lit(-1)).otherwise(F.lit(1)) 18: gross = F.abs(F.col("gross_premium")) * sign 19: commission = F.round(gross * F.coalesce(F.col("commission_pct"), F.lit(0)) / 100, 2) 20: return (df.withColumn("gross_premium_signed", gross) 21: .withColumn("commission_amount", commission) 22: .withColumn("net_premium", gross - commission))
The notebook (or job) becomes a thin wrapper:
1: from transforms.premium import latest_submission, derive_premium_amounts 2: 3: raw = spark.read.table("bronze.premium_bdx") 4: clean = derive_premium_amounts(latest_submission(raw)) 5: clean.write.mode("overwrite").saveAsTable("silver.premium_transaction")
In Fabric you can package src/ as a wheel and attach it to an Environment. In Glue, use --extra-py-files. On Databricks, a workspace file or wheel. The tests stay the same everywhere.
Step 2: A SparkSession fixture
1: # tests/conftest.py 2: import pytest 3: from pyspark.sql import SparkSession 4: 5: 6: @pytest.fixture(scope="session") 7: def spark(): 8: spark = (SparkSession.builder 9: .master("local[1]") 10: .appName("unit-tests") 11: .config("spark.sql.shuffle.partitions", "1") 12: .config("spark.default.parallelism", "1") 13: .config("spark.ui.enabled", "false") 14: .config("spark.sql.session.timeZone", "UTC") 15: .getOrCreate()) 16: yield spark 17: spark.stop()
Three settings do most of the work:
scope="session": start Spark once for the whole test run. Starting a session takes a few seconds, and doing it per test makes the suite painfully slow.shuffle.partitions = 1: the default of 200 partitions is designed for clusters. On tiny test data it just creates 200 empty tasks.- A fixed session time zone: otherwise submission timestamp tests pass on your laptop and fail on the CI runner.
Step 3: Tests with assertDataFrameEqual
Spark 3.5 added pyspark.testing with assertDataFrameEqual and assertSchemaEqual. They give readable diffs when something doesn't match, which is a big step up from comparing collect() output by hand.
1: # tests/test_premium.py 2: from datetime import datetime 3: from pyspark.testing import assertDataFrameEqual 4: from transforms.premium import latest_submission, derive_premium_amounts 5: 6: BDX_SCHEMA = ("umr STRING, section_no STRING, certificate_ref STRING, " 7: "transaction_seq INT, gross_premium DOUBLE, submitted_at TIMESTAMP, " 8: "bdx_version INT") 9: 10: 11: def test_resubmission_replaces_original(spark): 12: src = spark.createDataFrame( 13: [ 14: ("B0999CH001", "S1", "CERT-1001", 1, 1200.0, datetime(2025, 9, 5, 10, 0), 1), 15: ("B0999CH001", "S1", "CERT-1001", 1, 1250.0, datetime(2025, 9, 12, 9, 0), 2), # corrected 16: ("B0999CH001", "S1", "CERT-1002", 1, 800.0, datetime(2025, 9, 5, 10, 0), 1), 17: ], 18: BDX_SCHEMA, 19: ) 20: 21: expected = spark.createDataFrame( 22: [ 23: ("B0999CH001", "S1", "CERT-1001", 1, 1250.0, datetime(2025, 9, 12, 9, 0), 2), 24: ("B0999CH001", "S1", "CERT-1002", 1, 800.0, datetime(2025, 9, 5, 10, 0), 1), 25: ], 26: BDX_SCHEMA, 27: ) 28: 29: assertDataFrameEqual(latest_submission(src), expected) 30: 31: 32: def test_same_timestamp_uses_highest_version(spark): 33: ts = datetime(2025, 9, 5, 10, 0) 34: src = spark.createDataFrame( 35: [ 36: ("B0999CH001", "S1", "CERT-1001", 1, 1200.0, ts, 1), 37: ("B0999CH001", "S1", "CERT-1001", 1, 1300.0, ts, 2), 38: ], 39: BDX_SCHEMA, 40: ) 41: result = latest_submission(src).select("gross_premium") 42: assertDataFrameEqual(result, [(1300.0,)]) 43: 44: 45: def test_cancellation_is_negative_and_commission_nets_off(spark): 46: src = spark.createDataFrame( 47: [ 48: ("CERT-1", "NEW", 1000.0, 25.0), 49: ("CERT-2", "CANCELLATION", 400.0, 25.0), # some coverholders send positives 50: ("CERT-3", "CANCELLATION", -400.0, 25.0), # others send negatives 51: ("CERT-4", "MTA", 100.0, None), # missing commission 52: ], 53: "certificate_ref STRING, transaction_type STRING, gross_premium DOUBLE, commission_pct DOUBLE", 54: ) 55: 56: result = derive_premium_amounts(src).select( 57: "certificate_ref", "gross_premium_signed", "commission_amount", "net_premium") 58: 59: expected = spark.createDataFrame( 60: [ 61: ("CERT-1", 1000.0, 250.0, 750.0), 62: ("CERT-2", -400.0, -100.0, -300.0), 63: ("CERT-3", -400.0, -100.0, -300.0), 64: ("CERT-4", 100.0, 0.0, 100.0), 65: ], 66: "certificate_ref STRING, gross_premium_signed DOUBLE, " 67: "commission_amount DOUBLE, net_premium DOUBLE", 68: ) 69: 70: assertDataFrameEqual(result, expected)
Some useful behaviour:
- Row order is ignored by default. Pass
checkRowOrder=Truewhen order matters, for example after a sort. - Floating-point values are compared with a tolerance (
rtol/atol), so0.1 + 0.2doesn't fail your test. For real money amounts, useDECIMAL, which the helpers compare exactly. - The expected value can be a list of rows instead of a DataFrame, which is handy for small cases like the version tie-break test.
What to test
Don't test Spark itself. You don't need a test proving filter works. Focus on your business rules, especially:
- Sign conventions: cancellations and return premiums. Coverholders are inconsistent about whether they send them as positive or negative amounts, as the test above shows.
- Nulls: missing commission percentages, missing currencies, missing inception dates. Every
CASE WHENneeds a null test. - Resubmissions and ties: two versions with the same timestamp. Without a tie-breaker, the result isn't deterministic, which is why
bdx_versionis in the window ordering. - Boundaries: risks incepting on the first or last day of the contract section's period. Off-by-one on
<vs<=is a classic bug that assigns a risk to the wrong year of account. - Schema contracts: use
assertSchemaEqualon outputs that the premium fact or Power BI depends on, so a renamed column fails in CI rather than in an underwriter's report. - Empty input: a coverholder with no business this month should produce an empty DataFrame with the correct schema, not an error.
Running it in CI
If you already run build and deployment workflows in GitHub Actions, adding a test step is a small change:
1: # requirements-dev.txt 2: pyspark==3.5.* 3: pytest
1: # .github/workflows/tests.yml 2: name: tests 3: on: [push, pull_request] 4: jobs: 5: test: 6: runs-on: ubuntu-latest 7: steps: 8: - uses: actions/checkout@v4 9: - uses: actions/setup-java@v4 10: with: { distribution: temurin, java-version: '17' } 11: - uses: actions/setup-python@v5 12: with: { python-version: '3.11' } 13: - run: pip install -r requirements-dev.txt -e . 14: - run: pytest -q
Match the PySpark version to your target runtime, whether that's the Fabric runtime, the Glue version or the Databricks Runtime, so behaviour is the same in tests and production. A suite of 50–100 tests like these runs in well under a minute.
Integration tests are still needed
Unit tests catch logic bugs. They won't catch a permissions problem, a missing table, or a coverholder who suddenly changes their bordereau layout. You still want an end-to-end run in a test workspace with a sample month of bordereaux before release. With unit tests in place, that run should fail far less often, and when it does, the cause is more likely to be environmental than a business rule. That's the kind of failure you want to be left with.