From 6fca0e7bf55850ece57aefad026af77e5efda89b Mon Sep 17 00:00:00 2001 From: srielau Date: Fri, 4 Sep 2026 17:59:47 +0000 Subject: [PATCH 01/12] [SPARK-59255][SQL] Extend parse_sql to parse SQL batches --- .../parser/SqlStatementSplitter.scala | 95 ++++++++--- .../parser/SqlStatementSplitterSuite.scala | 14 ++ .../sql/catalyst/expressions/ParseSql.scala | 31 ++-- .../sql/catalyst/parser/ParseSqlResult.scala | 50 ++++-- .../spark/sql/execution/SparkSqlParser.scala | 7 + .../sql-functions/sql-expression-schema.md | 2 +- .../analyzer-results/parse-sql.sql.out | 69 +++++--- .../resources/sql-tests/inputs/parse-sql.sql | 55 +++--- .../sql-tests/results/parse-sql.sql.out | 159 ++++++++++-------- .../catalyst/expressions/ParseSqlSuite.scala | 27 ++- .../catalyst/parser/ParseSqlResultSuite.scala | 57 ++++++- 11 files changed, 398 insertions(+), 168 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala index 16112402350ef..83ef5ab72e6d9 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala @@ -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 +} + +/** 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`. @@ -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 @@ -121,7 +141,18 @@ 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") @@ -142,8 +173,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. @@ -158,6 +190,32 @@ object SqlStatementSplitter { // interpretation (e.g. `double_quoted_identifiers`). val conf = SqlApiConf.get + def appendToken(token: Token): Unit = { + if (buffer.isEmpty) bufferStart = 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 statement = raw.trim + if (statement.isEmpty) { + None + } else { + val leadingWhitespace = raw.indexOf(statement) + assert(bufferStart >= 0 && leadingWhitespace >= 0) + Some(PositionedSqlStatement( + statement, + terminator, + bufferStart + leadingWhitespace)) + } + } + while (!stopOuter && index < numTokens) { val startIdx = nextSignificantTokenIndex(tokenStream, index) if (startIdx < 0) { @@ -167,7 +225,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) { @@ -216,17 +274,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) { @@ -253,17 +307,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 += _) } - buffer.setLength(0) - bufferHasContent = false + resetBuffer() stopInner = true } else { if (token.getChannel != Token.HIDDEN_CHANNEL) bufferHasContent = true - buffer.append(token.getText) + appendToken(token) } } } else { @@ -275,7 +325,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 @@ -285,8 +335,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. */ diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala index 2e091aeff974f..eec2db0d53071 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala @@ -71,6 +71,20 @@ 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) + } + // ---------------------------------------------------------------------------------- // Error tolerance (mirrors Trino behavior) // ---------------------------------------------------------------------------------- diff --git a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala index 649feeee9adf7..ed406385c8822 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala @@ -27,10 +27,10 @@ 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. @@ -38,23 +38,26 @@ import org.apache.spark.unsafe.types.UTF8String */ // scalastyle:off line.size.limit @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 start + and 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 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_('SELEC'), '$[0].error.errorClass'); PARSE_SYNTAX_ERROR """, group = "misc_funcs", diff --git a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResult.scala b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResult.scala index fbc4b96f9dd71..8f86ceb0900ed 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResult.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResult.scala @@ -33,8 +33,8 @@ import org.apache.spark.sql.execution.command.{CreateViewCommand, DescribeQueryC import org.apache.spark.sql.execution.datasources.CreateTempViewUsing /** - * Parses a SQL statement string and returns a compact JSON description of the - * unresolved plan (parse-only; no catalog resolution). + * Parses a SQL batch string and returns a compact JSON array describing its + * unresolved statements (parse-only; no catalog resolution). * * Uses a stock [[SparkSqlParser]] (ThreadLocal) so statement coverage matches * the default production parser (EXPLAIN / SET / ADD JAR / temp views / etc.). @@ -43,12 +43,17 @@ import org.apache.spark.sql.execution.datasources.CreateTempViewUsing * executors without a session, so only the stock parser is available under * distributed eval. * - * On success the JSON always includes `parse_success`, the statement + * Every statement object includes its 1-based `start` in the original batch + * and its `length`, excluding surrounding whitespace and the terminating + * semicolon. On success it also includes `parse_success`, the statement * identifier/code (ISO/IEC 9075-2:2023 Table 39), and omits unused optional * fields (`target_table_references`, `source_table_references`, - * `function_references`, `select_list`, `parameter_markers`) when empty. On parse - * failure it returns `parse_success: false` with source location and a nested - * STANDARD-format error object, and does not throw. Only [[ParseException]] / + * `function_references`, `select_list`, `parameter_markers`) when empty. On + * parse failure the statement object contains `parse_success: false` with + * source location and a nested STANDARD-format error object, and parsing + * continues with later statements. Nested error locations are relative to the + * individual statement, while `start` is relative to the original batch. An + * empty or comment-only batch produces an empty array. Only [[ParseException]] / * [[SqlScriptingException]] are converted to JSON; unexpected / internal * failures propagate so the function fails. */ @@ -57,11 +62,18 @@ object ParseSqlResult { private val parser: ThreadLocal[SparkSqlParser] = ThreadLocal.withInitial(() => new SparkSqlParser()) - /** Parse `sql` and render the JSON result string. */ + /** Parse `sql` as a batch and render the JSON result array. */ def fromSql(sql: String): String = { + val split = parser.get().splitStatementsWithPositions(sql) + val statements = split.completeStatements ++ split.partialStatement + compact(render(JArray(statements.map(parseStatement).toList))) + } + + private def parseStatement(segment: PositionedSqlStatement): JObject = { + val sql = segment.statement try { // Do not inherit the outer query's origin from the parse_sql expression. - // Errors and parsed nodes must refer to the SQL string passed to this function. + // Errors and parsed nodes refer to the individual statement. val origin = if (sql.nonEmpty) { Origin(startIndex = Some(0), stopIndex = Some(sql.length - 1), sqlText = Some(sql)) } else { @@ -69,21 +81,23 @@ object ParseSqlResult { } CurrentOrigin.withOrigin(origin) { val plan = parser.get().parsePlan(sql) - fromPlan(plan) + successJson(plan, segment) } } catch { // User-facing parse / scripting failures become JSON; everything else fails. case e: ParseException => - errorJson(e) + errorJson(e, segment) case e: SqlScriptingException => - errorJson(e) + errorJson(e, segment) } } /** Build success JSON from an already-parsed unresolved plan. */ - def fromPlan(plan: LogicalPlan): String = { + private def successJson(plan: LogicalPlan, segment: PositionedSqlStatement): JObject = { val classification = SqlStatementCodes.classify(plan) val fields = mutable.ListBuffer.empty[JField] + fields += "start" -> JInt(segment.start + 1) + fields += "length" -> JInt(segment.length) fields += "parse_success" -> JBool(true) fields += "statement_identifier" -> JString(classification.statementIdentifier) fields += "statement_code" -> JInt(classification.statementCode) @@ -105,10 +119,12 @@ object ParseSqlResult { fields += "select_list" -> JArray(selectList.toList) } refs.parameterMarkers.foreach(markers => fields += "parameter_markers" -> markers) - compact(render(JObject(fields.toList))) + JObject(fields.toList) } - private def errorJson(e: SparkThrowable with Throwable): String = { + private def errorJson( + e: SparkThrowable with Throwable, + segment: PositionedSqlStatement): JObject = { val errorObj = parseJson( SparkThrowableHelper.getMessage(e, ErrorMessageFormat.STANDARD)).asInstanceOf[JObject] val origin = e match { @@ -122,10 +138,12 @@ object ParseSqlResult { } else { origin.toSeq.flatMap(queryContextField) } - compact(render(JObject( + JObject( + "start" -> JInt(segment.start + 1), + "length" -> JInt(segment.length), "parse_success" -> JBool(false), "error" -> JObject(errorObj.obj ++ contextFields ++ locationFields) - ))) + ) } private def queryContextField(origin: Origin): Option[JField] = origin.context match { diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala index 5070a40259e37..bce4250c18975 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala @@ -123,6 +123,13 @@ class SparkSqlParser extends AbstractSqlParser { override def splitStatements(sqlText: String): SqlStatementSplitResult = SqlStatementSplitter.split(sqlText, SparkSqlParser.substituteVariablesForValidation) + /** Split statements while retaining their positions in the original SQL text. */ + private[sql] def splitStatementsWithPositions( + sqlText: String): PositionedSqlStatementSplitResult = + SqlStatementSplitter.splitWithPositions( + sqlText, + SparkSqlParser.substituteVariablesForValidation) + /** * Internal parse method that handles both parameter substitution and regular parsing. * diff --git a/sql/core/src/test/resources/sql-functions/sql-expression-schema.md b/sql/core/src/test/resources/sql-functions/sql-expression-schema.md index 6f11ca0435ae9..b80684843fa60 100644 --- a/sql/core/src/test/resources/sql-functions/sql-expression-schema.md +++ b/sql/core/src/test/resources/sql-functions/sql-expression-schema.md @@ -279,7 +279,7 @@ | org.apache.spark.sql.catalyst.expressions.OctetLength | octet_length | SELECT octet_length('Spark SQL') | struct | | org.apache.spark.sql.catalyst.expressions.Or | or | SELECT true or false | struct<(true OR false):boolean> | | org.apache.spark.sql.catalyst.expressions.Overlay | overlay | SELECT overlay('Spark SQL' PLACING '_' FROM 6) | struct | -| org.apache.spark.sql.catalyst.expressions.ParseSql | parse_sql | SELECT parse_sql('SELECT a, b FROM t') | struct | +| org.apache.spark.sql.catalyst.expressions.ParseSql | parse_sql | SELECT parse_sql('SELECT 1;SELECT 2') | struct | | org.apache.spark.sql.catalyst.expressions.ParseToDate | to_date | SELECT to_date('2009-07-30 04:17:52') | struct | | org.apache.spark.sql.catalyst.expressions.ParseToTimestamp | to_timestamp | SELECT to_timestamp('2016-12-31 00:12:00') | struct | | org.apache.spark.sql.catalyst.expressions.ParseToTimestampLTZExpressionBuilder | to_timestamp_ltz | SELECT to_timestamp_ltz('2016-12-31 00:12:00') | struct | diff --git a/sql/core/src/test/resources/sql-tests/analyzer-results/parse-sql.sql.out b/sql/core/src/test/resources/sql-tests/analyzer-results/parse-sql.sql.out index bc9bbba7c3c71..30d0d993176f9 100644 --- a/sql/core/src/test/resources/sql-tests/analyzer-results/parse-sql.sql.out +++ b/sql/core/src/test/resources/sql-tests/analyzer-results/parse-sql.sql.out @@ -20,14 +20,35 @@ Project [parse_sql(SELECT db.my_func(a), count(b) FROM cat.ns.t1 JOIN t2) AS par +- OneRowRelation +-- !query +SELECT parse_sql('select 1; select 2') +-- !query analysis +Project [parse_sql(select 1; select 2) AS parse_sql(select 1; select 2)#x] ++- OneRowRelation + + +-- !query +SELECT parse_sql('select 1; select 2;') +-- !query analysis +Project [parse_sql(select 1; select 2;) AS parse_sql(select 1; select 2;)#x] ++- OneRowRelation + + +-- !query +SELECT parse_sql('SELECT 1; SELEC 2; SELECT 3') +-- !query analysis +Project [parse_sql(SELECT 1; SELEC 2; SELECT 3) AS parse_sql(SELECT 1; SELEC 2; SELECT 3)#x] ++- OneRowRelation + + -- !query SELECT - get_json_object(result, '$.statement_identifier') AS statement_identifier, - get_json_object(result, '$.source_table_references[0][0]') AS first_table, - get_json_object(result, '$.select_list[1].name[0]') AS second_column + get_json_object(result, '$[0].statement_identifier') AS statement_identifier, + get_json_object(result, '$[0].source_table_references[0][0]') AS first_table, + get_json_object(result, '$[0].select_list[1].name[0]') AS second_column FROM (SELECT parse_sql('SELECT a, b FROM t') AS result) -- !query analysis -Project [get_json_object(result#x, $.statement_identifier) AS statement_identifier#x, get_json_object(result#x, $.source_table_references[0][0]) AS first_table#x, get_json_object(result#x, $.select_list[1].name[0]) AS second_column#x] +Project [get_json_object(result#x, $[0].statement_identifier) AS statement_identifier#x, get_json_object(result#x, $[0].source_table_references[0][0]) AS first_table#x, get_json_object(result#x, $[0].select_list[1].name[0]) AS second_column#x] +- SubqueryAlias __auto_generated_subquery_name +- Project [parse_sql(SELECT a, b FROM t) AS result#x] +- OneRowRelation @@ -251,12 +272,12 @@ Project [parse_sql(SELEC FROM t) AS parse_sql(SELEC FROM t)#x] -- !query SELECT - get_json_object(result, '$.parse_success') AS parse_success, - get_json_object(result, '$.error.errorClass') AS error_class, - get_json_object(result, '$.error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[0].parse_success') AS parse_success, + get_json_object(result, '$[0].error.errorClass') AS error_class, + get_json_object(result, '$[0].error.queryContext[0].fragment') AS fragment FROM (SELECT parse_sql('SELEC FROM t') AS result) -- !query analysis -Project [get_json_object(result#x, $.parse_success) AS parse_success#x, get_json_object(result#x, $.error.errorClass) AS error_class#x, get_json_object(result#x, $.error.queryContext[0].fragment) AS fragment#x] +Project [get_json_object(result#x, $[0].parse_success) AS parse_success#x, get_json_object(result#x, $[0].error.errorClass) AS error_class#x, get_json_object(result#x, $[0].error.queryContext[0].fragment) AS fragment#x] +- SubqueryAlias __auto_generated_subquery_name +- Project [parse_sql(SELEC FROM t) AS result#x] +- OneRowRelation @@ -281,10 +302,10 @@ Project [parse_sql(SELECT * -- !query SELECT - get_json_object(result, '$.error.errorClass') AS error_class, - get_json_object(result, '$.error.line') AS line, - get_json_object(result, '$.error.position') AS position, - get_json_object(result, '$.error.queryContext[0].startIndex') AS start_index + get_json_object(result, '$[0].error.errorClass') AS error_class, + get_json_object(result, '$[0].error.line') AS line, + get_json_object(result, '$[0].error.position') AS position, + get_json_object(result, '$[0].error.queryContext[0].startIndex') AS start_index FROM ( SELECT parse_sql( 'SELECT * @@ -293,7 +314,7 @@ FROM ( CLUSTER BY b') AS result ) -- !query analysis -Project [get_json_object(result#x, $.error.errorClass) AS error_class#x, get_json_object(result#x, $.error.line) AS line#x, get_json_object(result#x, $.error.position) AS position#x, get_json_object(result#x, $.error.queryContext[0].startIndex) AS start_index#x] +Project [get_json_object(result#x, $[0].error.errorClass) AS error_class#x, get_json_object(result#x, $[0].error.line) AS line#x, get_json_object(result#x, $[0].error.position) AS position#x, get_json_object(result#x, $[0].error.queryContext[0].startIndex) AS start_index#x] +- SubqueryAlias __auto_generated_subquery_name +- Project [parse_sql(SELECT * FROM t @@ -391,10 +412,12 @@ Project [parse_sql(BEGIN -- !query SELECT - get_json_object(result, '$.error.errorClass') AS error_class, - get_json_object(result, '$.error.line') AS line, - get_json_object(result, '$.error.position') AS position, - get_json_object(result, '$.error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[1].start') AS start, + get_json_object(result, '$[1].length') AS length, + get_json_object(result, '$[1].error.errorClass') AS error_class, + get_json_object(result, '$[1].error.line') AS line, + get_json_object(result, '$[1].error.position') AS position, + get_json_object(result, '$[1].error.queryContext[0].fragment') AS fragment FROM ( SELECT parse_sql( 'BEGIN @@ -403,7 +426,7 @@ FROM ( END') AS result ) -- !query analysis -Project [get_json_object(result#x, $.error.errorClass) AS error_class#x, get_json_object(result#x, $.error.line) AS line#x, get_json_object(result#x, $.error.position) AS position#x, get_json_object(result#x, $.error.queryContext[0].fragment) AS fragment#x] +Project [get_json_object(result#x, $[1].start) AS start#x, get_json_object(result#x, $[1].length) AS length#x, get_json_object(result#x, $[1].error.errorClass) AS error_class#x, get_json_object(result#x, $[1].error.line) AS line#x, get_json_object(result#x, $[1].error.position) AS position#x, get_json_object(result#x, $[1].error.queryContext[0].fragment) AS fragment#x] +- SubqueryAlias __auto_generated_subquery_name +- Project [parse_sql(BEGIN SELECT 1; @@ -434,10 +457,10 @@ Project [parse_sql(BEGIN -- !query SELECT - get_json_object(result, '$.error.errorClass') AS error_class, - get_json_object(result, '$.error.line') AS line, - get_json_object(result, '$.error.position') AS position, - get_json_object(result, '$.error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[0].error.errorClass') AS error_class, + get_json_object(result, '$[0].error.line') AS line, + get_json_object(result, '$[0].error.position') AS position, + get_json_object(result, '$[0].error.queryContext[0].fragment') AS fragment FROM ( SELECT parse_sql( 'BEGIN @@ -447,7 +470,7 @@ FROM ( END') AS result ) -- !query analysis -Project [get_json_object(result#x, $.error.errorClass) AS error_class#x, get_json_object(result#x, $.error.line) AS line#x, get_json_object(result#x, $.error.position) AS position#x, get_json_object(result#x, $.error.queryContext[0].fragment) AS fragment#x] +Project [get_json_object(result#x, $[0].error.errorClass) AS error_class#x, get_json_object(result#x, $[0].error.line) AS line#x, get_json_object(result#x, $[0].error.position) AS position#x, get_json_object(result#x, $[0].error.queryContext[0].fragment) AS fragment#x] +- SubqueryAlias __auto_generated_subquery_name +- Project [parse_sql(BEGIN lbl_begin: BEGIN diff --git a/sql/core/src/test/resources/sql-tests/inputs/parse-sql.sql b/sql/core/src/test/resources/sql-tests/inputs/parse-sql.sql index b4eb15ffbc38e..6b73c8722c6bb 100644 --- a/sql/core/src/test/resources/sql-tests/inputs/parse-sql.sql +++ b/sql/core/src/test/resources/sql-tests/inputs/parse-sql.sql @@ -1,5 +1,5 @@ --- End-to-end coverage for parse_sql (SPARK-58738). --- Returns compact JSON for parse-only statement analysis via SparkSqlParser. +-- End-to-end coverage for parse_sql (SPARK-59255). +-- Returns a compact JSON array for parse-only SQL batch analysis via SparkSqlParser. -- Off by default while the JSON contract is still evolving. --SET spark.sql.function.parseSql.enabled=true @@ -10,11 +10,24 @@ SELECT parse_sql(NULL); SELECT parse_sql('SELECT a, b FROM t'); SELECT parse_sql('SELECT db.my_func(a), count(b) FROM cat.ns.t1 JOIN t2'); +-- batches with and without a final delimiter +--QUERY-DELIMITER-START +SELECT parse_sql('select 1; select 2'); +--QUERY-DELIMITER-END +--QUERY-DELIMITER-START +SELECT parse_sql('select 1; select 2;'); +--QUERY-DELIMITER-END + +-- a failing statement does not prevent later statements from being parsed +--QUERY-DELIMITER-START +SELECT parse_sql('SELECT 1; SELEC 2; SELECT 3'); +--QUERY-DELIMITER-END + -- JSON-path access over one shared successful parse result SELECT - get_json_object(result, '$.statement_identifier') AS statement_identifier, - get_json_object(result, '$.source_table_references[0][0]') AS first_table, - get_json_object(result, '$.select_list[1].name[0]') AS second_column + get_json_object(result, '$[0].statement_identifier') AS statement_identifier, + get_json_object(result, '$[0].source_table_references[0][0]') AS first_table, + get_json_object(result, '$[0].select_list[1].name[0]') AS second_column FROM (SELECT parse_sql('SELECT a, b FROM t') AS result); -- DML @@ -93,9 +106,9 @@ SELECT parse_sql('SELEC FROM t'); -- JSON-path access over one shared parse result SELECT - get_json_object(result, '$.parse_success') AS parse_success, - get_json_object(result, '$.error.errorClass') AS error_class, - get_json_object(result, '$.error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[0].parse_success') AS parse_success, + get_json_object(result, '$[0].error.errorClass') AS error_class, + get_json_object(result, '$[0].error.queryContext[0].fragment') AS fragment FROM (SELECT parse_sql('SELEC FROM t') AS result); -- full multiline parse-time validation error, including context and location @@ -107,10 +120,10 @@ SELECT parse_sql( -- JSON-path access over one shared multiline parse result SELECT - get_json_object(result, '$.error.errorClass') AS error_class, - get_json_object(result, '$.error.line') AS line, - get_json_object(result, '$.error.position') AS position, - get_json_object(result, '$.error.queryContext[0].startIndex') AS start_index + get_json_object(result, '$[0].error.errorClass') AS error_class, + get_json_object(result, '$[0].error.line') AS line, + get_json_object(result, '$[0].error.position') AS position, + get_json_object(result, '$[0].error.queryContext[0].startIndex') AS start_index FROM ( SELECT parse_sql( 'SELECT * @@ -143,10 +156,12 @@ SELECT parse_sql( -- JSON-path access over one shared scripting parse result --QUERY-DELIMITER-START SELECT - get_json_object(result, '$.error.errorClass') AS error_class, - get_json_object(result, '$.error.line') AS line, - get_json_object(result, '$.error.position') AS position, - get_json_object(result, '$.error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[1].start') AS start, + get_json_object(result, '$[1].length') AS length, + get_json_object(result, '$[1].error.errorClass') AS error_class, + get_json_object(result, '$[1].error.line') AS line, + get_json_object(result, '$[1].error.position') AS position, + get_json_object(result, '$[1].error.queryContext[0].fragment') AS fragment FROM ( SELECT parse_sql( 'BEGIN @@ -169,10 +184,10 @@ SELECT parse_sql( -- JSON-path access over one shared scripting validation result --QUERY-DELIMITER-START SELECT - get_json_object(result, '$.error.errorClass') AS error_class, - get_json_object(result, '$.error.line') AS line, - get_json_object(result, '$.error.position') AS position, - get_json_object(result, '$.error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[0].error.errorClass') AS error_class, + get_json_object(result, '$[0].error.line') AS line, + get_json_object(result, '$[0].error.position') AS position, + get_json_object(result, '$[0].error.queryContext[0].fragment') AS fragment FROM ( SELECT parse_sql( 'BEGIN diff --git a/sql/core/src/test/resources/sql-tests/results/parse-sql.sql.out b/sql/core/src/test/resources/sql-tests/results/parse-sql.sql.out index 986ce7c4ae945..1310f7c30fffa 100644 --- a/sql/core/src/test/resources/sql-tests/results/parse-sql.sql.out +++ b/sql/core/src/test/resources/sql-tests/results/parse-sql.sql.out @@ -12,7 +12,7 @@ SELECT parse_sql('SELECT a, b FROM t') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["t"]],"select_list":[{"name":["a"]},{"name":["b"]}]} +[{"start":1,"length":18,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["t"]],"select_list":[{"name":["a"]},{"name":["b"]}]}] -- !query @@ -20,14 +20,38 @@ SELECT parse_sql('SELECT db.my_func(a), count(b) FROM cat.ns.t1 JOIN t2') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["cat","ns","t1"],["t2"]],"function_references":[["db","my_func"],["count"]],"select_list":[{"name":[]},{"name":[]}]} +[{"start":1,"length":53,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["cat","ns","t1"],["t2"]],"function_references":[["db","my_func"],["count"]],"select_list":[{"name":[]},{"name":[]}]}] + + +-- !query +SELECT parse_sql('select 1; select 2') +-- !query schema +struct +-- !query output +[{"start":1,"length":8,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"select_list":[{"name":[]}]},{"start":11,"length":8,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"select_list":[{"name":[]}]}] + + +-- !query +SELECT parse_sql('select 1; select 2;') +-- !query schema +struct +-- !query output +[{"start":1,"length":8,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"select_list":[{"name":[]}]},{"start":11,"length":8,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"select_list":[{"name":[]}]}] + + +-- !query +SELECT parse_sql('SELECT 1; SELEC 2; SELECT 3') +-- !query schema +struct +-- !query output +[{"start":1,"length":8,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"select_list":[{"name":[]}]},{"start":11,"length":7,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"'SELEC'","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":7,"fragment":"SELEC 2"}],"line":1,"position":0}},{"start":20,"length":8,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"select_list":[{"name":[]}]}] -- !query SELECT - get_json_object(result, '$.statement_identifier') AS statement_identifier, - get_json_object(result, '$.source_table_references[0][0]') AS first_table, - get_json_object(result, '$.select_list[1].name[0]') AS second_column + get_json_object(result, '$[0].statement_identifier') AS statement_identifier, + get_json_object(result, '$[0].source_table_references[0][0]') AS first_table, + get_json_object(result, '$[0].select_list[1].name[0]') AS second_column FROM (SELECT parse_sql('SELECT a, b FROM t') AS result) -- !query schema struct @@ -40,7 +64,7 @@ SELECT parse_sql('INSERT INTO t SELECT 1') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"INSERT","statement_code":50,"target_table_references":[["t"]],"select_list":[{"name":[]}]} +[{"start":1,"length":22,"parse_success":true,"statement_identifier":"INSERT","statement_code":50,"target_table_references":[["t"]],"select_list":[{"name":[]}]}] -- !query @@ -48,7 +72,7 @@ SELECT parse_sql('DELETE FROM t WHERE a = 1') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"DELETE WHERE","statement_code":19,"target_table_references":[["t"]]} +[{"start":1,"length":25,"parse_success":true,"statement_identifier":"DELETE WHERE","statement_code":19,"target_table_references":[["t"]]}] -- !query @@ -56,7 +80,7 @@ SELECT parse_sql('UPDATE t SET a = 1 WHERE b = 2') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"UPDATE WHERE","statement_code":82,"target_table_references":[["t"]]} +[{"start":1,"length":30,"parse_success":true,"statement_identifier":"UPDATE WHERE","statement_code":82,"target_table_references":[["t"]]}] -- !query @@ -64,7 +88,7 @@ SELECT parse_sql('MERGE INTO t USING s ON t.id = s.id WHEN MATCHED THEN DELETE') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"MERGE","statement_code":128,"target_table_references":[["t"]],"source_table_references":[["s"]]} +[{"start":1,"length":60,"parse_success":true,"statement_identifier":"MERGE","statement_code":128,"target_table_references":[["t"]],"source_table_references":[["s"]]}] -- !query @@ -72,7 +96,7 @@ SELECT parse_sql('CREATE TABLE t (a INT)') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"CREATE TABLE","statement_code":77,"target_table_references":[["t"]]} +[{"start":1,"length":22,"parse_success":true,"statement_identifier":"CREATE TABLE","statement_code":77,"target_table_references":[["t"]]}] -- !query @@ -80,7 +104,7 @@ SELECT parse_sql('CREATE TABLE t AS SELECT 1 AS a') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"CREATE TABLE","statement_code":77,"target_table_references":[["t"]],"select_list":[{"name":["a"]}]} +[{"start":1,"length":31,"parse_success":true,"statement_identifier":"CREATE TABLE","statement_code":77,"target_table_references":[["t"]],"select_list":[{"name":["a"]}]}] -- !query @@ -88,7 +112,7 @@ SELECT parse_sql('DROP TABLE t') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"DROP TABLE","statement_code":32,"target_table_references":[["t"]]} +[{"start":1,"length":12,"parse_success":true,"statement_identifier":"DROP TABLE","statement_code":32,"target_table_references":[["t"]]}] -- !query @@ -96,7 +120,7 @@ SELECT parse_sql('CACHE TABLE t') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"CACHE TABLE","statement_code":-1,"target_table_references":[["t"]]} +[{"start":1,"length":13,"parse_success":true,"statement_identifier":"CACHE TABLE","statement_code":-1,"target_table_references":[["t"]]}] -- !query @@ -104,7 +128,7 @@ SELECT parse_sql('TABLE t') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["t"]]} +[{"start":1,"length":7,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["t"]]}] -- !query @@ -112,7 +136,7 @@ SELECT parse_sql('VALUES (1), (2)') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"SELECT","statement_code":21} +[{"start":1,"length":15,"parse_success":true,"statement_identifier":"SELECT","statement_code":21}] -- !query @@ -120,7 +144,7 @@ SELECT parse_sql('CREATE FUNCTION f AS ''x'' USING JAR ''y.jar''') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"CREATE ROUTINE","statement_code":14} +[{"start":1,"length":42,"parse_success":true,"statement_identifier":"CREATE ROUTINE","statement_code":14}] -- !query @@ -128,7 +152,7 @@ SELECT parse_sql('DECLARE VARIABLE x INT') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"DECLARE VARIABLE","statement_code":-8} +[{"start":1,"length":22,"parse_success":true,"statement_identifier":"DECLARE VARIABLE","statement_code":-8}] -- !query @@ -136,7 +160,7 @@ SELECT parse_sql('SELECT * FROM t WHERE a = :foo AND b = ?') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["t"]],"select_list":[{"name":["*"]}],"parameter_markers":{"named":["foo"],"unnamed_count":1}} +[{"start":1,"length":40,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["t"]],"select_list":[{"name":["*"]}],"parameter_markers":{"named":["foo"],"unnamed_count":1}}] -- !query @@ -144,7 +168,7 @@ SELECT parse_sql('WITH cte AS (SELECT a FROM hidden_base) SELECT a FROM cte') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["hidden_base"]],"select_list":[{"name":["a"]}]} +[{"start":1,"length":57,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["hidden_base"]],"select_list":[{"name":["a"]}]}] -- !query @@ -152,7 +176,7 @@ SELECT parse_sql('SELECT * FROM real_t WHERE EXISTS (WITH real_t AS (SELECT * FR -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["inner_base"],["real_t"]],"select_list":[{"name":["*"]}]} +[{"start":1,"length":98,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["inner_base"],["real_t"]],"select_list":[{"name":["*"]}]}] -- !query @@ -160,7 +184,7 @@ SELECT parse_sql('WITH a AS (SELECT * FROM b), b AS (SELECT 1 AS x) SELECT * FRO -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["b"]],"select_list":[{"name":["*"]}]} +[{"start":1,"length":65,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["b"]],"select_list":[{"name":["*"]}]}] -- !query @@ -168,7 +192,7 @@ SELECT parse_sql('SELECT (SELECT max(v) FROM scalar_src) AS m, t.a FROM outer_t -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["scalar_src"],["exists_src"],["outer_t"]],"function_references":[["max"]],"select_list":[{"name":["m"]},{"name":["t","a"]}]} +[{"start":1,"length":123,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["scalar_src"],["exists_src"],["outer_t"]],"function_references":[["max"]],"select_list":[{"name":["m"]},{"name":["t","a"]}]}] -- !query @@ -195,7 +219,7 @@ struct 0) > 0 ORDER BY greatest(t.a, 1)):string> -- !query output -{"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["scalar_t"],["left_t"],["right_t"]],"function_references":[["greatest"],["count_if"],["coalesce"],["sum"],["abs"],["lower"],["length"],["startswith"],["max"],["range"],["hash"]],"select_list":[{"name":[]},{"name":[]}]} +[{"start":1,"length":390,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"source_table_references":[["scalar_t"],["left_t"],["right_t"]],"function_references":[["greatest"],["count_if"],["coalesce"],["sum"],["abs"],["lower"],["length"],["startswith"],["max"],["range"],["hash"]],"select_list":[{"name":[]},{"name":[]}]}] -- !query @@ -224,7 +248,7 @@ struct -- !query output -{"parse_success":true,"statement_identifier":"MERGE","statement_code":128,"target_table_references":[["target"]],"source_table_references":[["source"]],"function_references":[["hash"],["should_update"],["coalesce"],["upper"],["lower"],["normalize_name"],["is_valid"]]} +[{"start":1,"length":320,"parse_success":true,"statement_identifier":"MERGE","statement_code":128,"target_table_references":[["target"]],"source_table_references":[["source"]],"function_references":[["hash"],["should_update"],["coalesce"],["upper"],["lower"],["normalize_name"],["is_valid"]]}] -- !query @@ -239,7 +263,7 @@ struct -- !query output -{"parse_success":true,"statement_identifier":"CREATE TABLE","statement_code":77,"target_table_references":[["defaults"]],"function_references":[["current_date"],["upper"]]} +[{"start":1,"length":106,"parse_success":true,"statement_identifier":"CREATE TABLE","statement_code":77,"target_table_references":[["defaults"]],"function_references":[["current_date"],["upper"]]}] -- !query @@ -247,14 +271,14 @@ SELECT parse_sql('SELEC FROM t') -- !query schema struct -- !query output -{"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"'SELEC'","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":12,"fragment":"SELEC FROM t"}],"line":1,"position":0}} +[{"start":1,"length":12,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"'SELEC'","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":12,"fragment":"SELEC FROM t"}],"line":1,"position":0}}] -- !query SELECT - get_json_object(result, '$.parse_success') AS parse_success, - get_json_object(result, '$.error.errorClass') AS error_class, - get_json_object(result, '$.error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[0].parse_success') AS parse_success, + get_json_object(result, '$[0].error.errorClass') AS error_class, + get_json_object(result, '$[0].error.queryContext[0].fragment') AS fragment FROM (SELECT parse_sql('SELEC FROM t') AS result) -- !query schema struct @@ -274,15 +298,15 @@ struct -- !query output -{"parse_success":false,"error":{"errorClass":"UNSUPPORTED_FEATURE.COMBINATION_QUERY_RESULT_CLAUSES","messageTemplate":"The feature is not supported: Combination of ORDER BY/SORT BY/DISTRIBUTE BY/CLUSTER BY.","sqlState":"0A000","queryContext":[{"objectType":"","objectName":"","startIndex":19,"stopIndex":42,"fragment":"ORDER BY a\n CLUSTER BY b"}],"line":3,"position":1}} +[{"start":1,"length":42,"parse_success":false,"error":{"errorClass":"UNSUPPORTED_FEATURE.COMBINATION_QUERY_RESULT_CLAUSES","messageTemplate":"The feature is not supported: Combination of ORDER BY/SORT BY/DISTRIBUTE BY/CLUSTER BY.","sqlState":"0A000","queryContext":[{"objectType":"","objectName":"","startIndex":19,"stopIndex":42,"fragment":"ORDER BY a\n CLUSTER BY b"}],"line":3,"position":1}}] -- !query SELECT - get_json_object(result, '$.error.errorClass') AS error_class, - get_json_object(result, '$.error.line') AS line, - get_json_object(result, '$.error.position') AS position, - get_json_object(result, '$.error.queryContext[0].startIndex') AS start_index + get_json_object(result, '$[0].error.errorClass') AS error_class, + get_json_object(result, '$[0].error.line') AS line, + get_json_object(result, '$[0].error.position') AS position, + get_json_object(result, '$[0].error.queryContext[0].startIndex') AS start_index FROM ( SELECT parse_sql( 'SELECT * @@ -301,7 +325,7 @@ SELECT parse_sql('') -- !query schema struct -- !query output -{"parse_success":false,"error":{"errorClass":"PARSE_EMPTY_STATEMENT","messageTemplate":"Syntax error, unexpected empty statement.","sqlState":"42617","line":1,"position":0}} +[] -- !query @@ -309,7 +333,7 @@ SELECT parse_sql('USE bad-name') -- !query schema struct -- !query output -{"parse_success":false,"error":{"errorClass":"INVALID_IDENTIFIER","messageTemplate":"The unquoted identifier is invalid and must be back quoted as: ``.\nUnquoted identifiers can only contain ASCII letters ('a' - 'z', 'A' - 'Z'), digits ('0' - '9'), and underbar ('_').\nUnquoted identifiers must also not start with a digit.\nDifferent data sources and meta stores may impose additional restrictions on valid identifiers.","sqlState":"42602","messageParameters":{"ident":"bad-name"},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":12,"fragment":"USE bad-name"}],"line":1,"position":7}} +[{"start":1,"length":12,"parse_success":false,"error":{"errorClass":"INVALID_IDENTIFIER","messageTemplate":"The unquoted identifier is invalid and must be back quoted as: ``.\nUnquoted identifiers can only contain ASCII letters ('a' - 'z', 'A' - 'Z'), digits ('0' - '9'), and underbar ('_').\nUnquoted identifiers must also not start with a digit.\nDifferent data sources and meta stores may impose additional restrictions on valid identifiers.","sqlState":"42602","messageParameters":{"ident":"bad-name"},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":12,"fragment":"USE bad-name"}],"line":1,"position":7}}] -- !query @@ -317,7 +341,7 @@ SELECT parse_sql('WITH c AS (SELECT 1), c AS (SELECT 2) SELECT * FROM c') -- !query schema struct -- !query output -{"parse_success":false,"error":{"errorClass":"DUPLICATED_CTE_NAMES","messageTemplate":"CTE definition can't have duplicate names: .","sqlState":"42602","messageParameters":{"duplicateNames":"`c`"},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":53,"fragment":"WITH c AS (SELECT 1), c AS (SELECT 2) SELECT * FROM c"}],"line":1,"position":0}} +[{"start":1,"length":53,"parse_success":false,"error":{"errorClass":"DUPLICATED_CTE_NAMES","messageTemplate":"CTE definition can't have duplicate names: .","sqlState":"42602","messageParameters":{"duplicateNames":"`c`"},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":53,"fragment":"WITH c AS (SELECT 1), c AS (SELECT 2) SELECT * FROM c"}],"line":1,"position":0}}] -- !query @@ -325,7 +349,7 @@ SELECT parse_sql('MERGE INTO target USING source ON target.id = source.id') -- !query schema struct -- !query output -{"parse_success":false,"error":{"errorClass":"MERGE_WITHOUT_WHEN","messageTemplate":"There must be at least one WHEN clause in a MERGE statement.","sqlState":"42601","queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":55,"fragment":"MERGE INTO target USING source ON target.id = source.id"}],"line":1,"position":0}} +[{"start":1,"length":55,"parse_success":false,"error":{"errorClass":"MERGE_WITHOUT_WHEN","messageTemplate":"There must be at least one WHEN clause in a MERGE statement.","sqlState":"42601","queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":55,"fragment":"MERGE INTO target USING source ON target.id = source.id"}],"line":1,"position":0}}] -- !query @@ -333,7 +357,7 @@ SELECT parse_sql('EXPLAIN SELECT 1') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"EXPLAIN","statement_code":-23,"select_list":[{"name":[]}]} +[{"start":1,"length":16,"parse_success":true,"statement_identifier":"EXPLAIN","statement_code":-23,"select_list":[{"name":[]}]}] -- !query @@ -341,7 +365,7 @@ SELECT parse_sql('SET spark.sql.adaptive.enabled=true') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"SET","statement_code":-24} +[{"start":1,"length":35,"parse_success":true,"statement_identifier":"SET","statement_code":-24}] -- !query @@ -349,7 +373,7 @@ SELECT parse_sql('ADD JAR /tmp/x.jar') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"ADD JAR","statement_code":-26} +[{"start":1,"length":18,"parse_success":true,"statement_identifier":"ADD JAR","statement_code":-26}] -- !query @@ -357,7 +381,7 @@ SELECT parse_sql('CREATE VIEW v AS SELECT a, b FROM t') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"CREATE VIEW","statement_code":84,"target_table_references":[["v"]],"source_table_references":[["t"]],"select_list":[{"name":["a"]},{"name":["b"]}]} +[{"start":1,"length":35,"parse_success":true,"statement_identifier":"CREATE VIEW","statement_code":84,"target_table_references":[["v"]],"source_table_references":[["t"]],"select_list":[{"name":["a"]},{"name":["b"]}]}] -- !query @@ -365,7 +389,7 @@ SELECT parse_sql('SELECT 1 AS IDENTIFIER(''alias.field'')') -- !query schema struct -- !query output -{"parse_success":false,"error":{"errorClass":"IDENTIFIER_TOO_MANY_NAME_PARTS","messageTemplate":" is not a valid identifier as it has more than name parts.","sqlState":"42601","messageParameters":{"identifier":"`alias`.`field`","limit":"1"},"queryContext":[{"objectType":"","objectName":"","startIndex":8,"stopIndex":37,"fragment":"1 AS IDENTIFIER('alias.field')"}],"line":1,"position":12}} +[{"start":1,"length":37,"parse_success":false,"error":{"errorClass":"IDENTIFIER_TOO_MANY_NAME_PARTS","messageTemplate":" is not a valid identifier as it has more than name parts.","sqlState":"42601","messageParameters":{"identifier":"`alias`.`field`","limit":"1"},"queryContext":[{"objectType":"","objectName":"","startIndex":8,"stopIndex":37,"fragment":"1 AS IDENTIFIER('alias.field')"}],"line":1,"position":12}}] -- !query @@ -373,7 +397,7 @@ SELECT parse_sql('SELECT DATE ''not-a-date''') -- !query schema struct -- !query output -{"parse_success":false,"error":{"errorClass":"INVALID_TYPED_LITERAL","messageTemplate":"The value of the typed literal is invalid: .","sqlState":"42604","messageParameters":{"value":"'not-a-date'","valueType":"\"DATE\""},"queryContext":[{"objectType":"","objectName":"","startIndex":8,"stopIndex":24,"fragment":"DATE 'not-a-date'"}],"line":1,"position":7}} +[{"start":1,"length":24,"parse_success":false,"error":{"errorClass":"INVALID_TYPED_LITERAL","messageTemplate":"The value of the typed literal is invalid: .","sqlState":"42604","messageParameters":{"value":"'not-a-date'","valueType":"\"DATE\""},"queryContext":[{"objectType":"","objectName":"","startIndex":8,"stopIndex":24,"fragment":"DATE 'not-a-date'"}],"line":1,"position":7}}] -- !query @@ -388,15 +412,17 @@ struct -- !query output -{"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"'2'","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":35,"fragment":"BEGIN\n SELECT 1;\n SELEC 2;\n END"}],"line":3,"position":9}} +[{"start":1,"length":17,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"end of input","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":17,"fragment":"BEGIN\n SELECT 1"}],"line":2,"position":11}},{"start":23,"length":7,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"'SELEC'","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":7,"fragment":"SELEC 2"}],"line":1,"position":0}},{"start":33,"length":3,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"'END'","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":3,"fragment":"END"}],"line":1,"position":0}}] -- !query SELECT - get_json_object(result, '$.error.errorClass') AS error_class, - get_json_object(result, '$.error.line') AS line, - get_json_object(result, '$.error.position') AS position, - get_json_object(result, '$.error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[1].start') AS start, + get_json_object(result, '$[1].length') AS length, + get_json_object(result, '$[1].error.errorClass') AS error_class, + get_json_object(result, '$[1].error.line') AS line, + get_json_object(result, '$[1].error.position') AS position, + get_json_object(result, '$[1].error.queryContext[0].fragment') AS fragment FROM ( SELECT parse_sql( 'BEGIN @@ -405,12 +431,9 @@ FROM ( END') AS result ) -- !query schema -struct +struct -- !query output -PARSE_SYNTAX_ERROR 3 9 BEGIN - SELECT 1; - SELEC 2; - END +23 7 PARSE_SYNTAX_ERROR 1 0 SELEC 2 -- !query @@ -427,15 +450,15 @@ struct -- !query output -{"parse_success":false,"error":{"errorClass":"LABELS_MISMATCH","messageTemplate":"Begin label does not match the end label .","sqlState":"42K0L","messageParameters":{"beginLabel":"`lbl_begin`","endLabel":"`lbl_end`"},"queryContext":[{"objectType":"","objectName":"","startIndex":10,"stopIndex":19,"fragment":"lbl_begin:"}],"line":2,"position":3}} +[{"start":1,"length":61,"parse_success":false,"error":{"errorClass":"LABELS_MISMATCH","messageTemplate":"Begin label does not match the end label .","sqlState":"42K0L","messageParameters":{"beginLabel":"`lbl_begin`","endLabel":"`lbl_end`"},"queryContext":[{"objectType":"","objectName":"","startIndex":10,"stopIndex":19,"fragment":"lbl_begin:"}],"line":2,"position":3}}] -- !query SELECT - get_json_object(result, '$.error.errorClass') AS error_class, - get_json_object(result, '$.error.line') AS line, - get_json_object(result, '$.error.position') AS position, - get_json_object(result, '$.error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[0].error.errorClass') AS error_class, + get_json_object(result, '$[0].error.line') AS line, + get_json_object(result, '$[0].error.position') AS position, + get_json_object(result, '$[0].error.queryContext[0].fragment') AS fragment FROM ( SELECT parse_sql( 'BEGIN @@ -459,9 +482,9 @@ AS t(sql_text) -- !query schema struct -- !query output -CACHE TABLE t {"parse_success":true,"statement_identifier":"CACHE TABLE","statement_code":-1,"target_table_references":[["t"]]} -INSERT INTO t SELECT 1 {"parse_success":true,"statement_identifier":"INSERT","statement_code":50,"target_table_references":[["t"]],"select_list":[{"name":[]}]} -SELECT 1 {"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"select_list":[{"name":[]}]} +CACHE TABLE t [{"start":1,"length":13,"parse_success":true,"statement_identifier":"CACHE TABLE","statement_code":-1,"target_table_references":[["t"]]}] +INSERT INTO t SELECT 1 [{"start":1,"length":22,"parse_success":true,"statement_identifier":"INSERT","statement_code":50,"target_table_references":[["t"]],"select_list":[{"name":[]}]}] +SELECT 1 [{"start":1,"length":8,"parse_success":true,"statement_identifier":"SELECT","statement_code":21,"select_list":[{"name":[]}]}] -- !query @@ -469,7 +492,7 @@ SELECT parse_sql('BEGIN SELECT 1; END') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"BEGIN END","statement_code":-22} +[{"start":1,"length":19,"parse_success":true,"statement_identifier":"BEGIN END","statement_code":-22}] -- !query @@ -477,7 +500,7 @@ SELECT parse_sql('BEGIN SELECT count(a) FROM script_t WHERE c = :p; END') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"BEGIN END","statement_code":-22,"source_table_references":[["script_t"]],"function_references":[["count"]],"parameter_markers":{"named":["p"]}} +[{"start":1,"length":53,"parse_success":true,"statement_identifier":"BEGIN END","statement_code":-22,"source_table_references":[["script_t"]],"function_references":[["count"]],"parameter_markers":{"named":["p"]}}] -- !query @@ -485,7 +508,7 @@ SELECT parse_sql('BEGIN SELECT * FROM t WHERE a = ?; END') -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"BEGIN END","statement_code":-22,"source_table_references":[["t"]],"parameter_markers":{"unnamed_count":1}} +[{"start":1,"length":38,"parse_success":true,"statement_identifier":"BEGIN END","statement_code":-22,"source_table_references":[["t"]],"parameter_markers":{"unnamed_count":1}}] -- !query @@ -493,7 +516,7 @@ SELECT parse_sql('BEGIN IF (SELECT flag FROM gate) THEN INSERT INTO dest SELECT -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"BEGIN END","statement_code":-22,"target_table_references":[["dest"],["src_else"]],"source_table_references":[["gate"],["src_if"]]} +[{"start":1,"length":115,"parse_success":true,"statement_identifier":"BEGIN END","statement_code":-22,"target_table_references":[["dest"],["src_else"]],"source_table_references":[["gate"],["src_if"]]}] -- !query @@ -501,7 +524,7 @@ SELECT parse_sql('BEGIN DECLARE EXIT HANDLER FOR SQLEXCEPTION BEGIN INSERT INTO -- !query schema struct -- !query output -{"parse_success":true,"statement_identifier":"BEGIN END","statement_code":-22,"target_table_references":[["err_log"]],"source_table_references":[["failing_row"],["main_t"]]} +[{"start":1,"length":127,"parse_success":true,"statement_identifier":"BEGIN END","statement_code":-22,"target_table_references":[["err_log"]],"source_table_references":[["failing_row"],["main_t"]]}] -- !query @@ -568,4 +591,4 @@ struct -- !query output -{"parse_success":true,"statement_identifier":"BEGIN END","statement_code":-22,"target_table_references":[["error_log"],["output_names"],["update_target"],["delete_target"]],"source_table_references":[["error_source"],["input_names"],["control_flags"],["update_source"],["delete_source"],["loop_source"],["loop_body"]],"function_references":[["format_string"],["normalize_name"],["is_valid"],["upper"],["enabled"],["coalesce"],["should_update"],["max"],["expired"],["ready"],["audit"],["count"]]} +[{"start":1,"length":783,"parse_success":true,"statement_identifier":"BEGIN END","statement_code":-22,"target_table_references":[["error_log"],["output_names"],["update_target"],["delete_target"]],"source_table_references":[["error_source"],["input_names"],["control_flags"],["update_source"],["delete_source"],["loop_source"],["loop_body"]],"function_references":[["format_string"],["normalize_name"],["is_valid"],["upper"],["enabled"],["coalesce"],["should_update"],["max"],["expired"],["ready"],["audit"],["count"]]}] diff --git a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/expressions/ParseSqlSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/expressions/ParseSqlSuite.scala index f44a409a0d641..0322c01c4dbca 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/expressions/ParseSqlSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/expressions/ParseSqlSuite.scala @@ -36,6 +36,16 @@ class ParseSqlSuite extends SparkFunSuite with ExpressionEvalHelper with SQLHelp parse(result) } + private def evalStatements(sql: String): List[JObject] = evalJson(sql) match { + case JArray(values) => values.map(_.asInstanceOf[JObject]) + case other => fail(s"expected JSON array, got: $other") + } + + private def evalStatement(sql: String): JObject = evalStatements(sql) match { + case value :: Nil => value + case other => fail(s"expected one statement, got ${other.size}: $other") + } + test("parse_sql is disabled by default") { assert(!SQLConf.get.parseSqlEnabled) checkError( @@ -58,7 +68,9 @@ class ParseSqlSuite extends SparkFunSuite with ExpressionEvalHelper with SQLHelp test("parse_sql returns JSON for a valid SELECT") { withSQLConf(SQLConf.PARSE_SQL_ENABLED.key -> "true") { - val j = evalJson("SELECT 1 AS a") + val j = evalStatement("SELECT 1 AS a") + assert(j \ "start" === JInt(1)) + assert(j \ "length" === JInt(13)) assert(j \ "parse_success" === JBool(true)) assert(j \ "statement_identifier" === JString("SELECT")) assert(j \ "statement_code" === JInt(21)) @@ -73,7 +85,7 @@ class ParseSqlSuite extends SparkFunSuite with ExpressionEvalHelper with SQLHelp test("parse_sql does not throw on syntax error") { withSQLConf(SQLConf.PARSE_SQL_ENABLED.key -> "true") { - val j = evalJson("NOT A STATEMENT !!!") + val j = evalStatement("NOT A STATEMENT !!!") assert(j \ "parse_success" === JBool(false)) assert(j \ "error" \ "errorClass" === JString("PARSE_SYNTAX_ERROR")) } @@ -83,8 +95,17 @@ class ParseSqlSuite extends SparkFunSuite with ExpressionEvalHelper with SQLHelp withSQLConf(SQLConf.PARSE_SQL_ENABLED.key -> "true") { val expr = ParseSql(Literal("INSERT INTO t SELECT 1")) assert(expr.isInstanceOf[CodegenFallback]) - val j = evalJson("INSERT INTO t SELECT 1") + val j = evalStatement("INSERT INTO t SELECT 1") assert(j \ "statement_identifier" === JString("INSERT")) } } + + test("parse_sql returns a JSON array for a SQL batch") { + withSQLConf(SQLConf.PARSE_SQL_ENABLED.key -> "true") { + val statements = evalStatements("SELECT 1; SELECT 2") + assert(statements.size === 2) + assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(11))) + assert(statements.map(_ \ "length") === Seq(JInt(8), JInt(8))) + } + } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala index 397584cb0c8db..69e608a4f7d3f 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala @@ -29,8 +29,16 @@ import org.apache.spark.sql.internal.SQLConf */ class ParseSqlResultSuite extends SparkFunSuite { - private def obj(sql: String): JObject = - parse(ParseSqlResult.fromSql(sql)).asInstanceOf[JObject] + private def objs(sql: String): List[JObject] = + parse(ParseSqlResult.fromSql(sql)) match { + case JArray(values) => values.map(_.asInstanceOf[JObject]) + case other => fail(s"expected JSON array, got: $other") + } + + private def obj(sql: String): JObject = objs(sql) match { + case value :: Nil => value + case other => fail(s"expected one statement, got ${other.size}: $other") + } private def tableRefs(sql: String, field: String): Set[Seq[String]] = obj(sql) \ field match { @@ -67,6 +75,51 @@ class ParseSqlResultSuite extends SparkFunSuite { assert(SqlStatementCodes.CreateMetricViewStmt.statementCode === -37) } + test("batches return one statement object per statement with source spans") { + Seq("select 1; select 2", "select 1; select 2;").foreach { sql => + val statements = objs(sql) + assert(statements.size === 2) + assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(11))) + assert(statements.map(_ \ "length") === Seq(JInt(8), JInt(8))) + assert(statements.map(_ \ "parse_success") === Seq(JBool(true), JBool(true))) + assert(statements.map(_ \ "statement_identifier") === + Seq(JString("SELECT"), JString("SELECT"))) + } + + val statements = objs(" SELECT 1 ;\n SELECT 2; ") + assert(statements.map(_ \ "start") === Seq(JInt(3), JInt(15))) + assert(statements.map(_ \ "length") === Seq(JInt(8), JInt(8))) + + val sqlWithDroppedComment = "SELECT 1; /* SELECT 2 */; SELECT 2" + val commentStatements = objs(sqlWithDroppedComment) + assert(commentStatements.map(_ \ "start") === + Seq(JInt(1), JInt(sqlWithDroppedComment.lastIndexOf("SELECT 2") + 1))) + } + + test("batch errors are isolated and preserve statement order") { + val statements = objs("SELECT 1; SELEC 2; SELECT 3") + assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(11), JInt(20))) + assert(statements.map(_ \ "length") === Seq(JInt(8), JInt(7), JInt(8))) + assert(statements.map(_ \ "parse_success") === + Seq(JBool(true), JBool(false), JBool(true))) + assert(statements(1) \ "error" \ "errorClass" === JString("PARSE_SYNTAX_ERROR")) + } + + test("SQL scripts remain one statement in a batch") { + val statements = objs("BEGIN SELECT 1; SELECT 2; END; SELECT 3") + assert(statements.size === 2) + assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(32))) + assert(statements.map(_ \ "length") === Seq(JInt(29), JInt(8))) + assert(statements.map(_ \ "statement_identifier") === + Seq(JString("BEGIN END"), JString("SELECT"))) + } + + test("empty batches contain no statements") { + Seq("", " ", ";;", "-- comment").foreach { sql => + assert(objs(sql).isEmpty, sql) + } + } + test("TABLE and VALUES classify as SELECT") { val table = obj("TABLE t") assert(table \ "statement_identifier" === JString("SELECT")) From 003156797de302ec8c3051117037f04f010c9725 Mon Sep 17 00:00:00 2001 From: srielau Date: Sat, 5 Sep 2026 01:16:32 +0000 Subject: [PATCH 02/12] [SPARK-59255][SQL] Fix Unicode offsets in SQL batch splitter --- .../parser/SqlStatementSplitter.scala | 25 +++++++++++-------- .../parser/SqlStatementSplitterSuite.scala | 16 ++++++++++++ .../catalyst/parser/ParseSqlResultSuite.scala | 13 ++++++++++ 3 files changed, 44 insertions(+), 10 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala index 83ef5ab72e6d9..98af6b6ffe750 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala @@ -191,7 +191,11 @@ object SqlStatementSplitter { val conf = SqlApiConf.get def appendToken(token: Token): Unit = { - if (buffer.isEmpty) bufferStart = token.getStartIndex + if (buffer.isEmpty) { + // CodePointCharStream token offsets count Unicode code points, while + // String offsets and lengths count UTF-16 code units. + bufferStart = sqlText.offsetByCodePoints(0, token.getStartIndex) + } buffer.append(token.getText) } @@ -354,13 +358,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 to UTF-16 String offsets. + * `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). @@ -386,9 +391,9 @@ object SqlStatementSplitter { conf: SqlApiConf): ParseOutcome = { val firstTok = stream.get(startIdx) val lastTok = stream.get(endIdx) - val regionStart = firstTok.getStartIndex + val regionStart = sqlText.offsetByCodePoints(0, firstTok.getStartIndex) // Token.getStopIndex is inclusive, substring's upper bound is exclusive. - val regionEnd = lastTok.getStopIndex + 1 + val regionEnd = sqlText.offsetByCodePoints(0, lastTok.getStopIndex + 1) val original = sqlText.substring(regionStart, regionEnd) val preprocessed = validationPreprocess(original) diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala index eec2db0d53071..4044823f24e7a 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala @@ -85,6 +85,22 @@ class SqlStatementSplitterSuite extends SparkFunSuite { assert(partial.length == 8) } + test("source positions use UTF-16 offsets for supplementary characters") { + val emoji = "\uD83D\uDE00" + 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) + } + // ---------------------------------------------------------------------------------- // Error tolerance (mirrors Trino behavior) // ---------------------------------------------------------------------------------- diff --git a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala index 69e608a4f7d3f..c1fcfc47d5f40 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala @@ -94,6 +94,19 @@ class ParseSqlResultSuite extends SparkFunSuite { val commentStatements = objs(sqlWithDroppedComment) assert(commentStatements.map(_ \ "start") === Seq(JInt(1), JInt(sqlWithDroppedComment.lastIndexOf("SELECT 2") + 1))) + + val emoji = "\uD83D\uDE00" + val unicodeSql = s"SELECT '$emoji$emoji'; SELECT 2;" + val unicodeStatements = objs(unicodeSql) + val spans = unicodeStatements.map { statement => + val JInt(start) = statement \ "start" + val JInt(length) = statement \ "length" + (start.toInt, length.toInt) + } + assert(spans === Seq((1, 13), (16, 8))) + assert(spans.map { case (start, length) => + unicodeSql.substring(start - 1, start - 1 + length) + } === Seq(s"SELECT '$emoji$emoji'", "SELECT 2")) } test("batch errors are isolated and preserve statement order") { From c15a60657c5923bf83d6a061e1110294705fdc24 Mon Sep 17 00:00:00 2001 From: srielau Date: Sat, 5 Sep 2026 19:28:10 +0000 Subject: [PATCH 03/12] [SPARK-59255][SQL] Document UTF-16 spans and trim lexer whitespace --- .../parser/SqlStatementSplitter.scala | 65 +++++++++++++++---- .../parser/SqlStatementSplitterSuite.scala | 11 ++++ .../sql/catalyst/expressions/ParseSql.scala | 11 +++- .../sql/catalyst/parser/ParseSqlResult.scala | 7 +- .../catalyst/parser/ParseSqlResultSuite.scala | 8 +++ 5 files changed, 84 insertions(+), 18 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala index 98af6b6ffe750..6bdcab8c58675 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala @@ -156,6 +156,11 @@ object SqlStatementSplitter { 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) @@ -191,11 +196,7 @@ object SqlStatementSplitter { val conf = SqlApiConf.get def appendToken(token: Token): Unit = { - if (buffer.isEmpty) { - // CodePointCharStream token offsets count Unicode code points, while - // String offsets and lengths count UTF-16 code units. - bufferStart = sqlText.offsetByCodePoints(0, token.getStartIndex) - } + if (buffer.isEmpty) bufferStart = toUtf16(token.getStartIndex) buffer.append(token.getText) } @@ -207,12 +208,11 @@ object SqlStatementSplitter { def positionedStatement(terminator: String): Option[PositionedSqlStatement] = { val raw = buffer.toString - val statement = raw.trim + val (leadingWhitespace, statement) = trimSqlWhitespace(raw) if (statement.isEmpty) { None } else { - val leadingWhitespace = raw.indexOf(statement) - assert(bufferStart >= 0 && leadingWhitespace >= 0) + assert(bufferStart >= 0) Some(PositionedSqlStatement( statement, terminator, @@ -254,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 @@ -359,7 +360,7 @@ object SqlStatementSplitter { * region) as a complete top-level Spark SQL statement. * * The region is extracted from the original source by converting ANTLR's - * Unicode code-point token offsets to UTF-16 String offsets. + * 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 @@ -384,6 +385,7 @@ object SqlStatementSplitter { */ private def tryParseRegion( sqlText: String, + toUtf16: Array[Int], stream: CommonTokenStream, startIdx: Int, endIdx: Int, @@ -391,9 +393,9 @@ object SqlStatementSplitter { conf: SqlApiConf): ParseOutcome = { val firstTok = stream.get(startIdx) val lastTok = stream.get(endIdx) - val regionStart = sqlText.offsetByCodePoints(0, firstTok.getStartIndex) + val regionStart = toUtf16(firstTok.getStartIndex) // Token.getStopIndex is inclusive, substring's upper bound is exclusive. - val regionEnd = sqlText.offsetByCodePoints(0, lastTok.getStopIndex + 1) + val regionEnd = toUtf16(lastTok.getStopIndex + 1) val original = sqlText.substring(regionStart, regionEnd) val preprocessed = validationPreprocess(original) @@ -460,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 diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala index 4044823f24e7a..56b8fc4f4f795 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala @@ -101,6 +101,17 @@ class SqlStatementSplitterSuite extends SparkFunSuite { assert(result.partialStatement.isEmpty) } + test("source positions trim Spark SQL Unicode whitespace") { + val sql = "\u00A0SELECT 1\u00A0;" + 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) // ---------------------------------------------------------------------------------- diff --git a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala index ed406385c8822..9bffb3ffb518a 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala @@ -37,11 +37,13 @@ import org.apache.spark.unsafe.types.UTF8String * User-facing parse errors become JSON; unexpected internal failures propagate. */ // scalastyle:off line.size.limit +// scalastyle:off nonascii @ExpressionDescription( 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 start - and length, parse success, Table 39 statement identifier/code, target and source - table references for lineage, select-list column names, and parameter markers). + 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 @@ -57,11 +59,14 @@ import org.apache.spark.unsafe.types.UTF8String Examples: > 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 diff --git a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResult.scala b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResult.scala index 8f86ceb0900ed..cd17bbc30162c 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResult.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResult.scala @@ -43,9 +43,10 @@ import org.apache.spark.sql.execution.datasources.CreateTempViewUsing * executors without a session, so only the stock parser is available under * distributed eval. * - * Every statement object includes its 1-based `start` in the original batch - * and its `length`, excluding surrounding whitespace and the terminating - * semicolon. On success it also includes `parse_success`, the statement + * Every statement object includes its 1-based UTF-16 code-unit `start` in the + * original batch and its UTF-16 code-unit `length`, excluding surrounding + * whitespace and the terminating semicolon. On success it also includes + * `parse_success`, the statement * identifier/code (ISO/IEC 9075-2:2023 Table 39), and omits unused optional * fields (`target_table_references`, `source_table_references`, * `function_references`, `select_list`, `parameter_markers`) when empty. On diff --git a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala index c1fcfc47d5f40..b4334cca76de7 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala @@ -107,6 +107,14 @@ class ParseSqlResultSuite extends SparkFunSuite { assert(spans.map { case (start, length) => unicodeSql.substring(start - 1, start - 1 + length) } === Seq(s"SELECT '$emoji$emoji'", "SELECT 2")) + + val nbspSql = "\u00A0SELECT 1\u00A0;" + val nbsp = objs(nbspSql).head + val JInt(nbspStart) = nbsp \ "start" + val JInt(nbspLength) = nbsp \ "length" + assert((nbspStart.toInt, nbspLength.toInt) === (2, 8)) + assert(nbspSql.substring(nbspStart.toInt - 1, nbspStart.toInt - 1 + nbspLength.toInt) === + "SELECT 1") } test("batch errors are isolated and preserve statement order") { From 9e8b87cda6cd10d02c75e4fbc4228f705de73d2b Mon Sep 17 00:00:00 2001 From: srielau Date: Sun, 6 Sep 2026 23:41:19 +0000 Subject: [PATCH 04/12] [SPARK-59255][SQL] Avoid non-ASCII Unicode escapes in parse_sql tests --- .../catalyst/parser/SqlStatementSplitterSuite.scala | 5 +++-- .../sql/catalyst/parser/ParseSqlResultSuite.scala | 11 ++++++----- 2 files changed, 9 insertions(+), 7 deletions(-) diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala index 56b8fc4f4f795..dfa07229901d5 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala @@ -86,7 +86,7 @@ class SqlStatementSplitterSuite extends SparkFunSuite { } test("source positions use UTF-16 offsets for supplementary characters") { - val emoji = "\uD83D\uDE00" + val emoji = new String(Character.toChars(0x1F600)) val sql = s"SELECT '$emoji$emoji'; SELECT 2;" val result = SqlStatementSplitter.splitWithPositions(sql, identity) @@ -102,7 +102,8 @@ class SqlStatementSplitterSuite extends SparkFunSuite { } test("source positions trim Spark SQL Unicode whitespace") { - val sql = "\u00A0SELECT 1\u00A0;" + 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") diff --git a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala index b4334cca76de7..af277692e766e 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala @@ -95,7 +95,7 @@ class ParseSqlResultSuite extends SparkFunSuite { assert(commentStatements.map(_ \ "start") === Seq(JInt(1), JInt(sqlWithDroppedComment.lastIndexOf("SELECT 2") + 1))) - val emoji = "\uD83D\uDE00" + val emoji = new String(Character.toChars(0x1F600)) val unicodeSql = s"SELECT '$emoji$emoji'; SELECT 2;" val unicodeStatements = objs(unicodeSql) val spans = unicodeStatements.map { statement => @@ -108,10 +108,11 @@ class ParseSqlResultSuite extends SparkFunSuite { unicodeSql.substring(start - 1, start - 1 + length) } === Seq(s"SELECT '$emoji$emoji'", "SELECT 2")) - val nbspSql = "\u00A0SELECT 1\u00A0;" - val nbsp = objs(nbspSql).head - val JInt(nbspStart) = nbsp \ "start" - val JInt(nbspLength) = nbsp \ "length" + val nbsp = 0xA0.toChar + val nbspSql = s"${nbsp}SELECT 1$nbsp;" + val nbspStmt = objs(nbspSql).head + val JInt(nbspStart) = nbspStmt \ "start" + val JInt(nbspLength) = nbspStmt \ "length" assert((nbspStart.toInt, nbspLength.toInt) === (2, 8)) assert(nbspSql.substring(nbspStart.toInt - 1, nbspStart.toInt - 1 + nbspLength.toInt) === "SELECT 1") From 94894f6941eb8d753818c1b683128a6e158682e6 Mon Sep 17 00:00:00 2001 From: srielau Date: Tue, 8 Sep 2026 18:40:26 +0000 Subject: [PATCH 05/12] [SPARK-59255][SQL] Preserve malformed compound statement boundaries in parse_sql --- .../parser/SqlStatementSplitter.scala | 105 ++++++++++++++++-- .../parser/SqlStatementSplitterSuite.scala | 23 ++++ .../spark/sql/execution/SparkSqlParser.scala | 3 +- .../analyzer-results/parse-sql.sql.out | 14 +-- .../resources/sql-tests/inputs/parse-sql.sql | 14 +-- .../sql-tests/results/parse-sql.sql.out | 19 ++-- .../catalyst/parser/ParseSqlResultSuite.scala | 8 ++ 7 files changed, 156 insertions(+), 30 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala index 6bdcab8c58675..d6f1ae24628c3 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala @@ -148,11 +148,14 @@ object SqlStatementSplitter { /** * Split SQL while retaining each statement's source position. This is used by - * parse-only tooling that reports spans into the original input. + * parse-only tooling that reports spans into the original input. Such tooling + * can preserve a balanced BEGIN ... END boundary even when its body is malformed; + * generic splitting leaves that disabled to retain extension fallback behavior. */ private[sql] def splitWithPositions( sqlText: String, - validationPreprocess: String => String): PositionedSqlStatementSplitResult = { + validationPreprocess: String => String, + preserveMalformedCompoundBoundaries: Boolean = false): PositionedSqlStatementSplitResult = { require(sqlText != null, "sqlText must not be null") require(validationPreprocess != null, "validationPreprocess must not be null") @@ -264,8 +267,31 @@ object SqlStatementSplitter { // prefix that includes the next `;`. d += 1 case FailedNonEof => - // Structurally invalid; stop extending. - failedNonEof = true + // A malformed but balanced compound body still belongs to its enclosing + // BEGIN ... END statement. Use error-recovering parsing to find that real outer + // boundary before falling back to ordinary delimiter splitting. + val recovered = if (preserveMalformedCompoundBoundaries) { + findMalformedCompoundEnd( + sqlText, + toUtf16, + tokenStream, + startIdx, + delimiterPositions, + d, + validationPreprocess, + conf) + } else { + None + } + recovered match { + case Some((recoveredDelimiter, recoveredEnd)) => + parsedOk = true + d = recoveredDelimiter + matchedDelimIdx = recoveredEnd + case None => + // Structurally invalid ordinary statement; stop extending. + failedNonEof = true + } } } @@ -275,7 +301,11 @@ object SqlStatementSplitter { // `startIdx` are leading whitespace/comments that we preserve in the // buffer too (so the emitted statement text matches the original // input shape, modulo trimming). - val terminator = tokenStream.get(matchedDelimIdx).getText + val terminator = if (tokenStream.get(matchedDelimIdx).getType == Token.EOF) { + "" + } else { + tokenStream.get(matchedDelimIdx).getText + } while (index < matchedDelimIdx) { val tok = tokenStream.get(index) if (tok.getChannel != Token.HIDDEN_CHANNEL) bufferHasContent = true @@ -347,6 +377,62 @@ object SqlStatementSplitter { unclosed && partial.nonEmpty) } + /** + * Returns the delimiter-array index and ending token index of a real outer END for a malformed + * compound statement. Error recovery may repair the body, but a missing END is synthetic and + * has token index -1. + */ + private def findMalformedCompoundEnd( + sqlText: String, + toUtf16: Array[Int], + stream: CommonTokenStream, + startIdx: Int, + delimiterPositions: Array[Int], + fromDelimiter: Int, + validationPreprocess: String => String, + conf: SqlApiConf): Option[(Int, Int)] = { + if (stream.get(startIdx).getType != SqlBaseLexer.BEGIN) { + return None + } + + var delimiter = fromDelimiter + while (delimiter <= delimiterPositions.length) { + val endIdx = if (delimiter < delimiterPositions.length) { + delimiterPositions(delimiter) + } else { + stream.size() - 1 + } + val firstTok = stream.get(startIdx) + val lastTok = stream.get(endIdx) + val regionStart = toUtf16(firstTok.getStartIndex) + val regionEnd = if (lastTok.getType == Token.EOF) { + sqlText.length + } else { + toUtf16(lastTok.getStopIndex + 1) + } + val candidate = validationPreprocess(sqlText.substring(regionStart, regionEnd)) + val lexer = new SqlBaseLexer( + new UpperCaseCharStream(CharStreams.fromString(candidate))) + lexer.removeErrorListeners() + val tokens = new CommonTokenStream(lexer) + tokens.fill() + val parser = new SqlBaseParser(tokens) + configureSplitterParser(parser, conf, bailOnError = false) + parser.getInterpreter.setPredictionMode(PredictionMode.LL) + try { + val context = parser.singleCompoundStatement() + val end = context.END() + if (end != null && end.getSymbol.getTokenIndex >= 0 && tokens.LA(1) == Token.EOF) { + return Some((delimiter, endIdx)) + } + } catch { + case _: StackOverflowError => return None + } + delimiter += 1 + } + None + } + /** Outcome of attempting to parse one statement candidate. */ private sealed trait ParseOutcome private case object ParsedOk extends ParseOutcome @@ -447,7 +533,10 @@ object SqlStatementSplitter { * exceptions, but the splitter surfaces them via the * [[SqlStatementSplitResult.hasUnclosedComment]] flag instead. */ - private def configureSplitterParser(parser: SqlBaseParser, conf: SqlApiConf): Unit = { + private def configureSplitterParser( + parser: SqlBaseParser, + conf: SqlApiConf, + bailOnError: Boolean = true): Unit = { if (conf.manageParserCaches) AbstractParser.installCaches(parser) parser.legacy_setops_precedence_enabled = conf.setOpsPrecedenceEnforced @@ -459,7 +548,9 @@ object SqlStatementSplitter { parser.single_character_pipe_operator_enabled = conf.singleCharacterPipeOperatorEnabled parser.removeErrorListeners() - parser.setErrorHandler(new BailErrorStrategy) + if (bailOnError) { + parser.setErrorHandler(new BailErrorStrategy) + } } /** diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala index dfa07229901d5..ae18e2a82036c 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala @@ -516,6 +516,29 @@ class SqlStatementSplitterSuite extends SparkFunSuite { assert(result.partialStatement.isEmpty) } + test("malformed balanced BEGIN block remains one statement") { + val block = "BEGIN SELECT 1; SELEC 2; END" + val result = SqlStatementSplitter + .splitWithPositions( + s"$block; SELECT 3;", + identity, + preserveMalformedCompoundBoundaries = true) + .withoutPositions + assert(result.completeStatements == Seq( + statement(block), + statement("SELECT 3"))) + assert(result.partialStatement.isEmpty) + + val withoutTerminator = SqlStatementSplitter + .splitWithPositions( + block, + identity, + preserveMalformedCompoundBoundaries = true) + .withoutPositions + assert(withoutTerminator.completeStatements == Seq(SqlStatement(block, ""))) + assert(withoutTerminator.partialStatement.isEmpty) + } + test("Valid BEGIN..END block is never split at internal ;") { // A valid `BEGIN ... END` block must be confirmed in full -- the splitter // must never emit at an internal `;`, even if a shorter prefix happens to diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala index bce4250c18975..3ad8b28c4e365 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala @@ -128,7 +128,8 @@ class SparkSqlParser extends AbstractSqlParser { sqlText: String): PositionedSqlStatementSplitResult = SqlStatementSplitter.splitWithPositions( sqlText, - SparkSqlParser.substituteVariablesForValidation) + SparkSqlParser.substituteVariablesForValidation, + preserveMalformedCompoundBoundaries = true) /** * Internal parse method that handles both parameter substitution and regular parsing. diff --git a/sql/core/src/test/resources/sql-tests/analyzer-results/parse-sql.sql.out b/sql/core/src/test/resources/sql-tests/analyzer-results/parse-sql.sql.out index 30d0d993176f9..4cc03527df591 100644 --- a/sql/core/src/test/resources/sql-tests/analyzer-results/parse-sql.sql.out +++ b/sql/core/src/test/resources/sql-tests/analyzer-results/parse-sql.sql.out @@ -412,12 +412,12 @@ Project [parse_sql(BEGIN -- !query SELECT - get_json_object(result, '$[1].start') AS start, - get_json_object(result, '$[1].length') AS length, - get_json_object(result, '$[1].error.errorClass') AS error_class, - get_json_object(result, '$[1].error.line') AS line, - get_json_object(result, '$[1].error.position') AS position, - get_json_object(result, '$[1].error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[0].start') AS start, + get_json_object(result, '$[0].length') AS length, + get_json_object(result, '$[0].error.errorClass') AS error_class, + get_json_object(result, '$[0].error.line') AS line, + get_json_object(result, '$[0].error.position') AS position, + get_json_object(result, '$[0].error.queryContext[0].fragment') AS fragment FROM ( SELECT parse_sql( 'BEGIN @@ -426,7 +426,7 @@ FROM ( END') AS result ) -- !query analysis -Project [get_json_object(result#x, $[1].start) AS start#x, get_json_object(result#x, $[1].length) AS length#x, get_json_object(result#x, $[1].error.errorClass) AS error_class#x, get_json_object(result#x, $[1].error.line) AS line#x, get_json_object(result#x, $[1].error.position) AS position#x, get_json_object(result#x, $[1].error.queryContext[0].fragment) AS fragment#x] +Project [get_json_object(result#x, $[0].start) AS start#x, get_json_object(result#x, $[0].length) AS length#x, get_json_object(result#x, $[0].error.errorClass) AS error_class#x, get_json_object(result#x, $[0].error.line) AS line#x, get_json_object(result#x, $[0].error.position) AS position#x, get_json_object(result#x, $[0].error.queryContext[0].fragment) AS fragment#x] +- SubqueryAlias __auto_generated_subquery_name +- Project [parse_sql(BEGIN SELECT 1; diff --git a/sql/core/src/test/resources/sql-tests/inputs/parse-sql.sql b/sql/core/src/test/resources/sql-tests/inputs/parse-sql.sql index 6b73c8722c6bb..5dcd9e42bbab6 100644 --- a/sql/core/src/test/resources/sql-tests/inputs/parse-sql.sql +++ b/sql/core/src/test/resources/sql-tests/inputs/parse-sql.sql @@ -144,7 +144,7 @@ SELECT parse_sql('CREATE VIEW v AS SELECT a, b FROM t'); SELECT parse_sql('SELECT 1 AS IDENTIFIER(''alias.field'')'); SELECT parse_sql('SELECT DATE ''not-a-date'''); --- location for an error inside a multiline script +-- a malformed balanced script remains one failed statement --QUERY-DELIMITER-START SELECT parse_sql( 'BEGIN @@ -156,12 +156,12 @@ SELECT parse_sql( -- JSON-path access over one shared scripting parse result --QUERY-DELIMITER-START SELECT - get_json_object(result, '$[1].start') AS start, - get_json_object(result, '$[1].length') AS length, - get_json_object(result, '$[1].error.errorClass') AS error_class, - get_json_object(result, '$[1].error.line') AS line, - get_json_object(result, '$[1].error.position') AS position, - get_json_object(result, '$[1].error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[0].start') AS start, + get_json_object(result, '$[0].length') AS length, + get_json_object(result, '$[0].error.errorClass') AS error_class, + get_json_object(result, '$[0].error.line') AS line, + get_json_object(result, '$[0].error.position') AS position, + get_json_object(result, '$[0].error.queryContext[0].fragment') AS fragment FROM ( SELECT parse_sql( 'BEGIN diff --git a/sql/core/src/test/resources/sql-tests/results/parse-sql.sql.out b/sql/core/src/test/resources/sql-tests/results/parse-sql.sql.out index 1310f7c30fffa..4fefe284ce9c5 100644 --- a/sql/core/src/test/resources/sql-tests/results/parse-sql.sql.out +++ b/sql/core/src/test/resources/sql-tests/results/parse-sql.sql.out @@ -412,17 +412,17 @@ struct -- !query output -[{"start":1,"length":17,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"end of input","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":17,"fragment":"BEGIN\n SELECT 1"}],"line":2,"position":11}},{"start":23,"length":7,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"'SELEC'","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":7,"fragment":"SELEC 2"}],"line":1,"position":0}},{"start":33,"length":3,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"'END'","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":3,"fragment":"END"}],"line":1,"position":0}}] +[{"start":1,"length":35,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"'2'","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":35,"fragment":"BEGIN\n SELECT 1;\n SELEC 2;\n END"}],"line":3,"position":9}}] -- !query SELECT - get_json_object(result, '$[1].start') AS start, - get_json_object(result, '$[1].length') AS length, - get_json_object(result, '$[1].error.errorClass') AS error_class, - get_json_object(result, '$[1].error.line') AS line, - get_json_object(result, '$[1].error.position') AS position, - get_json_object(result, '$[1].error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[0].start') AS start, + get_json_object(result, '$[0].length') AS length, + get_json_object(result, '$[0].error.errorClass') AS error_class, + get_json_object(result, '$[0].error.line') AS line, + get_json_object(result, '$[0].error.position') AS position, + get_json_object(result, '$[0].error.queryContext[0].fragment') AS fragment FROM ( SELECT parse_sql( 'BEGIN @@ -433,7 +433,10 @@ FROM ( -- !query schema struct -- !query output -23 7 PARSE_SYNTAX_ERROR 1 0 SELEC 2 +1 35 PARSE_SYNTAX_ERROR 3 9 BEGIN + SELECT 1; + SELEC 2; + END -- !query diff --git a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala index af277692e766e..d8e60e3c038ab 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala @@ -127,6 +127,14 @@ class ParseSqlResultSuite extends SparkFunSuite { assert(statements(1) \ "error" \ "errorClass" === JString("PARSE_SYNTAX_ERROR")) } + test("malformed balanced SQL scripts remain one statement in a batch") { + val statements = objs("BEGIN SELECT 1; SELEC 2; END; SELECT 3") + assert(statements.size === 2) + assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(31))) + assert(statements.map(_ \ "length") === Seq(JInt(28), JInt(8))) + assert(statements.map(_ \ "parse_success") === Seq(JBool(false), JBool(true))) + } + test("SQL scripts remain one statement in a batch") { val statements = objs("BEGIN SELECT 1; SELECT 2; END; SELECT 3") assert(statements.size === 2) From afbd833ef615459959e3fba4a1c7001208ab261c Mon Sep 17 00:00:00 2001 From: srielau Date: Tue, 8 Sep 2026 21:45:10 +0000 Subject: [PATCH 06/12] [SPARK-59255][SQL] Validate malformed compound END boundaries --- .../parser/SqlStatementSplitter.scala | 32 +++++++++++++++++-- .../parser/SqlStatementSplitterSuite.scala | 13 ++++++++ .../catalyst/parser/ParseSqlResultSuite.scala | 10 ++++++ 3 files changed, 52 insertions(+), 3 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala index d6f1ae24628c3..84121ff2c9a55 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala @@ -422,7 +422,7 @@ object SqlStatementSplitter { try { val context = parser.singleCompoundStatement() val end = context.END() - if (end != null && end.getSymbol.getTokenIndex >= 0 && tokens.LA(1) == Token.EOF) { + if (end != null && isOuterCompoundEnd(tokens, end.getSymbol)) { return Some((delimiter, endIdx)) } } catch { @@ -433,6 +433,30 @@ object SqlStatementSplitter { None } + /** + * Returns true only when recovery found a real END and the candidate itself + * ends in a real END, the optional outer semicolon, and EOF. Error recovery + * can bind the context's END node to an inner control terminator and skip its + * suffix, so the candidate suffix is checked independently. + */ + private def isOuterCompoundEnd(tokens: CommonTokenStream, recoveredEnd: Token): Boolean = { + if (recoveredEnd.getTokenIndex < 0) return false + + var index = tokens.size() - 1 + if (index < 0 || tokens.get(index).getType != Token.EOF) return false + index -= 1 + while (index >= 0 && tokens.get(index).getChannel == Token.HIDDEN_CHANNEL) { + index -= 1 + } + if (index >= 0 && tokens.get(index).getType == SqlBaseLexer.SEMICOLON) { + index -= 1 + while (index >= 0 && tokens.get(index).getChannel == Token.HIDDEN_CHANNEL) { + index -= 1 + } + } + index >= 0 && tokens.get(index).getType == SqlBaseLexer.END + } + /** Outcome of attempting to parse one statement candidate. */ private sealed trait ParseOutcome private case object ParsedOk extends ParseOutcome @@ -524,8 +548,10 @@ object SqlStatementSplitter { * Configure a fresh [[SqlBaseParser]] for splitter use: install the managed * caches so candidate parses share ANTLR DFA state across calls, apply the * session's behavior flags so the splitter agrees with the session parser - * on grammar interpretation (e.g. `double_quoted_identifiers`), and install - * a bail error strategy so failures throw immediately. + * on grammar interpretation (e.g. `double_quoted_identifiers`), and, when + * `bailOnError` is true, install a bail error strategy so failures throw + * immediately. Malformed-compound recovery keeps the default error strategy + * so it can inspect the recovered outer END boundary. * * Notably, the splitter does NOT install [[PostProcessor]] or * [[UnclosedCommentProcessor]] -- the former mutates the parse tree (which diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala index ae18e2a82036c..a27309df78937 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala @@ -539,6 +539,19 @@ class SqlStatementSplitterSuite extends SparkFunSuite { assert(withoutTerminator.partialStatement.isEmpty) } + test("malformed nested control END does not end the outer block") { + val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; SELECT 2; END" + val result = SqlStatementSplitter + .splitWithPositions( + s"$block; SELECT 3", + identity, + preserveMalformedCompoundBoundaries = true) + .withoutPositions + + assert(result.completeStatements == Seq(statement(block))) + assert(result.partialStatement == "SELECT 3") + } + test("Valid BEGIN..END block is never split at internal ;") { // A valid `BEGIN ... END` block must be confirmed in full -- the splitter // must never emit at an internal `;`, even if a shorter prefix happens to diff --git a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala index d8e60e3c038ab..bcc3d6d8eee76 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala @@ -135,6 +135,16 @@ class ParseSqlResultSuite extends SparkFunSuite { assert(statements.map(_ \ "parse_success") === Seq(JBool(false), JBool(true))) } + test("malformed nested control END does not end the outer SQL script") { + val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; SELECT 2; END" + val statements = objs(s"$block; SELECT 3") + + assert(statements.size === 2) + assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(block.length + 3))) + assert(statements.map(_ \ "length") === Seq(JInt(block.length), JInt(8))) + assert(statements.map(_ \ "parse_success") === Seq(JBool(false), JBool(true))) + } + test("SQL scripts remain one statement in a batch") { val statements = objs("BEGIN SELECT 1; SELECT 2; END; SELECT 3") assert(statements.size === 2) From df8249674c46e513d488cd6e46d8ba00a9f96e5a Mon Sep 17 00:00:00 2001 From: srielau Date: Wed, 9 Sep 2026 13:54:18 +0000 Subject: [PATCH 07/12] [SPARK-59255][SQL] Require recovered END to match the compound suffix --- .../parser/SqlStatementSplitter.scala | 79 ++++++++++++++++--- .../parser/SqlStatementSplitterSuite.scala | 13 +++ .../catalyst/parser/ParseSqlResultSuite.scala | 10 +++ 3 files changed, 89 insertions(+), 13 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala index 84121ff2c9a55..21ba45b14e4e7 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala @@ -434,27 +434,80 @@ object SqlStatementSplitter { } /** - * Returns true only when recovery found a real END and the candidate itself - * ends in a real END, the optional outer semicolon, and EOF. Error recovery - * can bind the context's END node to an inner control terminator and skip its - * suffix, so the candidate suffix is checked independently. + * Returns true only when a real recovered END is the candidate's trailing + * END, or when recovery bound that node to an earlier control terminator + * (`END IF`, `END WHILE`, ...) and a later statement-level END is the + * actual suffix. A statement such as `SELECT END` can end in an END token + * that is not the outer compound terminator. */ private def isOuterCompoundEnd(tokens: CommonTokenStream, recoveredEnd: Token): Boolean = { if (recoveredEnd.getTokenIndex < 0) return false + val suffixEnd = trailingEndToken(tokens) + if (suffixEnd == null) return false + val recoveredIndex = recoveredEnd.getTokenIndex + val suffixIndex = suffixEnd.getTokenIndex + recoveredIndex == suffixIndex || + (recoveredIndex < suffixIndex && + isControlTerminatorEnd(tokens, recoveredEnd) && + isStatementLevelEnd(tokens, suffixEnd)) + } + + /** + * Returns the candidate's trailing END when the default-channel suffix is + * END, an optional semicolon, and EOF. Otherwise null. + */ + private def trailingEndToken(tokens: CommonTokenStream): Token = { + var index = skipHiddenLeft(tokens, tokens.size() - 1) + if (index < 0 || tokens.get(index).getType != Token.EOF) return null + index = skipHiddenLeft(tokens, index - 1) + if (index >= 0 && tokens.get(index).getType == SqlBaseLexer.SEMICOLON) { + index = skipHiddenLeft(tokens, index - 1) + } + if (index >= 0 && tokens.get(index).getType == SqlBaseLexer.END) { + tokens.get(index) + } else { + null + } + } + + /** True when `end` is the END of END IF / WHILE / LOOP / REPEAT / FOR / CASE. */ + private def isControlTerminatorEnd(tokens: CommonTokenStream, end: Token): Boolean = { + val next = skipHiddenRight(tokens, end.getTokenIndex + 1) + next >= 0 && (tokens.get(next).getType match { + case SqlBaseLexer.IF | SqlBaseLexer.WHILE | SqlBaseLexer.LOOP | + SqlBaseLexer.REPEAT | SqlBaseLexer.FOR | SqlBaseLexer.CASE => true + case _ => false + }) + } - var index = tokens.size() - 1 - if (index < 0 || tokens.get(index).getType != Token.EOF) return false - index -= 1 + /** + * True when `end` starts a compound closer rather than an identifier in a + * statement such as `SELECT END`. The previous default-channel token must be + * BEGIN, ATOMIC, or a semicolon. + */ + private def isStatementLevelEnd(tokens: CommonTokenStream, end: Token): Boolean = { + val prev = skipHiddenLeft(tokens, end.getTokenIndex - 1) + prev >= 0 && (tokens.get(prev).getType match { + case SqlBaseLexer.BEGIN | SqlBaseLexer.ATOMIC | SqlBaseLexer.SEMICOLON => true + case _ => false + }) + } + + private def skipHiddenLeft(tokens: CommonTokenStream, from: Int): Int = { + var index = from while (index >= 0 && tokens.get(index).getChannel == Token.HIDDEN_CHANNEL) { index -= 1 } - if (index >= 0 && tokens.get(index).getType == SqlBaseLexer.SEMICOLON) { - index -= 1 - while (index >= 0 && tokens.get(index).getChannel == Token.HIDDEN_CHANNEL) { - index -= 1 - } + index + } + + private def skipHiddenRight(tokens: CommonTokenStream, from: Int): Int = { + var index = from + while (index < tokens.size() && + tokens.get(index).getChannel == Token.HIDDEN_CHANNEL) { + index += 1 } - index >= 0 && tokens.get(index).getType == SqlBaseLexer.END + if (index < tokens.size()) index else -1 } /** Outcome of attempting to parse one statement candidate. */ diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala index a27309df78937..36e19f4a6be95 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala @@ -552,6 +552,19 @@ class SqlStatementSplitterSuite extends SparkFunSuite { assert(result.partialStatement == "SELECT 3") } + test("statement-final END before the outer END does not end the block") { + val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; SELECT END; END" + val result = SqlStatementSplitter + .splitWithPositions( + s"$block; SELECT 3", + identity, + preserveMalformedCompoundBoundaries = true) + .withoutPositions + + assert(result.completeStatements == Seq(statement(block))) + assert(result.partialStatement == "SELECT 3") + } + test("Valid BEGIN..END block is never split at internal ;") { // A valid `BEGIN ... END` block must be confirmed in full -- the splitter // must never emit at an internal `;`, even if a shorter prefix happens to diff --git a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala index bcc3d6d8eee76..e91e3e2d8cbab 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala @@ -145,6 +145,16 @@ class ParseSqlResultSuite extends SparkFunSuite { assert(statements.map(_ \ "parse_success") === Seq(JBool(false), JBool(true))) } + test("statement-final END before the outer END does not end the SQL script") { + val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; SELECT END; END" + val statements = objs(s"$block; SELECT 3") + + assert(statements.size === 2) + assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(block.length + 3))) + assert(statements.map(_ \ "length") === Seq(JInt(block.length), JInt(8))) + assert(statements.map(_ \ "parse_success") === Seq(JBool(false), JBool(true))) + } + test("SQL scripts remain one statement in a batch") { val statements = objs("BEGIN SELECT 1; SELECT 2; END; SELECT 3") assert(statements.size === 2) From f0e7b73896cd6b3e96423151706ea307158c3b85 Mon Sep 17 00:00:00 2001 From: srielau Date: Wed, 9 Sep 2026 14:45:07 +0000 Subject: [PATCH 08/12] fix: [SPARK-59255] track nested compound depth during recovery --- .../parser/SqlStatementSplitter.scala | 39 ++++++++++++------- .../parser/SqlStatementSplitterSuite.scala | 13 +++++++ .../catalyst/parser/ParseSqlResultSuite.scala | 10 +++++ 3 files changed, 48 insertions(+), 14 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala index 21ba45b14e4e7..93c93cd387ad5 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala @@ -433,23 +433,34 @@ object SqlStatementSplitter { None } - /** - * Returns true only when a real recovered END is the candidate's trailing - * END, or when recovery bound that node to an earlier control terminator - * (`END IF`, `END WHILE`, ...) and a later statement-level END is the - * actual suffix. A statement such as `SELECT END` can end in an END token - * that is not the outer compound terminator. - */ + /** Returns true only when a real recovered END closes the candidate's outer BEGIN. */ private def isOuterCompoundEnd(tokens: CommonTokenStream, recoveredEnd: Token): Boolean = { if (recoveredEnd.getTokenIndex < 0) return false val suffixEnd = trailingEndToken(tokens) - if (suffixEnd == null) return false - val recoveredIndex = recoveredEnd.getTokenIndex - val suffixIndex = suffixEnd.getTokenIndex - recoveredIndex == suffixIndex || - (recoveredIndex < suffixIndex && - isControlTerminatorEnd(tokens, recoveredEnd) && - isStatementLevelEnd(tokens, suffixEnd)) + suffixEnd != null && closesOuterBegin(tokens, suffixEnd) + } + + private def closesOuterBegin(tokens: CommonTokenStream, suffixEnd: Token): Boolean = { + var depth = 0 + var index = 0 + val limit = suffixEnd.getTokenIndex + while (index <= limit) { + val token = tokens.get(index) + if (token.getChannel != Token.HIDDEN_CHANNEL) { + token.getType match { + case SqlBaseLexer.BEGIN => + depth += 1 + case SqlBaseLexer.END + if isStatementLevelEnd(tokens, token) && + !isControlTerminatorEnd(tokens, token) => + depth -= 1 + if (depth < 0) return false + case _ => + } + } + index += 1 + } + depth == 0 } /** diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala index 36e19f4a6be95..ac4ab1f157550 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala @@ -565,6 +565,19 @@ class SqlStatementSplitterSuite extends SparkFunSuite { assert(result.partialStatement == "SELECT 3") } + test("nested compound END does not end the outer malformed block") { + val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; BEGIN SELECT 2; END; END" + val result = SqlStatementSplitter + .splitWithPositions( + s"$block; SELECT 3", + identity, + preserveMalformedCompoundBoundaries = true) + .withoutPositions + + assert(result.completeStatements == Seq(statement(block))) + assert(result.partialStatement == "SELECT 3") + } + test("Valid BEGIN..END block is never split at internal ;") { // A valid `BEGIN ... END` block must be confirmed in full -- the splitter // must never emit at an internal `;`, even if a shorter prefix happens to diff --git a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala index e91e3e2d8cbab..97d054626fb70 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala @@ -155,6 +155,16 @@ class ParseSqlResultSuite extends SparkFunSuite { assert(statements.map(_ \ "parse_success") === Seq(JBool(false), JBool(true))) } + test("nested compound END does not end the outer malformed SQL script") { + val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; BEGIN SELECT 2; END; END" + val statements = objs(s"$block; SELECT 3") + + assert(statements.size === 2) + assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(block.length + 3))) + assert(statements.map(_ \ "length") === Seq(JInt(block.length), JInt(8))) + assert(statements.map(_ \ "parse_success") === Seq(JBool(false), JBool(true))) + } + test("SQL scripts remain one statement in a batch") { val statements = objs("BEGIN SELECT 1; SELECT 2; END; SELECT 3") assert(statements.size === 2) From 95b1cf5ebce6cf04e9da19c091671e992bd1f115 Mon Sep 17 00:00:00 2001 From: srielau Date: Wed, 9 Sep 2026 21:13:15 +0000 Subject: [PATCH 09/12] [SPARK-59255][SQL] Parse SQL batch boundaries in one grammar pass --- .../sql/catalyst/parser/SqlBaseParser.g4 | 93 ++++++ .../parser/SqlStatementSplitter.scala | 293 +++++++----------- .../parser/SqlStatementSplitterSuite.scala | 126 ++++++-- .../sql/catalyst/expressions/ParseSql.scala | 4 +- .../spark/sql/execution/SparkSqlParser.scala | 5 +- 5 files changed, 301 insertions(+), 220 deletions(-) diff --git a/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 b/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 index 7865f08423254..d5a9a53b28980 100644 --- a/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 +++ b/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 @@ -18,6 +18,8 @@ parser grammar SqlBaseParser; options { tokenVocab = SqlBaseLexer; } +tokens { PARSE_SQL_BATCH_DELIMITER } + @members { /** * When false, INTERSECT is given the greater precedence over the other set @@ -87,6 +89,97 @@ compoundOrSingleStatement | singleCompoundStatement ; +// Boundary-only grammar for parse_sql batches. Leaf statements deliberately accept arbitrary +// tokens: ParseSqlResult parses each emitted segment with the full grammar and records any error. +// BEGIN is excluded from the terminated fallback, so only a grammar context can own the +// semicolons inside a compound statement. BEGIN and END remain unrestricted inside leaf +// statements. The caller appends PARSE_SQL_BATCH_DELIMITER; the trailing BEGIN fallback consumes +// a structurally unclosed compound through that token without synthesizing an END token. +parseSqlBatch + : SEMICOLON* (items+=parseSqlBatchItem SEMICOLON*)* + PARSE_SQL_BATCH_DELIMITER? EOF + ; + +parseSqlBatchItem + : batchStatement=parseSqlBatchStatement + terminator=(SEMICOLON | PARSE_SQL_BATCH_DELIMITER) + | partialStatement=parseSqlBatchPartialCompoundStatement + terminator=PARSE_SQL_BATCH_DELIMITER + ; + +parseSqlBatchStatement + : parseSqlBatchCompoundStatement + | parseSqlBatchLeafStatement + ; + +parseSqlBatchPartialCompoundStatement + : BEGIN .*? + ; + +parseSqlBatchCompoundStatement + : BEGIN (NOT ATOMIC)? parseSqlBatchCompoundBody? END + ; + +parseSqlBatchBeginEndCompoundBlock + : beginLabel? BEGIN (NOT ATOMIC)? parseSqlBatchCompoundBody? END endLabel? + ; + +parseSqlBatchCompoundBody + : (parseSqlBatchCompoundBodyStatement SEMICOLON)+ + ; + +parseSqlBatchCompoundBodyStatement + : parseSqlBatchBeginEndCompoundBlock + | parseSqlBatchDeclareHandlerStatement + | parseSqlBatchIfElseStatement + | parseSqlBatchCaseStatement + | parseSqlBatchWhileStatement + | parseSqlBatchRepeatStatement + | parseSqlBatchLoopStatement + | parseSqlBatchForStatement + | parseSqlBatchLeafStatement + ; + +parseSqlBatchDeclareHandlerStatement + : DECLARE (CONTINUE | EXIT) HANDLER FOR conditionValues + (parseSqlBatchBeginEndCompoundBlock | parseSqlBatchLeafStatement) + ; + +parseSqlBatchWhileStatement + : beginLabel? WHILE booleanExpression DO parseSqlBatchCompoundBody END WHILE endLabel? + ; + +parseSqlBatchIfElseStatement + : IF booleanExpression THEN parseSqlBatchCompoundBody + (ELSEIF booleanExpression THEN parseSqlBatchCompoundBody)* + (ELSE parseSqlBatchCompoundBody)? END IF + ; + +parseSqlBatchRepeatStatement + : beginLabel? REPEAT parseSqlBatchCompoundBody UNTIL booleanExpression END REPEAT endLabel? + ; + +parseSqlBatchCaseStatement + : CASE (WHEN booleanExpression THEN parseSqlBatchCompoundBody)+ + (ELSE parseSqlBatchCompoundBody)? END CASE + | CASE expression (WHEN expression THEN parseSqlBatchCompoundBody)+ + (ELSE parseSqlBatchCompoundBody)? END CASE + ; + +parseSqlBatchLoopStatement + : beginLabel? LOOP parseSqlBatchCompoundBody END LOOP endLabel? + ; + +parseSqlBatchForStatement + : beginLabel? FOR (strictIdentifier AS)? query DO + parseSqlBatchCompoundBody END FOR endLabel? + ; + +parseSqlBatchLeafStatement + : {_input.LA(1) != BEGIN}? + (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ + ; + singleCompoundStatement : BEGIN (NOT ATOMIC)? compoundBody? END SEMICOLON? EOF ; diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala index 93c93cd387ad5..137d9b7d5a444 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala @@ -117,15 +117,16 @@ private[sql] case class PositionedSqlStatementSplitResult( * for real at execution time. When `validationPreprocess` is `identity` * (the default), the splitter behaves as a pure original-text splitter. * - * Performance note: for a single `BEGIN ... END` block with k internal `;`, - * the splitter calls `tryParseRegion` O(k) times on growing prefixes -- an + * Performance note: the generic splitter calls `tryParseRegion` O(k) times on + * growing prefixes for a single `BEGIN ... END` block with k internal `;` -- an * O(k^2) cost in the worst case (incomplete block on every keystroke in * interactive mode). Ordinary non-scripting SQL is O(n). A non-EOF terminated * single-statement rule (read `ctx.getStop` once per region) would make this * O(n), but Spark's `setResetStatement` has `SET .*?` / `RESET .*?` wildcards * that need an EOF anchor to terminate deterministically, so such a * single-statement rule-rewrite does not drop in cleanly. Tracked as a - * follow-up. + * follow-up. The parse_sql-only path uses [[splitForParseSql]] and performs one + * linear boundary parse instead. */ object SqlStatementSplitter { @@ -133,6 +134,100 @@ object SqlStatementSplitter { def split(sqlText: String): SqlStatementSplitResult = splitWithPositions(sqlText, identity).withoutPositions + /** + * Split a parse_sql batch in one grammar-owned pass while retaining source positions. + * Unlike the generic splitter, this boundary-only grammar accepts malformed leaf statements + * and uses scripting grammar contexts to assign internal semicolons to compound statements. + */ + private[sql] def splitForParseSql(sqlText: String): PositionedSqlStatementSplitResult = { + require(sqlText != null, "sqlText must not be null") + + val toUtf16 = utf16Offsets(sqlText) + val sourceLexer = new SqlBaseLexer(new UpperCaseCharStream(CharStreams.fromString(sqlText))) + sourceLexer.removeErrorListeners() + val sourceTokens = new CommonTokenStream(sourceLexer) + sourceTokens.fill() + val boundaryTokens = new java.util.ArrayList[Token](sourceTokens.size() + 1) + var sourceIndex = 0 + while (sourceIndex < sourceTokens.size() - 1) { + boundaryTokens.add(sourceTokens.get(sourceIndex)) + sourceIndex += 1 + } + val boundary = new CommonToken(SqlBaseParser.PARSE_SQL_BATCH_DELIMITER, "") + boundary.setStartIndex(toUtf16.length - 1) + boundary.setStopIndex(toUtf16.length - 2) + boundaryTokens.add(boundary) + boundaryTokens.add(sourceTokens.get(sourceTokens.size() - 1)) + val tokens = new CommonTokenStream(new ListTokenSource(boundaryTokens)) + tokens.fill() + val parser = new SqlBaseParser(tokens) + configureSplitterParser(parser, SqlApiConf.get) + parser.getInterpreter.setPredictionMode(PredictionMode.LL) + val batch = try { + parser.parseSqlBatch() + } catch { + case _: StackOverflowError => + return splitWithPositions(sqlText, identity) + } + + def statementStart(context: ParserRuleContext): Int = { + var index = context.getStart.getTokenIndex + while (index > 0 && tokens.get(index - 1).getChannel == Token.HIDDEN_CHANNEL) { + index -= 1 + } + toUtf16(tokens.get(index).getStartIndex) + } + + def positioned( + context: ParserRuleContext, + end: Int, + terminator: String): PositionedSqlStatement = { + val start = statementStart(context) + val (leadingWhitespace, statement) = trimSqlWhitespace(sqlText.substring(start, end)) + PositionedSqlStatement(statement, terminator, start + leadingWhitespace) + } + + val complete = mutable.ArrayBuffer.empty[PositionedSqlStatement] + var partial: Option[PositionedSqlStatement] = None + var i = 0 + while (i < batch.items.size()) { + val item = batch.items.get(i) + val context = Option(item.batchStatement).getOrElse(item.partialStatement) + if (item.terminator.getType == SqlBaseParser.PARSE_SQL_BATCH_DELIMITER) { + partial = Some(positioned(context, sqlText.length, "")) + } else { + complete += positioned( + context, + toUtf16(item.terminator.getStartIndex), + item.terminator.getText) + } + i += 1 + } + + partial = partial.orElse { + if (sourceLexer.has_unclosed_bracketed_comment) { + var index = sourceTokens.size() - 1 + while (index >= 0 && sourceTokens.get(index).getType != SqlBaseLexer.SEMICOLON) { + index -= 1 + } + val start = if (index < 0) { + 0 + } else { + toUtf16(sourceTokens.get(index).getStopIndex + 1) + } + val (leadingWhitespace, statement) = trimSqlWhitespace(sqlText.substring(start)) + Some(PositionedSqlStatement(statement, "", start + leadingWhitespace)) + } else { + None + } + } + + PositionedSqlStatementSplitResult( + complete.toSeq, + partial, + sourceLexer.has_unclosed_bracketed_comment && partial.nonEmpty) + } + /** * Split the given SQL text, applying `validationPreprocess` to each candidate * region before parser validation. The emitted [[SqlStatement]] text is @@ -147,15 +242,11 @@ object SqlStatementSplitter { 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. Such tooling - * can preserve a balanced BEGIN ... END boundary even when its body is malformed; - * generic splitting leaves that disabled to retain extension fallback behavior. + * Split SQL while retaining each statement's source position. */ private[sql] def splitWithPositions( sqlText: String, - validationPreprocess: String => String, - preserveMalformedCompoundBoundaries: Boolean = false): PositionedSqlStatementSplitResult = { + validationPreprocess: String => String): PositionedSqlStatementSplitResult = { require(sqlText != null, "sqlText must not be null") require(validationPreprocess != null, "validationPreprocess must not be null") @@ -267,31 +358,8 @@ object SqlStatementSplitter { // prefix that includes the next `;`. d += 1 case FailedNonEof => - // A malformed but balanced compound body still belongs to its enclosing - // BEGIN ... END statement. Use error-recovering parsing to find that real outer - // boundary before falling back to ordinary delimiter splitting. - val recovered = if (preserveMalformedCompoundBoundaries) { - findMalformedCompoundEnd( - sqlText, - toUtf16, - tokenStream, - startIdx, - delimiterPositions, - d, - validationPreprocess, - conf) - } else { - None - } - recovered match { - case Some((recoveredDelimiter, recoveredEnd)) => - parsedOk = true - d = recoveredDelimiter - matchedDelimIdx = recoveredEnd - case None => - // Structurally invalid ordinary statement; stop extending. - failedNonEof = true - } + // Structurally invalid ordinary statement; stop extending. + failedNonEof = true } } @@ -377,150 +445,6 @@ object SqlStatementSplitter { unclosed && partial.nonEmpty) } - /** - * Returns the delimiter-array index and ending token index of a real outer END for a malformed - * compound statement. Error recovery may repair the body, but a missing END is synthetic and - * has token index -1. - */ - private def findMalformedCompoundEnd( - sqlText: String, - toUtf16: Array[Int], - stream: CommonTokenStream, - startIdx: Int, - delimiterPositions: Array[Int], - fromDelimiter: Int, - validationPreprocess: String => String, - conf: SqlApiConf): Option[(Int, Int)] = { - if (stream.get(startIdx).getType != SqlBaseLexer.BEGIN) { - return None - } - - var delimiter = fromDelimiter - while (delimiter <= delimiterPositions.length) { - val endIdx = if (delimiter < delimiterPositions.length) { - delimiterPositions(delimiter) - } else { - stream.size() - 1 - } - val firstTok = stream.get(startIdx) - val lastTok = stream.get(endIdx) - val regionStart = toUtf16(firstTok.getStartIndex) - val regionEnd = if (lastTok.getType == Token.EOF) { - sqlText.length - } else { - toUtf16(lastTok.getStopIndex + 1) - } - val candidate = validationPreprocess(sqlText.substring(regionStart, regionEnd)) - val lexer = new SqlBaseLexer( - new UpperCaseCharStream(CharStreams.fromString(candidate))) - lexer.removeErrorListeners() - val tokens = new CommonTokenStream(lexer) - tokens.fill() - val parser = new SqlBaseParser(tokens) - configureSplitterParser(parser, conf, bailOnError = false) - parser.getInterpreter.setPredictionMode(PredictionMode.LL) - try { - val context = parser.singleCompoundStatement() - val end = context.END() - if (end != null && isOuterCompoundEnd(tokens, end.getSymbol)) { - return Some((delimiter, endIdx)) - } - } catch { - case _: StackOverflowError => return None - } - delimiter += 1 - } - None - } - - /** Returns true only when a real recovered END closes the candidate's outer BEGIN. */ - private def isOuterCompoundEnd(tokens: CommonTokenStream, recoveredEnd: Token): Boolean = { - if (recoveredEnd.getTokenIndex < 0) return false - val suffixEnd = trailingEndToken(tokens) - suffixEnd != null && closesOuterBegin(tokens, suffixEnd) - } - - private def closesOuterBegin(tokens: CommonTokenStream, suffixEnd: Token): Boolean = { - var depth = 0 - var index = 0 - val limit = suffixEnd.getTokenIndex - while (index <= limit) { - val token = tokens.get(index) - if (token.getChannel != Token.HIDDEN_CHANNEL) { - token.getType match { - case SqlBaseLexer.BEGIN => - depth += 1 - case SqlBaseLexer.END - if isStatementLevelEnd(tokens, token) && - !isControlTerminatorEnd(tokens, token) => - depth -= 1 - if (depth < 0) return false - case _ => - } - } - index += 1 - } - depth == 0 - } - - /** - * Returns the candidate's trailing END when the default-channel suffix is - * END, an optional semicolon, and EOF. Otherwise null. - */ - private def trailingEndToken(tokens: CommonTokenStream): Token = { - var index = skipHiddenLeft(tokens, tokens.size() - 1) - if (index < 0 || tokens.get(index).getType != Token.EOF) return null - index = skipHiddenLeft(tokens, index - 1) - if (index >= 0 && tokens.get(index).getType == SqlBaseLexer.SEMICOLON) { - index = skipHiddenLeft(tokens, index - 1) - } - if (index >= 0 && tokens.get(index).getType == SqlBaseLexer.END) { - tokens.get(index) - } else { - null - } - } - - /** True when `end` is the END of END IF / WHILE / LOOP / REPEAT / FOR / CASE. */ - private def isControlTerminatorEnd(tokens: CommonTokenStream, end: Token): Boolean = { - val next = skipHiddenRight(tokens, end.getTokenIndex + 1) - next >= 0 && (tokens.get(next).getType match { - case SqlBaseLexer.IF | SqlBaseLexer.WHILE | SqlBaseLexer.LOOP | - SqlBaseLexer.REPEAT | SqlBaseLexer.FOR | SqlBaseLexer.CASE => true - case _ => false - }) - } - - /** - * True when `end` starts a compound closer rather than an identifier in a - * statement such as `SELECT END`. The previous default-channel token must be - * BEGIN, ATOMIC, or a semicolon. - */ - private def isStatementLevelEnd(tokens: CommonTokenStream, end: Token): Boolean = { - val prev = skipHiddenLeft(tokens, end.getTokenIndex - 1) - prev >= 0 && (tokens.get(prev).getType match { - case SqlBaseLexer.BEGIN | SqlBaseLexer.ATOMIC | SqlBaseLexer.SEMICOLON => true - case _ => false - }) - } - - private def skipHiddenLeft(tokens: CommonTokenStream, from: Int): Int = { - var index = from - while (index >= 0 && tokens.get(index).getChannel == Token.HIDDEN_CHANNEL) { - index -= 1 - } - index - } - - private def skipHiddenRight(tokens: CommonTokenStream, from: Int): Int = { - var index = from - while (index < tokens.size() && - tokens.get(index).getChannel == Token.HIDDEN_CHANNEL) { - index += 1 - } - if (index < tokens.size()) index else -1 - } - /** Outcome of attempting to parse one statement candidate. */ private sealed trait ParseOutcome private case object ParsedOk extends ParseOutcome @@ -612,10 +536,8 @@ object SqlStatementSplitter { * Configure a fresh [[SqlBaseParser]] for splitter use: install the managed * caches so candidate parses share ANTLR DFA state across calls, apply the * session's behavior flags so the splitter agrees with the session parser - * on grammar interpretation (e.g. `double_quoted_identifiers`), and, when - * `bailOnError` is true, install a bail error strategy so failures throw - * immediately. Malformed-compound recovery keeps the default error strategy - * so it can inspect the recovered outer END boundary. + * on grammar interpretation (e.g. `double_quoted_identifiers`), and install + * a bail error strategy so failures throw immediately. * * Notably, the splitter does NOT install [[PostProcessor]] or * [[UnclosedCommentProcessor]] -- the former mutates the parse tree (which @@ -625,8 +547,7 @@ object SqlStatementSplitter { */ private def configureSplitterParser( parser: SqlBaseParser, - conf: SqlApiConf, - bailOnError: Boolean = true): Unit = { + conf: SqlApiConf): Unit = { if (conf.manageParserCaches) AbstractParser.installCaches(parser) parser.legacy_setops_precedence_enabled = conf.setOpsPrecedenceEnforced @@ -638,9 +559,7 @@ object SqlStatementSplitter { parser.single_character_pipe_operator_enabled = conf.singleCharacterPipeOperatorEnabled parser.removeErrorListeners() - if (bailOnError) { - parser.setErrorHandler(new BailErrorStrategy) - } + parser.setErrorHandler(new BailErrorStrategy) } /** diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala index ac4ab1f157550..d42f744db84d9 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala @@ -18,6 +18,7 @@ package org.apache.spark.sql.catalyst.parser import org.apache.spark.SparkFunSuite +import org.apache.spark.sql.internal.SQLConf class SqlStatementSplitterSuite extends SparkFunSuite { @@ -518,34 +519,22 @@ class SqlStatementSplitterSuite extends SparkFunSuite { test("malformed balanced BEGIN block remains one statement") { val block = "BEGIN SELECT 1; SELEC 2; END" - val result = SqlStatementSplitter - .splitWithPositions( - s"$block; SELECT 3;", - identity, - preserveMalformedCompoundBoundaries = true) + val result = SqlStatementSplitter.splitForParseSql(s"$block; SELECT 3;") .withoutPositions assert(result.completeStatements == Seq( statement(block), statement("SELECT 3"))) assert(result.partialStatement.isEmpty) - val withoutTerminator = SqlStatementSplitter - .splitWithPositions( - block, - identity, - preserveMalformedCompoundBoundaries = true) + val withoutTerminator = SqlStatementSplitter.splitForParseSql(block) .withoutPositions - assert(withoutTerminator.completeStatements == Seq(SqlStatement(block, ""))) - assert(withoutTerminator.partialStatement.isEmpty) + assert(withoutTerminator.completeStatements.isEmpty) + assert(withoutTerminator.partialStatement == block) } test("malformed nested control END does not end the outer block") { val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; SELECT 2; END" - val result = SqlStatementSplitter - .splitWithPositions( - s"$block; SELECT 3", - identity, - preserveMalformedCompoundBoundaries = true) + val result = SqlStatementSplitter.splitForParseSql(s"$block; SELECT 3") .withoutPositions assert(result.completeStatements == Seq(statement(block))) @@ -554,11 +543,7 @@ class SqlStatementSplitterSuite extends SparkFunSuite { test("statement-final END before the outer END does not end the block") { val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; SELECT END; END" - val result = SqlStatementSplitter - .splitWithPositions( - s"$block; SELECT 3", - identity, - preserveMalformedCompoundBoundaries = true) + val result = SqlStatementSplitter.splitForParseSql(s"$block; SELECT 3") .withoutPositions assert(result.completeStatements == Seq(statement(block))) @@ -567,17 +552,104 @@ class SqlStatementSplitterSuite extends SparkFunSuite { test("nested compound END does not end the outer malformed block") { val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; BEGIN SELECT 2; END; END" - val result = SqlStatementSplitter - .splitWithPositions( - s"$block; SELECT 3", - identity, - preserveMalformedCompoundBoundaries = true) + val result = SqlStatementSplitter.splitForParseSql(s"$block; SELECT 3") .withoutPositions assert(result.completeStatements == Seq(statement(block))) assert(result.partialStatement == "SELECT 3") } + test("parse_sql batch boundary leaves an unclosed compound as partial") { + val partial = "BEGIN SELECT 1; SELECT 2;" + val result = SqlStatementSplitter + .splitForParseSql(s"SELECT 0; $partial") + .withoutPositions + + assert(result.completeStatements == Seq(statement("SELECT 0"))) + assert(result.partialStatement == partial) + } + + test("parse_sql batch boundary preserves unclosed comments") { + val result = SqlStatementSplitter + .splitForParseSql("SELECT 1; /* unclosed") + .withoutPositions + assert(result.completeStatements == Seq(statement("SELECT 1"))) + assert(result.partialStatement == "/* unclosed") + assert(result.hasUnclosedComment) + + val commentOnly = SqlStatementSplitter + .splitForParseSql("/* unclosed") + .withoutPositions + assert(commentOnly.completeStatements.isEmpty) + assert(commentOnly.partialStatement == "/* unclosed") + assert(commentOnly.hasUnclosedComment) + } + + test("parse_sql batch boundary follows scripting grammar contexts") { + val blocks = Seq( + "BEGIN IF TRUE THEN SELEC 1; END IF; END", + "BEGIN CASE WHEN TRUE THEN SELEC 1; END CASE; END", + "BEGIN WHILE TRUE DO SELEC 1; END WHILE; END", + "BEGIN REPEAT SELEC 1; UNTIL TRUE END REPEAT; END", + "BEGIN LOOP SELEC 1; END LOOP; END", + "BEGIN FOR x AS SELECT 1 DO SELEC 1; END FOR; END", + "BEGIN lbl: BEGIN SELEC 1; END lbl; END", + "BEGIN IF ${flag} THEN BEGIN SELEC 1; END; SELECT 1; END IF; END", + "BEGIN DECLARE CONTINUE HANDLER FOR SQLEXCEPTION BEGIN SELEC 1; END; END", + "BEGIN DECLARE EXIT HANDLER FOR SQLEXCEPTION BEGIN SELEC 1; END; SELECT 1; END") + + blocks.foreach { block => + val result = SqlStatementSplitter + .splitForParseSql(s"$block; SELECT 9") + .withoutPositions + assert(result.completeStatements == Seq(statement(block)), block) + assert(result.partialStatement == "SELECT 9", block) + + val trailing = SqlStatementSplitter.splitForParseSql(block).withoutPositions + assert(trailing.completeStatements.isEmpty, block) + assert(trailing.partialStatement == block, block) + } + } + + test("parse_sql batch boundary tolerates malformed control headers") { + val blocks = Seq( + "BEGIN IF TRUE THEN SELECT 1; END; END", + "BEGIN FOR x AS SELEC 1 DO SELECT 1; END FOR; END") + + blocks.foreach { block => + val result = SqlStatementSplitter + .splitForParseSql(s"$block; SELECT 9") + .withoutPositions + assert(result.completeStatements == Seq(statement(block)), block) + assert(result.partialStatement == "SELECT 9", block) + } + } + + test("parse_sql batch boundary handles SET, RESET, and non-reserved identifiers") { + Seq(false, true).foreach { ansi => + SQLConf.withExistingConf(new SQLConf) { + SQLConf.get.setConf(SQLConf.ANSI_ENABLED, ansi) + SQLConf.get.setConf(SQLConf.ENFORCE_RESERVED_KEYWORDS, ansi) + val result = SqlStatementSplitter + .splitForParseSql( + "SET spark.sql.ansi.enabled=true; RESET spark.sql.ansi.enabled; " + + "SELECT begin, end FROM t") + .withoutPositions + + assert(result.completeStatements == Seq( + statement("SET spark.sql.ansi.enabled=true"), + statement("RESET spark.sql.ansi.enabled"))) + assert(result.partialStatement == "SELECT begin, end FROM t") + } + } + + val malformed = SqlStatementSplitter + .splitForParseSql("SELECT FROM; SELECT 2") + .withoutPositions + assert(malformed.completeStatements == Seq(statement("SELECT FROM"))) + assert(malformed.partialStatement == "SELECT 2") + } + test("Valid BEGIN..END block is never split at internal ;") { // A valid `BEGIN ... END` block must be confirmed in full -- the splitter // must never emit at an internal `;`, even if a shorter prefix happens to diff --git a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala index 9bffb3ffb518a..0f165cb0ac9e2 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala @@ -29,8 +29,8 @@ import org.apache.spark.unsafe.types.UTF8String /** * 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. + * select-list names, parameters). A statement that does not parse produces a + * result object containing a nested STANDARD-format error object. * * Behind [[SQLConf.PARSE_SQL_ENABLED]] while the JSON contract is still * evolving. Designed for batch evaluation over DataFrames of SQL text. diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala index 3ad8b28c4e365..84c0ffb1ce890 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala @@ -126,10 +126,7 @@ class SparkSqlParser extends AbstractSqlParser { /** Split statements while retaining their positions in the original SQL text. */ private[sql] def splitStatementsWithPositions( sqlText: String): PositionedSqlStatementSplitResult = - SqlStatementSplitter.splitWithPositions( - sqlText, - SparkSqlParser.substituteVariablesForValidation, - preserveMalformedCompoundBoundaries = true) + SqlStatementSplitter.splitForParseSql(sqlText) /** * Internal parse method that handles both parameter substitution and regular parsing. From e7f6b1a09a3fb04065fcbfaf6910280ae9c845db Mon Sep 17 00:00:00 2001 From: srielau Date: Thu, 10 Sep 2026 14:35:17 +0000 Subject: [PATCH 10/12] [SPARK-59255][SQL] Preserve malformed batch statement boundaries --- .../sql/catalyst/parser/SqlBaseParser.g4 | 110 +++++++++++++++--- .../parser/SqlStatementSplitterSuite.scala | 78 ++++++++++++- .../sql/catalyst/expressions/ParseSql.scala | 2 +- .../sql/catalyst/parser/ParseSqlResult.scala | 6 +- .../catalyst/parser/ParseSqlResultSuite.scala | 54 ++++++++- 5 files changed, 230 insertions(+), 20 deletions(-) diff --git a/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 b/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 index d5a9a53b28980..1ed63b15f76f3 100644 --- a/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 +++ b/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 @@ -109,9 +109,15 @@ parseSqlBatchItem parseSqlBatchStatement : parseSqlBatchCompoundStatement + | parseSqlBatchMalformedEmptyCompoundBlock + | parseSqlBatchMalformedBeginStatement | parseSqlBatchLeafStatement ; +parseSqlBatchMalformedBeginStatement + : BEGIN + ; + parseSqlBatchPartialCompoundStatement : BEGIN .*? ; @@ -124,12 +130,27 @@ parseSqlBatchBeginEndCompoundBlock : beginLabel? BEGIN (NOT ATOMIC)? parseSqlBatchCompoundBody? END endLabel? ; +parseSqlBatchMalformedEmptyCompoundBlock + : beginLabel? BEGIN (NOT ATOMIC)? SEMICOLON END endLabel? + ; + +parseSqlBatchMalformedBodyBeginStatement + : BEGIN (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ + ; + parseSqlBatchCompoundBody : (parseSqlBatchCompoundBodyStatement SEMICOLON)+ ; parseSqlBatchCompoundBodyStatement + : parseSqlBatchNestedStatement + | parseSqlBatchOrphanControlEndStatement + | parseSqlBatchBodyLeafStatement + ; + +parseSqlBatchNestedStatement : parseSqlBatchBeginEndCompoundBlock + | parseSqlBatchMalformedEmptyCompoundBlock | parseSqlBatchDeclareHandlerStatement | parseSqlBatchIfElseStatement | parseSqlBatchCaseStatement @@ -137,42 +158,100 @@ parseSqlBatchCompoundBodyStatement | parseSqlBatchRepeatStatement | parseSqlBatchLoopStatement | parseSqlBatchForStatement - | parseSqlBatchLeafStatement + | parseSqlBatchMalformedBodyBeginStatement + | parseSqlBatchMalformedBeginStatement + ; + +parseSqlBatchOrphanControlEndStatement + : END (IF | WHILE | LOOP | REPEAT | FOR | CASE) ; parseSqlBatchDeclareHandlerStatement : DECLARE (CONTINUE | EXIT) HANDLER FOR conditionValues - (parseSqlBatchBeginEndCompoundBlock | parseSqlBatchLeafStatement) + (parseSqlBatchBeginEndCompoundBlock + | parseSqlBatchMalformedEmptyCompoundBlock + | parseSqlBatchMalformedBodyBeginStatement + | parseSqlBatchBodyLeafStatement) ; parseSqlBatchWhileStatement - : beginLabel? WHILE booleanExpression DO parseSqlBatchCompoundBody END WHILE endLabel? + : beginLabel? WHILE booleanExpression DO parseSqlBatchCompoundBody + parseSqlBatchPrematureEnds? END WHILE endLabel? ; parseSqlBatchIfElseStatement - : IF booleanExpression THEN parseSqlBatchCompoundBody - (ELSEIF booleanExpression THEN parseSqlBatchCompoundBody)* - (ELSE parseSqlBatchCompoundBody)? END IF + : IF parseSqlBatchIfCondition THEN parseSqlBatchConditionalBody + (ELSEIF parseSqlBatchIfCondition THEN parseSqlBatchConditionalBody)* + (ELSE parseSqlBatchConditionalBody)? parseSqlBatchPrematureEnds? END IF + ; + +parseSqlBatchIfCondition + : (~(THEN | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))* + ; + +parseSqlBatchConditionalBody + : (parseSqlBatchConditionalBodyStatement SEMICOLON)+ + ; + +parseSqlBatchConditionalBodyStatement + : parseSqlBatchNestedStatement + | parseSqlBatchConditionalBodyLeafStatement + ; + +parseSqlBatchConditionalBodyLeafStatement + : {_input.LA(1) != BEGIN && _input.LA(1) != END && + _input.LA(1) != ELSEIF && _input.LA(1) != ELSE}? + (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ ; parseSqlBatchRepeatStatement - : beginLabel? REPEAT parseSqlBatchCompoundBody UNTIL booleanExpression END REPEAT endLabel? + : beginLabel? REPEAT parseSqlBatchCompoundBody parseSqlBatchPrematureEnds? + UNTIL booleanExpression END REPEAT endLabel? ; parseSqlBatchCaseStatement - : CASE (WHEN booleanExpression THEN parseSqlBatchCompoundBody)+ - (ELSE parseSqlBatchCompoundBody)? END CASE - | CASE expression (WHEN expression THEN parseSqlBatchCompoundBody)+ - (ELSE parseSqlBatchCompoundBody)? END CASE + : CASE (WHEN parseSqlBatchCaseWhenCondition THEN parseSqlBatchCaseBody)+ + (ELSE parseSqlBatchCaseBody)? parseSqlBatchPrematureEnds? END CASE + | CASE parseSqlBatchCaseOperand + (WHEN parseSqlBatchCaseWhenCondition THEN parseSqlBatchCaseBody)+ + (ELSE parseSqlBatchCaseBody)? parseSqlBatchPrematureEnds? END CASE + ; + +parseSqlBatchCaseOperand + : (~(WHEN | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ + ; + +parseSqlBatchCaseWhenCondition + : (~(THEN | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))* + ; + +parseSqlBatchCaseBody + : (parseSqlBatchCaseBodyStatement SEMICOLON)+ + ; + +parseSqlBatchCaseBodyStatement + : parseSqlBatchNestedStatement + | parseSqlBatchCaseBodyLeafStatement + ; + +parseSqlBatchCaseBodyLeafStatement + : {_input.LA(1) != BEGIN && _input.LA(1) != END && + _input.LA(1) != WHEN && _input.LA(1) != ELSE}? + (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ ; parseSqlBatchLoopStatement - : beginLabel? LOOP parseSqlBatchCompoundBody END LOOP endLabel? + : beginLabel? LOOP parseSqlBatchCompoundBody + parseSqlBatchPrematureEnds? END LOOP endLabel? ; parseSqlBatchForStatement : beginLabel? FOR (strictIdentifier AS)? query DO - parseSqlBatchCompoundBody END FOR endLabel? + parseSqlBatchCompoundBody parseSqlBatchPrematureEnds? END FOR endLabel? + ; + +parseSqlBatchPrematureEnds + : (END SEMICOLON)+ ; parseSqlBatchLeafStatement @@ -180,6 +259,11 @@ parseSqlBatchLeafStatement (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ ; +parseSqlBatchBodyLeafStatement + : {_input.LA(1) != BEGIN && _input.LA(1) != END}? + (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ + ; + singleCompoundStatement : BEGIN (NOT ATOMIC)? compoundBody? END SEMICOLON? EOF ; diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala index d42f744db84d9..766afd29cbb4f 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala @@ -567,6 +567,13 @@ class SqlStatementSplitterSuite extends SparkFunSuite { assert(result.completeStatements == Seq(statement("SELECT 0"))) assert(result.partialStatement == partial) + + val malformedPartial = "BEGIN SELECT 1; SELEC 2;" + val malformed = SqlStatementSplitter + .splitForParseSql(s"SELECT 0; $malformedPartial") + .withoutPositions + assert(malformed.completeStatements == Seq(statement("SELECT 0"))) + assert(malformed.partialStatement == malformedPartial) } test("parse_sql batch boundary preserves unclosed comments") { @@ -613,7 +620,7 @@ class SqlStatementSplitterSuite extends SparkFunSuite { test("parse_sql batch boundary tolerates malformed control headers") { val blocks = Seq( - "BEGIN IF TRUE THEN SELECT 1; END; END", + "BEGIN IF THEN SELECT 1; END IF; END", "BEGIN FOR x AS SELEC 1 DO SELECT 1; END FOR; END") blocks.foreach { block => @@ -625,6 +632,75 @@ class SqlStatementSplitterSuite extends SparkFunSuite { } } + test("parse_sql malformed BEGIN boundaries preserve following statements") { + val malformedBegin = SqlStatementSplitter + .splitForParseSql("BEGIN; SELECT 1;") + .withoutPositions + assert(malformedBegin.completeStatements == Seq( + statement("BEGIN"), + statement("SELECT 1"))) + assert(malformedBegin.partialStatement.isEmpty) + + val strayEnd = SqlStatementSplitter + .splitForParseSql("BEGIN SELECT 1; END; END; SELECT 2") + .withoutPositions + assert(strayEnd.completeStatements == Seq( + statement("BEGIN SELECT 1; END"), + statement("END"))) + assert(strayEnd.partialStatement == "SELECT 2") + + val malformedNestedBegin = SqlStatementSplitter + .splitForParseSql("BEGIN BEGIN; END; END; SELECT 3") + .withoutPositions + assert(malformedNestedBegin.completeStatements == Seq( + statement("BEGIN BEGIN; END; END"))) + assert(malformedNestedBegin.partialStatement == "SELECT 3") + + val malformedIf = "BEGIN IF TRUE THEN BEGIN SELECT 1; END IF; END" + val nestedIf = SqlStatementSplitter + .splitForParseSql(s"SELECT 0; $malformedIf; SELECT 9") + .withoutPositions + assert(nestedIf.completeStatements == Seq( + statement("SELECT 0"), + statement(malformedIf))) + assert(nestedIf.partialStatement == "SELECT 9") + + val malformedHandler = + "BEGIN DECLARE CONTINUE HANDLER FOR SQLEXCEPTION BEGIN; END; END" + val handler = SqlStatementSplitter + .splitForParseSql(s"SELECT 0; $malformedHandler; SELECT 9") + .withoutPositions + assert(handler.completeStatements == Seq( + statement("SELECT 0"), + statement(malformedHandler))) + assert(handler.partialStatement == "SELECT 9") + + val malformedBlocks = Seq( + "BEGIN lbl: BEGIN; END lbl; END", + "BEGIN BEGIN; END lbl; END", + "BEGIN IF TRUE THEN lbl: BEGIN; END lbl; END IF; END", + "BEGIN IF TRUE THEN SELECT 1; ELSEIF THEN lbl: BEGIN; END lbl; END IF; END", + "BEGIN CASE WHEN TRUE THEN SELECT 1; ELSE BEGIN; END; END CASE; END", + "BEGIN CASE x WHEN 1 THEN SELECT 1; ELSE BEGIN; END; END CASE; END", + "BEGIN CASE WHEN THEN BEGIN; END; END CASE; END", + "BEGIN IF TRUE THEN SELECT 1; END; END IF; END", + "BEGIN CASE WHEN TRUE THEN SELECT 1; END; END CASE; END", + "BEGIN WHILE TRUE DO SELECT 1; END; END WHILE; END", + "BEGIN LOOP SELECT 1; END; END LOOP; END", + "BEGIN FOR x AS SELECT 1 DO SELECT 1; END; END FOR; END", + "BEGIN REPEAT SELECT 1; END; UNTIL TRUE END REPEAT; END", + "BEGIN DECLARE CONTINUE HANDLER FOR SQLEXCEPTION lbl: BEGIN; END lbl; END") + malformedBlocks.foreach { block => + val result = try { + SqlStatementSplitter.splitForParseSql(s"$block; SELECT 9").withoutPositions + } catch { + case e: Exception => fail(block, e) + } + assert(result.completeStatements == Seq(statement(block)), block) + assert(result.partialStatement == "SELECT 9", block) + } + } + test("parse_sql batch boundary handles SET, RESET, and non-reserved identifiers") { Seq(false, true).foreach { ansi => SQLConf.withExistingConf(new SQLConf) { diff --git a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala index 0f165cb0ac9e2..f2e5a2896bea8 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala @@ -49,7 +49,7 @@ import org.apache.spark.unsafe.types.UTF8String Requires spark.sql.function.parseSql.enabled=true. On syntax / parse error returns JSON 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 comment-only batch returns `[]`.""", + locations are statement-relative. An empty or closed-comment-only batch returns `[]`.""", arguments = """ Arguments: * sqlStmt - A SQL batch string to split and parse. diff --git a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResult.scala b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResult.scala index cd17bbc30162c..d246bdf3faf4c 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResult.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResult.scala @@ -54,9 +54,9 @@ import org.apache.spark.sql.execution.datasources.CreateTempViewUsing * source location and a nested STANDARD-format error object, and parsing * continues with later statements. Nested error locations are relative to the * individual statement, while `start` is relative to the original batch. An - * empty or comment-only batch produces an empty array. Only [[ParseException]] / - * [[SqlScriptingException]] are converted to JSON; unexpected / internal - * failures propagate so the function fails. + * empty or closed-comment-only batch produces an empty array. Only + * [[ParseException]] / [[SqlScriptingException]] are converted to JSON; + * unexpected / internal failures propagate so the function fails. */ object ParseSqlResult { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala index 97d054626fb70..b3223bb066c13 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala @@ -127,6 +127,52 @@ class ParseSqlResultSuite extends SparkFunSuite { assert(statements(1) \ "error" \ "errorClass" === JString("PARSE_SYNTAX_ERROR")) } + test("malformed BEGIN statements preserve later batch results") { + val cases = Seq( + "BEGIN; SELECT 1;" -> Seq( + "BEGIN" -> false, + "SELECT 1" -> true), + "BEGIN SELECT 1; END; END; SELECT 2" -> Seq( + "BEGIN SELECT 1; END" -> true, + "END" -> false, + "SELECT 2" -> true), + "BEGIN BEGIN; END; END; SELECT 3" -> Seq( + "BEGIN BEGIN; END; END" -> false, + "SELECT 3" -> true), + "SELECT 0; BEGIN IF TRUE THEN BEGIN SELECT 1; END IF; END; SELECT 9" -> Seq( + "SELECT 0" -> true, + "BEGIN IF TRUE THEN BEGIN SELECT 1; END IF; END" -> false, + "SELECT 9" -> true), + "SELECT 0; BEGIN DECLARE CONTINUE HANDLER FOR SQLEXCEPTION BEGIN; END; END; SELECT 9" -> + Seq( + "SELECT 0" -> true, + "BEGIN DECLARE CONTINUE HANDLER FOR SQLEXCEPTION BEGIN; END; END" -> false, + "SELECT 9" -> true), + "BEGIN IF TRUE THEN lbl: BEGIN; END lbl; END IF; END; SELECT 9" -> Seq( + "BEGIN IF TRUE THEN lbl: BEGIN; END lbl; END IF; END" -> false, + "SELECT 9" -> true), + "BEGIN IF TRUE THEN SELECT 1; ELSEIF THEN BEGIN; END; END IF; END; SELECT 9" -> Seq( + "BEGIN IF TRUE THEN SELECT 1; ELSEIF THEN BEGIN; END; END IF; END" -> false, + "SELECT 9" -> true), + "BEGIN CASE WHEN TRUE THEN SELECT 1; ELSE BEGIN; END; END CASE; END; SELECT 9" -> Seq( + "BEGIN CASE WHEN TRUE THEN SELECT 1; ELSE BEGIN; END; END CASE; END" -> false, + "SELECT 9" -> true), + "BEGIN WHILE TRUE DO SELECT 1; END; END WHILE; END; SELECT 9" -> Seq( + "BEGIN WHILE TRUE DO SELECT 1; END; END WHILE; END" -> false, + "SELECT 9" -> true)) + + cases.foreach { case (sql, expected) => + val statements = objs(sql) + assert(statements.size === expected.size, sql) + statements.zip(expected).foreach { case (statement, (text, success)) => + val JInt(start) = statement \ "start" + val JInt(length) = statement \ "length" + assert(sql.substring(start.toInt - 1, start.toInt - 1 + length.toInt) === text) + assert(statement \ "parse_success" === JBool(success)) + } + } + } + test("malformed balanced SQL scripts remain one statement in a batch") { val statements = objs("BEGIN SELECT 1; SELEC 2; END; SELECT 3") assert(statements.size === 2) @@ -174,10 +220,14 @@ class ParseSqlResultSuite extends SparkFunSuite { Seq(JString("BEGIN END"), JString("SELECT"))) } - test("empty batches contain no statements") { - Seq("", " ", ";;", "-- comment").foreach { sql => + test("empty and closed-comment-only batches contain no statements") { + Seq("", " ", ";;", "-- comment", "/* closed */").foreach { sql => assert(objs(sql).isEmpty, sql) } + + val unclosed = objs("/* unclosed") + assert(unclosed.size === 1) + assert(unclosed.head \ "parse_success" === JBool(false)) } test("TABLE and VALUES classify as SELECT") { From 6a76d8114e4e4bbb8d064b9abc348cae8c7a43dd Mon Sep 17 00:00:00 2001 From: srielau Date: Sat, 12 Sep 2026 02:02:24 +0000 Subject: [PATCH 11/12] [SPARK-59255][SQL] Harden parse_sql batch boundaries --- .../sql/catalyst/parser/SqlBaseParser.g4 | 33 +++++++++++++++++-- .../parser/SqlStatementSplitter.scala | 4 +-- .../parser/SqlStatementSplitterSuite.scala | 14 ++++++++ .../catalyst/parser/ParseSqlResultSuite.scala | 28 +++++++++++++++- 4 files changed, 73 insertions(+), 6 deletions(-) diff --git a/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 b/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 index 1ed63b15f76f3..056fac7f3329c 100644 --- a/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 +++ b/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 @@ -124,6 +124,18 @@ parseSqlBatchPartialCompoundStatement parseSqlBatchCompoundStatement : BEGIN (NOT ATOMIC)? parseSqlBatchCompoundBody? END + parseSqlBatchMalformedEndSuffix? + | BEGIN (NOT ATOMIC)? parseSqlBatchCompoundBody? + parseSqlBatchFinalBodyLeafStatement END parseSqlBatchMalformedEndSuffix? + ; + +parseSqlBatchFinalBodyLeafStatement + : {_input.LA(1) != BEGIN && _input.LA(1) != END}? + (~(END | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ + ; + +parseSqlBatchMalformedEndSuffix + : (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ ; parseSqlBatchBeginEndCompoundBlock @@ -175,10 +187,15 @@ parseSqlBatchDeclareHandlerStatement ; parseSqlBatchWhileStatement - : beginLabel? WHILE booleanExpression DO parseSqlBatchCompoundBody + : beginLabel? WHILE parseSqlBatchWhileCondition DO parseSqlBatchCompoundBody parseSqlBatchPrematureEnds? END WHILE endLabel? ; +parseSqlBatchWhileCondition + : booleanExpression + | (~(DO | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))* + ; + parseSqlBatchIfElseStatement : IF parseSqlBatchIfCondition THEN parseSqlBatchConditionalBody (ELSEIF parseSqlBatchIfCondition THEN parseSqlBatchConditionalBody)* @@ -206,7 +223,12 @@ parseSqlBatchConditionalBodyLeafStatement parseSqlBatchRepeatStatement : beginLabel? REPEAT parseSqlBatchCompoundBody parseSqlBatchPrematureEnds? - UNTIL booleanExpression END REPEAT endLabel? + UNTIL parseSqlBatchRepeatCondition END REPEAT endLabel? + ; + +parseSqlBatchRepeatCondition + : booleanExpression + | (~(END | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))* ; parseSqlBatchCaseStatement @@ -246,10 +268,15 @@ parseSqlBatchLoopStatement ; parseSqlBatchForStatement - : beginLabel? FOR (strictIdentifier AS)? query DO + : beginLabel? FOR parseSqlBatchForHeader DO parseSqlBatchCompoundBody parseSqlBatchPrematureEnds? END FOR endLabel? ; +parseSqlBatchForHeader + : (strictIdentifier AS)? query + | (~(DO | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))* + ; + parseSqlBatchPrematureEnds : (END SEMICOLON)+ ; diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala index 137d9b7d5a444..662a927a4ecfa 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala @@ -125,8 +125,8 @@ private[sql] case class PositionedSqlStatementSplitResult( * O(n), but Spark's `setResetStatement` has `SET .*?` / `RESET .*?` wildcards * that need an EOF anchor to terminate deterministically, so such a * single-statement rule-rewrite does not drop in cleanly. Tracked as a - * follow-up. The parse_sql-only path uses [[splitForParseSql]] and performs one - * linear boundary parse instead. + * follow-up. The normal parse_sql-only path uses [[splitForParseSql]] and performs + * one linear boundary parse instead; a parser stack overflow falls back to the generic path. */ object SqlStatementSplitter { diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala index 766afd29cbb4f..f80be1be98add 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala @@ -602,6 +602,9 @@ class SqlStatementSplitterSuite extends SparkFunSuite { "BEGIN FOR x AS SELECT 1 DO SELEC 1; END FOR; END", "BEGIN lbl: BEGIN SELEC 1; END lbl; END", "BEGIN IF ${flag} THEN BEGIN SELEC 1; END; SELECT 1; END IF; END", + "BEGIN WHILE ${flag} DO BEGIN SELEC 1; END; END WHILE; END", + "BEGIN REPEAT BEGIN SELEC 1; END; UNTIL ${flag} END REPEAT; END", + "BEGIN FOR x AS SELECT * FROM ${table} DO BEGIN SELEC 1; END; END FOR; END", "BEGIN DECLARE CONTINUE HANDLER FOR SQLEXCEPTION BEGIN SELEC 1; END; END", "BEGIN DECLARE EXIT HANDLER FOR SQLEXCEPTION BEGIN SELEC 1; END; SELECT 1; END") @@ -699,6 +702,17 @@ class SqlStatementSplitterSuite extends SparkFunSuite { assert(result.completeStatements == Seq(statement(block)), block) assert(result.partialStatement == "SELECT 9", block) } + + val malformedBalanced = Seq( + "BEGIN SELECT 1; END bad", + "BEGIN SELECT 1 END") + malformedBalanced.foreach { block => + val result = SqlStatementSplitter + .splitForParseSql(s"$block; SELECT 2") + .withoutPositions + assert(result.completeStatements == Seq(statement(block)), block) + assert(result.partialStatement == "SELECT 2", block) + } } test("parse_sql batch boundary handles SET, RESET, and non-reserved identifiers") { diff --git a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala index b3223bb066c13..e3237ce4cc81f 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala @@ -159,7 +159,13 @@ class ParseSqlResultSuite extends SparkFunSuite { "SELECT 9" -> true), "BEGIN WHILE TRUE DO SELECT 1; END; END WHILE; END; SELECT 9" -> Seq( "BEGIN WHILE TRUE DO SELECT 1; END; END WHILE; END" -> false, - "SELECT 9" -> true)) + "SELECT 9" -> true), + "BEGIN SELECT 1; END bad; SELECT 2" -> Seq( + "BEGIN SELECT 1; END bad" -> false, + "SELECT 2" -> true), + "BEGIN SELECT 1 END; SELECT 2" -> Seq( + "BEGIN SELECT 1 END" -> false, + "SELECT 2" -> true)) cases.foreach { case (sql, expected) => val statements = objs(sql) @@ -220,6 +226,26 @@ class ParseSqlResultSuite extends SparkFunSuite { Seq(JString("BEGIN END"), JString("SELECT"))) } + test("loop boundary parsing preserves raw variable references for stock parsing") { + SQLConf.withExistingConf(new SQLConf) { + SQLConf.get.setConfString("flag", "TRUE") + SQLConf.get.setConfString("query", "SELECT 1") + val blocks = Seq( + "BEGIN WHILE ${flag} DO SELECT 1; END WHILE; END", + "BEGIN REPEAT SELECT 1; UNTIL ${flag} END REPEAT; END", + "BEGIN FOR x AS ${query} DO SELECT x; END FOR; END") + + blocks.foreach { block => + val sql = s"$block; SELECT 2" + val statements = objs(sql) + assert(statements.size === 2, block) + assert(statements.map(_ \ "parse_success") === Seq(JBool(true), JBool(true)), block) + assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(block.length + 3)), block) + assert(statements.map(_ \ "length") === Seq(JInt(block.length), JInt(8)), block) + } + } + } + test("empty and closed-comment-only batches contain no statements") { Seq("", " ", ";;", "-- comment", "/* closed */").foreach { sql => assert(objs(sql).isEmpty, sql) From 8761060206539b0296a8223f3ef88c4ad159c3a4 Mon Sep 17 00:00:00 2001 From: srielau Date: Sat, 12 Sep 2026 02:20:47 +0000 Subject: [PATCH 12/12] [SPARK-59255][SQL] Reuse shared SQL statement splitter --- .../sql/catalyst/parser/SqlBaseParser.g4 | 204 ---------------- .../parser/SqlStatementSplitter.scala | 116 +-------- .../parser/SqlStatementSplitterSuite.scala | 224 ------------------ .../sql/catalyst/expressions/ParseSql.scala | 4 +- .../spark/sql/execution/SparkSqlParser.scala | 4 +- .../analyzer-results/parse-sql.sql.out | 14 +- .../resources/sql-tests/inputs/parse-sql.sql | 14 +- .../sql-tests/results/parse-sql.sql.out | 19 +- .../catalyst/parser/ParseSqlResultSuite.scala | 110 --------- 9 files changed, 35 insertions(+), 674 deletions(-) diff --git a/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 b/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 index 056fac7f3329c..7865f08423254 100644 --- a/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 +++ b/sql/api/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 @@ -18,8 +18,6 @@ parser grammar SqlBaseParser; options { tokenVocab = SqlBaseLexer; } -tokens { PARSE_SQL_BATCH_DELIMITER } - @members { /** * When false, INTERSECT is given the greater precedence over the other set @@ -89,208 +87,6 @@ compoundOrSingleStatement | singleCompoundStatement ; -// Boundary-only grammar for parse_sql batches. Leaf statements deliberately accept arbitrary -// tokens: ParseSqlResult parses each emitted segment with the full grammar and records any error. -// BEGIN is excluded from the terminated fallback, so only a grammar context can own the -// semicolons inside a compound statement. BEGIN and END remain unrestricted inside leaf -// statements. The caller appends PARSE_SQL_BATCH_DELIMITER; the trailing BEGIN fallback consumes -// a structurally unclosed compound through that token without synthesizing an END token. -parseSqlBatch - : SEMICOLON* (items+=parseSqlBatchItem SEMICOLON*)* - PARSE_SQL_BATCH_DELIMITER? EOF - ; - -parseSqlBatchItem - : batchStatement=parseSqlBatchStatement - terminator=(SEMICOLON | PARSE_SQL_BATCH_DELIMITER) - | partialStatement=parseSqlBatchPartialCompoundStatement - terminator=PARSE_SQL_BATCH_DELIMITER - ; - -parseSqlBatchStatement - : parseSqlBatchCompoundStatement - | parseSqlBatchMalformedEmptyCompoundBlock - | parseSqlBatchMalformedBeginStatement - | parseSqlBatchLeafStatement - ; - -parseSqlBatchMalformedBeginStatement - : BEGIN - ; - -parseSqlBatchPartialCompoundStatement - : BEGIN .*? - ; - -parseSqlBatchCompoundStatement - : BEGIN (NOT ATOMIC)? parseSqlBatchCompoundBody? END - parseSqlBatchMalformedEndSuffix? - | BEGIN (NOT ATOMIC)? parseSqlBatchCompoundBody? - parseSqlBatchFinalBodyLeafStatement END parseSqlBatchMalformedEndSuffix? - ; - -parseSqlBatchFinalBodyLeafStatement - : {_input.LA(1) != BEGIN && _input.LA(1) != END}? - (~(END | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ - ; - -parseSqlBatchMalformedEndSuffix - : (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ - ; - -parseSqlBatchBeginEndCompoundBlock - : beginLabel? BEGIN (NOT ATOMIC)? parseSqlBatchCompoundBody? END endLabel? - ; - -parseSqlBatchMalformedEmptyCompoundBlock - : beginLabel? BEGIN (NOT ATOMIC)? SEMICOLON END endLabel? - ; - -parseSqlBatchMalformedBodyBeginStatement - : BEGIN (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ - ; - -parseSqlBatchCompoundBody - : (parseSqlBatchCompoundBodyStatement SEMICOLON)+ - ; - -parseSqlBatchCompoundBodyStatement - : parseSqlBatchNestedStatement - | parseSqlBatchOrphanControlEndStatement - | parseSqlBatchBodyLeafStatement - ; - -parseSqlBatchNestedStatement - : parseSqlBatchBeginEndCompoundBlock - | parseSqlBatchMalformedEmptyCompoundBlock - | parseSqlBatchDeclareHandlerStatement - | parseSqlBatchIfElseStatement - | parseSqlBatchCaseStatement - | parseSqlBatchWhileStatement - | parseSqlBatchRepeatStatement - | parseSqlBatchLoopStatement - | parseSqlBatchForStatement - | parseSqlBatchMalformedBodyBeginStatement - | parseSqlBatchMalformedBeginStatement - ; - -parseSqlBatchOrphanControlEndStatement - : END (IF | WHILE | LOOP | REPEAT | FOR | CASE) - ; - -parseSqlBatchDeclareHandlerStatement - : DECLARE (CONTINUE | EXIT) HANDLER FOR conditionValues - (parseSqlBatchBeginEndCompoundBlock - | parseSqlBatchMalformedEmptyCompoundBlock - | parseSqlBatchMalformedBodyBeginStatement - | parseSqlBatchBodyLeafStatement) - ; - -parseSqlBatchWhileStatement - : beginLabel? WHILE parseSqlBatchWhileCondition DO parseSqlBatchCompoundBody - parseSqlBatchPrematureEnds? END WHILE endLabel? - ; - -parseSqlBatchWhileCondition - : booleanExpression - | (~(DO | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))* - ; - -parseSqlBatchIfElseStatement - : IF parseSqlBatchIfCondition THEN parseSqlBatchConditionalBody - (ELSEIF parseSqlBatchIfCondition THEN parseSqlBatchConditionalBody)* - (ELSE parseSqlBatchConditionalBody)? parseSqlBatchPrematureEnds? END IF - ; - -parseSqlBatchIfCondition - : (~(THEN | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))* - ; - -parseSqlBatchConditionalBody - : (parseSqlBatchConditionalBodyStatement SEMICOLON)+ - ; - -parseSqlBatchConditionalBodyStatement - : parseSqlBatchNestedStatement - | parseSqlBatchConditionalBodyLeafStatement - ; - -parseSqlBatchConditionalBodyLeafStatement - : {_input.LA(1) != BEGIN && _input.LA(1) != END && - _input.LA(1) != ELSEIF && _input.LA(1) != ELSE}? - (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ - ; - -parseSqlBatchRepeatStatement - : beginLabel? REPEAT parseSqlBatchCompoundBody parseSqlBatchPrematureEnds? - UNTIL parseSqlBatchRepeatCondition END REPEAT endLabel? - ; - -parseSqlBatchRepeatCondition - : booleanExpression - | (~(END | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))* - ; - -parseSqlBatchCaseStatement - : CASE (WHEN parseSqlBatchCaseWhenCondition THEN parseSqlBatchCaseBody)+ - (ELSE parseSqlBatchCaseBody)? parseSqlBatchPrematureEnds? END CASE - | CASE parseSqlBatchCaseOperand - (WHEN parseSqlBatchCaseWhenCondition THEN parseSqlBatchCaseBody)+ - (ELSE parseSqlBatchCaseBody)? parseSqlBatchPrematureEnds? END CASE - ; - -parseSqlBatchCaseOperand - : (~(WHEN | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ - ; - -parseSqlBatchCaseWhenCondition - : (~(THEN | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))* - ; - -parseSqlBatchCaseBody - : (parseSqlBatchCaseBodyStatement SEMICOLON)+ - ; - -parseSqlBatchCaseBodyStatement - : parseSqlBatchNestedStatement - | parseSqlBatchCaseBodyLeafStatement - ; - -parseSqlBatchCaseBodyLeafStatement - : {_input.LA(1) != BEGIN && _input.LA(1) != END && - _input.LA(1) != WHEN && _input.LA(1) != ELSE}? - (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ - ; - -parseSqlBatchLoopStatement - : beginLabel? LOOP parseSqlBatchCompoundBody - parseSqlBatchPrematureEnds? END LOOP endLabel? - ; - -parseSqlBatchForStatement - : beginLabel? FOR parseSqlBatchForHeader DO - parseSqlBatchCompoundBody parseSqlBatchPrematureEnds? END FOR endLabel? - ; - -parseSqlBatchForHeader - : (strictIdentifier AS)? query - | (~(DO | SEMICOLON | PARSE_SQL_BATCH_DELIMITER))* - ; - -parseSqlBatchPrematureEnds - : (END SEMICOLON)+ - ; - -parseSqlBatchLeafStatement - : {_input.LA(1) != BEGIN}? - (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ - ; - -parseSqlBatchBodyLeafStatement - : {_input.LA(1) != BEGIN && _input.LA(1) != END}? - (~(SEMICOLON | PARSE_SQL_BATCH_DELIMITER))+ - ; - singleCompoundStatement : BEGIN (NOT ATOMIC)? compoundBody? END SEMICOLON? EOF ; diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala index 662a927a4ecfa..6bdcab8c58675 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitter.scala @@ -117,16 +117,15 @@ private[sql] case class PositionedSqlStatementSplitResult( * for real at execution time. When `validationPreprocess` is `identity` * (the default), the splitter behaves as a pure original-text splitter. * - * Performance note: the generic splitter calls `tryParseRegion` O(k) times on - * growing prefixes for a single `BEGIN ... END` block with k internal `;` -- an + * Performance note: for a single `BEGIN ... END` block with k internal `;`, + * the splitter calls `tryParseRegion` O(k) times on growing prefixes -- an * O(k^2) cost in the worst case (incomplete block on every keystroke in * interactive mode). Ordinary non-scripting SQL is O(n). A non-EOF terminated * single-statement rule (read `ctx.getStop` once per region) would make this * O(n), but Spark's `setResetStatement` has `SET .*?` / `RESET .*?` wildcards * that need an EOF anchor to terminate deterministically, so such a * single-statement rule-rewrite does not drop in cleanly. Tracked as a - * follow-up. The normal parse_sql-only path uses [[splitForParseSql]] and performs - * one linear boundary parse instead; a parser stack overflow falls back to the generic path. + * follow-up. */ object SqlStatementSplitter { @@ -134,100 +133,6 @@ object SqlStatementSplitter { def split(sqlText: String): SqlStatementSplitResult = splitWithPositions(sqlText, identity).withoutPositions - /** - * Split a parse_sql batch in one grammar-owned pass while retaining source positions. - * Unlike the generic splitter, this boundary-only grammar accepts malformed leaf statements - * and uses scripting grammar contexts to assign internal semicolons to compound statements. - */ - private[sql] def splitForParseSql(sqlText: String): PositionedSqlStatementSplitResult = { - require(sqlText != null, "sqlText must not be null") - - val toUtf16 = utf16Offsets(sqlText) - val sourceLexer = new SqlBaseLexer(new UpperCaseCharStream(CharStreams.fromString(sqlText))) - sourceLexer.removeErrorListeners() - val sourceTokens = new CommonTokenStream(sourceLexer) - sourceTokens.fill() - val boundaryTokens = new java.util.ArrayList[Token](sourceTokens.size() + 1) - var sourceIndex = 0 - while (sourceIndex < sourceTokens.size() - 1) { - boundaryTokens.add(sourceTokens.get(sourceIndex)) - sourceIndex += 1 - } - val boundary = new CommonToken(SqlBaseParser.PARSE_SQL_BATCH_DELIMITER, "") - boundary.setStartIndex(toUtf16.length - 1) - boundary.setStopIndex(toUtf16.length - 2) - boundaryTokens.add(boundary) - boundaryTokens.add(sourceTokens.get(sourceTokens.size() - 1)) - val tokens = new CommonTokenStream(new ListTokenSource(boundaryTokens)) - tokens.fill() - val parser = new SqlBaseParser(tokens) - configureSplitterParser(parser, SqlApiConf.get) - parser.getInterpreter.setPredictionMode(PredictionMode.LL) - val batch = try { - parser.parseSqlBatch() - } catch { - case _: StackOverflowError => - return splitWithPositions(sqlText, identity) - } - - def statementStart(context: ParserRuleContext): Int = { - var index = context.getStart.getTokenIndex - while (index > 0 && tokens.get(index - 1).getChannel == Token.HIDDEN_CHANNEL) { - index -= 1 - } - toUtf16(tokens.get(index).getStartIndex) - } - - def positioned( - context: ParserRuleContext, - end: Int, - terminator: String): PositionedSqlStatement = { - val start = statementStart(context) - val (leadingWhitespace, statement) = trimSqlWhitespace(sqlText.substring(start, end)) - PositionedSqlStatement(statement, terminator, start + leadingWhitespace) - } - - val complete = mutable.ArrayBuffer.empty[PositionedSqlStatement] - var partial: Option[PositionedSqlStatement] = None - var i = 0 - while (i < batch.items.size()) { - val item = batch.items.get(i) - val context = Option(item.batchStatement).getOrElse(item.partialStatement) - if (item.terminator.getType == SqlBaseParser.PARSE_SQL_BATCH_DELIMITER) { - partial = Some(positioned(context, sqlText.length, "")) - } else { - complete += positioned( - context, - toUtf16(item.terminator.getStartIndex), - item.terminator.getText) - } - i += 1 - } - - partial = partial.orElse { - if (sourceLexer.has_unclosed_bracketed_comment) { - var index = sourceTokens.size() - 1 - while (index >= 0 && sourceTokens.get(index).getType != SqlBaseLexer.SEMICOLON) { - index -= 1 - } - val start = if (index < 0) { - 0 - } else { - toUtf16(sourceTokens.get(index).getStopIndex + 1) - } - val (leadingWhitespace, statement) = trimSqlWhitespace(sqlText.substring(start)) - Some(PositionedSqlStatement(statement, "", start + leadingWhitespace)) - } else { - None - } - } - - PositionedSqlStatementSplitResult( - complete.toSeq, - partial, - sourceLexer.has_unclosed_bracketed_comment && partial.nonEmpty) - } - /** * Split the given SQL text, applying `validationPreprocess` to each candidate * region before parser validation. The emitted [[SqlStatement]] text is @@ -242,7 +147,8 @@ object SqlStatementSplitter { splitWithPositions(sqlText, validationPreprocess).withoutPositions /** - * Split SQL while retaining each statement's source position. + * 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, @@ -358,7 +264,7 @@ object SqlStatementSplitter { // prefix that includes the next `;`. d += 1 case FailedNonEof => - // Structurally invalid ordinary statement; stop extending. + // Structurally invalid; stop extending. failedNonEof = true } } @@ -369,11 +275,7 @@ object SqlStatementSplitter { // `startIdx` are leading whitespace/comments that we preserve in the // buffer too (so the emitted statement text matches the original // input shape, modulo trimming). - val terminator = if (tokenStream.get(matchedDelimIdx).getType == Token.EOF) { - "" - } else { - tokenStream.get(matchedDelimIdx).getText - } + val terminator = tokenStream.get(matchedDelimIdx).getText while (index < matchedDelimIdx) { val tok = tokenStream.get(index) if (tok.getChannel != Token.HIDDEN_CHANNEL) bufferHasContent = true @@ -545,9 +447,7 @@ object SqlStatementSplitter { * exceptions, but the splitter surfaces them via the * [[SqlStatementSplitResult.hasUnclosedComment]] flag instead. */ - private def configureSplitterParser( - parser: SqlBaseParser, - conf: SqlApiConf): Unit = { + private def configureSplitterParser(parser: SqlBaseParser, conf: SqlApiConf): Unit = { if (conf.manageParserCaches) AbstractParser.installCaches(parser) parser.legacy_setops_precedence_enabled = conf.setOpsPrecedenceEnforced diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala index f80be1be98add..dfa07229901d5 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/SqlStatementSplitterSuite.scala @@ -18,7 +18,6 @@ package org.apache.spark.sql.catalyst.parser import org.apache.spark.SparkFunSuite -import org.apache.spark.sql.internal.SQLConf class SqlStatementSplitterSuite extends SparkFunSuite { @@ -517,229 +516,6 @@ class SqlStatementSplitterSuite extends SparkFunSuite { assert(result.partialStatement.isEmpty) } - test("malformed balanced BEGIN block remains one statement") { - val block = "BEGIN SELECT 1; SELEC 2; END" - val result = SqlStatementSplitter.splitForParseSql(s"$block; SELECT 3;") - .withoutPositions - assert(result.completeStatements == Seq( - statement(block), - statement("SELECT 3"))) - assert(result.partialStatement.isEmpty) - - val withoutTerminator = SqlStatementSplitter.splitForParseSql(block) - .withoutPositions - assert(withoutTerminator.completeStatements.isEmpty) - assert(withoutTerminator.partialStatement == block) - } - - test("malformed nested control END does not end the outer block") { - val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; SELECT 2; END" - val result = SqlStatementSplitter.splitForParseSql(s"$block; SELECT 3") - .withoutPositions - - assert(result.completeStatements == Seq(statement(block))) - assert(result.partialStatement == "SELECT 3") - } - - test("statement-final END before the outer END does not end the block") { - val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; SELECT END; END" - val result = SqlStatementSplitter.splitForParseSql(s"$block; SELECT 3") - .withoutPositions - - assert(result.completeStatements == Seq(statement(block))) - assert(result.partialStatement == "SELECT 3") - } - - test("nested compound END does not end the outer malformed block") { - val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; BEGIN SELECT 2; END; END" - val result = SqlStatementSplitter.splitForParseSql(s"$block; SELECT 3") - .withoutPositions - - assert(result.completeStatements == Seq(statement(block))) - assert(result.partialStatement == "SELECT 3") - } - - test("parse_sql batch boundary leaves an unclosed compound as partial") { - val partial = "BEGIN SELECT 1; SELECT 2;" - val result = SqlStatementSplitter - .splitForParseSql(s"SELECT 0; $partial") - .withoutPositions - - assert(result.completeStatements == Seq(statement("SELECT 0"))) - assert(result.partialStatement == partial) - - val malformedPartial = "BEGIN SELECT 1; SELEC 2;" - val malformed = SqlStatementSplitter - .splitForParseSql(s"SELECT 0; $malformedPartial") - .withoutPositions - assert(malformed.completeStatements == Seq(statement("SELECT 0"))) - assert(malformed.partialStatement == malformedPartial) - } - - test("parse_sql batch boundary preserves unclosed comments") { - val result = SqlStatementSplitter - .splitForParseSql("SELECT 1; /* unclosed") - .withoutPositions - assert(result.completeStatements == Seq(statement("SELECT 1"))) - assert(result.partialStatement == "/* unclosed") - assert(result.hasUnclosedComment) - - val commentOnly = SqlStatementSplitter - .splitForParseSql("/* unclosed") - .withoutPositions - assert(commentOnly.completeStatements.isEmpty) - assert(commentOnly.partialStatement == "/* unclosed") - assert(commentOnly.hasUnclosedComment) - } - - test("parse_sql batch boundary follows scripting grammar contexts") { - val blocks = Seq( - "BEGIN IF TRUE THEN SELEC 1; END IF; END", - "BEGIN CASE WHEN TRUE THEN SELEC 1; END CASE; END", - "BEGIN WHILE TRUE DO SELEC 1; END WHILE; END", - "BEGIN REPEAT SELEC 1; UNTIL TRUE END REPEAT; END", - "BEGIN LOOP SELEC 1; END LOOP; END", - "BEGIN FOR x AS SELECT 1 DO SELEC 1; END FOR; END", - "BEGIN lbl: BEGIN SELEC 1; END lbl; END", - "BEGIN IF ${flag} THEN BEGIN SELEC 1; END; SELECT 1; END IF; END", - "BEGIN WHILE ${flag} DO BEGIN SELEC 1; END; END WHILE; END", - "BEGIN REPEAT BEGIN SELEC 1; END; UNTIL ${flag} END REPEAT; END", - "BEGIN FOR x AS SELECT * FROM ${table} DO BEGIN SELEC 1; END; END FOR; END", - "BEGIN DECLARE CONTINUE HANDLER FOR SQLEXCEPTION BEGIN SELEC 1; END; END", - "BEGIN DECLARE EXIT HANDLER FOR SQLEXCEPTION BEGIN SELEC 1; END; SELECT 1; END") - - blocks.foreach { block => - val result = SqlStatementSplitter - .splitForParseSql(s"$block; SELECT 9") - .withoutPositions - assert(result.completeStatements == Seq(statement(block)), block) - assert(result.partialStatement == "SELECT 9", block) - - val trailing = SqlStatementSplitter.splitForParseSql(block).withoutPositions - assert(trailing.completeStatements.isEmpty, block) - assert(trailing.partialStatement == block, block) - } - } - - test("parse_sql batch boundary tolerates malformed control headers") { - val blocks = Seq( - "BEGIN IF THEN SELECT 1; END IF; END", - "BEGIN FOR x AS SELEC 1 DO SELECT 1; END FOR; END") - - blocks.foreach { block => - val result = SqlStatementSplitter - .splitForParseSql(s"$block; SELECT 9") - .withoutPositions - assert(result.completeStatements == Seq(statement(block)), block) - assert(result.partialStatement == "SELECT 9", block) - } - } - - test("parse_sql malformed BEGIN boundaries preserve following statements") { - val malformedBegin = SqlStatementSplitter - .splitForParseSql("BEGIN; SELECT 1;") - .withoutPositions - assert(malformedBegin.completeStatements == Seq( - statement("BEGIN"), - statement("SELECT 1"))) - assert(malformedBegin.partialStatement.isEmpty) - - val strayEnd = SqlStatementSplitter - .splitForParseSql("BEGIN SELECT 1; END; END; SELECT 2") - .withoutPositions - assert(strayEnd.completeStatements == Seq( - statement("BEGIN SELECT 1; END"), - statement("END"))) - assert(strayEnd.partialStatement == "SELECT 2") - - val malformedNestedBegin = SqlStatementSplitter - .splitForParseSql("BEGIN BEGIN; END; END; SELECT 3") - .withoutPositions - assert(malformedNestedBegin.completeStatements == Seq( - statement("BEGIN BEGIN; END; END"))) - assert(malformedNestedBegin.partialStatement == "SELECT 3") - - val malformedIf = "BEGIN IF TRUE THEN BEGIN SELECT 1; END IF; END" - val nestedIf = SqlStatementSplitter - .splitForParseSql(s"SELECT 0; $malformedIf; SELECT 9") - .withoutPositions - assert(nestedIf.completeStatements == Seq( - statement("SELECT 0"), - statement(malformedIf))) - assert(nestedIf.partialStatement == "SELECT 9") - - val malformedHandler = - "BEGIN DECLARE CONTINUE HANDLER FOR SQLEXCEPTION BEGIN; END; END" - val handler = SqlStatementSplitter - .splitForParseSql(s"SELECT 0; $malformedHandler; SELECT 9") - .withoutPositions - assert(handler.completeStatements == Seq( - statement("SELECT 0"), - statement(malformedHandler))) - assert(handler.partialStatement == "SELECT 9") - - val malformedBlocks = Seq( - "BEGIN lbl: BEGIN; END lbl; END", - "BEGIN BEGIN; END lbl; END", - "BEGIN IF TRUE THEN lbl: BEGIN; END lbl; END IF; END", - "BEGIN IF TRUE THEN SELECT 1; ELSEIF THEN lbl: BEGIN; END lbl; END IF; END", - "BEGIN CASE WHEN TRUE THEN SELECT 1; ELSE BEGIN; END; END CASE; END", - "BEGIN CASE x WHEN 1 THEN SELECT 1; ELSE BEGIN; END; END CASE; END", - "BEGIN CASE WHEN THEN BEGIN; END; END CASE; END", - "BEGIN IF TRUE THEN SELECT 1; END; END IF; END", - "BEGIN CASE WHEN TRUE THEN SELECT 1; END; END CASE; END", - "BEGIN WHILE TRUE DO SELECT 1; END; END WHILE; END", - "BEGIN LOOP SELECT 1; END; END LOOP; END", - "BEGIN FOR x AS SELECT 1 DO SELECT 1; END; END FOR; END", - "BEGIN REPEAT SELECT 1; END; UNTIL TRUE END REPEAT; END", - "BEGIN DECLARE CONTINUE HANDLER FOR SQLEXCEPTION lbl: BEGIN; END lbl; END") - malformedBlocks.foreach { block => - val result = try { - SqlStatementSplitter.splitForParseSql(s"$block; SELECT 9").withoutPositions - } catch { - case e: Exception => fail(block, e) - } - assert(result.completeStatements == Seq(statement(block)), block) - assert(result.partialStatement == "SELECT 9", block) - } - - val malformedBalanced = Seq( - "BEGIN SELECT 1; END bad", - "BEGIN SELECT 1 END") - malformedBalanced.foreach { block => - val result = SqlStatementSplitter - .splitForParseSql(s"$block; SELECT 2") - .withoutPositions - assert(result.completeStatements == Seq(statement(block)), block) - assert(result.partialStatement == "SELECT 2", block) - } - } - - test("parse_sql batch boundary handles SET, RESET, and non-reserved identifiers") { - Seq(false, true).foreach { ansi => - SQLConf.withExistingConf(new SQLConf) { - SQLConf.get.setConf(SQLConf.ANSI_ENABLED, ansi) - SQLConf.get.setConf(SQLConf.ENFORCE_RESERVED_KEYWORDS, ansi) - val result = SqlStatementSplitter - .splitForParseSql( - "SET spark.sql.ansi.enabled=true; RESET spark.sql.ansi.enabled; " + - "SELECT begin, end FROM t") - .withoutPositions - - assert(result.completeStatements == Seq( - statement("SET spark.sql.ansi.enabled=true"), - statement("RESET spark.sql.ansi.enabled"))) - assert(result.partialStatement == "SELECT begin, end FROM t") - } - } - - val malformed = SqlStatementSplitter - .splitForParseSql("SELECT FROM; SELECT 2") - .withoutPositions - assert(malformed.completeStatements == Seq(statement("SELECT FROM"))) - assert(malformed.partialStatement == "SELECT 2") - } - test("Valid BEGIN..END block is never split at internal ;") { // A valid `BEGIN ... END` block must be confirmed in full -- the splitter // must never emit at an internal `;`, even if a shorter prefix happens to diff --git a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala index f2e5a2896bea8..9551b9e10acd5 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/catalyst/expressions/ParseSql.scala @@ -29,8 +29,8 @@ import org.apache.spark.unsafe.types.UTF8String /** * 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 produces a - * result object containing a nested STANDARD-format error object. + * 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. diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala index 84c0ffb1ce890..bce4250c18975 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkSqlParser.scala @@ -126,7 +126,9 @@ class SparkSqlParser extends AbstractSqlParser { /** Split statements while retaining their positions in the original SQL text. */ private[sql] def splitStatementsWithPositions( sqlText: String): PositionedSqlStatementSplitResult = - SqlStatementSplitter.splitForParseSql(sqlText) + SqlStatementSplitter.splitWithPositions( + sqlText, + SparkSqlParser.substituteVariablesForValidation) /** * Internal parse method that handles both parameter substitution and regular parsing. diff --git a/sql/core/src/test/resources/sql-tests/analyzer-results/parse-sql.sql.out b/sql/core/src/test/resources/sql-tests/analyzer-results/parse-sql.sql.out index 4cc03527df591..30d0d993176f9 100644 --- a/sql/core/src/test/resources/sql-tests/analyzer-results/parse-sql.sql.out +++ b/sql/core/src/test/resources/sql-tests/analyzer-results/parse-sql.sql.out @@ -412,12 +412,12 @@ Project [parse_sql(BEGIN -- !query SELECT - get_json_object(result, '$[0].start') AS start, - get_json_object(result, '$[0].length') AS length, - get_json_object(result, '$[0].error.errorClass') AS error_class, - get_json_object(result, '$[0].error.line') AS line, - get_json_object(result, '$[0].error.position') AS position, - get_json_object(result, '$[0].error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[1].start') AS start, + get_json_object(result, '$[1].length') AS length, + get_json_object(result, '$[1].error.errorClass') AS error_class, + get_json_object(result, '$[1].error.line') AS line, + get_json_object(result, '$[1].error.position') AS position, + get_json_object(result, '$[1].error.queryContext[0].fragment') AS fragment FROM ( SELECT parse_sql( 'BEGIN @@ -426,7 +426,7 @@ FROM ( END') AS result ) -- !query analysis -Project [get_json_object(result#x, $[0].start) AS start#x, get_json_object(result#x, $[0].length) AS length#x, get_json_object(result#x, $[0].error.errorClass) AS error_class#x, get_json_object(result#x, $[0].error.line) AS line#x, get_json_object(result#x, $[0].error.position) AS position#x, get_json_object(result#x, $[0].error.queryContext[0].fragment) AS fragment#x] +Project [get_json_object(result#x, $[1].start) AS start#x, get_json_object(result#x, $[1].length) AS length#x, get_json_object(result#x, $[1].error.errorClass) AS error_class#x, get_json_object(result#x, $[1].error.line) AS line#x, get_json_object(result#x, $[1].error.position) AS position#x, get_json_object(result#x, $[1].error.queryContext[0].fragment) AS fragment#x] +- SubqueryAlias __auto_generated_subquery_name +- Project [parse_sql(BEGIN SELECT 1; diff --git a/sql/core/src/test/resources/sql-tests/inputs/parse-sql.sql b/sql/core/src/test/resources/sql-tests/inputs/parse-sql.sql index 5dcd9e42bbab6..6b73c8722c6bb 100644 --- a/sql/core/src/test/resources/sql-tests/inputs/parse-sql.sql +++ b/sql/core/src/test/resources/sql-tests/inputs/parse-sql.sql @@ -144,7 +144,7 @@ SELECT parse_sql('CREATE VIEW v AS SELECT a, b FROM t'); SELECT parse_sql('SELECT 1 AS IDENTIFIER(''alias.field'')'); SELECT parse_sql('SELECT DATE ''not-a-date'''); --- a malformed balanced script remains one failed statement +-- location for an error inside a multiline script --QUERY-DELIMITER-START SELECT parse_sql( 'BEGIN @@ -156,12 +156,12 @@ SELECT parse_sql( -- JSON-path access over one shared scripting parse result --QUERY-DELIMITER-START SELECT - get_json_object(result, '$[0].start') AS start, - get_json_object(result, '$[0].length') AS length, - get_json_object(result, '$[0].error.errorClass') AS error_class, - get_json_object(result, '$[0].error.line') AS line, - get_json_object(result, '$[0].error.position') AS position, - get_json_object(result, '$[0].error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[1].start') AS start, + get_json_object(result, '$[1].length') AS length, + get_json_object(result, '$[1].error.errorClass') AS error_class, + get_json_object(result, '$[1].error.line') AS line, + get_json_object(result, '$[1].error.position') AS position, + get_json_object(result, '$[1].error.queryContext[0].fragment') AS fragment FROM ( SELECT parse_sql( 'BEGIN diff --git a/sql/core/src/test/resources/sql-tests/results/parse-sql.sql.out b/sql/core/src/test/resources/sql-tests/results/parse-sql.sql.out index 4fefe284ce9c5..1310f7c30fffa 100644 --- a/sql/core/src/test/resources/sql-tests/results/parse-sql.sql.out +++ b/sql/core/src/test/resources/sql-tests/results/parse-sql.sql.out @@ -412,17 +412,17 @@ struct -- !query output -[{"start":1,"length":35,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"'2'","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":35,"fragment":"BEGIN\n SELECT 1;\n SELEC 2;\n END"}],"line":3,"position":9}}] +[{"start":1,"length":17,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"end of input","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":17,"fragment":"BEGIN\n SELECT 1"}],"line":2,"position":11}},{"start":23,"length":7,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"'SELEC'","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":7,"fragment":"SELEC 2"}],"line":1,"position":0}},{"start":33,"length":3,"parse_success":false,"error":{"errorClass":"PARSE_SYNTAX_ERROR","messageTemplate":"Syntax error at or near .","sqlState":"42601","messageParameters":{"error":"'END'","hint":""},"queryContext":[{"objectType":"","objectName":"","startIndex":1,"stopIndex":3,"fragment":"END"}],"line":1,"position":0}}] -- !query SELECT - get_json_object(result, '$[0].start') AS start, - get_json_object(result, '$[0].length') AS length, - get_json_object(result, '$[0].error.errorClass') AS error_class, - get_json_object(result, '$[0].error.line') AS line, - get_json_object(result, '$[0].error.position') AS position, - get_json_object(result, '$[0].error.queryContext[0].fragment') AS fragment + get_json_object(result, '$[1].start') AS start, + get_json_object(result, '$[1].length') AS length, + get_json_object(result, '$[1].error.errorClass') AS error_class, + get_json_object(result, '$[1].error.line') AS line, + get_json_object(result, '$[1].error.position') AS position, + get_json_object(result, '$[1].error.queryContext[0].fragment') AS fragment FROM ( SELECT parse_sql( 'BEGIN @@ -433,10 +433,7 @@ FROM ( -- !query schema struct -- !query output -1 35 PARSE_SYNTAX_ERROR 3 9 BEGIN - SELECT 1; - SELEC 2; - END +23 7 PARSE_SYNTAX_ERROR 1 0 SELEC 2 -- !query diff --git a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala index e3237ce4cc81f..a04c8b7753854 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/catalyst/parser/ParseSqlResultSuite.scala @@ -127,96 +127,6 @@ class ParseSqlResultSuite extends SparkFunSuite { assert(statements(1) \ "error" \ "errorClass" === JString("PARSE_SYNTAX_ERROR")) } - test("malformed BEGIN statements preserve later batch results") { - val cases = Seq( - "BEGIN; SELECT 1;" -> Seq( - "BEGIN" -> false, - "SELECT 1" -> true), - "BEGIN SELECT 1; END; END; SELECT 2" -> Seq( - "BEGIN SELECT 1; END" -> true, - "END" -> false, - "SELECT 2" -> true), - "BEGIN BEGIN; END; END; SELECT 3" -> Seq( - "BEGIN BEGIN; END; END" -> false, - "SELECT 3" -> true), - "SELECT 0; BEGIN IF TRUE THEN BEGIN SELECT 1; END IF; END; SELECT 9" -> Seq( - "SELECT 0" -> true, - "BEGIN IF TRUE THEN BEGIN SELECT 1; END IF; END" -> false, - "SELECT 9" -> true), - "SELECT 0; BEGIN DECLARE CONTINUE HANDLER FOR SQLEXCEPTION BEGIN; END; END; SELECT 9" -> - Seq( - "SELECT 0" -> true, - "BEGIN DECLARE CONTINUE HANDLER FOR SQLEXCEPTION BEGIN; END; END" -> false, - "SELECT 9" -> true), - "BEGIN IF TRUE THEN lbl: BEGIN; END lbl; END IF; END; SELECT 9" -> Seq( - "BEGIN IF TRUE THEN lbl: BEGIN; END lbl; END IF; END" -> false, - "SELECT 9" -> true), - "BEGIN IF TRUE THEN SELECT 1; ELSEIF THEN BEGIN; END; END IF; END; SELECT 9" -> Seq( - "BEGIN IF TRUE THEN SELECT 1; ELSEIF THEN BEGIN; END; END IF; END" -> false, - "SELECT 9" -> true), - "BEGIN CASE WHEN TRUE THEN SELECT 1; ELSE BEGIN; END; END CASE; END; SELECT 9" -> Seq( - "BEGIN CASE WHEN TRUE THEN SELECT 1; ELSE BEGIN; END; END CASE; END" -> false, - "SELECT 9" -> true), - "BEGIN WHILE TRUE DO SELECT 1; END; END WHILE; END; SELECT 9" -> Seq( - "BEGIN WHILE TRUE DO SELECT 1; END; END WHILE; END" -> false, - "SELECT 9" -> true), - "BEGIN SELECT 1; END bad; SELECT 2" -> Seq( - "BEGIN SELECT 1; END bad" -> false, - "SELECT 2" -> true), - "BEGIN SELECT 1 END; SELECT 2" -> Seq( - "BEGIN SELECT 1 END" -> false, - "SELECT 2" -> true)) - - cases.foreach { case (sql, expected) => - val statements = objs(sql) - assert(statements.size === expected.size, sql) - statements.zip(expected).foreach { case (statement, (text, success)) => - val JInt(start) = statement \ "start" - val JInt(length) = statement \ "length" - assert(sql.substring(start.toInt - 1, start.toInt - 1 + length.toInt) === text) - assert(statement \ "parse_success" === JBool(success)) - } - } - } - - test("malformed balanced SQL scripts remain one statement in a batch") { - val statements = objs("BEGIN SELECT 1; SELEC 2; END; SELECT 3") - assert(statements.size === 2) - assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(31))) - assert(statements.map(_ \ "length") === Seq(JInt(28), JInt(8))) - assert(statements.map(_ \ "parse_success") === Seq(JBool(false), JBool(true))) - } - - test("malformed nested control END does not end the outer SQL script") { - val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; SELECT 2; END" - val statements = objs(s"$block; SELECT 3") - - assert(statements.size === 2) - assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(block.length + 3))) - assert(statements.map(_ \ "length") === Seq(JInt(block.length), JInt(8))) - assert(statements.map(_ \ "parse_success") === Seq(JBool(false), JBool(true))) - } - - test("statement-final END before the outer END does not end the SQL script") { - val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; SELECT END; END" - val statements = objs(s"$block; SELECT 3") - - assert(statements.size === 2) - assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(block.length + 3))) - assert(statements.map(_ \ "length") === Seq(JInt(block.length), JInt(8))) - assert(statements.map(_ \ "parse_success") === Seq(JBool(false), JBool(true))) - } - - test("nested compound END does not end the outer malformed SQL script") { - val block = "BEGIN IFF TRUE THEN SELECT 1; END IF; BEGIN SELECT 2; END; END" - val statements = objs(s"$block; SELECT 3") - - assert(statements.size === 2) - assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(block.length + 3))) - assert(statements.map(_ \ "length") === Seq(JInt(block.length), JInt(8))) - assert(statements.map(_ \ "parse_success") === Seq(JBool(false), JBool(true))) - } - test("SQL scripts remain one statement in a batch") { val statements = objs("BEGIN SELECT 1; SELECT 2; END; SELECT 3") assert(statements.size === 2) @@ -226,26 +136,6 @@ class ParseSqlResultSuite extends SparkFunSuite { Seq(JString("BEGIN END"), JString("SELECT"))) } - test("loop boundary parsing preserves raw variable references for stock parsing") { - SQLConf.withExistingConf(new SQLConf) { - SQLConf.get.setConfString("flag", "TRUE") - SQLConf.get.setConfString("query", "SELECT 1") - val blocks = Seq( - "BEGIN WHILE ${flag} DO SELECT 1; END WHILE; END", - "BEGIN REPEAT SELECT 1; UNTIL ${flag} END REPEAT; END", - "BEGIN FOR x AS ${query} DO SELECT x; END FOR; END") - - blocks.foreach { block => - val sql = s"$block; SELECT 2" - val statements = objs(sql) - assert(statements.size === 2, block) - assert(statements.map(_ \ "parse_success") === Seq(JBool(true), JBool(true)), block) - assert(statements.map(_ \ "start") === Seq(JInt(1), JInt(block.length + 3)), block) - assert(statements.map(_ \ "length") === Seq(JInt(block.length), JInt(8)), block) - } - } - } - test("empty and closed-comment-only batches contain no statements") { Seq("", " ", ";;", "-- comment", "/* closed */").foreach { sql => assert(objs(sql).isEmpty, sql)