Skip to content

Commit d0c083a

Browse files
Refactor regex group parsing and explicitly reject unsupported group types (#15331)
Refactors regex group handling to use an explicit `RegexGroup.Type` representation instead of separate `capture` and `lookahead` fields. The regex parser now recognizes additional valid Java group constructs: - Positive and negative lookbehind groups - Independent groups - Named capture groups These constructs remain unsupported by cuDF and are explicitly rejected during transpilation so they reliably fall back to CPU. The change also prevents unsupported groups adjacent to `$` or `\Z` from being silently discarded. Compatibility documentation now explicitly lists all unsupported group types, including lookahead groups. Testing includes: - Parser AST coverage for each group type - Transpiler fallback coverage for standalone, quantified, and anchor-adjacent unsupported groups - Integration coverage confirming `RLike` falls back to CPU Signed-off-by: Igor Peshansky <ipeshansky@nvidia.com>
1 parent b207233 commit d0c083a

7 files changed

Lines changed: 261 additions & 124 deletions

File tree

docs/compatibility.md

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -507,6 +507,9 @@ The following regular expression patterns are not yet supported on the GPU and w
507507
- Possessive quantifiers, such as `a*+`
508508
- Character classes that use union, intersection, or subtraction semantics, such as `[a-d[m-p]]`, `[a-z&&[def]]`,
509509
or `[a-z&&[^bc]]`
510+
- Lookahead/lookbehind groups: `(?=a)`, `(?!a)`, `(?<=a)`, `(?<!a)`
511+
- Independent groups: `(?>a)`
512+
- Named capture groups: `(?<n>a)`
510513
- Empty groups: `()`
511514
- Empty pattern: `""`
512515

integration_tests/src/main/python/regexp_test.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1066,9 +1066,9 @@ def test_rlike_fallback_possessive_quantifier():
10661066
conf=_regexp_conf)
10671067

