WiliamRosa
Databricks Partner

See if this helps:


Import org.apache.spark.sql.DataFrame

import org.apache.spark.sql.catalyst.plans.logical._
import org.apache.spark.sql.catalyst.expressions._

def extractFiltersAndJoins(df: DataFrame): (Set[String], Set[(String,String)]) = {
val plan = df.queryExecution.optimizedPlan
var filterCols = Set.empty[String]
var joinCols = Set.empty[(String,String)]

plan.foreach {
case Filter(condition, _) =>
condition.references.foreach(ref => filterCols += ref.sql)
case Join(left, right, _, condOpt, _) =>
condOpt.foreach { cond =>
cond.references.toSeq.combinations(2).foreach {
case Seq(a,b) => joinCols += (a.sql -> b.sql)
case _ =>
}
}
case _ =>
}
(filterCols, joinCols)
}

// use it
val df = spark.sql("""
SELECT c.id, SUM(s.val) v
FROM sales s JOIN customers c ON s.cid = c.id
WHERE s.dt >= DATE '2025-08-01'
GROUP BY c.id
""")
val (filters, joins) = extractFiltersAndJoins(df)
println("Filter columns: " + filters)
println("Join columns: " + joins)

Wiliam Rosa
Data Engineer | Machine Learning Engineer
LinkedIn: linkedin.com/in/wiliamrosa