diff --git a/python/docs/source/reference/pyspark.sql/functions.rst b/python/docs/source/reference/pyspark.sql/functions.rst index 94c5a9a87e13a..379393e26e57d 100644 --- a/python/docs/source/reference/pyspark.sql/functions.rst +++ b/python/docs/source/reference/pyspark.sql/functions.rst @@ -454,6 +454,7 @@ Aggregate Functions bit_and bit_or bit_xor + bitmap_and_agg bitmap_construct_agg bitmap_or_agg bitmap_xor_agg @@ -461,6 +462,7 @@ Aggregate Functions bool_or collect_list collect_set + collect_union corr count count_distinct diff --git a/python/pyspark/sql/functions/__init__.py b/python/pyspark/sql/functions/__init__.py index 861f5902f977a..36750573d5846 100644 --- a/python/pyspark/sql/functions/__init__.py +++ b/python/pyspark/sql/functions/__init__.py @@ -21,7 +21,7 @@ from pyspark.sql.functions.builtin import * # noqa: F403 __all__ = [ # noqa: F405 - # Normal functions + # Normal Functions "broadcast", "call_function", "col", @@ -620,7 +620,7 @@ "vector_normalize", "vector_avg", "vector_sum", - # Call Functions + # UDF, UDTF and UDT "call_udf", "pandas_udf", "udaf", diff --git a/python/pyspark/sql/functions/builtin.py b/python/pyspark/sql/functions/builtin.py index e54927c36e651..f22235f6d80df 100644 --- a/python/pyspark/sql/functions/builtin.py +++ b/python/pyspark/sql/functions/builtin.py @@ -109,6 +109,8 @@ # even though there might be few exceptions for legacy or inevitable reasons. # If you are fixing other language APIs together, also please note that Scala side is not the case # since it requires making every single overridden definition. +# Public function groups are defined by pyspark.sql.functions.__all__ and mirrored in the API +# reference. def _get_jvm_function(name: str, sc: "SparkContext") -> Callable: @@ -9213,9 +9215,6 @@ def factorial(col: "ColumnOrName") -> Column: return _invoke_function_over_columns("factorial", col) -# --------------- Window functions ------------------------ - - @_try_remote_functions def lag(col: "ColumnOrName", offset: int = 1, default: Optional[Any] = None) -> Column: """ @@ -9816,9 +9815,6 @@ def ntile(n: int) -> Column: return _invoke_function("ntile", int(_enum_to_value(n))) -# ---------------------- Date/Timestamp functions ------------------------------ - - @_try_remote_functions def curdate() -> Column: """ @@ -14551,9 +14547,6 @@ def to_timestamp_ntz( return _invoke_function_over_columns("to_timestamp_ntz", timestamp) -# ---------------------------- misc functions ---------------------------------- - - @_try_remote_functions def current_catalog() -> Column: """Returns the current catalog. @@ -15237,9 +15230,6 @@ def raise_error(errMsg: Union[Column, str]) -> Column: return _invoke_function_over_columns("raise_error", lit(errMsg)) -# ---------------------- String/Binary functions ------------------------------ - - @_try_remote_functions def upper(col: "ColumnOrName") -> Column: """ @@ -19939,9 +19929,6 @@ def quote(col: "ColumnOrName") -> Column: return _invoke_function_over_columns("quote", col) -# ---------------------- Collection functions ------------------------------ - - @overload def create_map(*cols: "ColumnOrName") -> Column: ... @@ -26848,9 +26835,6 @@ def str_to_map( return _invoke_function_over_columns("str_to_map", text, pairDelim, keyValueDelim) -# ---------------------- Partition transform functions -------------------------------- - - @_try_remote_functions def years(col: "ColumnOrName") -> Column: """ @@ -28910,9 +28894,6 @@ def bucket(numBuckets: Union[Column, int], col: "ColumnOrName") -> Column: return partitioning.bucket(numBuckets, col) -# Geospatial ST Functions - - @_try_remote_functions def st_asbinary(geo: "ColumnOrName", endianness: Optional["ColumnOrName"] = None) -> Column: """Returns the input GEOGRAPHY or GEOMETRY value in WKB format. @@ -29080,9 +29061,6 @@ def st_srid(geo: "ColumnOrName") -> Column: return _invoke_function_over_columns("st_srid", geo) -# Call Functions - - @_try_remote_functions def call_udf(udfName: str, *cols: "ColumnOrName") -> Column: """ @@ -29369,9 +29347,6 @@ def wrap_udt(col: "ColumnOrName", udt: "Union[UserDefinedType, Column]") -> Colu return _invoke_function("wrap_udt", _to_java_column(col), _to_java_column(udt_col)) -# ---------------------- Datasketch functions ------------------------------ - - @_try_remote_functions def hll_sketch_agg( col: "ColumnOrName", @@ -31915,9 +31890,6 @@ def tuple_union_theta_integer( return _invoke_function_over_columns(fn, col1, col2, _lgNomEntries, _mode) -# ---------------------- Predicates functions ------------------------------ - - @_try_remote_functions def ifnull(col1: "ColumnOrName", col2: "ColumnOrName") -> Column: """ @@ -33462,9 +33434,6 @@ def bitmap_xor_agg(col: "ColumnOrName") -> Column: return _invoke_function_over_columns("bitmap_xor_agg", col) -# ---------------------------- User Defined Function ---------------------------------- - - def udaf(agg: "Aggregator") -> "UserDefinedFunctionLike": """Turn an :class:`~pyspark.sql.aggregator.Aggregator` instance into a callable usable in ``groupBy().agg(...)`` (and as a window function), the Python counterpart of Scala's @@ -34038,9 +34007,6 @@ def arrow_udtf( return _create_pyarrow_udtf(cls=cls, returnType=returnType) -# ---------------------- Vector Functions ---------------------- - - @_try_remote_functions def vector_cosine_similarity(left: "ColumnOrName", right: "ColumnOrName") -> Column: """Returns the cosine similarity between two float vectors. diff --git a/sql/api/src/main/scala/org/apache/spark/sql/functions.scala b/sql/api/src/main/scala/org/apache/spark/sql/functions.scala index 307a82e5a44a6..589edde9c171c 100644 --- a/sql/api/src/main/scala/org/apache/spark/sql/functions.scala +++ b/sql/api/src/main/scala/org/apache/spark/sql/functions.scala @@ -53,33 +53,33 @@ import org.apache.spark.util.SparkClassUtils * only `Column` but also other types such as a native string. The other variants currently exist * for historical reasons. * - * @groupname udf_funcs UDF, UDAF and UDT - * @groupname agg_funcs Aggregate functions - * @groupname datetime_funcs Date and Timestamp functions - * @groupname sort_funcs Sort functions * @groupname normal_funcs Normal functions + * @groupname conditional_funcs Conditional functions + * @groupname predicate_funcs Predicate functions + * @groupname sort_funcs Sort functions * @groupname math_funcs Mathematical functions + * @groupname string_funcs String functions * @groupname bitwise_funcs Bitwise functions - * @groupname predicate_funcs Predicate functions - * @groupname conditional_funcs Conditional functions + * @groupname datetime_funcs Date and Timestamp functions * @groupname hash_funcs Hash functions - * @groupname misc_funcs Misc functions - * @groupname sketch_funcs Datasketch functions - * @groupname window_funcs Window functions - * @groupname generator_funcs Generator functions - * @groupname string_funcs String functions * @groupname collection_funcs Collection functions * @groupname array_funcs Array functions - * @groupname map_funcs Map functions * @groupname struct_funcs Struct functions - * @groupname st_funcs ST geospatial functions + * @groupname map_funcs Map functions + * @groupname agg_funcs Aggregate functions + * @groupname window_funcs Window functions + * @groupname generator_funcs Generator functions + * @groupname partition_transforms Partition transform functions * @groupname csv_funcs CSV functions * @groupname json_funcs JSON functions * @groupname variant_funcs VARIANT functions - * @groupname vector_funcs Vector functions * @groupname xml_funcs XML functions * @groupname url_funcs URL functions - * @groupname partition_transforms Partition transform functions + * @groupname misc_funcs Misc functions + * @groupname sketch_funcs Datasketch functions + * @groupname st_funcs ST geospatial functions + * @groupname vector_funcs Vector functions + * @groupname udf_funcs UDF, UDAF and UDT * @groupname Ungrouped Support functions for DataFrames * @since 1.3.0 */ @@ -88,6 +88,9 @@ import org.apache.spark.util.SparkClassUtils object functions { // scalastyle:on + // Function groups are defined by the @group tags above each function and the corresponding + // @groupname declarations. + /** * Returns a [[Column]] based on the given column name. * @@ -171,10 +174,6 @@ object functions { } } - ////////////////////////////////////////////////////////////////////////////////////////////// - // Sort functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Returns a sort expression based on ascending order of the column. * {{{ @@ -245,10 +244,6 @@ object functions { */ def desc_nulls_last(columnName: String): Column = Column(columnName).desc_nulls_last - ////////////////////////////////////////////////////////////////////////////////////////////// - // Aggregate functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * @group agg_funcs * @since 1.3.0 @@ -3829,10 +3824,6 @@ object functions { */ def bit_xor(e: Column): Column = Column.fn("bit_xor", e) - ////////////////////////////////////////////////////////////////////////////////////////////// - // Window functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Window function: computes the differences between consecutive cumulative counter values in a * time series, thereby converting the counter from the cumulative to the delta format. @@ -4243,10 +4234,6 @@ object functions { */ def row_number(): Column = Column.fn("row_number") - ////////////////////////////////////////////////////////////////////////////////////////////// - // Non-aggregate functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Creates a new array column. The input columns must all have the same data type. * @@ -4686,7 +4673,7 @@ object functions { * * @param e * the value to compute the mean of. A column that evaluates to a numeric or interval. - * @group math_funcs + * @group agg_funcs * @since 3.5.0 * @return * Returns a column that evaluates to a double. @@ -4757,7 +4744,7 @@ object functions { * * @param e * the value to compute the sum of. A column that evaluates to a numeric or interval. - * @group math_funcs + * @group agg_funcs * @since 3.5.0 * @return * Returns a column that evaluates to a numeric. @@ -4908,10 +4895,6 @@ object functions { */ def expr(expr: String): Column = Column(internal.SqlExpression(expr)) - ////////////////////////////////////////////////////////////////////////////////////////////// - // Math Functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Computes the absolute value of a numeric value. * @@ -6499,10 +6482,6 @@ object functions { def width_bucket(v: Column, min: Column, max: Column, numBucket: Column): Column = Column.fn("width_bucket", v, min, max, numBucket) - ////////////////////////////////////////////////////////////////////////////////////////////// - // Misc functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Returns the current catalog. * @@ -7465,10 +7444,6 @@ object functions { */ def bitmap_xor_agg(col: Column): Column = Column.fn("bitmap_xor_agg", col) - ////////////////////////////////////////////////////////////////////////////////////////////// - // String functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Computes the numeric value of the first character of the string column, and returns the * result as an int column. @@ -9512,10 +9487,6 @@ object functions { */ def quote(str: Column): Column = Column.fn("quote", str) - ////////////////////////////////////////////////////////////////////////////////////////////// - // Datasketch functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Returns the estimated number of unique values given the binary representation of a * Datasketches HllSketch. @@ -11515,10 +11486,6 @@ object functions { def kll_sketch_get_rank_double(sketch: Column, quantile: Column): Column = Column.fn("kll_sketch_get_rank_double", sketch, quantile) - ////////////////////////////////////////////////////////////////////////////////////////////// - // DateTime functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Returns the date that is `numMonths` after `startDate`. * @@ -13229,10 +13196,6 @@ object functions { def dayname(timeExp: Column): Column = Column.fn("dayname", timeExp) - ////////////////////////////////////////////////////////////////////////////////////////////// - // Collection functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Returns true if the array contains `value`, false if not. Returns null if the array or * `value` is null, or if `value` is not found and the array contains a null element. @@ -16976,10 +16939,6 @@ object functions { */ def bucket(numBuckets: Int, e: Column): Column = partitioning.bucket(numBuckets, e) - ////////////////////////////////////////////////////////////////////////////////////////////// - // Predicates functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Returns `col2` if `col1` is null, or `col1` otherwise. * @@ -17133,10 +17092,6 @@ object functions { */ - ////////////////////////////////////////////////////////////////////////////////////////////// - // ST geospatial functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Returns the input GEOGRAPHY or GEOMETRY value in WKB format. * @@ -17275,10 +17230,6 @@ object functions { def st_srid(geo: Column): Column = Column.fn("st_srid", geo) - ////////////////////////////////////////////////////////////////////////////////////////////// - // Scala UDF functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Obtains a `UserDefinedFunction` that wraps the given `Aggregator` so that it may be used with * untyped Data Frames. @@ -17621,10 +17572,6 @@ object functions { implicitly[TypeTag[A10]]) } - ////////////////////////////////////////////////////////////////////////////////////////////// - // Java UDF functions - ////////////////////////////////////////////////////////////////////////////////////////////// - /** * Defines a Java UDF0 instance as user-defined function (UDF). The caller must specify the * output data type, and there is no automatic input type coercion. By default the returned UDF @@ -17882,8 +17829,6 @@ object functions { Column.internalFn("wrap_udt", column, udt) } - // ---------------------- Vector Functions ---------------------- - /** * Returns the cosine similarity between two float vectors. * @param left diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala index 147a980eeb098..0ed6cee205da7 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala @@ -405,40 +405,55 @@ object FunctionRegistry { // anymore. See `AesEncrypt`/`AesDecrypt` as an example. private type FunctionRegistryEntry = (String, (ExpressionInfo, FunctionBuilder)) - private def miscNonAggregateExpressions: Seq[FunctionRegistryEntry] = Seq( - // misc non-aggregate functions - expression[Abs]("abs"), + private def conditionalExpressions: Seq[FunctionRegistryEntry] = Seq( + // conditional functions expression[Coalesce]("coalesce"), - expressionBuilder("explode", ExplodeExpressionBuilder), - expressionGeneratorBuilderOuter("explode_outer", ExplodeExpressionBuilder), - expression[Greatest]("greatest"), expression[If]("if"), - expressionBuilder("inline", InlineExpressionBuilder), - expressionGeneratorBuilderOuter("inline_outer", InlineExpressionBuilder), - expression[IsNaN]("isnan"), expression[Nvl]("ifnull", setAlias = true), - expression[IsNull]("isnull"), - expression[IsNotNull]("isnotnull"), - expression[Least]("least"), expression[NaNvl]("nanvl"), expression[NullIf]("nullif"), expression[NullIfZero]("nullifzero"), expression[Nvl]("nvl"), expression[Nvl2]("nvl2"), - expressionBuilder("posexplode", PosExplodeExpressionBuilder), - expressionGeneratorBuilderOuter("posexplode_outer", PosExplodeExpressionBuilder), - expression[Rand]("rand"), - expression[Rand]("random", true, Some("3.0.0")), - expression[Randn]("randn"), - expression[RandStr]("randstr"), - expression[Stack]("stack"), - expression[Uniform]("uniform"), expression[ZeroIfNull]("zeroifnull"), - CaseWhen.registryEntry + CaseWhen.registryEntry, + expression[Between]("between") + ) + + private def predicateExpressions: Seq[FunctionRegistryEntry] = Seq( + // predicate functions + expression[IsNaN]("isnan"), + expression[IsNull]("isnull"), + expression[IsNotNull]("isnotnull"), + expression[Like]("like"), + expression[ILike]("ilike"), + expression[RLike]("rlike"), + expression[RLike]("regexp_like", true, Some("3.2.0")), + expression[RLike]("regexp", true, Some("3.2.0")), + expression[EqualNull]("equal_null"), + expression[And]("and"), + expression[In]("in"), + expression[Not]("not"), + expression[Or]("or"), + expression[EqualNullSafe]("<=>"), + expression[EqualTo]("="), + expression[EqualTo]("=="), + expression[GreaterThan](">"), + expression[GreaterThanOrEqual](">="), + expression[LessThan]("<"), + expression[LessThanOrEqual]("<="), + expression[Not]("!") ) private def mathExpressions: Seq[FunctionRegistryEntry] = Seq( // math functions + expression[Abs]("abs"), + expression[Greatest]("greatest"), + expression[Least]("least"), + expression[Rand]("rand"), + expression[Rand]("random", true, Some("3.0.0")), + expression[Randn]("randn"), + expression[Uniform]("uniform"), expression[Acos]("acos"), expression[Acosh]("acosh"), expression[Asin]("asin"), @@ -479,142 +494,34 @@ object FunctionRegistry { expression[Rint]("rint"), expression[Round]("round"), expression[Truncate]("truncate"), - expression[ShiftLeft]("shiftleft"), - expression[ShiftRight]("shiftright"), - expression[ShiftRightUnsigned]("shiftrightunsigned"), expression[Signum]("sign", true), expression[Signum]("signum"), expression[Sin]("sin"), expression[Csc]("csc"), expression[Sinh]("sinh"), - expression[StringToMap]("str_to_map"), expression[Sqrt]("sqrt"), expression[Tan]("tan"), expression[Cot]("cot"), expression[Tanh]("tanh"), expression[WidthBucket]("width_bucket"), - expression[Add]("+"), expression[Subtract]("-"), expression[Multiply]("*"), expression[Divide]("/"), expression[IntegralDivide]("div"), - expression[Remainder]("%") - ) - - private def tryExpressions: Seq[FunctionRegistryEntry] = Seq( - // "try_*" function which always return Null instead of runtime error. + expression[Remainder]("%"), expression[TryAdd]("try_add"), expression[TryDivide]("try_divide"), expression[TryMod]("try_mod"), expression[TrySubtract]("try_subtract"), expression[TryMultiply]("try_multiply"), - expression[TryElementAt]("try_element_at"), - expressionBuilder("try_avg", TryAverageExpressionBuilder, setAlias = true), - expressionBuilder("try_sum", TrySumExpressionBuilder, setAlias = true), - expression[TryToBinary]("try_to_binary"), - expressionBuilder("try_to_timestamp", TryToTimestampExpressionBuilder, setAlias = true), - expressionBuilder("try_to_date", TryToDateExpressionBuilder, setAlias = true), - expressionBuilder("try_to_time", TryToTimeExpressionBuilder, setAlias = true), - expression[TryAesDecrypt]("try_aes_decrypt"), - expression[TryReflect]("try_reflect"), - expression[TryUrlDecode]("try_url_decode"), - expression[TryMakeInterval]("try_make_interval") - ) - - private def aggregateExpressions: Seq[FunctionRegistryEntry] = Seq( - // aggregate functions - expression[HyperLogLogPlusPlus]("approx_count_distinct"), - expression[Average]("avg"), - expression[Corr]("corr"), - expression[Count]("count"), - expression[CountIf]("count_if"), - expression[CovPopulation]("covar_pop"), - expression[CovSample]("covar_samp"), - expression[First]("first"), - expression[First]("first_value", true), - expression[AnyValue]("any_value"), - expression[Kurtosis]("kurtosis"), - expression[Last]("last"), - expression[Last]("last_value", true), - expression[Max]("max"), - expressionBuilder("max_by", MaxByBuilder), - expression[Average]("mean", true), - expression[Min]("min"), - expressionBuilder("min_by", MinByBuilder), - expression[Percentile]("percentile"), - expressionBuilder("percentile_cont", PercentileContBuilder), - expressionBuilder("percentile_disc", PercentileDiscBuilder), - expression[Median]("median"), - expression[Skewness]("skewness"), - expression[ApproximatePercentile]("percentile_approx"), - expression[ApproximatePercentile]("approx_percentile", true), - expression[HistogramNumeric]("histogram_numeric"), - expression[StddevSamp]("std", true), - expression[StddevSamp]("stddev", true), - expression[StddevPop]("stddev_pop"), - expression[StddevSamp]("stddev_samp"), - expression[Sum]("sum"), - expression[VarianceSamp]("variance", true), - expression[VariancePop]("var_pop"), - expression[VarianceSamp]("var_samp"), - expression[CollectList]("collect_list"), - expression[CollectList]("array_agg", true, Some("3.3.0")), - expression[CollectSet]("collect_set"), - expression[CollectUnion]("collect_union"), - expression[ListAgg]("listagg"), - expression[ListAgg]("string_agg", setAlias = true), - expressionBuilder("count_min_sketch", CountMinSketchAggExpressionBuilder), - expression[BoolAnd]("every", true), - expression[BoolAnd]("bool_and"), - expression[BoolOr]("any", true), - expression[BoolOr]("some", true), - expression[BoolOr]("bool_or"), - expression[RegrCount]("regr_count"), - expression[RegrAvgX]("regr_avgx"), - expression[RegrAvgY]("regr_avgy"), - expression[RegrR2]("regr_r2"), - expression[RegrSXX]("regr_sxx"), - expression[RegrSXY]("regr_sxy"), - expression[RegrSYY]("regr_syy"), - expression[RegrSlope]("regr_slope"), - expression[RegrIntercept]("regr_intercept"), - expressionBuilder("mode", ModeBuilder), - expression[HllSketchAgg]("hll_sketch_agg"), - expression[HllUnionAgg]("hll_union_agg"), - expression[ApproxTopK]("approx_top_k"), - expression[ThetaSketchAgg]("theta_sketch_agg"), - expression[ThetaUnionAgg]("theta_union_agg"), - expression[ThetaIntersectionAgg]("theta_intersection_agg"), - expression[ApproxTopKAccumulate]("approx_top_k_accumulate"), - expression[ApproxTopKCombine]("approx_top_k_combine"), - expression[KllSketchAggBigint]("kll_sketch_agg_bigint"), - expression[KllSketchAggFloat]("kll_sketch_agg_float"), - expression[KllSketchAggDouble]("kll_sketch_agg_double"), - expression[KllMergeAggBigint]("kll_merge_agg_bigint"), - expression[KllMergeAggFloat]("kll_merge_agg_float"), - expression[KllMergeAggDouble]("kll_merge_agg_double"), - expression[TupleIntersectionAggDouble]("tuple_intersection_agg_double"), - expression[TupleIntersectionAggInteger]("tuple_intersection_agg_integer"), - expressionBuilder("tuple_sketch_agg_double", TupleSketchAggDoubleExpressionBuilder), - expressionBuilder("tuple_sketch_agg_integer", TupleSketchAggIntegerExpressionBuilder), - expressionBuilder("tuple_union_agg_double", TupleUnionAggDoubleExpressionBuilder), - expressionBuilder("tuple_union_agg_integer", TupleUnionAggIntegerExpressionBuilder) - ) - - private def vectorExpressions: Seq[FunctionRegistryEntry] = Seq( - // vector functions - expression[VectorCosineSimilarity]("vector_cosine_similarity"), - expression[VectorInnerProduct]("vector_inner_product"), - expression[VectorL2Distance]("vector_l2_distance"), - expression[VectorNorm]("vector_norm"), - expression[VectorNormalize]("vector_normalize"), - expression[VectorAvg]("vector_avg"), - expression[VectorSum]("vector_sum") + expression[Unhex]("unhex") ) private def stringExpressions: Seq[FunctionRegistryEntry] = Seq( // string functions + expression[RandStr]("randstr"), + expression[TryToBinary]("try_to_binary"), expression[Ascii]("ascii"), expression[Chr]("char", true), expression[Chr]("chr"), @@ -639,7 +546,6 @@ object FunctionRegistry { expression[TryToNumber]("try_to_number"), expressionBuilder("to_char", ToCharacterBuilder), expressionBuilder("to_varchar", ToCharacterBuilder, setAlias = true, Some("3.5.0")), - expression[GetJsonObject]("get_json_object"), expression[InitCap]("initcap"), expressionBuilder("instr", StringInstrExpressionBuilder), expression[Lower]("lcase", true), @@ -648,14 +554,11 @@ object FunctionRegistry { expression[Levenshtein]("levenshtein"), expression[JaroWinkler]("jaro_winkler_similarity"), expression[Luhncheck]("luhn_check"), - expression[Like]("like"), - expression[ILike]("ilike"), expression[Lower]("lower"), expression[OctetLength]("octet_length"), expression[StringLocate]("locate"), expressionBuilder("lpad", LPadExpressionBuilder), expression[StringTrimLeft]("ltrim"), - expression[JsonTuple]("json_tuple"), expression[StringLocate]("position", true, Some("2.3.0")), expression[FormatString]("printf", true), expression[RegExpExtract]("regexp_extract"), @@ -664,9 +567,6 @@ object FunctionRegistry { expression[StringRepeat]("repeat"), expression[StringReplace]("replace"), expression[Overlay]("overlay"), - expression[RLike]("rlike"), - expression[RLike]("regexp_like", true, Some("3.2.0")), - expression[RLike]("regexp", true, Some("3.2.0")), expressionBuilder("rpad", RPadExpressionBuilder), expression[StringTrimRight]("rtrim"), expression[Sentences]("sentences"), @@ -685,17 +585,7 @@ object FunctionRegistry { expression[Upper]("ucase", true), expression[UnBase64]("unbase64"), expression[UnBase32]("from_base32"), - expression[Unhex]("unhex"), expression[Upper]("upper"), - expression[XPathList]("xpath"), - expression[XPathBoolean]("xpath_boolean"), - expression[XPathDouble]("xpath_double"), - expression[XPathDouble]("xpath_number", true), - expression[XPathFloat]("xpath_float"), - expression[XPathInt]("xpath_int"), - expression[XPathLong]("xpath_long"), - expression[XPathShort]("xpath_short"), - expression[XPathString]("xpath_string"), expression[RegExpCount]("regexp_count"), expression[RegExpSubStr]("regexp_substr"), expression[RegExpInStr]("regexp_instr"), @@ -704,19 +594,34 @@ object FunctionRegistry { expression[ValidateUTF8]("validate_utf8"), expression[TryValidateUTF8]("try_validate_utf8"), expression[Quote]("quote"), - expression[Normalize]("normalize") + expression[Normalize]("normalize"), + expression[ToBinary]("to_binary"), + expressionBuilder("mask", MaskExpressionBuilder) ) - private def urlExpressions: Seq[FunctionRegistryEntry] = Seq( - // url functions - expression[UrlEncode]("url_encode"), - expression[UrlDecode]("url_decode"), - expression[ParseUrl]("parse_url"), - expression[TryParseUrl]("try_parse_url") + private def bitwiseExpressions: Seq[FunctionRegistryEntry] = Seq( + // bitwise functions + expression[ShiftLeft]("shiftleft"), + expression[ShiftRight]("shiftright"), + expression[ShiftRightUnsigned]("shiftrightunsigned"), + expression[BitwiseAnd]("&"), + expression[BitwiseNot]("~"), + expression[BitwiseOr]("|"), + expression[BitwiseXor]("^"), + expression[ShiftLeft]("<<", true, Some("4.0.0")), + expression[ShiftRight](">>", true, Some("4.0.0")), + expression[ShiftRightUnsigned](">>>", true, Some("4.0.0")), + expression[BitwiseCount]("bit_count"), + expression[BitwiseGet]("bit_get"), + expression[BitwiseGet]("getbit", true) ) private def datetimeExpressions: Seq[FunctionRegistryEntry] = Seq( // datetime functions + expressionBuilder("try_to_timestamp", TryToTimestampExpressionBuilder, setAlias = true), + expressionBuilder("try_to_date", TryToDateExpressionBuilder, setAlias = true), + expressionBuilder("try_to_time", TryToTimeExpressionBuilder, setAlias = true), + expression[TryMakeInterval]("try_make_interval"), expression[AddMonths]("add_months"), expression[CurrentDate]("current_date"), expressionBuilder("curdate", CurDateExpressionBuilder, setAlias = true), @@ -749,7 +654,6 @@ object FunctionRegistry { expression[ParseToDate]("to_date"), expression[TimeDiff]("time_diff"), expression[ToTime]("to_time"), - expression[ToBinary]("to_binary"), expression[ToUnixTimestamp]("to_unix_timestamp"), expression[ToUTCTimestamp]("to_utc_timestamp"), // We keep the 2 expression builders below to have different function docs. @@ -806,8 +710,47 @@ object FunctionRegistry { expressionBuilder("time_bucket", TimeBucketExpressionBuilder) ) + private def hashExpressions: Seq[FunctionRegistryEntry] = Seq( + // hash functions + expression[Crc32]("crc32"), + expression[Md5]("md5"), + expression[Murmur3Hash]("hash"), + expression[XxHash64]("xxhash64"), + expression[Xxh364]("xxh3_64"), + expression[Xxh3128]("xxh3_128"), + expression[Sha1]("sha", true), + expression[Sha1]("sha1"), + expression[Sha2]("sha2") + ) + private def collectionExpressions: Seq[FunctionRegistryEntry] = Seq( // collection functions + expression[TryElementAt]("try_element_at"), + expression[ElementAt]("element_at"), + expression[Size]("size"), + expression[Size]("cardinality", true, Some("2.4.0")), + expression[Reverse]("reverse"), + expression[Concat]("concat") + ) + + private def lambdaExpressions: Seq[FunctionRegistryEntry] = Seq( + // lambda functions + expression[ArraySort]("array_sort"), + expression[ArrayTransform]("transform"), + expression[MapFilter]("map_filter"), + expression[ArrayFilter]("filter"), + expression[ArrayExists]("exists"), + expression[ArrayForAll]("forall"), + expression[ArrayAggregate]("aggregate"), + expression[ArrayAggregate]("reduce", setAlias = true, Some("3.4.0")), + expression[TransformValues]("transform_values"), + expression[TransformKeys]("transform_keys"), + expression[MapZipWith]("map_zip_with"), + expression[ZipWith]("zip_with") + ) + + private def arrayExpressions: Seq[FunctionRegistryEntry] = Seq( + // array functions expression[CreateArray]("array"), expression[ArrayContains]("array_contains"), expression[ArraysOverlap]("arrays_overlap"), @@ -816,68 +759,255 @@ object FunctionRegistry { expression[ArrayJoin]("array_join"), expression[ArrayPosition]("array_position"), expression[ArraySize]("array_size"), - expression[ArraySort]("array_sort"), expression[ArrayExcept]("array_except"), expression[ArrayUnion]("array_union"), expression[ArrayCompact]("array_compact"), - expression[CreateMap]("map"), - expression[CreateNamedStruct]("named_struct"), - expression[ElementAt]("element_at"), - expression[MapContainsKey]("map_contains_key"), - expression[MapFromArrays]("map_from_arrays"), - expression[MapKeys]("map_keys"), - expression[MapValues]("map_values"), - expression[MapEntries]("map_entries"), - expression[MapFromEntries]("map_from_entries"), - expression[MapConcat]("map_concat"), - expression[Size]("size"), expression[Slice]("slice"), expression[TrimArray]("trim_array"), - expression[Size]("cardinality", true, Some("2.4.0")), expression[ArraysZip]("arrays_zip"), expression[SortArray]("sort_array"), expression[Shuffle]("shuffle"), expression[ArrayMin]("array_min"), expression[ArrayMax]("array_max"), expression[ArrayAppend]("array_append"), - expression[Reverse]("reverse"), - expression[Concat]("concat"), expression[Flatten]("flatten"), expression[Sequence]("sequence"), expression[ArrayRepeat]("array_repeat"), expression[ArrayRemove]("array_remove"), expression[ArrayPrepend]("array_prepend"), expression[ArrayDistinct]("array_distinct"), - expression[ArrayTransform]("transform"), - expression[MapFilter]("map_filter"), - expression[ArrayFilter]("filter"), - expression[ArrayExists]("exists"), - expression[ArrayForAll]("forall"), - expression[ArrayAggregate]("aggregate"), - expression[ArrayAggregate]("reduce", setAlias = true, Some("3.4.0")), - expression[TransformValues]("transform_values"), - expression[TransformKeys]("transform_keys"), - expression[MapZipWith]("map_zip_with"), - expression[ZipWith]("zip_with"), - expression[Get]("get"), + expression[Get]("get") + ) + private def structExpressions: Seq[FunctionRegistryEntry] = Seq( + // struct functions + expression[CreateNamedStruct]("named_struct"), CreateStruct.registryEntry ) - private def miscExpressions: Seq[FunctionRegistryEntry] = Seq( - // misc functions - expression[AssertTrue]("assert_true"), - expressionBuilder("raise_error", RaiseErrorExpressionBuilder), - expression[Crc32]("crc32"), - expression[Md5]("md5"), + private def mapExpressions: Seq[FunctionRegistryEntry] = Seq( + // map functions + expression[StringToMap]("str_to_map"), + expression[CreateMap]("map"), + expression[MapContainsKey]("map_contains_key"), + expression[MapFromArrays]("map_from_arrays"), + expression[MapKeys]("map_keys"), + expression[MapValues]("map_values"), + expression[MapEntries]("map_entries"), + expression[MapFromEntries]("map_from_entries"), + expression[MapConcat]("map_concat") + ) + + private def aggregateExpressions: Seq[FunctionRegistryEntry] = Seq( + // aggregate functions + expressionBuilder("try_avg", TryAverageExpressionBuilder, setAlias = true), + expressionBuilder("try_sum", TrySumExpressionBuilder, setAlias = true), + expression[HyperLogLogPlusPlus]("approx_count_distinct"), + expression[Average]("avg"), + expression[Corr]("corr"), + expression[Count]("count"), + expression[CountIf]("count_if"), + expression[CovPopulation]("covar_pop"), + expression[CovSample]("covar_samp"), + expression[First]("first"), + expression[First]("first_value", true), + expression[AnyValue]("any_value"), + expression[Kurtosis]("kurtosis"), + expression[Last]("last"), + expression[Last]("last_value", true), + expression[Max]("max"), + expressionBuilder("max_by", MaxByBuilder), + expression[Average]("mean", true), + expression[Min]("min"), + expressionBuilder("min_by", MinByBuilder), + expression[Percentile]("percentile"), + expressionBuilder("percentile_cont", PercentileContBuilder), + expressionBuilder("percentile_disc", PercentileDiscBuilder), + expression[Median]("median"), + expression[Skewness]("skewness"), + expression[ApproximatePercentile]("percentile_approx"), + expression[ApproximatePercentile]("approx_percentile", true), + expression[HistogramNumeric]("histogram_numeric"), + expression[StddevSamp]("std", true), + expression[StddevSamp]("stddev", true), + expression[StddevPop]("stddev_pop"), + expression[StddevSamp]("stddev_samp"), + expression[Sum]("sum"), + expression[VarianceSamp]("variance", true), + expression[VariancePop]("var_pop"), + expression[VarianceSamp]("var_samp"), + expression[CollectList]("collect_list"), + expression[CollectList]("array_agg", true, Some("3.3.0")), + expression[CollectSet]("collect_set"), + expression[CollectUnion]("collect_union"), + expression[ListAgg]("listagg"), + expression[ListAgg]("string_agg", setAlias = true), + expressionBuilder("count_min_sketch", CountMinSketchAggExpressionBuilder), + expression[BoolAnd]("every", true), + expression[BoolAnd]("bool_and"), + expression[BoolOr]("any", true), + expression[BoolOr]("some", true), + expression[BoolOr]("bool_or"), + expression[RegrCount]("regr_count"), + expression[RegrAvgX]("regr_avgx"), + expression[RegrAvgY]("regr_avgy"), + expression[RegrR2]("regr_r2"), + expression[RegrSXX]("regr_sxx"), + expression[RegrSXY]("regr_sxy"), + expression[RegrSYY]("regr_syy"), + expression[RegrSlope]("regr_slope"), + expression[RegrIntercept]("regr_intercept"), + expressionBuilder("mode", ModeBuilder), + expression[HllSketchAgg]("hll_sketch_agg"), + expression[HllUnionAgg]("hll_union_agg"), + expression[ApproxTopK]("approx_top_k"), + expression[ThetaSketchAgg]("theta_sketch_agg"), + expression[ThetaUnionAgg]("theta_union_agg"), + expression[ThetaIntersectionAgg]("theta_intersection_agg"), + expression[ApproxTopKAccumulate]("approx_top_k_accumulate"), + expression[ApproxTopKCombine]("approx_top_k_combine"), + expression[KllSketchAggBigint]("kll_sketch_agg_bigint"), + expression[KllSketchAggFloat]("kll_sketch_agg_float"), + expression[KllSketchAggDouble]("kll_sketch_agg_double"), + expression[KllMergeAggBigint]("kll_merge_agg_bigint"), + expression[KllMergeAggFloat]("kll_merge_agg_float"), + expression[KllMergeAggDouble]("kll_merge_agg_double"), + expression[TupleIntersectionAggDouble]("tuple_intersection_agg_double"), + expression[TupleIntersectionAggInteger]("tuple_intersection_agg_integer"), + expressionBuilder("tuple_sketch_agg_double", TupleSketchAggDoubleExpressionBuilder), + expressionBuilder("tuple_sketch_agg_integer", TupleSketchAggIntegerExpressionBuilder), + expressionBuilder("tuple_union_agg_double", TupleUnionAggDoubleExpressionBuilder), + expressionBuilder("tuple_union_agg_integer", TupleUnionAggIntegerExpressionBuilder), + expression[Measure]("measure"), + expression[Grouping]("grouping"), + expression[GroupingID]("grouping_id"), + expression[BitAndAgg]("bit_and"), + expression[BitOrAgg]("bit_or"), + expression[BitXorAgg]("bit_xor"), + expression[BitmapConstructAgg]("bitmap_construct_agg"), + expression[BitmapOrAgg]("bitmap_or_agg"), + expression[BitmapAndAgg]("bitmap_and_agg"), + expression[BitmapXorAgg]("bitmap_xor_agg") + ) + + private def windowExpressions: Seq[FunctionRegistryEntry] = Seq( + // window functions + expression[Lead]("lead"), + expression[Lag]("lag"), + expression[RowNumber]("row_number"), + expression[CumeDist]("cume_dist"), + expression[NthValue]("nth_value"), + expression[NTile]("ntile"), + expression[Rank]("rank"), + expression[DenseRank]("dense_rank"), + expression[PercentRank]("percent_rank"), + expressionBuilder("counter_diff", CounterDiffExpressionBuilder) + ) + + private def generatorExpressions: Seq[FunctionRegistryEntry] = Seq( + // generator functions + expressionBuilder("explode", ExplodeExpressionBuilder), + expressionGeneratorBuilderOuter("explode_outer", ExplodeExpressionBuilder), + expressionBuilder("inline", InlineExpressionBuilder), + expressionGeneratorBuilderOuter("inline_outer", InlineExpressionBuilder), + expressionBuilder("posexplode", PosExplodeExpressionBuilder), + expressionGeneratorBuilderOuter("posexplode_outer", PosExplodeExpressionBuilder), + expression[Stack]("stack") + ) + + private def conversionExpressions: Seq[FunctionRegistryEntry] = Seq( + // conversion functions + expression[Cast]("cast"), + // Cast aliases (SPARK-16730) + castAlias("boolean", BooleanType), + castAlias("tinyint", ByteType), + castAlias("smallint", ShortType), + castAlias("int", IntegerType), + castAlias("bigint", LongType), + castAlias("float", FloatType), + castAlias("double", DoubleType), + castAlias("decimal", DecimalType.USER_DEFAULT), + castAlias("date", DateType), + castAlias("timestamp", TimestampType), + castAlias("time", TimeType()), + castAlias("binary", BinaryType), + castAlias("string", StringType) + ) + + private def csvExpressions: Seq[FunctionRegistryEntry] = Seq( + // CSV functions + expression[CsvToStructs]("from_csv"), + expression[SchemaOfCsv]("schema_of_csv"), + expression[StructsToCsv]("to_csv") + ) + + private def jsonExpressions: Seq[FunctionRegistryEntry] = Seq( + // JSON functions + expression[GetJsonObject]("get_json_object"), + expression[JsonTuple]("json_tuple"), + expression[StructsToJson]("to_json"), + expression[JsonToStructs]("from_json"), + expression[SchemaOfJson]("schema_of_json"), + expression[LengthOfJsonArray]("json_array_length"), + expression[JsonObjectKeys]("json_object_keys"), + expression[JsonTypeof]("json_typeof") + ) + + private def variantExpressions: Seq[FunctionRegistryEntry] = Seq( + // variant functions + expressionBuilder("parse_json", ParseJsonExpressionBuilder), + expressionBuilder("try_parse_json", TryParseJsonExpressionBuilder), + expression[IsVariantNull]("is_variant_null"), + expressionBuilder("variant_get", VariantGetExpressionBuilder), + expressionBuilder("try_variant_get", TryVariantGetExpressionBuilder), + expression[SchemaOfVariant]("schema_of_variant"), + expression[SchemaOfVariantAgg]("schema_of_variant_agg"), + expression[ToVariantObject]("to_variant_object"), + expression[VariantFromArrays]("variant_from_arrays"), + expression[VariantFromEntries]("variant_from_entries"), + expression[IsValidVariant]("is_valid_variant"), + expression[VariantDelete]("variant_delete"), + expressionBuilder("variant_insert", VariantInsertExpressionBuilder), + expressionBuilder("try_variant_insert", TryVariantInsertExpressionBuilder), + expressionBuilder("variant_set", VariantSetExpressionBuilder), + expressionBuilder("try_variant_set", TryVariantSetExpressionBuilder), + expressionBuilder("variant_array_append", VariantArrayAppendExpressionBuilder), + expressionBuilder("try_variant_array_append", TryVariantArrayAppendExpressionBuilder), + expressionBuilder("variant_strip_nulls", VariantStripNullsExpressionBuilder) + ) + + private def xmlExpressions: Seq[FunctionRegistryEntry] = Seq( + // XML functions + expression[XPathList]("xpath"), + expression[XPathBoolean]("xpath_boolean"), + expression[XPathDouble]("xpath_double"), + expression[XPathDouble]("xpath_number", true), + expression[XPathFloat]("xpath_float"), + expression[XPathInt]("xpath_int"), + expression[XPathLong]("xpath_long"), + expression[XPathShort]("xpath_short"), + expression[XPathString]("xpath_string"), + expression[XmlToStructs]("from_xml"), + expression[SchemaOfXml]("schema_of_xml"), + expression[StructsToXml]("to_xml") + ) + + private def urlExpressions: Seq[FunctionRegistryEntry] = Seq( + // URL functions + expression[TryUrlDecode]("try_url_decode"), + expression[UrlEncode]("url_encode"), + expression[UrlDecode]("url_decode"), + expression[ParseUrl]("parse_url"), + expression[TryParseUrl]("try_parse_url") + ) + + private def miscExpressions: Seq[FunctionRegistryEntry] = Seq( + // misc functions + expression[TryAesDecrypt]("try_aes_decrypt"), + expression[TryReflect]("try_reflect"), + expression[AssertTrue]("assert_true"), + expressionBuilder("raise_error", RaiseErrorExpressionBuilder), expression[Uuid]("uuid"), - expression[Murmur3Hash]("hash"), - expression[XxHash64]("xxhash64"), - expression[Xxh364]("xxh3_64"), - expression[Xxh3128]("xxh3_128"), - expression[Sha1]("sha", true), - expression[Sha1]("sha1"), - expression[Sha2]("sha2"), expression[AesEncrypt]("aes_encrypt"), expression[AesDecrypt]("aes_decrypt"), expression[Hmac]("hmac"), @@ -897,12 +1027,17 @@ object FunctionRegistry { expression[CallMethodViaReflection]("java_method", true), expression[SparkVersion]("version"), expression[TypeOf]("typeof"), - expression[EqualNull]("equal_null"), - expression[Measure]("measure") + expression[BitmapBucketNumber]("bitmap_bucket_number"), + expression[BitmapBitPosition]("bitmap_bit_position"), + expression[BitmapCount]("bitmap_count"), + expression[BitmapAnd]("bitmap_and"), + expression[BitmapOr]("bitmap_or"), + expression[BitmapAndNot]("bitmap_andnot"), + expression[BitmapXor]("bitmap_xor") ) private def dataSketchExpressions: Seq[FunctionRegistryEntry] = Seq( - // datasketch functions + // Datasketch functions expression[HllSketchEstimate]("hll_sketch_estimate"), expression[HllUnion]("hll_union"), expression[ThetaSketchEstimate]("theta_sketch_estimate"), @@ -945,114 +1080,21 @@ object FunctionRegistry { expression[KllSketchGetRankDouble]("kll_sketch_get_rank_double") ) - private def groupingExpressions: Seq[FunctionRegistryEntry] = Seq( - // grouping sets - expression[Grouping]("grouping"), - expression[GroupingID]("grouping_id") - ) - - private def windowExpressions: Seq[FunctionRegistryEntry] = Seq( - // window functions - expression[Lead]("lead"), - expression[Lag]("lag"), - expression[RowNumber]("row_number"), - expression[CumeDist]("cume_dist"), - expression[NthValue]("nth_value"), - expression[NTile]("ntile"), - expression[Rank]("rank"), - expression[DenseRank]("dense_rank"), - expression[PercentRank]("percent_rank"), - expressionBuilder("counter_diff", CounterDiffExpressionBuilder) - ) - - private def predicateExpressions: Seq[FunctionRegistryEntry] = Seq( - // predicates - expression[Between]("between"), - expression[And]("and"), - expression[In]("in"), - expression[Not]("not"), - expression[Or]("or") - ) - - private def comparisonExpressions: Seq[FunctionRegistryEntry] = Seq( - // comparison operators - expression[EqualNullSafe]("<=>"), - expression[EqualTo]("="), - expression[EqualTo]("=="), - expression[GreaterThan](">"), - expression[GreaterThanOrEqual](">="), - expression[LessThan]("<"), - expression[LessThanOrEqual]("<="), - expression[Not]("!") - ) - - private def bitwiseExpressions: Seq[FunctionRegistryEntry] = Seq( - // bitwise - expression[BitwiseAnd]("&"), - expression[BitwiseNot]("~"), - expression[BitwiseOr]("|"), - expression[BitwiseXor]("^"), - expression[ShiftLeft]("<<", true, Some("4.0.0")), - expression[ShiftRight](">>", true, Some("4.0.0")), - expression[ShiftRightUnsigned](">>>", true, Some("4.0.0")), - expression[BitwiseCount]("bit_count"), - expression[BitAndAgg]("bit_and"), - expression[BitOrAgg]("bit_or"), - expression[BitXorAgg]("bit_xor"), - expression[BitwiseGet]("bit_get"), - expression[BitwiseGet]("getbit", true) - ) - - private def bitmapExpressions: Seq[FunctionRegistryEntry] = Seq( - // bitmap functions and aggregates - expression[BitmapBucketNumber]("bitmap_bucket_number"), - expression[BitmapBitPosition]("bitmap_bit_position"), - expression[BitmapConstructAgg]("bitmap_construct_agg"), - expression[BitmapCount]("bitmap_count"), - expression[BitmapAnd]("bitmap_and"), - expression[BitmapOr]("bitmap_or"), - expression[BitmapAndNot]("bitmap_andnot"), - expression[BitmapXor]("bitmap_xor"), - expression[BitmapOrAgg]("bitmap_or_agg"), - expression[BitmapAndAgg]("bitmap_and_agg"), - expression[BitmapXorAgg]("bitmap_xor_agg") - ) - - private def jsonExpressions: Seq[FunctionRegistryEntry] = Seq( - // json - expression[StructsToJson]("to_json"), - expression[JsonToStructs]("from_json"), - expression[SchemaOfJson]("schema_of_json"), - expression[LengthOfJsonArray]("json_array_length"), - expression[JsonObjectKeys]("json_object_keys"), - expression[JsonTypeof]("json_typeof") + private def avroExpressions: Seq[FunctionRegistryEntry] = Seq( + // Avro functions + expression[FromAvro]("from_avro"), + expression[ToAvro]("to_avro"), + expression[SchemaOfAvro]("schema_of_avro") ) - private def variantExpressions: Seq[FunctionRegistryEntry] = Seq( - // Variant - expressionBuilder("parse_json", ParseJsonExpressionBuilder), - expressionBuilder("try_parse_json", TryParseJsonExpressionBuilder), - expression[IsVariantNull]("is_variant_null"), - expressionBuilder("variant_get", VariantGetExpressionBuilder), - expressionBuilder("try_variant_get", TryVariantGetExpressionBuilder), - expression[SchemaOfVariant]("schema_of_variant"), - expression[SchemaOfVariantAgg]("schema_of_variant_agg"), - expression[ToVariantObject]("to_variant_object"), - expression[VariantFromArrays]("variant_from_arrays"), - expression[VariantFromEntries]("variant_from_entries"), - expression[IsValidVariant]("is_valid_variant"), - expression[VariantDelete]("variant_delete"), - expressionBuilder("variant_insert", VariantInsertExpressionBuilder), - expressionBuilder("try_variant_insert", TryVariantInsertExpressionBuilder), - expressionBuilder("variant_set", VariantSetExpressionBuilder), - expressionBuilder("try_variant_set", TryVariantSetExpressionBuilder), - expressionBuilder("variant_array_append", VariantArrayAppendExpressionBuilder), - expressionBuilder("try_variant_array_append", TryVariantArrayAppendExpressionBuilder), - expressionBuilder("variant_strip_nulls", VariantStripNullsExpressionBuilder) + private def protobufExpressions: Seq[FunctionRegistryEntry] = Seq( + // Protobuf functions + expression[FromProtobuf]("from_protobuf"), + expression[ToProtobuf]("to_protobuf") ) private def spatialExpressions: Seq[FunctionRegistryEntry] = Seq( - // Spatial + // ST geospatial functions expression[ST_AsBinary]("st_asbinary"), expression[ST_GeogFromWKB]("st_geogfromwkb"), expression[ST_GeomFromWKB]("st_geomfromwkb"), @@ -1060,86 +1102,49 @@ object FunctionRegistry { expression[ST_SetSrid]("st_setsrid") ) - private def castExpressions: Seq[FunctionRegistryEntry] = Seq( - // cast - expression[Cast]("cast"), - // Cast aliases (SPARK-16730) - castAlias("boolean", BooleanType), - castAlias("tinyint", ByteType), - castAlias("smallint", ShortType), - castAlias("int", IntegerType), - castAlias("bigint", LongType), - castAlias("float", FloatType), - castAlias("double", DoubleType), - castAlias("decimal", DecimalType.USER_DEFAULT), - castAlias("date", DateType), - castAlias("timestamp", TimestampType), - castAlias("time", TimeType()), - castAlias("binary", BinaryType), - castAlias("string", StringType) - ) - - private def maskExpressions: Seq[FunctionRegistryEntry] = Seq( - // mask functions - expressionBuilder("mask", MaskExpressionBuilder) - ) - - private def csvExpressions: Seq[FunctionRegistryEntry] = Seq( - // csv - expression[CsvToStructs]("from_csv"), - expression[SchemaOfCsv]("schema_of_csv"), - expression[StructsToCsv]("to_csv") - ) - - private def xmlExpressions: Seq[FunctionRegistryEntry] = Seq( - // Xml - expression[XmlToStructs]("from_xml"), - expression[SchemaOfXml]("schema_of_xml"), - expression[StructsToXml]("to_xml") - ) - - private def avroExpressions: Seq[FunctionRegistryEntry] = Seq( - // Avro - expression[FromAvro]("from_avro"), - expression[ToAvro]("to_avro"), - expression[SchemaOfAvro]("schema_of_avro") - ) - - private def protobufExpressions: Seq[FunctionRegistryEntry] = Seq( - // Protobuf - expression[FromProtobuf]("from_protobuf"), - expression[ToProtobuf]("to_protobuf") + private def vectorExpressions: Seq[FunctionRegistryEntry] = Seq( + // vector functions + expression[VectorCosineSimilarity]("vector_cosine_similarity"), + expression[VectorInnerProduct]("vector_inner_product"), + expression[VectorL2Distance]("vector_l2_distance"), + expression[VectorNorm]("vector_norm"), + expression[VectorNormalize]("vector_normalize"), + expression[VectorAvg]("vector_avg"), + expression[VectorSum]("vector_sum") ) - // Keep expression groups in separate methods to limit the bytecode size of the static - // initializer and leave enough headroom for coverage instrumentation. + // Keep registry entries aligned with their ExpressionInfo groups. Public APIs and documentation + // use the corresponding taxonomy where applicable. Separate methods also limit the bytecode size + // of the static initializer and leave enough headroom for coverage instrumentation. val expressions: Map[String, (ExpressionInfo, FunctionBuilder)] = Seq( - miscNonAggregateExpressions, + conditionalExpressions, + predicateExpressions, mathExpressions, - tryExpressions, - aggregateExpressions, - vectorExpressions, stringExpressions, - urlExpressions, + bitwiseExpressions, datetimeExpressions, + hashExpressions, collectionExpressions, - miscExpressions, - dataSketchExpressions, - groupingExpressions, + lambdaExpressions, + arrayExpressions, + structExpressions, + mapExpressions, + aggregateExpressions, windowExpressions, - predicateExpressions, - comparisonExpressions, - bitwiseExpressions, - bitmapExpressions, + generatorExpressions, + conversionExpressions, + csvExpressions, jsonExpressions, variantExpressions, - spatialExpressions, - castExpressions, - maskExpressions, - csvExpressions, xmlExpressions, + urlExpressions, + miscExpressions, + dataSketchExpressions, avroExpressions, - protobufExpressions).flatten.toMap + protobufExpressions, + spatialExpressions, + vectorExpressions + ).flatten.toMap // BuiltinRegistryMixin normalizes any name to the builtin 3-part key (system.builtin.name). val builtin: SimpleFunctionRegistry = {