Skip to content
10 changes: 10 additions & 0 deletions common/utils/src/main/resources/error/error-conditions.json
Original file line number Diff line number Diff line change
Expand Up @@ -9145,11 +9145,21 @@
"Purge table."
]
},
"PYTHON_STATE_CHAR_VARCHAR_SCHEMA" : {
"message" : [
"Python state does not support CHAR/VARCHAR in <schemaKind> schema <dataType>."
]
},
"PYTHON_UDF_IN_ON_CLAUSE" : {
"message" : [
"Python UDF in the ON clause of a <joinType> 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 <dataType>."
]
},
"QUERY_ONLY_CORRUPT_RECORD_COLUMN" : {
"message" : [
"Queries from raw JSON/CSV/XML files are disallowed when the",
Expand Down
36 changes: 32 additions & 4 deletions python/pyspark/sql/connect/udtf.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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":
Expand Down Expand Up @@ -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
)
Expand Down
8 changes: 5 additions & 3 deletions python/pyspark/sql/pandas/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@
BinaryType,
BooleanType,
ByteType,
CharType,
DataType,
DateType,
DayTimeIntervalType,
Expand All @@ -57,6 +58,7 @@
TimestampType,
TimeType,
UserDefinedType,
VarcharType,
VariantType,
VariantVal,
YearMonthIntervalType,
Expand Down Expand Up @@ -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)):
Comment thread
srielau marked this conversation as resolved.
Comment thread
srielau marked this conversation as resolved.
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()
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
76 changes: 76 additions & 0 deletions python/pyspark/sql/tests/arrow/test_arrow.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
BinaryType,
BooleanType,
ByteType,
CharType,
DateType,
DayTimeIntervalType,
DecimalType,
Expand All @@ -54,6 +55,7 @@
TimestampNTZType,
TimestampType,
TimeType,
VarcharType,
VariantType,
)
from pyspark.testing.objects import ExamplePoint, ExamplePointUDT
Expand Down Expand Up @@ -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):
Expand Down
92 changes: 80 additions & 12 deletions python/pyspark/sql/tests/arrow/test_arrow_python_udf.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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",
Comment thread
srielau marked this conversation as resolved.
"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):
Expand Down
Loading