10681068
@allow_non_gpu('ProjectExec', 'RLike')
1069-
def test_rlike_fallback_lookaheads():
1069+
def test_rlike_fallback_lookarounds_independent_named():
10701070
gen = mk_str_gen('(\u20ac|\\w){0,3}a[|b*.$\r\n]{0,2}c\\w{0,3}')
1071-
for pattern in ['a(?=a*)', 'a(?!a*)']:
1071+
for pattern in ['a(?=a*)', 'a(?!a*)', 'a(?<=a)', 'a(?<!a)', 'a(?>a*)', 'a(?<n>a*)']:
10721072
assert_gpu_fallback_collect(
10731073
lambda spark, pattern=pattern: unary_op_df(spark, gen).selectExpr(
10741074
f'a rlike "{pattern}"'),

sql-plugin/src/main/scala/com/nvidia/spark/rapids/RegexParser.scala

Lines changed: 128 additions & 80 deletions
Large diffs are not rendered by default.

sql-plugin/src/main/scala/org/apache/spark/sql/rapids/stringFunctions.scala

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1211,7 +1211,7 @@ object GpuRegExpUtils {
12111211
case QuantifierVariableLength(0, _) => true
12121212
case _ => false
12131213
}
1214-
case RegexGroup(_, term, _) =>
1214+
case RegexGroup(_, term) =>
12151215
isASTEmptyRepetition(term)
12161216
case RegexSequence(parts) =>
12171217
parts.forall(isASTEmptyRepetition)
@@ -1234,8 +1234,8 @@ object GpuRegExpUtils {
12341234
def countGroups(pattern: String): Int = {
12351235
def countGroups(regexp: RegexAST): Int = {
12361236
regexp match {
1237-
case RegexGroup(capture, term, _) =>
1238-
(if (capture) 1 else 0) + countGroups(term)
1237+
case RegexGroup(groupType, term) =>
1238+
(if (groupType == RegexGroup.Capturing) 1 else 0) + countGroups(term)
12391239
case other => other.children().map(countGroups).sum
12401240
}
12411241
}
@@ -1244,7 +1244,7 @@ object GpuRegExpUtils {
12441244

12451245
def getChoicesFromRegex(regex: RegexAST): Option[Seq[String]] = {
12461246
regex match {
1247-
case RegexGroup(_, t, None) =>
1247+
case RegexGroup(RegexGroup.Capturing | RegexGroup.NonCapturing, t) =>
12481248
getChoicesFromRegex(t)
12491249
case RegexChoice(a, b) =>
12501250
getChoicesFromRegex(a) match {

tests/src/test/scala/com/nvidia/spark/rapids/RegularExpressionParserSuite.scala

Lines changed: 46 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -76,8 +76,20 @@ class RegularExpressionParserSuite extends AnyFunSuite {
7676
test("group") {
7777
assert(parse("(a)(b)") ===
7878
RegexSequence(ListBuffer(
79-
RegexGroup(capture = true, RegexSequence(ListBuffer(RegexChar('a'))), None),
80-
RegexGroup(capture = true, RegexSequence(ListBuffer(RegexChar('b'))), None))))
79+
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(RegexChar('a')))),
80+
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(RegexChar('b')))))))
81+
assert(parse("(?:a)(?:b)") ===
82+
RegexSequence(ListBuffer(
83+
RegexGroup(RegexGroup.NonCapturing, RegexSequence(ListBuffer(RegexChar('a')))),
84+
RegexGroup(RegexGroup.NonCapturing, RegexSequence(ListBuffer(RegexChar('b')))))))
85+
assert(parse("(?=a)(?!b)(?<=c)(?<!d)(?>e)(?<n>f)") ===
86+
RegexSequence(ListBuffer(
87+
RegexGroup(RegexGroup.PositiveLookahead, RegexSequence(ListBuffer(RegexChar('a')))),
88+
RegexGroup(RegexGroup.NegativeLookahead, RegexSequence(ListBuffer(RegexChar('b')))),
89+
RegexGroup(RegexGroup.PositiveLookbehind, RegexSequence(ListBuffer(RegexChar('c')))),
90+
RegexGroup(RegexGroup.NegativeLookbehind, RegexSequence(ListBuffer(RegexChar('d')))),
91+
RegexGroup(RegexGroup.Independent, RegexSequence(ListBuffer(RegexChar('e')))),
92+
RegexGroup(RegexGroup.Named("n"), RegexSequence(ListBuffer(RegexChar('f')))))))
8193
}
8294

8395
test("character class") {
@@ -120,14 +132,14 @@ class RegularExpressionParserSuite extends AnyFunSuite {
120132
assert(parse("\\[([A-Z]+)\\]") ===
121133
RegexSequence(ListBuffer(
122134
RegexEscaped('['),
123-
RegexGroup(capture = true,
135+
RegexGroup(RegexGroup.Capturing,
124136
RegexSequence(ListBuffer(
125137
RegexRepetition(
126138
RegexCharacterClass(negated = false, ListBuffer(
127139
RegexCharacterRange(RegexChar('A'), RegexChar('Z')))),
128140
SimpleQuantifier('+')
129141
)
130-
)), None
142+
))
131143
),
132144
RegexEscaped(']')
133145
))
@@ -203,24 +215,29 @@ class RegularExpressionParserSuite extends AnyFunSuite {
203215

204216
test("repetition with group containing simple repetition") {
205217
assert(parse("(3?)+") ===
206-
RegexSequence(ListBuffer(RegexRepetition(RegexGroup(capture = true,
218+
RegexSequence(ListBuffer(RegexRepetition(RegexGroup(RegexGroup.Capturing,
207219
RegexSequence(ListBuffer(RegexRepetition(RegexChar('3'),
208-
SimpleQuantifier('?')))), None),SimpleQuantifier('+')))))
220+
SimpleQuantifier('?'))))),SimpleQuantifier('+')))))
209221
}
210222

