From 15c56ba54c85276c372a1569cd2fde68953f745c Mon Sep 17 00:00:00 2001 From: srielau Date: Sun, 6 Sep 2026 03:16:42 +0000 Subject: [PATCH 1/9] [SPARK-59275][PYTHON] Complete CHAR/VARCHAR support for Python UDFs and Arrow --- python/pyspark/sql/pandas/types.py | 8 ++- python/pyspark/sql/tests/arrow/test_arrow.py | 45 ++++++++++++++ .../sql/tests/arrow/test_arrow_python_udf.py | 61 +++++++++++++++---- python/pyspark/sql/tests/test_creation.py | 18 ++++++ python/pyspark/sql/tests/test_udf.py | 31 ++++++++++ .../sql/catalyst/util/CharVarcharUtils.scala | 22 +++++++ .../sql/execution/arrow/ArrowConverters.scala | 8 ++- .../python/ArrowEvalPythonExec.scala | 10 ++- ...umnarArrowEvalPythonEvaluatorFactory.scala | 28 ++++++--- .../python/EvalPythonEvaluatorFactory.scala | 7 ++- .../sql/execution/python/EvaluatePython.scala | 14 ++++- .../python/EvaluatePythonSuite.scala | 17 +++++- 12 files changed, 238 insertions(+), 31 deletions(-) diff --git a/python/pyspark/sql/pandas/types.py b/python/pyspark/sql/pandas/types.py index c4facb3e3a8b4..1e52b499c6d38 100644 --- a/python/pyspark/sql/pandas/types.py +++ b/python/pyspark/sql/pandas/types.py @@ -34,6 +34,7 @@ BinaryType, BooleanType, ByteType, + CharType, DataType, DateType, DayTimeIntervalType, @@ -57,6 +58,7 @@ TimestampType, TimeType, UserDefinedType, + VarcharType, VariantType, VariantVal, YearMonthIntervalType, @@ -133,7 +135,7 @@ def to_arrow_type( arrow_type = pa.float64() elif isinstance(dt, DecimalType): arrow_type = pa.decimal128(dt.precision, dt.scale) - elif isinstance(dt, StringType): + elif isinstance(dt, (StringType, CharType, VarcharType)): arrow_type = pa.large_string() if prefers_large_types else pa.string() elif isinstance(dt, BinaryType): arrow_type = pa.large_binary() if prefers_large_types else pa.binary() @@ -904,7 +906,7 @@ def _to_corrected_pandas_type(dt: DataType) -> Optional[Any]: return np.dtype("timedelta64[ns]") else: return np.dtype("timedelta64[us]") - elif isinstance(dt, StringType): + elif isinstance(dt, (StringType, CharType, VarcharType)): if LooseVersion(pd.__version__) < "3.0.0": return None else: @@ -933,7 +935,7 @@ def _to_corrected_pandas_ext_type(dt: DataType) -> Optional[Any]: return pd.Float64Dtype() elif isinstance(dt, BooleanType): return pd.BooleanDtype() - elif isinstance(dt, StringType): + elif isinstance(dt, (StringType, CharType, VarcharType)): return pd.StringDtype() else: return None diff --git a/python/pyspark/sql/tests/arrow/test_arrow.py b/python/pyspark/sql/tests/arrow/test_arrow.py index df4239651d344..582b58732fe28 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow.py +++ b/python/pyspark/sql/tests/arrow/test_arrow.py @@ -38,6 +38,7 @@ BinaryType, BooleanType, ByteType, + CharType, DateType, DayTimeIntervalType, DecimalType, @@ -54,6 +55,7 @@ TimestampNTZType, TimestampType, TimeType, + VarcharType, VariantType, ) from pyspark.testing.objects import ExamplePoint, ExamplePointUDT @@ -1433,6 +1435,49 @@ def test_toArrow_duplicate_field_names(self): ): df.limit(0).toArrow() + def test_char_varchar_explicit_schema_and_to_arrow(self): + schema = StructType( + [ + StructField("c", CharType(4)), + StructField("s", StructType([StructField("v", VarcharType(3))])), + StructField("a", ArrayType(CharType(2))), + ] + ) + values = [{"c": "ab", "s": {"v": "xyz"}, "a": ["z"]}] + inputs = [ + pd.DataFrame(values), + pa.Table.from_pylist(values), + ] + + with self.sql_conf( + { + "spark.sql.charVarchar.standardSemantics.enabled": "true", + "spark.sql.execution.arrow.pyspark.enabled": "true", + } + ): + for data in inputs: + with self.subTest(input_type=type(data).__name__): + df = self.spark.createDataFrame(data, schema) + self.assertEqual( + df.first(), + Row(c="ab ", s=Row(v="xyz"), a=["z "]), + ) + + table = df.toArrow() + self.assertEqual(table.schema.field("c").type, pa.string()) + self.assertEqual(table.schema.field("s").type.field("v").type, pa.string()) + self.assertEqual(table.schema.field("a").type.value_type, pa.string()) + self.assertEqual( + table.to_pylist(), + [{"c": "ab ", "s": {"v": "xyz"}, "a": ["z "]}], + ) + + invalid = pa.table({"c": ["abcd"]}) + with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): + self.spark.createDataFrame( + invalid, StructType([StructField("c", VarcharType(3))]) + ).collect() + def test_createDataFrame_pandas_duplicate_field_names(self): for arrow_enabled in [True, False]: with self.subTest(arrow_enabled=arrow_enabled): diff --git a/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py b/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py index e56fc33ae1630..1eb3558f66978 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py @@ -15,17 +15,19 @@ # limitations under the License. # +import tempfile import unittest from decimal import Decimal -from pyspark.errors import AnalysisException, PySparkNotImplementedError, PythonException +from pyspark.errors import AnalysisException, PythonException from pyspark.loose_version import LooseVersion from pyspark.sql import Row -from pyspark.sql.functions import col, udf +from pyspark.sql.functions import col, lit, pandas_udf, udf from pyspark.sql.tests.test_udf import BaseUDFTestsMixin from pyspark.sql.types import ( ArrayType, BinaryType, + CharType, DayTimeIntervalType, DecimalType, MapType, @@ -270,18 +272,53 @@ def f(v: float): rounded = df.select(f("v").alias("d")).first().d self.assertEqual(rounded, Decimal("1.233999999999999986")) - def test_err_return_type(self): - with self.assertRaises(PySparkNotImplementedError) as pe: - udf(lambda x: x, VarcharType(10), useArrow=True) - - self.check_error( - exception=pe.exception, - errorClass="NOT_IMPLEMENTED", - messageParameters={ - "feature": "Invalid return type with Arrow-optimized Python UDF: VarcharType(10)" - }, + def test_char_varchar_results(self): + schema = StructType( + [ + StructField("c", CharType(4)), + StructField("v", VarcharType(3)), + StructField("nested", ArrayType(CharType(2))), + StructField("m", MapType(CharType(2), VarcharType(3))), + ] ) + with self.sql_conf( + { + "spark.sql.charVarchar.standardSemantics.enabled": "true", + "spark.sql.execution.arrow.pythonUDF.columnarInput.enabled": "true", + } + ): + result = self.spark.range(1).select( + udf( + lambda _: ("ab", "xyz", ["z"], {"k": "xy"}), + schema, + useArrow=True, + )("id").alias("s") + ) + self.assertEqual( + result.first().s, + Row(c="ab ", v="xyz", nested=["z "], m={"k ": "xy"}), + ) + + pandas_result = self.spark.range(1).select( + pandas_udf(lambda values: values, CharType(4))(lit("ab")).alias("c") + ) + self.assertEqual(pandas_result.first().c, "ab ") + + invalid = self.spark.range(1).select( + udf(lambda _: "abcd", VarcharType(3), useArrow=True)("id") + ) + with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): + invalid.collect() + + with tempfile.TemporaryDirectory() as path: + self.spark.range(1).write.parquet(path) + columnar_input = self.spark.read.parquet(path) + columnar_result = columnar_input.select( + udf(lambda _: "ab", CharType(4), useArrow=True)("id") + ) + self.assertEqual(columnar_result.first()[0], "ab ") + def test_named_arguments_negative(self): @udf("int") def test_udf(a, b): diff --git a/python/pyspark/sql/tests/test_creation.py b/python/pyspark/sql/tests/test_creation.py index 0b4542206001d..008570f420713 100644 --- a/python/pyspark/sql/tests/test_creation.py +++ b/python/pyspark/sql/tests/test_creation.py @@ -27,6 +27,8 @@ ) from pyspark.sql import Row from pyspark.sql.types import ( + ArrayType, + CharType, DateType, DecimalType, IntegerType, @@ -36,6 +38,7 @@ TimestampNTZType, TimestampType, TimeType, + VarcharType, ) from pyspark.testing import assertDataFrameEqual from pyspark.testing.sqlutils import ReusedSQLTestCase @@ -48,6 +51,21 @@ class DataFrameCreationTestsMixin: + def test_char_varchar_explicit_schema(self): + schema = StructType( + [ + StructField("c", CharType(4)), + StructField("nested", ArrayType(VarcharType(3))), + ] + ) + + with self.sql_conf({"spark.sql.charVarchar.standardSemantics.enabled": "true"}): + df = self.spark.createDataFrame([("ab", ["xyz"])], schema) + self.assertEqual(df.first(), Row(c="ab ", nested=["xyz"])) + + with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): + self.spark.createDataFrame([("ab", ["abcd"])], schema).collect() + def test_create_str_from_dict(self): data = [ {"broker": {"teamId": 3398, "contactEmail": "abc.xyz@123.ca"}}, diff --git a/python/pyspark/sql/tests/test_udf.py b/python/pyspark/sql/tests/test_udf.py index b1a7f9214fbbf..b1545dc67a873 100644 --- a/python/pyspark/sql/tests/test_udf.py +++ b/python/pyspark/sql/tests/test_udf.py @@ -35,6 +35,7 @@ ArrayType, BinaryType, BooleanType, + CharType, DayTimeIntervalType, DoubleType, IntegerType, @@ -44,6 +45,7 @@ StructField, StructType, TimestampNTZType, + VarcharType, VariantType, VariantVal, ) @@ -59,6 +61,35 @@ class BaseUDFTestsMixin: + def test_char_varchar_results(self): + schema = StructType( + [ + StructField("c", CharType(4)), + StructField("v", VarcharType(3)), + StructField("nested", ArrayType(CharType(2))), + StructField("m", MapType(CharType(2), VarcharType(3))), + ] + ) + + with self.sql_conf({"spark.sql.charVarchar.standardSemantics.enabled": "true"}): + result = self.spark.range(1).select( + udf( + lambda _: ("ab", "xyz", ["z"], {"k": "xy"}), + schema, + useArrow=False, + )("id").alias("s") + ) + self.assertEqual( + result.first().s, + Row(c="ab ", v="xyz", nested=["z "], m={"k ": "xy"}), + ) + + invalid = self.spark.range(1).select( + udf(lambda _: "abcd", VarcharType(3), useArrow=False)("id") + ) + with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): + invalid.collect() + def test_udf_with_callable(self): data = self.spark.createDataFrame([(i, i**2) for i in range(10)], ["number", "squared"]) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/CharVarcharUtils.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/CharVarcharUtils.scala index 00adcfe69bc56..f1a891c17e28b 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/CharVarcharUtils.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/CharVarcharUtils.scala @@ -33,6 +33,28 @@ object CharVarcharUtils extends Logging with SparkCharVarcharUtils { // visible for testing private[sql] val CHAR_VARCHAR_TYPE_STRING_METADATA_KEY = "__CHAR_VARCHAR_TYPE_STRING" + /** + * Replaces CHAR/VARCHAR with their unconstrained string representation regardless of session + * configuration. Use this only at physical boundaries, such as Arrow, that encode all character + * string types as UTF8. + */ + private[sql] def replaceCharVarcharWithStringForPhysicalType(dt: DataType): DataType = dt match { + case ArrayType(elementType, containsNull) => + ArrayType(replaceCharVarcharWithStringForPhysicalType(elementType), containsNull) + case MapType(keyType, valueType, valueContainsNull) => + MapType( + replaceCharVarcharWithStringForPhysicalType(keyType), + replaceCharVarcharWithStringForPhysicalType(valueType), + valueContainsNull) + case StructType(fields) => + StructType(fields.map { field => + field.copy(dataType = replaceCharVarcharWithStringForPhysicalType(field.dataType)) + }) + case c: CharType => c.toStringType + case v: VarcharType => v.toStringType + case other => other + } + /** * Creates a StringRPad expression with the pad literal inheriting the collation from the * str expression's data type. This is necessary because StringRPad may be created after diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala index 464cac157b25c..00985ab90a524 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala @@ -37,6 +37,7 @@ import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions.{UnsafeProjection, UnsafeRow} import org.apache.spark.sql.catalyst.plans.logical.LocalRelation import org.apache.spark.sql.catalyst.types.DataTypeUtils.toAttributes +import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.classic.{DataFrame, Dataset, SparkSession} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ @@ -548,6 +549,7 @@ private[sql] object ArrowConverters extends Logging { errorOnDuplicatedFieldNames: Boolean, largeVarTypes: Boolean): DataFrame = { val attrs = toAttributes(schema) + val checkedAttrs = attrs.map(attr => CharVarcharUtils.stringLengthCheck(attr, attr.dataType)) val batchesInDriver = arrowBatches.toArray val shouldUseRDD = session.sessionState.conf .arrowLocalRelationThreshold < batchesInDriver.map(_.length.toLong).sum @@ -557,13 +559,15 @@ private[sql] object ArrowConverters extends Logging { val rdd = session.sparkContext .parallelize(batchesInDriver.toImmutableArraySeq, batchesInDriver.length) .mapPartitions { batchesInExecutors => - ArrowConverters.fromBatchIterator( + val rows = ArrowConverters.fromBatchIterator( batchesInExecutors, schema, timeZoneId, errorOnDuplicatedFieldNames, largeVarTypes, TaskContext.get()) + val projection = UnsafeProjection.create(checkedAttrs, attrs) + rows.map(row => projection(row).copy(): InternalRow) } session.internalCreateDataFrame(rdd.setName("arrow"), schema) } else { @@ -577,7 +581,7 @@ private[sql] object ArrowConverters extends Logging { TaskContext.get()) // Project/copy it. Otherwise, the Arrow column vectors will be closed and released out. - val proj = UnsafeProjection.create(attrs, attrs) + val proj = UnsafeProjection.create(checkedAttrs, attrs) Dataset.ofRows(session, LocalRelation(attrs, data.map(r => proj(r).copy()).toArray.toImmutableArraySeq)) } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ArrowEvalPythonExec.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ArrowEvalPythonExec.scala index 0170d4354a6a9..71b2eb35e41fe 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ArrowEvalPythonExec.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ArrowEvalPythonExec.scala @@ -24,6 +24,7 @@ import org.apache.spark.api.python.{ChainedPythonFunctions, PythonEvalType} import org.apache.spark.rdd.RDD import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions._ +import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.errors.QueryExecutionErrors import org.apache.spark.sql.execution.SparkPlan import org.apache.spark.sql.execution.metric.SQLMetric @@ -200,9 +201,12 @@ class ArrowEvalPythonEvaluatorFactory( schema: StructType, context: TaskContext): Iterator[InternalRow] = { - val outputTypes = output.drop(childOutput.length).map(_.dataType.transformRecursively { - case udt: UserDefinedType[_] => udt.sqlType - }) + val outputTypes = output.drop(childOutput.length).map { attr => + CharVarcharUtils.replaceCharVarcharWithStringForPhysicalType( + attr.dataType.transformRecursively { + case udt: UserDefinedType[_] => udt.sqlType + }) + } val batchIter = Iterator(iter) diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala index 3981602875adb..20694abe065ad 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala @@ -27,6 +27,7 @@ import org.apache.spark.api.python.ChainedPythonFunctions import org.apache.spark.internal.config.Python.PYTHON_UDF_PIPELINED_EXECUTION import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions._ +import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.errors.QueryExecutionErrors import org.apache.spark.sql.execution.RowToColumnConverter import org.apache.spark.sql.execution.metric.SQLMetric @@ -79,6 +80,15 @@ private[python] class ColumnarArrowEvalPythonEvaluatorFactory( sessionUUID: Option[String]) extends PartitionEvaluatorFactory[ColumnarBatch, ColumnarBatch] { + private val checkedOutput = childOutput ++ output.drop(childOutput.length).map { attr => + CharVarcharUtils.stringLengthCheck(attr, attr.dataType) + } + private val hasCharVarcharOutput = + output.drop(childOutput.length).exists(attr => CharVarcharUtils.hasCharVarchar(attr.dataType)) + private val physicalOutputSchema = CharVarcharUtils + .replaceCharVarcharWithStringForPhysicalType(outputSchema) + .asInstanceOf[StructType] + override def createEvaluator() : PartitionEvaluator[ColumnarBatch, ColumnarBatch] = new ColumnarArrowEvalPythonPartitionEvaluator @@ -137,10 +147,12 @@ private[python] class ColumnarArrowEvalPythonEvaluatorFactory( StructField(s"_$i", dt) }.toArray) - val outputTypes = output.drop(childOutput.length).map( - _.dataType.transformRecursively { - case udt: UserDefinedType[_] => udt.sqlType - }) + val outputTypes = output.drop(childOutput.length).map { attr => + CharVarcharUtils.replaceCharVarcharWithStringForPhysicalType( + attr.dataType.transformRecursively { + case udt: UserDefinedType[_] => udt.sqlType + }) + } val inputColumnIndices = resolveColumnIndices(allInputs.toSeq) @@ -151,7 +163,7 @@ private[python] class ColumnarArrowEvalPythonEvaluatorFactory( batch.column(0).isInstanceOf[ArrowColumnVector] } - if (inputColumnIndices.isDefined && isArrow) { + if (inputColumnIndices.isDefined && isArrow && !hasCharVarcharOutput) { // Path 1: Arrow columnar -- full optimization. evalArrowColumnar(peekIter, context, pyFuncs, argMetas, udfInputSchema, outputTypes, inputColumnIndices.get) @@ -287,7 +299,7 @@ private[python] class ColumnarArrowEvalPythonEvaluatorFactory( } val joined = new JoinedRow - val resultProj = UnsafeProjection.create(output, output) + val resultProj = UnsafeProjection.create(checkedOutput, output) val rowIter = resultIter.flatMap { batch => validateOutputTypes(batch, outputTypes) @@ -315,9 +327,9 @@ private[python] class ColumnarArrowEvalPythonEvaluatorFactory( private def rowsToColumnarBatches( rowIter: Iterator[InternalRow], context: TaskContext): Iterator[ColumnarBatch] = { - val converters = new RowToColumnConverter(outputSchema) + val converters = new RowToColumnConverter(physicalOutputSchema) val vectors = OnHeapColumnVector - .allocateColumns(batchSize, outputSchema).toSeq + .allocateColumns(batchSize, physicalOutputSchema).toSeq val cb = new ColumnarBatch(vectors.toArray) context.addTaskCompletionListener[Unit] { _ => cb.close() } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala index c0a983c60afa8..b0c8a50b7b1da 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala @@ -26,6 +26,7 @@ import org.apache.spark.api.python.ChainedPythonFunctions import org.apache.spark.internal.config.Python.PYTHON_UDF_PIPELINED_EXECUTION import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions._ +import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata import org.apache.spark.sql.types.{DataType, StructField, StructType} import org.apache.spark.util.Utils @@ -36,6 +37,10 @@ abstract class EvalPythonEvaluatorFactory( output: Seq[Attribute]) extends PartitionEvaluatorFactory[InternalRow, InternalRow] { + private val checkedOutput = childOutput ++ output.drop(childOutput.length).map { attr => + CharVarcharUtils.stringLengthCheck(attr, attr.dataType) + } + protected def evaluate( funcs: Seq[(ChainedPythonFunctions, Long)], argMetas: Array[Array[ArgumentMetadata]], @@ -119,7 +124,7 @@ abstract class EvalPythonEvaluatorFactory( evaluate(pyFuncs, argMetas, projectedRowIter, schema, context) val joined = new JoinedRow - val resultProj = UnsafeProjection.create(output, output) + val resultProj = UnsafeProjection.create(checkedOutput, output) outputRowIterator.map { outputRow => resultProj(joined(queue.remove(), outputRow)) diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala index 3b9c2a3e69cd2..916ba2765527f 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala @@ -30,7 +30,7 @@ import org.apache.spark.rdd.RDD import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.types.ops.TypeApiOps -import org.apache.spark.sql.catalyst.util.{ArrayBasedMapData, ArrayData, GenericArrayData, MapData, STUtils} +import org.apache.spark.sql.catalyst.util.{ArrayBasedMapData, ArrayData, CharVarcharCodegenUtils, GenericArrayData, MapData, STUtils} import org.apache.spark.sql.types._ import org.apache.spark.unsafe.types.{BinaryView, UTF8String, VariantVal} @@ -207,6 +207,18 @@ object EvaluatePython { case c: Int => c.toLong } + case c: CharType => (obj: Any) => nullSafeConvert(obj) { + case _ => + CharVarcharCodegenUtils.charTypeWriteSideCheck( + UTF8String.fromString(obj.toString), c.length) + } + + case v: VarcharType => (obj: Any) => nullSafeConvert(obj) { + case _ => + CharVarcharCodegenUtils.varcharTypeWriteSideCheck( + UTF8String.fromString(obj.toString), v.length) + } + case _: StringType => (obj: Any) => nullSafeConvert(obj) { case _ => UTF8String.fromString(obj.toString) } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/EvaluatePythonSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/EvaluatePythonSuite.scala index ec26f1b2a865f..8852bf609ceba 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/EvaluatePythonSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/EvaluatePythonSuite.scala @@ -20,10 +20,25 @@ package org.apache.spark.sql.execution.python import org.apache.spark.{SparkFunSuite, SparkIllegalArgumentException, SparkRuntimeException} import org.apache.spark.sql.catalyst.util.STUtils import org.apache.spark.sql.types._ -import org.apache.spark.unsafe.types.BinaryView +import org.apache.spark.unsafe.types.{BinaryView, UTF8String} class EvaluatePythonSuite extends SparkFunSuite { + test("SPARK-59275: makeFromJava enforces CHAR/VARCHAR results") { + val charResult = EvaluatePython.makeFromJava(CharType(4))("ab") + assert(charResult === UTF8String.fromString("ab ")) + + val varcharResult = EvaluatePython.makeFromJava(VarcharType(4))("abcd ") + assert(varcharResult === UTF8String.fromString("abcd")) + + checkError( + exception = intercept[SparkRuntimeException] { + EvaluatePython.makeFromJava(VarcharType(4))("abcde") + }, + condition = "EXCEED_LIMIT_LENGTH", + parameters = Map("limit" -> "4")) + } + // POINT(1 2) in WKB, little-endian. private val pointWkb: Array[Byte] = "010100000000000000000031400000000000001C40" .grouped(2).map(Integer.parseInt(_, 16).toByte).toArray From 5ca8ceae9f94a8e26b4f2728b41cc6dd6d492756 Mon Sep 17 00:00:00 2001 From: srielau Date: Mon, 7 Sep 2026 00:00:19 +0000 Subject: [PATCH 2/9] fix: [SPARK-59275] address CHAR/VARCHAR review feedback --- python/pyspark/sql/tests/arrow/test_arrow.py | 20 ++++ .../sql/tests/arrow/test_arrow_python_udf.py | 9 -- python/pyspark/sql/tests/test_creation.py | 10 ++ python/pyspark/sql/tests/test_udf.py | 106 ++++++++++-------- python/pyspark/sql/udf.py | 38 +++++-- .../sql/execution/arrow/ArrowConverters.scala | 17 ++- ...umnarArrowEvalPythonEvaluatorFactory.scala | 14 ++- .../python/EvalPythonEvaluatorFactory.scala | 11 +- .../sql/execution/python/EvaluatePython.scala | 32 ++++-- .../spark/sql/IntegratedUDFTestUtils.scala | 30 +++++ .../python/ArrowColumnarPythonUDFSuite.scala | 36 +++++- 11 files changed, 240 insertions(+), 83 deletions(-) diff --git a/python/pyspark/sql/tests/arrow/test_arrow.py b/python/pyspark/sql/tests/arrow/test_arrow.py index 582b58732fe28..b5d9c9527066d 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow.py +++ b/python/pyspark/sql/tests/arrow/test_arrow.py @@ -1478,6 +1478,26 @@ def test_char_varchar_explicit_schema_and_to_arrow(self): invalid, StructType([StructField("c", VarcharType(3))]) ).collect() + legacy_schema = StructType( + [ + StructField("c", CharType(3)), + StructField("v", VarcharType(3)), + ] + ) + with self.sql_conf( + { + "spark.sql.legacy.charVarcharAsString": "true", + "spark.sql.preserveCharVarcharTypeInfo": "false", + "spark.sql.charVarchar.standardSemantics.enabled": "false", + "spark.sql.execution.arrow.pyspark.enabled": "true", + "spark.sql.execution.arrow.localRelationThreshold": "0", + } + ): + df = self.spark.createDataFrame( + pa.table({"c": ["a"], "v": ["abcd"]}), legacy_schema + ) + self.assertEqual(df.first(), Row(c="a", v="abcd")) + def test_createDataFrame_pandas_duplicate_field_names(self): for arrow_enabled in [True, False]: with self.subTest(arrow_enabled=arrow_enabled): diff --git a/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py b/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py index 1eb3558f66978..6a84ce2628f5f 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py @@ -15,7 +15,6 @@ # limitations under the License. # -import tempfile import unittest from decimal import Decimal @@ -311,14 +310,6 @@ def test_char_varchar_results(self): with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): invalid.collect() - with tempfile.TemporaryDirectory() as path: - self.spark.range(1).write.parquet(path) - columnar_input = self.spark.read.parquet(path) - columnar_result = columnar_input.select( - udf(lambda _: "ab", CharType(4), useArrow=True)("id") - ) - self.assertEqual(columnar_result.first()[0], "ab ") - def test_named_arguments_negative(self): @udf("int") def test_udf(a, b): diff --git a/python/pyspark/sql/tests/test_creation.py b/python/pyspark/sql/tests/test_creation.py index 008570f420713..1dbc966e29281 100644 --- a/python/pyspark/sql/tests/test_creation.py +++ b/python/pyspark/sql/tests/test_creation.py @@ -66,6 +66,16 @@ def test_char_varchar_explicit_schema(self): with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): self.spark.createDataFrame([("ab", ["abcd"])], schema).collect() + with self.sql_conf( + { + "spark.sql.legacy.charVarcharAsString": "true", + "spark.sql.preserveCharVarcharTypeInfo": "false", + "spark.sql.charVarchar.standardSemantics.enabled": "false", + } + ): + df = self.spark.createDataFrame([("ab", ["abcd"])], schema) + self.assertEqual(df.first(), Row(c="ab", nested=["abcd"])) + def test_create_str_from_dict(self): data = [ {"broker": {"teamId": 3398, "contactEmail": "abc.xyz@123.ca"}}, diff --git a/python/pyspark/sql/tests/test_udf.py b/python/pyspark/sql/tests/test_udf.py index b1545dc67a873..270357a293de4 100644 --- a/python/pyspark/sql/tests/test_udf.py +++ b/python/pyspark/sql/tests/test_udf.py @@ -27,7 +27,12 @@ import unittest from contextlib import redirect_stdout -from pyspark.errors import AnalysisException, PySparkTypeError, PythonException +from pyspark.errors import ( + AnalysisException, + PySparkNotImplementedError, + PySparkTypeError, + PythonException, +) from pyspark.logger import PySparkLogger from pyspark.sql import Column, Row, SparkSession from pyspark.sql.functions import assert_true, col, lit, rand, udf @@ -57,7 +62,7 @@ test_not_compiled_message, ) from pyspark.testing.utils import assertDataFrameEqual, eventually, timeout -from pyspark.util import is_remote_only +from pyspark.util import PythonEvalType, is_remote_only class BaseUDFTestsMixin: @@ -90,6 +95,44 @@ def test_char_varchar_results(self): with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): invalid.collect() + def test_char_varchar_legacy_as_string(self): + with self.sql_conf( + { + "spark.sql.legacy.charVarcharAsString": "true", + "spark.sql.preserveCharVarcharTypeInfo": "false", + "spark.sql.charVarchar.standardSemantics.enabled": "false", + } + ): + result = self.spark.range(1).select( + udf(lambda _: "a", CharType(3), useArrow=False)("id").alias("c"), + udf(lambda _: "abcd", VarcharType(3), useArrow=False)("id").alias("v"), + ) + self.assertEqual(result.first(), Row(c="a", v="abcd")) + + def test_char_varchar_non_scalar_return_types_unsupported(self): + nested_return_type = StructType( + [StructField("nested", ArrayType(CharType(3)))] + ) + struct_eval_types = [ + PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF, + PythonEvalType.SQL_GROUPED_MAP_ARROW_UDF, + PythonEvalType.SQL_MAP_PANDAS_ITER_UDF, + PythonEvalType.SQL_MAP_ARROW_ITER_UDF, + PythonEvalType.SQL_COGROUPED_MAP_PANDAS_UDF, + PythonEvalType.SQL_COGROUPED_MAP_ARROW_UDF, + ] + aggregate_eval_types = [ + PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF, + PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF, + ] + + for eval_type in struct_eval_types: + with self.assertRaisesRegex(PySparkNotImplementedError, "Invalid return type"): + UserDefinedFunction._check_return_type(nested_return_type, eval_type) + for eval_type in aggregate_eval_types: + with self.assertRaisesRegex(PySparkNotImplementedError, "Invalid return type"): + UserDefinedFunction._check_return_type(VarcharType(3), eval_type) + def test_udf_with_callable(self): data = self.spark.createDataFrame([(i, i**2) for i in range(10)], ["number", "squared"]) @@ -1482,49 +1525,24 @@ def my_udf(input_val): result_type = df_result.schema["result"].dataType self.assertEqual(result_type, StringType("fr")) - def test_udf_with_char_varchar_return_type(self): - char_type, char_value = ("char(10)", "a") - varchar_type, varchar_value = ("varchar(8)", "a") - array_with_char_type, array_with_char_type_value = ("array", ["a", "b"]) - array_with_varchar_type, array_with_varchar_value = ("array", ["a", "b"]) - map_type, map_value = (f"map<{char_type}, {varchar_type}>", {"a": "b"}) - struct_type, struct_value = ( - f"struct", - {"f1": "a", "f2": "b"}, + def test_udf_with_char_varchar_return_type_legacy(self): + schema = StructType( + [ + StructField("chars", ArrayType(CharType(3))), + StructField("values", MapType(CharType(2), VarcharType(3))), + ] ) - - pairs = [ - (char_type, char_value), - (varchar_type, varchar_value), - (array_with_char_type, array_with_char_type_value), - (array_with_varchar_type, array_with_varchar_value), - (map_type, map_value), - (struct_type, struct_value), - ( - f"struct", - f"{{'f1': {array_with_char_type_value}, 'f2': {array_with_varchar_value}, " - f"'f3': {map_value}}}", - ), - ( - f"map<{array_with_char_type}, {array_with_varchar_type}>", - f"{{{array_with_char_type_value}: {array_with_varchar_value}}}", - ), - (f"array<{struct_type}>", [struct_value, struct_value]), - ] - - for return_type, return_value in pairs: - with self.assertRaisesRegex( - Exception, - "(Please use a different output data type for your UDF or DataFrame|" - "Invalid return type with Arrow-optimized Python UDF)", - ): - - @udf(return_type) - def my_udf(): - return return_value - - self.spark.range(1).select(my_udf().alias("result")).show() + with self.sql_conf( + { + "spark.sql.legacy.charVarcharAsString": "true", + "spark.sql.preserveCharVarcharTypeInfo": "false", + "spark.sql.charVarchar.standardSemantics.enabled": "false", + } + ): + result = self.spark.range(1).select( + udf(lambda _: (["a"], {"k": "abcd"}), schema, useArrow=False)("id") + ) + self.assertEqual(result.first()[0], Row(chars=["a"], values={"k": "abcd"})) def test_udf_binary_type(self): def get_binary_type(x): diff --git a/python/pyspark/sql/udf.py b/python/pyspark/sql/udf.py index 2618ecb7d4b04..ff618e4f2197e 100644 --- a/python/pyspark/sql/udf.py +++ b/python/pyspark/sql/udf.py @@ -29,9 +29,12 @@ from pyspark.sql.pandas.types import to_arrow_type from pyspark.sql.pandas.utils import require_minimum_pandas_version, require_minimum_pyarrow_version from pyspark.sql.types import ( + CharType, DataType, StringType, StructType, + VarcharType, + _has_type, _parse_datatype_string, ) from pyspark.sql.utils import get_active_spark_context @@ -314,9 +317,24 @@ def _conf_is_true(key: str, default: Optional[str] = None) -> bool: @staticmethod def _check_return_type(returnType: DataType, evalType: int) -> None: + char_varchar_supported_eval_types = ( + PythonEvalType.SQL_ARROW_BATCHED_UDF, + PythonEvalType.SQL_SCALAR_PANDAS_UDF, + PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF, + PythonEvalType.SQL_SCALAR_ARROW_UDF, + PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF, + ) + + def check_arrow_type() -> None: + if evalType not in char_varchar_supported_eval_types and _has_type( + returnType, (CharType, VarcharType) + ): + raise TypeError + to_arrow_type(returnType, timezone="UTC") + if evalType == PythonEvalType.SQL_ARROW_BATCHED_UDF: try: - to_arrow_type(returnType, timezone="UTC") + check_arrow_type() except TypeError: raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", @@ -330,7 +348,7 @@ def _check_return_type(returnType: DataType, evalType: int) -> None: or evalType == PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF ): try: - to_arrow_type(returnType, timezone="UTC") + check_arrow_type() except TypeError: raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", @@ -343,7 +361,7 @@ def _check_return_type(returnType: DataType, evalType: int) -> None: or evalType == PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF ): try: - to_arrow_type(returnType, timezone="UTC") + check_arrow_type() except TypeError: raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", @@ -358,7 +376,7 @@ def _check_return_type(returnType: DataType, evalType: int) -> None: ): if isinstance(returnType, StructType): try: - to_arrow_type(returnType, timezone="UTC") + check_arrow_type() except TypeError: raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", @@ -383,7 +401,7 @@ def _check_return_type(returnType: DataType, evalType: int) -> None: ): if isinstance(returnType, StructType): try: - to_arrow_type(returnType, timezone="UTC") + check_arrow_type() except TypeError: raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", @@ -405,7 +423,7 @@ def _check_return_type(returnType: DataType, evalType: int) -> None: ): if isinstance(returnType, StructType): try: - to_arrow_type(returnType, timezone="UTC") + check_arrow_type() except TypeError: raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", @@ -425,7 +443,7 @@ def _check_return_type(returnType: DataType, evalType: int) -> None: elif evalType == PythonEvalType.SQL_COGROUPED_MAP_PANDAS_UDF: if isinstance(returnType, StructType): try: - to_arrow_type(returnType, timezone="UTC") + check_arrow_type() except TypeError: raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", @@ -444,7 +462,7 @@ def _check_return_type(returnType: DataType, evalType: int) -> None: elif evalType == PythonEvalType.SQL_COGROUPED_MAP_ARROW_UDF: if isinstance(returnType, StructType): try: - to_arrow_type(returnType, timezone="UTC") + check_arrow_type() except TypeError: raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", @@ -471,7 +489,7 @@ def _check_return_type(returnType: DataType, evalType: int) -> None: f"{returnType}" }, ) - to_arrow_type(returnType, timezone="UTC") + check_arrow_type() except TypeError: raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", @@ -483,7 +501,7 @@ def _check_return_type(returnType: DataType, evalType: int) -> None: elif evalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF: try: # Different from SQL_GROUPED_AGG_PANDAS_UDF, StructType is allowed here - to_arrow_type(returnType, timezone="UTC") + check_arrow_type() except TypeError: raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala index 00985ab90a524..590c061f777c3 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala @@ -549,7 +549,14 @@ private[sql] object ArrowConverters extends Logging { errorOnDuplicatedFieldNames: Boolean, largeVarTypes: Boolean): DataFrame = { val attrs = toAttributes(schema) - val checkedAttrs = attrs.map(attr => CharVarcharUtils.stringLengthCheck(attr, attr.dataType)) + val applyCharVarcharChecks = + CharVarcharUtils.hasCharVarchar(schema) && + CharVarcharUtils.shouldApplyWriteSideLengthCheck(session.sessionState.conf) + val checkedAttrs = if (applyCharVarcharChecks) { + attrs.map(attr => CharVarcharUtils.stringLengthCheck(attr, attr.dataType)) + } else { + attrs + } val batchesInDriver = arrowBatches.toArray val shouldUseRDD = session.sessionState.conf .arrowLocalRelationThreshold < batchesInDriver.map(_.length.toLong).sum @@ -566,8 +573,12 @@ private[sql] object ArrowConverters extends Logging { errorOnDuplicatedFieldNames, largeVarTypes, TaskContext.get()) - val projection = UnsafeProjection.create(checkedAttrs, attrs) - rows.map(row => projection(row).copy(): InternalRow) + if (applyCharVarcharChecks) { + val projection = UnsafeProjection.create(checkedAttrs, attrs) + rows.map(row => projection(row).copy(): InternalRow) + } else { + rows + } } session.internalCreateDataFrame(rdd.setName("arrow"), schema) } else { diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala index 20694abe065ad..2847ff6fd1cc7 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala @@ -33,6 +33,7 @@ import org.apache.spark.sql.execution.RowToColumnConverter import org.apache.spark.sql.execution.metric.SQLMetric import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata import org.apache.spark.sql.execution.vectorized.OnHeapColumnVector +import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{DataType, StructField, StructType, UserDefinedType} import org.apache.spark.sql.types.DataType.equalsIgnoreCompatibleCollation import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} @@ -80,11 +81,18 @@ private[python] class ColumnarArrowEvalPythonEvaluatorFactory( sessionUUID: Option[String]) extends PartitionEvaluatorFactory[ColumnarBatch, ColumnarBatch] { - private val checkedOutput = childOutput ++ output.drop(childOutput.length).map { attr => - CharVarcharUtils.stringLengthCheck(attr, attr.dataType) + private val applyCharVarcharChecks = + CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get) + private val checkedOutput = if (applyCharVarcharChecks) { + childOutput ++ output.drop(childOutput.length).map { attr => + CharVarcharUtils.stringLengthCheck(attr, attr.dataType) + } + } else { + output } private val hasCharVarcharOutput = - output.drop(childOutput.length).exists(attr => CharVarcharUtils.hasCharVarchar(attr.dataType)) + applyCharVarcharChecks && + output.drop(childOutput.length).exists(attr => CharVarcharUtils.hasCharVarchar(attr.dataType)) private val physicalOutputSchema = CharVarcharUtils .replaceCharVarcharWithStringForPhysicalType(outputSchema) .asInstanceOf[StructType] diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala index b0c8a50b7b1da..ef64897f9f9d7 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala @@ -28,6 +28,7 @@ import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata +import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{DataType, StructField, StructType} import org.apache.spark.util.Utils @@ -37,8 +38,14 @@ abstract class EvalPythonEvaluatorFactory( output: Seq[Attribute]) extends PartitionEvaluatorFactory[InternalRow, InternalRow] { - private val checkedOutput = childOutput ++ output.drop(childOutput.length).map { attr => - CharVarcharUtils.stringLengthCheck(attr, attr.dataType) + private val applyCharVarcharChecks = + CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get) + private val checkedOutput = if (applyCharVarcharChecks) { + childOutput ++ output.drop(childOutput.length).map { attr => + CharVarcharUtils.stringLengthCheck(attr, attr.dataType) + } + } else { + output } protected def evaluate( diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala index 916ba2765527f..3eefb77d751b5 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala @@ -30,7 +30,8 @@ import org.apache.spark.rdd.RDD import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.types.ops.TypeApiOps -import org.apache.spark.sql.catalyst.util.{ArrayBasedMapData, ArrayData, CharVarcharCodegenUtils, GenericArrayData, MapData, STUtils} +import org.apache.spark.sql.catalyst.util.{ArrayBasedMapData, ArrayData, CharVarcharCodegenUtils, CharVarcharUtils, GenericArrayData, MapData, STUtils} +import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ import org.apache.spark.unsafe.types.{BinaryView, UTF8String, VariantVal} @@ -146,10 +147,19 @@ object EvaluatePython { * Make a converter that converts `obj` to the type specified by the data type, or returns * null if the type of obj is unexpected. Because Python doesn't enforce the type. */ - def makeFromJava(dataType: DataType): Any => Any = - TypeApiOps(dataType).flatMap(_.makeFromJava).getOrElse(makeFromJavaDefault(dataType)) + def makeFromJava(dataType: DataType): Any => Any = { + val applyCharVarcharChecks = + CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get) + makeFromJava(dataType, applyCharVarcharChecks) + } + + private def makeFromJava(dataType: DataType, applyCharVarcharChecks: Boolean): Any => Any = + TypeApiOps(dataType).flatMap(_.makeFromJava) + .getOrElse(makeFromJavaDefault(dataType, applyCharVarcharChecks)) - private def makeFromJavaDefault(dataType: DataType): Any => Any = dataType match { + private def makeFromJavaDefault( + dataType: DataType, + applyCharVarcharChecks: Boolean): Any => Any = dataType match { case BooleanType => (obj: Any) => nullSafeConvert(obj) { case b: Boolean => b } @@ -207,13 +217,13 @@ object EvaluatePython { case c: Int => c.toLong } - case c: CharType => (obj: Any) => nullSafeConvert(obj) { + case c: CharType if applyCharVarcharChecks => (obj: Any) => nullSafeConvert(obj) { case _ => CharVarcharCodegenUtils.charTypeWriteSideCheck( UTF8String.fromString(obj.toString), c.length) } - case v: VarcharType => (obj: Any) => nullSafeConvert(obj) { + case v: VarcharType if applyCharVarcharChecks => (obj: Any) => nullSafeConvert(obj) { case _ => CharVarcharCodegenUtils.varcharTypeWriteSideCheck( UTF8String.fromString(obj.toString), v.length) @@ -229,7 +239,7 @@ object EvaluatePython { } case ArrayType(elementType, _) => - val elementFromJava = makeFromJava(elementType) + val elementFromJava = makeFromJava(elementType, applyCharVarcharChecks) (obj: Any) => nullSafeConvert(obj) { case c: java.util.List[_] => @@ -239,8 +249,8 @@ object EvaluatePython { } case MapType(keyType, valueType, _) => - val keyFromJava = makeFromJava(keyType) - val valueFromJava = makeFromJava(valueType) + val keyFromJava = makeFromJava(keyType, applyCharVarcharChecks) + val valueFromJava = makeFromJava(valueType, applyCharVarcharChecks) (obj: Any) => nullSafeConvert(obj) { case javaMap: java.util.Map[_, _] => @@ -251,7 +261,7 @@ object EvaluatePython { } case StructType(fields) => - val fieldsFromJava = fields.map(f => makeFromJava(f.dataType)) + val fieldsFromJava = fields.map(f => makeFromJava(f.dataType, applyCharVarcharChecks)) (obj: Any) => nullSafeConvert(obj) { case c if c.getClass.isArray => @@ -273,7 +283,7 @@ object EvaluatePython { row } - case udt: UserDefinedType[_] => makeFromJava(udt.sqlType) + case udt: UserDefinedType[_] => makeFromJava(udt.sqlType, applyCharVarcharChecks) case VariantType => (obj: Any) => nullSafeConvert(obj) { case s: java.util.HashMap[_, _] => diff --git a/sql/core/src/test/scala/org/apache/spark/sql/IntegratedUDFTestUtils.scala b/sql/core/src/test/scala/org/apache/spark/sql/IntegratedUDFTestUtils.scala index a5f3e72f47b7d..ad24ed3c23375 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/IntegratedUDFTestUtils.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/IntegratedUDFTestUtils.scala @@ -1410,6 +1410,35 @@ object IntegratedUDFTestUtils extends SQLHelper { val prettyName: String = "Scalar Pandas UDF" } + case class TestTypedScalarPandasUDF( + name: String, + returnType: DataType) extends TestUDF { + private[IntegratedUDFTestUtils] lazy val udf = new UserDefinedPythonFunction( + name = name, + func = SimplePythonFunction( + command = pandasFunc.toImmutableArraySeq, + envVars = workerEnv.clone().asInstanceOf[java.util.Map[String, String]], + pythonIncludes = List.empty[String].asJava, + pythonExec = pythonExec, + pythonVer = pythonVer, + broadcastVars = List.empty[Broadcast[PythonBroadcast]].asJava, + accumulator = null), + dataType = returnType, + pythonEvalType = PythonEvalType.SQL_SCALAR_PANDAS_UDF, + udfDeterministic = true) { + + override def builder(e: Seq[Expression]): Expression = { + assert(e.length == 1, "Defined UDF only has one column") + new PythonUDFWithoutId( + super.builder(e).asInstanceOf[PythonUDF]) + } + } + + def apply(exprs: Column*): Column = udf(exprs: _*) + + val prettyName: String = "Typed Scalar Pandas UDF" + } + /** * A Grouped Aggregate Pandas UDF that takes one column, executes the * Python native function calculating the count of the column using pandas. @@ -1606,6 +1635,7 @@ object IntegratedUDFTestUtils extends SQLHelper { def registerTestUDF(testUDF: TestUDF, session: classic.SparkSession): Unit = testUDF match { case udf: TestPythonUDF => session.udf.registerPython(udf.name, udf.udf) case udf: TestScalarPandasUDF => session.udf.registerPython(udf.name, udf.udf) + case udf: TestTypedScalarPandasUDF => session.udf.registerPython(udf.name, udf.udf) case udf: TestGroupedAggPandasUDF => session.udf.registerPython(udf.name, udf.udf) case udf: TestScalaUDF => val registry = session.sessionState.functionRegistry diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala index faf8c77678f3a..add3d133e56aa 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala @@ -17,11 +17,12 @@ package org.apache.spark.sql.execution.python +import org.apache.spark.SparkException import org.apache.spark.sql.IntegratedUDFTestUtils import org.apache.spark.sql.execution.SparkPlan import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.test.SharedSparkSession -import org.apache.spark.sql.types.StringType +import org.apache.spark.sql.types.{CharType, StringType, VarcharType} /** * End-to-end tests for the Arrow columnar Python UDF input path. @@ -103,6 +104,39 @@ class ArrowColumnarPythonUDFSuite extends SharedSparkSession { } } + test("Arrow-backed source: CHAR/VARCHAR output checks") { + assume(shouldTestPandasUDFs) + withSQLConf( + SQLConf.ARROW_PYSPARK_EXECUTION_ENABLED.key -> "true", + SQLConf.ARROW_PYSPARK_UDF_COLUMNAR_INPUT_ENABLED.key -> "true", + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + val charUDF = TestTypedScalarPandasUDF( + name = "arrow_char_udf", returnType = CharType(4)) + val varcharUDF = TestTypedScalarPandasUDF( + name = "arrow_varchar_udf", returnType = VarcharType(3)) + registerTestUDF(charUDF, spark) + registerTestUDF(varcharUDF, spark) + + val df = readArrowSource(numRows = 10) + val padded = df.selectExpr( + "id", "name", "value", "data", + "arrow_char_udf(id) as udf_id") + val arrowExec = collectNodes[ArrowEvalPythonExec]( + padded.queryExecution.executedPlan).head + assert(arrowExec.child.supportsColumnar, + "ArrowEvalPythonExec should retain its Arrow-backed columnar child") + assert(padded.select("udf_id").collect().map(_.getString(0)).toSeq === + (0 until 10).map(_.toString.padTo(4, ' ').mkString)) + + val exception = intercept[SparkException] { + df.selectExpr( + "id", "name", "value", "data", + "arrow_varchar_udf(name) as udf_name").collect() + } + assert(exception.getMessage.contains("EXCEED_LIMIT_LENGTH")) + } + } + test("Arrow-backed source: multiple UDF columns") { assume(shouldTestPandasUDFs) withSQLConf(SQLConf.ARROW_PYSPARK_EXECUTION_ENABLED.key -> "true") { From 157b0d247d12b20bbe10c2bf0b80d96d7dcbfdd9 Mon Sep 17 00:00:00 2001 From: srielau Date: Mon, 7 Sep 2026 17:04:08 +0000 Subject: [PATCH 3/9] fix: [SPARK-59275] address additional CHAR/VARCHAR review feedback --- .../sql/tests/arrow/test_arrow_python_udf.py | 14 ++++++ .../sql/tests/arrow/test_arrow_udtf.py | 27 ++++++++++-- python/pyspark/sql/tests/test_udf.py | 34 +++++++++++++++ python/pyspark/sql/udf.py | 10 ++++- python/pyspark/sql/udtf.py | 20 ++++++++- .../execution/python/ExtractPythonUDFs.scala | 3 ++ .../python/ArrowColumnarPythonUDFSuite.scala | 32 ++++++++++++++ .../python/ExtractPythonUDFsSuite.scala | 43 +++++++++++++++++++ 8 files changed, 176 insertions(+), 7 deletions(-) diff --git a/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py b/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py index 6a84ce2628f5f..d2e66b7355010 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py @@ -310,6 +310,20 @@ def test_char_varchar_results(self): with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): invalid.collect() + def test_char_varchar_results_legacy_as_string(self): + with self.sql_conf( + { + "spark.sql.legacy.charVarcharAsString": "true", + "spark.sql.preserveCharVarcharTypeInfo": "false", + "spark.sql.charVarchar.standardSemantics.enabled": "false", + } + ): + result = self.spark.range(1).select( + udf(lambda _: "a", CharType(3), useArrow=True)("id").alias("c"), + udf(lambda _: "abcd", VarcharType(3), useArrow=True)("id").alias("v"), + ) + self.assertEqual(result.first(), Row(c="a", v="abcd")) + def test_named_arguments_negative(self): @udf("int") def test_udf(a, b): diff --git a/python/pyspark/sql/tests/arrow/test_arrow_udtf.py b/python/pyspark/sql/tests/arrow/test_arrow_udtf.py index ac695a9f8001e..5f59bcbc97e66 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_udtf.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_udtf.py @@ -18,9 +18,9 @@ import unittest from typing import Iterator, Optional -from pyspark.errors import PySparkAttributeError, PythonException +from pyspark.errors import PySparkAttributeError, PySparkNotImplementedError, PythonException from pyspark.sql.functions import arrow_udtf, lit -from pyspark.sql.types import IntegerType, Row, StructField, StructType +from pyspark.sql.types import ArrayType, CharType, IntegerType, Row, StructField, StructType from pyspark.testing import assertDataFrameEqual from pyspark.testing.sqlutils import ReusedSQLTestCase from pyspark.testing.utils import have_pyarrow, pyarrow_requirement_message @@ -33,6 +33,25 @@ @unittest.skipIf(not have_pyarrow, pyarrow_requirement_message) class ArrowUDTFTestsMixin: + def test_char_varchar_return_types_unsupported(self): + @arrow_udtf(returnType="c CHAR(3)") + class DirectCharUDTF: + def eval(self) -> Iterator["pa.Table"]: + yield pa.table({"c": ["a"]}) + + nested_type = StructType([StructField("nested", ArrayType(CharType(3)))]) + + @arrow_udtf(returnType=nested_type) + class NestedCharUDTF: + def eval(self) -> Iterator["pa.Table"]: + yield pa.table({"nested": [["a"]]}) + + for function in (DirectCharUDTF, NestedCharUDTF): + with self.assertRaisesRegex( + PySparkNotImplementedError, "Invalid return type with Arrow UDTFs" + ): + function() + def test_arrow_udtf_data_conversion_error(self): from pyspark.sql.functions import udtf @@ -40,8 +59,8 @@ def test_arrow_udtf_data_conversion_error(self): class DataConversionErrorUDTF: def eval(self): # Return a non-tuple value when multiple return values are expected. - # This will cause LocalDataToArrowConversion.convert to fail with TypeError (len() on int), - # which should be wrapped in UDTF_ARROW_DATA_CONVERSION_ERROR. + # This causes LocalDataToArrowConversion.convert to fail with TypeError + # (len() on int), which should be wrapped in UDTF_ARROW_DATA_CONVERSION_ERROR. yield 1 # Enable Arrow optimization for regular UDTFs diff --git a/python/pyspark/sql/tests/test_udf.py b/python/pyspark/sql/tests/test_udf.py index 270357a293de4..4b4052ad11ff0 100644 --- a/python/pyspark/sql/tests/test_udf.py +++ b/python/pyspark/sql/tests/test_udf.py @@ -109,6 +109,38 @@ def test_char_varchar_legacy_as_string(self): ) self.assertEqual(result.first(), Row(c="a", v="abcd")) + def test_char_varchar_intermediate_udf_results(self): + for use_arrow in (False, True): + with self.subTest(use_arrow=use_arrow): + inner_char = udf(lambda _: "a", CharType(3), useArrow=use_arrow) + inner_varchar = udf(lambda _: "abcd", VarcharType(3), useArrow=use_arrow) + outer = udf(lambda value: value, StringType(), useArrow=use_arrow) + + with self.sql_conf( + {"spark.sql.charVarchar.standardSemantics.enabled": "true"} + ): + padded = self.spark.range(1).select( + outer(inner_char("id")).alias("result") + ) + self.assertEqual(padded.first().result, "a ") + + invalid = self.spark.range(1).select(outer(inner_varchar("id"))) + with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): + invalid.collect() + + with self.sql_conf( + { + "spark.sql.legacy.charVarcharAsString": "true", + "spark.sql.preserveCharVarcharTypeInfo": "false", + "spark.sql.charVarchar.standardSemantics.enabled": "false", + } + ): + result = self.spark.range(1).select( + outer(inner_char("id")).alias("c"), + outer(inner_varchar("id")).alias("v"), + ) + self.assertEqual(result.first(), Row(c="a", v="abcd")) + def test_char_varchar_non_scalar_return_types_unsupported(self): nested_return_type = StructType( [StructField("nested", ArrayType(CharType(3)))] @@ -123,7 +155,9 @@ def test_char_varchar_non_scalar_return_types_unsupported(self): ] aggregate_eval_types = [ PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF, + PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF, PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF, + PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF, ] for eval_type in struct_eval_types: diff --git a/python/pyspark/sql/udf.py b/python/pyspark/sql/udf.py index ff618e4f2197e..c529d7fb617db 100644 --- a/python/pyspark/sql/udf.py +++ b/python/pyspark/sql/udf.py @@ -478,7 +478,10 @@ def check_arrow_type() -> None: "return_type": str(returnType), }, ) - elif evalType == PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF: + elif ( + evalType == PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF + or evalType == PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF + ): try: # StructType is not yet allowed as a return type, explicitly check here to fail fast if isinstance(returnType, StructType): @@ -498,7 +501,10 @@ def check_arrow_type() -> None: f"{returnType}" }, ) - elif evalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF: + elif ( + evalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF + or evalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF + ): try: # Different from SQL_GROUPED_AGG_PANDAS_UDF, StructType is allowed here check_arrow_type() diff --git a/python/pyspark/sql/udtf.py b/python/pyspark/sql/udtf.py index e8dc25dc907ec..76cb5910268fa 100644 --- a/python/pyspark/sql/udtf.py +++ b/python/pyspark/sql/udtf.py @@ -28,11 +28,19 @@ from pyspark.errors import ( PySparkAttributeError, PySparkImportError, + PySparkNotImplementedError, PySparkPicklingError, PySparkTypeError, ) from pyspark.sql.pandas.utils import require_minimum_pandas_version, require_minimum_pyarrow_version -from pyspark.sql.types import DataType, StructType, _parse_datatype_string +from pyspark.sql.types import ( + CharType, + DataType, + StructType, + VarcharType, + _has_type, + _parse_datatype_string, +) from pyspark.sql.udf import _wrap_function from pyspark.util import PythonEvalType @@ -369,6 +377,16 @@ def returnType(self) -> Optional[StructType]: "return_type": f"{parsed}", }, ) + if self.evalType in ( + PythonEvalType.SQL_ARROW_TABLE_UDF, + PythonEvalType.SQL_ARROW_UDTF, + ) and _has_type(parsed, (CharType, VarcharType)): + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={ + "feature": f"Invalid return type with Arrow UDTFs: {parsed}" + }, + ) self._returnType_placeholder = parsed return self._returnType_placeholder diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFs.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFs.scala index c290ec04dedab..2a0e68da67581 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFs.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFs.scala @@ -30,6 +30,7 @@ import org.apache.spark.sql.catalyst.expressions.aggregate.AggregateExpression import org.apache.spark.sql.catalyst.plans.logical._ import org.apache.spark.sql.catalyst.rules.Rule import org.apache.spark.sql.catalyst.trees.TreePattern._ +import org.apache.spark.sql.catalyst.util.CharVarcharUtils /** @@ -200,6 +201,8 @@ object ExtractPythonUDFs extends Rule[LogicalPlan] with Logging { case Seq(child: PythonUDF) => correctEvalType(e, pythonUDFArrowFallbackOnUDT) == correctEvalType(child, pythonUDFArrowFallbackOnUDT) && + !(CharVarcharUtils.shouldApplyWriteSideLengthCheck(conf) && + CharVarcharUtils.hasCharVarchar(child.dataType)) && shouldExtractUDFExpressionTree(child, pythonUDFArrowFallbackOnUDT) // Python UDF can't be evaluated directly in JVM case children => !children.exists(hasScalarPythonUDF) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala index add3d133e56aa..7ce29543221a9 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala @@ -137,6 +137,38 @@ class ArrowColumnarPythonUDFSuite extends SharedSparkSession { } } + test("Arrow-backed source: legacy CHAR/VARCHAR output remains unchecked") { + assume(shouldTestPandasUDFs) + withSQLConf( + SQLConf.ARROW_PYSPARK_EXECUTION_ENABLED.key -> "true", + SQLConf.ARROW_PYSPARK_UDF_COLUMNAR_INPUT_ENABLED.key -> "true", + SQLConf.LEGACY_CHAR_VARCHAR_AS_STRING.key -> "true", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "false", + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false") { + val charUDF = TestTypedScalarPandasUDF( + name = "legacy_arrow_char_udf", returnType = CharType(4)) + val varcharUDF = TestTypedScalarPandasUDF( + name = "legacy_arrow_varchar_udf", returnType = VarcharType(3)) + registerTestUDF(charUDF, spark) + registerTestUDF(varcharUDF, spark) + + val result = readArrowSource(numRows = 10).selectExpr( + "id", "name", "value", "data", + "legacy_arrow_char_udf(id) as udf_id", + "legacy_arrow_varchar_udf(name) as udf_name") + val arrowExec = collectNodes[ArrowEvalPythonExec]( + result.queryExecution.executedPlan).head + assert(arrowExec.child.supportsColumnar, + "ArrowEvalPythonExec should retain its Arrow-backed columnar child") + + val rows = result.select("udf_id", "udf_name").collect() + rows.zipWithIndex.foreach { case (row, index) => + assert(row.getString(0) === index.toString) + assert(row.getString(1) === s"row_$index") + } + } + } + test("Arrow-backed source: multiple UDF columns") { assume(shouldTestPandasUDFs) withSQLConf(SQLConf.ARROW_PYSPARK_EXECUTION_ENABLED.key -> "true") { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFsSuite.scala index 4c4273006b925..284723f9e2c3b 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFsSuite.scala @@ -17,6 +17,7 @@ package org.apache.spark.sql.execution.python +import org.apache.spark.api.python.PythonEvalType import org.apache.spark.sql.catalyst.plans.logical.{ArrowEvalPython, BatchEvalPython, Limit, LocalLimit} import org.apache.spark.sql.execution.{FileSourceScanExec, SparkPlan} import org.apache.spark.sql.execution.datasources.v2.BatchScanExec @@ -24,6 +25,7 @@ import org.apache.spark.sql.execution.datasources.v2.parquet.ParquetScan import org.apache.spark.sql.functions.col import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.test.SharedSparkSession +import org.apache.spark.sql.types.{CharType, DataType, StringType} class ExtractPythonUDFsSuite extends SharedSparkSession { import testImplicits._ @@ -32,6 +34,14 @@ class ExtractPythonUDFsSuite extends SharedSparkSession { val batchedNondeterministicPythonUDF = new MyDummyNondeterministicPythonUDF val scalarPandasUDF = new MyDummyScalarPandasUDF + private def typedPythonUDF(dataType: DataType, evalType: Int) = + UserDefinedPythonFunction( + name = "typedPythonUDF", + func = new DummyUDF, + dataType = dataType, + pythonEvalType = evalType, + udfDeterministic = true) + private def collectBatchExec(plan: SparkPlan): Seq[BatchEvalPythonExec] = plan.collect { case b: BatchEvalPythonExec => b } @@ -56,6 +66,39 @@ class ExtractPythonUDFsSuite extends SharedSparkSession { assert(arrowEvalNodes.size == 1) } + test("CHAR/VARCHAR intermediate results split chained Python UDFs") { + withSQLConf(SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true") { + val df = Seq(("Hello", 4)).toDF("a", "b") + val batchInner = typedPythonUDF(CharType(3), PythonEvalType.SQL_BATCHED_UDF) + val batchOuter = typedPythonUDF(StringType, PythonEvalType.SQL_BATCHED_UDF) + val arrowInner = typedPythonUDF(CharType(3), PythonEvalType.SQL_SCALAR_PANDAS_UDF) + val arrowOuter = typedPythonUDF(StringType, PythonEvalType.SQL_SCALAR_PANDAS_UDF) + + assert(collectBatchExec( + df.select(batchOuter(batchInner(col("a")))).queryExecution.executedPlan).size == 2) + assert(collectArrowExec( + df.select(arrowOuter(arrowInner(col("a")))).queryExecution.executedPlan).size == 2) + } + } + + test("Legacy CHAR/VARCHAR intermediate results retain Python UDF chaining") { + withSQLConf( + SQLConf.LEGACY_CHAR_VARCHAR_AS_STRING.key -> "true", + SQLConf.PRESERVE_CHAR_VARCHAR_TYPE_INFO.key -> "false", + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "false") { + val df = Seq(("Hello", 4)).toDF("a", "b") + val batchInner = typedPythonUDF(CharType(3), PythonEvalType.SQL_BATCHED_UDF) + val batchOuter = typedPythonUDF(StringType, PythonEvalType.SQL_BATCHED_UDF) + val arrowInner = typedPythonUDF(CharType(3), PythonEvalType.SQL_SCALAR_PANDAS_UDF) + val arrowOuter = typedPythonUDF(StringType, PythonEvalType.SQL_SCALAR_PANDAS_UDF) + + assert(collectBatchExec( + df.select(batchOuter(batchInner(col("a")))).queryExecution.executedPlan).size == 1) + assert(collectArrowExec( + df.select(arrowOuter(arrowInner(col("a")))).queryExecution.executedPlan).size == 1) + } + } + test("Mixed Batched Python UDFs and Pandas UDF should be separate physical node") { val df = Seq(("Hello", 4)).toDF("a", "b") val df2 = df.withColumn("c", batchedPythonUDF(col("a"))) From b4582cf0532a2f76be9d1036eabfc89c1453445b Mon Sep 17 00:00:00 2001 From: srielau Date: Mon, 7 Sep 2026 17:16:22 +0000 Subject: [PATCH 4/9] fix: [SPARK-59275] validate Arrow UDTFs in Connect --- python/pyspark/sql/connect/udtf.py | 22 +++++++- .../sql/tests/arrow/test_arrow_python_udf.py | 28 +++++++++++ python/pyspark/sql/tests/test_udf.py | 50 +++++++++---------- python/pyspark/sql/udtf.py | 24 +++++---- 4 files changed, 85 insertions(+), 39 deletions(-) diff --git a/python/pyspark/sql/connect/udtf.py b/python/pyspark/sql/connect/udtf.py index aa7d86c0f72d7..e966a378a4ff4 100644 --- a/python/pyspark/sql/connect/udtf.py +++ b/python/pyspark/sql/connect/udtf.py @@ -32,8 +32,13 @@ from pyspark.sql.connect.types import UnparsedDataType from pyspark.sql.connect.utils import get_python_ver from pyspark.sql.pandas.utils import require_minimum_pandas_version, require_minimum_pyarrow_version -from pyspark.sql.types import DataType, StructType -from pyspark.sql.udtf import AnalyzeArgument, AnalyzeResult, _validate_udtf_handler # noqa: F401 +from pyspark.sql.types import DataType, StructType, _parse_datatype_string +from pyspark.sql.udtf import ( # noqa: F401 + AnalyzeArgument, + AnalyzeResult, + _check_arrow_udtf_return_type, + _validate_udtf_handler, +) from pyspark.sql.udtf import UDTFRegistration as PySparkUDTFRegistration from pyspark.util import PythonEvalType @@ -167,9 +172,21 @@ def __init__( self.evalType = evalType self.deterministic = deterministic + def _check_return_type(self) -> None: + if self.returnType is None: + return + return_type = ( + _parse_datatype_string(self.returnType.data_type_string) + if isinstance(self.returnType, UnparsedDataType) + else self.returnType + ) + _check_arrow_udtf_return_type(return_type, self.evalType) + def _build_common_inline_user_defined_table_function( self, *args: "ColumnOrName", **kwargs: "ColumnOrName" ) -> CommonInlineUserDefinedTableFunction: + self._check_return_type() + def to_expr(col: "ColumnOrName") -> Expression: if isinstance(col, Column): return col._expr @@ -245,6 +262,7 @@ def register( }, ) + f._check_return_type() self.sparkSession._client.register_udtf( f.func, f.returnType, name, f.evalType, f.deterministic ) diff --git a/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py b/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py index d2e66b7355010..e1c69000a39c4 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py @@ -324,6 +324,34 @@ def test_char_varchar_results_legacy_as_string(self): ) self.assertEqual(result.first(), Row(c="a", v="abcd")) + def test_char_varchar_intermediate_udf_results_arrow(self): + inner_char = udf(lambda _: "a", CharType(3), useArrow=True) + inner_varchar = udf(lambda _: "abcd", VarcharType(3), useArrow=True) + outer = udf(lambda value: value, StringType(), useArrow=True) + + with self.sql_conf({"spark.sql.charVarchar.standardSemantics.enabled": "true"}): + padded = self.spark.range(1).select( + outer(inner_char("id")).alias("result") + ) + self.assertEqual(padded.first().result, "a ") + + invalid = self.spark.range(1).select(outer(inner_varchar("id"))) + with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): + invalid.collect() + + with self.sql_conf( + { + "spark.sql.legacy.charVarcharAsString": "true", + "spark.sql.preserveCharVarcharTypeInfo": "false", + "spark.sql.charVarchar.standardSemantics.enabled": "false", + } + ): + result = self.spark.range(1).select( + outer(inner_char("id")).alias("c"), + outer(inner_varchar("id")).alias("v"), + ) + self.assertEqual(result.first(), Row(c="a", v="abcd")) + def test_named_arguments_negative(self): @udf("int") def test_udf(a, b): diff --git a/python/pyspark/sql/tests/test_udf.py b/python/pyspark/sql/tests/test_udf.py index 4b4052ad11ff0..72821153e821d 100644 --- a/python/pyspark/sql/tests/test_udf.py +++ b/python/pyspark/sql/tests/test_udf.py @@ -110,36 +110,32 @@ def test_char_varchar_legacy_as_string(self): self.assertEqual(result.first(), Row(c="a", v="abcd")) def test_char_varchar_intermediate_udf_results(self): - for use_arrow in (False, True): - with self.subTest(use_arrow=use_arrow): - inner_char = udf(lambda _: "a", CharType(3), useArrow=use_arrow) - inner_varchar = udf(lambda _: "abcd", VarcharType(3), useArrow=use_arrow) - outer = udf(lambda value: value, StringType(), useArrow=use_arrow) + inner_char = udf(lambda _: "a", CharType(3), useArrow=False) + inner_varchar = udf(lambda _: "abcd", VarcharType(3), useArrow=False) + outer = udf(lambda value: value, StringType(), useArrow=False) - with self.sql_conf( - {"spark.sql.charVarchar.standardSemantics.enabled": "true"} - ): - padded = self.spark.range(1).select( - outer(inner_char("id")).alias("result") - ) - self.assertEqual(padded.first().result, "a ") + with self.sql_conf({"spark.sql.charVarchar.standardSemantics.enabled": "true"}): + padded = self.spark.range(1).select( + outer(inner_char("id")).alias("result") + ) + self.assertEqual(padded.first().result, "a ") - invalid = self.spark.range(1).select(outer(inner_varchar("id"))) - with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): - invalid.collect() + invalid = self.spark.range(1).select(outer(inner_varchar("id"))) + with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): + invalid.collect() - with self.sql_conf( - { - "spark.sql.legacy.charVarcharAsString": "true", - "spark.sql.preserveCharVarcharTypeInfo": "false", - "spark.sql.charVarchar.standardSemantics.enabled": "false", - } - ): - result = self.spark.range(1).select( - outer(inner_char("id")).alias("c"), - outer(inner_varchar("id")).alias("v"), - ) - self.assertEqual(result.first(), Row(c="a", v="abcd")) + with self.sql_conf( + { + "spark.sql.legacy.charVarcharAsString": "true", + "spark.sql.preserveCharVarcharTypeInfo": "false", + "spark.sql.charVarchar.standardSemantics.enabled": "false", + } + ): + result = self.spark.range(1).select( + outer(inner_char("id")).alias("c"), + outer(inner_varchar("id")).alias("v"), + ) + self.assertEqual(result.first(), Row(c="a", v="abcd")) def test_char_varchar_non_scalar_return_types_unsupported(self): nested_return_type = StructType( diff --git a/python/pyspark/sql/udtf.py b/python/pyspark/sql/udtf.py index 76cb5910268fa..0850df082046d 100644 --- a/python/pyspark/sql/udtf.py +++ b/python/pyspark/sql/udtf.py @@ -325,6 +325,19 @@ def _validate_udtf_handler(cls: Any, returnType: Optional[Union[StructType, str] ) +def _check_arrow_udtf_return_type(return_type: DataType, eval_type: int) -> None: + if eval_type in ( + PythonEvalType.SQL_ARROW_TABLE_UDF, + PythonEvalType.SQL_ARROW_UDTF, + ) and _has_type(return_type, (CharType, VarcharType)): + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={ + "feature": f"Invalid return type with Arrow UDTFs: {return_type}" + }, + ) + + class UserDefinedTableFunction: """ User-defined table function in Python @@ -377,16 +390,7 @@ def returnType(self) -> Optional[StructType]: "return_type": f"{parsed}", }, ) - if self.evalType in ( - PythonEvalType.SQL_ARROW_TABLE_UDF, - PythonEvalType.SQL_ARROW_UDTF, - ) and _has_type(parsed, (CharType, VarcharType)): - raise PySparkNotImplementedError( - errorClass="NOT_IMPLEMENTED", - messageParameters={ - "feature": f"Invalid return type with Arrow UDTFs: {parsed}" - }, - ) + _check_arrow_udtf_return_type(parsed, self.evalType) self._returnType_placeholder = parsed return self._returnType_placeholder From 3234e913fb67566dee5127a9c7098188438856f5 Mon Sep 17 00:00:00 2001 From: srielau Date: Tue, 8 Sep 2026 17:56:49 +0000 Subject: [PATCH 5/9] fix: [SPARK-59275] address Python CHAR/VARCHAR review feedback --- .../resources/error/error-conditions.json | 5 + python/pyspark/sql/connect/udtf.py | 28 ++- .../sql/tests/arrow/test_arrow_udf_scalar.py | 31 +++ .../sql/tests/arrow/test_arrow_udtf.py | 20 +- .../sql/tests/connect/test_parity_udtf.py | 23 +++ .../tests/pandas/test_pandas_udf_scalar.py | 19 ++ python/pyspark/sql/tests/test_udf.py | 34 ++++ python/pyspark/sql/udf.py | 16 +- .../sql/catalyst/analysis/Analyzer.scala | 3 +- .../sql/catalyst/expressions/PythonUDF.scala | 9 +- .../sql/errors/QueryCompilationErrors.scala | 6 + .../spark/sql/api/python/PythonSQLUtils.scala | 15 +- .../sql/execution/arrow/ArrowConverters.scala | 25 ++- .../python/ArrowEvalPythonExec.scala | 2 +- .../python/BatchEvalPythonExec.scala | 10 +- .../python/BatchEvalPythonUDTFExec.scala | 2 +- ...umnarArrowEvalPythonEvaluatorFactory.scala | 34 ++-- .../python/EvalPythonEvaluatorFactory.scala | 24 +-- .../sql/execution/python/EvaluatePython.scala | 32 +++- .../python/UserDefinedPythonFunction.scala | 32 +++- ...nsformWithStateInPySparkDeserializer.scala | 59 ++++-- ...nsformWithStateInPySparkPythonRunner.scala | 11 +- ...ansformWithStateInPySparkStateServer.scala | 176 +++++++++++------- .../spark/sql/IntegratedUDFTestUtils.scala | 27 ++- .../arrow/ArrowConvertersSuite.scala | 35 +++- .../python/ArrowColumnarPythonUDFSuite.scala | 16 +- ...rmWithStateInPySparkStateServerSuite.scala | 62 +++++- 27 files changed, 607 insertions(+), 149 deletions(-) diff --git a/common/utils/src/main/resources/error/error-conditions.json b/common/utils/src/main/resources/error/error-conditions.json index ea8f77498e5d5..ea3f376f2940e 100644 --- a/common/utils/src/main/resources/error/error-conditions.json +++ b/common/utils/src/main/resources/error/error-conditions.json @@ -9145,6 +9145,11 @@ "Purge table." ] }, + "PYTHON_ARROW_UDTF_CHAR_VARCHAR_RETURN_TYPE" : { + "message" : [ + "Arrow-optimized Python UDTFs do not support CHAR/VARCHAR in return type ." + ] + }, "PYTHON_UDF_IN_ON_CLAUSE" : { "message" : [ "Python UDF in the ON clause of a JOIN. In case of an INNER JOIN consider rewriting to a CROSS JOIN with a WHERE clause." diff --git a/python/pyspark/sql/connect/udtf.py b/python/pyspark/sql/connect/udtf.py index e966a378a4ff4..471e80c7e02bb 100644 --- a/python/pyspark/sql/connect/udtf.py +++ b/python/pyspark/sql/connect/udtf.py @@ -19,7 +19,7 @@ """ import warnings -from typing import TYPE_CHECKING, Any, List, Optional, Type, Union +from typing import TYPE_CHECKING, Any, List, Optional, Set, Type, Union from pyspark.errors import PySparkAttributeError, PySparkRuntimeError, PySparkTypeError from pyspark.sql.connect.column import Column @@ -32,7 +32,7 @@ from pyspark.sql.connect.types import UnparsedDataType from pyspark.sql.connect.utils import get_python_ver from pyspark.sql.pandas.utils import require_minimum_pandas_version, require_minimum_pyarrow_version -from pyspark.sql.types import DataType, StructType, _parse_datatype_string +from pyspark.sql.types import DataType, StructType from pyspark.sql.udtf import ( # noqa: F401 AnalyzeArgument, AnalyzeResult, @@ -171,21 +171,31 @@ def __init__( self._name = name or func.__name__ self.evalType = evalType self.deterministic = deterministic + self._validated_return_type_session_ids: Set[str] = set() - def _check_return_type(self) -> None: - if self.returnType is None: + def _check_return_type(self, session: "SparkSession") -> None: + if self.returnType is None or self.evalType not in ( + PythonEvalType.SQL_ARROW_TABLE_UDF, + PythonEvalType.SQL_ARROW_UDTF, + ): + return + if session._session_id in self._validated_return_type_session_ids: return return_type = ( - _parse_datatype_string(self.returnType.data_type_string) + session._parse_ddl(self.returnType.data_type_string) if isinstance(self.returnType, UnparsedDataType) else self.returnType ) _check_arrow_udtf_return_type(return_type, self.evalType) + self._validated_return_type_session_ids.add(session._session_id) def _build_common_inline_user_defined_table_function( - self, *args: "ColumnOrName", **kwargs: "ColumnOrName" + self, + session: "SparkSession", + *args: "ColumnOrName", + **kwargs: "ColumnOrName", ) -> CommonInlineUserDefinedTableFunction: - self._check_return_type() + self._check_return_type(session) def to_expr(col: "ColumnOrName") -> Expression: if isinstance(col, Column): @@ -218,7 +228,7 @@ def __call__(self, *args: "ColumnOrName", **kwargs: "ColumnOrName") -> "DataFram session = SparkSession.active() - plan = self._build_common_inline_user_defined_table_function(*args, **kwargs) + plan = self._build_common_inline_user_defined_table_function(session, *args, **kwargs) return DataFrame(plan, session) def asDeterministic(self) -> "UserDefinedTableFunction": @@ -262,7 +272,7 @@ def register( }, ) - f._check_return_type() + f._check_return_type(self.sparkSession) self.sparkSession._client.register_udtf( f.func, f.returnType, name, f.evalType, f.deterministic ) diff --git a/python/pyspark/sql/tests/arrow/test_arrow_udf_scalar.py b/python/pyspark/sql/tests/arrow/test_arrow_udf_scalar.py index 761654455870e..3096d8a0f62f4 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_udf_scalar.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_udf_scalar.py @@ -32,6 +32,7 @@ BinaryType, BooleanType, ByteType, + CharType, DecimalType, DoubleType, FloatType, @@ -43,6 +44,7 @@ StringType, StructField, StructType, + VarcharType, YearMonthIntervalType, ) from pyspark.testing.sqlutils import ReusedSQLTestCase @@ -58,6 +60,35 @@ @unittest.skipIf(not have_pyarrow, pyarrow_requirement_message) class ScalarArrowUDFTestsMixin: + def test_char_varchar_scalar_results(self): + import pyarrow as pa + + @arrow_udf(CharType(3), ArrowUDFType.SCALAR) + def scalar_char(values): + return pa.array(["a"] * len(values)) + + @arrow_udf(CharType(3), ArrowUDFType.SCALAR_ITER) + def iterator_char(batches): + for values in batches: + yield pa.array(["a"] * len(values)) + + @arrow_udf(VarcharType(3), ArrowUDFType.SCALAR) + def scalar_varchar(values): + return pa.array(["abcd"] * len(values)) + + @arrow_udf(VarcharType(3), ArrowUDFType.SCALAR_ITER) + def iterator_varchar(batches): + for values in batches: + yield pa.array(["abcd"] * len(values)) + + with self.sql_conf({"spark.sql.charVarchar.standardSemantics.enabled": "true"}): + for function in (scalar_char, iterator_char): + rows = self.spark.range(2).select(function("id")).collect() + self.assertEqual([row[0] for row in rows], ["a ", "a "]) + for function in (scalar_varchar, iterator_varchar): + with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): + self.spark.range(1).select(function("id")).collect() + @property def nondeterministic_arrow_udf(self): import numpy as np diff --git a/python/pyspark/sql/tests/arrow/test_arrow_udtf.py b/python/pyspark/sql/tests/arrow/test_arrow_udtf.py index 5f59bcbc97e66..e54bcc658c874 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_udtf.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_udtf.py @@ -19,8 +19,9 @@ from typing import Iterator, Optional from pyspark.errors import PySparkAttributeError, PySparkNotImplementedError, PythonException -from pyspark.sql.functions import arrow_udtf, lit +from pyspark.sql.functions import arrow_udtf, lit, udtf from pyspark.sql.types import ArrayType, CharType, IntegerType, Row, StructField, StructType +from pyspark.sql.udtf import AnalyzeResult from pyspark.testing import assertDataFrameEqual from pyspark.testing.sqlutils import ReusedSQLTestCase from pyspark.testing.utils import have_pyarrow, pyarrow_requirement_message @@ -52,6 +53,23 @@ def eval(self) -> Iterator["pa.Table"]: ): function() + def test_analyze_char_varchar_return_types_unsupported(self): + @udtf(returnType=None, useArrow=True) + class DynamicNestedCharUDTF: + @staticmethod + def analyze() -> AnalyzeResult: + return AnalyzeResult( + StructType([StructField("nested", ArrayType(CharType(3)))]) + ) + + def eval(self): + yield (["a"],) + + with self.assertRaisesRegex( + Exception, "Arrow-optimized Python UDTFs do not support CHAR/VARCHAR" + ): + DynamicNestedCharUDTF().collect() + def test_arrow_udtf_data_conversion_error(self): from pyspark.sql.functions import udtf diff --git a/python/pyspark/sql/tests/connect/test_parity_udtf.py b/python/pyspark/sql/tests/connect/test_parity_udtf.py index 62a91c4822aa8..42a6f1b44ed77 100644 --- a/python/pyspark/sql/tests/connect/test_parity_udtf.py +++ b/python/pyspark/sql/tests/connect/test_parity_udtf.py @@ -16,7 +16,9 @@ # import os import unittest +from unittest.mock import patch +from pyspark.sql.connect.udtf import UserDefinedTableFunction from pyspark.sql.functions import lit, udtf from pyspark.sql.tests.test_udtf import ( BaseUDTFTestsMixin, @@ -24,6 +26,7 @@ UDTFArrowTestsMixin, ) from pyspark.testing.connectutils import ReusedConnectTestCase, should_test_connect +from pyspark.util import PythonEvalType if should_test_connect: from pyspark.errors.exceptions.connect import ( @@ -49,6 +52,26 @@ def tearDownClass(cls): def test_struct_output_type_casting_row(self): self.check_struct_output_type_casting_row(PickleException) + def test_return_type_validation_rpc_is_arrow_only_and_cached(self): + class TestUDTF: + def eval(self): + yield (1,) + + regular = UserDefinedTableFunction( + TestUDTF, "value INT", evalType=PythonEvalType.SQL_TABLE_UDF + ) + arrow = UserDefinedTableFunction( + TestUDTF, "value INT", evalType=PythonEvalType.SQL_ARROW_TABLE_UDF + ) + + with patch.object(self.spark, "_parse_ddl", wraps=self.spark._parse_ddl) as parse_ddl: + regular() + regular() + self.assertEqual(parse_ddl.call_count, 0) + arrow() + arrow() + self.assertEqual(parse_ddl.call_count, 1) + def test_udtf_with_invalid_return_type(self): @udtf(returnType="int") class TestUDTF: diff --git a/python/pyspark/sql/tests/pandas/test_pandas_udf_scalar.py b/python/pyspark/sql/tests/pandas/test_pandas_udf_scalar.py index 2e0d46bedcee0..88a65fd41ed89 100644 --- a/python/pyspark/sql/tests/pandas/test_pandas_udf_scalar.py +++ b/python/pyspark/sql/tests/pandas/test_pandas_udf_scalar.py @@ -44,6 +44,7 @@ BinaryType, BooleanType, ByteType, + CharType, DateType, DecimalType, DoubleType, @@ -56,6 +57,7 @@ StructField, StructType, TimestampType, + VarcharType, VariantType, VariantVal, YearMonthIntervalType, @@ -86,6 +88,23 @@ pandas_requirement_message or pyarrow_requirement_message, ) class ScalarPandasUDFTestsMixin: + def test_char_varchar_scalar_iterator_results(self): + @pandas_udf(CharType(3), PandasUDFType.SCALAR_ITER) + def char_udf(iterator): + for series in iterator: + yield pd.Series(["a"] * len(series)) + + @pandas_udf(VarcharType(3), PandasUDFType.SCALAR_ITER) + def varchar_udf(iterator): + for series in iterator: + yield pd.Series(["abcd"] * len(series)) + + with self.sql_conf({"spark.sql.charVarchar.standardSemantics.enabled": "true"}): + rows = self.spark.range(2).select(char_udf("id")).collect() + self.assertEqual([row[0] for row in rows], ["a ", "a "]) + with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): + self.spark.range(1).select(varchar_udf("id")).collect() + @property def nondeterministic_vectorized_udf(self): import numpy as np diff --git a/python/pyspark/sql/tests/test_udf.py b/python/pyspark/sql/tests/test_udf.py index 72821153e821d..e2141aaa9d924 100644 --- a/python/pyspark/sql/tests/test_udf.py +++ b/python/pyspark/sql/tests/test_udf.py @@ -137,6 +137,28 @@ def test_char_varchar_intermediate_udf_results(self): ) self.assertEqual(result.first(), Row(c="a", v="abcd")) + def test_char_varchar_view_keeps_resolved_semantics(self): + with self.temp_view("char_varchar_udf_view"): + with self.sql_conf({"spark.sql.charVarchar.standardSemantics.enabled": "true"}): + self.spark.range(1).select( + udf(lambda _: "a", CharType(3), useArrow=False)("id").alias("c"), + udf(lambda _: "abcd", VarcharType(3), useArrow=False)("id").alias("v"), + ).createOrReplaceTempView("char_varchar_udf_view") + + with self.sql_conf( + { + "spark.sql.legacy.charVarcharAsString": "true", + "spark.sql.preserveCharVarcharTypeInfo": "false", + "spark.sql.charVarchar.standardSemantics.enabled": "false", + } + ): + self.assertEqual( + self.spark.sql("SELECT c FROM char_varchar_udf_view").first().c, + "a ", + ) + with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): + self.spark.sql("SELECT v FROM char_varchar_udf_view").collect() + def test_char_varchar_non_scalar_return_types_unsupported(self): nested_return_type = StructType( [StructField("nested", ArrayType(CharType(3)))] @@ -155,6 +177,15 @@ def test_char_varchar_non_scalar_return_types_unsupported(self): PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF, PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF, ] + stateful_and_incremental_eval_types = [ + PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_UDF, + PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF, + PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_UDF, + PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_INIT_STATE_UDF, + PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL_UDF, + PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF, + PythonEvalType.SQL_WINDOW_AGG_ARROW_INCREMENTAL_UDF, + ] for eval_type in struct_eval_types: with self.assertRaisesRegex(PySparkNotImplementedError, "Invalid return type"): @@ -162,6 +193,9 @@ def test_char_varchar_non_scalar_return_types_unsupported(self): for eval_type in aggregate_eval_types: with self.assertRaisesRegex(PySparkNotImplementedError, "Invalid return type"): UserDefinedFunction._check_return_type(VarcharType(3), eval_type) + for eval_type in stateful_and_incremental_eval_types: + with self.assertRaisesRegex(PySparkNotImplementedError, "Invalid return type"): + UserDefinedFunction._check_return_type(nested_return_type, eval_type) def test_udf_with_callable(self): data = self.spark.createDataFrame([(i, i**2) for i in range(10)], ["number", "squared"]) diff --git a/python/pyspark/sql/udf.py b/python/pyspark/sql/udf.py index c529d7fb617db..89319bdcf933b 100644 --- a/python/pyspark/sql/udf.py +++ b/python/pyspark/sql/udf.py @@ -317,7 +317,11 @@ def _conf_is_true(key: str, default: Optional[str] = None) -> bool: @staticmethod def _check_return_type(returnType: DataType, evalType: int) -> None: + class _InvalidCharVarcharArrowTypeError(TypeError): + pass + char_varchar_supported_eval_types = ( + PythonEvalType.SQL_BATCHED_UDF, PythonEvalType.SQL_ARROW_BATCHED_UDF, PythonEvalType.SQL_SCALAR_PANDAS_UDF, PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF, @@ -329,7 +333,7 @@ def check_arrow_type() -> None: if evalType not in char_varchar_supported_eval_types and _has_type( returnType, (CharType, VarcharType) ): - raise TypeError + raise _InvalidCharVarcharArrowTypeError to_arrow_type(returnType, timezone="UTC") if evalType == PythonEvalType.SQL_ARROW_BATCHED_UDF: @@ -516,6 +520,16 @@ def check_arrow_type() -> None: f"{returnType}" }, ) + elif evalType not in char_varchar_supported_eval_types and _has_type( + returnType, (CharType, VarcharType) + ): + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={ + "feature": f"Invalid return type with Python UDF eval type {evalType}: " + f"{returnType}" + }, + ) @property def returnType(self) -> DataType: diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala index 35b9052686dcf..fffd384990404 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala @@ -2539,7 +2539,8 @@ class Analyzer( } PythonUDTF( u.name, u.func, analyzeResult.schema, Some(analyzeResult.pickledAnalyzeResult), - newChildren, u.evalType, u.udfDeterministic, u.resultId, None, u.tableArguments) + newChildren, u.evalType, u.udfDeterministic, u.resultId, None, u.tableArguments, + u.applyCharVarcharChecks) } } } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala index b820172bc2b53..d55d6fb93893f 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala @@ -337,7 +337,8 @@ case class PythonUDF( // single lambda, and one more for each enclosing lambda when the UDF is lifted out of a nested // lambda (e.g. `transform(arr, i -> transform(i, x -> f(x)))` lifts `f` to depth 2). Ignored // for every non-element-wise eval type, where it stays at its default of 1. - elementwiseNestingDepth: Int = 1) + elementwiseNestingDepth: Int = 1, + applyCharVarcharChecks: Boolean = false) extends Expression with PythonFuncExpression with Unevaluable { lazy val resultAttribute: Attribute = AttributeReference(toPrettySQL(this), dataType, nullable)( @@ -495,7 +496,8 @@ case class PythonUDTF( udfDeterministic: Boolean, resultId: ExprId = NamedExpression.newExprId, pythonUDTFPartitionColumnIndexes: Option[PythonUDTFPartitionColumnIndexes] = None, - tableArguments: Option[Seq[Boolean]] = None) + tableArguments: Option[Seq[Boolean]] = None, + applyCharVarcharChecks: Boolean = false) extends UnevaluableGenerator with PythonFuncExpression { override lazy val canonicalized: Expression = { @@ -525,7 +527,8 @@ case class UnresolvedPolymorphicPythonUDTF( udfDeterministic: Boolean, resolveElementMetadata: (PythonFunction, Seq[Expression]) => PythonUDTFAnalyzeResult, resultId: ExprId = NamedExpression.newExprId, - tableArguments: Option[Seq[Boolean]] = None) + tableArguments: Option[Seq[Boolean]] = None, + applyCharVarcharChecks: Boolean = false) extends UnevaluableGenerator with PythonFuncExpression { override lazy val resolved = false diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryCompilationErrors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryCompilationErrors.scala index 28f496ab1bd8b..06d1aef0b4983 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryCompilationErrors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryCompilationErrors.scala @@ -3022,6 +3022,12 @@ private[sql] object QueryCompilationErrors extends QueryErrorsBase with Compilat messageParameters = Map("config" -> toSQLConf(config))) } + def invalidPythonArrowUDTFReturnType(dataType: DataType): SparkUnsupportedOperationException = { + new SparkUnsupportedOperationException( + errorClass = "UNSUPPORTED_FEATURE.PYTHON_ARROW_UDTF_CHAR_VARCHAR_RETURN_TYPE", + messageParameters = Map("dataType" -> toSQLType(dataType))) + } + def externalUDFWithMultipleChildrenUnsupportedError(udf: Expression): Throwable = { new AnalysisException( errorClass = "UNSUPPORTED_FEATURE.EXTERNAL_UDF_WITH_MULTIPLE_CHILDREN", diff --git a/sql/core/src/main/scala/org/apache/spark/sql/api/python/PythonSQLUtils.scala b/sql/core/src/main/scala/org/apache/spark/sql/api/python/PythonSQLUtils.scala index a2aef54f1cd7b..4e463ee02873c 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/api/python/PythonSQLUtils.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/api/python/PythonSQLUtils.scala @@ -33,6 +33,7 @@ import org.apache.spark.sql.catalyst.analysis.{FunctionRegistry, TableFunctionRe import org.apache.spark.sql.catalyst.encoders.ExpressionEncoder import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.parser.CatalystSqlParser +import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.classic.{DataFrameReader => ClassicDataFrameReader} import org.apache.spark.sql.classic.ClassicConversions._ import org.apache.spark.sql.classic.ExpressionUtils.expression @@ -128,7 +129,19 @@ private[sql] object PythonSQLUtils extends Logging { arr: Array[Byte], returnType: StructType, deserializer: ExpressionEncoder.Deserializer[Row]): Row = { - val fromJava = EvaluatePython.makeFromJava(returnType) + toJVMRow( + arr, + returnType, + deserializer, + CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get)) + } + + def toJVMRow( + arr: Array[Byte], + returnType: StructType, + deserializer: ExpressionEncoder.Deserializer[Row], + applyCharVarcharChecks: Boolean): Row = { + val fromJava = EvaluatePython.makeFromJava(returnType, applyCharVarcharChecks) val internalRow = fromJava(withInternalRowUnpickler(_.loads(arr))).asInstanceOf[InternalRow] deserializer(internalRow) diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala index 590c061f777c3..c8f25bafb1c02 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala @@ -357,7 +357,7 @@ private[sql] object ArrowConverters extends Logging { private[sql] abstract class InternalRowIterator( arrowBatchIter: Iterator[Array[Byte]], context: TaskContext) - extends Iterator[InternalRow] { + extends CloseableIterator[InternalRow] { // Keep all the resources we have opened in order, should be closed in reverse order finally. val resources = new ArrayBuffer[AutoCloseable]() protected val allocator: BufferAllocator = ArrowUtils.rootAllocator.newChildAllocator( @@ -368,11 +368,12 @@ private[sql] object ArrowConverters extends Logging { private var rowIterAndSchema = if (arrowBatchIter.hasNext) nextBatch() else (Iterator.empty, null) + private var closed = false // We will ensure schemas parsed from every batch are the same. val schema: StructType = rowIterAndSchema._2 if (context != null) context.addTaskCompletionListener[Unit] { _ => - closeAll(resources.toSeq.reverse: _*) + close() } override def hasNext: Boolean = rowIterAndSchema._1.hasNext || { @@ -385,13 +386,20 @@ private[sql] object ArrowConverters extends Logging { } rowIterAndSchema._1.hasNext } else { - closeAll(resources.toSeq.reverse: _*) + close() false } } override def next(): InternalRow = rowIterAndSchema._1.next() + override def close(): Unit = { + if (!closed) { + closed = true + closeAll(resources.toSeq.reverse: _*) + } + } + def nextBatch(): (Iterator[InternalRow], StructType) } @@ -451,7 +459,7 @@ private[sql] object ArrowConverters extends Logging { timeZoneId: String, errorOnDuplicatedFieldNames: Boolean, largeVarTypes: Boolean, - context: TaskContext): Iterator[InternalRow] = { + context: TaskContext): CloseableIterator[InternalRow] = { new InternalRowIteratorWithoutSchema( arrowBatchIter, schema, timeZoneId, errorOnDuplicatedFieldNames, largeVarTypes, context ) @@ -593,8 +601,13 @@ private[sql] object ArrowConverters extends Logging { // Project/copy it. Otherwise, the Arrow column vectors will be closed and released out. val proj = UnsafeProjection.create(checkedAttrs, attrs) - Dataset.ofRows(session, - LocalRelation(attrs, data.map(r => proj(r).copy()).toArray.toImmutableArraySeq)) + val rows = + try { + data.map(r => proj(r).copy()).toArray.toImmutableArraySeq + } finally { + data.close() + } + Dataset.ofRows(session, LocalRelation(attrs, rows)) } } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ArrowEvalPythonExec.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ArrowEvalPythonExec.scala index 71b2eb35e41fe..66982a2b5b120 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ArrowEvalPythonExec.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ArrowEvalPythonExec.scala @@ -192,7 +192,7 @@ class ArrowEvalPythonEvaluatorFactory( pythonMetrics: Map[String, SQLMetric], jobArtifactUUID: Option[String], sessionUUID: Option[String]) - extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + extends EvalPythonEvaluatorFactory(childOutput, udfs, output, outputAlreadyChecked = false) { override def evaluate( funcs: Seq[(ChainedPythonFunctions, Long)], diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/BatchEvalPythonExec.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/BatchEvalPythonExec.scala index 4c39ce98cf551..89512ba7baa87 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/BatchEvalPythonExec.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/BatchEvalPythonExec.scala @@ -77,7 +77,7 @@ class BatchEvalPythonEvaluatorFactory( jobArtifactUUID: Option[String], sessionUUID: Option[String], binaryAsBytes: Boolean) - extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + extends EvalPythonEvaluatorFactory(childOutput, udfs, output, outputAlreadyChecked = true) { override def evaluate( funcs: Seq[(ChainedPythonFunctions, Long)], @@ -106,7 +106,13 @@ class BatchEvalPythonEvaluatorFactory( StructType(udfs.map(u => StructField("", u.dataType, u.nullable))) } - val fromJava = EvaluatePython.makeFromJava(resultType) + val fromJava = if (udfs.length == 1) { + EvaluatePython.makeFromJava(resultType, udfs.head.applyCharVarcharChecks) + } else { + EvaluatePython.makeFromJava( + resultType.asInstanceOf[StructType], + udfs.map(_.applyCharVarcharChecks)) + } outputIterator.flatMap { pickedResult => val unpickledBatch = unpickle.loads(pickedResult) diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/BatchEvalPythonUDTFExec.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/BatchEvalPythonUDTFExec.scala index a2d94c226b7f8..e3005a3d673fb 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/BatchEvalPythonUDTFExec.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/BatchEvalPythonUDTFExec.scala @@ -86,7 +86,7 @@ case class BatchEvalPythonUDTFExec( // The return type of a UDTF is an array of struct. val resultType = udtf.dataType - val fromJava = EvaluatePython.makeFromJava(resultType) + val fromJava = EvaluatePython.makeFromJava(resultType, udtf.applyCharVarcharChecks) outputIterator.flatMap { pickedResult => val unpickledBatch = unpickle.loads(pickedResult) diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala index 2847ff6fd1cc7..2f9de39ad3d78 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala @@ -33,12 +33,19 @@ import org.apache.spark.sql.execution.RowToColumnConverter import org.apache.spark.sql.execution.metric.SQLMetric import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata import org.apache.spark.sql.execution.vectorized.OnHeapColumnVector -import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{DataType, StructField, StructType, UserDefinedType} import org.apache.spark.sql.types.DataType.equalsIgnoreCompatibleCollation import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} import org.apache.spark.util.Utils +private[python] object ColumnarArrowEvalPythonEvaluatorFactory { + def toPhysicalType(dataType: DataType): DataType = { + CharVarcharUtils.replaceCharVarcharWithStringForPhysicalType(dataType.transformRecursively { + case udt: UserDefinedType[_] => udt.sqlType + }) + } +} + /** * Evaluator factory for Arrow Python UDFs: ColumnarBatch in, ColumnarBatch out. * @@ -81,20 +88,20 @@ private[python] class ColumnarArrowEvalPythonEvaluatorFactory( sessionUUID: Option[String]) extends PartitionEvaluatorFactory[ColumnarBatch, ColumnarBatch] { - private val applyCharVarcharChecks = - CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get) - private val checkedOutput = if (applyCharVarcharChecks) { - childOutput ++ output.drop(childOutput.length).map { attr => + private val udfOutput = output.drop(childOutput.length) + private val checkedOutput = childOutput ++ udfOutput.zip(udfs).map { case (attr, udf) => + if (udf.applyCharVarcharChecks) { CharVarcharUtils.stringLengthCheck(attr, attr.dataType) + } else { + attr } - } else { - output } private val hasCharVarcharOutput = - applyCharVarcharChecks && - output.drop(childOutput.length).exists(attr => CharVarcharUtils.hasCharVarchar(attr.dataType)) - private val physicalOutputSchema = CharVarcharUtils - .replaceCharVarcharWithStringForPhysicalType(outputSchema) + udfOutput.zip(udfs).exists { case (attr, udf) => + udf.applyCharVarcharChecks && CharVarcharUtils.hasCharVarchar(attr.dataType) + } + private val physicalOutputSchema = ColumnarArrowEvalPythonEvaluatorFactory + .toPhysicalType(outputSchema) .asInstanceOf[StructType] override def createEvaluator() @@ -156,10 +163,7 @@ private[python] class ColumnarArrowEvalPythonEvaluatorFactory( }.toArray) val outputTypes = output.drop(childOutput.length).map { attr => - CharVarcharUtils.replaceCharVarcharWithStringForPhysicalType( - attr.dataType.transformRecursively { - case udt: UserDefinedType[_] => udt.sqlType - }) + ColumnarArrowEvalPythonEvaluatorFactory.toPhysicalType(attr.dataType) } val inputColumnIndices = resolveColumnIndices(allInputs.toSeq) diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala index ef64897f9f9d7..b386bfaba7214 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala @@ -28,24 +28,26 @@ import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata -import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{DataType, StructField, StructType} import org.apache.spark.util.Utils abstract class EvalPythonEvaluatorFactory( childOutput: Seq[Attribute], udfs: Seq[PythonUDF], - output: Seq[Attribute]) - extends PartitionEvaluatorFactory[InternalRow, InternalRow] { - - private val applyCharVarcharChecks = - CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get) - private val checkedOutput = if (applyCharVarcharChecks) { - childOutput ++ output.drop(childOutput.length).map { attr => - CharVarcharUtils.stringLengthCheck(attr, attr.dataType) - } - } else { + output: Seq[Attribute], + outputAlreadyChecked: Boolean) + extends PartitionEvaluatorFactory[InternalRow, InternalRow] { + + private val checkedOutput = if (outputAlreadyChecked) { output + } else { + childOutput ++ output.drop(childOutput.length).zip(udfs).map { case (attr, udf) => + if (udf.applyCharVarcharChecks) { + CharVarcharUtils.stringLengthCheck(attr, attr.dataType) + } else { + attr + } + } } protected def evaluate( diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala index 3eefb77d751b5..85dfde9692f0c 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala @@ -153,10 +153,38 @@ object EvaluatePython { makeFromJava(dataType, applyCharVarcharChecks) } - private def makeFromJava(dataType: DataType, applyCharVarcharChecks: Boolean): Any => Any = - TypeApiOps(dataType).flatMap(_.makeFromJava) + private[sql] def makeFromJava(dataType: DataType, applyCharVarcharChecks: Boolean): Any => Any = + TypeApiOps(dataType) + .flatMap(_.makeFromJava) .getOrElse(makeFromJavaDefault(dataType, applyCharVarcharChecks)) + private[python] def makeFromJava( + dataType: StructType, + applyCharVarcharChecks: Seq[Boolean]): Any => Any = { + require(dataType.length == applyCharVarcharChecks.length) + val fieldsFromJava = dataType.fields.zip(applyCharVarcharChecks).map { + case (field, applyChecks) => makeFromJava(field.dataType, applyChecks) + } + (obj: Any) => + nullSafeConvert(obj) { + case values if values.getClass.isArray => + val array = values.asInstanceOf[Array[_]] + if (array.length != dataType.length) { + throw new SparkIllegalArgumentException( + errorClass = "STRUCT_ARRAY_LENGTH_MISMATCH", + messageParameters = + Map("expected" -> dataType.length.toString, "actual" -> array.length.toString)) + } + val row = new GenericInternalRow(dataType.length) + var index = 0 + while (index < dataType.length) { + row(index) = fieldsFromJava(index)(array(index)) + index += 1 + } + row + } + } + private def makeFromJavaDefault( dataType: DataType, applyCharVarcharChecks: Boolean): Any => Any = dataType match { diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala index 9f53f078b9176..c39679e71ccf4 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala @@ -31,6 +31,7 @@ import org.apache.spark.sql.catalyst.analysis.UnresolvedAttribute import org.apache.spark.sql.catalyst.expressions.{Alias, Ascending, Descending, Expression, FunctionTableSubqueryArgumentExpression, NamedArgumentExpression, NullsFirst, NullsLast, PythonAggregate, PythonUDAF, PythonUDF, PythonUDTF, PythonUDTFAnalyzeResult, PythonUDTFSelectedExpression, SortOrder, TranspiledPythonUDF, UnresolvedPolymorphicPythonUDTF, UnresolvedTableArgPlanId} import org.apache.spark.sql.catalyst.parser.ParserInterface import org.apache.spark.sql.catalyst.plans.logical.{Generate, LogicalPlan, NamedParametersSupport, OneRowRelation} +import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.classic.{DataFrame, Dataset, SparkSession} import org.apache.spark.sql.classic.ClassicConversions._ import org.apache.spark.sql.classic.ColumnConversions @@ -86,7 +87,6 @@ case class UserDefinedPythonFunction( val optionInputTypes: List[List[String]] = transpiledInputTypes.asScala.map(_.asScala.toList).toList - val udfExpr = if (pythonEvalType == PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF || pythonEvalType == PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF || pythonEvalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF @@ -108,7 +108,14 @@ case class UserDefinedPythonFunction( } PythonAggregate(name, func, dataType, e, udfDeterministic, bufferStruct) } else { - PythonUDF(name, func, dataType, e, pythonEvalType, udfDeterministic) + PythonUDF( + name, + func, + dataType, + e, + pythonEvalType, + udfDeterministic, + applyCharVarcharChecks = CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get)) } // The ``_udf_param_N`` substitution below is positional, so a UDF // call site that supplied named arguments (e.g. SQL ``name => val`` @@ -208,6 +215,14 @@ case class UserDefinedPythonTableFunction( pythonEvalType: Int, udfDeterministic: Boolean) { + private def validateArrowReturnType(schema: StructType): Unit = { + if ((pythonEvalType == PythonEvalType.SQL_ARROW_TABLE_UDF || + pythonEvalType == PythonEvalType.SQL_ARROW_UDTF) && + CharVarcharUtils.hasCharVarchar(schema)) { + throw QueryCompilationErrors.invalidPythonArrowUDTFReturnType(schema) + } + } + def this( name: String, func: PythonFunction, @@ -243,8 +258,11 @@ case class UserDefinedPythonTableFunction( case _ => false } + val applyCharVarcharChecks = + CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get) val udtf = returnType match { case Some(rt) => + validateArrowReturnType(rt) PythonUDTF( name = name, func = func, @@ -253,12 +271,15 @@ case class UserDefinedPythonTableFunction( children = exprs, evalType = pythonEvalType, udfDeterministic = udfDeterministic, - tableArguments = Some(tableArgs)) + tableArguments = Some(tableArgs), + applyCharVarcharChecks = applyCharVarcharChecks) case _ => val runAnalyzeInPython = (func: PythonFunction, exprs: Seq[Expression]) => { val runner = new UserDefinedPythonTableFunctionAnalyzeRunner(name, func, exprs, tableArgs, parser) - runner.runInPython() + val analyzeResult = runner.runInPython() + validateArrowReturnType(analyzeResult.schema) + analyzeResult } UnresolvedPolymorphicPythonUDTF( name = name, @@ -267,7 +288,8 @@ case class UserDefinedPythonTableFunction( evalType = pythonEvalType, udfDeterministic = udfDeterministic, resolveElementMetadata = runAnalyzeInPython, - tableArguments = Some(tableArgs)) + tableArguments = Some(tableArgs), + applyCharVarcharChecks = applyCharVarcharChecks) } Generate( udtf, diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkDeserializer.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkDeserializer.scala index 508595606592a..91bfe5d525962 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkDeserializer.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkDeserializer.scala @@ -28,6 +28,10 @@ import org.apache.spark.internal.Logging import org.apache.spark.sql.Row import org.apache.spark.sql.api.python.PythonSQLUtils import org.apache.spark.sql.catalyst.encoders.ExpressionEncoder +import org.apache.spark.sql.catalyst.expressions.UnsafeProjection +import org.apache.spark.sql.catalyst.types.DataTypeUtils.toAttributes +import org.apache.spark.sql.catalyst.util.CharVarcharUtils +import org.apache.spark.sql.types.StructType import org.apache.spark.sql.util.ArrowUtils import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} @@ -35,8 +39,22 @@ import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, Column * A helper class to deserialize state Arrow batches from the state socket in * TransformWithStateInPySpark. */ -class TransformWithStateInPySparkDeserializer(deserializer: ExpressionEncoder.Deserializer[Row]) - extends Logging { +class TransformWithStateInPySparkDeserializer( + schema: StructType, + deserializer: ExpressionEncoder.Deserializer[Row], + applyCharVarcharChecks: Boolean) + extends Logging { + private val attrs = toAttributes(schema) + private val checkedProjection = + if (applyCharVarcharChecks && CharVarcharUtils.hasCharVarchar(schema)) { + val checkedAttrs = attrs.map { attr => + CharVarcharUtils.stringLengthCheck(attr, attr.dataType) + } + Some(UnsafeProjection.create(checkedAttrs, attrs)) + } else { + None + } + private lazy val allocator = ArrowUtils.rootAllocator.newChildAllocator( s"stdin reader for transformWithStateInPySpark state socket", 0, Long.MaxValue) @@ -45,18 +63,26 @@ class TransformWithStateInPySparkDeserializer(deserializer: ExpressionEncoder.De */ def readArrowBatches(stream: DataInputStream): Seq[Row] = { val reader = new ArrowStreamReader(stream, allocator) - val root = reader.getVectorSchemaRoot - val vectors = root.getFieldVectors.asScala.map { vector => - new ArrowColumnVector(vector) - }.toArray[ColumnVector] - val rows = ArrayBuffer[Row]() - while (reader.loadNextBatch()) { - val batch = new ColumnarBatch(vectors) - batch.setNumRows(root.getRowCount) - rows.appendAll(batch.rowIterator().asScala.map(r => deserializer(r.copy()))) + try { + val root = reader.getVectorSchemaRoot + val vectors = root.getFieldVectors.asScala + .map { vector => + new ArrowColumnVector(vector) + } + .toArray[ColumnVector] + val rows = ArrayBuffer[Row]() + while (reader.loadNextBatch()) { + val batch = new ColumnarBatch(vectors) + batch.setNumRows(root.getRowCount) + rows.appendAll(batch.rowIterator().asScala.map { row => + val copied = row.copy() + deserializer(checkedProjection.map(_(copied)).getOrElse(copied)) + }) + } + rows.toSeq + } finally { + reader.close(false) } - reader.close(false) - rows.toSeq } def readListElements(stream: DataInputStream, listStateInfo: ListStateInfo): Seq[Row] = { @@ -70,8 +96,11 @@ class TransformWithStateInPySparkDeserializer(deserializer: ExpressionEncoder.De } else { val bytes = new Array[Byte](size) stream.read(bytes, 0, size) - val newRow = PythonSQLUtils.toJVMRow(bytes, listStateInfo.schema, - listStateInfo.deserializer) + val newRow = PythonSQLUtils.toJVMRow( + bytes, + listStateInfo.schema, + listStateInfo.deserializer, + applyCharVarcharChecks) rows.append(newRow) } } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkPythonRunner.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkPythonRunner.scala index 8f3392711946c..d72818687ace4 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkPythonRunner.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkPythonRunner.scala @@ -33,6 +33,7 @@ import org.apache.spark.internal.Logging import org.apache.spark.internal.config.Python.{PYTHON_UNIX_DOMAIN_SOCKET_DIR, PYTHON_UNIX_DOMAIN_SOCKET_ENABLED} import org.apache.spark.security.SocketAuthHelper import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.execution.metric.SQLMetric import org.apache.spark.sql.execution.python.{BasicPythonArrowOutput, PythonArrowInput, PythonUDFRunner} import org.apache.spark.sql.execution.python.streaming.TransformWithStateInPySparkPythonRunner.{GroupedInType, InType} @@ -278,8 +279,10 @@ abstract class TransformWithStateInPySparkPythonBaseRunner[I]( new TransformWithStateInPySparkStateServer(stateServerSocket, processorHandle, groupingKeySchema, sqlConf.arrowTransformWithStateInPySparkMaxStateRecordsPerBatch, - batchTimestampMs, eventTimeWatermarkForEviction, - authHelper = stateServerAuthHelper)) + batchTimestampMs, + eventTimeWatermarkForEviction, + authHelper = stateServerAuthHelper, + applyCharVarcharChecks = CharVarcharUtils.shouldApplyWriteSideLengthCheck(sqlConf))) context.addTaskCompletionListener[Unit] { _ => logInfo(log"completion listener called") @@ -362,7 +365,9 @@ class TransformWithStateInPySparkPythonPreInitRunner( new TransformWithStateInPySparkStateServer(stateServerSocket, processorHandleImpl, groupingKeySchema, sqlConf.arrowTransformWithStateInPySparkMaxStateRecordsPerBatch, - authHelper = stateServerAuthHelper).run() + authHelper = stateServerAuthHelper, + applyCharVarcharChecks = CharVarcharUtils.shouldApplyWriteSideLengthCheck(sqlConf)) + .run() } catch { case e: Exception => throw new SparkException("TransformWithStateInPySpark state server " + diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServer.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServer.scala index 6d085eb8980a3..e5938d257b458 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServer.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServer.scala @@ -36,10 +36,11 @@ import org.apache.spark.SparkEnv import org.apache.spark.internal.{Logging, LogKeys} import org.apache.spark.internal.config.Python.PYTHON_UNIX_DOMAIN_SOCKET_ENABLED import org.apache.spark.security.SocketAuthHelper -import org.apache.spark.sql.{Encoders, Row} +import org.apache.spark.sql.Row import org.apache.spark.sql.api.python.PythonSQLUtils import org.apache.spark.sql.catalyst.encoders.ExpressionEncoder import org.apache.spark.sql.catalyst.parser.CatalystSqlParser +import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.execution.streaming.operators.stateful.transformwithstate.StateVariableType import org.apache.spark.sql.execution.streaming.operators.stateful.transformwithstate.statefulprocessor.{ImplicitGroupingKeyTracker, StatefulProcessorHandleImpl, StatefulProcessorHandleImplBase, StatefulProcessorHandleState} import org.apache.spark.sql.execution.streaming.state.StateMessage.{HandleState, ImplicitGroupingKeyRequest, ListStateCall, MapStateCall, StatefulProcessorCall, StateRequest, StateResponse, StateResponseWithLongTypeVal, StateResponseWithMapIterator, StateResponseWithMapKeysOrValues, StateResponseWithStringTypeVal, StateResponseWithTimer, StateVariableRequest, TimerInfo, TimerRequest, TimerStateCallCommand, TimerValueRequest, UtilsRequest, ValueStateCall} @@ -75,13 +76,38 @@ class TransformWithStateInPySparkStateServer( keyValueIteratorMapForTest: mutable.HashMap[String, Iterator[(Row, Row)]] = null, expiryTimerIterForTest: mutable.HashMap[String, Iterator[(Row, Long)]] = null, listTimerMapForTest: mutable.HashMap[String, Iterator[Long]] = null, - authHelper: SocketAuthHelper = null) - extends Runnable with Logging { + authHelper: SocketAuthHelper = null, + applyCharVarcharChecks: Boolean = false) + extends Runnable + with Logging { import PythonResponseWriterUtils._ private val keyRowDeserializer: ExpressionEncoder.Deserializer[Row] = - ExpressionEncoder(groupingKeySchema).resolveAndBind().createDeserializer() + ExpressionEncoder(physicalSchema(groupingKeySchema)).resolveAndBind().createDeserializer() + + private def deserializeRow( + bytes: Array[Byte], + schema: StructType, + deserializer: ExpressionEncoder.Deserializer[Row]): Row = { + PythonSQLUtils.toJVMRow(bytes, schema, deserializer, applyCharVarcharChecks) + } + + private def conversionSchema(schema: StructType): StructType = { + if (applyCharVarcharChecks) { + schema + } else { + physicalSchema(schema) + } + } + + private def physicalSchema(schema: StructType): StructType = { + CharVarcharUtils.replaceCharVarcharWithStringForPhysicalType(schema).asInstanceOf[StructType] + } + + private def stateEncoder(schema: StructType): ExpressionEncoder[Row] = { + ExpressionEncoder(physicalSchema(schema)).resolveAndBind() + } private var inputStream: DataInputStream = _ private var outputStream: DataOutputStream = outputStreamForTest @@ -331,7 +357,8 @@ class TransformWithStateInPySparkStateServer( case ImplicitGroupingKeyRequest.MethodCase.SETIMPLICITKEY => val keyBytes = message.getSetImplicitKey.getKey.toByteArray // The key row is serialized as a byte array, we need to convert it back to a Row - val keyRow = PythonSQLUtils.toJVMRow(keyBytes, groupingKeySchema, keyRowDeserializer) + val keyRow = + deserializeRow(keyBytes, conversionSchema(groupingKeySchema), keyRowDeserializer) ImplicitGroupingKeyTracker.setImplicitKey(keyRow) // Reset the list/map state iterators for a new grouping key. iterators = new mutable.HashMap[String, Iterator[Row]]() @@ -489,8 +516,8 @@ class TransformWithStateInPySparkStateServer( case ValueStateCall.MethodCase.VALUESTATEUPDATE => val byteArray = message.getValueStateUpdate.getValue.toByteArray // The value row is serialized as a byte array, we need to convert it back to a Row - val valueRow = PythonSQLUtils.toJVMRow(byteArray, valueStateInfo.schema, - valueStateInfo.deserializer) + val valueRow = + deserializeRow(byteArray, valueStateInfo.schema, valueStateInfo.deserializer) valueStateInfo.valueState.update(valueRow) sendResponse(0) case ValueStateCall.MethodCase.CLEAR => @@ -514,7 +541,10 @@ class TransformWithStateInPySparkStateServer( deserializer = if (deserializerForTest != null) { deserializerForTest } else { - new TransformWithStateInPySparkDeserializer(listStateInfo.deserializer) + new TransformWithStateInPySparkDeserializer( + listStateInfo.schema, + listStateInfo.deserializer, + applyCharVarcharChecks) } message.getMethodCase match { case ListStateCall.MethodCase.EXISTS => @@ -534,10 +564,7 @@ class TransformWithStateInPySparkStateServer( } else { val elements = message.getListStatePut.getValueList.asScala elements.map { e => - PythonSQLUtils.toJVMRow( - e.toByteArray, - listStateInfo.schema, - listStateInfo.deserializer) + deserializeRow(e.toByteArray, listStateInfo.schema, listStateInfo.deserializer) } } listStateInfo.listState.put(rows.toArray) @@ -556,8 +583,7 @@ class TransformWithStateInPySparkStateServer( } case ListStateCall.MethodCase.APPENDVALUE => val byteArray = message.getAppendValue.getValue.toByteArray - val newRow = - PythonSQLUtils.toJVMRow(byteArray, listStateInfo.schema, listStateInfo.deserializer) + val newRow = deserializeRow(byteArray, listStateInfo.schema, listStateInfo.deserializer) listStateInfo.listState.appendValue(newRow) sendResponse(0) case ListStateCall.MethodCase.APPENDLIST => @@ -570,10 +596,7 @@ class TransformWithStateInPySparkStateServer( } else { val elements = message.getAppendList.getValueList.asScala elements.map { e => - PythonSQLUtils.toJVMRow( - e.toByteArray, - listStateInfo.schema, - listStateInfo.deserializer) + deserializeRow(e.toByteArray, listStateInfo.schema, listStateInfo.deserializer) } } listStateInfo.listState.appendList(rows.toArray) @@ -610,8 +633,8 @@ class TransformWithStateInPySparkStateServer( } case MapStateCall.MethodCase.GETVALUE => val keyBytes = message.getGetValue.getUserKey.toByteArray - val keyRow = PythonSQLUtils.toJVMRow(keyBytes, mapStateInfo.keySchema, - mapStateInfo.keyDeserializer) + val keyRow = + deserializeRow(keyBytes, mapStateInfo.keySchema, mapStateInfo.keyDeserializer) val value = mapStateInfo.mapState.getValue(keyRow) if (value != null) { val valueBytes = PythonSQLUtils.toPyRow(value) @@ -624,8 +647,8 @@ class TransformWithStateInPySparkStateServer( } case MapStateCall.MethodCase.CONTAINSKEY => val keyBytes = message.getContainsKey.getUserKey.toByteArray - val keyRow = PythonSQLUtils.toJVMRow(keyBytes, mapStateInfo.keySchema, - mapStateInfo.keyDeserializer) + val keyRow = + deserializeRow(keyBytes, mapStateInfo.keySchema, mapStateInfo.keyDeserializer) if (mapStateInfo.mapState.containsKey(keyRow)) { sendResponse(0) } else { @@ -633,11 +656,11 @@ class TransformWithStateInPySparkStateServer( } case MapStateCall.MethodCase.UPDATEVALUE => val keyBytes = message.getUpdateValue.getUserKey.toByteArray - val keyRow = PythonSQLUtils.toJVMRow(keyBytes, mapStateInfo.keySchema, - mapStateInfo.keyDeserializer) + val keyRow = + deserializeRow(keyBytes, mapStateInfo.keySchema, mapStateInfo.keyDeserializer) val valueBytes = message.getUpdateValue.getValue.toByteArray - val valueRow = PythonSQLUtils.toJVMRow(valueBytes, mapStateInfo.valueSchema, - mapStateInfo.valueDeserializer) + val valueRow = + deserializeRow(valueBytes, mapStateInfo.valueSchema, mapStateInfo.valueDeserializer) mapStateInfo.mapState.updateValue(keyRow, valueRow) sendResponse(0) case MapStateCall.MethodCase.ITERATOR => @@ -678,8 +701,8 @@ class TransformWithStateInPySparkStateServer( } case MapStateCall.MethodCase.REMOVEKEY => val keyBytes = message.getRemoveKey.getUserKey.toByteArray - val keyRow = PythonSQLUtils.toJVMRow(keyBytes, mapStateInfo.keySchema, - mapStateInfo.keyDeserializer) + val keyRow = + deserializeRow(keyBytes, mapStateInfo.keySchema, mapStateInfo.keyDeserializer) mapStateInfo.mapState.removeKey(keyRow) sendResponse(0) case MapStateCall.MethodCase.CLEAR => @@ -696,16 +719,20 @@ class TransformWithStateInPySparkStateServer( stateType: StateVariableType.StateVariableType, ttlDurationMs: Option[Long], mapStateValueSchemaString: String = null): Unit = { - val schema = StructType.fromString(schemaString) - val expressionEncoder = ExpressionEncoder(schema).resolveAndBind() + val logicalSchema = StructType.fromString(schemaString) + val schema = conversionSchema(logicalSchema) + val expressionEncoder = stateEncoder(logicalSchema) stateType match { - case StateVariableType.ValueState => if (!valueStates.contains(stateName)) { - val state = if (ttlDurationMs.isEmpty) { - statefulProcessorHandle.getValueState[Row](stateName, Encoders.row(schema), - TTLConfig.NONE) + case StateVariableType.ValueState => + if (!valueStates.contains(stateName)) { + val state = if (ttlDurationMs.isEmpty) { + statefulProcessorHandle + .getValueState[Row](stateName, expressionEncoder, TTLConfig.NONE) } else { statefulProcessorHandle.getValueState( - stateName, Encoders.row(schema), TTLConfig(Duration.ofMillis(ttlDurationMs.get))) + stateName, + expressionEncoder, + TTLConfig(Duration.ofMillis(ttlDurationMs.get))) } valueStates.put(stateName, ValueStateInfo(state, schema, expressionEncoder.createDeserializer())) @@ -714,40 +741,61 @@ class TransformWithStateInPySparkStateServer( sendResponse(1, s"Value state $stateName already exists") } - case StateVariableType.ListState => if (!listStates.contains(stateName)) { - val state = if (ttlDurationMs.isEmpty) { - statefulProcessorHandle.getListState[Row](stateName, Encoders.row(schema), - TTLConfig.NONE) + case StateVariableType.ListState => + if (!listStates.contains(stateName)) { + val state = if (ttlDurationMs.isEmpty) { + statefulProcessorHandle + .getListState[Row](stateName, expressionEncoder, TTLConfig.NONE) + } else { + statefulProcessorHandle.getListState( + stateName, + expressionEncoder, + TTLConfig(Duration.ofMillis(ttlDurationMs.get))) + } + listStates.put( + stateName, + ListStateInfo( + state, + schema, + expressionEncoder.createDeserializer(), + expressionEncoder.createSerializer())) + sendResponse(0) } else { - statefulProcessorHandle.getListState( - stateName, Encoders.row(schema), TTLConfig(Duration.ofMillis(ttlDurationMs.get))) + sendResponse(1, s"List state $stateName already exists") } - listStates.put(stateName, - ListStateInfo(state, schema, expressionEncoder.createDeserializer(), - expressionEncoder.createSerializer())) - sendResponse(0) - } else { - sendResponse(1, s"List state $stateName already exists") - } - case StateVariableType.MapState => if (!mapStates.contains(stateName)) { - val valueSchema = StructType.fromString(mapStateValueSchemaString) - val valueExpressionEncoder = ExpressionEncoder(valueSchema).resolveAndBind() - val state = if (ttlDurationMs.isEmpty) { - statefulProcessorHandle.getMapState[Row, Row](stateName, - Encoders.row(schema), Encoders.row(valueSchema), TTLConfig.NONE) + case StateVariableType.MapState => + if (!mapStates.contains(stateName)) { + val logicalValueSchema = StructType.fromString(mapStateValueSchemaString) + val valueSchema = conversionSchema(logicalValueSchema) + val valueExpressionEncoder = stateEncoder(logicalValueSchema) + val state = if (ttlDurationMs.isEmpty) { + statefulProcessorHandle.getMapState[Row, Row]( + stateName, + expressionEncoder, + valueExpressionEncoder, + TTLConfig.NONE) + } else { + statefulProcessorHandle.getMapState[Row, Row]( + stateName, + expressionEncoder, + valueExpressionEncoder, + TTLConfig(Duration.ofMillis(ttlDurationMs.get))) + } + mapStates.put( + stateName, + MapStateInfo( + state, + schema, + valueSchema, + expressionEncoder.createDeserializer(), + expressionEncoder.createSerializer(), + valueExpressionEncoder.createDeserializer(), + valueExpressionEncoder.createSerializer())) + sendResponse(0) } else { - statefulProcessorHandle.getMapState[Row, Row](stateName, Encoders.row(schema), - Encoders.row(valueSchema), TTLConfig(Duration.ofMillis(ttlDurationMs.get))) + sendResponse(1, s"Map state $stateName already exists") } - mapStates.put(stateName, - MapStateInfo(state, schema, valueSchema, expressionEncoder.createDeserializer(), - expressionEncoder.createSerializer(), valueExpressionEncoder.createDeserializer(), - valueExpressionEncoder.createSerializer())) - sendResponse(0) - } else { - sendResponse(1, s"Map state $stateName already exists") - } } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/IntegratedUDFTestUtils.scala b/sql/core/src/test/scala/org/apache/spark/sql/IntegratedUDFTestUtils.scala index ad24ed3c23375..c4f8f9d1cc657 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/IntegratedUDFTestUtils.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/IntegratedUDFTestUtils.scala @@ -443,12 +443,31 @@ object IntegratedUDFTestUtils extends SQLHelper { children: Seq[Expression], evalType: Int, udfDeterministic: Boolean, - resultId: ExprId) - extends PythonUDF(name, func, dataType, children, evalType, udfDeterministic, resultId) { + resultId: ExprId, + elementwiseNestingDepth: Int, + applyCharVarcharChecks: Boolean) + extends PythonUDF( + name, + func, + dataType, + children, + evalType, + udfDeterministic, + resultId, + elementwiseNestingDepth, + applyCharVarcharChecks) { def this(pudf: PythonUDF) = { - this(pudf.name, pudf.func, pudf.dataType, pudf.children, - pudf.evalType, pudf.udfDeterministic, pudf.resultId) + this( + pudf.name, + pudf.func, + pudf.dataType, + pudf.children, + pudf.evalType, + pudf.udfDeterministic, + pudf.resultId, + pudf.elementwiseNestingDepth, + pudf.applyCharVarcharChecks) } override def toString: String = s"$name(${children.mkString(", ")})" diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/arrow/ArrowConvertersSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/arrow/ArrowConvertersSuite.scala index d5fcbfdaa33b6..11c057d60655c 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/arrow/ArrowConvertersSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/arrow/ArrowConvertersSuite.scala @@ -36,7 +36,7 @@ import org.apache.spark.sql.catalyst.util.DateTimeUtils import org.apache.spark.sql.classic.DataFrame import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.test.SharedSparkSession -import org.apache.spark.sql.types.{ArrayType, BinaryType, DataType, Decimal, IntegerType, NullType, StringType, StructField, StructType, TimestampLTZNanosType, TimestampNTZNanosType} +import org.apache.spark.sql.types.{ArrayType, BinaryType, DataType, Decimal, IntegerType, NullType, StringType, StructField, StructType, TimestampLTZNanosType, TimestampNTZNanosType, VarcharType} import org.apache.spark.sql.util.ArrowUtils import org.apache.spark.unsafe.types.{TimestampNanosVal, UTF8String} import org.apache.spark.util.Utils @@ -1435,6 +1435,39 @@ class ArrowConvertersSuite extends SharedSparkSession { assert(count == inputRows.length) } + test("local Arrow DataFrame conversion closes resources when VARCHAR validation fails") { + val physicalSchema = StructType(Seq(StructField("value", StringType))) + val logicalSchema = StructType(Seq(StructField("value", VarcharType(3)))) + val rows = Iterator.single(InternalRow(UTF8String.fromString("abcd"))) + val batches = ArrowConverters + .toBatchIterator( + rows, + physicalSchema, + 1, + "UTC", + errorOnDuplicatedFieldNames = true, + largeVarTypes = false, + TaskContext.empty()) + .toArray + val allocatedBefore = ArrowUtils.rootAllocator.getAllocatedMemory + + withSQLConf( + SQLConf.CHAR_VARCHAR_STANDARD_SEMANTICS.key -> "true", + SQLConf.ARROW_LOCAL_RELATION_THRESHOLD.key -> Long.MaxValue.toString) { + val error = intercept[Exception] { + ArrowConverters.toDataFrame( + batches.iterator, + logicalSchema, + spark, + "UTC", + errorOnDuplicatedFieldNames = true, + largeVarTypes = false) + } + assert(error.getMessage.contains("EXCEED_LIMIT_LENGTH")) + } + assert(ArrowUtils.rootAllocator.getAllocatedMemory === allocatedBefore) + } + test("SPARK-57159: roundtrip arrow batches with nanosecond timestamps") { withSQLConf(SQLConf.TIMESTAMP_NANOS_TYPES_ENABLED.key -> "true") { Seq[(DataType, String)]( diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala index 7ce29543221a9..82d91c78af821 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala @@ -21,8 +21,8 @@ import org.apache.spark.SparkException import org.apache.spark.sql.IntegratedUDFTestUtils import org.apache.spark.sql.execution.SparkPlan import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.test.SharedSparkSession -import org.apache.spark.sql.types.{CharType, StringType, VarcharType} +import org.apache.spark.sql.test.{ExamplePointUDT, SharedSparkSession} +import org.apache.spark.sql.types._ /** * End-to-end tests for the Arrow columnar Python UDF input path. @@ -53,6 +53,18 @@ class ArrowColumnarPythonUDFSuite extends SharedSparkSession { } } + test("CHAR/VARCHAR output normalization also unwraps UDT siblings") { + val logicalSchema = + StructType(Seq(StructField("c", CharType(3)), StructField("point", new ExamplePointUDT()))) + val expectedSchema = StructType( + Seq( + StructField("c", StringType), + StructField("point", ArrayType(DoubleType, containsNull = false)))) + + assert( + ColumnarArrowEvalPythonEvaluatorFactory.toPhysicalType(logicalSchema) === expectedSchema) + } + test("Arrow-backed source: no ColumnarToRowExec before ArrowEvalPythonExec") { assume(shouldTestPandasUDFs) withSQLConf(SQLConf.ARROW_PYSPARK_EXECUTION_ENABLED.key -> "true") { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServerSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServerSuite.scala index ea1c8a0a5d8e0..7d709cf4fd297 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServerSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServerSuite.scala @@ -38,6 +38,8 @@ import org.mockito.invocation.InvocationOnMock import org.scalatest.BeforeAndAfterEach import org.scalatest.concurrent.Eventually.{eventually, timeout} +import net.razorvine.pickle.Pickler + import org.apache.spark.SparkFunSuite import org.apache.spark.sql.{Encoder, Row} import org.apache.spark.sql.catalyst.encoders.ExpressionEncoder @@ -46,7 +48,7 @@ import org.apache.spark.sql.execution.streaming.operators.stateful.transformwith import org.apache.spark.sql.execution.streaming.state.StateMessage import org.apache.spark.sql.execution.streaming.state.StateMessage.{AppendList, AppendValue, Clear, ContainsKey, DeleteTimer, Exists, ExpiryTimerRequest, Get, GetProcessingTime, GetValue, GetWatermark, HandleState, Keys, ListStateCall, ListStateGet, ListStatePut, ListTimers, MapStateCall, ParseStringSchema, RegisterTimer, RemoveKey, SetHandleState, StateCallCommand, StatefulProcessorCall, TimerRequest, TimerStateCallCommand, TimerValueRequest, UpdateValue, UtilsRequest, Values, ValueStateCall, ValueStateUpdate} import org.apache.spark.sql.streaming.{ListState, MapState, TTLConfig, ValueState} -import org.apache.spark.sql.types.{IntegerType, StructField, StructType} +import org.apache.spark.sql.types.{IntegerType, StringType, StructField, StructType, VarcharType} import org.apache.spark.tags.SlowSQLTest @SlowSQLTest @@ -293,6 +295,64 @@ class TransformWithStateInPySparkStateServerSuite extends SparkFunSuite with Bef verify(outputStream).writeInt(0) } + test("legacy CHAR/VARCHAR policy applies to value, list, and map state updates") { + val schema = StructType(Seq(StructField("value", VarcharType(3)))) + val encoder = + ExpressionEncoder(StructType(Seq(StructField("value", StringType)))).resolveAndBind() + val deserializer = encoder.createDeserializer() + val serializer = encoder.createSerializer() + val bytes = ByteString.copyFrom(new Pickler(true, false).dumps(Array[AnyRef]("abcd"))) + val valueStateInfo = + mutable.HashMap(stateName -> ValueStateInfo(valueState, schema, deserializer)) + val listStateInfo = + mutable.HashMap(stateName -> ListStateInfo(listState, schema, deserializer, serializer)) + val mapStateInfo = mutable.HashMap( + stateName -> MapStateInfo( + mapState, + schema, + schema, + deserializer, + serializer, + deserializer, + serializer)) + val legacyStateServer = new TransformWithStateInPySparkStateServer( + serverSocket, + statefulProcessorHandle, + groupingKeySchema, + 2, + outputStreamForTest = outputStream, + valueStateMapForTest = valueStateInfo, + deserializerForTest = transformWithStateInPySparkDeserializer, + listStatesMapForTest = listStateInfo, + mapStatesMapForTest = mapStateInfo, + applyCharVarcharChecks = false) + + legacyStateServer.handleValueStateRequest( + ValueStateCall + .newBuilder() + .setStateName(stateName) + .setValueStateUpdate(ValueStateUpdate.newBuilder().setValue(bytes)) + .build()) + legacyStateServer.handleListStateRequest( + ListStateCall + .newBuilder() + .setStateName(stateName) + .setAppendValue(AppendValue.newBuilder().setValue(bytes)) + .build()) + legacyStateServer.handleMapStateRequest( + MapStateCall + .newBuilder() + .setStateName(stateName) + .setUpdateValue(UpdateValue.newBuilder().setUserKey(bytes).setValue(bytes)) + .build()) + + verify(valueState).update(argThat((row: Row) => row.getString(0) == "abcd")) + verify(listState).appendValue(argThat((row: Row) => row.getString(0) == "abcd")) + verify(mapState).updateValue( + argThat((row: Row) => row.getString(0) == "abcd"), + argThat((row: Row) => row.getString(0) == "abcd")) + } + test("list state exists") { val message = ListStateCall.newBuilder().setStateName(stateName) .setExists(Exists.newBuilder().build()).build() From af66d51d5fcb4b8ac9163ccc76b2249f8d7635e4 Mon Sep 17 00:00:00 2001 From: srielau Date: Wed, 9 Sep 2026 14:34:02 +0000 Subject: [PATCH 6/9] fix: [SPARK-59275] close Python type boundary gaps --- .../resources/error/error-conditions.json | 9 +- python/pyspark/sql/tests/arrow/test_arrow.py | 13 ++ .../sql/tests/arrow/test_arrow_udtf.py | 6 +- .../sql/tests/test_python_datasource.py | 24 ++- python/pyspark/sql/tests/test_udf.py | 49 +++++ .../test_udf_in_higher_order_function.py | 15 +- python/pyspark/sql/tests/test_udtf.py | 31 +++ python/pyspark/sql/udf.py | 24 +++ python/pyspark/sql/udtf.py | 11 +- .../sql/worker/plan_data_source_read.py | 12 +- .../sql/errors/QueryCompilationErrors.scala | 12 +- .../sql/execution/arrow/ArrowConverters.scala | 17 +- .../sql/execution/python/EvaluatePython.scala | 11 +- .../python/ExtractPythonUDFFromLambda.scala | 3 +- .../execution/python/ExtractPythonUDFs.scala | 4 +- .../python/UserDefinedPythonFunction.scala | 12 +- ...nsformWithStateInPySparkDeserializer.scala | 59 ++---- ...nsformWithStateInPySparkPythonRunner.scala | 11 +- ...ansformWithStateInPySparkStateServer.scala | 180 +++++++----------- .../python/ArrowColumnarPythonUDFSuite.scala | 8 +- .../python/EvaluatePythonSuite.scala | 29 ++- ...rmWithStateInPySparkStateServerSuite.scala | 116 ++++++----- 22 files changed, 390 insertions(+), 266 deletions(-) diff --git a/common/utils/src/main/resources/error/error-conditions.json b/common/utils/src/main/resources/error/error-conditions.json index ea3f376f2940e..073bcaf0f8fb7 100644 --- a/common/utils/src/main/resources/error/error-conditions.json +++ b/common/utils/src/main/resources/error/error-conditions.json @@ -9145,9 +9145,9 @@ "Purge table." ] }, - "PYTHON_ARROW_UDTF_CHAR_VARCHAR_RETURN_TYPE" : { + "PYTHON_STATE_CHAR_VARCHAR_SCHEMA" : { "message" : [ - "Arrow-optimized Python UDTFs do not support CHAR/VARCHAR in return type ." + "Python state does not support CHAR/VARCHAR in schema ." ] }, "PYTHON_UDF_IN_ON_CLAUSE" : { @@ -9155,6 +9155,11 @@ "Python UDF in the ON clause of a JOIN. In case of an INNER JOIN consider rewriting to a CROSS JOIN with a WHERE clause." ] }, + "PYTHON_UDTF_CHAR_VARCHAR_RETURN_TYPE" : { + "message" : [ + "Python UDTFs do not support CHAR/VARCHAR in return type ." + ] + }, "QUERY_ONLY_CORRUPT_RECORD_COLUMN" : { "message" : [ "Queries from raw JSON/CSV/XML files are disallowed when the", diff --git a/python/pyspark/sql/tests/arrow/test_arrow.py b/python/pyspark/sql/tests/arrow/test_arrow.py index b5d9c9527066d..a16d5f23996bd 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow.py +++ b/python/pyspark/sql/tests/arrow/test_arrow.py @@ -1478,6 +1478,18 @@ def test_char_varchar_explicit_schema_and_to_arrow(self): invalid, StructType([StructField("c", VarcharType(3))]) ).collect() + with self.sql_conf({"spark.sql.execution.arrow.localRelationThreshold": "0"}): + rdd_df = self.spark.createDataFrame( + pa.table({"c": ["a"]}), + StructType([StructField("c", CharType(3))]), + ) + self.assertEqual(rdd_df.first(), Row(c="a ")) + with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): + self.spark.createDataFrame( + pa.table({"v": ["abcd"]}), + StructType([StructField("v", VarcharType(3))]), + ).collect() + legacy_schema = StructType( [ StructField("c", CharType(3)), @@ -1496,6 +1508,7 @@ def test_char_varchar_explicit_schema_and_to_arrow(self): df = self.spark.createDataFrame( pa.table({"c": ["a"], "v": ["abcd"]}), legacy_schema ) + self.assertEqual(df.schema, StructType().add("c", "string").add("v", "string")) self.assertEqual(df.first(), Row(c="a", v="abcd")) def test_createDataFrame_pandas_duplicate_field_names(self): diff --git a/python/pyspark/sql/tests/arrow/test_arrow_udtf.py b/python/pyspark/sql/tests/arrow/test_arrow_udtf.py index e54bcc658c874..e7aab5c2309cc 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_udtf.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_udtf.py @@ -49,7 +49,7 @@ def eval(self) -> Iterator["pa.Table"]: for function in (DirectCharUDTF, NestedCharUDTF): with self.assertRaisesRegex( - PySparkNotImplementedError, "Invalid return type with Arrow UDTFs" + PySparkNotImplementedError, "CHAR/VARCHAR return type in Python UDTFs" ): function() @@ -65,9 +65,7 @@ def analyze() -> AnalyzeResult: def eval(self): yield (["a"],) - with self.assertRaisesRegex( - Exception, "Arrow-optimized Python UDTFs do not support CHAR/VARCHAR" - ): + with self.assertRaisesRegex(Exception, "Python UDTFs do not support CHAR/VARCHAR"): DynamicNestedCharUDTF().collect() def test_arrow_udtf_data_conversion_error(self): diff --git a/python/pyspark/sql/tests/test_python_datasource.py b/python/pyspark/sql/tests/test_python_datasource.py index 5bb8e9df1e3b2..cc900455eade2 100644 --- a/python/pyspark/sql/tests/test_python_datasource.py +++ b/python/pyspark/sql/tests/test_python_datasource.py @@ -53,7 +53,16 @@ ) from pyspark.sql.functions import spark_partition_id from pyspark.sql.session import SparkSession -from pyspark.sql.types import DecimalType, IntegerType, Row, StructField, StructType, VariantVal +from pyspark.sql.types import ( + ArrayType, + CharType, + DecimalType, + IntegerType, + Row, + StructField, + StructType, + VariantVal, +) from pyspark.testing import assertDataFrameEqual from pyspark.testing.sqlutils import ( SPARK_HOME, @@ -193,6 +202,19 @@ def test_data_source_read_output_row(self): df = self.spark.read.format("test").load() assertDataFrameEqual(df, [Row(0, 1)]) + def test_data_source_char_varchar_return_type_is_unsupported(self): + schema = StructType([StructField("value", ArrayType(CharType(2)))]) + self.register_data_source( + read_func=lambda schema, partition: iter([(["a"],)]), + output=schema, + ) + with self.sql_conf({"spark.sql.charVarchar.standardSemantics.enabled": "true"}): + with self.assertRaisesRegex( + PythonException, + "CHAR/VARCHAR return types in Python DataSource", + ): + self.spark.read.format("test").load().collect() + def test_data_source_read_output_named_row(self): self.register_data_source( read_func=lambda schema, partition: iter([Row(j=1, i=0), Row(i=1, j=2)]) diff --git a/python/pyspark/sql/tests/test_udf.py b/python/pyspark/sql/tests/test_udf.py index e2141aaa9d924..fd91926c4f935 100644 --- a/python/pyspark/sql/tests/test_udf.py +++ b/python/pyspark/sql/tests/test_udf.py @@ -50,6 +50,7 @@ StructField, StructType, TimestampNTZType, + UserDefinedType, VarcharType, VariantType, VariantVal, @@ -159,6 +160,25 @@ def test_char_varchar_view_keeps_resolved_semantics(self): with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): self.spark.sql("SELECT v FROM char_varchar_udf_view").collect() + def test_char_varchar_mixed_captured_policies_in_one_batch(self): + char_udf = udf(lambda _: "a", CharType(3), useArrow=False) + varchar_udf = udf(lambda _: "abcd", VarcharType(3), useArrow=False) + with self.sql_conf({"spark.sql.charVarchar.standardSemantics.enabled": "true"}): + checked_char = char_udf("id").alias("c") + with self.sql_conf( + { + "spark.sql.legacy.charVarcharAsString": "true", + "spark.sql.preserveCharVarcharTypeInfo": "false", + "spark.sql.charVarchar.standardSemantics.enabled": "false", + } + ): + unchecked_varchar = varchar_udf("id").alias("v") + + self.assertEqual( + self.spark.range(1).select(checked_char, unchecked_varchar).first(), + Row(c="a ", v="abcd"), + ) + def test_char_varchar_non_scalar_return_types_unsupported(self): nested_return_type = StructType( [StructField("nested", ArrayType(CharType(3)))] @@ -197,6 +217,35 @@ def test_char_varchar_non_scalar_return_types_unsupported(self): with self.assertRaisesRegex(PySparkNotImplementedError, "Invalid return type"): UserDefinedFunction._check_return_type(nested_return_type, eval_type) + def test_char_varchar_inside_udt_return_type_is_unsupported(self): + class CharVarcharUDT(UserDefinedType): + @classmethod + def sqlType(cls): + return StructType([StructField("value", ArrayType(CharType(2)))]) + + @classmethod + def module(cls): + return __name__ + + @classmethod + def scalaUDT(cls): + return "" + + def serialize(self, obj): + return obj + + def deserialize(self, datum): + return datum + + with self.assertRaisesRegex( + PySparkNotImplementedError, + "CHAR/VARCHAR inside Python UDF UDT return type", + ): + UserDefinedFunction._check_return_type( + CharVarcharUDT(), + PythonEvalType.SQL_ARROW_BATCHED_UDF, + ) + def test_udf_with_callable(self): data = self.spark.createDataFrame([(i, i**2) for i in range(10)], ["number", "squared"]) diff --git a/python/pyspark/sql/tests/test_udf_in_higher_order_function.py b/python/pyspark/sql/tests/test_udf_in_higher_order_function.py index b3479f4a26609..5c8c971c4df29 100644 --- a/python/pyspark/sql/tests/test_udf_in_higher_order_function.py +++ b/python/pyspark/sql/tests/test_udf_in_higher_order_function.py @@ -20,7 +20,7 @@ from pyspark.errors import AnalysisException from pyspark.sql import functions as sf from pyspark.sql.functions import udf -from pyspark.sql.types import ArrayType, DoubleType, IntegerType, StringType +from pyspark.sql.types import ArrayType, CharType, DoubleType, IntegerType, StringType, VarcharType from pyspark.testing.sqlutils import ReusedSQLTestCase from pyspark.testing.utils import ( assertDataFrameEqual, @@ -52,6 +52,19 @@ def test_transform(self): df.select(sf.transform("values", lambda x: x + 1).alias("r")), ) + def test_transform_char_varchar_results(self): + df = self.spark.createDataFrame([([1, 2],)], "values array") + char_udf = udf(lambda _: "a", CharType(3)) + varchar_udf = udf(lambda _: "abcd", VarcharType(3)) + + with self.sql_conf({"spark.sql.charVarchar.standardSemantics.enabled": "true"}): + assertDataFrameEqual( + df.select(sf.transform("values", lambda x: char_udf(x)).alias("r")), + [(["a ", "a "],)], + ) + with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): + df.select(sf.transform("values", lambda x: varchar_udf(x))).collect() + def test_transform_null_array_and_null_elements(self): # A null array must stay null, and a null *element* must reach the UDF as None. df = self.spark.createDataFrame([([1, None, 3],), (None,), ([],)], "values array") diff --git a/python/pyspark/sql/tests/test_udtf.py b/python/pyspark/sql/tests/test_udtf.py index 10641b0769a0e..f5f10934946e5 100644 --- a/python/pyspark/sql/tests/test_udtf.py +++ b/python/pyspark/sql/tests/test_udtf.py @@ -30,6 +30,7 @@ AnalysisException, IllegalArgumentException, PySparkAttributeError, + PySparkNotImplementedError, PySparkPicklingError, PySparkTypeError, PythonException, @@ -53,6 +54,7 @@ from pyspark.sql.types import ( ArrayType, BooleanType, + CharType, DataType, IntegerType, LongType, @@ -77,6 +79,35 @@ class BaseUDTFTestsMixin: + def test_char_varchar_return_types_unsupported(self): + nested_type = StructType([StructField("nested", ArrayType(CharType(3)))]) + + @udtf(returnType=nested_type, useArrow=False) + class NestedCharUDTF: + def eval(self): + yield (["a"],) + + with self.assertRaisesRegex( + PySparkNotImplementedError, + "CHAR/VARCHAR return type in Python UDTFs", + ): + NestedCharUDTF() + + def test_analyze_char_varchar_return_types_unsupported(self): + @udtf(returnType=None, useArrow=False) + class DynamicNestedCharUDTF: + @staticmethod + def analyze() -> AnalyzeResult: + return AnalyzeResult( + StructType([StructField("nested", ArrayType(CharType(3)))]) + ) + + def eval(self): + yield (["a"],) + + with self.assertRaisesRegex(Exception, "Python UDTFs do not support CHAR/VARCHAR"): + DynamicNestedCharUDTF().collect() + def test_simple_udtf(self): class TestUDTF: def eval(self): diff --git a/python/pyspark/sql/udf.py b/python/pyspark/sql/udf.py index 89319bdcf933b..49308c97298d0 100644 --- a/python/pyspark/sql/udf.py +++ b/python/pyspark/sql/udf.py @@ -29,10 +29,13 @@ from pyspark.sql.pandas.types import to_arrow_type from pyspark.sql.pandas.utils import require_minimum_pandas_version, require_minimum_pyarrow_version from pyspark.sql.types import ( + ArrayType, CharType, DataType, + MapType, StringType, StructType, + UserDefinedType, VarcharType, _has_type, _parse_datatype_string, @@ -320,6 +323,27 @@ def _check_return_type(returnType: DataType, evalType: int) -> None: class _InvalidCharVarcharArrowTypeError(TypeError): pass + def has_char_varchar_in_udt(data_type: DataType) -> bool: + if isinstance(data_type, UserDefinedType): + return _has_type(data_type.sqlType(), (CharType, VarcharType)) + if isinstance(data_type, StructType): + return any(has_char_varchar_in_udt(f.dataType) for f in data_type.fields) + if isinstance(data_type, ArrayType): + return has_char_varchar_in_udt(data_type.elementType) + if isinstance(data_type, MapType): + return has_char_varchar_in_udt( + data_type.keyType + ) or has_char_varchar_in_udt(data_type.valueType) + return False + + if has_char_varchar_in_udt(returnType): + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={ + "feature": f"CHAR/VARCHAR inside Python UDF UDT return type: {returnType}" + }, + ) + char_varchar_supported_eval_types = ( PythonEvalType.SQL_BATCHED_UDF, PythonEvalType.SQL_ARROW_BATCHED_UDF, diff --git a/python/pyspark/sql/udtf.py b/python/pyspark/sql/udtf.py index 0850df082046d..b4035d7bdd612 100644 --- a/python/pyspark/sql/udtf.py +++ b/python/pyspark/sql/udtf.py @@ -325,15 +325,12 @@ def _validate_udtf_handler(cls: Any, returnType: Optional[Union[StructType, str] ) -def _check_arrow_udtf_return_type(return_type: DataType, eval_type: int) -> None: - if eval_type in ( - PythonEvalType.SQL_ARROW_TABLE_UDF, - PythonEvalType.SQL_ARROW_UDTF, - ) and _has_type(return_type, (CharType, VarcharType)): +def _check_udtf_return_type(return_type: DataType) -> None: + if _has_type(return_type, (CharType, VarcharType)): raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", messageParameters={ - "feature": f"Invalid return type with Arrow UDTFs: {return_type}" + "feature": f"CHAR/VARCHAR return type in Python UDTFs: {return_type}" }, ) @@ -390,7 +387,7 @@ def returnType(self) -> Optional[StructType]: "return_type": f"{parsed}", }, ) - _check_arrow_udtf_return_type(parsed, self.evalType) + _check_udtf_return_type(parsed) self._returnType_placeholder = parsed return self._returnType_placeholder diff --git a/python/pyspark/sql/worker/plan_data_source_read.py b/python/pyspark/sql/worker/plan_data_source_read.py index e2e66bdc0bd02..5014b7c6750f3 100644 --- a/python/pyspark/sql/worker/plan_data_source_read.py +++ b/python/pyspark/sql/worker/plan_data_source_read.py @@ -21,7 +21,7 @@ import pyarrow as pa -from pyspark.errors import PySparkAssertionError, PySparkRuntimeError +from pyspark.errors import PySparkAssertionError, PySparkNotImplementedError, PySparkRuntimeError from pyspark.logger.worker_io import capture_outputs from pyspark.serializers import ( read_bool, @@ -40,7 +40,10 @@ from pyspark.sql.pandas.types import to_arrow_schema from pyspark.sql.types import ( BinaryType, + CharType, StructType, + VarcharType, + _has_type, _parse_datatype_json_string, ) from pyspark.sql.worker.utils import check_pushdown_not_disabled, worker_run @@ -64,6 +67,13 @@ def records_to_arrow_batches( of pyarrow record batches. For each Python tuple, check the types of each field and append it to the records batch. """ + if _has_type(return_type, (CharType, VarcharType)): + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={ + "feature": f"CHAR/VARCHAR return types in Python DataSource: {return_type}" + }, + ) pa_schema = to_arrow_schema(return_type, timezone="UTC") column_names = return_type.fieldNames() diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryCompilationErrors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryCompilationErrors.scala index 06d1aef0b4983..e972a3253603c 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryCompilationErrors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryCompilationErrors.scala @@ -3022,12 +3022,20 @@ private[sql] object QueryCompilationErrors extends QueryErrorsBase with Compilat messageParameters = Map("config" -> toSQLConf(config))) } - def invalidPythonArrowUDTFReturnType(dataType: DataType): SparkUnsupportedOperationException = { + def invalidPythonUDTFReturnType(dataType: DataType): SparkUnsupportedOperationException = { new SparkUnsupportedOperationException( - errorClass = "UNSUPPORTED_FEATURE.PYTHON_ARROW_UDTF_CHAR_VARCHAR_RETURN_TYPE", + errorClass = "UNSUPPORTED_FEATURE.PYTHON_UDTF_CHAR_VARCHAR_RETURN_TYPE", messageParameters = Map("dataType" -> toSQLType(dataType))) } + def invalidPythonStateSchema( + dataType: DataType, + schemaKind: String): SparkUnsupportedOperationException = { + new SparkUnsupportedOperationException( + errorClass = "UNSUPPORTED_FEATURE.PYTHON_STATE_CHAR_VARCHAR_SCHEMA", + messageParameters = Map("dataType" -> toSQLType(dataType), "schemaKind" -> schemaKind)) + } + def externalUDFWithMultipleChildrenUnsupportedError(udf: Expression): Throwable = { new AnalysisException( errorClass = "UNSUPPORTED_FEATURE.EXTERNAL_UDF_WITH_MULTIPLE_CHILDREN", diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala index c8f25bafb1c02..b1434e57cd018 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/arrow/ArrowConverters.scala @@ -556,15 +556,20 @@ private[sql] object ArrowConverters extends Logging { timeZoneId: String, errorOnDuplicatedFieldNames: Boolean, largeVarTypes: Boolean): DataFrame = { - val attrs = toAttributes(schema) + val physicalSchema = + CharVarcharUtils.replaceCharVarcharWithStringForPhysicalType(schema).asInstanceOf[StructType] + val attrs = toAttributes(physicalSchema) val applyCharVarcharChecks = CharVarcharUtils.hasCharVarchar(schema) && CharVarcharUtils.shouldApplyWriteSideLengthCheck(session.sessionState.conf) val checkedAttrs = if (applyCharVarcharChecks) { - attrs.map(attr => CharVarcharUtils.stringLengthCheck(attr, attr.dataType)) + attrs.zip(schema.fields).map { case (attr, field) => + CharVarcharUtils.stringLengthCheck(attr, field.dataType) + } } else { attrs } + val outputSchema = if (applyCharVarcharChecks) schema else physicalSchema val batchesInDriver = arrowBatches.toArray val shouldUseRDD = session.sessionState.conf .arrowLocalRelationThreshold < batchesInDriver.map(_.length.toLong).sum @@ -576,7 +581,7 @@ private[sql] object ArrowConverters extends Logging { .mapPartitions { batchesInExecutors => val rows = ArrowConverters.fromBatchIterator( batchesInExecutors, - schema, + physicalSchema, timeZoneId, errorOnDuplicatedFieldNames, largeVarTypes, @@ -588,12 +593,12 @@ private[sql] object ArrowConverters extends Logging { rows } } - session.internalCreateDataFrame(rdd.setName("arrow"), schema) + session.internalCreateDataFrame(rdd.setName("arrow"), outputSchema) } else { logDebug("Using LocalRelation in createDataFrame with Arrow optimization.") val data = ArrowConverters.fromBatchIterator( batchesInDriver.iterator, - schema, + physicalSchema, timeZoneId, errorOnDuplicatedFieldNames, largeVarTypes, @@ -607,7 +612,7 @@ private[sql] object ArrowConverters extends Logging { } finally { data.close() } - Dataset.ofRows(session, LocalRelation(attrs, rows)) + Dataset.ofRows(session, LocalRelation(toAttributes(outputSchema), rows)) } } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala index 85dfde9692f0c..909f65c101ee0 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala @@ -30,7 +30,7 @@ import org.apache.spark.rdd.RDD import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.types.ops.TypeApiOps -import org.apache.spark.sql.catalyst.util.{ArrayBasedMapData, ArrayData, CharVarcharCodegenUtils, CharVarcharUtils, GenericArrayData, MapData, STUtils} +import org.apache.spark.sql.catalyst.util.{ArrayBasedMapBuilder, ArrayData, CharVarcharCodegenUtils, CharVarcharUtils, GenericArrayData, MapData, STUtils} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ import org.apache.spark.unsafe.types.{BinaryView, UTF8String, VariantVal} @@ -282,10 +282,11 @@ object EvaluatePython { (obj: Any) => nullSafeConvert(obj) { case javaMap: java.util.Map[_, _] => - ArrayBasedMapData( - javaMap, - (key: Any) => keyFromJava(key), - (value: Any) => valueFromJava(value)) + val builder = new ArrayBasedMapBuilder(keyType, valueType) + javaMap.asScala.foreach { case (key, value) => + builder.put(keyFromJava(key), valueFromJava(value)) + } + builder.build() } case StructType(fields) => diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFFromLambda.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFFromLambda.scala index dfaa4732ca9d7..151f0597ef8f2 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFFromLambda.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFFromLambda.scala @@ -441,7 +441,8 @@ object ExtractPythonUDFFromLambda extends Rule[LogicalPlan] { // `PythonUDF.liftedElementwiseEvalType`. PythonUDF.liftedElementwiseEvalType(udf.evalType), udf.udfDeterministic, - elementwiseNestingDepth = newDepth) + elementwiseNestingDepth = newDepth, + applyCharVarcharChecks = udf.applyCharVarcharChecks) val signature: Expression = if (udf.udfDeterministic) lifted.canonicalized else lifted val ordinal = ordinalBySignature.getOrElseUpdate(signature, { val o = distinctLifted.length diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFs.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFs.scala index 2a0e68da67581..12aa9674554f6 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFs.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractPythonUDFs.scala @@ -189,6 +189,8 @@ object ExtractPythonUDFs extends Rule[LogicalPlan] with Logging { * - PythonUDF (baz) * - if the eval types of the UDF expressions in the chain differ, return false. * - if a UDF has more than one child, e.g. foo(bar(), baz()), return false + * - if a child UDF has a CHAR/VARCHAR result whose captured policy requires assignment checks, + * return false so the checked JVM conversion boundary is preserved. * If we return false here, the expectation is that the recursive calls of * collectEvaluableUDFsFromExpressions will then visit the children and extract them first to * separate nodes. @@ -201,7 +203,7 @@ object ExtractPythonUDFs extends Rule[LogicalPlan] with Logging { case Seq(child: PythonUDF) => correctEvalType(e, pythonUDFArrowFallbackOnUDT) == correctEvalType(child, pythonUDFArrowFallbackOnUDT) && - !(CharVarcharUtils.shouldApplyWriteSideLengthCheck(conf) && + !(child.applyCharVarcharChecks && CharVarcharUtils.hasCharVarchar(child.dataType)) && shouldExtractUDFExpressionTree(child, pythonUDFArrowFallbackOnUDT) // Python UDF can't be evaluated directly in JVM diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala index c39679e71ccf4..a336136559b0c 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala @@ -215,11 +215,9 @@ case class UserDefinedPythonTableFunction( pythonEvalType: Int, udfDeterministic: Boolean) { - private def validateArrowReturnType(schema: StructType): Unit = { - if ((pythonEvalType == PythonEvalType.SQL_ARROW_TABLE_UDF || - pythonEvalType == PythonEvalType.SQL_ARROW_UDTF) && - CharVarcharUtils.hasCharVarchar(schema)) { - throw QueryCompilationErrors.invalidPythonArrowUDTFReturnType(schema) + private def validateReturnType(schema: StructType): Unit = { + if (CharVarcharUtils.hasCharVarchar(schema)) { + throw QueryCompilationErrors.invalidPythonUDTFReturnType(schema) } } @@ -262,7 +260,7 @@ case class UserDefinedPythonTableFunction( CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get) val udtf = returnType match { case Some(rt) => - validateArrowReturnType(rt) + validateReturnType(rt) PythonUDTF( name = name, func = func, @@ -278,7 +276,7 @@ case class UserDefinedPythonTableFunction( val runner = new UserDefinedPythonTableFunctionAnalyzeRunner(name, func, exprs, tableArgs, parser) val analyzeResult = runner.runInPython() - validateArrowReturnType(analyzeResult.schema) + validateReturnType(analyzeResult.schema) analyzeResult } UnresolvedPolymorphicPythonUDTF( diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkDeserializer.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkDeserializer.scala index 91bfe5d525962..508595606592a 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkDeserializer.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkDeserializer.scala @@ -28,10 +28,6 @@ import org.apache.spark.internal.Logging import org.apache.spark.sql.Row import org.apache.spark.sql.api.python.PythonSQLUtils import org.apache.spark.sql.catalyst.encoders.ExpressionEncoder -import org.apache.spark.sql.catalyst.expressions.UnsafeProjection -import org.apache.spark.sql.catalyst.types.DataTypeUtils.toAttributes -import org.apache.spark.sql.catalyst.util.CharVarcharUtils -import org.apache.spark.sql.types.StructType import org.apache.spark.sql.util.ArrowUtils import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} @@ -39,22 +35,8 @@ import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, Column * A helper class to deserialize state Arrow batches from the state socket in * TransformWithStateInPySpark. */ -class TransformWithStateInPySparkDeserializer( - schema: StructType, - deserializer: ExpressionEncoder.Deserializer[Row], - applyCharVarcharChecks: Boolean) - extends Logging { - private val attrs = toAttributes(schema) - private val checkedProjection = - if (applyCharVarcharChecks && CharVarcharUtils.hasCharVarchar(schema)) { - val checkedAttrs = attrs.map { attr => - CharVarcharUtils.stringLengthCheck(attr, attr.dataType) - } - Some(UnsafeProjection.create(checkedAttrs, attrs)) - } else { - None - } - +class TransformWithStateInPySparkDeserializer(deserializer: ExpressionEncoder.Deserializer[Row]) + extends Logging { private lazy val allocator = ArrowUtils.rootAllocator.newChildAllocator( s"stdin reader for transformWithStateInPySpark state socket", 0, Long.MaxValue) @@ -63,26 +45,18 @@ class TransformWithStateInPySparkDeserializer( */ def readArrowBatches(stream: DataInputStream): Seq[Row] = { val reader = new ArrowStreamReader(stream, allocator) - try { - val root = reader.getVectorSchemaRoot - val vectors = root.getFieldVectors.asScala - .map { vector => - new ArrowColumnVector(vector) - } - .toArray[ColumnVector] - val rows = ArrayBuffer[Row]() - while (reader.loadNextBatch()) { - val batch = new ColumnarBatch(vectors) - batch.setNumRows(root.getRowCount) - rows.appendAll(batch.rowIterator().asScala.map { row => - val copied = row.copy() - deserializer(checkedProjection.map(_(copied)).getOrElse(copied)) - }) - } - rows.toSeq - } finally { - reader.close(false) + val root = reader.getVectorSchemaRoot + val vectors = root.getFieldVectors.asScala.map { vector => + new ArrowColumnVector(vector) + }.toArray[ColumnVector] + val rows = ArrayBuffer[Row]() + while (reader.loadNextBatch()) { + val batch = new ColumnarBatch(vectors) + batch.setNumRows(root.getRowCount) + rows.appendAll(batch.rowIterator().asScala.map(r => deserializer(r.copy()))) } + reader.close(false) + rows.toSeq } def readListElements(stream: DataInputStream, listStateInfo: ListStateInfo): Seq[Row] = { @@ -96,11 +70,8 @@ class TransformWithStateInPySparkDeserializer( } else { val bytes = new Array[Byte](size) stream.read(bytes, 0, size) - val newRow = PythonSQLUtils.toJVMRow( - bytes, - listStateInfo.schema, - listStateInfo.deserializer, - applyCharVarcharChecks) + val newRow = PythonSQLUtils.toJVMRow(bytes, listStateInfo.schema, + listStateInfo.deserializer) rows.append(newRow) } } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkPythonRunner.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkPythonRunner.scala index d72818687ace4..8f3392711946c 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkPythonRunner.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkPythonRunner.scala @@ -33,7 +33,6 @@ import org.apache.spark.internal.Logging import org.apache.spark.internal.config.Python.{PYTHON_UNIX_DOMAIN_SOCKET_DIR, PYTHON_UNIX_DOMAIN_SOCKET_ENABLED} import org.apache.spark.security.SocketAuthHelper import org.apache.spark.sql.catalyst.InternalRow -import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.execution.metric.SQLMetric import org.apache.spark.sql.execution.python.{BasicPythonArrowOutput, PythonArrowInput, PythonUDFRunner} import org.apache.spark.sql.execution.python.streaming.TransformWithStateInPySparkPythonRunner.{GroupedInType, InType} @@ -279,10 +278,8 @@ abstract class TransformWithStateInPySparkPythonBaseRunner[I]( new TransformWithStateInPySparkStateServer(stateServerSocket, processorHandle, groupingKeySchema, sqlConf.arrowTransformWithStateInPySparkMaxStateRecordsPerBatch, - batchTimestampMs, - eventTimeWatermarkForEviction, - authHelper = stateServerAuthHelper, - applyCharVarcharChecks = CharVarcharUtils.shouldApplyWriteSideLengthCheck(sqlConf))) + batchTimestampMs, eventTimeWatermarkForEviction, + authHelper = stateServerAuthHelper)) context.addTaskCompletionListener[Unit] { _ => logInfo(log"completion listener called") @@ -365,9 +362,7 @@ class TransformWithStateInPySparkPythonPreInitRunner( new TransformWithStateInPySparkStateServer(stateServerSocket, processorHandleImpl, groupingKeySchema, sqlConf.arrowTransformWithStateInPySparkMaxStateRecordsPerBatch, - authHelper = stateServerAuthHelper, - applyCharVarcharChecks = CharVarcharUtils.shouldApplyWriteSideLengthCheck(sqlConf)) - .run() + authHelper = stateServerAuthHelper).run() } catch { case e: Exception => throw new SparkException("TransformWithStateInPySpark state server " + diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServer.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServer.scala index e5938d257b458..c5cdf142a1a9f 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServer.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServer.scala @@ -36,11 +36,12 @@ import org.apache.spark.SparkEnv import org.apache.spark.internal.{Logging, LogKeys} import org.apache.spark.internal.config.Python.PYTHON_UNIX_DOMAIN_SOCKET_ENABLED import org.apache.spark.security.SocketAuthHelper -import org.apache.spark.sql.Row +import org.apache.spark.sql.{Encoders, Row} import org.apache.spark.sql.api.python.PythonSQLUtils import org.apache.spark.sql.catalyst.encoders.ExpressionEncoder import org.apache.spark.sql.catalyst.parser.CatalystSqlParser import org.apache.spark.sql.catalyst.util.CharVarcharUtils +import org.apache.spark.sql.errors.QueryCompilationErrors import org.apache.spark.sql.execution.streaming.operators.stateful.transformwithstate.StateVariableType import org.apache.spark.sql.execution.streaming.operators.stateful.transformwithstate.statefulprocessor.{ImplicitGroupingKeyTracker, StatefulProcessorHandleImpl, StatefulProcessorHandleImplBase, StatefulProcessorHandleState} import org.apache.spark.sql.execution.streaming.state.StateMessage.{HandleState, ImplicitGroupingKeyRequest, ListStateCall, MapStateCall, StatefulProcessorCall, StateRequest, StateResponse, StateResponseWithLongTypeVal, StateResponseWithMapIterator, StateResponseWithMapKeysOrValues, StateResponseWithStringTypeVal, StateResponseWithTimer, StateVariableRequest, TimerInfo, TimerRequest, TimerStateCallCommand, TimerValueRequest, UtilsRequest, ValueStateCall} @@ -76,38 +77,21 @@ class TransformWithStateInPySparkStateServer( keyValueIteratorMapForTest: mutable.HashMap[String, Iterator[(Row, Row)]] = null, expiryTimerIterForTest: mutable.HashMap[String, Iterator[(Row, Long)]] = null, listTimerMapForTest: mutable.HashMap[String, Iterator[Long]] = null, - authHelper: SocketAuthHelper = null, - applyCharVarcharChecks: Boolean = false) - extends Runnable - with Logging { + authHelper: SocketAuthHelper = null) + extends Runnable with Logging { import PythonResponseWriterUtils._ - private val keyRowDeserializer: ExpressionEncoder.Deserializer[Row] = - ExpressionEncoder(physicalSchema(groupingKeySchema)).resolveAndBind().createDeserializer() - - private def deserializeRow( - bytes: Array[Byte], - schema: StructType, - deserializer: ExpressionEncoder.Deserializer[Row]): Row = { - PythonSQLUtils.toJVMRow(bytes, schema, deserializer, applyCharVarcharChecks) - } - - private def conversionSchema(schema: StructType): StructType = { - if (applyCharVarcharChecks) { - schema - } else { - physicalSchema(schema) + private def validateStateSchema(schema: StructType, schemaKind: String): Unit = { + if (CharVarcharUtils.hasCharVarchar(schema)) { + throw QueryCompilationErrors.invalidPythonStateSchema(schema, schemaKind) } } - private def physicalSchema(schema: StructType): StructType = { - CharVarcharUtils.replaceCharVarcharWithStringForPhysicalType(schema).asInstanceOf[StructType] - } + validateStateSchema(groupingKeySchema, "grouping key") - private def stateEncoder(schema: StructType): ExpressionEncoder[Row] = { - ExpressionEncoder(physicalSchema(schema)).resolveAndBind() - } + private val keyRowDeserializer: ExpressionEncoder.Deserializer[Row] = + ExpressionEncoder(groupingKeySchema).resolveAndBind().createDeserializer() private var inputStream: DataInputStream = _ private var outputStream: DataOutputStream = outputStreamForTest @@ -357,8 +341,7 @@ class TransformWithStateInPySparkStateServer( case ImplicitGroupingKeyRequest.MethodCase.SETIMPLICITKEY => val keyBytes = message.getSetImplicitKey.getKey.toByteArray // The key row is serialized as a byte array, we need to convert it back to a Row - val keyRow = - deserializeRow(keyBytes, conversionSchema(groupingKeySchema), keyRowDeserializer) + val keyRow = PythonSQLUtils.toJVMRow(keyBytes, groupingKeySchema, keyRowDeserializer) ImplicitGroupingKeyTracker.setImplicitKey(keyRow) // Reset the list/map state iterators for a new grouping key. iterators = new mutable.HashMap[String, Iterator[Row]]() @@ -516,8 +499,8 @@ class TransformWithStateInPySparkStateServer( case ValueStateCall.MethodCase.VALUESTATEUPDATE => val byteArray = message.getValueStateUpdate.getValue.toByteArray // The value row is serialized as a byte array, we need to convert it back to a Row - val valueRow = - deserializeRow(byteArray, valueStateInfo.schema, valueStateInfo.deserializer) + val valueRow = PythonSQLUtils.toJVMRow(byteArray, valueStateInfo.schema, + valueStateInfo.deserializer) valueStateInfo.valueState.update(valueRow) sendResponse(0) case ValueStateCall.MethodCase.CLEAR => @@ -541,10 +524,7 @@ class TransformWithStateInPySparkStateServer( deserializer = if (deserializerForTest != null) { deserializerForTest } else { - new TransformWithStateInPySparkDeserializer( - listStateInfo.schema, - listStateInfo.deserializer, - applyCharVarcharChecks) + new TransformWithStateInPySparkDeserializer(listStateInfo.deserializer) } message.getMethodCase match { case ListStateCall.MethodCase.EXISTS => @@ -564,7 +544,10 @@ class TransformWithStateInPySparkStateServer( } else { val elements = message.getListStatePut.getValueList.asScala elements.map { e => - deserializeRow(e.toByteArray, listStateInfo.schema, listStateInfo.deserializer) + PythonSQLUtils.toJVMRow( + e.toByteArray, + listStateInfo.schema, + listStateInfo.deserializer) } } listStateInfo.listState.put(rows.toArray) @@ -583,7 +566,8 @@ class TransformWithStateInPySparkStateServer( } case ListStateCall.MethodCase.APPENDVALUE => val byteArray = message.getAppendValue.getValue.toByteArray - val newRow = deserializeRow(byteArray, listStateInfo.schema, listStateInfo.deserializer) + val newRow = + PythonSQLUtils.toJVMRow(byteArray, listStateInfo.schema, listStateInfo.deserializer) listStateInfo.listState.appendValue(newRow) sendResponse(0) case ListStateCall.MethodCase.APPENDLIST => @@ -596,7 +580,10 @@ class TransformWithStateInPySparkStateServer( } else { val elements = message.getAppendList.getValueList.asScala elements.map { e => - deserializeRow(e.toByteArray, listStateInfo.schema, listStateInfo.deserializer) + PythonSQLUtils.toJVMRow( + e.toByteArray, + listStateInfo.schema, + listStateInfo.deserializer) } } listStateInfo.listState.appendList(rows.toArray) @@ -633,8 +620,8 @@ class TransformWithStateInPySparkStateServer( } case MapStateCall.MethodCase.GETVALUE => val keyBytes = message.getGetValue.getUserKey.toByteArray - val keyRow = - deserializeRow(keyBytes, mapStateInfo.keySchema, mapStateInfo.keyDeserializer) + val keyRow = PythonSQLUtils.toJVMRow(keyBytes, mapStateInfo.keySchema, + mapStateInfo.keyDeserializer) val value = mapStateInfo.mapState.getValue(keyRow) if (value != null) { val valueBytes = PythonSQLUtils.toPyRow(value) @@ -647,8 +634,8 @@ class TransformWithStateInPySparkStateServer( } case MapStateCall.MethodCase.CONTAINSKEY => val keyBytes = message.getContainsKey.getUserKey.toByteArray - val keyRow = - deserializeRow(keyBytes, mapStateInfo.keySchema, mapStateInfo.keyDeserializer) + val keyRow = PythonSQLUtils.toJVMRow(keyBytes, mapStateInfo.keySchema, + mapStateInfo.keyDeserializer) if (mapStateInfo.mapState.containsKey(keyRow)) { sendResponse(0) } else { @@ -656,11 +643,11 @@ class TransformWithStateInPySparkStateServer( } case MapStateCall.MethodCase.UPDATEVALUE => val keyBytes = message.getUpdateValue.getUserKey.toByteArray - val keyRow = - deserializeRow(keyBytes, mapStateInfo.keySchema, mapStateInfo.keyDeserializer) + val keyRow = PythonSQLUtils.toJVMRow(keyBytes, mapStateInfo.keySchema, + mapStateInfo.keyDeserializer) val valueBytes = message.getUpdateValue.getValue.toByteArray - val valueRow = - deserializeRow(valueBytes, mapStateInfo.valueSchema, mapStateInfo.valueDeserializer) + val valueRow = PythonSQLUtils.toJVMRow(valueBytes, mapStateInfo.valueSchema, + mapStateInfo.valueDeserializer) mapStateInfo.mapState.updateValue(keyRow, valueRow) sendResponse(0) case MapStateCall.MethodCase.ITERATOR => @@ -701,8 +688,8 @@ class TransformWithStateInPySparkStateServer( } case MapStateCall.MethodCase.REMOVEKEY => val keyBytes = message.getRemoveKey.getUserKey.toByteArray - val keyRow = - deserializeRow(keyBytes, mapStateInfo.keySchema, mapStateInfo.keyDeserializer) + val keyRow = PythonSQLUtils.toJVMRow(keyBytes, mapStateInfo.keySchema, + mapStateInfo.keyDeserializer) mapStateInfo.mapState.removeKey(keyRow) sendResponse(0) case MapStateCall.MethodCase.CLEAR => @@ -719,20 +706,17 @@ class TransformWithStateInPySparkStateServer( stateType: StateVariableType.StateVariableType, ttlDurationMs: Option[Long], mapStateValueSchemaString: String = null): Unit = { - val logicalSchema = StructType.fromString(schemaString) - val schema = conversionSchema(logicalSchema) - val expressionEncoder = stateEncoder(logicalSchema) + val schema = StructType.fromString(schemaString) + validateStateSchema(schema, s"$stateType state") + val expressionEncoder = ExpressionEncoder(schema).resolveAndBind() stateType match { - case StateVariableType.ValueState => - if (!valueStates.contains(stateName)) { - val state = if (ttlDurationMs.isEmpty) { - statefulProcessorHandle - .getValueState[Row](stateName, expressionEncoder, TTLConfig.NONE) + case StateVariableType.ValueState => if (!valueStates.contains(stateName)) { + val state = if (ttlDurationMs.isEmpty) { + statefulProcessorHandle.getValueState[Row](stateName, Encoders.row(schema), + TTLConfig.NONE) } else { statefulProcessorHandle.getValueState( - stateName, - expressionEncoder, - TTLConfig(Duration.ofMillis(ttlDurationMs.get))) + stateName, Encoders.row(schema), TTLConfig(Duration.ofMillis(ttlDurationMs.get))) } valueStates.put(stateName, ValueStateInfo(state, schema, expressionEncoder.createDeserializer())) @@ -741,61 +725,41 @@ class TransformWithStateInPySparkStateServer( sendResponse(1, s"Value state $stateName already exists") } - case StateVariableType.ListState => - if (!listStates.contains(stateName)) { - val state = if (ttlDurationMs.isEmpty) { - statefulProcessorHandle - .getListState[Row](stateName, expressionEncoder, TTLConfig.NONE) - } else { - statefulProcessorHandle.getListState( - stateName, - expressionEncoder, - TTLConfig(Duration.ofMillis(ttlDurationMs.get))) - } - listStates.put( - stateName, - ListStateInfo( - state, - schema, - expressionEncoder.createDeserializer(), - expressionEncoder.createSerializer())) - sendResponse(0) + case StateVariableType.ListState => if (!listStates.contains(stateName)) { + val state = if (ttlDurationMs.isEmpty) { + statefulProcessorHandle.getListState[Row](stateName, Encoders.row(schema), + TTLConfig.NONE) } else { - sendResponse(1, s"List state $stateName already exists") + statefulProcessorHandle.getListState( + stateName, Encoders.row(schema), TTLConfig(Duration.ofMillis(ttlDurationMs.get))) } + listStates.put(stateName, + ListStateInfo(state, schema, expressionEncoder.createDeserializer(), + expressionEncoder.createSerializer())) + sendResponse(0) + } else { + sendResponse(1, s"List state $stateName already exists") + } - case StateVariableType.MapState => - if (!mapStates.contains(stateName)) { - val logicalValueSchema = StructType.fromString(mapStateValueSchemaString) - val valueSchema = conversionSchema(logicalValueSchema) - val valueExpressionEncoder = stateEncoder(logicalValueSchema) - val state = if (ttlDurationMs.isEmpty) { - statefulProcessorHandle.getMapState[Row, Row]( - stateName, - expressionEncoder, - valueExpressionEncoder, - TTLConfig.NONE) - } else { - statefulProcessorHandle.getMapState[Row, Row]( - stateName, - expressionEncoder, - valueExpressionEncoder, - TTLConfig(Duration.ofMillis(ttlDurationMs.get))) - } - mapStates.put( - stateName, - MapStateInfo( - state, - schema, - valueSchema, - expressionEncoder.createDeserializer(), - expressionEncoder.createSerializer(), - valueExpressionEncoder.createDeserializer(), - valueExpressionEncoder.createSerializer())) - sendResponse(0) + case StateVariableType.MapState => if (!mapStates.contains(stateName)) { + val valueSchema = StructType.fromString(mapStateValueSchemaString) + validateStateSchema(valueSchema, "map state value") + val valueExpressionEncoder = ExpressionEncoder(valueSchema).resolveAndBind() + val state = if (ttlDurationMs.isEmpty) { + statefulProcessorHandle.getMapState[Row, Row](stateName, + Encoders.row(schema), Encoders.row(valueSchema), TTLConfig.NONE) } else { - sendResponse(1, s"Map state $stateName already exists") + statefulProcessorHandle.getMapState[Row, Row](stateName, Encoders.row(schema), + Encoders.row(valueSchema), TTLConfig(Duration.ofMillis(ttlDurationMs.get))) } + mapStates.put(stateName, + MapStateInfo(state, schema, valueSchema, expressionEncoder.createDeserializer(), + expressionEncoder.createSerializer(), valueExpressionEncoder.createDeserializer(), + valueExpressionEncoder.createSerializer())) + sendResponse(0) + } else { + sendResponse(1, s"Map state $stateName already exists") + } } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala index 82d91c78af821..6f5e1ba17ab34 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala @@ -137,7 +137,7 @@ class ArrowColumnarPythonUDFSuite extends SharedSparkSession { padded.queryExecution.executedPlan).head assert(arrowExec.child.supportsColumnar, "ArrowEvalPythonExec should retain its Arrow-backed columnar child") - assert(padded.select("udf_id").collect().map(_.getString(0)).toSeq === + assert(padded.collect().map(_.getString(4)).toSeq === (0 until 10).map(_.toString.padTo(4, ' ').mkString)) val exception = intercept[SparkException] { @@ -173,10 +173,10 @@ class ArrowColumnarPythonUDFSuite extends SharedSparkSession { assert(arrowExec.child.supportsColumnar, "ArrowEvalPythonExec should retain its Arrow-backed columnar child") - val rows = result.select("udf_id", "udf_name").collect() + val rows = result.collect() rows.zipWithIndex.foreach { case (row, index) => - assert(row.getString(0) === index.toString) - assert(row.getString(1) === s"row_$index") + assert(row.getString(4) === index.toString) + assert(row.getString(5) === s"row_$index") } } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/EvaluatePythonSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/EvaluatePythonSuite.scala index 8852bf609ceba..ee99ea750ce2d 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/EvaluatePythonSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/EvaluatePythonSuite.scala @@ -18,11 +18,13 @@ package org.apache.spark.sql.execution.python import org.apache.spark.{SparkFunSuite, SparkIllegalArgumentException, SparkRuntimeException} -import org.apache.spark.sql.catalyst.util.STUtils +import org.apache.spark.sql.catalyst.SQLConfHelper +import org.apache.spark.sql.catalyst.util.{MapData, STUtils} +import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ import org.apache.spark.unsafe.types.{BinaryView, UTF8String} -class EvaluatePythonSuite extends SparkFunSuite { +class EvaluatePythonSuite extends SparkFunSuite with SQLConfHelper { test("SPARK-59275: makeFromJava enforces CHAR/VARCHAR results") { val charResult = EvaluatePython.makeFromJava(CharType(4))("ab") @@ -39,6 +41,29 @@ class EvaluatePythonSuite extends SparkFunSuite { parameters = Map("limit" -> "4")) } + test("SPARK-59275: makeFromJava deduplicates normalized CHAR map keys") { + val input = new java.util.LinkedHashMap[String, Integer]() + input.put("a", 1) + input.put("a ", 2) + val convert = EvaluatePython.makeFromJava(MapType(CharType(2), IntegerType)) + + checkError( + exception = intercept[SparkRuntimeException] { + convert(input) + }, + condition = "DUPLICATED_MAP_KEY", + parameters = Map( + "key" -> "a ", + "mapKeyDedupPolicy" -> "\"spark.sql.mapKeyDedupPolicy\"")) + + withSQLConf(SQLConf.MAP_KEY_DEDUP_POLICY.key -> "LAST_WIN") { + val result = convert(input).asInstanceOf[MapData] + assert(result.numElements() === 1) + assert(result.keyArray().getUTF8String(0) === UTF8String.fromString("a ")) + assert(result.valueArray().getInt(0) === 2) + } + } + // POINT(1 2) in WKB, little-endian. private val pointWkb: Array[Byte] = "010100000000000000000031400000000000001C40" .grouped(2).map(Integer.parseInt(_, 16).toByte).toArray diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServerSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServerSuite.scala index 7d709cf4fd297..1a28c2c3dd809 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServerSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServerSuite.scala @@ -38,9 +38,7 @@ import org.mockito.invocation.InvocationOnMock import org.scalatest.BeforeAndAfterEach import org.scalatest.concurrent.Eventually.{eventually, timeout} -import net.razorvine.pickle.Pickler - -import org.apache.spark.SparkFunSuite +import org.apache.spark.{SparkFunSuite, SparkUnsupportedOperationException} import org.apache.spark.sql.{Encoder, Row} import org.apache.spark.sql.catalyst.encoders.ExpressionEncoder import org.apache.spark.sql.catalyst.expressions.GenericRowWithSchema @@ -48,7 +46,7 @@ import org.apache.spark.sql.execution.streaming.operators.stateful.transformwith import org.apache.spark.sql.execution.streaming.state.StateMessage import org.apache.spark.sql.execution.streaming.state.StateMessage.{AppendList, AppendValue, Clear, ContainsKey, DeleteTimer, Exists, ExpiryTimerRequest, Get, GetProcessingTime, GetValue, GetWatermark, HandleState, Keys, ListStateCall, ListStateGet, ListStatePut, ListTimers, MapStateCall, ParseStringSchema, RegisterTimer, RemoveKey, SetHandleState, StateCallCommand, StatefulProcessorCall, TimerRequest, TimerStateCallCommand, TimerValueRequest, UpdateValue, UtilsRequest, Values, ValueStateCall, ValueStateUpdate} import org.apache.spark.sql.streaming.{ListState, MapState, TTLConfig, ValueState} -import org.apache.spark.sql.types.{IntegerType, StringType, StructField, StructType, VarcharType} +import org.apache.spark.sql.types.{ArrayType, CharType, IntegerType, StructField, StructType} import org.apache.spark.tags.SlowSQLTest @SlowSQLTest @@ -233,6 +231,58 @@ class TransformWithStateInPySparkStateServerSuite extends SparkFunSuite with Bef } } + test("CHAR/VARCHAR grouping key schemas are rejected before state processing") { + val schema = StructType(StructField("key", ArrayType(CharType(3))) :: Nil) + val error = intercept[SparkUnsupportedOperationException] { + new TransformWithStateInPySparkStateServer( + serverSocket, + statefulProcessorHandle, + schema, + 2, + batchTimestampMs, + eventTimeWatermarkForEviction, + outputStream) + } + assert(error.getCondition === "UNSUPPORTED_FEATURE.PYTHON_STATE_CHAR_VARCHAR_SCHEMA") + assert(error.getMessageParameters.get("schemaKind") === "grouping key") + } + + test("CHAR/VARCHAR value, list, and map state schemas are rejected") { + val unsupportedSchema = + StructType(StructField("value", ArrayType(CharType(3))) :: Nil).toString + val supportedSchema = stateSchema.toString + val calls = Seq( + StatefulProcessorCall.newBuilder().setGetValueState( + StateCallCommand.newBuilder() + .setStateName("value") + .setSchema(unsupportedSchema) + .build()).build(), + StatefulProcessorCall.newBuilder().setGetListState( + StateCallCommand.newBuilder() + .setStateName("list") + .setSchema(unsupportedSchema) + .build()).build(), + StatefulProcessorCall.newBuilder().setGetMapState( + StateCallCommand.newBuilder() + .setStateName("map-key") + .setSchema(unsupportedSchema) + .setMapStateValueSchema(supportedSchema) + .build()).build(), + StatefulProcessorCall.newBuilder().setGetMapState( + StateCallCommand.newBuilder() + .setStateName("map-value") + .setSchema(supportedSchema) + .setMapStateValueSchema(unsupportedSchema) + .build()).build()) + + calls.foreach { call => + val error = intercept[SparkUnsupportedOperationException] { + stateServer.handleStatefulProcessorCall(call) + } + assert(error.getCondition === "UNSUPPORTED_FEATURE.PYTHON_STATE_CHAR_VARCHAR_SCHEMA") + } + } + test("delete if exists") { val stateCallCommandBuilder = StateCallCommand.newBuilder() .setStateName("stateName") @@ -295,64 +345,6 @@ class TransformWithStateInPySparkStateServerSuite extends SparkFunSuite with Bef verify(outputStream).writeInt(0) } - test("legacy CHAR/VARCHAR policy applies to value, list, and map state updates") { - val schema = StructType(Seq(StructField("value", VarcharType(3)))) - val encoder = - ExpressionEncoder(StructType(Seq(StructField("value", StringType)))).resolveAndBind() - val deserializer = encoder.createDeserializer() - val serializer = encoder.createSerializer() - val bytes = ByteString.copyFrom(new Pickler(true, false).dumps(Array[AnyRef]("abcd"))) - val valueStateInfo = - mutable.HashMap(stateName -> ValueStateInfo(valueState, schema, deserializer)) - val listStateInfo = - mutable.HashMap(stateName -> ListStateInfo(listState, schema, deserializer, serializer)) - val mapStateInfo = mutable.HashMap( - stateName -> MapStateInfo( - mapState, - schema, - schema, - deserializer, - serializer, - deserializer, - serializer)) - val legacyStateServer = new TransformWithStateInPySparkStateServer( - serverSocket, - statefulProcessorHandle, - groupingKeySchema, - 2, - outputStreamForTest = outputStream, - valueStateMapForTest = valueStateInfo, - deserializerForTest = transformWithStateInPySparkDeserializer, - listStatesMapForTest = listStateInfo, - mapStatesMapForTest = mapStateInfo, - applyCharVarcharChecks = false) - - legacyStateServer.handleValueStateRequest( - ValueStateCall - .newBuilder() - .setStateName(stateName) - .setValueStateUpdate(ValueStateUpdate.newBuilder().setValue(bytes)) - .build()) - legacyStateServer.handleListStateRequest( - ListStateCall - .newBuilder() - .setStateName(stateName) - .setAppendValue(AppendValue.newBuilder().setValue(bytes)) - .build()) - legacyStateServer.handleMapStateRequest( - MapStateCall - .newBuilder() - .setStateName(stateName) - .setUpdateValue(UpdateValue.newBuilder().setUserKey(bytes).setValue(bytes)) - .build()) - - verify(valueState).update(argThat((row: Row) => row.getString(0) == "abcd")) - verify(listState).appendValue(argThat((row: Row) => row.getString(0) == "abcd")) - verify(mapState).updateValue( - argThat((row: Row) => row.getString(0) == "abcd"), - argThat((row: Row) => row.getString(0) == "abcd")) - } - test("list state exists") { val message = ListStateCall.newBuilder().setStateName(stateName) .setExists(Exists.newBuilder().build()).build() From 835f5f87d4a2c9eb3d4c33897bff3a7977329889 Mon Sep 17 00:00:00 2001 From: srielau Date: Thu, 10 Sep 2026 00:52:45 +0000 Subject: [PATCH 7/9] fix: [SPARK-59275] address Python CI regressions --- python/pyspark/sql/connect/udtf.py | 4 ++-- python/pyspark/sql/tests/arrow/test_arrow.py | 4 +--- .../pyspark/sql/tests/arrow/test_arrow_python_udf.py | 4 +--- python/pyspark/sql/tests/arrow/test_arrow_udtf.py | 4 +--- python/pyspark/sql/tests/test_udf.py | 8 ++------ python/pyspark/sql/tests/test_udtf.py | 4 +--- python/pyspark/sql/udf.py | 6 +++--- .../execution/python/UserDefinedPythonFunction.scala | 11 +++++++++-- .../python/ArrowColumnarPythonUDFSuite.scala | 4 ++-- .../TransformWithStateInPySparkStateServerSuite.scala | 4 ++-- 10 files changed, 24 insertions(+), 29 deletions(-) diff --git a/python/pyspark/sql/connect/udtf.py b/python/pyspark/sql/connect/udtf.py index 471e80c7e02bb..12b0550842c9b 100644 --- a/python/pyspark/sql/connect/udtf.py +++ b/python/pyspark/sql/connect/udtf.py @@ -36,7 +36,7 @@ from pyspark.sql.udtf import ( # noqa: F401 AnalyzeArgument, AnalyzeResult, - _check_arrow_udtf_return_type, + _check_udtf_return_type, _validate_udtf_handler, ) from pyspark.sql.udtf import UDTFRegistration as PySparkUDTFRegistration @@ -186,7 +186,7 @@ def _check_return_type(self, session: "SparkSession") -> None: if isinstance(self.returnType, UnparsedDataType) else self.returnType ) - _check_arrow_udtf_return_type(return_type, self.evalType) + _check_udtf_return_type(return_type) self._validated_return_type_session_ids.add(session._session_id) def _build_common_inline_user_defined_table_function( diff --git a/python/pyspark/sql/tests/arrow/test_arrow.py b/python/pyspark/sql/tests/arrow/test_arrow.py index a16d5f23996bd..10369b0db8e7f 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow.py +++ b/python/pyspark/sql/tests/arrow/test_arrow.py @@ -1505,9 +1505,7 @@ def test_char_varchar_explicit_schema_and_to_arrow(self): "spark.sql.execution.arrow.localRelationThreshold": "0", } ): - df = self.spark.createDataFrame( - pa.table({"c": ["a"], "v": ["abcd"]}), legacy_schema - ) + df = self.spark.createDataFrame(pa.table({"c": ["a"], "v": ["abcd"]}), legacy_schema) self.assertEqual(df.schema, StructType().add("c", "string").add("v", "string")) self.assertEqual(df.first(), Row(c="a", v="abcd")) diff --git a/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py b/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py index e1c69000a39c4..06eab1af2f484 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py @@ -330,9 +330,7 @@ def test_char_varchar_intermediate_udf_results_arrow(self): outer = udf(lambda value: value, StringType(), useArrow=True) with self.sql_conf({"spark.sql.charVarchar.standardSemantics.enabled": "true"}): - padded = self.spark.range(1).select( - outer(inner_char("id")).alias("result") - ) + padded = self.spark.range(1).select(outer(inner_char("id")).alias("result")) self.assertEqual(padded.first().result, "a ") invalid = self.spark.range(1).select(outer(inner_varchar("id"))) diff --git a/python/pyspark/sql/tests/arrow/test_arrow_udtf.py b/python/pyspark/sql/tests/arrow/test_arrow_udtf.py index e7aab5c2309cc..ab5b32ef88fb5 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_udtf.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_udtf.py @@ -58,9 +58,7 @@ def test_analyze_char_varchar_return_types_unsupported(self): class DynamicNestedCharUDTF: @staticmethod def analyze() -> AnalyzeResult: - return AnalyzeResult( - StructType([StructField("nested", ArrayType(CharType(3)))]) - ) + return AnalyzeResult(StructType([StructField("nested", ArrayType(CharType(3)))])) def eval(self): yield (["a"],) diff --git a/python/pyspark/sql/tests/test_udf.py b/python/pyspark/sql/tests/test_udf.py index fd91926c4f935..6130bf3ee1351 100644 --- a/python/pyspark/sql/tests/test_udf.py +++ b/python/pyspark/sql/tests/test_udf.py @@ -116,9 +116,7 @@ def test_char_varchar_intermediate_udf_results(self): outer = udf(lambda value: value, StringType(), useArrow=False) with self.sql_conf({"spark.sql.charVarchar.standardSemantics.enabled": "true"}): - padded = self.spark.range(1).select( - outer(inner_char("id")).alias("result") - ) + padded = self.spark.range(1).select(outer(inner_char("id")).alias("result")) self.assertEqual(padded.first().result, "a ") invalid = self.spark.range(1).select(outer(inner_varchar("id"))) @@ -180,9 +178,7 @@ def test_char_varchar_mixed_captured_policies_in_one_batch(self): ) def test_char_varchar_non_scalar_return_types_unsupported(self): - nested_return_type = StructType( - [StructField("nested", ArrayType(CharType(3)))] - ) + nested_return_type = StructType([StructField("nested", ArrayType(CharType(3)))]) struct_eval_types = [ PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF, PythonEvalType.SQL_GROUPED_MAP_ARROW_UDF, diff --git a/python/pyspark/sql/tests/test_udtf.py b/python/pyspark/sql/tests/test_udtf.py index f5f10934946e5..4a34a752c9bf3 100644 --- a/python/pyspark/sql/tests/test_udtf.py +++ b/python/pyspark/sql/tests/test_udtf.py @@ -98,9 +98,7 @@ def test_analyze_char_varchar_return_types_unsupported(self): class DynamicNestedCharUDTF: @staticmethod def analyze() -> AnalyzeResult: - return AnalyzeResult( - StructType([StructField("nested", ArrayType(CharType(3)))]) - ) + return AnalyzeResult(StructType([StructField("nested", ArrayType(CharType(3)))])) def eval(self): yield (["a"],) diff --git a/python/pyspark/sql/udf.py b/python/pyspark/sql/udf.py index 49308c97298d0..38de2728faaa2 100644 --- a/python/pyspark/sql/udf.py +++ b/python/pyspark/sql/udf.py @@ -331,9 +331,9 @@ def has_char_varchar_in_udt(data_type: DataType) -> bool: if isinstance(data_type, ArrayType): return has_char_varchar_in_udt(data_type.elementType) if isinstance(data_type, MapType): - return has_char_varchar_in_udt( - data_type.keyType - ) or has_char_varchar_in_udt(data_type.valueType) + return has_char_varchar_in_udt(data_type.keyType) or has_char_varchar_in_udt( + data_type.valueType + ) return False if has_char_varchar_in_udt(returnType): diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala index a336136559b0c..4fbeaf6974773 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala @@ -108,14 +108,21 @@ case class UserDefinedPythonFunction( } PythonAggregate(name, func, dataType, e, udfDeterministic, bufferStruct) } else { + val applyCharVarcharChecks = + CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get) + val resolvedDataType = if (applyCharVarcharChecks) { + dataType + } else { + CharVarcharUtils.replaceCharVarcharWithStringForPhysicalType(dataType) + } PythonUDF( name, func, - dataType, + resolvedDataType, e, pythonEvalType, udfDeterministic, - applyCharVarcharChecks = CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get)) + applyCharVarcharChecks = applyCharVarcharChecks) } // The ``_udf_param_N`` substitution below is positional, so a UDF // call site that supplied named arguments (e.g. SQL ``name => val`` diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala index 6f5e1ba17ab34..51384af308021 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/ArrowColumnarPythonUDFSuite.scala @@ -17,7 +17,7 @@ package org.apache.spark.sql.execution.python -import org.apache.spark.SparkException +import org.apache.spark.SparkRuntimeException import org.apache.spark.sql.IntegratedUDFTestUtils import org.apache.spark.sql.execution.SparkPlan import org.apache.spark.sql.internal.SQLConf @@ -140,7 +140,7 @@ class ArrowColumnarPythonUDFSuite extends SharedSparkSession { assert(padded.collect().map(_.getString(4)).toSeq === (0 until 10).map(_.toString.padTo(4, ' ').mkString)) - val exception = intercept[SparkException] { + val exception = intercept[SparkRuntimeException] { df.selectExpr( "id", "name", "value", "data", "arrow_varchar_udf(name) as udf_name").collect() diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServerSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServerSuite.scala index 1a28c2c3dd809..2912ec97a0af7 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServerSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServerSuite.scala @@ -249,8 +249,8 @@ class TransformWithStateInPySparkStateServerSuite extends SparkFunSuite with Bef test("CHAR/VARCHAR value, list, and map state schemas are rejected") { val unsupportedSchema = - StructType(StructField("value", ArrayType(CharType(3))) :: Nil).toString - val supportedSchema = stateSchema.toString + StructType(StructField("value", ArrayType(CharType(3))) :: Nil).json + val supportedSchema = stateSchema.json val calls = Seq( StatefulProcessorCall.newBuilder().setGetValueState( StateCallCommand.newBuilder() From aa73a7fcc0df6be8e49de7e624bf2a8ddb87e590 Mon Sep 17 00:00:00 2001 From: srielau Date: Thu, 10 Sep 2026 04:52:45 +0000 Subject: [PATCH 8/9] fix: address CHAR/VARCHAR Python CI failures --- python/pyspark/sql/tests/test_udf.py | 4 ++-- .../sql/catalyst/expressions/PythonUDF.scala | 3 ++- .../sql/connect/planner/SparkConnectPlanner.scala | 15 ++++++++++++--- .../ColumnarArrowEvalPythonEvaluatorFactory.scala | 4 +--- .../python/UserDefinedPythonFunction.scala | 3 ++- 5 files changed, 19 insertions(+), 10 deletions(-) diff --git a/python/pyspark/sql/tests/test_udf.py b/python/pyspark/sql/tests/test_udf.py index 6130bf3ee1351..b6c6dc18ba6c1 100644 --- a/python/pyspark/sql/tests/test_udf.py +++ b/python/pyspark/sql/tests/test_udf.py @@ -152,7 +152,7 @@ def test_char_varchar_view_keeps_resolved_semantics(self): } ): self.assertEqual( - self.spark.sql("SELECT c FROM char_varchar_udf_view").first().c, + self.spark.sql("SELECT c FROM char_varchar_udf_view").collect()[0].c, "a ", ) with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"): @@ -173,7 +173,7 @@ def test_char_varchar_mixed_captured_policies_in_one_batch(self): unchecked_varchar = varchar_udf("id").alias("v") self.assertEqual( - self.spark.range(1).select(checked_char, unchecked_varchar).first(), + self.spark.range(1).select(checked_char, unchecked_varchar).collect()[0], Row(c="a ", v="abcd"), ) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala index d55d6fb93893f..8e81ddf5ff562 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala @@ -338,7 +338,8 @@ case class PythonUDF( // lambda (e.g. `transform(arr, i -> transform(i, x -> f(x)))` lifts `f` to depth 2). Ignored // for every non-element-wise eval type, where it stays at its default of 1. elementwiseNestingDepth: Int = 1, - applyCharVarcharChecks: Boolean = false) + applyCharVarcharChecks: Boolean = false, + hasCharVarcharResult: Boolean = false) extends Expression with PythonFuncExpression with Unevaluable { lazy val resultAttribute: Attribute = AttributeReference(toPrettySQL(this), dataType, nullable)( diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala index 4803c5865e145..e18b8d1d571f4 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala @@ -1517,7 +1517,15 @@ class SparkConnectPlanner( if (schema == null) { throw InvalidInputErrors.schemaRequiredForLocalRelation() } - LocalRelation(schema) + LocalRelation(localRelationOutputSchema(schema)) + } + } + + private def localRelationOutputSchema(schema: StructType): StructType = { + if (CharVarcharUtils.shouldApplyWriteSideLengthCheck(session.sessionState.conf)) { + schema + } else { + CharVarcharUtils.replaceCharVarcharWithStringForPhysicalType(schema).asInstanceOf[StructType] } } @@ -1613,6 +1621,7 @@ class SparkConnectPlanner( case None => logical.LocalRelation(attributes, data.map(_.copy()).toArray.toImmutableArraySeq) case Some(schema) => + val outputSchema = localRelationOutputSchema(schema) def normalize(dt: DataType): DataType = dt match { case udt: UserDefinedType[_] => normalize(udt.sqlType) case StructType(fields) => @@ -1628,7 +1637,7 @@ class SparkConnectPlanner( case _ => dt } - val normalized = normalize(schema).asInstanceOf[StructType] + val normalized = normalize(outputSchema).asInstanceOf[StructType] import org.apache.spark.util.ArrayImplicits._ val project = Dataset @@ -1642,7 +1651,7 @@ class SparkConnectPlanner( val proj = UnsafeProjection.create(project.projectList, project.child.output) logical.LocalRelation( - DataTypeUtils.toAttributes(schema), + DataTypeUtils.toAttributes(outputSchema), data.map(proj).map(_.copy()).toSeq) } } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala index 2f9de39ad3d78..d1d6433de4d39 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala @@ -97,9 +97,7 @@ private[python] class ColumnarArrowEvalPythonEvaluatorFactory( } } private val hasCharVarcharOutput = - udfOutput.zip(udfs).exists { case (attr, udf) => - udf.applyCharVarcharChecks && CharVarcharUtils.hasCharVarchar(attr.dataType) - } + udfs.exists(_.hasCharVarcharResult) private val physicalOutputSchema = ColumnarArrowEvalPythonEvaluatorFactory .toPhysicalType(outputSchema) .asInstanceOf[StructType] diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala index 4fbeaf6974773..fe6bd77ee8e42 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala @@ -122,7 +122,8 @@ case class UserDefinedPythonFunction( e, pythonEvalType, udfDeterministic, - applyCharVarcharChecks = applyCharVarcharChecks) + applyCharVarcharChecks = applyCharVarcharChecks, + hasCharVarcharResult = CharVarcharUtils.hasCharVarchar(dataType)) } // The ``_udf_param_N`` substitution below is positional, so a UDF // call site that supplied named arguments (e.g. SQL ``name => val`` From ab2f542e262c388bc69429ab080e59a263088474 Mon Sep 17 00:00:00 2001 From: srielau Date: Thu, 10 Sep 2026 16:43:57 +0000 Subject: [PATCH 9/9] fix: [SPARK-59275] address remaining CI failures --- python/pyspark/sql/connect/udf.py | 2 ++ python/pyspark/sql/tests/connect/test_parity_udf.py | 4 ++++ .../sql/connect/planner/SparkConnectPlanner.scala | 4 +++- .../org/apache/spark/sql/classic/SparkSession.scala | 12 +++++++++++- .../ColumnarArrowEvalPythonEvaluatorFactory.scala | 4 +++- 5 files changed, 23 insertions(+), 3 deletions(-) diff --git a/python/pyspark/sql/connect/udf.py b/python/pyspark/sql/connect/udf.py index 9eb2bf5e05107..adb085a218bf8 100644 --- a/python/pyspark/sql/connect/udf.py +++ b/python/pyspark/sql/connect/udf.py @@ -192,6 +192,8 @@ def __init__( # so it survives ``_wrapped()``, ``asNondeterministic()`` and ``spark.udf.register``. self.bufferSchema = bufferSchema + _check_return_type = staticmethod(PySparkUserDefinedFunction._check_return_type) + @property def returnType(self) -> DataType: # Make sure this is called after Connect Session is initialized. diff --git a/python/pyspark/sql/tests/connect/test_parity_udf.py b/python/pyspark/sql/tests/connect/test_parity_udf.py index d6e44759185de..4232ddf840b23 100644 --- a/python/pyspark/sql/tests/connect/test_parity_udf.py +++ b/python/pyspark/sql/tests/connect/test_parity_udf.py @@ -46,6 +46,10 @@ def test_udf_with_input_file_name_for_hadooprdd(self): def test_same_accumulator_in_udfs(self): super().test_same_accumulator_in_udfs() + @unittest.skip("Spark Connect resolves both UDFs when the plan is submitted.") + def test_char_varchar_mixed_captured_policies_in_one_batch(self): + super().test_char_varchar_mixed_captured_policies_in_one_batch() + @unittest.skip("Spark Connect does not support broadcast but the test depends on it.") def test_broadcast_in_udf(self): super().test_broadcast_in_udf() diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala index e18b8d1d571f4..317a48bacc335 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala @@ -1525,7 +1525,9 @@ class SparkConnectPlanner( if (CharVarcharUtils.shouldApplyWriteSideLengthCheck(session.sessionState.conf)) { schema } else { - CharVarcharUtils.replaceCharVarcharWithStringForPhysicalType(schema).asInstanceOf[StructType] + CharVarcharUtils + .replaceCharVarcharWithStringForPhysicalType(schema) + .asInstanceOf[StructType] } } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/classic/SparkSession.scala b/sql/core/src/main/scala/org/apache/spark/sql/classic/SparkSession.scala index ffbfed21bb29d..f1a054ed08ffd 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/classic/SparkSession.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/classic/SparkSession.scala @@ -887,11 +887,21 @@ class SparkSession private( private[sql] def applySchemaToPythonRDD( rdd: RDD[Array[Any]], schema: StructType): DataFrame = { + val applyCharVarcharChecks = + CharVarcharUtils.hasCharVarchar(schema) && + CharVarcharUtils.shouldApplyWriteSideLengthCheck(sessionState.conf) + val outputSchema = if (applyCharVarcharChecks) { + schema + } else { + CharVarcharUtils + .replaceCharVarcharWithStringForPhysicalType(schema) + .asInstanceOf[StructType] + } val rowRdd = rdd.mapPartitions { iter => val fromJava = python.EvaluatePython.makeFromJava(schema) iter.map(r => fromJava(r).asInstanceOf[InternalRow]) } - internalCreateDataFrame(rowRdd, schema) + internalCreateDataFrame(rowRdd, outputSchema) } /** diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala index d1d6433de4d39..a2dd06dec7784 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/ColumnarArrowEvalPythonEvaluatorFactory.scala @@ -179,9 +179,11 @@ private[python] class ColumnarArrowEvalPythonEvaluatorFactory( udfInputSchema, outputTypes, inputColumnIndices.get) } else { // Path 2 & 3: non-Arrow or complex expressions. + val directInputColumnIndices = + if (hasCharVarcharOutput) None else inputColumnIndices evalWithRowQueue(peekIter, context, pyFuncs, argMetas, allInputs.toSeq, udfInputSchema, outputTypes, - inputColumnIndices) + directInputColumnIndices) } }