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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions docs/compatibility.md
Original file line number Diff line number Diff line change
Expand Up @@ -507,6 +507,9 @@ The following regular expression patterns are not yet supported on the GPU and w
- Possessive quantifiers, such as `a*+`
- Character classes that use union, intersection, or subtraction semantics, such as `[a-d[m-p]]`, `[a-z&&[def]]`,
or `[a-z&&[^bc]]`
- Lookahead/lookbehind groups: `(?=a)`, `(?!a)`, `(?<=a)`, `(?<!a)`
- Independent groups: `(?>a)`
- Named capture groups: `(?<n>a)`
- Empty groups: `()`
- Empty pattern: `""`

Expand Down
4 changes: 2 additions & 2 deletions integration_tests/src/main/python/regexp_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -1042,9 +1042,9 @@ def test_rlike_fallback_possessive_quantifier():
conf=_regexp_conf)

@allow_non_gpu('ProjectExec', 'RLike')
def test_rlike_fallback_lookaheads():
def test_rlike_fallback_lookarounds_independent_named():
gen = mk_str_gen('(\u20ac|\\w){0,3}a[|b*.$\r\n]{0,2}c\\w{0,3}')
for pattern in ['a(?=a*)', 'a(?!a*)']:
for pattern in ['a(?=a*)', 'a(?!a*)', 'a(?<=a)', 'a(?<!a)', 'a(?>a*)', 'a(?<n>a*)']:
assert_gpu_fallback_collect(
lambda spark, pattern=pattern: unary_op_df(spark, gen).selectExpr(
f'a rlike "{pattern}"'),
Expand Down
208 changes: 128 additions & 80 deletions sql-plugin/src/main/scala/com/nvidia/spark/rapids/RegexParser.scala

Large diffs are not rendered by default.

Original file line number Diff line number Diff line change
Expand Up @@ -1211,7 +1211,7 @@ object GpuRegExpUtils {
case QuantifierVariableLength(0, _) => true
case _ => false
}
case RegexGroup(_, term, _) =>
case RegexGroup(_, term) =>
isASTEmptyRepetition(term)
case RegexSequence(parts) =>
parts.forall(isASTEmptyRepetition)
Expand All @@ -1234,8 +1234,8 @@ object GpuRegExpUtils {
def countGroups(pattern: String): Int = {
def countGroups(regexp: RegexAST): Int = {
regexp match {
case RegexGroup(capture, term, _) =>
(if (capture) 1 else 0) + countGroups(term)
case RegexGroup(groupType, term) =>
(if (groupType == RegexGroup.Capturing) 1 else 0) + countGroups(term)
case other => other.children().map(countGroups).sum
Comment thread
greptile-apps[bot] marked this conversation as resolved.
}
}
Expand All @@ -1244,7 +1244,7 @@ object GpuRegExpUtils {

def getChoicesFromRegex(regex: RegexAST): Option[Seq[String]] = {
regex match {
case RegexGroup(_, t, None) =>
case RegexGroup(RegexGroup.Capturing | RegexGroup.NonCapturing, t) =>
getChoicesFromRegex(t)
case RegexChoice(a, b) =>
getChoicesFromRegex(a) match {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -76,8 +76,20 @@ class RegularExpressionParserSuite extends AnyFunSuite {
test("group") {
assert(parse("(a)(b)") ===
RegexSequence(ListBuffer(
RegexGroup(capture = true, RegexSequence(ListBuffer(RegexChar('a'))), None),
RegexGroup(capture = true, RegexSequence(ListBuffer(RegexChar('b'))), None))))
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(RegexChar('a')))),
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(RegexChar('b')))))))
assert(parse("(?:a)(?:b)") ===
RegexSequence(ListBuffer(
RegexGroup(RegexGroup.NonCapturing, RegexSequence(ListBuffer(RegexChar('a')))),
RegexGroup(RegexGroup.NonCapturing, RegexSequence(ListBuffer(RegexChar('b')))))))
assert(parse("(?=a)(?!b)(?<=c)(?<!d)(?>e)(?<n>f)") ===
RegexSequence(ListBuffer(
RegexGroup(RegexGroup.PositiveLookahead, RegexSequence(ListBuffer(RegexChar('a')))),
RegexGroup(RegexGroup.NegativeLookahead, RegexSequence(ListBuffer(RegexChar('b')))),
RegexGroup(RegexGroup.PositiveLookbehind, RegexSequence(ListBuffer(RegexChar('c')))),
RegexGroup(RegexGroup.NegativeLookbehind, RegexSequence(ListBuffer(RegexChar('d')))),
RegexGroup(RegexGroup.Independent, RegexSequence(ListBuffer(RegexChar('e')))),
RegexGroup(RegexGroup.Named("n"), RegexSequence(ListBuffer(RegexChar('f')))))))
}

test("character class") {
Expand Down Expand Up @@ -120,14 +132,14 @@ class RegularExpressionParserSuite extends AnyFunSuite {
assert(parse("\\[([A-Z]+)\\]") ===
RegexSequence(ListBuffer(
RegexEscaped('['),
RegexGroup(capture = true,
RegexGroup(RegexGroup.Capturing,
RegexSequence(ListBuffer(
RegexRepetition(
RegexCharacterClass(negated = false, ListBuffer(
RegexCharacterRange(RegexChar('A'), RegexChar('Z')))),
SimpleQuantifier('+')
)
)), None
))
),
RegexEscaped(']')
))
Expand Down Expand Up @@ -203,24 +215,29 @@ class RegularExpressionParserSuite extends AnyFunSuite {

test("repetition with group containing simple repetition") {
assert(parse("(3?)+") ===
RegexSequence(ListBuffer(RegexRepetition(RegexGroup(capture = true,
RegexSequence(ListBuffer(RegexRepetition(RegexGroup(RegexGroup.Capturing,
RegexSequence(ListBuffer(RegexRepetition(RegexChar('3'),
SimpleQuantifier('?')))), None),SimpleQuantifier('+')))))
SimpleQuantifier('?'))))),SimpleQuantifier('+')))))
}