211223
test("repetition with group containing escape character") {
212224
assert(parse(raw"(\A)+") ===
213-
RegexSequence(ListBuffer(RegexRepetition(RegexGroup(capture = true,
214-
RegexSequence(ListBuffer(RegexEscaped('A'))), None),
225+
RegexSequence(ListBuffer(RegexRepetition(RegexGroup(RegexGroup.Capturing,
226+
RegexSequence(ListBuffer(RegexEscaped('A')))),
227+
SimpleQuantifier('+'))))
228+
)
229+
assert(parse(raw"(?:\A)+") ===
230+
RegexSequence(ListBuffer(RegexRepetition(RegexGroup(RegexGroup.NonCapturing,
231+
RegexSequence(ListBuffer(RegexEscaped('A')))),
215232
SimpleQuantifier('+'))))
216233
)
217234
}
218235

219236
test("group containing choice with repetition") {
220237
assert(parse("(\t+|a)") == RegexSequence(ListBuffer(
221-
RegexGroup(capture = true, RegexChoice(RegexSequence(ListBuffer(
238+
RegexGroup(RegexGroup.Capturing, RegexChoice(RegexSequence(ListBuffer(
222239
RegexRepetition(RegexChar('\t'),SimpleQuantifier('+')))),
223-
RegexSequence(ListBuffer(RegexChar('a')))), None))))
240+
RegexSequence(ListBuffer(RegexChar('a'))))))))
224241
}
225242

226243
test("multiple choice (2)") {
@@ -244,17 +261,17 @@ class RegularExpressionParserSuite extends AnyFunSuite {
244261
assert(e.getMessage.startsWith("Base expression cannot start with quantifier"))
245262

246263
assert(parse("(?:a?)") === RegexSequence(ListBuffer(
247-
RegexGroup(capture = false, RegexSequence(ListBuffer(
248-
RegexRepetition(RegexChar('a'), SimpleQuantifier('?')))), None))))
264+
RegexGroup(RegexGroup.NonCapturing, RegexSequence(ListBuffer(
265+
RegexRepetition(RegexChar('a'), SimpleQuantifier('?'))))))))
249266
}
250267

251268
test("group not starting with ? is a capturing group") {
252269
assert(parse("(=a)") === RegexSequence(ListBuffer(
253-
RegexGroup(true, RegexSequence(ListBuffer(
254-
RegexChar('='), RegexChar('a'))), None))))
270+
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
271+
RegexChar('='), RegexChar('a')))))))
255272
assert(parse("(!a)") === RegexSequence(ListBuffer(
256-
RegexGroup(true, RegexSequence(ListBuffer(
257-
RegexChar('!'), RegexChar('a'))), None))))
273+
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
274+
RegexChar('!'), RegexChar('a')))))))
258275
}
259276

260277
test("complex expression") {
@@ -280,42 +297,42 @@ class RegularExpressionParserSuite extends AnyFunSuite {
280297
RegexSequence(ListBuffer(RegexChar('^'),
281298
RegexRepetition(RegexCharacterClass(negated = false, ListBuffer(
282299
RegexChar('+'), RegexEscaped('-'))), SimpleQuantifier('?')),
283-
RegexGroup(capture = true, RegexChoice(RegexSequence(ListBuffer(
284-
RegexGroup(capture = true, RegexSequence(ListBuffer(
285-
RegexGroup(capture = true, RegexChoice(RegexSequence(ListBuffer(
286-
RegexGroup(capture = true, RegexSequence(ListBuffer(
300+
RegexGroup(RegexGroup.Capturing, RegexChoice(RegexSequence(ListBuffer(
301+
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
302+
RegexGroup(RegexGroup.Capturing, RegexChoice(RegexSequence(ListBuffer(
303+
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
287304
RegexRepetition(RegexCharacterClass(negated = false, ListBuffer(
288305
RegexCharacterRange(RegexChar('0'), RegexChar('9')))),
289-
SimpleQuantifier('+')))), None))),
306+
SimpleQuantifier('+'))))))),
290307
RegexChoice(RegexSequence(ListBuffer(
291-
RegexGroup(capture = true, RegexSequence(ListBuffer(
308+
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
292309
RegexRepetition(
293310
RegexCharacterClass(negated = false, ListBuffer(
294311
RegexCharacterRange(RegexChar('0'), RegexChar('9')))),
295312
SimpleQuantifier('*')), RegexEscaped('.'),
296313
RegexRepetition(
297314
RegexCharacterClass(negated = false, ListBuffer(
298315
RegexCharacterRange(RegexChar('0'), RegexChar('9')))),
299-
SimpleQuantifier('+')))), None))), RegexSequence(ListBuffer(
300-
RegexGroup(capture = true, RegexSequence(ListBuffer(
316+
SimpleQuantifier('+'))))))), RegexSequence(ListBuffer(
317+
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
301318
RegexRepetition(
302319
RegexCharacterClass(negated = false, ListBuffer(
303320
RegexCharacterRange(RegexChar('0'), RegexChar('9')))),
304321
SimpleQuantifier('+')), RegexEscaped('.'),
305322
RegexRepetition(RegexCharacterClass(negated = false,
306323
ListBuffer(RegexCharacterRange(RegexChar('0'), RegexChar('9')))),
307-
SimpleQuantifier('*')))), None))))), None),
324+
SimpleQuantifier('*')))))))))),
308325
RegexRepetition(
309-
RegexGroup(capture = true, RegexSequence(ListBuffer(
326+
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
310327
RegexCharacterClass(negated = false, ListBuffer(RegexChar('e'), RegexChar('E'))),
311328
RegexRepetition(RegexCharacterClass(negated = false,
312329
ListBuffer(RegexChar('+'), RegexEscaped('-'))),SimpleQuantifier('?')),
313330
RegexRepetition(RegexCharacterClass(negated = false,
314331
ListBuffer(RegexCharacterRange(RegexChar('0'), RegexChar('9')))),
315-
SimpleQuantifier('+')))), None), SimpleQuantifier('?')),
332+
SimpleQuantifier('+'))))), SimpleQuantifier('?')),
316333
RegexRepetition(RegexCharacterClass(negated = false, ListBuffer(
317334
RegexChar('f'), RegexChar('F'), RegexChar('d'), RegexChar('D'))),
318-
SimpleQuantifier('?')))), None))),
335+
SimpleQuantifier('?'))))))),
319336
RegexChoice(RegexSequence(ListBuffer(
320337
RegexChar('I'), RegexChar('n'), RegexChar('f'))),
321338
RegexSequence(ListBuffer(
@@ -324,7 +341,7 @@ class RegularExpressionParserSuite extends AnyFunSuite {
324341
RegexCharacterClass(negated = false,
325342
ListBuffer(RegexChar('a'), RegexChar('A'))),
326343
RegexCharacterClass(negated = false,
327-
ListBuffer(RegexChar('n'), RegexChar('N'))))))), None),
344+
ListBuffer(RegexChar('n'), RegexChar('N')))))))),
328345
RegexChar('$'))))
329346
}
330347

