srielau commented on code in PR #58549:
URL: https://github.com/apache/spark/pull/58549#discussion_r4049457057
##########
python/pyspark/sql/tests/arrow/test_arrow.py:
##########
@@ -1433,6 +1441,137 @@ def test_toArrow_duplicate_field_names(self):
):
df.limit(0).toArrow()
+ def test_char_varchar_explicit_schema_and_to_arrow(self):
Review Comment:
This test lives in `ArrowTestsMixin`, so `ArrowParityTests` inherits it.
That suite uses a Connect session, but this PR explicitly rejects this schema
in `connect/session.py`; it will fail at the first `createDataFrame` call. The
UDT-storage test below also cannot assign Connect's read-only `_schema`
property. Please move these Classic-only tests to `ArrowTests` or override/skip
both in `test_parity_arrow.py`.
##########
sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala:
##########
@@ -337,9 +337,14 @@ 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,
+ // Original CHAR/VARCHAR result type when write-side checks apply. Absent
for unconstrained
+ // results so CHAR/VARCHAR policy is not part of PythonUDF equality.
Review Comment:
These two sentences are fragments.
```suggestion
// The original CHAR/VARCHAR result type when write-side checks apply.
This is absent for
// unconstrained results, so the CHAR/VARCHAR policy is not part of
PythonUDF equality.
```
##########
sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala:
##########
@@ -337,9 +337,14 @@ 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,
+ // Original CHAR/VARCHAR result type when write-side checks apply. Absent
for unconstrained
+ // results so CHAR policy is not part of PythonUDF equality.
Review Comment:
Updated the comment to refer to CHAR/VARCHAR policy.
##########
python/pyspark/sql/connect/udtf.py:
##########
@@ -166,10 +171,33 @@ def __init__(
self._name = name or func.__name__
self.evalType = evalType
self.deterministic = deterministic
+ self._validated_return_type_session_ids: Set[str] = set()
Review Comment:
Replaced the monotonically growing set with a one-entry current-session
memo. The regression verifies same-session reuse and bounded replacement across
session IDs.
##########
python/pyspark/sql/tests/arrow/test_arrow_udf_scalar.py:
##########
@@ -58,6 +60,50 @@
@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.legacy.charVarcharAsString": "false",
+ "spark.sql.preserveCharVarcharTypeInfo": "false",
+ "spark.sql.charVarchar.standardSemantics.enabled": "false",
+ }
+ ):
+ for function in (scalar_char, iterator_char):
+ result = self.spark.range(1).select(function("id").alias("c"))
+ self.assertEqual(result.schema["c"].dataType, StringType())
+ self.assertEqual(result.first().c, "a ")
+ for function in (scalar_varchar, iterator_varchar):
+ with self.assertRaisesRegex(Exception, "EXCEED_LIMIT_LENGTH"):
+ self.spark.range(1).select(function("id")).collect()
+
+ 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()
Review Comment:
Added first-class schema assertions for native Arrow, pandas, higher-order
array elements, row UDFs, DataFrame creation, and Python-RDD conversion.
##########
sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/CharVarcharUtils.scala:
##########
@@ -33,6 +33,43 @@ object CharVarcharUtils extends Logging with
SparkCharVarcharUtils {
// visible for testing
private[sql] val CHAR_VARCHAR_TYPE_STRING_METADATA_KEY =
"__CHAR_VARCHAR_TYPE_STRING"
+ /**
+ * Returns whether a type contains CHAR/VARCHAR, including inside UDT
storage. This is intended
+ * for validation at boundaries that do not support CHAR/VARCHAR in UDT
storage.
+ */
+ private[sql] def physicalTypeHasCharVarchar(dt: DataType): Boolean = dt
match {
+ case ArrayType(elementType, _) => physicalTypeHasCharVarchar(elementType)
+ case MapType(keyType, valueType, _) =>
+ physicalTypeHasCharVarchar(keyType) ||
physicalTypeHasCharVarchar(valueType)
+ case StructType(fields) => fields.exists(f =>
physicalTypeHasCharVarchar(f.dataType))
+ case udt: UserDefinedType[_] => physicalTypeHasCharVarchar(udt.sqlType)
+ case _: CharType | _: VarcharType => true
+ case _ => false
+ }
+
+ /**
+ * Replaces CHAR/VARCHAR with their unconstrained string representation
regardless of session
Review Comment:
Narrowed the Scaladoc to document raw UDT storage and the requirement that
callers reject unsupported constrained UDT storage first.
##########
python/pyspark/sql/connect/session.py:
##########
@@ -501,6 +504,13 @@ def createDataFrame(
_num_cols = len(schema.fields)
else:
_num_cols = 1
+ if _has_physical_type(schema, (CharType, VarcharType)):
Review Comment:
Added the physical-type check after `_inferSchemaFromList` and before
`LocalDataToArrowConversion`, plus an inferred UDT rejection test.
##########
python/pyspark/sql/tests/test_udf.py:
##########
@@ -55,10 +63,198 @@
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.legacy.charVarcharAsString": "false",
+ "spark.sql.preserveCharVarcharTypeInfo": "false",
+ "spark.sql.charVarchar.standardSemantics.enabled": "false",
+ }
+ ):
+ default_result = self.spark.range(1).select(
+ udf(lambda _: "a", CharType(3),
useArrow=False)("id").alias("c")
+ )
+ self.assertEqual(default_result.schema["c"].dataType, StringType())
+ self.assertEqual(default_result.first().c, "a ")
+
+ 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 = [
Review Comment:
Completed the rejection matrix, including grouped-map iterator, stateful
grouped-map, and both window aggregate eval types. The focused matrix test
passes.
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala:
##########
@@ -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)
Review Comment:
Restricted `ArrayBasedMapBuilder` to active checks with logical CHAR/VARCHAR
in the key type. Other maps retain `ArrayBasedMapData`; a binary-key regression
verifies unchanged behavior.
##########
sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/util/CharVarcharUtils.scala:
##########
@@ -33,6 +33,43 @@ object CharVarcharUtils extends Logging with
SparkCharVarcharUtils {
// visible for testing
private[sql] val CHAR_VARCHAR_TYPE_STRING_METADATA_KEY =
"__CHAR_VARCHAR_TYPE_STRING"
+ /**
+ * Returns whether a type contains CHAR/VARCHAR, including inside UDT
storage. This is intended
+ * for validation at boundaries that do not support CHAR/VARCHAR in UDT
storage.
+ */
+ private[sql] def physicalTypeHasCharVarchar(dt: DataType): Boolean = dt
match {
+ case ArrayType(elementType, _) => physicalTypeHasCharVarchar(elementType)
+ case MapType(keyType, valueType, _) =>
+ physicalTypeHasCharVarchar(keyType) ||
physicalTypeHasCharVarchar(valueType)
+ case StructType(fields) => fields.exists(f =>
physicalTypeHasCharVarchar(f.dataType))
+ case udt: UserDefinedType[_] => physicalTypeHasCharVarchar(udt.sqlType)
+ case _: CharType | _: VarcharType => true
+ case _ => false
+ }
+
+ /**
+ * 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.
+ */
Review Comment:
Updated the Scaladoc as suggested: it now documents UTF-8, raw UDT
unwrapping, and the caller's rejection requirement.
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvaluatePython.scala:
##########
@@ -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 = {
Review Comment:
Made the generic `makeFromJava(DataType)` overload behavior-neutral by
passing `applyCharVarcharChecks = false`. Supported paths continue to pass
their captured policy explicitly.
##########
python/pyspark/sql/tests/test_udf.py:
##########
@@ -55,10 +63,198 @@
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.legacy.charVarcharAsString": "false",
+ "spark.sql.preserveCharVarcharTypeInfo": "false",
+ "spark.sql.charVarchar.standardSemantics.enabled": "false",
+ }
+ ):
+ default_result = self.spark.range(1).select(
+ udf(lambda _: "a", CharType(3),
useArrow=False)("id").alias("c")
+ )
+ self.assertEqual(default_result.schema["c"].dataType, StringType())
+ self.assertEqual(default_result.first().c, "a ")
+
+ 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):
Review Comment:
Added focused rejection coverage for row UDF UDT storage,
`applyInPandasWithState`, row-mode Connect UDTF DDL, Python DataSource UDT
storage, and `toArrow` UDT storage.
##########
python/pyspark/sql/udf.py:
##########
@@ -29,9 +29,16 @@
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,
Review Comment:
Removed the unused imports and switched the allow-list checks to the new
`_has_logical_type` helper.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]