diff --git a/common/utils/src/main/resources/error/error-conditions.json b/common/utils/src/main/resources/error/error-conditions.json index ea8f77498e5d5..073bcaf0f8fb7 100644 --- a/common/utils/src/main/resources/error/error-conditions.json +++ b/common/utils/src/main/resources/error/error-conditions.json @@ -9145,11 +9145,21 @@ "Purge table." ] }, + "PYTHON_STATE_CHAR_VARCHAR_SCHEMA" : { + "message" : [ + "Python state does not support CHAR/VARCHAR in schema ." + ] + }, "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." ] }, + "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/connect/udtf.py b/python/pyspark/sql/connect/udtf.py index aa7d86c0f72d7..12b0550842c9b 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 @@ -33,7 +33,12 @@ 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.udtf import ( # noqa: F401 + AnalyzeArgument, + AnalyzeResult, + _check_udtf_return_type, + _validate_udtf_handler, +) from pyspark.sql.udtf import UDTFRegistration as PySparkUDTFRegistration from pyspark.util import PythonEvalType @@ -166,10 +171,32 @@ 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, 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 = ( + session._parse_ddl(self.returnType.data_type_string) + if isinstance(self.returnType, UnparsedDataType) + else self.returnType + ) + _check_udtf_return_type(return_type) + 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(session) + def to_expr(col: "ColumnOrName") -> Expression: if isinstance(col, Column): return col._expr @@ -201,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": @@ -245,6 +272,7 @@ def register( }, ) + 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/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..10369b0db8e7f 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,80 @@ 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() + + 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)), + 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.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): 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..06eab1af2f484 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_udf.py @@ -18,14 +18,15 @@ 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 +271,85 @@ 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() + + 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_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/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 ac695a9f8001e..ab5b32ef88fb5 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_udtf.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_udtf.py @@ -18,9 +18,10 @@ import unittest from typing import Iterator, Optional -from pyspark.errors import PySparkAttributeError, PythonException -from pyspark.sql.functions import arrow_udtf, lit -from pyspark.sql.types import IntegerType, Row, StructField, StructType +from pyspark.errors import PySparkAttributeError, PySparkNotImplementedError, PythonException +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 @@ -33,6 +34,38 @@ @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, "CHAR/VARCHAR return type in Python UDTFs" + ): + 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, "Python UDTFs do not support CHAR/VARCHAR"): + DynamicNestedCharUDTF().collect() + def test_arrow_udtf_data_conversion_error(self): from pyspark.sql.functions import udtf @@ -40,8 +73,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/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_creation.py b/python/pyspark/sql/tests/test_creation.py index 0b4542206001d..1dbc966e29281 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,31 @@ 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() + + 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_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 b1a7f9214fbbf..b6c6dc18ba6c1 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 @@ -35,6 +40,7 @@ ArrayType, BinaryType, BooleanType, + CharType, DayTimeIntervalType, DoubleType, IntegerType, @@ -44,6 +50,8 @@ StructField, StructType, TimestampNTZType, + UserDefinedType, + VarcharType, VariantType, VariantVal, ) @@ -55,10 +63,185 @@ 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: + 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_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_intermediate_udf_results(self): + 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 ") + + 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_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").collect()[0].c, + "a ", + ) + 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).collect()[0], + 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_PANDAS_ITER_UDF, + 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"): + 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) + 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_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"]) @@ -1451,49 +1634,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/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..4a34a752c9bf3 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,33 @@ 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 2618ecb7d4b04..38de2728faaa2 100644 --- a/python/pyspark/sql/udf.py +++ b/python/pyspark/sql/udf.py @@ -29,9 +29,15 @@ 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, ) from pyspark.sql.utils import get_active_spark_context @@ -314,9 +320,49 @@ 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 + + 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, + 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 _InvalidCharVarcharArrowTypeError + 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 +376,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 +389,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 +404,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 +429,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 +451,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 +471,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 +490,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", @@ -460,7 +506,10 @@ def _check_return_type(returnType: DataType, evalType: int) -> 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): @@ -471,7 +520,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", @@ -480,10 +529,13 @@ def _check_return_type(returnType: DataType, evalType: int) -> 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 - to_arrow_type(returnType, timezone="UTC") + check_arrow_type() except TypeError: raise PySparkNotImplementedError( errorClass="NOT_IMPLEMENTED", @@ -492,6 +544,16 @@ def _check_return_type(returnType: DataType, evalType: int) -> 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/python/pyspark/sql/udtf.py b/python/pyspark/sql/udtf.py index e8dc25dc907ec..b4035d7bdd612 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 @@ -317,6 +325,16 @@ def _validate_udtf_handler(cls: Any, returnType: Optional[Union[StructType, str] ) +def _check_udtf_return_type(return_type: DataType) -> None: + if _has_type(return_type, (CharType, VarcharType)): + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={ + "feature": f"CHAR/VARCHAR return type in Python UDTFs: {return_type}" + }, + ) + + class UserDefinedTableFunction: """ User-defined table function in Python @@ -369,6 +387,7 @@ def returnType(self) -> Optional[StructType]: "return_type": f"{parsed}", }, ) + _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/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..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 @@ -337,7 +337,9 @@ 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, + hasCharVarcharResult: Boolean = false) extends Expression with PythonFuncExpression with Unevaluable { lazy val resultAttribute: Attribute = AttributeReference(toPrettySQL(this), dataType, nullable)( @@ -495,7 +497,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 +528,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/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/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..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,6 +3022,20 @@ private[sql] object QueryCompilationErrors extends QueryErrorsBase with Compilat messageParameters = Map("config" -> toSQLConf(config))) } + def invalidPythonUDTFReturnType(dataType: DataType): SparkUnsupportedOperationException = { + new SparkUnsupportedOperationException( + 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/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/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 464cac157b25c..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 @@ -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._ @@ -356,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( @@ -367,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 || { @@ -384,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) } @@ -450,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 ) @@ -547,7 +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.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 @@ -557,29 +579,40 @@ 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, + physicalSchema, timeZoneId, errorOnDuplicatedFieldNames, largeVarTypes, TaskContext.get()) + if (applyCharVarcharChecks) { + val projection = UnsafeProjection.create(checkedAttrs, attrs) + rows.map(row => projection(row).copy(): InternalRow) + } else { + 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, TaskContext.get()) // Project/copy it. Otherwise, the Arrow column vectors will be closed and released out. - val proj = UnsafeProjection.create(attrs, attrs) - Dataset.ofRows(session, - LocalRelation(attrs, data.map(r => proj(r).copy()).toArray.toImmutableArraySeq)) + val proj = UnsafeProjection.create(checkedAttrs, attrs) + val rows = + try { + data.map(r => proj(r).copy()).toArray.toImmutableArraySeq + } finally { + data.close() + } + Dataset.ofRows(session, LocalRelation(toAttributes(outputSchema), 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 0170d4354a6a9..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 @@ -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 @@ -191,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)], @@ -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/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 3981602875adb..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 @@ -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 @@ -37,6 +38,14 @@ 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. * @@ -79,6 +88,20 @@ private[python] class ColumnarArrowEvalPythonEvaluatorFactory( sessionUUID: Option[String]) extends PartitionEvaluatorFactory[ColumnarBatch, ColumnarBatch] { + 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 + } + } + private val hasCharVarcharOutput = + udfs.exists(_.hasCharVarcharResult) + private val physicalOutputSchema = ColumnarArrowEvalPythonEvaluatorFactory + .toPhysicalType(outputSchema) + .asInstanceOf[StructType] + override def createEvaluator() : PartitionEvaluator[ColumnarBatch, ColumnarBatch] = new ColumnarArrowEvalPythonPartitionEvaluator @@ -137,10 +160,9 @@ 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 => + ColumnarArrowEvalPythonEvaluatorFactory.toPhysicalType(attr.dataType) + } val inputColumnIndices = resolveColumnIndices(allInputs.toSeq) @@ -151,7 +173,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 +309,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 +337,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..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 @@ -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 @@ -33,8 +34,21 @@ import org.apache.spark.util.Utils abstract class EvalPythonEvaluatorFactory( childOutput: Seq[Attribute], udfs: Seq[PythonUDF], - output: Seq[Attribute]) - extends PartitionEvaluatorFactory[InternalRow, InternalRow] { + 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( funcs: Seq[(ChainedPythonFunctions, Long)], @@ -119,7 +133,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..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,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, 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} @@ -146,10 +147,47 @@ 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 makeFromJavaDefault(dataType: DataType): Any => Any = dataType match { + 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 { case BooleanType => (obj: Any) => nullSafeConvert(obj) { case b: Boolean => b } @@ -207,6 +245,18 @@ object EvaluatePython { case c: Int => c.toLong } + case c: CharType if applyCharVarcharChecks => (obj: Any) => nullSafeConvert(obj) { + case _ => + CharVarcharCodegenUtils.charTypeWriteSideCheck( + UTF8String.fromString(obj.toString), c.length) + } + + case v: VarcharType if applyCharVarcharChecks => (obj: Any) => nullSafeConvert(obj) { + case _ => + CharVarcharCodegenUtils.varcharTypeWriteSideCheck( + UTF8String.fromString(obj.toString), v.length) + } + case _: StringType => (obj: Any) => nullSafeConvert(obj) { case _ => UTF8String.fromString(obj.toString) } @@ -217,7 +267,7 @@ object EvaluatePython { } case ArrayType(elementType, _) => - val elementFromJava = makeFromJava(elementType) + val elementFromJava = makeFromJava(elementType, applyCharVarcharChecks) (obj: Any) => nullSafeConvert(obj) { case c: java.util.List[_] => @@ -227,19 +277,20 @@ 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[_, _] => - 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) => - 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 => @@ -261,7 +312,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/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 c290ec04dedab..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 @@ -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 /** @@ -188,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. @@ -200,6 +203,8 @@ object ExtractPythonUDFs extends Rule[LogicalPlan] with Logging { case Seq(child: PythonUDF) => correctEvalType(e, pythonUDFArrowFallbackOnUDT) == correctEvalType(child, pythonUDFArrowFallbackOnUDT) && + !(child.applyCharVarcharChecks && + 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/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..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 @@ -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,22 @@ case class UserDefinedPythonFunction( } PythonAggregate(name, func, dataType, e, udfDeterministic, bufferStruct) } else { - PythonUDF(name, func, dataType, e, pythonEvalType, udfDeterministic) + val applyCharVarcharChecks = + CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get) + val resolvedDataType = if (applyCharVarcharChecks) { + dataType + } else { + CharVarcharUtils.replaceCharVarcharWithStringForPhysicalType(dataType) + } + PythonUDF( + name, + func, + resolvedDataType, + e, + pythonEvalType, + udfDeterministic, + 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`` @@ -208,6 +223,12 @@ case class UserDefinedPythonTableFunction( pythonEvalType: Int, udfDeterministic: Boolean) { + private def validateReturnType(schema: StructType): Unit = { + if (CharVarcharUtils.hasCharVarchar(schema)) { + throw QueryCompilationErrors.invalidPythonUDTFReturnType(schema) + } + } + def this( name: String, func: PythonFunction, @@ -243,8 +264,11 @@ case class UserDefinedPythonTableFunction( case _ => false } + val applyCharVarcharChecks = + CharVarcharUtils.shouldApplyWriteSideLengthCheck(SQLConf.get) val udtf = returnType match { case Some(rt) => + validateReturnType(rt) PythonUDTF( name = name, func = func, @@ -253,12 +277,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() + validateReturnType(analyzeResult.schema) + analyzeResult } UnresolvedPolymorphicPythonUDTF( name = name, @@ -267,7 +294,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/TransformWithStateInPySparkStateServer.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/streaming/TransformWithStateInPySparkStateServer.scala index 6d085eb8980a3..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 @@ -40,6 +40,8 @@ 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} @@ -80,6 +82,14 @@ class TransformWithStateInPySparkStateServer( import PythonResponseWriterUtils._ + private def validateStateSchema(schema: StructType, schemaKind: String): Unit = { + if (CharVarcharUtils.hasCharVarchar(schema)) { + throw QueryCompilationErrors.invalidPythonStateSchema(schema, schemaKind) + } + } + + validateStateSchema(groupingKeySchema, "grouping key") + private val keyRowDeserializer: ExpressionEncoder.Deserializer[Row] = ExpressionEncoder(groupingKeySchema).resolveAndBind().createDeserializer() private var inputStream: DataInputStream = _ @@ -697,6 +707,7 @@ class TransformWithStateInPySparkStateServer( ttlDurationMs: Option[Long], mapStateValueSchemaString: String = null): Unit = { val schema = StructType.fromString(schemaString) + validateStateSchema(schema, s"$stateType state") val expressionEncoder = ExpressionEncoder(schema).resolveAndBind() stateType match { case StateVariableType.ValueState => if (!valueStates.contains(stateName)) { @@ -732,6 +743,7 @@ class TransformWithStateInPySparkStateServer( 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, 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..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(", ")})" @@ -1410,6 +1429,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 +1654,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/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 faf8c77678f3a..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,11 +17,12 @@ package org.apache.spark.sql.execution.python +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 -import org.apache.spark.sql.test.SharedSparkSession -import org.apache.spark.sql.types.StringType +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. @@ -52,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") { @@ -103,6 +116,71 @@ 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.collect().map(_.getString(4)).toSeq === + (0 until 10).map(_.toString.padTo(4, ' ').mkString)) + + val exception = intercept[SparkRuntimeException] { + 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: 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.collect() + rows.zipWithIndex.foreach { case (row, index) => + assert(row.getString(4) === index.toString) + assert(row.getString(5) === 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/EvaluatePythonSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/python/EvaluatePythonSuite.scala index ec26f1b2a865f..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,51 @@ 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 +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") + 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")) + } + + 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" 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"))) 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..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 @@ -38,7 +38,7 @@ import org.mockito.invocation.InvocationOnMock import org.scalatest.BeforeAndAfterEach import org.scalatest.concurrent.Eventually.{eventually, timeout} -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 @@ -46,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, StructField, StructType} +import org.apache.spark.sql.types.{ArrayType, CharType, IntegerType, StructField, StructType} import org.apache.spark.tags.SlowSQLTest @SlowSQLTest @@ -231,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).json + val supportedSchema = stateSchema.json + 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")