Skip to content
Closed
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
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,26 @@ case class SqlStatementSplitResult(
def isEmpty: Boolean = completeStatements.isEmpty && partialStatement.isEmpty
}

/** A split SQL statement together with its 0-based start in the original input. */
private[sql] case class PositionedSqlStatement(
statement: String,
terminator: String,
start: Int) {
def length: Int = statement.length
Comment thread
srielau marked this conversation as resolved.
}

/** Internal splitter result that retains source positions. */
private[sql] case class PositionedSqlStatementSplitResult(
completeStatements: Seq[PositionedSqlStatement],
partialStatement: Option[PositionedSqlStatement],
hasUnclosedComment: Boolean) {

def withoutPositions: SqlStatementSplitResult = SqlStatementSplitResult(
completeStatements.map(s => SqlStatement(s.statement, s.terminator)),
partialStatement.map(_.statement).getOrElse(""),
hasUnclosedComment)
}

/**
* A parser-based SQL statement splitter, inspired by Trino's
* `io.trino.cli.lexer.StatementSplitter`.
Expand Down Expand Up @@ -111,7 +131,7 @@ object SqlStatementSplitter {

/** Split the given SQL text into individual statements at `;` boundaries. */
def split(sqlText: String): SqlStatementSplitResult =
split(sqlText, identity)
splitWithPositions(sqlText, identity).withoutPositions

