From a3709b91062c70ab0efe8cdb56bfd29460c790e4 Mon Sep 17 00:00:00 2001 From: Takeshi Yamamuro Date: Fri, 22 Feb 2019 11:35:23 +0900 Subject: [PATCH 1/7] Fix --- .../expressions/collectionOperations.scala | 8 +- .../expressions/complexTypeExtractors.scala | 79 +++++++++++-------- .../CollectionExpressionsSuite.scala | 52 ++++++++++++ 3 files changed, 103 insertions(+), 36 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala index 67f6739b1e18f..61e7daa2bf62e 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala @@ -1929,7 +1929,8 @@ case class ArrayPosition(left: Expression, right: Expression) b """, since = "2.4.0") -case class ElementAt(left: Expression, right: Expression) extends GetMapValueUtil { +case class ElementAt(left: Expression, right: Expression) + extends GetMapValueUtil with GetArrayItemUtil { @transient private lazy val mapKeyType = left.dataType.asInstanceOf[MapType].keyType @@ -1974,7 +1975,10 @@ case class ElementAt(left: Expression, right: Expression) extends GetMapValueUti } } - override def nullable: Boolean = true + override def nullable: Boolean = left.dataType match { + case _: ArrayType => computeNullabilityFromArray + case _: MapType => computeNullabilityFromMap + } override def nullSafeEval(value: Any, ordinal: Any): Any = doElementAt(value, ordinal) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala index 55ed617e2904b..93e4690ae4c5a 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala @@ -216,24 +216,15 @@ case class GetArrayStructFields( } /** - * Returns the field at `ordinal` in the Array `child`. - * - * We need to do type checking here as `ordinal` expression maybe unresolved. + * Common trait for [[GetArrayItem]] and [[ElementAt]]. */ -case class GetArrayItem(child: Expression, ordinal: Expression) - extends BinaryExpression with ExpectsInputTypes with ExtractValue with NullIntolerant { - - // We have done type checking for child in `ExtractValue`, so only need to check the `ordinal`. - override def inputTypes: Seq[AbstractDataType] = Seq(AnyDataType, IntegralType) - - override def toString: String = s"$child[$ordinal]" - override def sql: String = s"${child.sql}[${ordinal.sql}]" +trait GetArrayItemUtil extends BinaryExpression { - override def left: Expression = child - override def right: Expression = ordinal + private val child = left + private val ordinal = right /** `Null` is returned for invalid ordinals. */ - override def nullable: Boolean = if (ordinal.foldable && !ordinal.nullable) { + protected def computeNullabilityFromArray: Boolean = if (ordinal.foldable && !ordinal.nullable) { val intOrdinal = ordinal.eval().asInstanceOf[Number].intValue() child match { case CreateArray(ar) if intOrdinal < ar.length => @@ -247,7 +238,25 @@ case class GetArrayItem(child: Expression, ordinal: Expression) } else { true } +} +/** + * Returns the field at `ordinal` in the Array `child`. + * + * We need to do type checking here as `ordinal` expression maybe unresolved. + */ +case class GetArrayItem(child: Expression, ordinal: Expression) + extends GetArrayItemUtil with ExpectsInputTypes with ExtractValue with NullIntolerant { + + // We have done type checking for child in `ExtractValue`, so only need to check the `ordinal`. + override def inputTypes: Seq[AbstractDataType] = Seq(AnyDataType, IntegralType) + + override def toString: String = s"$child[$ordinal]" + override def sql: String = s"${child.sql}[${ordinal.sql}]" + + override def left: Expression = child + override def right: Expression = ordinal + override def nullable: Boolean = computeNullabilityFromArray override def dataType: DataType = child.dataType.asInstanceOf[ArrayType].elementType protected override def nullSafeEval(value: Any, ordinal: Any): Any = { @@ -281,10 +290,29 @@ case class GetArrayItem(child: Expression, ordinal: Expression) } /** - * Common base class for [[GetMapValue]] and [[ElementAt]]. + * Common trait for [[GetMapValue]] and [[ElementAt]]. */ +trait GetMapValueUtil extends BinaryExpression with ImplicitCastInputTypes { + + private val child = left + private val key = right + + /** `Null` is returned for invalid ordinals. */ + protected def computeNullabilityFromMap: Boolean = if (key.foldable && !key.nullable) { + val keyObj = key.eval() + child match { + case m: CreateMap if m.resolved => + m.keys.zip(m.values).filter { case (k, _) => k.foldable && !k.nullable }.find { + case (k, _) if k.eval() == keyObj => true + case _ => false + }.map(_._2.nullable).getOrElse(true) + case _ => + true + } + } else { + true + } -abstract class GetMapValueUtil extends BinaryExpression with ImplicitCastInputTypes { // todo: current search is O(n), improve it. def getValueEval(value: Any, ordinal: Any, keyType: DataType, ordering: Ordering[Any]): Any = { val map = value.asInstanceOf[MapData] @@ -379,24 +407,7 @@ case class GetMapValue(child: Expression, key: Expression) override def left: Expression = child override def right: Expression = key - - /** `Null` is returned for invalid ordinals. */ - override def nullable: Boolean = if (key.foldable && !key.nullable) { - val keyObj = key.eval() - child match { - case m: CreateMap if m.resolved => - m.keys.zip(m.values).filter { case (k, _) => k.foldable && !k.nullable }.find { - case (k, _) if k.eval() == keyObj => true - case _ => false - }.map(_._2.nullable).getOrElse(true) - case _ => - true - } - } else { - true - } - - + override def nullable: Boolean = computeNullabilityFromMap override def dataType: DataType = child.dataType.asInstanceOf[MapType].valueType // todo: current search is O(n), improve it. diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/CollectionExpressionsSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/CollectionExpressionsSuite.scala index bed8547dbc83d..63f08eba00407 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/CollectionExpressionsSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/CollectionExpressionsSuite.scala @@ -1092,6 +1092,58 @@ class CollectionExpressionsSuite extends SparkFunSuite with ExpressionEvalHelper checkEvaluation(ElementAt(mb0, Literal(Array[Byte](3, 4))), null) } + test("SPARK-26965 correctly handles ElementAt nullability for arrays") { + // CreateArray case + val a = AttributeReference("a", IntegerType, nullable = false)() + val b = AttributeReference("b", IntegerType, nullable = true)() + val array = CreateArray(a :: b :: Nil) + assert(!ElementAt(array, Literal(0)).nullable) + assert(ElementAt(array, Literal(1)).nullable) + assert(!ElementAt(array, Subtract(Literal(2), Literal(2))).nullable) + assert(ElementAt(array, AttributeReference("ordinal", IntegerType)()).nullable) + + // GetArrayStructFields case + val f1 = StructField("a", IntegerType, nullable = false) + val f2 = StructField("b", IntegerType, nullable = true) + val structType = StructType(f1 :: f2 :: Nil) + val c = AttributeReference("c", structType, nullable = false)() + val inputArray1 = CreateArray(c :: Nil) + val inputArray1ContainsNull = c.nullable + val stArray1 = GetArrayStructFields(inputArray1, f1, 0, 2, inputArray1ContainsNull) + assert(!ElementAt(stArray1, Literal(0)).nullable) + val stArray2 = GetArrayStructFields(inputArray1, f2, 1, 2, inputArray1ContainsNull) + assert(ElementAt(stArray2, Literal(0)).nullable) + + val d = AttributeReference("d", structType, nullable = true)() + val inputArray2 = CreateArray(c :: d :: Nil) + val inputArray2ContainsNull = c.nullable || d.nullable + val stArray3 = GetArrayStructFields(inputArray2, f1, 0, 2, inputArray2ContainsNull) + assert(!ElementAt(stArray3, Literal(0)).nullable) + assert(ElementAt(stArray3, Literal(1)).nullable) + val stArray4 = GetArrayStructFields(inputArray2, f2, 1, 2, inputArray2ContainsNull) + assert(ElementAt(stArray4, Literal(0)).nullable) + assert(ElementAt(stArray4, Literal(1)).nullable) + } + + test("SPARK-26965 correctly handles ElementAt nullability for maps") { + // String key test + val k1 = Literal("k1") + val v1 = AttributeReference("v1", StringType, nullable = true)() + val k2 = Literal("k2") + val v2 = AttributeReference("v2", StringType, nullable = false)() + val map1 = CreateMap(k1 :: v1 :: k2 :: v2 :: Nil) + assert(ElementAt(map1, Literal("k1")).nullable) + assert(!ElementAt(map1, Literal("k2")).nullable) + assert(ElementAt(map1, Literal("non-existent-key")).nullable) + + // Complex type key test + val k3 = Literal.create((1, "a")) + val k4 = Literal.create((2, "b")) + val map2 = CreateMap(k3 :: v1 :: k4 :: v2 :: Nil) + assert(ElementAt(map2, Literal.create((1, "a"))).nullable) + assert(!ElementAt(map2, Literal.create((2, "b"))).nullable) + } + test("Concat") { // Primitive-type elements val ai0 = Literal.create(Seq(1, 2, 3), ArrayType(IntegerType, containsNull = false)) From c2d522a5f7029715ad5f16729c5c0668c1ad3bd7 Mon Sep 17 00:00:00 2001 From: Takeshi Yamamuro Date: Mon, 25 Feb 2019 10:17:55 +0900 Subject: [PATCH 2/7] Fix --- .../expressions/complexTypeExtractors.scala | 50 +++++++++---------- 1 file changed, 25 insertions(+), 25 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala index 93e4690ae4c5a..50e00bcdf6245 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala @@ -215,31 +215,6 @@ case class GetArrayStructFields( } } -/** - * Common trait for [[GetArrayItem]] and [[ElementAt]]. - */ -trait GetArrayItemUtil extends BinaryExpression { - - private val child = left - private val ordinal = right - - /** `Null` is returned for invalid ordinals. */ - protected def computeNullabilityFromArray: Boolean = if (ordinal.foldable && !ordinal.nullable) { - val intOrdinal = ordinal.eval().asInstanceOf[Number].intValue() - child match { - case CreateArray(ar) if intOrdinal < ar.length => - ar(intOrdinal).nullable - case GetArrayStructFields(CreateArray(elements), field, _, _, _) - if intOrdinal < elements.length => - elements(intOrdinal).nullable || field.nullable - case _ => - true - } - } else { - true - } -} - /** * Returns the field at `ordinal` in the Array `child`. * @@ -289,6 +264,31 @@ case class GetArrayItem(child: Expression, ordinal: Expression) } } +/** + * Common trait for [[GetArrayItem]] and [[ElementAt]]. + */ +trait GetArrayItemUtil extends BinaryExpression { + + private val child = left + private val ordinal = right + + /** `Null` is returned for invalid ordinals. */ + protected def computeNullabilityFromArray: Boolean = if (ordinal.foldable && !ordinal.nullable) { + val intOrdinal = ordinal.eval().asInstanceOf[Number].intValue() + child match { + case CreateArray(ar) if intOrdinal < ar.length => + ar(intOrdinal).nullable + case GetArrayStructFields(CreateArray(elements), field, _, _, _) + if intOrdinal < elements.length => + elements(intOrdinal).nullable || field.nullable + case _ => + true + } + } else { + true + } +} + /** * Common trait for [[GetMapValue]] and [[ElementAt]]. */ From 759095510eb77c1c1ee1a544014d40fc8e9555e1 Mon Sep 17 00:00:00 2001 From: Takeshi Yamamuro Date: Tue, 26 Feb 2019 11:38:50 +0900 Subject: [PATCH 3/7] Fix --- .../expressions/collectionOperations.scala | 2 +- .../expressions/complexTypeExtractors.scala | 34 +++++++++---------- 2 files changed, 18 insertions(+), 18 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala index 61e7daa2bf62e..9be01b6c2c609 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala @@ -1976,7 +1976,7 @@ case class ElementAt(left: Expression, right: Expression) } override def nullable: Boolean = left.dataType match { - case _: ArrayType => computeNullabilityFromArray + case _: ArrayType => computeNullabilityFromArray(left, right) case _: MapType => computeNullabilityFromMap } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala index 50e00bcdf6245..77bc40d86b98a 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala @@ -221,7 +221,8 @@ case class GetArrayStructFields( * We need to do type checking here as `ordinal` expression maybe unresolved. */ case class GetArrayItem(child: Expression, ordinal: Expression) - extends GetArrayItemUtil with ExpectsInputTypes with ExtractValue with NullIntolerant { + extends BinaryExpression with GetArrayItemUtil with ExpectsInputTypes with ExtractValue + with NullIntolerant { // We have done type checking for child in `ExtractValue`, so only need to check the `ordinal`. override def inputTypes: Seq[AbstractDataType] = Seq(AnyDataType, IntegralType) @@ -231,7 +232,7 @@ case class GetArrayItem(child: Expression, ordinal: Expression) override def left: Expression = child override def right: Expression = ordinal - override def nullable: Boolean = computeNullabilityFromArray + override def nullable: Boolean = computeNullabilityFromArray(left, right) override def dataType: DataType = child.dataType.asInstanceOf[ArrayType].elementType protected override def nullSafeEval(value: Any, ordinal: Any): Any = { @@ -267,25 +268,24 @@ case class GetArrayItem(child: Expression, ordinal: Expression) /** * Common trait for [[GetArrayItem]] and [[ElementAt]]. */ -trait GetArrayItemUtil extends BinaryExpression { - - private val child = left - private val ordinal = right +trait GetArrayItemUtil { /** `Null` is returned for invalid ordinals. */ - protected def computeNullabilityFromArray: Boolean = if (ordinal.foldable && !ordinal.nullable) { - val intOrdinal = ordinal.eval().asInstanceOf[Number].intValue() - child match { - case CreateArray(ar) if intOrdinal < ar.length => - ar(intOrdinal).nullable - case GetArrayStructFields(CreateArray(elements), field, _, _, _) + protected def computeNullabilityFromArray(child: Expression, ordinal: Expression): Boolean = { + if (ordinal.foldable && !ordinal.nullable) { + val intOrdinal = ordinal.eval().asInstanceOf[Number].intValue() + child match { + case CreateArray(ar) if intOrdinal < ar.length => + ar(intOrdinal).nullable + case GetArrayStructFields(CreateArray(elements), field, _, _, _) if intOrdinal < elements.length => - elements(intOrdinal).nullable || field.nullable - case _ => - true + elements(intOrdinal).nullable || field.nullable + case _ => + true + } + } else { + true } - } else { - true } } From 499a57ba4ff79bee745a2aa04df253bb78b905aa Mon Sep 17 00:00:00 2001 From: Takeshi Yamamuro Date: Fri, 1 Mar 2019 09:36:52 +0900 Subject: [PATCH 4/7] Fix --- .../catalyst/expressions/collectionOperations.scala | 2 +- .../catalyst/expressions/complexTypeExtractors.scala | 11 ++++------- .../expressions/CollectionExpressionsSuite.scala | 4 ++-- 3 files changed, 7 insertions(+), 10 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala index 9be01b6c2c609..a6a6355f26f0d 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala @@ -1930,7 +1930,7 @@ case class ArrayPosition(left: Expression, right: Expression) """, since = "2.4.0") case class ElementAt(left: Expression, right: Expression) - extends GetMapValueUtil with GetArrayItemUtil { + extends GetMapValueUtil with GetArrayItemUtil { @transient private lazy val mapKeyType = left.dataType.asInstanceOf[MapType].keyType diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala index 77bc40d86b98a..4b4c5cb845c25 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala @@ -222,7 +222,7 @@ case class GetArrayStructFields( */ case class GetArrayItem(child: Expression, ordinal: Expression) extends BinaryExpression with GetArrayItemUtil with ExpectsInputTypes with ExtractValue - with NullIntolerant { + with NullIntolerant { // We have done type checking for child in `ExtractValue`, so only need to check the `ordinal`. override def inputTypes: Seq[AbstractDataType] = Seq(AnyDataType, IntegralType) @@ -294,13 +294,10 @@ trait GetArrayItemUtil { */ trait GetMapValueUtil extends BinaryExpression with ImplicitCastInputTypes { - private val child = left - private val key = right - /** `Null` is returned for invalid ordinals. */ - protected def computeNullabilityFromMap: Boolean = if (key.foldable && !key.nullable) { - val keyObj = key.eval() - child match { + protected def computeNullabilityFromMap: Boolean = if (right.foldable && !right.nullable) { + val keyObj = right.eval() + left match { case m: CreateMap if m.resolved => m.keys.zip(m.values).filter { case (k, _) => k.foldable && !k.nullable }.find { case (k, _) if k.eval() == keyObj => true diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/CollectionExpressionsSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/CollectionExpressionsSuite.scala index 63f08eba00407..3b0cfe69730a5 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/CollectionExpressionsSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/CollectionExpressionsSuite.scala @@ -1092,7 +1092,7 @@ class CollectionExpressionsSuite extends SparkFunSuite with ExpressionEvalHelper checkEvaluation(ElementAt(mb0, Literal(Array[Byte](3, 4))), null) } - test("SPARK-26965 correctly handles ElementAt nullability for arrays") { + test("correctly handles ElementAt nullability for arrays") { // CreateArray case val a = AttributeReference("a", IntegerType, nullable = false)() val b = AttributeReference("b", IntegerType, nullable = true)() @@ -1125,7 +1125,7 @@ class CollectionExpressionsSuite extends SparkFunSuite with ExpressionEvalHelper assert(ElementAt(stArray4, Literal(1)).nullable) } - test("SPARK-26965 correctly handles ElementAt nullability for maps") { + test("correctly handles ElementAt nullability for maps") { // String key test val k1 = Literal("k1") val v1 = AttributeReference("v1", StringType, nullable = true)() From 7faf51a54aff455bd0ebf5f76babda2099bba165 Mon Sep 17 00:00:00 2001 From: Takeshi Yamamuro Date: Mon, 4 Mar 2019 09:12:01 +0900 Subject: [PATCH 5/7] Fix --- .../expressions/complexTypeExtractors.scala | 22 ++++++------------- .../CollectionExpressionsSuite.scala | 19 ---------------- .../expressions/ComplexTypeSuite.scala | 19 ---------------- 3 files changed, 7 insertions(+), 53 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala index 4b4c5cb845c25..5a091c9b3cb8a 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala @@ -294,21 +294,13 @@ trait GetArrayItemUtil { */ trait GetMapValueUtil extends BinaryExpression with ImplicitCastInputTypes { - /** `Null` is returned for invalid ordinals. */ - protected def computeNullabilityFromMap: Boolean = if (right.foldable && !right.nullable) { - val keyObj = right.eval() - left match { - case m: CreateMap if m.resolved => - m.keys.zip(m.values).filter { case (k, _) => k.foldable && !k.nullable }.find { - case (k, _) if k.eval() == keyObj => true - case _ => false - }.map(_._2.nullable).getOrElse(true) - case _ => - true - } - } else { - true - } + /** + * `Null` is returned for invalid ordinals. + * + * TODO: We could make nullability more precise in foldable cases (e.g., literal input). + * But, since the key search is O(n), revisit this. + */ + protected def computeNullabilityFromMap: Boolean = true // todo: current search is O(n), improve it. def getValueEval(value: Any, ordinal: Any, keyType: DataType, ordering: Ordering[Any]): Any = { diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/CollectionExpressionsSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/CollectionExpressionsSuite.scala index 3b0cfe69730a5..2ddad744cbab0 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/CollectionExpressionsSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/CollectionExpressionsSuite.scala @@ -1125,25 +1125,6 @@ class CollectionExpressionsSuite extends SparkFunSuite with ExpressionEvalHelper assert(ElementAt(stArray4, Literal(1)).nullable) } - test("correctly handles ElementAt nullability for maps") { - // String key test - val k1 = Literal("k1") - val v1 = AttributeReference("v1", StringType, nullable = true)() - val k2 = Literal("k2") - val v2 = AttributeReference("v2", StringType, nullable = false)() - val map1 = CreateMap(k1 :: v1 :: k2 :: v2 :: Nil) - assert(ElementAt(map1, Literal("k1")).nullable) - assert(!ElementAt(map1, Literal("k2")).nullable) - assert(ElementAt(map1, Literal("non-existent-key")).nullable) - - // Complex type key test - val k3 = Literal.create((1, "a")) - val k4 = Literal.create((2, "b")) - val map2 = CreateMap(k3 :: v1 :: k4 :: v2 :: Nil) - assert(ElementAt(map2, Literal.create((1, "a"))).nullable) - assert(!ElementAt(map2, Literal.create((2, "b"))).nullable) - } - test("Concat") { // Primitive-type elements val ai0 = Literal.create(Seq(1, 2, 3), ArrayType(IntegerType, containsNull = false)) diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ComplexTypeSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ComplexTypeSuite.scala index d65b49f11884d..d8d65715281d4 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ComplexTypeSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ComplexTypeSuite.scala @@ -110,25 +110,6 @@ class ComplexTypeSuite extends SparkFunSuite with ExpressionEvalHelper { checkEvaluation(GetMapValue(nestedMap, Literal("a")), Map("b" -> "c")) } - test("SPARK-26747 handles GetMapValue nullability correctly when input key is foldable") { - // String key test - val k1 = Literal("k1") - val v1 = AttributeReference("v1", StringType, nullable = true)() - val k2 = Literal("k2") - val v2 = AttributeReference("v2", StringType, nullable = false)() - val map1 = CreateMap(k1 :: v1 :: k2 :: v2 :: Nil) - assert(GetMapValue(map1, Literal("k1")).nullable) - assert(!GetMapValue(map1, Literal("k2")).nullable) - assert(GetMapValue(map1, Literal("non-existent-key")).nullable) - - // Complex type key test - val k3 = Literal.create((1, "a")) - val k4 = Literal.create((2, "b")) - val map2 = CreateMap(k3 :: v1 :: k4 :: v2 :: Nil) - assert(GetMapValue(map2, Literal.create((1, "a"))).nullable) - assert(!GetMapValue(map2, Literal.create((2, "b"))).nullable) - } - test("GetStructField") { val typeS = StructType(StructField("a", IntegerType) :: Nil) val struct = Literal.create(create_row(1), typeS) From b95a38910edab3f20c11c28a59fecdfca1f48d26 Mon Sep 17 00:00:00 2001 From: Takeshi Yamamuro Date: Mon, 4 Mar 2019 10:02:02 +0900 Subject: [PATCH 6/7] Fix --- .../catalyst/expressions/collectionOperations.scala | 2 +- .../catalyst/expressions/complexTypeExtractors.scala | 12 +++--------- 2 files changed, 4 insertions(+), 10 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala index a6a6355f26f0d..018b6b9c9c375 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/collectionOperations.scala @@ -1977,7 +1977,7 @@ case class ElementAt(left: Expression, right: Expression) override def nullable: Boolean = left.dataType match { case _: ArrayType => computeNullabilityFromArray(left, right) - case _: MapType => computeNullabilityFromMap + case _: MapType => true } override def nullSafeEval(value: Any, ordinal: Any): Any = doElementAt(value, ordinal) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala index 5a091c9b3cb8a..3c4a9eb2ec7e9 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala @@ -294,14 +294,6 @@ trait GetArrayItemUtil { */ trait GetMapValueUtil extends BinaryExpression with ImplicitCastInputTypes { - /** - * `Null` is returned for invalid ordinals. - * - * TODO: We could make nullability more precise in foldable cases (e.g., literal input). - * But, since the key search is O(n), revisit this. - */ - protected def computeNullabilityFromMap: Boolean = true - // todo: current search is O(n), improve it. def getValueEval(value: Any, ordinal: Any, keyType: DataType, ordering: Ordering[Any]): Any = { val map = value.asInstanceOf[MapData] @@ -396,7 +388,9 @@ case class GetMapValue(child: Expression, key: Expression) override def left: Expression = child override def right: Expression = key - override def nullable: Boolean = computeNullabilityFromMap + + /** `Null` is returned for invalid ordinals. */ + override def nullable: Boolean = true override def dataType: DataType = child.dataType.asInstanceOf[MapType].valueType // todo: current search is O(n), improve it. From 95f286237b1f08fc4f624f3ffab900c80ef6c8f8 Mon Sep 17 00:00:00 2001 From: Takeshi Yamamuro Date: Mon, 4 Mar 2019 16:13:05 +0900 Subject: [PATCH 7/7] Add comments --- .../sql/catalyst/expressions/complexTypeExtractors.scala | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala index 3c4a9eb2ec7e9..e9d60ed3a481f 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/complexTypeExtractors.scala @@ -389,7 +389,13 @@ case class GetMapValue(child: Expression, key: Expression) override def left: Expression = child override def right: Expression = key - /** `Null` is returned for invalid ordinals. */ + /** + * `Null` is returned for invalid ordinals. + * + * TODO: We could make nullability more precise in foldable cases (e.g., literal input). + * But, since the key search is O(n), it takes much time to compute nullability. + * If we find efficient key searches, revisit this. + */ override def nullable: Boolean = true override def dataType: DataType = child.dataType.asInstanceOf[MapType].valueType