test("repetition with group containing escape character") {
assert(parse(raw"(\A)+") ===
RegexSequence(ListBuffer(RegexRepetition(RegexGroup(capture = true,
RegexSequence(ListBuffer(RegexEscaped('A'))), None),
RegexSequence(ListBuffer(RegexRepetition(RegexGroup(RegexGroup.Capturing,
RegexSequence(ListBuffer(RegexEscaped('A')))),
SimpleQuantifier('+'))))
)
assert(parse(raw"(?:\A)+") ===
RegexSequence(ListBuffer(RegexRepetition(RegexGroup(RegexGroup.NonCapturing,
RegexSequence(ListBuffer(RegexEscaped('A')))),
SimpleQuantifier('+'))))
)
}

test("group containing choice with repetition") {
assert(parse("(\t+|a)") == RegexSequence(ListBuffer(
RegexGroup(capture = true, RegexChoice(RegexSequence(ListBuffer(
RegexGroup(RegexGroup.Capturing, RegexChoice(RegexSequence(ListBuffer(
RegexRepetition(RegexChar('\t'),SimpleQuantifier('+')))),
RegexSequence(ListBuffer(RegexChar('a')))), None))))
RegexSequence(ListBuffer(RegexChar('a'))))))))
}

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

assert(parse("(?:a?)") === RegexSequence(ListBuffer(
RegexGroup(capture = false, RegexSequence(ListBuffer(
RegexRepetition(RegexChar('a'), SimpleQuantifier('?')))), None))))
RegexGroup(RegexGroup.NonCapturing, RegexSequence(ListBuffer(
RegexRepetition(RegexChar('a'), SimpleQuantifier('?'))))))))
}

test("group not starting with ? is a capturing group") {
assert(parse("(=a)") === RegexSequence(ListBuffer(
RegexGroup(true, RegexSequence(ListBuffer(
RegexChar('='), RegexChar('a'))), None))))
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
RegexChar('='), RegexChar('a')))))))
assert(parse("(!a)") === RegexSequence(ListBuffer(
RegexGroup(true, RegexSequence(ListBuffer(
RegexChar('!'), RegexChar('a'))), None))))
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
RegexChar('!'), RegexChar('a')))))))
}

