uros-b commented on code in PR #58119:
URL: https://github.com/apache/spark/pull/58119#discussion_r3872870504
##########
sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/optimizer/RewriteWithExpression.scala:
##########
@@ -66,6 +66,105 @@ object RewriteWithExpression extends Rule[LogicalPlan] {
}
}
+ /**
+ * Rewrites the `With` expressions in a single expression tree by inlining
their common
+ * expressions. Uses `transformUp` to handle nested `With`.
+ *
+ * Inlining duplicates a definition at every reference, so a definition
referenced more than once
+ * must be safe to duplicate (see `isSafeToDuplicate`); anything else
throws, as there is no plan
+ * here to pre-evaluate it into. Run [[ComputeCurrentTime]],
[[ReplaceCurrentLike]] and
+ * [[SpecialDatetimeValues]] first to fold the current date/time and catalog
expressions to
+ * literals.
+ *
+ * Does not descend into subquery plans (e.g. `ScalarSubquery`). A caller
whose expression
+ * may contain a subquery must rewrite those plans separately.
+ */
+ private[sql] def applyForExpression(expression: Expression): Expression =
+ inlineWith(expression, checkDuplication = true)
+
+ // The plan-level rewrite shares this to inline `With` in conditional
branches, which may not be
+ // evaluated and so can't be pulled into a Project; it passes false to
inline unconditionally.
+ private def inlineWith(expression: Expression, checkDuplication: Boolean):
Expression = {
+ // Which definitions are safe to duplicate can only be decided
outer-first, so it is collected
+ // in a separate pass before inlining bottom-up.
+ val safeIds = if (checkDuplication) {
+ safeToDuplicateIds(expression)
+ } else {
+ Set.empty[CommonExpressionId]
+ }
+ expression.transformUpWithPruning(_.containsPattern(WITH_EXPRESSION)) {
+ case With(child, defs) =>
+ if (checkDuplication) {
+ // Nested `With` is already rewritten, so a ref left in `child`
belongs to `defs` or to
+ // an enclosing `With`. Only the ids defined here are checked.
+ val refCounts = child.collect { case ref: CommonExpressionRef =>
ref.id }
+ .groupBy(identity)
+ .transform((_, refs) => refs.size)
+ defs.foreach { commonExprDef =>
+ if (refCounts.getOrElse(commonExprDef.id, 0) > 1 &&
+ !safeIds.contains(commonExprDef.id)) {
+ throw SparkException.internalError(
+ "Cannot inline a common expression definition that is
referenced more than " +
+ s"once and is not safe to duplicate:
${commonExprDef.child.sql}")
+ }
+ }
+ }
+ val refToExpr = defs.map(commonExprDef => commonExprDef.id ->
commonExprDef.child).toMap
+ child.transformWithPruning(_.containsPattern(COMMON_EXPR_REF)) {
+ case ref: CommonExpressionRef if refToExpr.contains(ref.id) =>
refToExpr(ref.id)
+ }
+ }
+ }
+
+ /**
+ * Ids of the definitions that are safe to duplicate. Walks outer-first, so
a definition that
+ * references an enclosing one is decided after it, and visits a
definition's own nested `With`
+ * before the definition itself. Ids are globally unique, so one flat set is
enough.
+ */
+ private def safeToDuplicateIds(expression: Expression):
Set[CommonExpressionId] = {
+ val safeIds = mutable.Set.empty[CommonExpressionId]
+ def visit(e: Expression): Unit = if (e.containsPattern(WITH_EXPRESSION)) {
+ e match {
+ case With(child, defs) =>
+ defs.foreach { commonExprDef =>
+ visit(commonExprDef.child)
+ if (isSafeToDuplicate(commonExprDef.child, safeIds)) {
+ safeIds += commonExprDef.id
+ }
+ }
+ visit(child)
+ case other => other.children.foreach(visit)
+ }
+ }
+ visit(expression)
+ safeIds.toSet
+ }
+
+ /**
+ * Whether a copy of `e` at every reference always evaluates to the same
value. `foldable` is not
+ * enough: `current_timestamp()` re-reads the clock, a TIME -> TIMESTAMP
cast derives the current
+ * date, `aes_encrypt` is a foldable `StaticInvoke` with a random IV, and a
Hive UDF is foldable
+ * on a user-declared determinism flag.
+ */
+ private def isSafeToDuplicate(
+ e: Expression,
+ safeIds: collection.Set[CommonExpressionId]): Boolean = {
+ // A whitelist: an unrecognized expression is rejected rather than assumed
safe.
+ def isLiteralTree(e: Expression): Boolean = e match {
+ case _: Literal => true
+ // A ref is safe exactly when the definition it points to is.
+ case ref: CommonExpressionRef => safeIds.contains(ref.id)
+ // Any other leaf varies per row (attribute) or reads a clock.
+ case _: LeafExpression => false
+ // Arbitrary Java methods and UDFs can be foldable without being pure.
+ case _: NonSQLExpression | _: UserDefinedExpression => false
+ case _ => e.deterministic && e.children.forall(isLiteralTree)
+ }
+ // `current_time()` and a TIME -> TIMESTAMP cast are not leaves and their
only leaf is a
+ // literal, so the tree walk accepts them and only these patterns reject
them.
+ !e.containsAnyPattern(CURRENT_LIKE, CAST_TO_TIMESTAMP) && isLiteralTree(e)
Review Comment:
CAST_TO_TIMESTAMP rejects every timestamp-target cast, although only TIME ->
TIMESTAMP is unsafe. A repeated CAST('1970-01-01' AS TIMESTAMP) incorrectly
throws.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]