Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
150 changes: 71 additions & 79 deletions sql-plugin/src/main/scala/com/nvidia/spark/rapids/RegexParser.scala
Original file line number Diff line number Diff line change
Expand Up @@ -161,7 +161,9 @@ class RegexParser(pattern: String) {
* {n,}
* {n,m} (only valid if m >= n)
*/
private def tryParseBraceQuantifier(): Option[RegexQuantifier] = {
private def tryParseBraceQuantifier(): Option[RegexQuantifier.Base] = {
import RegexQuantifier._

// The caller restores its position when this is a literal brace rather than a quantifier.
consumeExpected('{')
consumeInt.flatMap { minLength =>

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[Really optional] Almost certainly a follow-up change, but just noting here that we still try to parse the quantifier and fall back on treating invalid strings as literals. At some point we might want to match Java's behavior, which commits to the quantifier given the { prefix and errors out for invalid ones (like a{, a{}, or a{x}).

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Filed #15962 to track this follow-up.

Expand All @@ -171,13 +173,13 @@ class RegexParser(pattern: String) {
val maxLength = consumeInt()
if (peek().contains('}') && maxLength.forall(_ >= minLength)) {
consumeExpected('}')
Some(QuantifierVariableLength(minLength, maxLength))
Some(Variable(minLength, maxLength))
} else {
None

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Wanted to see if I could also get an answer to #15478 (comment) here, which seems to have gotten lost in the noise… Quoting here for ease of access:

[Optional] Ideally, we would error out in this case. Java does:

  • Pattern.compile("a{3,4!") => PatternSyntaxException: Unclosed counted closure near index 5
  • Pattern.compile("a{3,2}") => PatternSyntaxException: Illegal repetition range near index 5

With the Pattern.compile guard before we even parse the regex, this is technically unreachable outside of parseUnchecked anyway…

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agreed that these should be errors on the production path. RegexParser.parse() calls Pattern.compile(pattern) before parseUnchecked(), so both examples already throw PatternSyntaxException before the custom parser reaches this branch. I kept tryParseBraceQuantifier returning None for malformed brace syntax because parseUnchecked() is the package-private path that deliberately bypasses Java validation. This preserves the existing validation boundary instead of duplicating Java syntax checks in the unchecked parser.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

FWIW, I think sticking closer to Java semantics helps maintain cleaner code and avoids questions about consistency… Since these edge cases are only reachable from tests, it would be easiest to match Java to decide on the expected behavior.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks, updated.

}
case Some('}') =>
consumeExpected('}')
Some(QuantifierFixedLength(minLength))
Some(Fixed(minLength))
case _ =>
None
}
Expand All @@ -191,9 +193,15 @@ class RegexParser(pattern: String) {
val baseQuantifier = peek() match {
case Some('{') =>
tryParseBraceQuantifier()
case Some(ch) if "*+?".contains(ch) =>
case Some('*') =>
consume()
Comment thread
igorpeshansky marked this conversation as resolved.
Outdated
Some(ZeroOrMore)
case Some('+') =>
consume()
Some(SimpleQuantifier(ch))
Some(OneOrMore)
case Some('?') =>
consume()
Some(ZeroOrOne)
case _ => None
}

Expand All @@ -208,7 +216,7 @@ class RegexParser(pattern: String) {
Possessive
case _ => Greedy
}
val quantifier = base.withMode(mode)
val quantifier = RegexQuantifier(base, mode)
// Point diagnostics at the modifier when present, otherwise at the base quantifier.
quantifier.position = Some(if (mode == Greedy) start else pos - 1)
Some(quantifier)
Comment thread
igorpeshansky marked this conversation as resolved.
Outdated
Expand Down Expand Up @@ -781,6 +789,8 @@ sealed case class RegexRewriteFlags(
RegexSplitMode if performing a split (string_split)
*/
class CudfRegexTranspiler(mode: RegexMode) {
import RegexQuantifier._

// cuDF reads at most three count digits for a repetition.
// https://github.com/NVIDIA/cudf/blob/7a6f5c1a/cpp/src/strings/regex/regcomp.cpp#L684
private val maxRepetitionCount = 999
Expand All @@ -789,11 +799,11 @@ class CudfRegexTranspiler(mode: RegexMode) {
'b' -> '\b', 'e' -> '\u001b')

private def exceedsCudfRepetitionCountLimit(quantifier: RegexQuantifier): Boolean = {
quantifier match {
case QuantifierFixedLength(length, _) => length > maxRepetitionCount
case QuantifierVariableLength(minLength, maxLength, _) =>
quantifier.base match {
case Fixed(length) => length > maxRepetitionCount
case Variable(minLength, maxLength) =>
minLength > maxRepetitionCount || maxLength.exists(_ > maxRepetitionCount)
case _: SimpleQuantifier => false
case ZeroOrOne | ZeroOrMore | OneOrMore => false
}
}

Expand Down Expand Up @@ -1580,7 +1590,7 @@ class CudfRegexTranspiler(mode: RegexMode) {
s"cuDF does not support repetition counts greater than $maxRepetitionCount",
repetition.position)

case (_, q) if q.isPossessive =>
case (_, q @ RegexQuantifier(_, Possessive)) =>
Comment thread
igorpeshansky marked this conversation as resolved.
Outdated
throw new RegexUnsupportedException(
s"Possessive quantifier ${q.toRegexString} not supported", q.position)

Expand All @@ -1594,12 +1604,12 @@ class CudfRegexTranspiler(mode: RegexMode) {
"regexp_split on GPU does not support empty match repetition consistently with Spark",
quantifier.position)

case (_, QuantifierVariableLength(0, Some(0), _)) if mode != RegexFindMode =>
case (_, RegexQuantifier(Variable(0, Some(0)), _)) if mode != RegexFindMode =>
Comment thread
igorpeshansky marked this conversation as resolved.
Outdated
throw new RegexUnsupportedException(
"regex_replace and regex_split on GPU do not support repetition with {0,0}",
quantifier.position)

case (_, QuantifierFixedLength(0, _)) if mode != RegexFindMode =>
case (_, RegexQuantifier(Fixed(0), _)) if mode != RegexFindMode =>
throw new RegexUnsupportedException(
"regex_replace and regex_split on GPU do not support repetition with {0}",
quantifier.position)
Expand All @@ -1609,12 +1619,13 @@ class CudfRegexTranspiler(mode: RegexMode) {
"Repetition of lookaround, independent, or named capture groups is not supported",
g.position)

case (RegexGroup(groupType, term), SimpleQuantifier(ch, _))
if "+*".contains(ch) && !isSupportedRepetitionBase(term) =>
(term, ch) match {
case (RegexGroup(groupType, term), RegexQuantifier(simpleBase, _))
Comment thread
igorpeshansky marked this conversation as resolved.
Outdated
if (simpleBase == OneOrMore || simpleBase == ZeroOrMore) &&
!isSupportedRepetitionBase(term) =>
(term, simpleBase) match {
// \Z is not supported in groups
case (RegexEscaped('A'), '+') |
(RegexSequence(ListBuffer(RegexEscaped('A'))), '+') =>
case (RegexEscaped('A'), OneOrMore) |
(RegexSequence(ListBuffer(RegexEscaped('A'))), OneOrMore) =>
// (\A)+ can be transpiled to (\A) (dropping the repetition)
// we use rewrite(...) here to handle logic regarding modes
// (\A is not supported in RegexSplitMode)
Expand All @@ -1627,7 +1638,7 @@ class CudfRegexTranspiler(mode: RegexMode) {
s"cuDF does not support repetition of group containing: " +
s"${unsupportedTerm.toRegexString}", term.position)
}
case (RegexGroup(groupType, term), QuantifierVariableLength(_, _, _))
case (RegexGroup(groupType, term), RegexQuantifier(Variable(_, _), _))
if !isSupportedRepetitionBase(term) =>
term match {
// \Z is not supported in groups
Expand All @@ -1646,7 +1657,7 @@ class CudfRegexTranspiler(mode: RegexMode) {
s"cuDF does not support repetition of group containing: " +
s"${unsupportedTerm.toRegexString}", term.position)
}
case (RegexGroup(groupType, term), QuantifierFixedLength(n, _))
case (RegexGroup(groupType, term), RegexQuantifier(Fixed(n), _))
if !isSupportedRepetitionBase(term) =>
term match {
// \Z is not supported in groups
Expand All @@ -1664,25 +1675,27 @@ class CudfRegexTranspiler(mode: RegexMode) {
s"cuDF does not support repetition of group containing: " +
s"${unsupportedTerm.toRegexString}", term.position)
}
case (RegexGroup(_, term), SimpleQuantifier('?', _)) =>
case (RegexGroup(_, term), RegexQuantifier(ZeroOrOne, _)) =>
if (isEntirelyWordBoundary(term) || isEntirelyLineAnchor(term)) {
throw new RegexUnsupportedException(
s"cuDF does not support repetition of: ${term.toRegexString}", term.position)
}
RegexRepetition(rewrite(base, None, flags), quantifier)
case (RegexEscaped(ch), SimpleQuantifier('+', _)) if "AZ".contains(ch) =>
case (RegexEscaped(ch), RegexQuantifier(OneOrMore, _))
Comment thread
igorpeshansky marked this conversation as resolved.
Outdated
if "AZ".contains(ch) =>
// \A+ can be transpiled to \A (dropping the repetition)
// \Z+ can be transpiled to \Z (dropping the repetition)
// we use rewrite(...) here to handle logic regarding modes
// (\A and \Z are not supported in RegexSplitMode)
rewrite(base, previous, flags)
// NOTE: \A* can be transpiled to \A?
// however, \A? is not supported in libcudf yet
case (RegexEscaped(ch), QuantifierFixedLength(n, _)) if n > 0 && "AZ".contains(ch) =>
case (RegexEscaped(ch), RegexQuantifier(Fixed(n), _))
if n > 0 && "AZ".contains(ch) =>
// \A{2} can be transpiled to \A (dropping the repetition)
// \Z{2} can be transpiled to \Z (dropping the repetition)
rewrite(base, previous, flags)
case (RegexEscaped(ch), QuantifierVariableLength(n, _, _))
case (RegexEscaped(ch), RegexQuantifier(Variable(n, _), _))
if n > 0 && "AZ".contains(ch) =>
// \A{1,5} can be transpiled to \A (dropping the repetition)
// \Z{1,} can be transpiled to \Z (dropping the repetition)
Expand Down Expand Up @@ -1742,14 +1755,14 @@ class CudfRegexTranspiler(mode: RegexMode) {
(ll, rr) match {
// ll = lazyQuantifier inside a choice
case (RegexSequence(ListBuffer(RegexRepetition(
_, SimpleQuantifier('?', RegexQuantifier.Reluctant)))), _) |
_, RegexQuantifier(ZeroOrOne, Reluctant)))), _) |
// rr = lazyQuantifier inside a choice
(_, RegexSequence(ListBuffer(RegexRepetition(
_, SimpleQuantifier('?', RegexQuantifier.Reluctant))))) =>
_, RegexQuantifier(ZeroOrOne, Reluctant))))) =>
throw new RegexUnsupportedException(
"cuDF does not support lazy quantifier inside choice", r.position)
case (_, RegexChoice(RegexSequence(_), RegexSequence(ListBuffer(RegexRepetition(
RegexEscaped('A'), SimpleQuantifier('?', _)), _)))) =>
RegexEscaped('A'), RegexQuantifier(ZeroOrOne, _)), _)))) =>
throw new RegexUnsupportedException("Invalid regex pattern at position", r.position)
case _ =>
}
Expand All @@ -1765,12 +1778,12 @@ class CudfRegexTranspiler(mode: RegexMode) {
}
part match {
case RegexRepetition(base, quantifier) => (base, quantifier) match {
case (_, QuantifierVariableLength(0, Some(0), _)) =>
case (_, RegexQuantifier(Variable(0, Some(0)), _)) =>
throw new RegexUnsupportedException(
"Repetition with {0,0} not supported in capture groups",
quantifier.position)

case (_, QuantifierFixedLength(0, _)) =>
case (_, RegexQuantifier(Fixed(0), _)) =>
throw new RegexUnsupportedException(
"Repetition with {0} not supported in capture groups",
quantifier.position)
Expand Down Expand Up @@ -2043,73 +2056,52 @@ sealed case class RegexRepetition(a: RegexAST, quantifier: RegexQuantifier) exte
}

object RegexQuantifier {
sealed trait Base
case object ZeroOrOne extends Base
case object ZeroOrMore extends Base
case object OneOrMore extends Base
final case class Fixed(length: Int) extends Base
final case class Variable(minLength: Int, maxLength: Option[Int]) extends Base
Comment thread
igorpeshansky marked this conversation as resolved.
Outdated

sealed trait Mode
case object Greedy extends Mode
case object Reluctant extends Mode
case object Possessive extends Mode
}

sealed trait RegexQuantifier {
sealed case class RegexQuantifier(
base: RegexQuantifier.Base,
mode: RegexQuantifier.Mode = RegexQuantifier.Greedy) {
Comment thread
igorpeshansky marked this conversation as resolved.
Outdated
import RegexQuantifier._

def mode: Mode
def minLength: Int
protected def baseToRegexString: String
protected def copyWithMode(newMode: Mode): RegexQuantifier

var position: Option[Int] = None

final def withMode(newMode: Mode): RegexQuantifier = {
val updated = copyWithMode(newMode)
updated.position = position
updated
def minLength: Int = base match {
case ZeroOrOne | ZeroOrMore => 0
case OneOrMore => 1
case Fixed(length) => length
case Variable(minLength, _) => minLength
}

final def isPossessive: Boolean = mode == Possessive
final def toRegexString: String = {
def toRegexString: String = {
val baseString = base match {
case ZeroOrOne => "?"
case ZeroOrMore => "*"
case OneOrMore => "+"
case Fixed(length) => s"{$length}"
case Variable(minLength, maxLength) =>
maxLength match {
case Some(max) => s"{$minLength,$max}"
case None => s"{$minLength,}"
}
Comment thread
igorpeshansky marked this conversation as resolved.
Outdated
}
val suffix = mode match {
case Greedy => ""
case Reluctant => "?"
case Possessive => "+"
}
s"$baseToRegexString$suffix"
}
}

sealed case class SimpleQuantifier(
ch: Char,
mode: RegexQuantifier.Mode = RegexQuantifier.Greedy) extends RegexQuantifier {
override def minLength: Int = if (ch == '+') 1 else 0
override protected def baseToRegexString: String = ch.toString
override protected def copyWithMode(newMode: RegexQuantifier.Mode): RegexQuantifier =
copy(mode = newMode)
}

sealed case class QuantifierFixedLength(
length: Int,
mode: RegexQuantifier.Mode = RegexQuantifier.Greedy)
extends RegexQuantifier {
override def minLength: Int = length
override protected def baseToRegexString: String = s"{$length}"
override protected def copyWithMode(newMode: RegexQuantifier.Mode): RegexQuantifier =
copy(mode = newMode)
}

sealed case class QuantifierVariableLength(
minLength: Int,
maxLength: Option[Int],
mode: RegexQuantifier.Mode = RegexQuantifier.Greedy)
extends RegexQuantifier {
override protected def baseToRegexString: String = {
maxLength match {
case Some(max) =>
s"{$minLength,$max}"
case _ =>
s"{$minLength,}"
}
s"$baseString$suffix"
}
override protected def copyWithMode(newMode: RegexQuantifier.Mode): RegexQuantifier =
copy(mode = newMode)
}

sealed trait RegexCharacterClassComponent extends RegexAST
Expand Down Expand Up @@ -2348,7 +2340,7 @@ object RegexRewrite {
private def isWildcard(ast: RegexAST): Boolean = {
ast match {
case RegexRepetition(RegexChar('.'),
SimpleQuantifier('*', RegexQuantifier.Greedy)) => true
RegexQuantifier(RegexQuantifier.ZeroOrMore, RegexQuantifier.Greedy)) => true
case RegexSequence(parts) if parts.forall(isWildcard) => true
case RegexGroup(groupType, term) if isTransparentGroup(groupType) && isWildcard(term) => true
case _ => false
Expand Down
Loading
Loading