spark - code

24m read · 4798 words

case

when, otherwise

from pyspark.sql.functions import when, col

df = df.withColumn(
    "age_group",
    when(col("age").isNull(),  "unknown")
    .when(col("age") < 0,      "invalid")
    .when(col("age") < 18,     "minor")
    .otherwise("adult")
)

filters - where

# multiple conditions
df.where((col("age") > 18) & (col("city") == "delhi"))
df.where("age > 18 AND city = 'delhi'") # SQL string - same result

# in
df.where(col("city").isin("pune", "mumbai"))
# not in
df.where(~col("city").isin("pune", "mumbai"))

# between - inclusive (lowerBound <= val <= upperbound)
df.where(col("age").between(18, 60))

# date
df.where(col("dt") >= "2024-01-01")
df.where(col("dt").between("2024-01-01", "2024-12-31"))

# filter JSON
df.where(col("address.city") == "pune")

# filter array — check if value exists in array
from pyspark.sql.functions import array_contains
df.where(array_contains(col("tags"), "premium"))

# after explode
df.where(col("order.amt") > 500)

join

df1.join(df2, df1.id == df2.id, "inner")
  1. inner
  2. left
  3. right
  4. full
  5. left_semi
  6. left_anti
  7. cross

Join with aliases

a = df1.alias("a")
b = df2.alias("b")

result = a.join(
    b,
    (a.id == b.id) & (a.country == b.country), # Use `&`, not Python `and`.
    "inner"
).select(
    F.col("a.*"),
    F.col("b.name").alias("city_name")
)

Hints

# use on df on either side of the join
df.hint("merge")

df1.hint("merge")
        .join(
            df2.hint("merge"),
            "id"
        )
  1. broadcast
  2. merge
  3. shuffle_hash
  4. shufflereplicatenl

window functions

  1. rank
    1. rownumber
    2. rank, denserank
    3. percentrank, ntile
    4. cumedist
  2. aggregation
    1. sum
    2. avg
    3. min
    4. max
    5. count
  3. lag, lead
  4. first, last
from pyspark.sql.window import Window

from pyspark.sql.functions import (
    row_number, rank, dense_rank, percent_rank, ntile,

    sum, avg, min, max, count,

    lag, lead,
    first, last,

    cume_dist
)

define window

# sample data - customer_id, order_date, amount, city


# running total frame
win = (Window
    .partitionBy("customer_id")
    .orderBy("order_date")
    .rowsBetween(Window.unboundedPreceding, Window.currentRow) # default window size
)

Frame bounds:

# sliding window - last 3 rows including current
w_slide = (Window
    .partitionBy("customer_id")
    .orderBy("order_date")
    .rowsBetween(-2, Window.currentRow)
)

