From 6b95d5517b9f4dc0f1d69a790aa3abef39b71c47 Mon Sep 17 00:00:00 2001 From: Xinyao Zhang Date: Sat, 10 Oct 2026 03:34:07 +0000 Subject: [PATCH] [SQL] Return zero for instr with zero start and empty substring --- .../util/CollationAwareUTF8String.java | 5 +++-- .../apache/spark/unsafe/types/UTF8String.java | 6 +++--- .../unsafe/types/CollationSupportSuite.java | 13 +++++++++++++ .../spark/unsafe/types/UTF8StringSuite.java | 4 ++++ .../expressions/StringExpressionsSuite.scala | 18 ++++++++++++++++++ 5 files changed, 41 insertions(+), 5 deletions(-) diff --git a/common/unsafe/src/main/java/org/apache/spark/sql/catalyst/util/CollationAwareUTF8String.java b/common/unsafe/src/main/java/org/apache/spark/sql/catalyst/util/CollationAwareUTF8String.java index 5e239212dfc66..441af9774fb17 100644 --- a/common/unsafe/src/main/java/org/apache/spark/sql/catalyst/util/CollationAwareUTF8String.java +++ b/common/unsafe/src/main/java/org/apache/spark/sql/catalyst/util/CollationAwareUTF8String.java @@ -825,8 +825,8 @@ private static int lowercaseIndexOfSlow(final UTF8String target, final UTF8Strin public static int lowercaseIndexOf(final UTF8String target, final UTF8String pattern, final int start, final int occurrence) { assert occurrence > 0; - if (pattern.numBytes() == 0) return target.indexOfEmpty(start); if (start == 0) return MATCH_NOT_FOUND; + if (pattern.numBytes() == 0) return target.indexOfEmpty(start); if (target.isFullAscii() && pattern.isFullAscii()) { return target.toLowerCase().indexOf(pattern.toLowerCase(), start, occurrence); } @@ -928,8 +928,9 @@ public static int indexOf(final UTF8String target, final UTF8String pattern, public static int indexOf(final UTF8String target, final UTF8String pattern, final int start, final int occurrence, final int collationId) { assert occurrence > 0; + if (start == 0) return MATCH_NOT_FOUND; if (pattern.numBytes() == 0) return target.indexOfEmpty(start); - if (target.numBytes() == 0 || start == 0) return MATCH_NOT_FOUND; + if (target.numBytes() == 0) return MATCH_NOT_FOUND; String targetStr = target.toValidString(); String patternStr = pattern.toValidString(); diff --git a/common/unsafe/src/main/java/org/apache/spark/unsafe/types/UTF8String.java b/common/unsafe/src/main/java/org/apache/spark/unsafe/types/UTF8String.java index 5318cd8d11c29..437faa399f5e8 100644 --- a/common/unsafe/src/main/java/org/apache/spark/unsafe/types/UTF8String.java +++ b/common/unsafe/src/main/java/org/apache/spark/unsafe/types/UTF8String.java @@ -1319,12 +1319,12 @@ public int indexOf(UTF8String v, int start) { */ public int indexOf(UTF8String pattern, int start, int occurrence) { assert occurrence > 0; - if (pattern.numBytes() == 0) { - return indexOfEmpty(start); - } if (start == 0) { return -1; } + if (pattern.numBytes() == 0) { + return indexOfEmpty(start); + } int charCount = 0; int byteIdx; diff --git a/common/unsafe/src/test/java/org/apache/spark/unsafe/types/CollationSupportSuite.java b/common/unsafe/src/test/java/org/apache/spark/unsafe/types/CollationSupportSuite.java index edd24671f1dc7..908de60074310 100644 --- a/common/unsafe/src/test/java/org/apache/spark/unsafe/types/CollationSupportSuite.java +++ b/common/unsafe/src/test/java/org/apache/spark/unsafe/types/CollationSupportSuite.java @@ -4047,6 +4047,19 @@ private void assertStringInstrWithOccurrence(String string, String substring, in assertEquals(expected, res); } + @Test + public void testStringInstrWithOccurrenceZeroStart() throws SparkException { + for (String collationName : testSupportedCollations) { + for (String string : new String[]{"", "abc", "\u4F60\u597D"}) { + for (String substring : new String[]{"", "a"}) { + for (int occurrence : new int[]{1, 2}) { + assertStringInstrWithOccurrence(string, substring, 0, occurrence, collationName, 0); + } + } + } + } + } + @Test public void testStringInstrWithOccurrence() throws SparkException { // Test start = 1 and occurrence = 1 (equivalent to StringInstr) diff --git a/common/unsafe/src/test/java/org/apache/spark/unsafe/types/UTF8StringSuite.java b/common/unsafe/src/test/java/org/apache/spark/unsafe/types/UTF8StringSuite.java index c871b785abb04..d6453daf8fb8b 100644 --- a/common/unsafe/src/test/java/org/apache/spark/unsafe/types/UTF8StringSuite.java +++ b/common/unsafe/src/test/java/org/apache/spark/unsafe/types/UTF8StringSuite.java @@ -388,6 +388,10 @@ public void indexOf() { assertEquals(0, fromString("hello").indexOf(EMPTY_UTF8, 1, 1)); assertEquals(0, fromString("hello").indexOf(EMPTY_UTF8, 5, 1)); assertEquals(0, fromString("hello").indexOf(EMPTY_UTF8, -1, 1)); + assertEquals(-1, fromString("hello").indexOf(EMPTY_UTF8, 0, 1)); + assertEquals(-1, fromString("hello").indexOf(EMPTY_UTF8, 0, 2)); + assertEquals(-1, EMPTY_UTF8.indexOf(EMPTY_UTF8, 0, 1)); + assertEquals(-1, EMPTY_UTF8.indexOf(EMPTY_UTF8, 0, 2)); // Boundary cases assertEquals(0, fromString("x").indexOf(fromString("x"), 1, 1)); diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/StringExpressionsSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/StringExpressionsSuite.scala index f3d479887ab27..886249042c86e 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/StringExpressionsSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/StringExpressionsSuite.scala @@ -1348,6 +1348,24 @@ class StringExpressionsSuite extends SparkFunSuite with ExpressionEvalHelper { // scalastyle:on } + test("INSTR with start = 0 returns 0 for empty substrings") { + for (collation <- Seq("UTF8_BINARY", "UTF8_LCASE", "UNICODE", "UNICODE_CI")) { + val stringType = StringType(collation) + val expression = StringInstrWithOccurrence( + BoundReference(0, stringType, nullable = false), + BoundReference(1, stringType, nullable = false), + BoundReference(2, IntegerType, nullable = false), + BoundReference(3, IntegerType, nullable = false)) + for { + string <- Seq("", "abc", "\u4F60\u597D") + substring <- Seq("", "a") + occurrence <- Seq(1, 2) + } { + checkEvaluation(expression, 0, create_row(string, substring, 0, occurrence)) + } + } + } + test("StringInstrExpressionBuilder") { val seq1 = Seq(Literal("abcabc"), Literal("a"), Literal(-1), Literal(2)) val seq2 = Seq(Literal("abcabc"), Literal("a"), Literal(-1))