tests/src/test/scala/com/nvidia/spark/rapids/RegularExpressionRewriteSuite.scala

Lines changed: 21 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
/*
2-
* Copyright (c) 2024-2025, NVIDIA CORPORATION.
2+
* Copyright (c) 2024-2026, NVIDIA CORPORATION.
33
*
44
* Licensed under the Apache License, Version 2.0 (the "License");
55
* you may not use this file except in compliance with the License.
@@ -97,4 +97,24 @@ class RegularExpressionRewriteSuite extends AnyFunSuite {
9797
)
9898
verifyRewritePattern(patterns, excepted)
9999
}
100+
101+
test("regex rewrite ignores unsupported group types") {
102+
import RegexOptimizationType._
103+
// Lookaround, independent, and named-capture groups are not supported on the GPU. They
104+
// now parse successfully, so the rewrite optimizer must treat them as opaque (only
105+
// capturing/non-capturing groups may be unwrapped) and never mistake them for a simple
106+
// contains/multiple-contains/prefix pattern.
107+
val patterns = Seq(
108+
"(?=abc)",
109+
"(?!abc)",
110+
"(?<=abc)",
111+
"(?<!abc)",
112+
"(?>abc)",
113+
"(?<n>abc)",
114+
"(?<n>abc|def)",
115+
"(?>.*)abc"
116+
)
117+
val excepted = Seq.fill(patterns.length)(NoOptimization)
118+
verifyRewritePattern(patterns, excepted)
119+
}
100120
}

tests/src/test/scala/com/nvidia/spark/rapids/RegularExpressionTranspilerSuite.scala

Lines changed: 57 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -158,17 +158,53 @@ class RegularExpressionTranspilerSuite extends AnyFunSuite {
158158
"\\z is not supported on GPU")
159159
}
160160

161-
test("cuDF does not support positive or negative lookahead") {
162-
val negPatterns = Seq("a(?!b)", "a(?!b)c?")
163-
negPatterns.foreach(pattern =>
161+
test("cuDF does not support positive or negative lookaround") {
162+
val posLookaheadPatterns = Seq("a(?=b)", "a(?=b)c?")
163+
posLookaheadPatterns.foreach(pattern =>
164+
assertUnsupported(pattern, RegexFindMode,
165+
"Positive lookahead groups are not supported")
166+
)
167+
168+
val negLookaheadPatterns = Seq("a(?!b)", "a(?!b)c?")
169+
negLookaheadPatterns.foreach(pattern =>
164170
assertUnsupported(pattern, RegexFindMode,
165171
"Negative lookahead groups are not supported")
166172
)
167173

168-
val posPatterns = Seq("a(?=b)", "a(?=b)c?")
169-
posPatterns.foreach(pattern =>
174+
val posLookbehindPatterns = Seq("a(?<=b)", "a(?<=b)c?")
175+
posLookbehindPatterns.foreach(pattern =>
170176
assertUnsupported(pattern, RegexFindMode,
171-
"Positive lookahead groups are not supported")
177+
"Positive lookbehind groups are not supported")
178+
)
179+
180+
val negLookbehindPatterns = Seq("a(?<!b)", "a(?<!b)c?")
181+
negLookbehindPatterns.foreach(pattern =>
182+
assertUnsupported(pattern, RegexFindMode,
183+
"Negative lookbehind groups are not supported")
184+
)
185+
}
186+
187+
test("cuDF does not support quantified lookaround, independent, or named capture groups") {
188+
val quantifiedPatterns =
189+
Seq(raw"a(?>\A)+", raw"a(?=\A)+", raw"a(?!\A)+", raw"a(?<=\A)+", raw"a(?<!\A){2}",
190+
raw"a(?<n>\A){1,}")
191+
quantifiedPatterns.foreach(pattern =>
192+
assertUnsupported(pattern, RegexFindMode,
193+
"Repetition of lookaround, independent, or named capture groups is not supported")
194+
)
195+
}
196+
197+
test("cuDF does not support independent or named capture groups") {
198+
val independentPatterns = Seq("a(?>b)", "a(?>b)c?")
199+
independentPatterns.foreach(pattern =>
200+
assertUnsupported(pattern, RegexFindMode,
201+
"Independent groups are not supported")
202+
)
203+
204+
val namedCapturePatterns = Seq("a(?<name>b)", "a(?<name>b)c?")
205+
namedCapturePatterns.foreach(pattern =>
206+
assertUnsupported(pattern, RegexFindMode,
207+
"Named capture groups are not supported")
172208
)
173209
}
174210

@@ -179,6 +215,17 @@ class RegularExpressionTranspilerSuite extends AnyFunSuite {
179215
)
180216
}
181217

218+
test("cuDF does not support $ followed by lookaround, independent, or named groups") {
219+
val patterns =
220+
Seq("$(?=a)", "$(?!a)", "$(?<=a)", "$(?<!a)", "$(?>a)", "$(?<n>a)",
221+
raw"\Z(?=a)", raw"\Z(?>a)")
222+
patterns.foreach(pattern =>
223+
assertUnsupported(pattern, RegexFindMode,
224+
"Regex sequence $ followed by a lookaround, independent, or named capture " +
225+
"group is not supported")
226+
)
227+
}
228+
182229
test("cuDF does not support quantifier syntax when not quantifying anything") {
183230
// note that we could choose to transpile and escape the '{' and '}' characters
184231
val patterns = Seq("{1,2}", "{1,}", "{1}")
@@ -380,7 +427,7 @@ class RegularExpressionTranspilerSuite extends AnyFunSuite {
380427
}
381428

382429
test("line anchor $ - find") {
383-
val patterns = Seq("a$", "a$b", "\f$", "$\f","TEST$")
430+
val patterns = Seq("a$", "a$b", "\f$", "$\f", "TEST$")
384431
val inputs = Seq("a", "a\n", "a\r", "a\r\n", "a\f", "\f", "\r", "\u0085", "\u2028",
385432
"\u2029", "\n", "\r\n", "\r\n\r", "\r\n\u0085", "\n\r",
386433
"\n\u0085", "\n\u2028", "\n\u2029", "2+|+??wD\n", "a\r\nb",
@@ -1394,7 +1441,9 @@ class FuzzRegExp(suggestedChars: String, skipKnownIssues: Boolean = true,
13941441
}
13951442

13961443
private def group(depth: Int) = {
1397-
RegexGroup(capture = rr.nextBoolean(), generate(depth + 1), None)
1444+
RegexGroup(
1445+
if (rr.nextBoolean()) RegexGroup.Capturing else RegexGroup.NonCapturing,
1446+
generate(depth + 1))
13981447
}
13991448

14001449
private def repetition(depth: Int) = {

0 commit comments

Comments
 (0)