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")
- inner
- left
- right
- full
- left_semi
- left_anti
- 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"
)
- broadcast
- merge
- shuffle_hash
- shufflereplicatenl
window functions
- rank
- rownumber
- rank, denserank
- percentrank, ntile
- cumedist
- aggregation
- sum
- avg
- min
- max
- count
- lag, lead
- 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
)
- rowsBetween - physical offset (by row)
- rangeBetween - value offset (needs numeric/date orderBy). Use to make frames based on values (eg - rows of last 7 days, irrespective of how many there are)
Frame bounds:
- unboundedPreceding
- unboundedFollowing
- currentRow
# 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
- No
partitionBy-> all data in one partition -> OOM on large tables. - Default frame with
orderBy=RANGE BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW— includes ties. - Default frame without
orderBy= whole partition. lastwithout explicitunboundedFollowingframe returns current row, not partition last.rank/dense_rank/row_numberignore frame — always rank within full partition order.- Window functions not allowed in WHERE/HAVING — wrap in subquery or CTE.
- Multiple windows on same partition+order -> Spark reuses one shuffle. Define them consistently.
percent_rankof first row = 0.0 always;cume_distof last row = 1.0 always.
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:
- write mode
- .mode("overwrite")
- overwrite specific partitions only
"replaceWhere", "dt >= '2024-01-01' AND dt < '2024-02-01'"
- schema management
"mergeSchema", "true""overwriteSchema", "true"
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
- Batch frequency -
trigger(...)processingTime="30 seconds"availableNow=Truecontinuous="1 second"- not exactly micro batch. This is v low latency mode- default - unspecified. It'll process as quickly as it can
- Amount admitted per batch
- Limit on number of messages
maxOffsetsPerTrigger- offsets of the topic per triggerminOffsetsPerTrigger- accumulate at least these many offsets before processing
- Limit on total size of data processed
maxBytesPerTrigger- use for messages with highly variable sizes"100MB"
- Limit on number of messages
- Starting positions
startingOffsets- only applies when query starts without existing checkpoint"earliest","latest",'{"events":{"0":100,"1":250}}'
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
explodeon a null/empty array drops the row silently — useexplode_outerto keep it.- After explode, the array column is gone — select what you need before exploding or re-join.
- DDL schema shorthand for the same schema:
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
- JDBC reads are single-partition by default — always use
numPartitionsfor large tables. lowerBound/upperBoundare splitting hints not filters — rows outside range still read.overwritewithouttruncate=truedrops and recreates table — kills indexes.batchsizedefault is 1000 — too low for large writes, set 10k–50k.- Driver jar must be on cluster — in Databricks add
com.microsoft.sqlserver:mssql-jdbc:12.4.2.jre11as Maven library. - Stored procs and DDL must go through a direct JDBC/pyodbc connection on the driver, not Spark.
- Never put credentials in code — always
dbutils.secrets.get. - Parallel writers (
numPartitionson write) can cause deadlocks on SQL Server — tune based on target table's lock behaviour.
scd
SCD
- scd 1 - corrections/current-state data
- overwrite existing
- scd 2 - audit history/reporting
- insert new row for change, mark older as expired
- scd 3 - only want immediate previous value
- store current and previous value in separate columns
- scd 4 - separate operational(current) and historical workloads
- update current table and push update to historical table (append only history)
- eg - track cart - to know what's currently in cart, but also know what was added and removed from the cart
- scd 6 - complex reporting requirements
- combination of 1, 2, 3
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
- cols - id, name, city
- added cols - surrogatekey, validfrom, validto, iscurrent
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
- added cols - previousname, previouscity
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
- cols - id, name, city
- added cols - validfrom, validto
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 |