Skip to content

Commit ed88558

Browse files
kyramanKyra Mangasarian
andauthored
fix: backtick-escape column names in pruneColumns to prevent V2 DataSource function-name collision (#748)
Co-authored-by: Kyra Mangasarian <kyraman@amazon.com>
1 parent abf22d6 commit ed88558

2 files changed

Lines changed: 29 additions & 1 deletion

File tree

src/main/scala/com/amazon/deequ/analyzers/runners/AnalysisRunner.scala

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@ import com.amazon.deequ.analyzers._
2020
import com.amazon.deequ.io.DfsUtils
2121
import com.amazon.deequ.metrics.{DoubleMetric, Metric}
2222
import com.amazon.deequ.repository.{MetricsRepository, ResultKey}
23+
import com.amazon.deequ.utilities.ColumnUtil.escapeColumn
2324
import org.apache.spark.sql.Column
2425
import org.apache.spark.sql.functions.col
2526
import org.apache.spark.sql.types.StructType
@@ -409,7 +410,7 @@ object AnalysisRunner {
409410
// All analyzers are dataset-level (e.g. Size), no column selection needed
410411
data
411412
} else {
412-
data.select(neededColumns.map(col): _*)
413+
data.select(neededColumns.map(c => col(escapeColumn(c))): _*)
413414
}
414415
}
415416
}

src/test/scala/com/amazon/deequ/analyzers/runners/AnalysisRunnerTests.scala

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -532,5 +532,32 @@ class AnalysisRunnerTests extends AnyWordSpec
532532
assert(result.metric(mean).get.value === mean.calculate(data).value)
533533
assert(result.metric(minimum).get.value === minimum.calculate(data).value)
534534
}
535+
536+
"handle column pruning when column name collides with Spark function name" in
537+
withSparkSession { session =>
538+
import session.implicits._
539+
val data = Seq(
540+
("setosa", 5),
541+
("setosa", 4),
542+
("versicolor", 6),
543+
("versicolor", 7),
544+
("virginica", 6)
545+
).toDF("flower_type", "length")
546+
547+
val completeness = Completeness("length")
548+
val mean = Mean("length")
549+
val maximum = Maximum("length")
550+
551+
val analysis = Analysis()
552+
.addAnalyzer(completeness)
553+
.addAnalyzer(mean)
554+
.addAnalyzer(maximum)
555+
556+
val result = AnalysisRunner.run(data, analysis)
557+
558+
assert(result.metric(completeness).get.value.isSuccess)
559+
assert(result.metric(mean).get.value.isSuccess)
560+
assert(result.metric(maximum).get.value.get === 7.0)
561+
}
535562
}
536563
}

0 commit comments

Comments
 (0)