Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions python/docs/source/reference/pyspark.sql/functions.rst
Original file line number Diff line number Diff line change
Expand Up @@ -611,6 +611,7 @@ VARIANT Functions
schema_of_variant
schema_of_variant_agg
try_variant_get
variant_array_length
variant_array_append
try_variant_array_append
variant_delete
Expand Down
7 changes: 7 additions & 0 deletions python/pyspark/sql/connect/functions/builtin.py
Original file line number Diff line number Diff line change
Expand Up @@ -2248,6 +2248,13 @@ def is_variant_null(v: "ColumnOrName") -> Column:
is_variant_null.__doc__ = pysparkfuncs.is_variant_null.__doc__


def variant_array_length(v: "ColumnOrName") -> Column:
return _invoke_function("variant_array_length", _to_col(v))


variant_array_length.__doc__ = pysparkfuncs.variant_array_length.__doc__


def is_valid_variant(v: "ColumnOrName") -> Column:
return _invoke_function("is_valid_variant", _to_col(v))

Expand Down
1 change: 1 addition & 0 deletions python/pyspark/sql/functions/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -486,6 +486,7 @@
# VARIANT Functions
"is_valid_variant",
"is_variant_null",
"variant_array_length",
"parse_json",
"schema_of_variant",
"schema_of_variant_agg",
Expand Down
31 changes: 31 additions & 0 deletions python/pyspark/sql/functions/builtin.py
Original file line number Diff line number Diff line change
Expand Up @@ -23025,6 +23025,37 @@ def is_variant_null(v: "ColumnOrName") -> Column:
return _invoke_function("is_variant_null", _to_java_column(v))


@_try_remote_functions
def variant_array_length(v: "ColumnOrName") -> Column:
"""
Returns the number of elements in a variant array. Returns NULL if the input is SQL NULL, a
variant null, or any non-array variant value.

.. versionadded:: 5.0.0

Parameters
----------
v : :class:`~pyspark.sql.Column` or str
a variant column or column name
A column that evaluates to a variant.

Returns
-------
:class:`~pyspark.sql.Column`
an integer column representing the array length, or NULL for non-array variant values
Returns a column that evaluates to an integer.

Examples
--------
>>> df = spark.createDataFrame([('''[1, 2, 3]''',), ('''{"a": 1}''',), ('null',)], ['json'])
>>> df.select(variant_array_length(parse_json(df.json)).alias("r")).collect()
[Row(r=3), Row(r=None), Row(r=None)]
"""
from pyspark.sql.classic.column import _to_java_column

return _invoke_function("variant_array_length", _to_java_column(v))