test("complex expression") {
Expand All @@ -280,42 +297,42 @@ class RegularExpressionParserSuite extends AnyFunSuite {
RegexSequence(ListBuffer(RegexChar('^'),
RegexRepetition(RegexCharacterClass(negated = false, ListBuffer(
RegexChar('+'), RegexEscaped('-'))), SimpleQuantifier('?')),
RegexGroup(capture = true, RegexChoice(RegexSequence(ListBuffer(
RegexGroup(capture = true, RegexSequence(ListBuffer(
RegexGroup(capture = true, RegexChoice(RegexSequence(ListBuffer(
RegexGroup(capture = true, RegexSequence(ListBuffer(
RegexGroup(RegexGroup.Capturing, RegexChoice(RegexSequence(ListBuffer(
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
RegexGroup(RegexGroup.Capturing, RegexChoice(RegexSequence(ListBuffer(
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
RegexRepetition(RegexCharacterClass(negated = false, ListBuffer(
RegexCharacterRange(RegexChar('0'), RegexChar('9')))),
SimpleQuantifier('+')))), None))),
SimpleQuantifier('+'))))))),
RegexChoice(RegexSequence(ListBuffer(
RegexGroup(capture = true, RegexSequence(ListBuffer(
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
RegexRepetition(
RegexCharacterClass(negated = false, ListBuffer(
RegexCharacterRange(RegexChar('0'), RegexChar('9')))),
SimpleQuantifier('*')), RegexEscaped('.'),
RegexRepetition(
RegexCharacterClass(negated = false, ListBuffer(
RegexCharacterRange(RegexChar('0'), RegexChar('9')))),
SimpleQuantifier('+')))), None))), RegexSequence(ListBuffer(
RegexGroup(capture = true, RegexSequence(ListBuffer(
SimpleQuantifier('+'))))))), RegexSequence(ListBuffer(
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
RegexRepetition(
RegexCharacterClass(negated = false, ListBuffer(
RegexCharacterRange(RegexChar('0'), RegexChar('9')))),
SimpleQuantifier('+')), RegexEscaped('.'),
RegexRepetition(RegexCharacterClass(negated = false,
ListBuffer(RegexCharacterRange(RegexChar('0'), RegexChar('9')))),
SimpleQuantifier('*')))), None))))), None),
SimpleQuantifier('*')))))))))),
RegexRepetition(
RegexGroup(capture = true, RegexSequence(ListBuffer(
RegexGroup(RegexGroup.Capturing, RegexSequence(ListBuffer(
RegexCharacterClass(negated = false, ListBuffer(RegexChar('e'), RegexChar('E'))),
RegexRepetition(RegexCharacterClass(negated = false,
ListBuffer(RegexChar('+'), RegexEscaped('-'))),SimpleQuantifier('?')),
RegexRepetition(RegexCharacterClass(negated = false,
ListBuffer(RegexCharacterRange(RegexChar('0'), RegexChar('9')))),
SimpleQuantifier('+')))), None), SimpleQuantifier('?')),
SimpleQuantifier('+'))))), SimpleQuantifier('?')),
RegexRepetition(RegexCharacterClass(negated = false, ListBuffer(
RegexChar('f'), RegexChar('F'), RegexChar('d'), RegexChar('D'))),
SimpleQuantifier('?')))), None))),
SimpleQuantifier('?'))))))),
RegexChoice(RegexSequence(ListBuffer(
RegexChar('I'), RegexChar('n'), RegexChar('f'))),
RegexSequence(ListBuffer(
Expand All @@ -324,7 +341,7 @@ class RegularExpressionParserSuite extends AnyFunSuite {
RegexCharacterClass(negated = false,
ListBuffer(RegexChar('a'), RegexChar('A'))),
RegexCharacterClass(negated = false,
ListBuffer(RegexChar('n'), RegexChar('N'))))))), None),
ListBuffer(RegexChar('n'), RegexChar('N')))))))),
RegexChar('$'))))
}

Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
/*
* Copyright (c) 2024-2025, NVIDIA CORPORATION.
* Copyright (c) 2024-2026, NVIDIA CORPORATION.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
Expand Down Expand Up @@ -97,4 +97,24 @@ class RegularExpressionRewriteSuite extends AnyFunSuite {
)
verifyRewritePattern(patterns, excepted)
}

test("regex rewrite ignores unsupported group types") {
import RegexOptimizationType._
// Lookaround, independent, and named-capture groups are not supported on the GPU. They
// now parse successfully, so the rewrite optimizer must treat them as opaque (only
// capturing/non-capturing groups may be unwrapped) and never mistake them for a simple
// contains/multiple-contains/prefix pattern.
val patterns = Seq(
"(?=abc)",
"(?!abc)",
"(?<=abc)",
"(?<!abc)",
"(?>abc)",
"(?<n>abc)",
"(?<n>abc|def)",
"(?>.*)abc"
)
val excepted = Seq.fill(patterns.length)(NoOptimization)
verifyRewritePattern(patterns, excepted)
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -158,17 +158,53 @@ class RegularExpressionTranspilerSuite extends AnyFunSuite {
"\\z is not supported on GPU")
}

test("cuDF does not support positive or negative lookahead") {
val negPatterns = Seq("a(?!b)", "a(?!b)c?")
negPatterns.foreach(pattern =>
test("cuDF does not support positive or negative lookaround") {
val posLookaheadPatterns = Seq("a(?=b)", "a(?=b)c?")
posLookaheadPatterns.foreach(pattern =>
assertUnsupported(pattern, RegexFindMode,
"Positive lookahead groups are not supported")
)

val negLookaheadPatterns = Seq("a(?!b)", "a(?!b)c?")
negLookaheadPatterns.foreach(pattern =>
assertUnsupported(pattern, RegexFindMode,
"Negative lookahead groups are not supported")
)

val posPatterns = Seq("a(?=b)", "a(?=b)c?")
posPatterns.foreach(pattern =>
val posLookbehindPatterns = Seq("a(?<=b)", "a(?<=b)c?")
posLookbehindPatterns.foreach(pattern =>
assertUnsupported(pattern, RegexFindMode,
"Positive lookahead groups are not supported")
"Positive lookbehind groups are not supported")
)

val negLookbehindPatterns = Seq("a(?<!b)", "a(?<!b)c?")
negLookbehindPatterns.foreach(pattern =>
assertUnsupported(pattern, RegexFindMode,
"Negative lookbehind groups are not supported")
)
}

test("cuDF does not support quantified lookaround, independent, or named capture groups") {
val quantifiedPatterns =
Seq(raw"a(?>\A)+", raw"a(?=\A)+", raw"a(?!\A)+", raw"a(?<=\A)+", raw"a(?<!\A){2}",
raw"a(?<n>\A){1,}")
quantifiedPatterns.foreach(pattern =>
assertUnsupported(pattern, RegexFindMode,
"Repetition of lookaround, independent, or named capture groups is not supported")
)
}

test("cuDF does not support independent or named capture groups") {
val independentPatterns = Seq("a(?>b)", "a(?>b)c?")
independentPatterns.foreach(pattern =>
assertUnsupported(pattern, RegexFindMode,
"Independent groups are not supported")
)

val namedCapturePatterns = Seq("a(?<name>b)", "a(?<name>b)c?")
namedCapturePatterns.foreach(pattern =>
assertUnsupported(pattern, RegexFindMode,
"Named capture groups are not supported")
)
}

Expand All @@ -179,6 +215,17 @@ class RegularExpressionTranspilerSuite extends AnyFunSuite {
)
}

test("cuDF does not support $ followed by lookaround, independent, or named groups") {
val patterns =
Seq("$(?=a)", "$(?!a)", "$(?<=a)", "$(?<!a)", "$(?>a)", "$(?<n>a)",
raw"\Z(?=a)", raw"\Z(?>a)")
patterns.foreach(pattern =>
assertUnsupported(pattern, RegexFindMode,
"Regex sequence $ followed by a lookaround, independent, or named capture " +
"group is not supported")
)
}

test("cuDF does not support quantifier syntax when not quantifying anything") {
// note that we could choose to transpile and escape the '{' and '}' characters
val patterns = Seq("{1,2}", "{1,}", "{1}")
Expand Down Expand Up @@ -380,7 +427,7 @@ class RegularExpressionTranspilerSuite extends AnyFunSuite {
}

test("line anchor $ - find") {
val patterns = Seq("a$", "a$b", "\f$", "$\f","TEST$")
val patterns = Seq("a$", "a$b", "\f$", "$\f", "TEST$")
val inputs = Seq("a", "a\n", "a\r", "a\r\n", "a\f", "\f", "\r", "\u0085", "\u2028",
"\u2029", "\n", "\r\n", "\r\n\r", "\r\n\u0085", "\n\r",
"\n\u0085", "\n\u2028", "\n\u2029", "2+|+??wD\n", "a\r\nb",
Expand Down Expand Up @@ -1381,7 +1428,9 @@ class FuzzRegExp(suggestedChars: String, skipKnownIssues: Boolean = true,
}

private def group(depth: Int) = {
RegexGroup(capture = rr.nextBoolean(), generate(depth + 1), None)
RegexGroup(
if (rr.nextBoolean()) RegexGroup.Capturing else RegexGroup.NonCapturing,
generate(depth + 1))
}

private def repetition(depth: Int) = {
Expand Down
Loading