From 6b42040db8c48a72ffb3b8ef9c60d41921abfb59 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Fri, 18 Nov 2022 11:32:17 +0300 Subject: [PATCH 01/38] Add UnboundParameter --- .../spark/sql/catalyst/parser/SqlBaseLexer.g4 | 4 +++ .../sql/catalyst/parser/SqlBaseParser.g4 | 5 +++ .../sql/catalyst/expressions/parameters.scala | 35 +++++++++++++++++++ .../sql/catalyst/parser/AstBuilder.scala | 8 +++++ 4 files changed, 52 insertions(+) create mode 100644 sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 index 41adbda7b101e..5d446dcc3e87a 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 @@ -507,3 +507,7 @@ WS UNRECOGNIZED : . ; + +NAMED_PARAMETER_MARKER + : '@' + ; diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 index a3c5f4a7b0709..bee64772c8db9 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 @@ -929,6 +929,7 @@ constant | number #numericLiteral | booleanValue #booleanLiteral | stringLit+ #stringLiteral + | namedParameter #unboundParameter ; comparisonOperator @@ -1161,6 +1162,10 @@ version | stringLit ; +namedParameter + : NAMED_PARAMETER_MARKER IDENTIFIER + ; + // When `SQL_standard_keyword_behavior=true`, there are 2 kinds of keywords in Spark SQL. // - Reserved keywords: // Keywords that are reserved and can't be used as identifiers for table, view, column, diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala new file mode 100644 index 0000000000000..86670c3771ce0 --- /dev/null +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -0,0 +1,35 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.catalyst.expressions + +import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan + +/** + * The logical node represents a named parameter that should be bound to a literal + * with concrete value and type later. + * + * @param name The identifier of the parameter without the marker '@'. + */ +case class UnboundParameter(name: String) extends LogicalPlan { + override def output: Seq[Attribute] = Seq.empty + + override def children: Seq[LogicalPlan] = Seq.empty + + override protected def withNewChildrenInternal( + newChildren: IndexedSeq[LogicalPlan]): LogicalPlan = this +} diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala index af2097b5d0ff5..ac5977038bda0 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala @@ -4822,4 +4822,12 @@ class AstBuilder extends SqlBaseParserBaseVisitor[AnyRef] with SQLConfHelper wit override def visitTimestampdiff(ctx: TimestampdiffContext): Expression = withOrigin(ctx) { TimestampDiff(ctx.unit.getText, expression(ctx.startTimestamp), expression(ctx.endTimestamp)) } + + /** + * Create a node in logical plan of named parameter which represents a literal with + * a non-bound value and unknown type. + * */ + override def visitUnboundParameter(ctx: UnboundParameterContext): LogicalPlan = withOrigin(ctx) { + UnboundParameter(ctx.getText) + } } From fe3ca7a1adb7c770bc1c032bc64e74a68ba2f4c2 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Fri, 18 Nov 2022 12:58:15 +0300 Subject: [PATCH 02/38] Add a test --- .../spark/sql/catalyst/parser/SqlBaseLexer.g4 | 4 ++-- .../sql/catalyst/parser/SqlBaseParser.g4 | 6 +----- .../sql/catalyst/expressions/parameters.scala | 21 ++++++++++++------- .../sql/catalyst/parser/AstBuilder.scala | 5 ++--- .../sql/catalyst/parser/PlanParserSuite.scala | 6 ++++++ 5 files changed, 25 insertions(+), 17 deletions(-) diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 index 5d446dcc3e87a..5d586dcb6b96c 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 @@ -508,6 +508,6 @@ UNRECOGNIZED : . ; -NAMED_PARAMETER_MARKER - : '@' +NAMED_PARAMETER + : '@' IDENTIFIER ; diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 index bee64772c8db9..a50aab9a8df6c 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 @@ -924,12 +924,12 @@ primaryExpression constant : NULL #nullLiteral + | NAMED_PARAMETER #unboundParameter | interval #intervalLiteral | identifier stringLit #typeConstructor | number #numericLiteral | booleanValue #booleanLiteral | stringLit+ #stringLiteral - | namedParameter #unboundParameter ; comparisonOperator @@ -1162,10 +1162,6 @@ version | stringLit ; -namedParameter - : NAMED_PARAMETER_MARKER IDENTIFIER - ; - // When `SQL_standard_keyword_behavior=true`, there are 2 kinds of keywords in Spark SQL. // - Reserved keywords: // Keywords that are reserved and can't be used as identifiers for table, view, column, diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index 86670c3771ce0..31c57d40a27f9 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -17,19 +17,26 @@ package org.apache.spark.sql.catalyst.expressions -import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan +import org.apache.spark.SparkException +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, ExprCode} +import org.apache.spark.sql.types.{DataType, NullType} /** - * The logical node represents a named parameter that should be bound to a literal + * The expression represents a named parameter that should be bound to a literal * with concrete value and type later. * * @param name The identifier of the parameter without the marker '@'. */ -case class UnboundParameter(name: String) extends LogicalPlan { - override def output: Seq[Attribute] = Seq.empty +case class UnboundParameter(name: String) extends LeafExpression { + override def dataType: DataType = NullType + override def nullable: Boolean = true - override def children: Seq[LogicalPlan] = Seq.empty + override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { + throw SparkException.internalError(s"Found an unbound parameter: $name") + } - override protected def withNewChildrenInternal( - newChildren: IndexedSeq[LogicalPlan]): LogicalPlan = this + def eval(input: InternalRow): Any = { + throw SparkException.internalError(s"Found an unbound parameter: $name") + } } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala index ac5977038bda0..01ad5c7604550 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala @@ -4824,10 +4824,9 @@ class AstBuilder extends SqlBaseParserBaseVisitor[AnyRef] with SQLConfHelper wit } /** - * Create a node in logical plan of named parameter which represents a literal with - * a non-bound value and unknown type. + * Create a named parameter which represents a literal with a non-bound value and unknown type. * */ - override def visitUnboundParameter(ctx: UnboundParameterContext): LogicalPlan = withOrigin(ctx) { + override def visitUnboundParameter(ctx: UnboundParameterContext): Expression = withOrigin(ctx) { UnboundParameter(ctx.getText) } } diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala index 968e22272341f..2d450f1ed8ca5 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala @@ -1542,4 +1542,10 @@ class PlanParserSuite extends AnalysisTest { .toAggregateExpression(false, Some(GreaterThan(UnresolvedAttribute("id"), Literal(10)))) ) } + + test("named parameters") { + comparePlans( + parsePlan("SELECT @param1"), + Project(UnresolvedAlias(UnboundParameter("@param1"), None) :: Nil, OneRowRelation())) + } } From c264476338c314c21bde947b5e659df222540781 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Mon, 21 Nov 2022 12:22:35 +0300 Subject: [PATCH 03/38] Refactoring + add more parsing tests --- .../spark/sql/catalyst/parser/SqlBaseLexer.g4 | 2 +- .../sql/catalyst/parser/SqlBaseParser.g4 | 2 +- .../sql/catalyst/expressions/parameters.scala | 10 ++++----- .../sql/catalyst/parser/AstBuilder.scala | 4 ++-- .../sql/catalyst/parser/PlanParserSuite.scala | 22 +++++++++++++++++-- 5 files changed, 29 insertions(+), 11 deletions(-) diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 index 5d586dcb6b96c..a380fa19c3424 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 @@ -508,6 +508,6 @@ UNRECOGNIZED : . ; -NAMED_PARAMETER +PARAMETER : '@' IDENTIFIER ; diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 index a50aab9a8df6c..ee88d81e5264d 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 @@ -924,7 +924,7 @@ primaryExpression constant : NULL #nullLiteral - | NAMED_PARAMETER #unboundParameter + | PARAMETER #parameterLiteral | interval #intervalLiteral | identifier stringLit #typeConstructor | number #numericLiteral diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index 31c57d40a27f9..f07aaa969735b 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -23,20 +23,20 @@ import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, ExprCo import org.apache.spark.sql.types.{DataType, NullType} /** - * The expression represents a named parameter that should be bound to a literal - * with concrete value and type later. + * The expression represents a named parameter that should be bound later + * to a literal with concrete value and type. * * @param name The identifier of the parameter without the marker '@'. */ -case class UnboundParameter(name: String) extends LeafExpression { +case class NamedParameter(name: String) extends LeafExpression { override def dataType: DataType = NullType override def nullable: Boolean = true override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { - throw SparkException.internalError(s"Found an unbound parameter: $name") + throw SparkException.internalError(s"Found the unbound parameter: $name") } def eval(input: InternalRow): Any = { - throw SparkException.internalError(s"Found an unbound parameter: $name") + throw SparkException.internalError(s"Found the unbound parameter: $name") } } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala index 562847f721fec..954adb5e9fd21 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala @@ -4830,7 +4830,7 @@ class AstBuilder extends SqlBaseParserBaseVisitor[AnyRef] with SQLConfHelper wit /** * Create a named parameter which represents a literal with a non-bound value and unknown type. * */ - override def visitUnboundParameter(ctx: UnboundParameterContext): Expression = withOrigin(ctx) { - UnboundParameter(ctx.getText) + override def visitParameterLiteral(ctx: ParameterLiteralContext): Expression = withOrigin(ctx) { + NamedParameter(ctx.getText.stripPrefix("@")) } } diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala index 70b8c59c0cea4..9bf3830406156 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala @@ -1571,7 +1571,25 @@ class PlanParserSuite extends AnalysisTest { test("named parameters") { comparePlans( - parsePlan("SELECT @param1"), - Project(UnresolvedAlias(UnboundParameter("@param1"), None) :: Nil, OneRowRelation())) + parsePlan("SELECT @param_1"), + Project(UnresolvedAlias(NamedParameter("param_1"), None) :: Nil, OneRowRelation())) + comparePlans( + parsePlan("SELECT abs(@1Abc)"), + Project(UnresolvedAlias( + UnresolvedFunction( + "abs" :: Nil, + NamedParameter("1Abc") :: Nil, + isDistinct = false), None) :: Nil, + OneRowRelation())) + comparePlans( + parsePlan("SELECT * FROM a LIMIT @limitA"), + table("a").select(star()).limit(NamedParameter("limitA"))) + // Invalid empty name and invalid symbol in a name + Seq("@", "@-").foreach { name => + checkError( + exception = parseException(s"SELECT $name"), + errorClass = "PARSE_SYNTAX_ERROR", + parameters = Map("error" -> "'@'", "hint" -> "")) + } } } From 0f32b753781cb11bc3fcc2db9ce20485a2af32dc Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Tue, 22 Nov 2022 21:03:07 +0300 Subject: [PATCH 04/38] NamedParameter should be bound --- core/src/main/resources/error/error-classes.json | 5 +++++ .../spark/sql/catalyst/analysis/CheckAnalysis.scala | 5 +++++ .../spark/sql/catalyst/expressions/parameters.scala | 4 ++-- .../sql/errors/QueryCompilationErrorsSuite.scala | 13 +++++++++++++ 4 files changed, 25 insertions(+), 2 deletions(-) diff --git a/core/src/main/resources/error/error-classes.json b/core/src/main/resources/error/error-classes.json index 4da9d2f9fbcac..90921f5b341e7 100644 --- a/core/src/main/resources/error/error-classes.json +++ b/core/src/main/resources/error/error-classes.json @@ -1054,6 +1054,11 @@ "Unable to convert SQL type to Protobuf type ." ] }, + "UNBOUND_PARAMETER" : { + "message" : [ + "Found the unbound parameter: . Use `bind()` to substitute the parameter by a literal." + ] + }, "UNCLOSED_BRACKETED_COMMENT" : { "message" : [ "Found an unclosed bracketed comment. Please, append */ at the end of the comment." diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/CheckAnalysis.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/CheckAnalysis.scala index aecf36660cd19..2c4d641c601ef 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/CheckAnalysis.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/CheckAnalysis.scala @@ -318,6 +318,11 @@ trait CheckAnalysis extends PredicateHelper with LookupCatalog with QueryErrorsB errorClass = "_LEGACY_ERROR_TEMP_2413", messageParameters = Map("argName" -> e.prettyName)) + case p: NamedParameter => + p.failAnalysis( + errorClass = "UNBOUND_PARAMETER", + messageParameters = Map("name" -> p.name)) + case _ => }) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index f07aaa969735b..523300cac255f 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -33,10 +33,10 @@ case class NamedParameter(name: String) extends LeafExpression { override def nullable: Boolean = true override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { - throw SparkException.internalError(s"Found the unbound parameter: $name") + throw SparkException.internalError(s"Found the unbound parameter: $name.") } def eval(input: InternalRow): Any = { - throw SparkException.internalError(s"Found the unbound parameter: $name") + throw SparkException.internalError(s"Found the unbound parameter: $name.") } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala index bed647ef49fcd..ce3e45cc07451 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala @@ -667,6 +667,19 @@ class QueryCompilationErrorsSuite errorClass = "DATATYPE_MISMATCH.INVALID_JSON_SCHEMA", parameters = Map("schema" -> "\"INT\"", "sqlExpr" -> "\"from_json(a)\"")) } + + test("UNBOUND_PARAMETER: named parameters should be substituted") { + checkError( + exception = intercept[AnalysisException] { + sql("select @abc").collect() + }, + errorClass = "UNBOUND_PARAMETER", + parameters = Map("name" -> "abc"), + context = ExpectedContext( + fragment = "@abc", + start = 7, + stop = 10)) + } } class MyCastToString extends SparkUserDefinedFunction( From 2806334d0092446ae8c921976028c867a5105275 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Wed, 23 Nov 2022 21:14:13 +0300 Subject: [PATCH 05/38] Add new method bind() to Dataset --- .../spark/sql/catalyst/plans/logical/object.scala | 10 ++++++++++ .../spark/sql/catalyst/trees/TreePatterns.scala | 1 + .../main/scala/org/apache/spark/sql/Dataset.scala | 13 ++++++++++++- 3 files changed, 23 insertions(+), 1 deletion(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/object.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/object.scala index e5fe07e2d950d..64355d6051c6e 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/object.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/object.scala @@ -690,3 +690,13 @@ case class CoGroup( override protected def withNewChildrenInternal( newLeft: LogicalPlan, newRight: LogicalPlan): CoGroup = copy(left = newLeft, right = newRight) } + +case class Bind(args: Map[String, Expression], child: LogicalPlan) extends UnaryNode { + + override def output: Seq[Attribute] = child.output + + final override val nodePatterns: Seq[TreePattern] = Seq(BIND) + + override protected def withNewChildInternal(newChild: LogicalPlan): Bind = + copy(child = newChild) +} diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala index 8fca9ec60cdff..34770723c169e 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala @@ -95,6 +95,7 @@ object TreePattern extends Enumeration { // Logical plan patterns (alphabetically ordered) val AGGREGATE: Value = Value val AS_OF_JOIN: Value = Value + val BIND: Value = Value val COMMAND: Value = Value val CTE: Value = Value val DISTINCT_LIKE: Value = Value diff --git a/sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala b/sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala index 5215858a9fd5c..64d2f8269cd69 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala @@ -28,7 +28,7 @@ import scala.util.control.NonFatal import org.apache.commons.lang3.StringUtils import org.apache.spark.TaskContext -import org.apache.spark.annotation.{DeveloperApi, Stable, Unstable} +import org.apache.spark.annotation.{DeveloperApi, Experimental, Stable, Unstable} import org.apache.spark.api.java.JavaRDD import org.apache.spark.api.java.function._ import org.apache.spark.api.python.{PythonRDD, SerDeUtil} @@ -3940,6 +3940,17 @@ class Dataset[T] private[sql]( files.toSet.toArray } + /** + * Bind query parameters to literal values. + * + * @param args A map of parameter names to their values and types. + * @since 3.4.0 + */ + @Experimental + def bind(args: Map[String, (Any, DataType)]): Dataset[T] = withTypedPlan { + Bind(args.mapValues { case (v, dt) => Literal.create(v, dt) }, logicalPlan) + } + /** * Returns `true` when the logical query plans inside both [[Dataset]]s are equal and * therefore return same results. From d16e85bac8c590d30a8c67d90e51fa90fef0d9ae Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Thu, 24 Nov 2022 00:36:41 +0300 Subject: [PATCH 06/38] Return back Literal and fix for scala 2.13 --- .../org/apache/spark/sql/catalyst/plans/logical/object.scala | 2 +- sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/object.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/object.scala index 64355d6051c6e..e988c02efeec8 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/object.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/object.scala @@ -691,7 +691,7 @@ case class CoGroup( newLeft: LogicalPlan, newRight: LogicalPlan): CoGroup = copy(left = newLeft, right = newRight) } -case class Bind(args: Map[String, Expression], child: LogicalPlan) extends UnaryNode { +case class Bind(args: Map[String, Literal], child: LogicalPlan) extends UnaryNode { override def output: Seq[Attribute] = child.output diff --git a/sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala b/sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala index 64d2f8269cd69..e2b24b437dae1 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala @@ -3948,7 +3948,7 @@ class Dataset[T] private[sql]( */ @Experimental def bind(args: Map[String, (Any, DataType)]): Dataset[T] = withTypedPlan { - Bind(args.mapValues { case (v, dt) => Literal.create(v, dt) }, logicalPlan) + Bind(args.mapValues { case (v, dt) => Literal.create(v, dt) }.toMap, logicalPlan) } /** From e48e2e1680c61590e62020bf06d498f80fad23a8 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Thu, 24 Nov 2022 19:40:28 +0300 Subject: [PATCH 07/38] Add the BindParameters rule --- .../sql/catalyst/analysis/Analyzer.scala | 19 ++++++++++++++++++- .../sql/catalyst/expressions/parameters.scala | 3 +++ .../sql/catalyst/rules/RuleIdCollection.scala | 1 + .../sql/catalyst/trees/TreePatterns.scala | 1 + .../sql/catalyst/analysis/AnalysisSuite.scala | 9 +++++++++ 5 files changed, 32 insertions(+), 1 deletion(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala index 104c5c1e08053..ef58a96466e27 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala @@ -267,7 +267,8 @@ class Analyzer(override val catalogManager: CatalogManager) CTESubstitution, WindowsSubstitution, EliminateUnions, - SubstituteUnresolvedOrdinals), + SubstituteUnresolvedOrdinals, + BindParameters), Batch("Disable Hints", Once, new ResolveHints.DisableHints), Batch("Hints", fixedPoint, @@ -4125,3 +4126,19 @@ object RemoveTempResolvedColumn extends Rule[LogicalPlan] { CurrentOrigin.withOrigin(t.origin)(UnresolvedAttribute(t.nameParts)) } } + +/** + * The rule `BindParameters` in the `Substitution` batch substitutes named parameters by + * theirs linked literals. + */ +object BindParameters extends Rule[LogicalPlan] { + override def apply(plan: LogicalPlan): LogicalPlan = { + plan.resolveOperatorsUpWithPruning(_.containsPattern(BIND), ruleId) { + case Bind(args, child) => + child.transformAllExpressionsWithPruning(_.containsPattern(PARAMETER), ruleId) { + case NamedParameter(name) if args.contains(name) => + args(name) + } + } + } +} diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index 523300cac255f..78fb9cf3c7440 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -20,6 +20,7 @@ package org.apache.spark.sql.catalyst.expressions import org.apache.spark.SparkException import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, ExprCode} +import org.apache.spark.sql.catalyst.trees.TreePattern.{PARAMETER, TreePattern} import org.apache.spark.sql.types.{DataType, NullType} /** @@ -32,6 +33,8 @@ case class NamedParameter(name: String) extends LeafExpression { override def dataType: DataType = NullType override def nullable: Boolean = true + final override val nodePatterns: Seq[TreePattern] = Seq(PARAMETER) + override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { throw SparkException.internalError(s"Found the unbound parameter: $name.") } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/rules/RuleIdCollection.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/rules/RuleIdCollection.scala index f6bef88ab868e..665bdc293f63d 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/rules/RuleIdCollection.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/rules/RuleIdCollection.scala @@ -79,6 +79,7 @@ object RuleIdCollection { "org.apache.spark.sql.catalyst.analysis.Analyzer$WindowsSubstitution" :: "org.apache.spark.sql.catalyst.analysis.AnsiTypeCoercion$AnsiCombinedTypeCoercionRule" :: "org.apache.spark.sql.catalyst.analysis.ApplyCharTypePadding" :: + "org.apache.spark.sql.catalyst.analysis.BindParameters" :: "org.apache.spark.sql.catalyst.analysis.DeduplicateRelations" :: "org.apache.spark.sql.catalyst.analysis.EliminateSubqueryAliases" :: "org.apache.spark.sql.catalyst.analysis.EliminateUnions" :: diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala index 34770723c169e..f43c4443d7ef5 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala @@ -70,6 +70,7 @@ object TreePattern extends Enumeration { val NULL_LITERAL: Value = Value val SERIALIZE_FROM_OBJECT: Value = Value val OUTER_REFERENCE: Value = Value + val PARAMETER: Value = Value val PIVOT: Value = Value val PLAN_EXPRESSION: Value = Value val PYTHON_UDF: Value = Value diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala index 8b303ec3bb1fe..52bd66e67dc1f 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala @@ -1295,4 +1295,13 @@ class AnalysisSuite extends AnalysisTest with Matchers { assertAnalysisSuccess(finalPlan) } + + test("bind named parameters to literals") { + val plan = Bind( + args = Map("limitA" -> Literal(10)), + parsePlan("SELECT * FROM a LIMIT @limitA")) + comparePlans( + BindParameters.apply(plan), + parsePlan("SELECT * FROM a LIMIT 10")) + } } From 693cd92adc3b562a9b59802cde0a916512ec47af Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Mon, 28 Nov 2022 22:31:54 +0300 Subject: [PATCH 08/38] Add an end-to-end test --- .../spark/sql/catalyst/analysis/CheckAnalysis.scala | 5 ----- .../spark/sql/catalyst/optimizer/Optimizer.scala | 3 ++- .../sql/catalyst/optimizer/finishAnalysis.scala | 13 +++++++++++++ .../spark/sql/catalyst/analysis/AnalysisSuite.scala | 2 +- .../spark/sql/catalyst/parser/PlanParserSuite.scala | 2 +- .../scala/org/apache/spark/sql/DatasetSuite.scala | 10 ++++++++++ .../sql/errors/QueryCompilationErrorsSuite.scala | 2 +- 7 files changed, 28 insertions(+), 9 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/CheckAnalysis.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/CheckAnalysis.scala index ed633f0ef48a9..12dac5c632a3b 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/CheckAnalysis.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/CheckAnalysis.scala @@ -325,11 +325,6 @@ trait CheckAnalysis extends PredicateHelper with LookupCatalog with QueryErrorsB errorClass = "_LEGACY_ERROR_TEMP_2413", messageParameters = Map("argName" -> e.prettyName)) - case p: NamedParameter => - p.failAnalysis( - errorClass = "UNBOUND_PARAMETER", - messageParameters = Map("name" -> p.name)) - case _ => }) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala index ecb93f6b2390f..aedaf94c35ac3 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala @@ -292,7 +292,8 @@ abstract class Optimizer(catalogManager: CatalogManager) ComputeCurrentTime, ReplaceCurrentLike(catalogManager), SpecialDatetimeValues, - RewriteAsOfJoin) + RewriteAsOfJoin, + CheckUnboundParameters) override def apply(plan: LogicalPlan): LogicalPlan = { rules.foldLeft(plan) { case (sp, rule) => rule.apply(sp) } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/finishAnalysis.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/finishAnalysis.scala index 466781fa1def7..cb16749b34441 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/finishAnalysis.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/finishAnalysis.scala @@ -19,6 +19,7 @@ package org.apache.spark.sql.catalyst.optimizer import java.time.{Instant, LocalDateTime} +import org.apache.spark.sql.AnalysisException import org.apache.spark.sql.catalyst.CurrentUserContext.CURRENT_USER import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.plans.logical._ @@ -142,3 +143,15 @@ object SpecialDatetimeValues extends Rule[LogicalPlan] { } } } + +object CheckUnboundParameters extends Rule[LogicalPlan] { + override def apply(plan: LogicalPlan): LogicalPlan = { + plan.transformAllExpressionsWithPruning(_.containsPattern(PARAMETER)) { + case param @ NamedParameter(name) => + throw new AnalysisException( + errorClass = "UNBOUND_PARAMETER", + messageParameters = Map("name" -> name), + origin = param.origin) + } + } +} diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala index 52bd66e67dc1f..923254328dcb6 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala @@ -1296,7 +1296,7 @@ class AnalysisSuite extends AnalysisTest with Matchers { assertAnalysisSuccess(finalPlan) } - test("bind named parameters to literals") { + test("SPARK-41271: bind named parameters to literals") { val plan = Bind( args = Map("limitA" -> Literal(10)), parsePlan("SELECT * FROM a LIMIT @limitA")) diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala index 9bf3830406156..9064998030bc9 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala @@ -1569,7 +1569,7 @@ class PlanParserSuite extends AnalysisTest { ) } - test("named parameters") { + test("SPARK-41271: parsing of named parameters") { comparePlans( parsePlan("SELECT @param_1"), Project(UnresolvedAlias(NamedParameter("param_1"), None) :: Nil, OneRowRelation())) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index d298d7129c70d..72f48bd0547f8 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -2245,6 +2245,16 @@ class DatasetSuite extends QueryTest assert(parquetFiles.size === 10) } } + + test("SPARK-41271: bind parameters") { + val input = spark.range(10) + .selectExpr("id", "id % @div as c0") + .where("c0 = @constA") + val df = input.bind(Map( + "div" -> (3, IntegerType), + "constA" -> (1L, LongType))) + checkAnswer(df, Row(1, 1) :: Row(4, 1) :: Row(7, 1) :: Nil) + } } class DatasetLargeResultCollectingSuite extends QueryTest diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala index ce3e45cc07451..e1cea2dc5d5b3 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala @@ -668,7 +668,7 @@ class QueryCompilationErrorsSuite parameters = Map("schema" -> "\"INT\"", "sqlExpr" -> "\"from_json(a)\"")) } - test("UNBOUND_PARAMETER: named parameters should be substituted") { + test("UNBOUND_PARAMETER - SPARK-41271: non-substituted parameters") { checkError( exception = intercept[AnalysisException] { sql("select @abc").collect() From 4e55269bf757b871042846f3148ed9f2d7e96cdc Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Wed, 30 Nov 2022 23:00:28 +0300 Subject: [PATCH 09/38] Add a config to control the feature --- .../sql/catalyst/parser/SqlBaseParser.g4 | 7 ++- .../sql/catalyst/analysis/Analyzer.scala | 15 +++--- .../sql/catalyst/parser/ParseDriver.scala | 1 + .../apache/spark/sql/internal/SQLConf.scala | 10 ++++ .../sql/catalyst/analysis/AnalysisSuite.scala | 14 +++--- .../sql/catalyst/parser/PlanParserSuite.scala | 46 +++++++++++-------- .../org/apache/spark/sql/DatasetSuite.scala | 16 ++++--- 7 files changed, 71 insertions(+), 38 deletions(-) diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 index 5d7480967825e..87c3e142abcc5 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 @@ -40,6 +40,11 @@ options { tokenVocab = SqlBaseLexer; } * When true, double quoted literals are identifiers rather than STRINGs. */ public boolean double_quoted_identifiers = false; + + /** + * When true, identifiers that begin from `@` are considered as named parameters. + */ + public boolean parameters_enabled = false; } singleStatement @@ -930,7 +935,7 @@ primaryExpression constant : NULL #nullLiteral - | PARAMETER #parameterLiteral + | {parameters_enabled}? PARAMETER #parameterLiteral | interval #intervalLiteral | identifier stringLit #typeConstructor | number #numericLiteral diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala index 89f707ad175d8..27d3fd1f4c489 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala @@ -4133,12 +4133,15 @@ object RemoveTempResolvedColumn extends Rule[LogicalPlan] { */ object BindParameters extends Rule[LogicalPlan] { override def apply(plan: LogicalPlan): LogicalPlan = { - plan.resolveOperatorsUpWithPruning(_.containsPattern(BIND), ruleId) { - case Bind(args, child) => - child.transformAllExpressionsWithPruning(_.containsPattern(PARAMETER), ruleId) { - case NamedParameter(name) if args.contains(name) => - args(name) - } + if (SQLConf.get.parametersEnabled) { + plan.resolveOperatorsUpWithPruning(_.containsPattern(BIND), ruleId) { + case Bind(args, child) => + child.transformAllExpressionsWithPruning(_.containsPattern(PARAMETER), ruleId) { + case NamedParameter(name) if args.contains(name) => args(name) + } + } + } else { + plan } } } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseDriver.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseDriver.scala index 727d35d5c9152..31aa4ba537991 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseDriver.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseDriver.scala @@ -119,6 +119,7 @@ abstract class AbstractSqlParser extends ParserInterface with SQLConfHelper with parser.legacy_exponent_literal_as_decimal_enabled = conf.exponentLiteralAsDecimalEnabled parser.SQL_standard_keyword_behavior = conf.enforceReservedKeywords parser.double_quoted_identifiers = conf.doubleQuotedIdentifiers + parser.parameters_enabled = conf.parametersEnabled try { try { diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala index 84d78f365acbc..eae3841c64649 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala @@ -4027,6 +4027,14 @@ object SQLConf { .checkValues(ErrorMessageFormat.values.map(_.toString)) .createWithDefault(ErrorMessageFormat.PRETTY.toString) + val PARAMETERS_ENABLED = buildConf("spark.sql.parameters.enabled") + .doc("When set to true, queries can have named parameters that should be substituted " + + "by literal values later using `bind()`. If set to false, Spark handles constants " + + "with the `@` prefix as regular identifiers and does not consider them as parameters.") + .version("3.4.0") + .booleanConf + .createWithDefault(true) + /** * Holds information about keys that have been deprecated. * @@ -4838,6 +4846,8 @@ class SQLConf extends Serializable with Logging { def allowsTempViewCreationWithMultipleNameparts: Boolean = getConf(SQLConf.ALLOW_TEMP_VIEW_CREATION_WITH_MULTIPLE_NAME_PARTS) + def parametersEnabled: Boolean = getConf(SQLConf.PARAMETERS_ENABLED) + /** ********************** SQLConf functionality methods ************ */ /** Set Spark SQL configuration properties. */ diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala index 923254328dcb6..f409643602c56 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala @@ -1297,11 +1297,13 @@ class AnalysisSuite extends AnalysisTest with Matchers { } test("SPARK-41271: bind named parameters to literals") { - val plan = Bind( - args = Map("limitA" -> Literal(10)), - parsePlan("SELECT * FROM a LIMIT @limitA")) - comparePlans( - BindParameters.apply(plan), - parsePlan("SELECT * FROM a LIMIT 10")) + withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { + val plan = Bind( + args = Map("limitA" -> Literal(10)), + parsePlan("SELECT * FROM a LIMIT @limitA")) + comparePlans( + BindParameters.apply(plan), + parsePlan("SELECT * FROM a LIMIT 10")) + } } } diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala index 9064998030bc9..dced97ceb2389 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala @@ -1570,26 +1570,36 @@ class PlanParserSuite extends AnalysisTest { } test("SPARK-41271: parsing of named parameters") { - comparePlans( - parsePlan("SELECT @param_1"), - Project(UnresolvedAlias(NamedParameter("param_1"), None) :: Nil, OneRowRelation())) - comparePlans( - parsePlan("SELECT abs(@1Abc)"), - Project(UnresolvedAlias( - UnresolvedFunction( - "abs" :: Nil, - NamedParameter("1Abc") :: Nil, - isDistinct = false), None) :: Nil, - OneRowRelation())) - comparePlans( - parsePlan("SELECT * FROM a LIMIT @limitA"), - table("a").select(star()).limit(NamedParameter("limitA"))) - // Invalid empty name and invalid symbol in a name - Seq("@", "@-").foreach { name => + withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { + comparePlans( + parsePlan("SELECT @param_1"), + Project(UnresolvedAlias(NamedParameter("param_1"), None) :: Nil, OneRowRelation())) + comparePlans( + parsePlan("SELECT abs(@1Abc)"), + Project(UnresolvedAlias( + UnresolvedFunction( + "abs" :: Nil, + NamedParameter("1Abc") :: Nil, + isDistinct = false), None) :: Nil, + OneRowRelation())) + comparePlans( + parsePlan("SELECT * FROM a LIMIT @limitA"), + table("a").select(star()).limit(NamedParameter("limitA"))) + // Invalid empty name and invalid symbol in a name + Seq("@", "@-").foreach { name => + checkError( + exception = parseException(s"SELECT $name"), + errorClass = "PARSE_SYNTAX_ERROR", + parameters = Map("error" -> "'@'", "hint" -> "")) + } + } + withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "false") { checkError( - exception = parseException(s"SELECT $name"), + exception = intercept[ParseException] { + parsePlan("SELECT @param_1") + }, errorClass = "PARSE_SYNTAX_ERROR", - parameters = Map("error" -> "'@'", "hint" -> "")) + parameters = Map("error" -> "'@param_1'", "hint" -> "")) } } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index 72f48bd0547f8..6e625fdbed8a3 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -2247,13 +2247,15 @@ class DatasetSuite extends QueryTest } test("SPARK-41271: bind parameters") { - val input = spark.range(10) - .selectExpr("id", "id % @div as c0") - .where("c0 = @constA") - val df = input.bind(Map( - "div" -> (3, IntegerType), - "constA" -> (1L, LongType))) - checkAnswer(df, Row(1, 1) :: Row(4, 1) :: Row(7, 1) :: Nil) + withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { + val input = spark.range(10) + .selectExpr("id", "id % @div as c0") + .where("c0 = @constA") + val df = input.bind(Map( + "div" -> (3, IntegerType), + "constA" -> (1L, LongType))) + checkAnswer(df, Row(1, 1) :: Row(4, 1) :: Row(7, 1) :: Nil) + } } } From fcdf5f249f91c828bac5de046a8a16295b44af4e Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Thu, 1 Dec 2022 22:03:24 +0300 Subject: [PATCH 10/38] Support parametrized SQL queries by sql() --- .../sql/catalyst/analysis/Analyzer.scala | 34 ++++++++++++------- .../sql/catalyst/optimizer/Optimizer.scala | 3 +- .../catalyst/optimizer/finishAnalysis.scala | 13 ------- .../sql/catalyst/analysis/AnalysisSuite.scala | 7 ++-- .../scala/org/apache/spark/sql/Dataset.scala | 13 +------ .../org/apache/spark/sql/SparkSession.scala | 27 +++++++++++---- .../org/apache/spark/sql/DatasetSuite.scala | 17 ++++++---- .../errors/QueryCompilationErrorsSuite.scala | 10 +++--- .../apache/spark/sql/test/SQLTestUtils.scala | 2 +- .../InsertIntoHiveTableBenchmark.scala | 4 +-- .../ObjectHashAggregateExecBenchmark.scala | 4 +-- 11 files changed, 68 insertions(+), 66 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala index a2fb182cfa132..86d9c3532146b 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala @@ -48,7 +48,7 @@ import org.apache.spark.sql.connector.catalog.CatalogV2Implicits._ import org.apache.spark.sql.connector.catalog.TableChange.{After, ColumnPosition} import org.apache.spark.sql.connector.catalog.functions.{AggregateFunction => V2AggregateFunction, ScalarFunction, UnboundFunction} import org.apache.spark.sql.connector.expressions.{FieldReference, IdentityTransform, Transform} -import org.apache.spark.sql.errors.{QueryCompilationErrors, QueryExecutionErrors} +import org.apache.spark.sql.errors.{QueryCompilationErrors, QueryErrorsBase, QueryExecutionErrors} import org.apache.spark.sql.execution.datasources.v2.DataSourceV2Relation import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.internal.SQLConf.{PartitionOverwriteMode, StoreAssignmentPolicy} @@ -268,8 +268,7 @@ class Analyzer(override val catalogManager: CatalogManager) CTESubstitution, WindowsSubstitution, EliminateUnions, - SubstituteUnresolvedOrdinals, - BindParameters), + SubstituteUnresolvedOrdinals), Batch("Disable Hints", Once, new ResolveHints.DisableHints), Batch("Hints", fixedPoint, @@ -4136,16 +4135,27 @@ object RemoveTempResolvedColumn extends Rule[LogicalPlan] { } /** - * The rule `BindParameters` in the `Substitution` batch substitutes named parameters by - * theirs linked literals. + * Finds all named parameters in the given plan and substitudes them by literal values + * evaluated from `args` values. */ -object BindParameters extends Rule[LogicalPlan] { - override def apply(plan: LogicalPlan): LogicalPlan = { - if (SQLConf.get.parametersEnabled) { - plan.resolveOperatorsUpWithPruning(_.containsPattern(BIND), ruleId) { - case Bind(args, child) => - child.transformAllExpressionsWithPruning(_.containsPattern(PARAMETER), ruleId) { - case NamedParameter(name) if args.contains(name) => args(name) +object BindParameters extends QueryErrorsBase { + def apply(plan: LogicalPlan, args: Map[String, Expression]): LogicalPlan = { + if (!args.isEmpty && SQLConf.get.parametersEnabled) { + args.filter(!_._2.foldable).headOption.foreach { case (name, expr) => + expr.failAnalysis( + errorClass = "NON_FOLDABLE_SQL_ARG", + messageParameters = Map( + "name" -> name, + "expr" -> toSQLExpr(expr))) + } + plan.transformAllExpressionsWithPruning(_.containsPattern(PARAMETER)) { + case param @ NamedParameter(name) => + if (args.contains(name)) { + args(name) + } else { + param.failAnalysis( + errorClass = "UNBOUND_PARAMETER", + messageParameters = Map("name" -> name)) } } } else { diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala index 0b3ab9289fc04..1f0fb66775366 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/Optimizer.scala @@ -292,8 +292,7 @@ abstract class Optimizer(catalogManager: CatalogManager) ComputeCurrentTime, ReplaceCurrentLike(catalogManager), SpecialDatetimeValues, - RewriteAsOfJoin, - CheckUnboundParameters) + RewriteAsOfJoin) override def apply(plan: LogicalPlan): LogicalPlan = { rules.foldLeft(plan) { case (sp, rule) => rule.apply(sp) } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/finishAnalysis.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/finishAnalysis.scala index cb16749b34441..466781fa1def7 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/finishAnalysis.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/finishAnalysis.scala @@ -19,7 +19,6 @@ package org.apache.spark.sql.catalyst.optimizer import java.time.{Instant, LocalDateTime} -import org.apache.spark.sql.AnalysisException import org.apache.spark.sql.catalyst.CurrentUserContext.CURRENT_USER import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.plans.logical._ @@ -143,15 +142,3 @@ object SpecialDatetimeValues extends Rule[LogicalPlan] { } } } - -object CheckUnboundParameters extends Rule[LogicalPlan] { - override def apply(plan: LogicalPlan): LogicalPlan = { - plan.transformAllExpressionsWithPruning(_.containsPattern(PARAMETER)) { - case param @ NamedParameter(name) => - throw new AnalysisException( - errorClass = "UNBOUND_PARAMETER", - messageParameters = Map("name" -> name), - origin = param.origin) - } - } -} diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala index f409643602c56..45117ced14adc 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala @@ -1298,11 +1298,10 @@ class AnalysisSuite extends AnalysisTest with Matchers { test("SPARK-41271: bind named parameters to literals") { withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { - val plan = Bind( - args = Map("limitA" -> Literal(10)), - parsePlan("SELECT * FROM a LIMIT @limitA")) comparePlans( - BindParameters.apply(plan), + BindParameters( + plan = parsePlan("SELECT * FROM a LIMIT @limitA"), + args = Map("limitA" -> Literal(10))), parsePlan("SELECT * FROM a LIMIT 10")) } } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala b/sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala index 6010afa3936fc..5f6512d4e4b07 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/Dataset.scala @@ -28,7 +28,7 @@ import scala.util.control.NonFatal import org.apache.commons.lang3.StringUtils import org.apache.spark.TaskContext -import org.apache.spark.annotation.{DeveloperApi, Experimental, Stable, Unstable} +import org.apache.spark.annotation.{DeveloperApi, Stable, Unstable} import org.apache.spark.api.java.JavaRDD import org.apache.spark.api.java.function._ import org.apache.spark.api.python.{PythonRDD, SerDeUtil} @@ -3938,17 +3938,6 @@ class Dataset[T] private[sql]( files.toSet.toArray } - /** - * Bind query parameters to literal values. - * - * @param args A map of parameter names to their values and types. - * @since 3.4.0 - */ - @Experimental - def bind(args: Map[String, (Any, DataType)]): Dataset[T] = withTypedPlan { - Bind(args.mapValues { case (v, dt) => Literal.create(v, dt) }.toMap, logicalPlan) - } - /** * Returns `true` when the logical query plans inside both [[Dataset]]s are equal and * therefore return same results. diff --git a/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala b/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala index 3d9f16799577a..bdd323f3f43f8 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala @@ -35,7 +35,7 @@ import org.apache.spark.rdd.RDD import org.apache.spark.scheduler.{SparkListener, SparkListenerApplicationEnd} import org.apache.spark.sql.catalog.Catalog import org.apache.spark.sql.catalyst._ -import org.apache.spark.sql.catalyst.analysis.UnresolvedRelation +import org.apache.spark.sql.catalyst.analysis.{BindParameters, UnresolvedRelation} import org.apache.spark.sql.catalyst.encoders._ import org.apache.spark.sql.catalyst.expressions.AttributeReference import org.apache.spark.sql.catalyst.plans.logical.{LocalRelation, Range} @@ -608,20 +608,35 @@ class SparkSession private( | Everything else | * ----------------- */ - /** - * Executes a SQL query using Spark, returning the result as a `DataFrame`. + /** + * Executes a SQL query substituting named parameters by the given arguments, + * returning the result as a `DataFrame`. * This API eagerly runs DDL/DML commands, but not for SELECT queries. * - * @since 2.0.0 + * @param sqlText A SQL statement with named parameters to execute. + * @param args A map of parameter names to typed literals. + * + * @since 3.4.0 */ - def sql(sqlText: String): DataFrame = withActive { + @Experimental + def sql(sqlText: String, args: Map[String, String]): DataFrame = withActive { val tracker = new QueryPlanningTracker val plan = tracker.measurePhase(QueryPlanningTracker.PARSING) { - sessionState.sqlParser.parsePlan(sqlText) + val parser = sessionState.sqlParser + val parsedArgs = args.mapValues(parser.parseExpression).toMap + BindParameters(parser.parsePlan(sqlText), parsedArgs) } Dataset.ofRows(self, plan, tracker) } + /** + * Executes a SQL query using Spark, returning the result as a `DataFrame`. + * This API eagerly runs DDL/DML commands, but not for SELECT queries. + * + * @since 2.0.0 + */ + def sql(sqlText: String): DataFrame = sql(sqlText, Map.empty) + /** * Execute an arbitrary string command inside an external execution engine rather than Spark. * This could be useful when user wants to execute some commands out of Spark. For diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index 6e625fdbed8a3..6e454918e9b63 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -2248,13 +2248,16 @@ class DatasetSuite extends QueryTest test("SPARK-41271: bind parameters") { withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { - val input = spark.range(10) - .selectExpr("id", "id % @div as c0") - .where("c0 = @constA") - val df = input.bind(Map( - "div" -> (3, IntegerType), - "constA" -> (1L, LongType))) - checkAnswer(df, Row(1, 1) :: Row(4, 1) :: Row(7, 1) :: Nil) + val sqlText = + """ + |SELECT id, id % @div as c0 + |FROM VALUES (0), (1), (2), (3), (4), (5), (6), (7), (8), (9) AS t(id) + |WHERE id < @constA + |""".stripMargin + val args = Map("div" -> "3", "constA" -> "4L") + checkAnswer( + spark.sql(sqlText, args), + Row(0, 0) :: Row(1, 1) :: Row(2, 2) :: Row(3, 0) :: Nil) } } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala index e1cea2dc5d5b3..56b90c263d031 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala @@ -671,14 +671,14 @@ class QueryCompilationErrorsSuite test("UNBOUND_PARAMETER - SPARK-41271: non-substituted parameters") { checkError( exception = intercept[AnalysisException] { - sql("select @abc").collect() + spark.sql("select @abc, @def", Map("abc" -> "1")) }, errorClass = "UNBOUND_PARAMETER", - parameters = Map("name" -> "abc"), + parameters = Map("name" -> "def"), context = ExpectedContext( - fragment = "@abc", - start = 7, - stop = 10)) + fragment = "@def", + start = 13, + stop = 16)) } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/test/SQLTestUtils.scala b/sql/core/src/test/scala/org/apache/spark/sql/test/SQLTestUtils.scala index ae425419c5407..dd55fcfe42cac 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/test/SQLTestUtils.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/test/SQLTestUtils.scala @@ -229,7 +229,7 @@ private[sql] trait SQLTestUtilsBase protected def sparkContext = spark.sparkContext // Shorthand for running a query using our SQLContext - protected lazy val sql = spark.sql _ + protected lazy val sql: String => DataFrame = spark.sql _ /** * A helper object for importing SQL implicits. diff --git a/sql/hive/src/test/scala/org/apache/spark/sql/execution/benchmark/InsertIntoHiveTableBenchmark.scala b/sql/hive/src/test/scala/org/apache/spark/sql/execution/benchmark/InsertIntoHiveTableBenchmark.scala index 7634598569891..b64b7823acd54 100644 --- a/sql/hive/src/test/scala/org/apache/spark/sql/execution/benchmark/InsertIntoHiveTableBenchmark.scala +++ b/sql/hive/src/test/scala/org/apache/spark/sql/execution/benchmark/InsertIntoHiveTableBenchmark.scala @@ -18,7 +18,7 @@ package org.apache.spark.sql.execution.benchmark import org.apache.spark.benchmark.Benchmark -import org.apache.spark.sql.SparkSession +import org.apache.spark.sql.{DataFrame, SparkSession} import org.apache.spark.sql.hive.test.TestHive /** @@ -40,7 +40,7 @@ object InsertIntoHiveTableBenchmark extends SqlBasedBenchmark { val tempView = "temp" val numRows = 1024 * 10 - val sql = spark.sql _ + val sql: String => DataFrame = spark.sql _ // scalastyle:off hadoopconfiguration private val hadoopConf = spark.sparkContext.hadoopConfiguration diff --git a/sql/hive/src/test/scala/org/apache/spark/sql/execution/benchmark/ObjectHashAggregateExecBenchmark.scala b/sql/hive/src/test/scala/org/apache/spark/sql/execution/benchmark/ObjectHashAggregateExecBenchmark.scala index 5d0a5ce09571a..1a4700e7445b6 100644 --- a/sql/hive/src/test/scala/org/apache/spark/sql/execution/benchmark/ObjectHashAggregateExecBenchmark.scala +++ b/sql/hive/src/test/scala/org/apache/spark/sql/execution/benchmark/ObjectHashAggregateExecBenchmark.scala @@ -22,7 +22,7 @@ import scala.concurrent.duration._ import org.apache.hadoop.hive.ql.udf.generic.GenericUDAFPercentileApprox import org.apache.spark.benchmark.Benchmark -import org.apache.spark.sql.{Column, SparkSession} +import org.apache.spark.sql.{Column, DataFrame, SparkSession} import org.apache.spark.sql.catalyst.expressions.Literal import org.apache.spark.sql.catalyst.expressions.aggregate.ApproximatePercentile import org.apache.spark.sql.hive.execution.TestingTypedCount @@ -46,7 +46,7 @@ object ObjectHashAggregateExecBenchmark extends SqlBasedBenchmark { override def getSparkSession: SparkSession = TestHive.sparkSession - private val sql = spark.sql _ + private val sql: String => DataFrame = spark.sql _ import spark.implicits._ private def hiveUDAFvsSparkAF(N: Int): Unit = { From 10c76801613a67db48772c735f6cd5e100f70e77 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Thu, 1 Dec 2022 22:57:51 +0300 Subject: [PATCH 11/38] Just strip any marker --- .../org/apache/spark/sql/catalyst/parser/AstBuilder.scala | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala index 21e7f16b62398..16153228289e1 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala @@ -4842,6 +4842,8 @@ class AstBuilder extends SqlBaseParserBaseVisitor[AnyRef] with SQLConfHelper wit * Create a named parameter which represents a literal with a non-bound value and unknown type. * */ override def visitParameterLiteral(ctx: ParameterLiteralContext): Expression = withOrigin(ctx) { - NamedParameter(ctx.getText.stripPrefix("@")) + val name = ctx.getText + assert(name.length > 1) + NamedParameter(name.substring(1)) } } From cddff4acca493fc55b356288fd93d6a6cfac19d1 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Thu, 1 Dec 2022 23:34:20 +0300 Subject: [PATCH 12/38] Refactoring --- .../sql/catalyst/analysis/Analyzer.scala | 32 +-------------- .../sql/catalyst/expressions/parameters.scala | 39 ++++++++++++++++++- .../sql/catalyst/parser/AstBuilder.scala | 2 +- .../sql/catalyst/analysis/AnalysisSuite.scala | 2 +- .../sql/catalyst/parser/PlanParserSuite.scala | 6 +-- .../org/apache/spark/sql/SparkSession.scala | 6 +-- 6 files changed, 46 insertions(+), 41 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala index 86d9c3532146b..7f66ddaa8942a 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/Analyzer.scala @@ -48,7 +48,7 @@ import org.apache.spark.sql.connector.catalog.CatalogV2Implicits._ import org.apache.spark.sql.connector.catalog.TableChange.{After, ColumnPosition} import org.apache.spark.sql.connector.catalog.functions.{AggregateFunction => V2AggregateFunction, ScalarFunction, UnboundFunction} import org.apache.spark.sql.connector.expressions.{FieldReference, IdentityTransform, Transform} -import org.apache.spark.sql.errors.{QueryCompilationErrors, QueryErrorsBase, QueryExecutionErrors} +import org.apache.spark.sql.errors.{QueryCompilationErrors, QueryExecutionErrors} import org.apache.spark.sql.execution.datasources.v2.DataSourceV2Relation import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.internal.SQLConf.{PartitionOverwriteMode, StoreAssignmentPolicy} @@ -4133,33 +4133,3 @@ object RemoveTempResolvedColumn extends Rule[LogicalPlan] { CurrentOrigin.withOrigin(t.origin)(UnresolvedAttribute(t.nameParts)) } } - -/** - * Finds all named parameters in the given plan and substitudes them by literal values - * evaluated from `args` values. - */ -object BindParameters extends QueryErrorsBase { - def apply(plan: LogicalPlan, args: Map[String, Expression]): LogicalPlan = { - if (!args.isEmpty && SQLConf.get.parametersEnabled) { - args.filter(!_._2.foldable).headOption.foreach { case (name, expr) => - expr.failAnalysis( - errorClass = "NON_FOLDABLE_SQL_ARG", - messageParameters = Map( - "name" -> name, - "expr" -> toSQLExpr(expr))) - } - plan.transformAllExpressionsWithPruning(_.containsPattern(PARAMETER)) { - case param @ NamedParameter(name) => - if (args.contains(name)) { - args(name) - } else { - param.failAnalysis( - errorClass = "UNBOUND_PARAMETER", - messageParameters = Map("name" -> name)) - } - } - } else { - plan - } - } -} diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index 78fb9cf3c7440..b8b69ad5305a3 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -19,17 +19,21 @@ package org.apache.spark.sql.catalyst.expressions import org.apache.spark.SparkException import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.analysis.AnalysisErrorAt import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, ExprCode} +import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan import org.apache.spark.sql.catalyst.trees.TreePattern.{PARAMETER, TreePattern} +import org.apache.spark.sql.errors.QueryErrorsBase +import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{DataType, NullType} /** * The expression represents a named parameter that should be bound later * to a literal with concrete value and type. * - * @param name The identifier of the parameter without the marker '@'. + * @param name The identifier of the parameter without the marker. */ -case class NamedParameter(name: String) extends LeafExpression { +case class Parameter(name: String) extends LeafExpression { override def dataType: DataType = NullType override def nullable: Boolean = true @@ -43,3 +47,34 @@ case class NamedParameter(name: String) extends LeafExpression { throw SparkException.internalError(s"Found the unbound parameter: $name.") } } + + +/** + * Finds all named parameters in the given plan and substitutes them by literal values + * evaluated from `args` values. + */ +object Parameter extends QueryErrorsBase { + def bind(plan: LogicalPlan, args: Map[String, Expression]): LogicalPlan = { + if (!args.isEmpty && SQLConf.get.parametersEnabled) { + args.filter(!_._2.foldable).headOption.foreach { case (name, expr) => + expr.failAnalysis( + errorClass = "NON_FOLDABLE_SQL_ARG", + messageParameters = Map( + "name" -> name, + "expr" -> toSQLExpr(expr))) + } + plan.transformAllExpressionsWithPruning(_.containsPattern(PARAMETER)) { + case param @ Parameter(name) => + if (args.contains(name)) { + args(name) + } else { + param.failAnalysis( + errorClass = "UNBOUND_PARAMETER", + messageParameters = Map("name" -> name)) + } + } + } else { + plan + } + } +} diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala index 16153228289e1..88b6fe7255905 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala @@ -4844,6 +4844,6 @@ class AstBuilder extends SqlBaseParserBaseVisitor[AnyRef] with SQLConfHelper wit override def visitParameterLiteral(ctx: ParameterLiteralContext): Expression = withOrigin(ctx) { val name = ctx.getText assert(name.length > 1) - NamedParameter(name.substring(1)) + Parameter(name.substring(1)) } } diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala index 45117ced14adc..a33721b8b418e 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala @@ -1299,7 +1299,7 @@ class AnalysisSuite extends AnalysisTest with Matchers { test("SPARK-41271: bind named parameters to literals") { withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { comparePlans( - BindParameters( + Parameter.bind( plan = parsePlan("SELECT * FROM a LIMIT @limitA"), args = Map("limitA" -> Literal(10))), parsePlan("SELECT * FROM a LIMIT 10")) diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala index dced97ceb2389..72b4357621f88 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala @@ -1573,18 +1573,18 @@ class PlanParserSuite extends AnalysisTest { withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { comparePlans( parsePlan("SELECT @param_1"), - Project(UnresolvedAlias(NamedParameter("param_1"), None) :: Nil, OneRowRelation())) + Project(UnresolvedAlias(Parameter("param_1"), None) :: Nil, OneRowRelation())) comparePlans( parsePlan("SELECT abs(@1Abc)"), Project(UnresolvedAlias( UnresolvedFunction( "abs" :: Nil, - NamedParameter("1Abc") :: Nil, + Parameter("1Abc") :: Nil, isDistinct = false), None) :: Nil, OneRowRelation())) comparePlans( parsePlan("SELECT * FROM a LIMIT @limitA"), - table("a").select(star()).limit(NamedParameter("limitA"))) + table("a").select(star()).limit(Parameter("limitA"))) // Invalid empty name and invalid symbol in a name Seq("@", "@-").foreach { name => checkError( diff --git a/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala b/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala index bdd323f3f43f8..b5331020d688b 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala @@ -35,9 +35,9 @@ import org.apache.spark.rdd.RDD import org.apache.spark.scheduler.{SparkListener, SparkListenerApplicationEnd} import org.apache.spark.sql.catalog.Catalog import org.apache.spark.sql.catalyst._ -import org.apache.spark.sql.catalyst.analysis.{BindParameters, UnresolvedRelation} +import org.apache.spark.sql.catalyst.analysis.UnresolvedRelation import org.apache.spark.sql.catalyst.encoders._ -import org.apache.spark.sql.catalyst.expressions.AttributeReference +import org.apache.spark.sql.catalyst.expressions.{AttributeReference, Parameter} import org.apache.spark.sql.catalyst.plans.logical.{LocalRelation, Range} import org.apache.spark.sql.catalyst.util.CharVarcharUtils import org.apache.spark.sql.connector.ExternalCommandRunner @@ -624,7 +624,7 @@ class SparkSession private( val plan = tracker.measurePhase(QueryPlanningTracker.PARSING) { val parser = sessionState.sqlParser val parsedArgs = args.mapValues(parser.parseExpression).toMap - BindParameters(parser.parsePlan(sqlText), parsedArgs) + Parameter.bind(parser.parsePlan(sqlText), parsedArgs) } Dataset.ofRows(self, plan, tracker) } From a0f568adccc81f44f923dadfd74fd20821014697 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Fri, 2 Dec 2022 00:07:36 +0300 Subject: [PATCH 13/38] Remove the Bind node --- .../spark/sql/catalyst/plans/logical/object.scala | 10 ---------- 1 file changed, 10 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/object.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/object.scala index e988c02efeec8..e5fe07e2d950d 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/object.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/plans/logical/object.scala @@ -690,13 +690,3 @@ case class CoGroup( override protected def withNewChildrenInternal( newLeft: LogicalPlan, newRight: LogicalPlan): CoGroup = copy(left = newLeft, right = newRight) } - -case class Bind(args: Map[String, Literal], child: LogicalPlan) extends UnaryNode { - - override def output: Seq[Attribute] = child.output - - final override val nodePatterns: Seq[TreePattern] = Seq(BIND) - - override protected def withNewChildInternal(newChild: LogicalPlan): Bind = - copy(child = newChild) -} From 2e93bec447054ffd8e683da47b0ef1fae78bdcc9 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Fri, 2 Dec 2022 13:15:36 +0300 Subject: [PATCH 14/38] Add the error class NON_FOLDABLE_SQL_ARG --- core/src/main/resources/error/error-classes.json | 5 +++++ .../spark/sql/catalyst/expressions/parameters.scala | 4 +--- .../sql/errors/QueryCompilationErrorsSuite.scala | 13 +++++++++++++ 3 files changed, 19 insertions(+), 3 deletions(-) diff --git a/core/src/main/resources/error/error-classes.json b/core/src/main/resources/error/error-classes.json index b376a20720e05..491ddb98104e1 100644 --- a/core/src/main/resources/error/error-classes.json +++ b/core/src/main/resources/error/error-classes.json @@ -852,6 +852,11 @@ "It is not allowed to use an aggregate function in the argument of another aggregate function. Please use the inner aggregate function in a sub-query." ] }, + "NON_FOLDABLE_SQL_ARG" : { + "message" : [ + "The argument is not foldable. Consider to replace it by a literal value." + ] + }, "NON_LAST_MATCHED_CLAUSE_OMIT_CONDITION" : { "message" : [ "When there are more than one MATCHED clauses in a MERGE statement, only the last MATCHED clause can omit the condition." diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index b8b69ad5305a3..291f86e7278ba 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -59,9 +59,7 @@ object Parameter extends QueryErrorsBase { args.filter(!_._2.foldable).headOption.foreach { case (name, expr) => expr.failAnalysis( errorClass = "NON_FOLDABLE_SQL_ARG", - messageParameters = Map( - "name" -> name, - "expr" -> toSQLExpr(expr))) + messageParameters = Map("name" -> name)) } plan.transformAllExpressionsWithPruning(_.containsPattern(PARAMETER)) { case param @ Parameter(name) => diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala index 56b90c263d031..85f0a623ba336 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala @@ -680,6 +680,19 @@ class QueryCompilationErrorsSuite start = 13, stop = 16)) } + + test("NON_FOLDABLE_SQL_ARG - SPARK-41271: non-foldable argument of `sql()`") { + checkError( + exception = intercept[AnalysisException] { + spark.sql("SELECT @param1 FROM VALUES (1) AS t(col1)", Map("param1" -> "col1 + 1")) + }, + errorClass = "NON_FOLDABLE_SQL_ARG", + parameters = Map("name" -> "param1"), + context = ExpectedContext( + fragment = "col1 + 1", + start = 0, + stop = 7)) + } } class MyCastToString extends SparkUserDefinedFunction( From cf03c3bb2e496668b8cc67ed581b9ee2d6d640e8 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Fri, 2 Dec 2022 13:26:58 +0300 Subject: [PATCH 15/38] Improve comments --- .../apache/spark/sql/catalyst/expressions/parameters.scala | 7 +++---- .../apache/spark/sql/catalyst/rules/RuleIdCollection.scala | 1 - .../main/scala/org/apache/spark/sql/internal/SQLConf.scala | 4 ++-- .../src/main/scala/org/apache/spark/sql/SparkSession.scala | 2 +- 4 files changed, 6 insertions(+), 8 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index 291f86e7278ba..b9aa3648e5438 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -28,8 +28,7 @@ import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{DataType, NullType} /** - * The expression represents a named parameter that should be bound later - * to a literal with concrete value and type. + * The expression represents a named parameter that should be replaces by a foldable expression. * * @param name The identifier of the parameter without the marker. */ @@ -50,8 +49,8 @@ case class Parameter(name: String) extends LeafExpression { /** - * Finds all named parameters in the given plan and substitutes them by literal values - * evaluated from `args` values. + * Finds all named parameters in the given plan and substitutes them by + * foldable expressions of `args` values. */ object Parameter extends QueryErrorsBase { def bind(plan: LogicalPlan, args: Map[String, Expression]): LogicalPlan = { diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/rules/RuleIdCollection.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/rules/RuleIdCollection.scala index 665bdc293f63d..f6bef88ab868e 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/rules/RuleIdCollection.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/rules/RuleIdCollection.scala @@ -79,7 +79,6 @@ object RuleIdCollection { "org.apache.spark.sql.catalyst.analysis.Analyzer$WindowsSubstitution" :: "org.apache.spark.sql.catalyst.analysis.AnsiTypeCoercion$AnsiCombinedTypeCoercionRule" :: "org.apache.spark.sql.catalyst.analysis.ApplyCharTypePadding" :: - "org.apache.spark.sql.catalyst.analysis.BindParameters" :: "org.apache.spark.sql.catalyst.analysis.DeduplicateRelations" :: "org.apache.spark.sql.catalyst.analysis.EliminateSubqueryAliases" :: "org.apache.spark.sql.catalyst.analysis.EliminateUnions" :: diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala index eae3841c64649..cf6fcacc04bb3 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala @@ -4028,8 +4028,8 @@ object SQLConf { .createWithDefault(ErrorMessageFormat.PRETTY.toString) val PARAMETERS_ENABLED = buildConf("spark.sql.parameters.enabled") - .doc("When set to true, queries can have named parameters that should be substituted " + - "by literal values later using `bind()`. If set to false, Spark handles constants " + + .doc("When set to true, `spark.sql()` executes input SQL queries by substituting named " + + "parameters by given literal values. If set to false, Spark handles constants " + "with the `@` prefix as regular identifiers and does not consider them as parameters.") .version("3.4.0") .booleanConf diff --git a/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala b/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala index b5331020d688b..71df5e724d71d 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala @@ -614,7 +614,7 @@ class SparkSession private( * This API eagerly runs DDL/DML commands, but not for SELECT queries. * * @param sqlText A SQL statement with named parameters to execute. - * @param args A map of parameter names to typed literals. + * @param args A map of parameter names to literal values. * * @since 3.4.0 */ From 0605f2227aff71bf3aa56336181ffe0dda208447 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Fri, 2 Dec 2022 13:30:11 +0300 Subject: [PATCH 16/38] Add unboundError() --- .../spark/sql/catalyst/expressions/parameters.scala | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index b9aa3648e5438..be0d25ed60c34 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -38,13 +38,12 @@ case class Parameter(name: String) extends LeafExpression { final override val nodePatterns: Seq[TreePattern] = Seq(PARAMETER) - override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { + private def unboundError() = throw SparkException.internalError(s"Found the unbound parameter: $name.") - } - def eval(input: InternalRow): Any = { - throw SparkException.internalError(s"Found the unbound parameter: $name.") - } + override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = unboundError() + + def eval(input: InternalRow): Any = unboundError() } From 210d23c901ae8c1151f680f3b224f2d4a1fb7551 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Fri, 2 Dec 2022 13:35:18 +0300 Subject: [PATCH 17/38] Remove the BIND tag --- .../scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala | 1 - 1 file changed, 1 deletion(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala index f43c4443d7ef5..dd5d173500657 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/trees/TreePatterns.scala @@ -96,7 +96,6 @@ object TreePattern extends Enumeration { // Logical plan patterns (alphabetically ordered) val AGGREGATE: Value = Value val AS_OF_JOIN: Value = Value - val BIND: Value = Value val COMMAND: Value = Value val CTE: Value = Value val DISTINCT_LIKE: Value = Value From 16966c9de766e95f8a24610bba42959eefb2606e Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Fri, 2 Dec 2022 17:49:39 +0300 Subject: [PATCH 18/38] Fix UNBOUND_PARAMETER --- core/src/main/resources/error/error-classes.json | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/core/src/main/resources/error/error-classes.json b/core/src/main/resources/error/error-classes.json index 491ddb98104e1..74edc23a9b90a 100644 --- a/core/src/main/resources/error/error-classes.json +++ b/core/src/main/resources/error/error-classes.json @@ -1137,7 +1137,7 @@ }, "UNBOUND_PARAMETER" : { "message" : [ - "Found the unbound parameter: . Use `bind()` to substitute the parameter by a literal." + "Found the unbound parameter: . Please, fix `args` and provide a mapping of the parameter to a literal value." ] }, "UNCLOSED_BRACKETED_COMMENT" : { From b8a6fd9d641b0d1f7949a93d2986a7022970fd93 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Sat, 3 Dec 2022 10:22:33 +0300 Subject: [PATCH 19/38] Add Java-specific method --- .../org/apache/spark/sql/SparkSession.scala | 19 +++++++++++++++++-- .../spark/sql/JavaSparkSessionSuite.java | 16 ++++++++++++++++ 2 files changed, 33 insertions(+), 2 deletions(-) diff --git a/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala b/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala index 71df5e724d71d..adbe593ac56fc 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/SparkSession.scala @@ -608,7 +608,7 @@ class SparkSession private( | Everything else | * ----------------- */ - /** + /** * Executes a SQL query substituting named parameters by the given arguments, * returning the result as a `DataFrame`. * This API eagerly runs DDL/DML commands, but not for SELECT queries. @@ -629,13 +629,28 @@ class SparkSession private( Dataset.ofRows(self, plan, tracker) } + /** + * Executes a SQL query substituting named parameters by the given arguments, + * returning the result as a `DataFrame`. + * This API eagerly runs DDL/DML commands, but not for SELECT queries. + * + * @param sqlText A SQL statement with named parameters to execute. + * @param args A map of parameter names to literal values. + * + * @since 3.4.0 + */ + @Experimental + def sql(sqlText: String, args: java.util.Map[String, String]): DataFrame = { + sql(sqlText, args.asScala.toMap) + } + /** * Executes a SQL query using Spark, returning the result as a `DataFrame`. * This API eagerly runs DDL/DML commands, but not for SELECT queries. * * @since 2.0.0 */ - def sql(sqlText: String): DataFrame = sql(sqlText, Map.empty) + def sql(sqlText: String): DataFrame = sql(sqlText, Map.empty[String, String]) /** * Execute an arbitrary string command inside an external execution engine rather than Spark. diff --git a/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java b/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java index b1df377936dfa..c362caa4fdeff 100644 --- a/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java +++ b/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java @@ -23,6 +23,7 @@ import org.junit.Test; import java.util.HashMap; +import java.util.List; import java.util.Map; public class JavaSparkSessionSuite { @@ -54,4 +55,19 @@ public void config() { Assert.assertEquals(spark.conf().get(e.getKey()), e.getValue().toString()); } } + + @Test + public void sqlParameters() { + spark = SparkSession.builder().master("local[*]").appName("testing").getOrCreate(); + Map params = new HashMap(); + params.put("_i1", "INTERVAL '1-1' YEAR TO MONTH"); + params.put("p2", "'abc'"); + Dataset ds = spark.sql( + "SELECT @p2, i FROM VALUES (INTERVAL '2-2' YEAR TO MONTH) AS t(i) WHERE i > @_i1", + params); + List rows = ds.collectAsList(); + Assert.assertEquals(1, rows.size()); + Assert.assertEquals("abc", rows.get(0).getString(0)); + Assert.assertEquals(java.time.Period.of(2, 2, 0), rows.get(0).get(1)); + } } From 89c5371abafee8dd70950bbc176edc77b56b2ad1 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Tue, 6 Dec 2022 11:09:23 +0300 Subject: [PATCH 20/38] Add tests for double quotes --- .../java/test/org/apache/spark/sql/JavaSparkSessionSuite.java | 4 ++-- .../src/test/scala/org/apache/spark/sql/DatasetSuite.scala | 4 ++++ 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java b/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java index c362caa4fdeff..61d2739458379 100644 --- a/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java +++ b/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java @@ -61,13 +61,13 @@ public void sqlParameters() { spark = SparkSession.builder().master("local[*]").appName("testing").getOrCreate(); Map params = new HashMap(); params.put("_i1", "INTERVAL '1-1' YEAR TO MONTH"); - params.put("p2", "'abc'"); + params.put("p2", "'a\"bc'"); Dataset ds = spark.sql( "SELECT @p2, i FROM VALUES (INTERVAL '2-2' YEAR TO MONTH) AS t(i) WHERE i > @_i1", params); List rows = ds.collectAsList(); Assert.assertEquals(1, rows.size()); - Assert.assertEquals("abc", rows.get(0).getString(0)); + Assert.assertEquals("a\"bc", rows.get(0).getString(0)); Assert.assertEquals(java.time.Period.of(2, 2, 0), rows.get(0).get(1)); } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index 6e454918e9b63..28cbbb2260774 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -2258,6 +2258,10 @@ class DatasetSuite extends QueryTest checkAnswer( spark.sql(sqlText, args), Row(0, 0) :: Row(1, 1) :: Row(2, 2) :: Row(3, 0) :: Nil) + + checkAnswer( + spark.sql("""SELECT contains('Spark \'SQL\'', @subStr)""", Map("subStr" -> "'SQL'")), + Row(true)) } } } From e883aade64b520d7bd7df7735d2691f9c9bf0654 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Tue, 6 Dec 2022 18:26:18 +0300 Subject: [PATCH 21/38] Add one more test for a foldable expr - cast --- .../src/test/scala/org/apache/spark/sql/DatasetSuite.scala | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index 28cbbb2260774..d4644b3710ae3 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -2262,6 +2262,11 @@ class DatasetSuite extends QueryTest checkAnswer( spark.sql("""SELECT contains('Spark \'SQL\'', @subStr)""", Map("subStr" -> "'SQL'")), Row(true)) + checkAnswer( + spark.sql( + """SELECT 1 + @castInt""", + Map("castInt" -> "CAST('100' AS INT)")), + Row(101)) } } } From d05526cca6e8e31a4dc10cb1e13b5466f93cb5c5 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Tue, 6 Dec 2022 21:03:07 +0300 Subject: [PATCH 22/38] Convert the internal error to an user-facing one --- .../sql/catalyst/expressions/parameters.scala | 14 +++++++------- .../spark/sql/errors/QueryExecutionErrors.scala | 6 ++++++ .../sql/errors/QueryExecutionErrorsSuite.scala | 9 +++++++++ 3 files changed, 22 insertions(+), 7 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index be0d25ed60c34..7d864049bf302 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -17,13 +17,12 @@ package org.apache.spark.sql.catalyst.expressions -import org.apache.spark.SparkException import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.analysis.AnalysisErrorAt import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, ExprCode} import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan import org.apache.spark.sql.catalyst.trees.TreePattern.{PARAMETER, TreePattern} -import org.apache.spark.sql.errors.QueryErrorsBase +import org.apache.spark.sql.errors.{QueryErrorsBase, QueryExecutionErrors} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{DataType, NullType} @@ -38,12 +37,13 @@ case class Parameter(name: String) extends LeafExpression { final override val nodePatterns: Seq[TreePattern] = Seq(PARAMETER) - private def unboundError() = - throw SparkException.internalError(s"Found the unbound parameter: $name.") - - override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = unboundError() + override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { + throw QueryExecutionErrors.unboundParameterError(name) + } - def eval(input: InternalRow): Any = unboundError() + def eval(input: InternalRow): Any = { + throw QueryExecutionErrors.unboundParameterError(name) + } } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryExecutionErrors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryExecutionErrors.scala index 15dfa581c5976..c441fba1f4beb 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryExecutionErrors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryExecutionErrors.scala @@ -2827,4 +2827,10 @@ private[sql] object QueryExecutionErrors extends QueryErrorsBase { "location" -> toSQLValue(location.toString, StringType), "identifier" -> toSQLId(tableId.nameParts))) } + + def unboundParameterError(name: String): Throwable = { + new SparkRuntimeException( + errorClass = "UNBOUND_PARAMETER", + messageParameters = Map("name" -> name)) + } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala index e01ff56752c3a..83056cc457d11 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala @@ -769,6 +769,15 @@ class QueryExecutionErrorsSuite assert(e.getErrorClass === "STREAM_FAILED") assert(e.getCause.isInstanceOf[NullPointerException]) } + + test("UNBOUND_PARAMETER - SPARK-41271: non-substituted parameters") { + checkError( + exception = intercept[SparkRuntimeException] { + sql("select @abc").collect() + }, + errorClass = "UNBOUND_PARAMETER", + parameters = Map("name" -> "abc")) + } } class FakeFileSystemSetPermission extends LocalFileSystem { From ce19640563f9db9d1003461dd55be125dac357fc Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Wed, 7 Dec 2022 11:08:02 +0300 Subject: [PATCH 23/38] Add a test for ignored args --- .../apache/spark/sql/catalyst/analysis/AnalysisSuite.scala | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala index a33721b8b418e..ff11a2bdceaba 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala @@ -1303,6 +1303,12 @@ class AnalysisSuite extends AnalysisTest with Matchers { plan = parsePlan("SELECT * FROM a LIMIT @limitA"), args = Map("limitA" -> Literal(10))), parsePlan("SELECT * FROM a LIMIT 10")) + // Ignore unused arguments + comparePlans( + Parameter.bind( + plan = parsePlan("SELECT c FROM a WHERE c < @param2"), + args = Map("param1" -> Literal(10), "param2" -> Literal(20))), + parsePlan("SELECT c FROM a WHERE c < 20")) } } } From 5a1f33d7f258206eb7cfa0e8cd706d31b8ebbdb5 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Wed, 7 Dec 2022 14:51:07 +0300 Subject: [PATCH 24/38] Allow literals only. --- .../main/resources/error/error-classes.json | 4 ++-- .../sql/catalyst/expressions/parameters.scala | 6 ++--- .../org/apache/spark/sql/DatasetSuite.scala | 5 ---- .../errors/QueryCompilationErrorsSuite.scala | 24 ++++++++++--------- 4 files changed, 18 insertions(+), 21 deletions(-) diff --git a/core/src/main/resources/error/error-classes.json b/core/src/main/resources/error/error-classes.json index 74edc23a9b90a..eac8ed185b730 100644 --- a/core/src/main/resources/error/error-classes.json +++ b/core/src/main/resources/error/error-classes.json @@ -852,9 +852,9 @@ "It is not allowed to use an aggregate function in the argument of another aggregate function. Please use the inner aggregate function in a sub-query." ] }, - "NON_FOLDABLE_SQL_ARG" : { + "INVALID_SQL_ARG" : { "message" : [ - "The argument is not foldable. Consider to replace it by a literal value." + "The argument of `sql()` is invalid. Consider to replace it by a literal value." ] }, "NON_LAST_MATCHED_CLAUSE_OMIT_CONDITION" : { diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index 7d864049bf302..dc9a9d6bbc38c 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -54,10 +54,10 @@ case class Parameter(name: String) extends LeafExpression { object Parameter extends QueryErrorsBase { def bind(plan: LogicalPlan, args: Map[String, Expression]): LogicalPlan = { if (!args.isEmpty && SQLConf.get.parametersEnabled) { - args.filter(!_._2.foldable).headOption.foreach { case (name, expr) => + args.filter(!_._2.isInstanceOf[Literal]).headOption.foreach { case (name, expr) => expr.failAnalysis( - errorClass = "NON_FOLDABLE_SQL_ARG", - messageParameters = Map("name" -> name)) + errorClass = "INVALID_SQL_ARG", + messageParameters = Map("name" -> toSQLId(name))) } plan.transformAllExpressionsWithPruning(_.containsPattern(PARAMETER)) { case param @ Parameter(name) => diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index d4644b3710ae3..28cbbb2260774 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -2262,11 +2262,6 @@ class DatasetSuite extends QueryTest checkAnswer( spark.sql("""SELECT contains('Spark \'SQL\'', @subStr)""", Map("subStr" -> "'SQL'")), Row(true)) - checkAnswer( - spark.sql( - """SELECT 1 + @castInt""", - Map("castInt" -> "CAST('100' AS INT)")), - Row(101)) } } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala index 85f0a623ba336..5789f807f59a5 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala @@ -681,17 +681,19 @@ class QueryCompilationErrorsSuite stop = 16)) } - test("NON_FOLDABLE_SQL_ARG - SPARK-41271: non-foldable argument of `sql()`") { - checkError( - exception = intercept[AnalysisException] { - spark.sql("SELECT @param1 FROM VALUES (1) AS t(col1)", Map("param1" -> "col1 + 1")) - }, - errorClass = "NON_FOLDABLE_SQL_ARG", - parameters = Map("name" -> "param1"), - context = ExpectedContext( - fragment = "col1 + 1", - start = 0, - stop = 7)) + test("INVALID_SQL_ARG - SPARK-41271: non-literal argument of `sql()`") { + Seq("col1 + 1", "CAST('100' AS INT)", "map('a', 1, 'b', 2)", "array(1)").foreach { arg => + checkError( + exception = intercept[AnalysisException] { + spark.sql("SELECT @param1 FROM VALUES (1) AS t(col1)", Map("param1" -> arg)) + }, + errorClass = "INVALID_SQL_ARG", + parameters = Map("name" -> "`param1`"), + context = ExpectedContext( + fragment = arg, + start = 0, + stop = arg.length - 1)) + } } } From 43fd1b9513c400bf2dda8a6ec7807d8401be54be Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Wed, 7 Dec 2022 16:26:52 +0300 Subject: [PATCH 25/38] Fix the order of error classes in error-classes.json --- core/src/main/resources/error/error-classes.json | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/core/src/main/resources/error/error-classes.json b/core/src/main/resources/error/error-classes.json index eac8ed185b730..725eee3bc7f1f 100644 --- a/core/src/main/resources/error/error-classes.json +++ b/core/src/main/resources/error/error-classes.json @@ -795,6 +795,11 @@ } } }, + "INVALID_SQL_ARG" : { + "message" : [ + "The argument of `sql()` is invalid. Consider to replace it by a literal value." + ] + }, "INVALID_SQL_SYNTAX" : { "message" : [ "Invalid SQL syntax: " @@ -852,11 +857,6 @@ "It is not allowed to use an aggregate function in the argument of another aggregate function. Please use the inner aggregate function in a sub-query." ] }, - "INVALID_SQL_ARG" : { - "message" : [ - "The argument of `sql()` is invalid. Consider to replace it by a literal value." - ] - }, "NON_LAST_MATCHED_CLAUSE_OMIT_CONDITION" : { "message" : [ "When there are more than one MATCHED clauses in a MERGE statement, only the last MATCHED clause can omit the condition." From 5ad46f82cdfb08c2af6079bba04a238f3ec7fd26 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Thu, 8 Dec 2022 10:14:35 +0300 Subject: [PATCH 26/38] Switch the parameter marker from @ to : --- .../spark/sql/catalyst/parser/SqlBaseLexer.g4 | 2 +- .../spark/sql/catalyst/parser/SqlBaseParser.g4 | 2 +- .../org/apache/spark/sql/internal/SQLConf.scala | 2 +- .../sql/catalyst/analysis/AnalysisSuite.scala | 4 ++-- .../sql/catalyst/parser/PlanParserSuite.scala | 14 +++++++------- .../apache/spark/sql/JavaSparkSessionSuite.java | 2 +- .../scala/org/apache/spark/sql/DatasetSuite.scala | 6 +++--- .../sql/errors/QueryCompilationErrorsSuite.scala | 6 +++--- .../sql/errors/QueryExecutionErrorsSuite.scala | 2 +- 9 files changed, 20 insertions(+), 20 deletions(-) diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 index a380fa19c3424..679d309f30b79 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 @@ -509,5 +509,5 @@ UNRECOGNIZED ; PARAMETER - : '@' IDENTIFIER + : ':' IDENTIFIER ; diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 index 87c3e142abcc5..3ce6d8cc27645 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 @@ -42,7 +42,7 @@ options { tokenVocab = SqlBaseLexer; } public boolean double_quoted_identifiers = false; /** - * When true, identifiers that begin from `@` are considered as named parameters. + * When true, identifiers that begin from `:` are considered as named parameters. */ public boolean parameters_enabled = false; } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala index cf6fcacc04bb3..c130b342e5174 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala @@ -4030,7 +4030,7 @@ object SQLConf { val PARAMETERS_ENABLED = buildConf("spark.sql.parameters.enabled") .doc("When set to true, `spark.sql()` executes input SQL queries by substituting named " + "parameters by given literal values. If set to false, Spark handles constants " + - "with the `@` prefix as regular identifiers and does not consider them as parameters.") + "with the `:` prefix as regular identifiers and does not consider them as parameters.") .version("3.4.0") .booleanConf .createWithDefault(true) diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala index ff11a2bdceaba..89c2963961bda 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala @@ -1300,13 +1300,13 @@ class AnalysisSuite extends AnalysisTest with Matchers { withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { comparePlans( Parameter.bind( - plan = parsePlan("SELECT * FROM a LIMIT @limitA"), + plan = parsePlan("SELECT * FROM a LIMIT :limitA"), args = Map("limitA" -> Literal(10))), parsePlan("SELECT * FROM a LIMIT 10")) // Ignore unused arguments comparePlans( Parameter.bind( - plan = parsePlan("SELECT c FROM a WHERE c < @param2"), + plan = parsePlan("SELECT c FROM a WHERE c < :param2"), args = Map("param1" -> Literal(10), "param2" -> Literal(20))), parsePlan("SELECT c FROM a WHERE c < 20")) } diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala index 72b4357621f88..b8a97cb2ee3a0 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala @@ -1572,10 +1572,10 @@ class PlanParserSuite extends AnalysisTest { test("SPARK-41271: parsing of named parameters") { withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { comparePlans( - parsePlan("SELECT @param_1"), + parsePlan("SELECT :param_1"), Project(UnresolvedAlias(Parameter("param_1"), None) :: Nil, OneRowRelation())) comparePlans( - parsePlan("SELECT abs(@1Abc)"), + parsePlan("SELECT abs(:1Abc)"), Project(UnresolvedAlias( UnresolvedFunction( "abs" :: Nil, @@ -1583,23 +1583,23 @@ class PlanParserSuite extends AnalysisTest { isDistinct = false), None) :: Nil, OneRowRelation())) comparePlans( - parsePlan("SELECT * FROM a LIMIT @limitA"), + parsePlan("SELECT * FROM a LIMIT :limitA"), table("a").select(star()).limit(Parameter("limitA"))) // Invalid empty name and invalid symbol in a name - Seq("@", "@-").foreach { name => + Seq(":", ":-").foreach { name => checkError( exception = parseException(s"SELECT $name"), errorClass = "PARSE_SYNTAX_ERROR", - parameters = Map("error" -> "'@'", "hint" -> "")) + parameters = Map("error" -> "':'", "hint" -> "")) } } withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "false") { checkError( exception = intercept[ParseException] { - parsePlan("SELECT @param_1") + parsePlan("SELECT :param_1") }, errorClass = "PARSE_SYNTAX_ERROR", - parameters = Map("error" -> "'@param_1'", "hint" -> "")) + parameters = Map("error" -> "':param_1'", "hint" -> "")) } } } diff --git a/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java b/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java index 61d2739458379..b9e29897c26ef 100644 --- a/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java +++ b/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java @@ -63,7 +63,7 @@ public void sqlParameters() { params.put("_i1", "INTERVAL '1-1' YEAR TO MONTH"); params.put("p2", "'a\"bc'"); Dataset ds = spark.sql( - "SELECT @p2, i FROM VALUES (INTERVAL '2-2' YEAR TO MONTH) AS t(i) WHERE i > @_i1", + "SELECT :p2, i FROM VALUES (INTERVAL '2-2' YEAR TO MONTH) AS t(i) WHERE i > :_i1", params); List rows = ds.collectAsList(); Assert.assertEquals(1, rows.size()); diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index 28cbbb2260774..d00518e4fc18b 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -2250,9 +2250,9 @@ class DatasetSuite extends QueryTest withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { val sqlText = """ - |SELECT id, id % @div as c0 + |SELECT id, id % :div as c0 |FROM VALUES (0), (1), (2), (3), (4), (5), (6), (7), (8), (9) AS t(id) - |WHERE id < @constA + |WHERE id < :constA |""".stripMargin val args = Map("div" -> "3", "constA" -> "4L") checkAnswer( @@ -2260,7 +2260,7 @@ class DatasetSuite extends QueryTest Row(0, 0) :: Row(1, 1) :: Row(2, 2) :: Row(3, 0) :: Nil) checkAnswer( - spark.sql("""SELECT contains('Spark \'SQL\'', @subStr)""", Map("subStr" -> "'SQL'")), + spark.sql("""SELECT contains('Spark \'SQL\'', :subStr)""", Map("subStr" -> "'SQL'")), Row(true)) } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala index 5789f807f59a5..be5b2c22f957a 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala @@ -671,12 +671,12 @@ class QueryCompilationErrorsSuite test("UNBOUND_PARAMETER - SPARK-41271: non-substituted parameters") { checkError( exception = intercept[AnalysisException] { - spark.sql("select @abc, @def", Map("abc" -> "1")) + spark.sql("select :abc, :def", Map("abc" -> "1")) }, errorClass = "UNBOUND_PARAMETER", parameters = Map("name" -> "def"), context = ExpectedContext( - fragment = "@def", + fragment = ":def", start = 13, stop = 16)) } @@ -685,7 +685,7 @@ class QueryCompilationErrorsSuite Seq("col1 + 1", "CAST('100' AS INT)", "map('a', 1, 'b', 2)", "array(1)").foreach { arg => checkError( exception = intercept[AnalysisException] { - spark.sql("SELECT @param1 FROM VALUES (1) AS t(col1)", Map("param1" -> arg)) + spark.sql("SELECT :param1 FROM VALUES (1) AS t(col1)", Map("param1" -> arg)) }, errorClass = "INVALID_SQL_ARG", parameters = Map("name" -> "`param1`"), diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala index 83056cc457d11..d80975d121535 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala @@ -773,7 +773,7 @@ class QueryExecutionErrorsSuite test("UNBOUND_PARAMETER - SPARK-41271: non-substituted parameters") { checkError( exception = intercept[SparkRuntimeException] { - sql("select @abc").collect() + sql("select :abc").collect() }, errorClass = "UNBOUND_PARAMETER", parameters = Map("name" -> "abc")) From 637778ce30e7da386b723ecfdc31bd2946d28ce0 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Thu, 8 Dec 2022 16:26:53 +0300 Subject: [PATCH 27/38] Fix parsing error --- .../org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 | 4 ---- .../org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 | 2 +- 2 files changed, 1 insertion(+), 5 deletions(-) diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 index 679d309f30b79..41adbda7b101e 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseLexer.g4 @@ -507,7 +507,3 @@ WS UNRECOGNIZED : . ; - -PARAMETER - : ':' IDENTIFIER - ; diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 index 3ce6d8cc27645..c1711910ace44 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 @@ -935,7 +935,7 @@ primaryExpression constant : NULL #nullLiteral - | {parameters_enabled}? PARAMETER #parameterLiteral + | {parameters_enabled}? ':' identifier #parameterLiteral | interval #intervalLiteral | identifier stringLit #typeConstructor | number #numericLiteral From d2ce0965dc63e3ce28590affd1dd39313f688ec8 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Thu, 8 Dec 2022 19:22:54 +0300 Subject: [PATCH 28/38] Use COLON --- .../spark/sql/catalyst/parser/SqlBaseParser.g4 | 2 +- .../sql/catalyst/parser/PlanParserSuite.scala | 16 +++++++++------- 2 files changed, 10 insertions(+), 8 deletions(-) diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 index c1711910ace44..79c934286cfd9 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 @@ -935,7 +935,7 @@ primaryExpression constant : NULL #nullLiteral - | {parameters_enabled}? ':' identifier #parameterLiteral + | {parameters_enabled}? COLON identifier #parameterLiteral | interval #intervalLiteral | identifier stringLit #typeConstructor | number #numericLiteral diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala index b8a97cb2ee3a0..2e109c59cf4ed 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala @@ -1586,12 +1586,14 @@ class PlanParserSuite extends AnalysisTest { parsePlan("SELECT * FROM a LIMIT :limitA"), table("a").select(star()).limit(Parameter("limitA"))) // Invalid empty name and invalid symbol in a name - Seq(":", ":-").foreach { name => - checkError( - exception = parseException(s"SELECT $name"), - errorClass = "PARSE_SYNTAX_ERROR", - parameters = Map("error" -> "':'", "hint" -> "")) - } + checkError( + exception = parseException(s"SELECT :-"), + errorClass = "PARSE_SYNTAX_ERROR", + parameters = Map("error" -> "'-'", "hint" -> "")) + checkError( + exception = parseException(s"SELECT :"), + errorClass = "PARSE_SYNTAX_ERROR", + parameters = Map("error" -> "end of input", "hint" -> "")) } withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "false") { checkError( @@ -1599,7 +1601,7 @@ class PlanParserSuite extends AnalysisTest { parsePlan("SELECT :param_1") }, errorClass = "PARSE_SYNTAX_ERROR", - parameters = Map("error" -> "':param_1'", "hint" -> "")) + parameters = Map("error" -> "':'", "hint" -> "")) } } } From d3ac69c86e4af55e06ebf61576ea95df2b92af3e Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Thu, 8 Dec 2022 22:11:28 +0300 Subject: [PATCH 29/38] a literal value -> a SQL literal statement --- core/src/main/resources/error/error-classes.json | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/core/src/main/resources/error/error-classes.json b/core/src/main/resources/error/error-classes.json index 725eee3bc7f1f..a95cea9a3b436 100644 --- a/core/src/main/resources/error/error-classes.json +++ b/core/src/main/resources/error/error-classes.json @@ -797,7 +797,7 @@ }, "INVALID_SQL_ARG" : { "message" : [ - "The argument of `sql()` is invalid. Consider to replace it by a literal value." + "The argument of `sql()` is invalid. Consider to replace it by a SQL literal statement." ] }, "INVALID_SQL_SYNTAX" : { @@ -1137,7 +1137,7 @@ }, "UNBOUND_PARAMETER" : { "message" : [ - "Found the unbound parameter: . Please, fix `args` and provide a mapping of the parameter to a literal value." + "Found the unbound parameter: . Please, fix `args` and provide a mapping of the parameter to a SQL literal statement." ] }, "UNCLOSED_BRACKETED_COMMENT" : { From 08e6dfa2f6415e7578e49e628a76a3affa3b2f22 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Thu, 8 Dec 2022 22:20:52 +0300 Subject: [PATCH 30/38] Address comments related to error classes --- core/src/main/resources/error/error-classes.json | 2 +- .../apache/spark/sql/catalyst/expressions/parameters.scala | 2 +- .../org/apache/spark/sql/errors/QueryExecutionErrors.scala | 2 +- .../apache/spark/sql/errors/QueryCompilationErrorsSuite.scala | 4 ++-- .../apache/spark/sql/errors/QueryExecutionErrorsSuite.scala | 4 ++-- 5 files changed, 7 insertions(+), 7 deletions(-) diff --git a/core/src/main/resources/error/error-classes.json b/core/src/main/resources/error/error-classes.json index a95cea9a3b436..c4f9ea32a8c00 100644 --- a/core/src/main/resources/error/error-classes.json +++ b/core/src/main/resources/error/error-classes.json @@ -1135,7 +1135,7 @@ "Unable to convert SQL type to Protobuf type ." ] }, - "UNBOUND_PARAMETER" : { + "UNBOUND_SQL_PARAMETER" : { "message" : [ "Found the unbound parameter: . Please, fix `args` and provide a mapping of the parameter to a SQL literal statement." ] diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index dc9a9d6bbc38c..561a41a631526 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -65,7 +65,7 @@ object Parameter extends QueryErrorsBase { args(name) } else { param.failAnalysis( - errorClass = "UNBOUND_PARAMETER", + errorClass = "UNBOUND_SQL_PARAMETER", messageParameters = Map("name" -> name)) } } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryExecutionErrors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryExecutionErrors.scala index c441fba1f4beb..ea63f276b95b9 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryExecutionErrors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryExecutionErrors.scala @@ -2830,7 +2830,7 @@ private[sql] object QueryExecutionErrors extends QueryErrorsBase { def unboundParameterError(name: String): Throwable = { new SparkRuntimeException( - errorClass = "UNBOUND_PARAMETER", + errorClass = "UNBOUND_SQL_PARAMETER", messageParameters = Map("name" -> name)) } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala index be5b2c22f957a..706ba3e4199ce 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala @@ -668,12 +668,12 @@ class QueryCompilationErrorsSuite parameters = Map("schema" -> "\"INT\"", "sqlExpr" -> "\"from_json(a)\"")) } - test("UNBOUND_PARAMETER - SPARK-41271: non-substituted parameters") { + test("UNBOUND_SQL_PARAMETER - SPARK-41271: non-substituted parameters") { checkError( exception = intercept[AnalysisException] { spark.sql("select :abc, :def", Map("abc" -> "1")) }, - errorClass = "UNBOUND_PARAMETER", + errorClass = "UNBOUND_SQL_PARAMETER", parameters = Map("name" -> "def"), context = ExpectedContext( fragment = ":def", diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala index d80975d121535..f1979a3163473 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala @@ -770,12 +770,12 @@ class QueryExecutionErrorsSuite assert(e.getCause.isInstanceOf[NullPointerException]) } - test("UNBOUND_PARAMETER - SPARK-41271: non-substituted parameters") { + test("UNBOUND_SQL_PARAMETER - SPARK-41271: non-substituted parameters") { checkError( exception = intercept[SparkRuntimeException] { sql("select :abc").collect() }, - errorClass = "UNBOUND_PARAMETER", + errorClass = "UNBOUND_SQL_PARAMETER", parameters = Map("name" -> "abc")) } } From 27927fb74cfa0b80f649c38781a9a82fa0f216f3 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Fri, 9 Dec 2022 09:02:14 +0300 Subject: [PATCH 31/38] Add ParametersSuite --- .../org/apache/spark/sql/DatasetSuite.scala | 19 ----- .../apache/spark/sql/ParametersSuite.scala | 78 +++++++++++++++++++ .../errors/QueryCompilationErrorsSuite.scala | 28 ------- .../errors/QueryExecutionErrorsSuite.scala | 9 --- 4 files changed, 78 insertions(+), 56 deletions(-) create mode 100644 sql/core/src/test/scala/org/apache/spark/sql/ParametersSuite.scala diff --git a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala index d00518e4fc18b..d298d7129c70d 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/DatasetSuite.scala @@ -2245,25 +2245,6 @@ class DatasetSuite extends QueryTest assert(parquetFiles.size === 10) } } - - test("SPARK-41271: bind parameters") { - withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { - val sqlText = - """ - |SELECT id, id % :div as c0 - |FROM VALUES (0), (1), (2), (3), (4), (5), (6), (7), (8), (9) AS t(id) - |WHERE id < :constA - |""".stripMargin - val args = Map("div" -> "3", "constA" -> "4L") - checkAnswer( - spark.sql(sqlText, args), - Row(0, 0) :: Row(1, 1) :: Row(2, 2) :: Row(3, 0) :: Nil) - - checkAnswer( - spark.sql("""SELECT contains('Spark \'SQL\'', :subStr)""", Map("subStr" -> "'SQL'")), - Row(true)) - } - } } class DatasetLargeResultCollectingSuite extends QueryTest diff --git a/sql/core/src/test/scala/org/apache/spark/sql/ParametersSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/ParametersSuite.scala new file mode 100644 index 0000000000000..6490c9ad683cf --- /dev/null +++ b/sql/core/src/test/scala/org/apache/spark/sql/ParametersSuite.scala @@ -0,0 +1,78 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql + +import org.apache.spark.SparkRuntimeException +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.test.SharedSparkSession + +class ParametersSuite extends QueryTest with SharedSparkSession { + + test("bind parameters") { + withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { + val sqlText = + """ + |SELECT id, id % :div as c0 + |FROM VALUES (0), (1), (2), (3), (4), (5), (6), (7), (8), (9) AS t(id) + |WHERE id < :constA + |""".stripMargin + val args = Map("div" -> "3", "constA" -> "4L") + checkAnswer( + spark.sql(sqlText, args), + Row(0, 0) :: Row(1, 1) :: Row(2, 2) :: Row(3, 0) :: Nil) + + checkAnswer( + spark.sql("""SELECT contains('Spark \'SQL\'', :subStr)""", Map("subStr" -> "'SQL'")), + Row(true)) + } + } + + test("non-substituted parameters") { + checkError( + exception = intercept[AnalysisException] { + spark.sql("select :abc, :def", Map("abc" -> "1")) + }, + errorClass = "UNBOUND_SQL_PARAMETER", + parameters = Map("name" -> "def"), + context = ExpectedContext( + fragment = ":def", + start = 13, + stop = 16)) + checkError( + exception = intercept[SparkRuntimeException] { + sql("select :abc").collect() + }, + errorClass = "UNBOUND_SQL_PARAMETER", + parameters = Map("name" -> "abc")) + } + + test("non-literal argument of `sql()`") { + Seq("col1 + 1", "CAST('100' AS INT)", "map('a', 1, 'b', 2)", "array(1)").foreach { arg => + checkError( + exception = intercept[AnalysisException] { + spark.sql("SELECT :param1 FROM VALUES (1) AS t(col1)", Map("param1" -> arg)) + }, + errorClass = "INVALID_SQL_ARG", + parameters = Map("name" -> "`param1`"), + context = ExpectedContext( + fragment = arg, + start = 0, + stop = arg.length - 1)) + } + } +} diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala index 706ba3e4199ce..bed647ef49fcd 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryCompilationErrorsSuite.scala @@ -667,34 +667,6 @@ class QueryCompilationErrorsSuite errorClass = "DATATYPE_MISMATCH.INVALID_JSON_SCHEMA", parameters = Map("schema" -> "\"INT\"", "sqlExpr" -> "\"from_json(a)\"")) } - - test("UNBOUND_SQL_PARAMETER - SPARK-41271: non-substituted parameters") { - checkError( - exception = intercept[AnalysisException] { - spark.sql("select :abc, :def", Map("abc" -> "1")) - }, - errorClass = "UNBOUND_SQL_PARAMETER", - parameters = Map("name" -> "def"), - context = ExpectedContext( - fragment = ":def", - start = 13, - stop = 16)) - } - - test("INVALID_SQL_ARG - SPARK-41271: non-literal argument of `sql()`") { - Seq("col1 + 1", "CAST('100' AS INT)", "map('a', 1, 'b', 2)", "array(1)").foreach { arg => - checkError( - exception = intercept[AnalysisException] { - spark.sql("SELECT :param1 FROM VALUES (1) AS t(col1)", Map("param1" -> arg)) - }, - errorClass = "INVALID_SQL_ARG", - parameters = Map("name" -> "`param1`"), - context = ExpectedContext( - fragment = arg, - start = 0, - stop = arg.length - 1)) - } - } } class MyCastToString extends SparkUserDefinedFunction( diff --git a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala index f1979a3163473..e01ff56752c3a 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/errors/QueryExecutionErrorsSuite.scala @@ -769,15 +769,6 @@ class QueryExecutionErrorsSuite assert(e.getErrorClass === "STREAM_FAILED") assert(e.getCause.isInstanceOf[NullPointerException]) } - - test("UNBOUND_SQL_PARAMETER - SPARK-41271: non-substituted parameters") { - checkError( - exception = intercept[SparkRuntimeException] { - sql("select :abc").collect() - }, - errorClass = "UNBOUND_SQL_PARAMETER", - parameters = Map("name" -> "abc")) - } } class FakeFileSystemSetPermission extends LocalFileSystem { From c390260ff4a88027852c25cec7c657d0c0bd11fb Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Fri, 9 Dec 2022 15:02:15 +0300 Subject: [PATCH 32/38] Address Wenchen's review comments --- .../sql/catalyst/analysis/CheckAnalysis.scala | 5 ++++ .../sql/catalyst/expressions/parameters.scala | 25 ++++++++----------- .../sql/errors/QueryExecutionErrors.scala | 6 ----- .../apache/spark/sql/ParametersSuite.scala | 11 +++++--- 4 files changed, 23 insertions(+), 24 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/CheckAnalysis.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/CheckAnalysis.scala index 12dac5c632a3b..ff19ce3de2c32 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/CheckAnalysis.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/CheckAnalysis.scala @@ -325,6 +325,11 @@ trait CheckAnalysis extends PredicateHelper with LookupCatalog with QueryErrorsB errorClass = "_LEGACY_ERROR_TEMP_2413", messageParameters = Map("argName" -> e.prettyName)) + case p: Parameter => + p.failAnalysis( + errorClass = "UNBOUND_SQL_PARAMETER", + messageParameters = Map("name" -> toSQLId(p.name))) + case _ => }) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index 561a41a631526..26f70af95a517 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -17,12 +17,13 @@ package org.apache.spark.sql.catalyst.expressions +import org.apache.spark.SparkException import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.analysis.AnalysisErrorAt import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, ExprCode} import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan import org.apache.spark.sql.catalyst.trees.TreePattern.{PARAMETER, TreePattern} -import org.apache.spark.sql.errors.{QueryErrorsBase, QueryExecutionErrors} +import org.apache.spark.sql.errors.QueryErrorsBase import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{DataType, NullType} @@ -32,18 +33,21 @@ import org.apache.spark.sql.types.{DataType, NullType} * @param name The identifier of the parameter without the marker. */ case class Parameter(name: String) extends LeafExpression { + + override lazy val resolved: Boolean = false override def dataType: DataType = NullType override def nullable: Boolean = true final override val nodePatterns: Seq[TreePattern] = Seq(PARAMETER) - override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { - throw QueryExecutionErrors.unboundParameterError(name) + private def unboundError(): Nothing = { + throw SparkException.internalError( + s"The parameter `$name` must be bound at the analysis phase.") } - def eval(input: InternalRow): Any = { - throw QueryExecutionErrors.unboundParameterError(name) - } + override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = unboundError() + + def eval(input: InternalRow): Any = unboundError() } @@ -60,14 +64,7 @@ object Parameter extends QueryErrorsBase { messageParameters = Map("name" -> toSQLId(name))) } plan.transformAllExpressionsWithPruning(_.containsPattern(PARAMETER)) { - case param @ Parameter(name) => - if (args.contains(name)) { - args(name) - } else { - param.failAnalysis( - errorClass = "UNBOUND_SQL_PARAMETER", - messageParameters = Map("name" -> name)) - } + case Parameter(name) if args.contains(name) => args(name) } } else { plan diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryExecutionErrors.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryExecutionErrors.scala index ea63f276b95b9..15dfa581c5976 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryExecutionErrors.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/errors/QueryExecutionErrors.scala @@ -2827,10 +2827,4 @@ private[sql] object QueryExecutionErrors extends QueryErrorsBase { "location" -> toSQLValue(location.toString, StringType), "identifier" -> toSQLId(tableId.nameParts))) } - - def unboundParameterError(name: String): Throwable = { - new SparkRuntimeException( - errorClass = "UNBOUND_SQL_PARAMETER", - messageParameters = Map("name" -> name)) - } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/ParametersSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/ParametersSuite.scala index 6490c9ad683cf..6e4150c56b580 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/ParametersSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/ParametersSuite.scala @@ -17,7 +17,6 @@ package org.apache.spark.sql -import org.apache.spark.SparkRuntimeException import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.test.SharedSparkSession @@ -48,17 +47,21 @@ class ParametersSuite extends QueryTest with SharedSparkSession { spark.sql("select :abc, :def", Map("abc" -> "1")) }, errorClass = "UNBOUND_SQL_PARAMETER", - parameters = Map("name" -> "def"), + parameters = Map("name" -> "`def`"), context = ExpectedContext( fragment = ":def", start = 13, stop = 16)) checkError( - exception = intercept[SparkRuntimeException] { + exception = intercept[AnalysisException] { sql("select :abc").collect() }, errorClass = "UNBOUND_SQL_PARAMETER", - parameters = Map("name" -> "abc")) + parameters = Map("name" -> "`abc`"), + context = ExpectedContext( + fragment = ":abc", + start = 7, + stop = 10)) } test("non-literal argument of `sql()`") { From 81cb619184cf6a2269237b09d4dc3f93594ca532 Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Mon, 12 Dec 2022 18:06:32 +0300 Subject: [PATCH 33/38] Use identifier() to get parameter name --- .../org/apache/spark/sql/catalyst/parser/AstBuilder.scala | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala index 88b6fe7255905..96a0460aa53af 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/AstBuilder.scala @@ -4842,8 +4842,6 @@ class AstBuilder extends SqlBaseParserBaseVisitor[AnyRef] with SQLConfHelper wit * Create a named parameter which represents a literal with a non-bound value and unknown type. * */ override def visitParameterLiteral(ctx: ParameterLiteralContext): Expression = withOrigin(ctx) { - val name = ctx.getText - assert(name.length > 1) - Parameter(name.substring(1)) + Parameter(ctx.identifier().getText) } } From 4b7cb237c02067d6a1a826f514df8711ceb7bd2b Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Mon, 12 Dec 2022 18:29:53 +0300 Subject: [PATCH 34/38] Extend Unevaluable --- .../sql/catalyst/expressions/parameters.scala | 16 +--------------- 1 file changed, 1 insertion(+), 15 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index 26f70af95a517..a8ca53b8196f9 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -17,10 +17,7 @@ package org.apache.spark.sql.catalyst.expressions -import org.apache.spark.SparkException -import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.analysis.AnalysisErrorAt -import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, ExprCode} import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan import org.apache.spark.sql.catalyst.trees.TreePattern.{PARAMETER, TreePattern} import org.apache.spark.sql.errors.QueryErrorsBase @@ -32,22 +29,11 @@ import org.apache.spark.sql.types.{DataType, NullType} * * @param name The identifier of the parameter without the marker. */ -case class Parameter(name: String) extends LeafExpression { - +case class Parameter(name: String) extends LeafExpression with Unevaluable { override lazy val resolved: Boolean = false override def dataType: DataType = NullType override def nullable: Boolean = true - final override val nodePatterns: Seq[TreePattern] = Seq(PARAMETER) - - private def unboundError(): Nothing = { - throw SparkException.internalError( - s"The parameter `$name` must be bound at the analysis phase.") - } - - override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = unboundError() - - def eval(input: InternalRow): Any = unboundError() } From 0a3e500351c1513454b1da3f2f561a7dd75b042c Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Mon, 12 Dec 2022 19:06:13 +0300 Subject: [PATCH 35/38] Throw an internal error for dataType() and nullable() --- .../sql/catalyst/expressions/parameters.scala | 18 ++++++++++++------ 1 file changed, 12 insertions(+), 6 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index a8ca53b8196f9..788119b9b5575 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -17,29 +17,35 @@ package org.apache.spark.sql.catalyst.expressions +import org.apache.spark.SparkException import org.apache.spark.sql.catalyst.analysis.AnalysisErrorAt import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan import org.apache.spark.sql.catalyst.trees.TreePattern.{PARAMETER, TreePattern} import org.apache.spark.sql.errors.QueryErrorsBase import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.{DataType, NullType} +import org.apache.spark.sql.types.DataType /** - * The expression represents a named parameter that should be replaces by a foldable expression. + * The expression represents a named parameter that should be replaces by a literal. * * @param name The identifier of the parameter without the marker. */ case class Parameter(name: String) extends LeafExpression with Unevaluable { override lazy val resolved: Boolean = false - override def dataType: DataType = NullType - override def nullable: Boolean = true + + private def unboundError(methodName: String): Nothing = { + throw SparkException.internalError( + s"Cannot call `$methodName()` of the unbound parameter `$name`.") + } + override def dataType: DataType = unboundError("dataType") + override def nullable: Boolean = unboundError("nullable") + final override val nodePatterns: Seq[TreePattern] = Seq(PARAMETER) } /** - * Finds all named parameters in the given plan and substitutes them by - * foldable expressions of `args` values. + * Finds all named parameters in the given plan and substitutes them by literals of `args` values. */ object Parameter extends QueryErrorsBase { def bind(plan: LogicalPlan, args: Map[String, Expression]): LogicalPlan = { From 15e2b27ad77daf565beaf0a7cb556b7c74ba617b Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Mon, 12 Dec 2022 19:31:12 +0300 Subject: [PATCH 36/38] Remove the SQL config --- .../sql/catalyst/parser/SqlBaseParser.g4 | 7 +-- .../sql/catalyst/expressions/parameters.scala | 3 +- .../sql/catalyst/parser/ParseDriver.scala | 1 - .../apache/spark/sql/internal/SQLConf.scala | 10 ---- .../sql/catalyst/analysis/AnalysisSuite.scala | 24 ++++---- .../sql/catalyst/parser/PlanParserSuite.scala | 56 ++++++++----------- .../apache/spark/sql/ParametersSuite.scala | 29 +++++----- 7 files changed, 49 insertions(+), 81 deletions(-) diff --git a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 index 79c934286cfd9..078a993911698 100644 --- a/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 +++ b/sql/catalyst/src/main/antlr4/org/apache/spark/sql/catalyst/parser/SqlBaseParser.g4 @@ -40,11 +40,6 @@ options { tokenVocab = SqlBaseLexer; } * When true, double quoted literals are identifiers rather than STRINGs. */ public boolean double_quoted_identifiers = false; - - /** - * When true, identifiers that begin from `:` are considered as named parameters. - */ - public boolean parameters_enabled = false; } singleStatement @@ -935,7 +930,7 @@ primaryExpression constant : NULL #nullLiteral - | {parameters_enabled}? COLON identifier #parameterLiteral + | COLON identifier #parameterLiteral | interval #intervalLiteral | identifier stringLit #typeConstructor | number #numericLiteral diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index 788119b9b5575..b7fc003b222bc 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -22,7 +22,6 @@ import org.apache.spark.sql.catalyst.analysis.AnalysisErrorAt import org.apache.spark.sql.catalyst.plans.logical.LogicalPlan import org.apache.spark.sql.catalyst.trees.TreePattern.{PARAMETER, TreePattern} import org.apache.spark.sql.errors.QueryErrorsBase -import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.DataType /** @@ -49,7 +48,7 @@ case class Parameter(name: String) extends LeafExpression with Unevaluable { */ object Parameter extends QueryErrorsBase { def bind(plan: LogicalPlan, args: Map[String, Expression]): LogicalPlan = { - if (!args.isEmpty && SQLConf.get.parametersEnabled) { + if (!args.isEmpty) { args.filter(!_._2.isInstanceOf[Literal]).headOption.foreach { case (name, expr) => expr.failAnalysis( errorClass = "INVALID_SQL_ARG", diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseDriver.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseDriver.scala index 31aa4ba537991..727d35d5c9152 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseDriver.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/parser/ParseDriver.scala @@ -119,7 +119,6 @@ abstract class AbstractSqlParser extends ParserInterface with SQLConfHelper with parser.legacy_exponent_literal_as_decimal_enabled = conf.exponentLiteralAsDecimalEnabled parser.SQL_standard_keyword_behavior = conf.enforceReservedKeywords parser.double_quoted_identifiers = conf.doubleQuotedIdentifiers - parser.parameters_enabled = conf.parametersEnabled try { try { diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala index c130b342e5174..84d78f365acbc 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala @@ -4027,14 +4027,6 @@ object SQLConf { .checkValues(ErrorMessageFormat.values.map(_.toString)) .createWithDefault(ErrorMessageFormat.PRETTY.toString) - val PARAMETERS_ENABLED = buildConf("spark.sql.parameters.enabled") - .doc("When set to true, `spark.sql()` executes input SQL queries by substituting named " + - "parameters by given literal values. If set to false, Spark handles constants " + - "with the `:` prefix as regular identifiers and does not consider them as parameters.") - .version("3.4.0") - .booleanConf - .createWithDefault(true) - /** * Holds information about keys that have been deprecated. * @@ -4846,8 +4838,6 @@ class SQLConf extends Serializable with Logging { def allowsTempViewCreationWithMultipleNameparts: Boolean = getConf(SQLConf.ALLOW_TEMP_VIEW_CREATION_WITH_MULTIPLE_NAME_PARTS) - def parametersEnabled: Boolean = getConf(SQLConf.PARAMETERS_ENABLED) - /** ********************** SQLConf functionality methods ************ */ /** Set Spark SQL configuration properties. */ diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala index 89c2963961bda..b7cb7fa59ca1a 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/analysis/AnalysisSuite.scala @@ -1297,18 +1297,16 @@ class AnalysisSuite extends AnalysisTest with Matchers { } test("SPARK-41271: bind named parameters to literals") { - withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { - comparePlans( - Parameter.bind( - plan = parsePlan("SELECT * FROM a LIMIT :limitA"), - args = Map("limitA" -> Literal(10))), - parsePlan("SELECT * FROM a LIMIT 10")) - // Ignore unused arguments - comparePlans( - Parameter.bind( - plan = parsePlan("SELECT c FROM a WHERE c < :param2"), - args = Map("param1" -> Literal(10), "param2" -> Literal(20))), - parsePlan("SELECT c FROM a WHERE c < 20")) - } + comparePlans( + Parameter.bind( + plan = parsePlan("SELECT * FROM a LIMIT :limitA"), + args = Map("limitA" -> Literal(10))), + parsePlan("SELECT * FROM a LIMIT 10")) + // Ignore unused arguments + comparePlans( + Parameter.bind( + plan = parsePlan("SELECT c FROM a WHERE c < :param2"), + args = Map("param1" -> Literal(10), "param2" -> Literal(20))), + parsePlan("SELECT c FROM a WHERE c < 20")) } } diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala index 2e109c59cf4ed..035e623117853 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/parser/PlanParserSuite.scala @@ -1570,38 +1570,28 @@ class PlanParserSuite extends AnalysisTest { } test("SPARK-41271: parsing of named parameters") { - withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { - comparePlans( - parsePlan("SELECT :param_1"), - Project(UnresolvedAlias(Parameter("param_1"), None) :: Nil, OneRowRelation())) - comparePlans( - parsePlan("SELECT abs(:1Abc)"), - Project(UnresolvedAlias( - UnresolvedFunction( - "abs" :: Nil, - Parameter("1Abc") :: Nil, - isDistinct = false), None) :: Nil, - OneRowRelation())) - comparePlans( - parsePlan("SELECT * FROM a LIMIT :limitA"), - table("a").select(star()).limit(Parameter("limitA"))) - // Invalid empty name and invalid symbol in a name - checkError( - exception = parseException(s"SELECT :-"), - errorClass = "PARSE_SYNTAX_ERROR", - parameters = Map("error" -> "'-'", "hint" -> "")) - checkError( - exception = parseException(s"SELECT :"), - errorClass = "PARSE_SYNTAX_ERROR", - parameters = Map("error" -> "end of input", "hint" -> "")) - } - withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "false") { - checkError( - exception = intercept[ParseException] { - parsePlan("SELECT :param_1") - }, - errorClass = "PARSE_SYNTAX_ERROR", - parameters = Map("error" -> "':'", "hint" -> "")) - } + comparePlans( + parsePlan("SELECT :param_1"), + Project(UnresolvedAlias(Parameter("param_1"), None) :: Nil, OneRowRelation())) + comparePlans( + parsePlan("SELECT abs(:1Abc)"), + Project(UnresolvedAlias( + UnresolvedFunction( + "abs" :: Nil, + Parameter("1Abc") :: Nil, + isDistinct = false), None) :: Nil, + OneRowRelation())) + comparePlans( + parsePlan("SELECT * FROM a LIMIT :limitA"), + table("a").select(star()).limit(Parameter("limitA"))) + // Invalid empty name and invalid symbol in a name + checkError( + exception = parseException(s"SELECT :-"), + errorClass = "PARSE_SYNTAX_ERROR", + parameters = Map("error" -> "'-'", "hint" -> "")) + checkError( + exception = parseException(s"SELECT :"), + errorClass = "PARSE_SYNTAX_ERROR", + parameters = Map("error" -> "end of input", "hint" -> "")) } } diff --git a/sql/core/src/test/scala/org/apache/spark/sql/ParametersSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/ParametersSuite.scala index 6e4150c56b580..668a1e4ad7d96 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/ParametersSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/ParametersSuite.scala @@ -17,28 +17,25 @@ package org.apache.spark.sql -import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.test.SharedSparkSession class ParametersSuite extends QueryTest with SharedSparkSession { test("bind parameters") { - withSQLConf(SQLConf.PARAMETERS_ENABLED.key -> "true") { - val sqlText = - """ - |SELECT id, id % :div as c0 - |FROM VALUES (0), (1), (2), (3), (4), (5), (6), (7), (8), (9) AS t(id) - |WHERE id < :constA - |""".stripMargin - val args = Map("div" -> "3", "constA" -> "4L") - checkAnswer( - spark.sql(sqlText, args), - Row(0, 0) :: Row(1, 1) :: Row(2, 2) :: Row(3, 0) :: Nil) + val sqlText = + """ + |SELECT id, id % :div as c0 + |FROM VALUES (0), (1), (2), (3), (4), (5), (6), (7), (8), (9) AS t(id) + |WHERE id < :constA + |""".stripMargin + val args = Map("div" -> "3", "constA" -> "4L") + checkAnswer( + spark.sql(sqlText, args), + Row(0, 0) :: Row(1, 1) :: Row(2, 2) :: Row(3, 0) :: Nil) - checkAnswer( - spark.sql("""SELECT contains('Spark \'SQL\'', :subStr)""", Map("subStr" -> "'SQL'")), - Row(true)) - } + checkAnswer( + spark.sql("""SELECT contains('Spark \'SQL\'', :subStr)""", Map("subStr" -> "'SQL'")), + Row(true)) } test("non-substituted parameters") { From 165a21dfc96e1474da3ffed0e0ccf5bdc957f77c Mon Sep 17 00:00:00 2001 From: Max Gekk Date: Mon, 12 Dec 2022 19:34:54 +0300 Subject: [PATCH 37/38] Remove a test from JavaSparkSessionSuite.java --- .../apache/spark/sql/JavaSparkSessionSuite.java | 16 ---------------- 1 file changed, 16 deletions(-) diff --git a/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java b/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java index b9e29897c26ef..b1df377936dfa 100644 --- a/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java +++ b/sql/core/src/test/java/test/org/apache/spark/sql/JavaSparkSessionSuite.java @@ -23,7 +23,6 @@ import org.junit.Test; import java.util.HashMap; -import java.util.List; import java.util.Map; public class JavaSparkSessionSuite { @@ -55,19 +54,4 @@ public void config() { Assert.assertEquals(spark.conf().get(e.getKey()), e.getValue().toString()); } } - - @Test - public void sqlParameters() { - spark = SparkSession.builder().master("local[*]").appName("testing").getOrCreate(); - Map params = new HashMap(); - params.put("_i1", "INTERVAL '1-1' YEAR TO MONTH"); - params.put("p2", "'a\"bc'"); - Dataset ds = spark.sql( - "SELECT :p2, i FROM VALUES (INTERVAL '2-2' YEAR TO MONTH) AS t(i) WHERE i > :_i1", - params); - List rows = ds.collectAsList(); - Assert.assertEquals(1, rows.size()); - Assert.assertEquals("a\"bc", rows.get(0).getString(0)); - Assert.assertEquals(java.time.Period.of(2, 2, 0), rows.get(0).get(1)); - } } From 2857350451e3d2df2410ff21e223659ec837f1f2 Mon Sep 17 00:00:00 2001 From: Maxim Gekk Date: Thu, 15 Dec 2022 08:08:54 +0300 Subject: [PATCH 38/38] Update sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala Co-authored-by: Wenchen Fan --- .../org/apache/spark/sql/catalyst/expressions/parameters.scala | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala index b7fc003b222bc..fae2b9a1a9f46 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/parameters.scala @@ -25,7 +25,7 @@ import org.apache.spark.sql.errors.QueryErrorsBase import org.apache.spark.sql.types.DataType /** - * The expression represents a named parameter that should be replaces by a literal. + * The expression represents a named parameter that should be replaced by a literal. * * @param name The identifier of the parameter without the marker. */