@_try_remote_functions
def is_valid_variant(v: "ColumnOrName") -> Column:
"""
Expand Down
2 changes: 2 additions & 0 deletions python/pyspark/sql/tests/test_functions.py
Original file line number Diff line number Diff line change
Expand Up @@ -3595,6 +3595,7 @@ def check(resultDf, expected):

check(df.select(F.is_variant_null(v)), [False, False])
check(df.select(F.is_valid_variant(v)), [True, True])
check(df.select(F.variant_array_length(v)), [None, None])
check(df.select(F.to_json(F.variant_delete(v, "$.a"))), ["{}", '{"b":2}'])
check(df.select(F.to_json(F.variant_delete(v, df.path))), ["{}", "{}"])
check(
Expand Down Expand Up @@ -3643,6 +3644,7 @@ def check(resultDf, expected):
['{"a":1,"z":9}', '{"b":2,"z":9}'],
)
arr = F.parse_json(df.arr)
check(df.select(F.variant_array_length(arr)), [2, 2])
check(
df.select(F.to_json(F.variant_array_append(arr, "$", F.lit(9)))),
["[1,2,9]", "[[3],4,9]"],
Expand Down
13 changes: 13 additions & 0 deletions sql/api/src/main/scala/org/apache/spark/sql/functions.scala
Original file line number Diff line number Diff line change
Expand Up @@ -14426,6 +14426,19 @@ object functions {
*/
def is_variant_null(v: Column): Column = Column.fn("is_variant_null", v)

/**
* Returns the number of elements in a variant array. Returns NULL if the input is SQL NULL, a
* variant null, or any non-array variant value.
*
* @param v
* a variant column. A column that evaluates to a variant.
* @group variant_funcs
* @since 5.0.0
* @return
* Returns a column that evaluates to an integer.
*/
def variant_array_length(v: Column): Column = Column.fn("variant_array_length", v)

/**
* Check if a variant value is valid. Returns true if the variant is valid, false if it is
* malformed, and NULL if the input is NULL.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -995,6 +995,7 @@ object FunctionRegistry {
expressionBuilder("parse_json", ParseJsonExpressionBuilder),
expressionBuilder("try_parse_json", TryParseJsonExpressionBuilder),
expression[IsVariantNull]("is_variant_null"),
expression[VariantArrayLength]("variant_array_length"),
expressionBuilder("variant_get", VariantGetExpressionBuilder),
expressionBuilder("try_variant_get", TryVariantGetExpressionBuilder),
expression[SchemaOfVariant]("schema_of_variant"),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,19 @@ object VariantExpressionEvalUtils {
def isValidVariant(input: VariantVal): Boolean =
VariantUtil.isValidVariant(input.getValue, input.getMetadata)

def variantArrayLength(input: VariantVal): Integer = {
if (input == null) {
null
} else {
val v = new Variant(input.getValue, input.getMetadata)
if (v.getType == VariantUtil.Type.ARRAY) {
v.arraySize()
} else {
null
}
}
}

/**
* Parse a JSONPath for a variant manipulation function. Throws `INVALID_VARIANT_PATH` on a
* malformed path, or on the empty (root `$`) path unless `allowRoot` is set (as
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -123,6 +123,50 @@ case class IsVariantNull(child: Expression) extends UnaryExpression
copy(child = newChild)
}

// scalastyle:off line.size.limit
@ExpressionDescription(
usage = "_FUNC_(expr) - Returns the number of elements in a variant array. Returns NULL if " +
"the input is SQL NULL, a variant null, or any non-array variant value.",
arguments = """
Arguments:
* expr - A variant value to inspect.
""",
examples = """
Examples:
> SELECT _FUNC_(parse_json('[1, 2, 3]'));
3
> SELECT _FUNC_(parse_json('{"a": 1}'));
NULL
> SELECT _FUNC_(parse_json('null'));
NULL
""",
since = "5.0.0",
group = "variant_funcs")
// scalastyle:on line.size.limit
case class VariantArrayLength(child: Expression) extends UnaryExpression
with ExpectsInputTypes with RuntimeReplaceable {

override lazy val replacement: Expression = StaticInvoke(
VariantExpressionEvalUtils.getClass,
IntegerType,
"variantArrayLength",
Seq(child),
inputTypes,
propagateNull = false,
returnNullable = true)

override def inputTypes: Seq[AbstractDataType] = Seq(VariantType)

override def dataType: DataType = IntegerType

override def nullable: Boolean = true

override def prettyName: String = "variant_array_length"

override protected def withNewChildInternal(newChild: Expression): VariantArrayLength =
copy(child = newChild)
}

// scalastyle:off line.size.limit
@ExpressionDescription(
usage = "_FUNC_(expr) - Convert a nested input (array/map/struct) into a variant where maps and structs are converted to variant objects which are unordered unlike SQL structs. Input maps can only have string keys.",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2801,6 +2801,10 @@ class PlanGenerationTestSuite extends ConnectFunSuite with Logging {
fn.is_variant_null(fn.parse_json(fn.col("g")))
}

functionTest("variant_array_length") {
fn.variant_array_length(fn.parse_json(fn.col("g")))
}

functionTest("is_valid_variant") {
fn.is_valid_variant(fn.parse_json(fn.col("g")))
}
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,2 @@
Project [static_invoke(VariantExpressionEvalUtils.variantArrayLength(static_invoke(VariantExpressionEvalUtils.parseJson(g#0, false, true, true)))) AS variant_array_length(parse_json(g))#0]
+- LocalRelation <empty>, [id#0L, a#0, b#0, d#0, e#0, f#0, g#0]
Original file line number Diff line number Diff line change
@@ -0,0 +1,83 @@
{
"common": {
"planId": "1"
},
"project": {
"input": {
"common": {
"planId": "0"
},
"localRelation": {
"schema": "struct\u003cid:bigint,a:int,b:double,d:struct\u003cid:bigint,a:int,b:double\u003e,e:array\u003cint\u003e,f:map\u003cstring,struct\u003cid:bigint,a:int,b:double\u003e\u003e,g:string\u003e"
}
},
"expressions": [{
"unresolvedFunction": {
"functionName": "variant_array_length",
"arguments": [{
"unresolvedFunction": {
"functionName": "parse_json",
"arguments": [{
"unresolvedAttribute": {
"unparsedIdentifier": "g"
},
"common": {
"origin": {
"jvmOrigin": {
"stackTrace": [{
"classLoaderName": "app",
"declaringClass": "org.apache.spark.sql.functions$",
"methodName": "col",
"fileName": "functions.scala"
}, {
"classLoaderName": "app",
"declaringClass": "org.apache.spark.sql.PlanGenerationTestSuite",
"methodName": "~~trimmed~anonfun~~",
"fileName": "PlanGenerationTestSuite.scala"
}]
}
}
}
}],
"isInternal": false
},
"common": {
"origin": {
"jvmOrigin": {
"stackTrace": [{
"classLoaderName": "app",
"declaringClass": "org.apache.spark.sql.functions$",
"methodName": "parse_json",
"fileName": "functions.scala"
}, {
"classLoaderName": "app",
"declaringClass": "org.apache.spark.sql.PlanGenerationTestSuite",
"methodName": "~~trimmed~anonfun~~",
"fileName": "PlanGenerationTestSuite.scala"
}]
}
}
}
}],
"isInternal": false
},
"common": {
"origin": {
"jvmOrigin": {
"stackTrace": [{
"classLoaderName": "app",
"declaringClass": "org.apache.spark.sql.functions$",
"methodName": "variant_array_length",
"fileName": "functions.scala"
}, {
"classLoaderName": "app",
"declaringClass": "org.apache.spark.sql.PlanGenerationTestSuite",
"methodName": "~~trimmed~anonfun~~",
"fileName": "PlanGenerationTestSuite.scala"
}]
}
}
}
}]
}
}
Binary file not shown.
Original file line number Diff line number Diff line change
Expand Up @@ -575,6 +575,7 @@
| org.apache.spark.sql.catalyst.expressions.variant.TryVariantInsertExpressionBuilder | try_variant_insert | SELECT try_variant_insert(parse_json('{"a": 1}'), '$.b', 2) | struct<try_variant_insert(parse_json({"a": 1}), $.b, 2):variant> |
| org.apache.spark.sql.catalyst.expressions.variant.TryVariantSetExpressionBuilder | try_variant_set | SELECT try_variant_set(parse_json('{"a": 1}'), '$.a', 2) | struct<try_variant_set(parse_json({"a": 1}), $.a, 2, true):variant> |
| org.apache.spark.sql.catalyst.expressions.variant.VariantArrayAppendExpressionBuilder | variant_array_append | SELECT variant_array_append(parse_json('[1, 2, 3]'), '$', 4) | struct<variant_array_append(parse_json([1, 2, 3]), $, 4):variant> |
| org.apache.spark.sql.catalyst.expressions.variant.VariantArrayLength | variant_array_length | SELECT variant_array_length(parse_json('[1, 2, 3]')) | struct<variant_array_length(parse_json([1, 2, 3])):int> |
| org.apache.spark.sql.catalyst.expressions.variant.VariantDelete | variant_delete | SELECT variant_delete(parse_json('{"a": 1, "b": 2, "c": 3, "items": [1, 2, 3]}'), NULL, '$.a', '$.c') | struct<variant_delete(parse_json({"a": 1, "b": 2, "c": 3, "items": [1, 2, 3]}), NULL, $.a, $.c):variant> |
| org.apache.spark.sql.catalyst.expressions.variant.VariantFromArrays | variant_from_arrays | SELECT variant_from_arrays(array('a', 'b'), array(1, 2)) | struct<variant_from_arrays(array(a, b), array(1, 2)):variant> |
| org.apache.spark.sql.catalyst.expressions.variant.VariantFromEntries | variant_from_entries | SELECT variant_from_entries(array(struct('a', 1), struct('b', 2))) | struct<variant_from_entries(array(struct(a, 1), struct(b, 2))):variant> |
Expand All @@ -590,4 +591,4 @@
| org.apache.spark.sql.catalyst.expressions.xml.XPathList | xpath | SELECT xpath('<a><b>b1</b><b>b2</b><b>b3</b><c>c1</c><c>c2</c></a>','a/b/text()') | struct<xpath(<a><b>b1</b><b>b2</b><b>b3</b><c>c1</c><c>c2</c></a>, a/b/text()):array<string>> |
| org.apache.spark.sql.catalyst.expressions.xml.XPathLong | xpath_long | SELECT xpath_long('<a><b>1</b><b>2</b></a>', 'sum(a/b)') | struct<xpath_long(<a><b>1</b><b>2</b></a>, sum(a/b)):bigint> |
| org.apache.spark.sql.catalyst.expressions.xml.XPathShort | xpath_short | SELECT xpath_short('<a><b>1</b><b>2</b></a>', 'sum(a/b)') | struct<xpath_short(<a><b>1</b><b>2</b></a>, sum(a/b)):smallint> |
| org.apache.spark.sql.catalyst.expressions.xml.XPathString | xpath_string | SELECT xpath_string('<a><b>b</b><c>cc</c></a>','a/c') | struct<xpath_string(<a><b>b</b><c>cc</c></a>, a/c):string> |
| org.apache.spark.sql.catalyst.expressions.xml.XPathString | xpath_string | SELECT xpath_string('<a><b>b</b><c>cc</c></a>','a/c') | struct<xpath_string(<a><b>b</b><c>cc</c></a>, a/c):string> |
Original file line number Diff line number Diff line change
Expand Up @@ -342,6 +342,14 @@ class VariantEndToEndSuite extends SharedSparkSession {
check("{ \"a\": null, \"b\": {\"c\": null, \"d\": [13, null]} }", "$.b.d[2]", expected = false)
}

test("variant_array_length") {
val df = Seq("[1, 2, 3]", "[]", """{"a": 1}""", "null", null).toDF("j")
val expected = Seq(Row(3), Row(0), Row(null), Row(null), Row(null))

checkAnswer(df.selectExpr("variant_array_length(parse_json(j))"), expected)
checkAnswer(df.select(variant_array_length(parse_json($"j"))), expected)
}

test("schema_of_variant_agg") {
// Literal input.
checkAnswer(
Expand Down