# range frame - rows within 7 days of current row
# irrespective of how many rows the window will have
w_range = (Window
    .partitionBy("customer_id")
    .orderBy("order_date_long")
    .rangeBetween(-7 * 86400, Window.currentRow # epoch seconds
)

Ranking

python

df = df.withColumn("row_number",    row_number().over(w))   # no ties, always unique

df = df.withColumn("rank",       rank().over(w))         # ties get same rank, next rank skips
df = df.withColumn("dense_rank", dense_rank().over(w))   # ties get same rank, no skip

df = df.withColumn("ntile",    ntile(4).over(w))       # quartile bucket 1..4
df = df.withColumn("percent_rank",   percent_rank().over(w)) # 0.0 to 1.0
df = df.withColumn("cume_dist",  cume_dist().over(w))    # fraction of rows <= current

Lag / lead

lag(col, offset, default)

df = df.withColumn("lag_2",lag("amount",  2, 0).over(w)) # 2 rows back, default 0

lead(col, offset, default)
df = df.withColumn("lead_2",lead("amount",  2, 0).over(w)) # 2 rows ahead, default 0

# example - diff between current and previous rows
df = df.withColumn("diff",         col("amount") - lag("amount", 1).over(w)) # change

Aggregations

python

# running - unboundedPreceding & currentRow
df = df.withColumn("sum",  sum("amount").over(win))
df = df.withColumn("avg",    avg("amount").over(win))
df = df.withColumn("count",  count("amount").over(win))

First / last

python

# first/last need ignorNulls and explicit frame for reliable results
win_full = (
    Window
    .partitionBy("customer_id")
    .orderBy("order_date")
   .rowsBetween(Window.unboundedPreceding, Window.unboundedFollowing)
)

df = df.withColumn("first", first("amount", ignoreNulls=True).over(win_full))
df = df.withColumn("last",  last("amount",  ignoreNulls=True).over(win_full))

Common patterns

latest record per customer (dedup)

from pyspark.sql.functions import col

deduped = (
    df
    .withColumn("rn", row_number().over(win))
    .filter("rn = 1")
    .drop("rn")
)

top N per group

# top N per group
top2 = (
    df
    .withColumn("rn", row_number().over(win))
    .filter(col("rn") <= 2)
    .drop("rn")
)

Percent of total per partition

# percent of total per partition
df = (df
    .withColumn("pct", col("amount") / sum("amount").over(win_full) * 100
    )
)

Other stuff

Date and time handling

All fns are part of pyspark.sql.functions()

Current time

current_date()
current_timestamp()

Convert

# str to object
to_date("col_str", "dd-MM-yyyy")                    # "18-08-2026" -> 2026-08-18
to_timestamp("col_str", "yyyy-MM-dd HH:mm:ss")       # "2026-08-18 14:30:00" -> timestamp

# object to str
date_format("col_date", "yyyy-MM")                      # 2026-08-18 -> "2026-08"

Format specifier strings are usual python placeholders:

yyyy   → year
MM     → month
dd     → day
HH     → hour (24h)
mm     → minute
ss     → second
SSS    → milliseconds
a      → AM/PM
E      → day name
MMM    → month name

Arithmetic

# add/sub
date_add("col", 7)             # add days (can be -ve)
add_months("col", 3)           # add months (can be -ve)

# diff
datediff("col_end_date", "col_start_date")  # diff in days
months_between("col_end", "col_start")

Extract

from pyspark.sql.functions import extract, col, lit

extract(
    field=lit('YEAR'), 
    source=date_col
) 
# fields - "YEAR", "MONTH", "DAY", "HOUR", "MINUTE", "SECOND"

quarter()
dayofmonth()
dayofweek()

Boundaries / truncation

last_day("date")              # date of last day of month - 2026-08-18 -> 2026-08-31

# truncate timestamp
date_trunc("month", "ts") # 2026-08-18 14:30:45 -> 2026-08-01 00:00:00
# takes - "year", "quarter", "month", "day", "week", "hour", "minute", "second", "microsecond", "millisecond"

# truncate date
trunc("date", "month") #2026-08-18 -> 2026-08-01
# takes - "year", "quarter", "month", "week"

time_trunc (unit:str, col:col)
# takes - "hour", "minute", "second", "millisecond", "microsecond"

I/O

Delta lake

read

# managed unity catalog table
df = spark.read.table("catalog.schema.orders")

# external tabale
df = spark.read.format("delta").load("/mnt/delta/orders")

# specific version
df = spark.read
    .option("versionAsOf", 5)
    .table("catalog.schema.orders")
df = spark.read
    .option("timestampAsOf", "2024-01-01")
    .table("catalog.schema.orders")

write

# append
(
df
    .write
    .format("delta")
    .mode("append")
    .option("mergeSchema", "true") # schema evolution
    .saveAsTable("catalog.schema.orders")
)

options:

create table

# DDL
spark.sql("""
  CREATE TABLE IF NOT EXISTS catalog.schema.orders (
    order_id  BIGINT,
    customer  STRING,
    amt       DECIMAL(10,2),
    dt        DATE
  )
  USING DELTA
  CLUSTER BY (customer_id)
  TBLPROPERTIES (
    'delta.enableDeletionVectors' = 'true',
    'delta.enableChangeDataFeed'  = 'true'
  )
""")

maintenance

-- compaction
OPTIMIZE catalog.schema.orders

OPTIMIZE catalog.schema.orders ZORDER BY (customer_id)

-- 7 days history cleanup
VACUUM  catalog.schema.orders RETAIN 168 HOURS

-- CBO stats
ANALYZE TABLE catalog.schema.orders COMPUTE STATISTICS FOR ALL COLUMNS

Inspection

DESCRIBE HISTORY catalog.schema.orders
DESCRIBE DETAIL  catalog.schema.orders

kafka

Options to control batch frequency, read sizes and positions

If spark processing capacity is less than kafka arrivals, there will be backlog

read

raw = (spark.readStream
       .format("kafka")
       .option("kafka.bootstrap.servers", "broker1:9092,broker2:9092")
       .option("subscribe", "orders")
       .option("startingOffsets", "latest")        # latest / earliest / json offset
       .option("maxOffsetsPerTrigger", 100_000)    # pace ingestion
       .option("failOnDataLoss", "false")          # topic deleted / offsets expired
       .load())

batch read (backfill/one-time)

df = (spark.read                               # read not readStream
      .format("kafka")
      .option("kafka.bootstrap.servers", "broker1:9092,broker2:9092")
      .option("subscribe", "orders")
      .option("startingOffsets", """{"orders":{"0":1000,"1":1000}}""")  # exact offsets
      .option("endingOffsets",   """{"orders":{"0":2000,"1":2000}}""")
      .load())

columns you get:

Column Type Notes
key binary cast to STRING if text key
value binary your payload — always cast/parse
topic string useful when subscribing to multiple
partition int
offset long
timestamp timestamp producer or broker timestamp
timestampType int 0=CreateTime, 1=LogAppendTime

parse value

from pyspark.sql.functions import col, from_json, from_avro
from pyspark.sql.types import StructType, StructField, StringType, IntegerType, TimestampType

schema = StructType([
    StructField("order_id", StringType()),
    StructField("customer", StringType()),
    StructField("amt",      IntegerType()),
    StructField("ts",       TimestampType())
])

# JSON value
parsed = (raw
          .select(
              from_json(col("value").cast("string"), schema).alias("d"),
              col("timestamp").alias("kafka_ts")
          )
          .select("d.*", "kafka_ts"))

# Avro value (with schema registry)
parsed = (raw
          .select(from_avro(col("value"), "<avro_schema_json_string>").alias("d"))
          .select("d.*"))

write

streaming write

(parsed
 .selectExpr(
     "CAST(order_id AS STRING) AS key",   # key must be STRING or BINARY
     "to_json(struct(*)) AS value"         # value must be STRING or BINARY
 )
 .writeStream
 .format("kafka")
 .option("kafka.bootstrap.servers", "broker1:9092,broker2:9092")
 .option("topic", "orders_enriched")
 .option("checkpointLocation", "/mnt/checkpoints/orders_enriched")
 .outputMode("append")
 .trigger(processingTime="30 seconds")
 .start())

batch write

(df.selectExpr("CAST(id AS STRING) AS key", "to_json(struct(*)) AS value")
   .write
   .format("kafka")
   .option("kafka.bootstrap.servers", "broker1:9092,broker2:9092")
   .option("topic", "orders_out")
   .save())

csv

Read all files in subdirs too? use option("recursiveFileLookup", "true") (disables partition inference)

read

df = (spark.read
      .format("csv")
      .option("header", "true")
      .option("inferSchema", "true")       # avoid in prod — triggers extra scan
      .load("/mnt/data/orders.csv"))

# explicit schema (prod way)
schema = "order_id INT, customer STRING, amt DECIMAL(10,2), dt DATE"

df = (spark.read
      .format("csv")
      .schema(schema)
      .option("header", "true")
      .option("dateFormat", "yyyy-MM-dd")
      .option("timestampFormat", "yyyy-MM-dd HH:mm:ss")
      .option("nullValue", "NULL")         # treat this string as null
      .option("emptyValue", "")
      .option("mode", "PERMISSIVE")        # PERMISSIVE / DROPMALFORMED / FAILFAST
      .option("columnNameOfCorruptRecord", "_corrupt_record")
      .load("/mnt/data/orders.csv"))

# read folder — all CSVs in directory
df = spark.read.schema(schema).option("header","true").csv("/mnt/data/orders/")

# read multiple explicit paths
df = spark.read.schema(schema).option("header","true").csv(
    "/mnt/data/orders_jan.csv",
    "/mnt/data/orders_feb.csv"
)

options

Option Default Notes
header|false use first row as column names
sep / delimiter|,|any single char; \t for TSV
quote|"|quoting character
escape|\|escape character inside quotes
multiLine|false fields with newlines inside quotes
encoding|UTF-8 ISO-8859-1 for legacy files
ignoreLeadingWhiteSpace|false
ignoreTrailingWhiteSpace|false
nanValue|NaN|string to treat as NaN
positiveInf / negativeInf|Inf / -Inf|
comment|disabled skip lines starting with this char

write

# basic
(df.write
   .format("csv")
   .mode("overwrite")               # overwrite / append / error / ignore
   .option("header", "true")
   .save("/mnt/data/output/orders"))

# options
(df.write
   .format("csv")
   .mode("overwrite")
   .option("header", "true")
   .option("sep", ",")
   .option("quote", '"')
   .option("escape", "\\")
   .option("nullValue", "")
   .option("dateFormat", "yyyy-MM-dd")
   .option("timestampFormat", "yyyy-MM-dd HH:mm:ss")
   .option("compression", "none")   # none / gzip / bz2 / deflate / snappy
   .save("/mnt/data/output/orders"))

# control number of output files
df.coalesce(1).write.csv(...)       # single file — small data only
df.repartition(10).write.csv(...)   # 10 files

json

{"id": 1, "name": "alice", "address": {"city": "pune", "pin": "411001"}, "orders": [{"oid": 101, "amt": 500}, {"oid": 102, "amt": 300}]}

read

from pyspark.sql.types import *
from pyspark.sql.functions import col, explode, explode_outer

schema = StructType([
    StructField("id",      IntegerType()),
    StructField("name",    StringType()),
    StructField("address", StructType([
        StructField("city", StringType()),
        StructField("pin",  StringType())
    ])),
    StructField("orders",  ArrayType(StructType([
        StructField("oid", IntegerType()),
        StructField("amt", IntegerType())
    ])))
])

df = spark.read.schema(schema).json("/mnt/data/users.json")
df.show()

flatten and explode in 1 step

and also flatten exploded json

final = (df
    .select(
        "id",
        "name",
        col("address.city").alias("city"),
        col("address.pin").alias("pin"),
        explode("orders").alias("order")
    )
    .select(
        "id", "name", "city", "pin",
        col("order.oid").alias("order_id"),
        col("order.amt").alias("amount")
    ))

Gotchas

schema = "id INT, name STRING, address STRUCT<city:STRING, pin:STRING>, orders ARRAY<STRUCT<oid:INT, amt:INT>>"

sqlserver

jdbc_url = "jdbc:sqlserver://server.database.windows.net:1433;databaseName=mydb"

connection_props = {
    "user":     dbutils.secrets.get("scope", "sql-user"),
    "password": dbutils.secrets.get("scope", "sql-password"),
    "driver":   "com.microsoft.sqlserver.jdbc.SQLServerDriver"
}

azure sql with aad token (no password)

import struct, pyodbc
from azure.identity import ManagedIdentityCredential

cred  = ManagedIdentityCredential()
token = cred.get_token("https://database.windows.net/.default").token

# pack token for pyodbc
token_bytes  = token.encode("utf-16-le")
token_struct = struct.pack(f"<I{len(token_bytes)}s", len(token_bytes), token_bytes)

conn = pyodbc.connect(
    "DRIVER={ODBC Driver 17 for SQL Server};"
    "SERVER=server.database.windows.net;"
    "DATABASE=mydb",
    attrs_before={1256: token_struct}   # SQL_COPT_SS_ACCESS_TOKEN = 1256
)

read - full table

df = (spark.read
      .jdbc(url=jdbc_url, table="dbo.orders", properties=connection_props))

read - query

df = (spark.read
      .jdbc(
          url        = jdbc_url,
          table      = "(SELECT * FROM dbo.orders WHERE dt >= '2024-01-01') AS t",
          properties = connection_props
      ))

read - parallel

manual

predicates = [
    "order_id BETWEEN 1       AND 250000",
    "order_id BETWEEN 250001  AND 500000",
    "order_id BETWEEN 500001  AND 750000",
    "order_id BETWEEN 750001  AND 1000000",
]

df = (spark.read
      .jdbc(
      url=jdbc_url,
      table="dbo.orders",
      predicates=predicates,
      properties=connection_props
      )
)

automatic - only on partitionable cols like int, float, date, timestamp

df = (spark.read
      .jdbc(
          url               = jdbc_url,
          table             = "(SELECT * FROM dbo.orders) AS t",
          column            = "order_id",        # numeric / date column
          lowerBound        = 1,
          upperBound        = 10_000_000,
          numPartitions     = 16,
          properties        = connection_props
      ))

write

append

(df.write
   .jdbc(url=jdbc_url, table="dbo.orders_out",
         mode="append", properties=connection_props))

overwrite

(df.write
   .option("truncate", "true")        # truncate instead of DROP + recreate
   .jdbc(url=jdbc_url, table="dbo.orders_out",
         mode="overwrite", properties=connection_props))

write - batch size + isolation

(df.write
   .option("batchsize",        10_000)
   .option("isolationLevel",   "READ_COMMITTED")
   .option("numPartitions",    8)              # parallel writers
   .jdbc(url=jdbc_url, table="dbo.orders_out",
         mode="append", properties=connection_props))

Gotchas

scd

SCD

scd 1

Overwrite existing data

from delta.tables import DeltaTable

target = DeltaTable.forName(spark, "dim_customer")

(
    target.alias("t")
    .merge(
        source_df.alias("s"),
        "t.customer_id = s.customer_id"
    )
    .whenMatchedUpdateAll()
    .whenNotMatchedInsertAll()
    .execute()
)

scd 2

single merge

Create a staged source containing 2 tows - one to expire existing row, one to insert new version

from delta.tables import DeltaTable
from pyspark.sql import functions as F

target = DeltaTable.forName(spark, "dim_customer")

# Create source rows:
# 1. A "matched" row to expire the existing version
# 2. An "unmatched" row to insert the new version
staged = (
    source_df.alias("s")
    .join(
        spark.table("dim_customer").alias("t"),
        (F.col("s.id") == F.col("t.id")) &
        (F.col("t.is_current") == True),
        "left"
    )
    .withColumn(
        "changed",
        F.col("t.id").isNotNull() &
        (
            (F.col("s.name") != F.col("t.name")) |
            (F.col("s.city") != F.col("t.city"))
        )
    )
    .select(
        F.col("s.id"),
        F.col("s.name"),
        F.col("s.city"),
        F.col("changed")
    )
    .withColumn(
        "merge_id",
        F.when(F.col("changed"), F.col("id"))
         .otherwise(F.lit(None))
    )
)

# Duplicate changed rows:
#   merge_id = id      -> matches existing row and expires it
#   merge_id = NULL    -> doesn't match, so inserts new version
expire_rows = staged.filter("changed").withColumn(
    "merge_id", F.col("id")
)

insert_rows = staged.withColumn(
    "merge_id",
    F.lit(None).cast("long")
)

staged = expire_rows.unionByName(insert_rows)

(
    target.alias("t")
    .merge(
        staged.alias("s"),
        "t.id = s.merge_id AND t.is_current = true"
    )
    .whenMatchedUpdate(
        set={
            "is_current": "false",
            "valid_to": "current_timestamp()"
        }
    )
    .whenNotMatchedInsert(
        values={
            "id": "s.id",
            "name": "s.name",
            "city": "s.city",
            "valid_from": "current_timestamp()",
            "valid_to": "NULL",
            "is_current": "true"
        }
    )
    .execute()
)

2 merges

from delta.tables import DeltaTable
from pyspark.sql import functions as F

target = DeltaTable.forName(spark, "dim_customer")

# 1. Find existing records whose attributes changed
changed = (
    source_df.alias("s")
    .join(
        spark.table("dim_customer").alias("t"),
        (F.col("s.id") == F.col("t.id")) &
        (F.col("t.is_current") == True),
        "inner"
    )
    .where(
        (F.col("s.name") != F.col("t.name")) |
        (F.col("s.city") != F.col("t.city"))
    )
    .select("s.id")
)

# 2. Expire old versions
(
    target.alias("t")
    .merge(
        changed.alias("s"),
        "t.id = s.id AND t.is_current = true"
    )
    .whenMatchedUpdate(
        set={
            "is_current": "false",
            "valid_to": "current_timestamp()"
        }
    )
    .execute()
)

# 3. Insert new versions + genuinely new records
new_rows = (
    source_df
    .withColumn("is_current", F.lit(True))
    .withColumn("valid_from", F.current_timestamp())
    .withColumn("valid_to", F.lit(None).cast("timestamp"))
)

(
    target.alias("t")
    .merge(
        new_rows.alias("s"),
        "t.id = s.id AND t.is_current = true"
    )
    .whenNotMatchedInsertAll()
    .execute()
)

scd 3

cols - id, name, city

from delta.tables import DeltaTable
from pyspark.sql import functions as F

target = DeltaTable.forName(spark, "dim_customer")

(
    target.alias("t")
    .merge(
        source_df.alias("s"),
        "t.id = s.id"
    )
    .whenMatchedUpdate(
        set={
            "name": "s.name",
            "previous_name": "t.name",
            "city": "s.city",
            "previous_city": "t.city"
        }
    )
    .whenNotMatchedInsert(
        values={
            "id": "s.id",
            "name": "s.name",
            "previous_name": "NULL",
            "city": "s.city",
            "previous_city": "NULL"
        }
    )
    .execute()
)

scd 4

from delta.tables import DeltaTable
from pyspark.sql import functions as F

current = DeltaTable.forName(spark, "dim_customer")
history = DeltaTable.forName(spark, "dim_customer_history")

# Find changed existing records
changed = (
    source_df.alias("s")
    .join(
        spark.table("dim_customer").alias("t"),
        F.col("s.id") == F.col("t.id"),
        "inner"
    )
    .where(
        (F.col("s.name") != F.col("t.name")) |
        (F.col("s.city") != F.col("t.city"))
    )
    .select(
        "t.id",
        "t.name",
        "t.city"
    )
)

# 1. Copy old versions to history
(
    changed
    .withColumn("valid_from", F.current_timestamp())
    .withColumn("valid_to", F.current_timestamp())
    .write
    .format("delta")
    .mode("append")
    .saveAsTable("dim_customer_history")
)

# 2. Update current table
(
    current.alias("t")
    .merge(
        source_df.alias("s"),
        "t.id = s.id"
    )
    .whenMatchedUpdate(
        set={
            "name": "s.name",
            "city": "s.city"
        }
    )
    .whenNotMatchedInsert(
        values={
            "id": "s.id",
            "name": "s.name",
            "city": "s.city"
        }
    )
    .execute()
)

scd 5

--

Workflows

Handle nulls

# Coalesce
df.withColumn("col_processed", F.coalesce(col("col_processed"), col("other_col")))

# Fill nulls
fill_na_values = {"col1": val1, "col2": val2}
df.fillna(fill_na_values) # set col's values through dict

# Drop nulls
df.na.drop(subset=["col1", "col2"]) # specifying subset is optional

# Using case
df.withColumn("some_column",
    when(
        col("some_column")
        .isNull(), 0)
        .otherwise(col("some_column")
    )
)

Deduplicate

# all cols
fully_unique_rows_df = df.distinct()
# specific cols
unique_names_df = df.dropDuplicates(subset=["name", "surname"])

Functions

group_concat

In PySpark, the equivalent of SQL GROUP_CONCAT is typically collect_list() + concat_ws().

concat_ws means concat with separator

from pyspark.sql import functions as F

result = df.groupBy("customerId").agg(
    F.concat_ws(", ", F.collect_list("productName")).alias("products")
)

Note: for unique values, use collectset() instead of collectlist()

Input:

customerId productName
1 Laptop
1 Mouse
1 Keyboard
2 Monitor
2 Webcam

Output:

customerId products
1 Laptop, Mouse, Keyboard
2 Monitor, Webcam
Colophon 4798 words · 24m read
Written as a markdown note in Obsidian. Built into this page by a Python script on 2026-10-02.

Pages