/**
* Split the given SQL text, applying `validationPreprocess` to each candidate
Expand All @@ -121,10 +141,26 @@ object SqlStatementSplitter {
*
* Pass `identity` for a pure original-text splitter (the default).
*/
def split(sqlText: String, validationPreprocess: String => String): SqlStatementSplitResult = {
def split(
sqlText: String,
validationPreprocess: String => String): SqlStatementSplitResult =
splitWithPositions(sqlText, validationPreprocess).withoutPositions

/**
* Split SQL while retaining each statement's source position. This is used by
* parse-only tooling that reports spans into the original input.
*/
private[sql] def splitWithPositions(
sqlText: String,
validationPreprocess: String => String): PositionedSqlStatementSplitResult = {
require(sqlText != null, "sqlText must not be null")
require(validationPreprocess != null, "validationPreprocess must not be null")

// CodePointCharStream token offsets count Unicode code points, while
// String offsets and lengths count UTF-16 code units. Map once so each
// token-boundary lookup is O(1) rather than rescanning the prefix.
val toUtf16 = utf16Offsets(sqlText)

val lexer = new SqlBaseLexer(new UpperCaseCharStream(CharStreams.fromString(sqlText)))
lexer.removeErrorListeners()
val tokenStream = new CommonTokenStream(lexer)
Expand All @@ -142,8 +178,9 @@ object SqlStatementSplitter {
acc.toArray
}

val completeStatements = mutable.ArrayBuffer.empty[SqlStatement]
val completeStatements = mutable.ArrayBuffer.empty[PositionedSqlStatement]
val buffer = new StringBuilder()
var bufferStart = -1
// Whether `buffer` contains any non-hidden token (i.e. any actual SQL content
// beyond whitespace and comments). Chunks that only contain whitespace/comments
// are dropped, matching the spark-sql CLI's long-standing behavior.
Expand All @@ -158,6 +195,31 @@ object SqlStatementSplitter {
// interpretation (e.g. `double_quoted_identifiers`).
val conf = SqlApiConf.get

def appendToken(token: Token): Unit = {
if (buffer.isEmpty) bufferStart = toUtf16(token.getStartIndex)
buffer.append(token.getText)
}

def resetBuffer(): Unit = {
buffer.setLength(0)
bufferStart = -1
bufferHasContent = false
}

def positionedStatement(terminator: String): Option[PositionedSqlStatement] = {
val raw = buffer.toString
val (leadingWhitespace, statement) = trimSqlWhitespace(raw)
if (statement.isEmpty) {
None
} else {
assert(bufferStart >= 0)
Some(PositionedSqlStatement(
statement,
terminator,
bufferStart + leadingWhitespace))
}
}

while (!stopOuter && index < numTokens) {
val startIdx = nextSignificantTokenIndex(tokenStream, index)
if (startIdx < 0) {
Expand All @@ -167,7 +229,7 @@ object SqlStatementSplitter {
while (index < numTokens) {
val tok = tokenStream.get(index)
index += 1
if (tok.getType != Token.EOF) buffer.append(tok.getText)
if (tok.getType != Token.EOF) appendToken(tok)
}
stopOuter = true
} else if (tokenStream.get(startIdx).getType == SqlBaseLexer.SEMICOLON) {
Expand All @@ -192,7 +254,8 @@ object SqlStatementSplitter {
while (!parsedOk && !failedNonEof && d < delimiterPositions.length) {
val candidateEnd = delimiterPositions(d)
tryParseRegion(
sqlText, tokenStream, startIdx, candidateEnd, validationPreprocess, conf) match {
sqlText, toUtf16, tokenStream, startIdx, candidateEnd,
validationPreprocess, conf) match {
case ParsedOk =>
parsedOk = true
matchedDelimIdx = candidateEnd
Expand All @@ -216,17 +279,13 @@ object SqlStatementSplitter {
while (index < matchedDelimIdx) {
val tok = tokenStream.get(index)
if (tok.getChannel != Token.HIDDEN_CHANNEL) bufferHasContent = true
buffer.append(tok.getText)
appendToken(tok)
index += 1
}
if (bufferHasContent) {
val stmt = buffer.toString.trim
if (stmt.nonEmpty) {
completeStatements += SqlStatement(stmt, terminator)
}
positionedStatement(terminator).foreach(completeStatements += _)
}
buffer.setLength(0)
bufferHasContent = false
resetBuffer()
index = matchedDelimIdx + 1
delimSearchStart = d + 1
} else if (failedNonEof) {
Expand All @@ -253,17 +312,13 @@ object SqlStatementSplitter {
stopInner = true
} else if (token.getType == SqlBaseLexer.SEMICOLON) {
if (bufferHasContent) {
val stmt = buffer.toString.trim
if (stmt.nonEmpty) {
completeStatements += SqlStatement(stmt, token.getText)
}
positionedStatement(token.getText).foreach(completeStatements += _)
Comment thread
srielau marked this conversation as resolved.
}
buffer.setLength(0)
bufferHasContent = false
resetBuffer()
stopInner = true
} else {
if (token.getChannel != Token.HIDDEN_CHANNEL) bufferHasContent = true
buffer.append(token.getText)
appendToken(token)
}
}
} else {
Expand All @@ -275,7 +330,7 @@ object SqlStatementSplitter {
index += 1
if (tok.getType != Token.EOF) {
if (tok.getChannel != Token.HIDDEN_CHANNEL) bufferHasContent = true
buffer.append(tok.getText)
appendToken(tok)
}
}
stopOuter = true
Expand All @@ -285,8 +340,11 @@ object SqlStatementSplitter {

val unclosed = lexer.has_unclosed_bracketed_comment
val partial =
if (bufferHasContent || unclosed) buffer.toString.trim else ""
SqlStatementSplitResult(completeStatements.toSeq, partial, unclosed && partial.nonEmpty)
if (bufferHasContent || unclosed) positionedStatement("") else None
PositionedSqlStatementSplitResult(
completeStatements.toSeq,
partial,
unclosed && partial.nonEmpty)
}

/** Outcome of attempting to parse one statement candidate. */
Expand All @@ -301,13 +359,14 @@ object SqlStatementSplitter {
* position of the trailing `;` token whose char range belongs to the
* region) as a complete top-level Spark SQL statement.
*
* The region is extracted from the original source by char-offset
* (`Token.getStartIndex` / `getStopIndex`), `validationPreprocess` is
* applied to it, and the result is re-lexed and parsed with a fresh
* [[SqlBaseParser]]. This isolation means the splitter's parser sees a
* sub-stream whose EOF lands right after the trailing `;`, so the existing
* `compoundOrSingleStatement` rule (which requires `SEMICOLON* EOF`) acts
* as the per-statement validator without any custom grammar rule.
* The region is extracted from the original source by converting ANTLR's
* Unicode code-point token offsets through a one-time UTF-16 map.
* `validationPreprocess` is applied to it, and the result is re-lexed and
* parsed with a fresh [[SqlBaseParser]]. This isolation means the splitter's
* parser sees a sub-stream whose EOF lands right after the trailing `;`, so
* the existing `compoundOrSingleStatement` rule (which requires
* `SEMICOLON* EOF`) acts as the per-statement validator without any custom
* grammar rule.
*
* Uses the same two-stage SLL -> LL prediction strategy as the main parser
* for performance (most statements parse cleanly with the faster SLL stage).
Expand All @@ -326,16 +385,17 @@ object SqlStatementSplitter {
*/
private def tryParseRegion(
sqlText: String,
toUtf16: Array[Int],
stream: CommonTokenStream,
startIdx: Int,
endIdx: Int,
validationPreprocess: String => String,
conf: SqlApiConf): ParseOutcome = {
val firstTok = stream.get(startIdx)
val lastTok = stream.get(endIdx)
val regionStart = firstTok.getStartIndex
val regionStart = toUtf16(firstTok.getStartIndex)
// Token.getStopIndex is inclusive, substring's upper bound is exclusive.
val regionEnd = lastTok.getStopIndex + 1
val regionEnd = toUtf16(lastTok.getStopIndex + 1)
val original = sqlText.substring(regionStart, regionEnd)
val preprocessed = validationPreprocess(original)

Expand Down Expand Up @@ -402,6 +462,45 @@ object SqlStatementSplitter {
parser.setErrorHandler(new BailErrorStrategy)
}

/**
* Maps each Unicode code-point index to a UTF-16 code-unit offset. The last
* entry is `sqlText.length`, so an inclusive ANTLR stop index converts with
* `toUtf16(stopIndex + 1)`.
*/
private def utf16Offsets(sqlText: String): Array[Int] = {
val cuLen = sqlText.length
val offsets = new Array[Int](sqlText.codePointCount(0, cuLen) + 1)
var cu = 0
var cp = 0
while (cu < cuLen) {
offsets(cp) = cu
cu += Character.charCount(sqlText.codePointAt(cu))
cp += 1
}
offsets(cp) = cuLen
offsets
}

// Spark SQL WS token (SqlBaseLexer): ASCII space plus Unicode spaces the
// lexer hides. String.trim only strips characters <= U+0020.
// Hex literals keep the source ASCII (Scala unicode escapes are still non-ASCII).
private def isSqlWhitespace(c: Char): Boolean = c.toInt match {
case 0x20 | 0x09 | 0x0A | 0x0C | 0x0D | 0x0B | 0xA0 | 0x1680 |
0x2000 | 0x2001 | 0x2002 | 0x2003 | 0x2004 | 0x2005 |
0x2006 | 0x2007 | 0x2008 | 0x2009 | 0x200A | 0x2028 |
0x202F | 0x205F | 0x3000 => true
case _ => false
}

/** Trim lexer whitespace; return (leading UTF-16 count, trimmed text). */
private def trimSqlWhitespace(s: String): (Int, String) = {
var start = 0
var end = s.length
while (start < end && isSqlWhitespace(s.charAt(start))) start += 1
while (end > start && isSqlWhitespace(s.charAt(end - 1))) end -= 1
(start, s.substring(start, end))
}

/** Returns the index of the next non-hidden, non-EOF token at or after `from`, or -1. */
private def nextSignificantTokenIndex(stream: CommonTokenStream, from: Int): Int = {
var i = from
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,48 @@ class SqlStatementSplitterSuite extends SparkFunSuite {
assert(result.partialStatement == "select * from")
}

test("source positions skip comments attached to empty statements") {
val sql = " select 1 ; /* select 2 */; select 2"
val result = SqlStatementSplitter.splitWithPositions(sql, identity)
val complete = result.completeStatements.head
assert(complete.statement == "select 1")
assert(complete.start == 2)
assert(complete.length == 8)

val partial = result.partialStatement.get
assert(partial.statement == "select 2")
assert(partial.start == sql.lastIndexOf("select 2"))
assert(partial.length == 8)
}

test("source positions use UTF-16 offsets for supplementary characters") {
val emoji = new String(Character.toChars(0x1F600))
val sql = s"SELECT '$emoji$emoji'; SELECT 2;"
val result = SqlStatementSplitter.splitWithPositions(sql, identity)

assert(result.completeStatements.map(_.statement) ==
Seq(s"SELECT '$emoji$emoji'", "SELECT 2"))
assert(result.completeStatements.map(_.start) == Seq(0, 15))
assert(result.completeStatements.map(_.length) == Seq(13, 8))
result.completeStatements.foreach { statement =>
assert(sql.substring(statement.start, statement.start + statement.length) ==
statement.statement)
}
assert(result.partialStatement.isEmpty)
}

test("source positions trim Spark SQL Unicode whitespace") {
val nbsp = 0xA0.toChar
val sql = s"${nbsp}SELECT 1$nbsp;"
val result = SqlStatementSplitter.splitWithPositions(sql, identity)
val complete = result.completeStatements.head
assert(complete.statement == "SELECT 1")
assert(complete.start == 1)
assert(complete.length == 8)
assert(sql.substring(complete.start, complete.start + complete.length) ==
complete.statement)
}

// ----------------------------------------------------------------------------------
// Error tolerance (mirrors Trino behavior)
// ----------------------------------------------------------------------------------
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,38 +27,46 @@ import org.apache.spark.sql.types.{AbstractDataType, DataType, StringType}
import org.apache.spark.unsafe.types.UTF8String

/**
* Parses a SQL statement string and returns a compact JSON description of the
* unresolved statement (identifier/code, lineage references, select-list names,
* parameters), or a STANDARD-format error object when the statement does not
* parse.
* Parses a SQL batch string and returns a compact JSON array describing its
* unresolved statements (source position, identifier/code, lineage references,
* select-list names, parameters). A statement that does not parse is represented
* by a STANDARD-format error object at its position in the array.
*
* Behind [[SQLConf.PARSE_SQL_ENABLED]] while the JSON contract is still
* evolving. Designed for batch evaluation over DataFrames of SQL text.
* User-facing parse errors become JSON; unexpected internal failures propagate.
*/
// scalastyle:off line.size.limit
// scalastyle:off nonascii
@ExpressionDescription(
usage = """_FUNC_(sqlStmt) - Parses `sqlStmt` with the stock Spark SQL parser and
returns a JSON string describing the statement (parse success, Table 39 statement
identifier/code, target and source table references for lineage, select-list column
names, and parameter markers). Session parser extensions are not applied.
usage = """_FUNC_(sqlStmt) - Splits `sqlStmt` into SQL statements, parses each with
the stock Spark SQL parser, and returns a JSON array describing them (1-based
UTF-16 code-unit `start` and UTF-16 code-unit `length`, parse success, Table 39
statement identifier/code, target and source table references for lineage,
select-list column names, and parameter markers).
Statement length excludes surrounding whitespace and the terminating semicolon.
Session parser extensions are not applied.
Requires spark.sql.function.parseSql.enabled=true. On syntax / parse error returns JSON
with `parse_success` false, source location, and a nested STANDARD error object
instead of throwing.""",
for that statement with `parse_success` false, source location, and a nested STANDARD
error object instead of throwing or stopping the remaining statements. Nested error
locations are statement-relative. An empty or closed-comment-only batch returns `[]`.""",
arguments = """
Arguments:
* sqlStmt - A SQL statement string to parse.
* sqlStmt - A SQL batch string to split and parse.
An expression that evaluates to a string.
""",
examples = """
Examples:
> SELECT _FUNC_('SELECT a, b FROM t');
{"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["t"]],"select_list":[{"name":["a"]},{"name":["b"]}]}
> SELECT get_json_object(_FUNC_('SELEC'), '$.error.errorClass');
> SELECT _FUNC_('SELECT 1;SELECT 2');
[{"start":1,"length":8,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"select_list":[{"name":[]}]},{"start":10,"length":8,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"select_list":[{"name":[]}]}]
> SELECT get_json_object(_FUNC_('SELECT ''😀'';SELECT 2'), '$[1].start');
13
> SELECT get_json_object(_FUNC_('SELEC'), '$[0].error.errorClass');
PARSE_SYNTAX_ERROR
""",
group = "misc_funcs",
since = "4.4.0")
// scalastyle:on nonascii
// scalastyle:on line.size.limit
case class ParseSql(child: Expression)
extends UnaryExpression
Expand Down
Loading