diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/higherOrderFunctions.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/higherOrderFunctions.scala index f87b4f70d298a..5d7cef9785d72 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/higherOrderFunctions.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/higherOrderFunctions.scala @@ -86,9 +86,15 @@ case class NamedLambdaVariable( override def qualifier: Seq[String] = Seq.empty + override def stateful: Boolean = true + override def newInstance(): NamedExpression = copy(exprId = NamedExpression.newExprId, value = new AtomicReference()) + override def withNewChildrenInternal( + newChildren: IndexedSeq[Expression]): NamedLambdaVariable = + copy(value = new AtomicReference()) + override def toAttribute: Attribute = { AttributeReference(name, dataType, nullable, Metadata.empty)(exprId, Seq.empty) } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/regexpExpressions.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/regexpExpressions.scala index e99c6b2b29e79..114b8afa9f051 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/regexpExpressions.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/regexpExpressions.scala @@ -723,6 +723,7 @@ case class RegExpReplace(subject: Expression, regexp: Expression, rep: Expressio // last replacement string, we don't want to convert a UTF8String => java.langString every time. @transient private var lastReplacement: String = _ @transient private var lastReplacementInUTF8: UTF8String = _ + override def stateful: Boolean = true final override val nodePatterns: Seq[TreePattern] = Seq(REGEXP_REPLACE) override def nullSafeEval(s: Any, p: Any, r: Any, i: Any): Any = { @@ -840,6 +841,7 @@ abstract class RegExpExtractBase @transient private var lastRegex: UTF8String = _ // last regex pattern, we cache it for performance concern @transient private var pattern: Pattern = _ + override def stateful: Boolean = true final override val nodePatterns: Seq[TreePattern] = Seq(REGEXP_EXTRACT_FAMILY) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/stringExpressions.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/stringExpressions.scala index e41562f6f1d89..9b1b33c55cb92 100755 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/stringExpressions.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/stringExpressions.scala @@ -1184,6 +1184,7 @@ case class StringTranslate(srcExpr: Expression, matchingExpr: Expression, replac @transient private var lastMatching: UTF8String = _ @transient private var lastReplace: UTF8String = _ @transient private var dict: JMap[String, String] = _ + override def stateful: Boolean = true final lazy val collationId: Int = first.dataType.asInstanceOf[StringType].collationId @@ -3489,6 +3490,7 @@ case class FormatNumber(x: Expression, d: Expression) // as a decimal separator. @transient private lazy val numberFormat = new DecimalFormat("", new DecimalFormatSymbols(Locale.US)) + override def stateful: Boolean = true override protected def nullSafeEval(xObject: Any, dObject: Any): Any = { right.dataType match { diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/HigherOrderFunctionsSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/HigherOrderFunctionsSuite.scala index c33d258ac4de6..b9c870f310906 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/HigherOrderFunctionsSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/HigherOrderFunctionsSuite.scala @@ -940,6 +940,16 @@ class HigherOrderFunctionsSuite extends SparkFunSuite with ExpressionEvalHelper "actualType" -> toSQLType(StringType) ))) } + + test("NamedLambdaVariable is stateful and produces a fresh copy") { + val lv = NamedLambdaVariable("x", IntegerType, nullable = false) + assert(lv.stateful, "NamedLambdaVariable.stateful should be true") + val copy = lv.freshCopyIfContainsStatefulExpression() + assert(copy ne lv, + "freshCopyIfContainsStatefulExpression should return a new instance for NamedLambdaVariable") + assert(copy.asInstanceOf[NamedLambdaVariable].value ne lv.value, + "fresh copy should have an independent AtomicReference value") + } } case class CodegenFallbackExpr(child: Expression) extends UnaryExpression with CodegenFallback { diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/RegexpExpressionsSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/RegexpExpressionsSuite.scala index 0bf29553ea33d..5377b9c7fc0b9 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/RegexpExpressionsSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/RegexpExpressionsSuite.scala @@ -711,4 +711,28 @@ class RegexpExpressionsSuite extends SparkFunSuite with ExpressionEvalHelper { ) ) } + + test("RegExpReplace and RegExpExtractBase are stateful and produce fresh copies") { + val s = Literal("hello world") + val p = Literal("(\\w+)") + val r = Literal("X") + + val replace = RegExpReplace(s, p, r) + assert(replace.stateful, "RegExpReplace.stateful should be true") + val replaceCopy = replace.freshCopyIfContainsStatefulExpression() + assert(replaceCopy ne replace, + "freshCopyIfContainsStatefulExpression should return a new instance for RegExpReplace") + + val extract = RegExpExtract(s, p, Literal(1)) + assert(extract.stateful, "RegExpExtract.stateful should be true") + val extractCopy = extract.freshCopyIfContainsStatefulExpression() + assert(extractCopy ne extract, + "freshCopyIfContainsStatefulExpression should return a new instance for RegExpExtract") + + val extractAll = RegExpExtractAll(s, p, Literal(1)) + assert(extractAll.stateful, "RegExpExtractAll.stateful should be true") + val extractAllCopy = extractAll.freshCopyIfContainsStatefulExpression() + assert(extractAllCopy ne extractAll, + "freshCopyIfContainsStatefulExpression should return a new instance for RegExpExtractAll") + } } diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/StringExpressionsSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/StringExpressionsSuite.scala index 711b5edd72ad1..050a0245343ff 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/StringExpressionsSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/StringExpressionsSuite.scala @@ -2268,4 +2268,23 @@ class StringExpressionsSuite extends SparkFunSuite with ExpressionEvalHelper { } } } + + test("StringTranslate and FormatNumber are stateful and produce fresh copies") { + val src = Literal("aeiou") + val matching = Literal("aeiou") + val replace = Literal("12345") + val translate = StringTranslate(src, matching, replace) + assert(translate.stateful, "StringTranslate.stateful should be true") + val translateCopy = translate.freshCopyIfContainsStatefulExpression() + assert(translateCopy ne translate, + "freshCopyIfContainsStatefulExpression should return a new instance for StringTranslate") + + val num = Literal(1234567.89) + val fmt = Literal(2) + val formatNumber = FormatNumber(num, fmt) + assert(formatNumber.stateful, "FormatNumber.stateful should be true") + val formatNumberCopy = formatNumber.freshCopyIfContainsStatefulExpression() + assert(formatNumberCopy ne formatNumber, + "freshCopyIfContainsStatefulExpression should return a new instance for FormatNumber") + } }