From 9fe45df403356f8b09c9c9313aeabc912ce843ed Mon Sep 17 00:00:00 2001 From: rich7420 Date: Sun, 13 Sep 2026 23:24:51 +0800 Subject: [PATCH 1/2] test: cover string and collection collation routing --- .../routing_collection_collation_disabled.sql | 33 ++++++++++++++ .../routing_collection_collation_enabled.sql | 33 ++++++++++++++ .../routing_string_collation_disabled.sql | 45 +++++++++++++++++++ .../routing_string_collation_enabled.sql | 45 +++++++++++++++++++ .../comet/CometStringExpressionSuite.scala | 27 ++++++++++- 5 files changed, 182 insertions(+), 1 deletion(-) create mode 100644 spark/src/test/resources/sql-tests/expressions/array/routing_collection_collation_disabled.sql create mode 100644 spark/src/test/resources/sql-tests/expressions/array/routing_collection_collation_enabled.sql create mode 100644 spark/src/test/resources/sql-tests/expressions/string/routing_string_collation_disabled.sql create mode 100644 spark/src/test/resources/sql-tests/expressions/string/routing_string_collation_enabled.sql diff --git a/spark/src/test/resources/sql-tests/expressions/array/routing_collection_collation_disabled.sql b/spark/src/test/resources/sql-tests/expressions/array/routing_collection_collation_disabled.sql new file mode 100644 index 00000000000..9aff3d0e57d --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/array/routing_collection_collation_disabled.sql @@ -0,0 +1,33 @@ +-- Licensed to the Apache Software Foundation (ASF) under one +-- or more contributor license agreements. See the NOTICE file +-- distributed with this work for additional information +-- regarding copyright ownership. The ASF licenses this file +-- to you under the Apache License, Version 2.0 (the +-- "License"); you may not use this file except in compliance +-- with the License. You may obtain a copy of the License at +-- +-- http://www.apache.org/licenses/LICENSE-2.0 +-- +-- Unless required by applicable law or agreed to in writing, +-- software distributed under the License is distributed on an +-- "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +-- KIND, either express or implied. See the License for the +-- specific language governing permissions and limitations +-- under the License. + +-- MinSparkVersion: 4.0 +-- Config: spark.comet.exec.scalaUDF.codegen.enabled=false +-- Config: spark.comet.expression.ArrayJoin.allowIncompatible=false +-- Config: spark.comet.expression.StringToMap.allowIncompatible=false + +statement +CREATE TABLE routing_collection_collation(s STRING, a ARRAY) USING parquet + +statement +INSERT INTO routing_collection_collation VALUES ('a:1,b:2', array('a', 'B')), ('', array()), (NULL, NULL) + +query expect_fallback(array_join: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT array_join(transform(a, x -> x COLLATE UTF8_LCASE), ',') FROM routing_collection_collation + +query expect_fallback(str_to_map: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT str_to_map(s COLLATE UTF8_LCASE) FROM routing_collection_collation diff --git a/spark/src/test/resources/sql-tests/expressions/array/routing_collection_collation_enabled.sql b/spark/src/test/resources/sql-tests/expressions/array/routing_collection_collation_enabled.sql new file mode 100644 index 00000000000..ce6187a506f --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/array/routing_collection_collation_enabled.sql @@ -0,0 +1,33 @@ +-- Licensed to the Apache Software Foundation (ASF) under one +-- or more contributor license agreements. See the NOTICE file +-- distributed with this work for additional information +-- regarding copyright ownership. The ASF licenses this file +-- to you under the Apache License, Version 2.0 (the +-- "License"); you may not use this file except in compliance +-- with the License. You may obtain a copy of the License at +-- +-- http://www.apache.org/licenses/LICENSE-2.0 +-- +-- Unless required by applicable law or agreed to in writing, +-- software distributed under the License is distributed on an +-- "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +-- KIND, either express or implied. See the License for the +-- specific language governing permissions and limitations +-- under the License. + +-- MinSparkVersion: 4.0 +-- Config: spark.comet.exec.scalaUDF.codegen.enabled=true +-- Config: spark.comet.expression.ArrayJoin.allowIncompatible=false +-- Config: spark.comet.expression.StringToMap.allowIncompatible=false + +statement +CREATE TABLE routing_collection_collation(s STRING, a ARRAY) USING parquet + +statement +INSERT INTO routing_collection_collation VALUES ('a:1,b:2', array('a', 'B')), ('', array()), (NULL, NULL) + +query expect_dispatch(array_join) +SELECT array_join(transform(a, x -> x COLLATE UTF8_LCASE), ',') FROM routing_collection_collation + +query expect_dispatch(str_to_map) +SELECT str_to_map(s COLLATE UTF8_LCASE) FROM routing_collection_collation diff --git a/spark/src/test/resources/sql-tests/expressions/string/routing_string_collation_disabled.sql b/spark/src/test/resources/sql-tests/expressions/string/routing_string_collation_disabled.sql new file mode 100644 index 00000000000..132371c1f66 --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/string/routing_string_collation_disabled.sql @@ -0,0 +1,45 @@ +-- Licensed to the Apache Software Foundation (ASF) under one +-- or more contributor license agreements. See the NOTICE file +-- distributed with this work for additional information +-- regarding copyright ownership. The ASF licenses this file +-- to you under the Apache License, Version 2.0 (the +-- "License"); you may not use this file except in compliance +-- with the License. You may obtain a copy of the License at +-- +-- http://www.apache.org/licenses/LICENSE-2.0 +-- +-- Unless required by applicable law or agreed to in writing, +-- software distributed under the License is distributed on an +-- "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +-- KIND, either express or implied. See the License for the +-- specific language governing permissions and limitations +-- under the License. + +-- MinSparkVersion: 4.0 +-- Config: spark.comet.exec.scalaUDF.codegen.enabled=false +-- Config: spark.comet.expression.Concat.allowIncompatible=false +-- Config: spark.comet.expression.Reverse.allowIncompatible=false + +statement +CREATE TABLE routing_string_collation(s STRING, plain STRING) USING parquet + +statement +INSERT INTO routing_string_collation VALUES ('Hello', 'hello'), ('', ''), (NULL, NULL) + +query expect_native(reverse) +SELECT reverse(plain) FROM routing_string_collation + +query expect_native(levenshtein) +SELECT levenshtein(plain, 'hello') FROM routing_string_collation + +query expect_fallback(concat: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT concat(s COLLATE UTF8_LCASE, '!') FROM routing_string_collation + +query expect_fallback(reverse: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT reverse(s COLLATE UTF8_LCASE) FROM routing_string_collation + +query expect_fallback(like: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT (s COLLATE UTF8_LCASE) LIKE 'H_llo' FROM routing_string_collation + +query expect_fallback(levenshtein: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT levenshtein(s COLLATE UTF8_LCASE, 'HELLO') FROM routing_string_collation diff --git a/spark/src/test/resources/sql-tests/expressions/string/routing_string_collation_enabled.sql b/spark/src/test/resources/sql-tests/expressions/string/routing_string_collation_enabled.sql new file mode 100644 index 00000000000..5dbe954f2d1 --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/string/routing_string_collation_enabled.sql @@ -0,0 +1,45 @@ +-- Licensed to the Apache Software Foundation (ASF) under one +-- or more contributor license agreements. See the NOTICE file +-- distributed with this work for additional information +-- regarding copyright ownership. The ASF licenses this file +-- to you under the Apache License, Version 2.0 (the +-- "License"); you may not use this file except in compliance +-- with the License. You may obtain a copy of the License at +-- +-- http://www.apache.org/licenses/LICENSE-2.0 +-- +-- Unless required by applicable law or agreed to in writing, +-- software distributed under the License is distributed on an +-- "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +-- KIND, either express or implied. See the License for the +-- specific language governing permissions and limitations +-- under the License. + +-- MinSparkVersion: 4.0 +-- Config: spark.comet.exec.scalaUDF.codegen.enabled=true +-- Config: spark.comet.expression.Concat.allowIncompatible=false +-- Config: spark.comet.expression.Reverse.allowIncompatible=false + +statement +CREATE TABLE routing_string_collation(s STRING, plain STRING) USING parquet + +statement +INSERT INTO routing_string_collation VALUES ('Hello', 'hello'), ('', ''), (NULL, NULL) + +query expect_native(reverse) +SELECT reverse(plain) FROM routing_string_collation + +query expect_native(levenshtein) +SELECT levenshtein(plain, 'hello') FROM routing_string_collation + +query expect_dispatch(concat) +SELECT concat(s COLLATE UTF8_LCASE, '!') FROM routing_string_collation + +query expect_dispatch(reverse) +SELECT reverse(s COLLATE UTF8_LCASE) FROM routing_string_collation + +query expect_dispatch(like) +SELECT (s COLLATE UTF8_LCASE) LIKE 'H_llo' FROM routing_string_collation + +query expect_dispatch(levenshtein) +SELECT levenshtein(s COLLATE UTF8_LCASE, 'HELLO') FROM routing_string_collation diff --git a/spark/src/test/scala/org/apache/comet/CometStringExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometStringExpressionSuite.scala index c58403487ac..081f5e9323e 100644 --- a/spark/src/test/scala/org/apache/comet/CometStringExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometStringExpressionSuite.scala @@ -23,8 +23,9 @@ import scala.util.Random import org.apache.parquet.hadoop.ParquetOutputFormat import org.apache.spark.sql.{CometTestBase, DataFrame} +import org.apache.spark.sql.catalyst.expressions.{Concat, Literal, Reverse} import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.{DataTypes, StructField, StructType} +import org.apache.spark.sql.types.{DataType, DataTypes, StructField, StructType} import org.apache.comet.CometSparkSessionExtensions.isSpark40Plus import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator} @@ -38,6 +39,30 @@ class CometStringExpressionSuite extends CometTestBase with CometCodegenAssertio "తెలుగు") // scalastyle:on + if (isSpark40Plus) { + test("collated strings preserve native opt-in routing") { + withParquetTable(Seq(("abc", 1), ("", 2), (null, 3)), "tbl") { + // Build typed literals directly: Collate and casts of columns are not native, while + // ordinary constant folding would remove the expression we want to test. + val text = Literal.create("abc", DataType.fromDDL("STRING COLLATE UTF8_LCASE")) + withSQLConf( + SQLConf.OPTIMIZER_EXCLUDED_RULES.key -> + "org.apache.spark.sql.catalyst.optimizer.ConstantFolding", + CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> "false", + CometConf.getExprAllowIncompatConfigKey("Concat") -> "true", + CometConf.getExprAllowIncompatConfigKey("Reverse") -> "true") { + for ((name, expression) <- Seq( + "concat" -> Concat(Seq(text, text)), + "reverse" -> Reverse(text))) { + checkSparkAnswerAndImpl( + sql("SELECT _1 FROM tbl").select(getColumnFromExpression(expression)), + native = Seq(name)) + } + } + } + } + } + test("lpad string") { testStringPadding("lpad") } From 89b5669f3a7e9c880a50b2857dc118e9d56fa89d Mon Sep 17 00:00:00 2001 From: rich7420 Date: Sun, 13 Sep 2026 23:24:52 +0800 Subject: [PATCH 2/2] test: cover predicate expression routing configurations --- .../routing_predicates_disabled.sql | 92 +++++++++++++++++++ .../routing_predicates_enabled.sql | 92 +++++++++++++++++++ .../apache/comet/CometExpressionSuite.scala | 51 +++++++--- 3 files changed, 224 insertions(+), 11 deletions(-) create mode 100644 spark/src/test/resources/sql-tests/expressions/conditional/routing_predicates_disabled.sql create mode 100644 spark/src/test/resources/sql-tests/expressions/conditional/routing_predicates_enabled.sql diff --git a/spark/src/test/resources/sql-tests/expressions/conditional/routing_predicates_disabled.sql b/spark/src/test/resources/sql-tests/expressions/conditional/routing_predicates_disabled.sql new file mode 100644 index 00000000000..b906a29b7c2 --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/conditional/routing_predicates_disabled.sql @@ -0,0 +1,92 @@ +-- Licensed to the Apache Software Foundation (ASF) under one +-- or more contributor license agreements. See the NOTICE file +-- distributed with this work for additional information +-- regarding copyright ownership. The ASF licenses this file +-- to you under the Apache License, Version 2.0 (the +-- "License"); you may not use this file except in compliance +-- with the License. You may obtain a copy of the License at +-- +-- http://www.apache.org/licenses/LICENSE-2.0 +-- +-- Unless required by applicable law or agreed to in writing, +-- software distributed under the License is distributed on an +-- "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +-- KIND, either express or implied. See the License for the +-- specific language governing permissions and limitations +-- under the License. + +-- MinSparkVersion: 4.0 +-- Config: spark.comet.exec.scalaUDF.codegen.enabled=false +-- Config: spark.sql.optimizer.inSetConversionThreshold=0 + +statement +CREATE TABLE routing_predicates(s STRING, t STRING, a STRING, b STRING) USING parquet + +statement +INSERT INTO routing_predicates VALUES ('Hello', 'HELLO', 'Hello', 'HELLO'), ('', '', '', ''), (NULL, NULL, NULL, NULL) + +query expect_native(equalto) +SELECT a = b FROM routing_predicates + +query expect_native(equalnullsafe) +SELECT a <=> b FROM routing_predicates + +query expect_native(lessthan) +SELECT a < b FROM routing_predicates + +query expect_native(lessthanorequal) +SELECT a <= b FROM routing_predicates + +query expect_native(greaterthan) +SELECT a > b FROM routing_predicates + +query expect_native(greaterthanorequal) +SELECT a >= b FROM routing_predicates + +query expect_native(contains) +SELECT contains(a, b) FROM routing_predicates + +query expect_native(startswith) +SELECT startswith(a, b) FROM routing_predicates + +query expect_native(endswith) +SELECT endswith(a, b) FROM routing_predicates + +query expect_native(in) +SELECT a IN (b, 'other') FROM routing_predicates + +query expect_native(inset) +SELECT a IN ('HELLO', 'a', 'b') FROM routing_predicates + +query expect_fallback(equalto: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT (s COLLATE UTF8_LCASE) = (t COLLATE UTF8_LCASE) FROM routing_predicates + +query expect_fallback(equalnullsafe: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT (s COLLATE UTF8_LCASE) <=> (t COLLATE UTF8_LCASE) FROM routing_predicates + +query expect_fallback(lessthan: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT (s COLLATE UTF8_LCASE) < (t COLLATE UTF8_LCASE) FROM routing_predicates + +query expect_fallback(lessthanorequal: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT (s COLLATE UTF8_LCASE) <= (t COLLATE UTF8_LCASE) FROM routing_predicates + +query expect_fallback(greaterthan: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT (s COLLATE UTF8_LCASE) > (t COLLATE UTF8_LCASE) FROM routing_predicates + +query expect_fallback(greaterthanorequal: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT (s COLLATE UTF8_LCASE) >= (t COLLATE UTF8_LCASE) FROM routing_predicates + +query expect_fallback(contains: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT contains((s COLLATE UTF8_LCASE), (t COLLATE UTF8_LCASE)) FROM routing_predicates + +query expect_fallback(startswith: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT startswith((s COLLATE UTF8_LCASE), (t COLLATE UTF8_LCASE)) FROM routing_predicates + +query expect_fallback(endswith: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT endswith((s COLLATE UTF8_LCASE), (t COLLATE UTF8_LCASE)) FROM routing_predicates + +query expect_fallback(in: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT (s COLLATE UTF8_LCASE) IN ((t COLLATE UTF8_LCASE), 'other') FROM routing_predicates + +query expect_fallback(inset: spark.comet.exec.scalaUDF.codegen.enabled=false) +SELECT (s COLLATE UTF8_LCASE) IN ('HELLO', 'a', 'b') FROM routing_predicates diff --git a/spark/src/test/resources/sql-tests/expressions/conditional/routing_predicates_enabled.sql b/spark/src/test/resources/sql-tests/expressions/conditional/routing_predicates_enabled.sql new file mode 100644 index 00000000000..83dadeafb3a --- /dev/null +++ b/spark/src/test/resources/sql-tests/expressions/conditional/routing_predicates_enabled.sql @@ -0,0 +1,92 @@ +-- Licensed to the Apache Software Foundation (ASF) under one +-- or more contributor license agreements. See the NOTICE file +-- distributed with this work for additional information +-- regarding copyright ownership. The ASF licenses this file +-- to you under the Apache License, Version 2.0 (the +-- "License"); you may not use this file except in compliance +-- with the License. You may obtain a copy of the License at +-- +-- http://www.apache.org/licenses/LICENSE-2.0 +-- +-- Unless required by applicable law or agreed to in writing, +-- software distributed under the License is distributed on an +-- "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +-- KIND, either express or implied. See the License for the +-- specific language governing permissions and limitations +-- under the License. + +-- MinSparkVersion: 4.0 +-- Config: spark.comet.exec.scalaUDF.codegen.enabled=true +-- Config: spark.sql.optimizer.inSetConversionThreshold=0 + +statement +CREATE TABLE routing_predicates(s STRING, t STRING, a STRING, b STRING) USING parquet + +statement +INSERT INTO routing_predicates VALUES ('Hello', 'HELLO', 'Hello', 'HELLO'), ('', '', '', ''), (NULL, NULL, NULL, NULL) + +query expect_native(equalto) +SELECT a = b FROM routing_predicates + +query expect_native(equalnullsafe) +SELECT a <=> b FROM routing_predicates + +query expect_native(lessthan) +SELECT a < b FROM routing_predicates + +query expect_native(lessthanorequal) +SELECT a <= b FROM routing_predicates + +query expect_native(greaterthan) +SELECT a > b FROM routing_predicates + +query expect_native(greaterthanorequal) +SELECT a >= b FROM routing_predicates + +query expect_native(contains) +SELECT contains(a, b) FROM routing_predicates + +query expect_native(startswith) +SELECT startswith(a, b) FROM routing_predicates + +query expect_native(endswith) +SELECT endswith(a, b) FROM routing_predicates + +query expect_native(in) +SELECT a IN (b, 'other') FROM routing_predicates + +query expect_native(inset) +SELECT a IN ('HELLO', 'a', 'b') FROM routing_predicates + +query expect_dispatch(equalto) +SELECT (s COLLATE UTF8_LCASE) = (t COLLATE UTF8_LCASE) FROM routing_predicates + +query expect_dispatch(equalnullsafe) +SELECT (s COLLATE UTF8_LCASE) <=> (t COLLATE UTF8_LCASE) FROM routing_predicates + +query expect_dispatch(lessthan) +SELECT (s COLLATE UTF8_LCASE) < (t COLLATE UTF8_LCASE) FROM routing_predicates + +query expect_dispatch(lessthanorequal) +SELECT (s COLLATE UTF8_LCASE) <= (t COLLATE UTF8_LCASE) FROM routing_predicates + +query expect_dispatch(greaterthan) +SELECT (s COLLATE UTF8_LCASE) > (t COLLATE UTF8_LCASE) FROM routing_predicates + +query expect_dispatch(greaterthanorequal) +SELECT (s COLLATE UTF8_LCASE) >= (t COLLATE UTF8_LCASE) FROM routing_predicates + +query expect_dispatch(contains) +SELECT contains((s COLLATE UTF8_LCASE), (t COLLATE UTF8_LCASE)) FROM routing_predicates + +query expect_dispatch(startswith) +SELECT startswith((s COLLATE UTF8_LCASE), (t COLLATE UTF8_LCASE)) FROM routing_predicates + +query expect_dispatch(endswith) +SELECT endswith((s COLLATE UTF8_LCASE), (t COLLATE UTF8_LCASE)) FROM routing_predicates + +query expect_dispatch(in) +SELECT (s COLLATE UTF8_LCASE) IN ((t COLLATE UTF8_LCASE), 'other') FROM routing_predicates + +query expect_dispatch(inset) +SELECT (s COLLATE UTF8_LCASE) IN ('HELLO', 'a', 'b') FROM routing_predicates diff --git a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala index 4ab86b4f373..620e75b43e3 100644 --- a/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometExpressionSuite.scala @@ -23,7 +23,7 @@ import scala.util.Random import org.apache.hadoop.fs.Path import org.apache.spark.sql.{Column, CometTestBase, DataFrame, Row} -import org.apache.spark.sql.catalyst.expressions.{Alias, Cast, FromUnixTime, Literal, StructsToJson, TruncDate, TruncTimestamp} +import org.apache.spark.sql.catalyst.expressions.{Alias, Cast, FromUnixTime, InSet, Literal, StructsToJson, TruncDate, TruncTimestamp} import org.apache.spark.sql.catalyst.optimizer.{ConvertToLocalRelation, OptimizeIn, SimplifyExtractValueOps} import org.apache.spark.sql.comet.CometProjectExec import org.apache.spark.sql.execution.{ProjectExec, SparkPlan} @@ -1685,10 +1685,11 @@ class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper { test("test in(set)/not in(set)") { Seq("100", "0").foreach { inSetThreshold => - Seq(false, true).foreach { dictionary => + for (dictionary <- Seq(false, true); codegen <- Seq("false", "true")) { withSQLConf( SQLConf.OPTIMIZER_INSET_CONVERSION_THRESHOLD.key -> inSetThreshold, - "parquet.enable.dictionary" -> dictionary.toString) { + "parquet.enable.dictionary" -> dictionary.toString, + CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> codegen) { val table = "names" withTable(table) { sql(s"create table $table(id int, name varchar(20)) using parquet") @@ -1696,9 +1697,15 @@ class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper { s"insert into $table values(1, 'James'), (1, 'Jones'), (2, 'Smith'), (3, 'Smith')," + "(NULL, 'Jones'), (4, NULL)") - checkSparkAnswerAndOperator(s"SELECT * FROM $table WHERE id in (1, 2, 4, NULL)") - checkSparkAnswerAndOperator( - s"SELECT * FROM $table WHERE name in ('Smith', 'Brown', NULL)") + val nativeName = if (inSetThreshold == "0") "inset" else "in" + checkSparkAnswerAndImpl( + s"SELECT * FROM $table WHERE id in (1, 2, 4, NULL)", + native = Seq(nativeName), + dispatched = Seq.empty) + checkSparkAnswerAndImpl( + s"SELECT * FROM $table WHERE name in ('Smith', 'Brown', NULL)", + native = Seq(nativeName), + dispatched = Seq.empty) // TODO: why with not in, the plan is only `LocalTableScan`? checkSparkAnswerAndOperator(s"SELECT * FROM $table WHERE id not in (1)") @@ -1723,12 +1730,34 @@ class CometExpressionSuite extends CometTestBase with AdaptiveSparkPlanHelper { withParquetTable(data, "tbl") { // An unset config exercises the version-dependent default, which follows ANSI mode on // Spark 4.0+ and is always the legacy behavior on Spark 3.x. - for (legacy <- Seq(Some("true"), Some("false"), None); ansi <- Seq("true", "false")) { + for { + legacy <- Seq(Some("true"), Some("false"), None) + ansi <- Seq("true", "false") + codegen <- Seq("true", "false") + } { val legacyConf = legacy.map("spark.sql.legacy.nullInEmptyListBehavior" -> _).toSeq - withSQLConf(Seq(SQLConf.ANSI_ENABLED.key -> ansi) ++ legacyConf: _*) { - val df = sql("SELECT _1 AS a FROM tbl") - .select(col("a"), col("a").isin(), !col("a").isin()) - checkSparkAnswer(df) + withSQLConf( + Seq( + SQLConf.ANSI_ENABLED.key -> ansi, + CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> codegen) ++ legacyConf: _*) { + val input = sql("SELECT _1 AS a FROM tbl") + val emptySet = InSet(input.queryExecution.analyzed.output.head, Set.empty[Any]) + val df = input.select( + col("a"), + col("a").isin(), + !col("a").isin(), + getColumnFromExpression(emptySet)) + val legacyEnabled = !CometSparkSessionExtensions.isSpark35Plus || + legacy.map(_.toBoolean).getOrElse(!isSpark40Plus || !ansi.toBoolean) + if (!legacyEnabled) { + checkSparkAnswerAndImpl(df, native = Seq("in", "inset")) + } else if (codegen.toBoolean) { + checkSparkAnswerAndImpl(df, dispatched = Seq("in", "inset")) + } else { + checkSparkAnswerAndFallbackReason( + df, + s"in: ${CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key}=false") + } } } }