This is an automated email from the ASF dual-hosted git repository.

hongze pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/incubator-gluten.git


The following commit(s) were added to refs/heads/main by this push:
     new 13d9a17078 [GLUTEN-8497][CORE] A unified CallInfo API to replace 
AdaptiveContext (#8551)
13d9a17078 is described below

commit 13d9a17078dc9aa293c13a01e14ea52fe33bac97
Author: Hongze Zhang <[email protected]>
AuthorDate: Fri Jan 17 12:00:35 2025 +0800

    [GLUTEN-8497][CORE] A unified CallInfo API to replace AdaptiveContext 
(#8551)
---
 .../gluten/backendsapi/clickhouse/CHRuleApi.scala  |   5 +-
 .../gluten/backendsapi/velox/VeloxRuleApi.scala    |   8 +-
 .../gluten/extension/caller/CallerInfo.scala       |  71 +++++++++++++
 .../extension/columnar/ColumnarRuleApplier.scala   |   4 +-
 .../extension/columnar/cost/GlutenCostModel.scala  |  18 +++-
 .../columnar/enumerated/EnumeratedApplier.scala    |  19 +---
 .../columnar/enumerated/EnumeratedTransform.scala  |   4 +-
 .../columnar/heuristic/HeuristicApplier.scala      |  55 ++++------
 .../columnar/heuristic/HeuristicTransform.scala    |   4 +-
 .../gluten/extension/injector/GlutenInjector.scala |   5 +-
 .../gluten/extension/util/AdaptiveContext.scala    |  86 ----------------
 .../sql/execution/FallbackStrategiesSuite.scala    | 113 +++++++++++----------
 .../sql/execution/FallbackStrategiesSuite.scala    | 113 +++++++++++----------
 .../sql/execution/FallbackStrategiesSuite.scala    | 113 +++++++++++----------
 .../sql/execution/FallbackStrategiesSuite.scala    | 113 +++++++++++----------
 15 files changed, 355 insertions(+), 376 deletions(-)

diff --git 
a/backends-clickhouse/src/main/scala/org/apache/gluten/backendsapi/clickhouse/CHRuleApi.scala
 
b/backends-clickhouse/src/main/scala/org/apache/gluten/backendsapi/clickhouse/CHRuleApi.scala
index 21ae342a22..c79931fa4e 100644
--- 
a/backends-clickhouse/src/main/scala/org/apache/gluten/backendsapi/clickhouse/CHRuleApi.scala
+++ 
b/backends-clickhouse/src/main/scala/org/apache/gluten/backendsapi/clickhouse/CHRuleApi.scala
@@ -120,11 +120,10 @@ object CHRuleApi {
     injector.injectPostTransform(c => 
AddPreProjectionForHashJoin.apply(c.session))
 
     // Gluten columnar: Fallback policies.
-    injector.injectFallbackPolicy(
-      c => ExpandFallbackPolicy(c.ac.isAdaptiveContext(), c.ac.originalPlan()))
+    injector.injectFallbackPolicy(c => p => 
ExpandFallbackPolicy(c.caller.isAqe(), p))
 
     // Gluten columnar: Post rules.
-    injector.injectPost(c => RemoveTopmostColumnarToRow(c.session, 
c.ac.isAdaptiveContext()))
+    injector.injectPost(c => RemoveTopmostColumnarToRow(c.session, 
c.caller.isAqe()))
     SparkShimLoader.getSparkShims
       .getExtendedColumnarPostRules()
       .foreach(each => injector.injectPost(c => intercept(each(c.session))))
diff --git 
a/backends-velox/src/main/scala/org/apache/gluten/backendsapi/velox/VeloxRuleApi.scala
 
b/backends-velox/src/main/scala/org/apache/gluten/backendsapi/velox/VeloxRuleApi.scala
index 0cf6ac6713..9825ae1d31 100644
--- 
a/backends-velox/src/main/scala/org/apache/gluten/backendsapi/velox/VeloxRuleApi.scala
+++ 
b/backends-velox/src/main/scala/org/apache/gluten/backendsapi/velox/VeloxRuleApi.scala
@@ -101,11 +101,10 @@ object VeloxRuleApi {
     injector.injectPostTransform(c => 
InsertTransitions.create(c.outputsColumnar, VeloxBatch))
 
     // Gluten columnar: Fallback policies.
-    injector.injectFallbackPolicy(
-      c => ExpandFallbackPolicy(c.ac.isAdaptiveContext(), c.ac.originalPlan()))
+    injector.injectFallbackPolicy(c => p => 
ExpandFallbackPolicy(c.caller.isAqe(), p))
 
     // Gluten columnar: Post rules.
-    injector.injectPost(c => RemoveTopmostColumnarToRow(c.session, 
c.ac.isAdaptiveContext()))
+    injector.injectPost(c => RemoveTopmostColumnarToRow(c.session, 
c.caller.isAqe()))
     SparkShimLoader.getSparkShims
       .getExtendedColumnarPostRules()
       .foreach(each => injector.injectPost(c => each(c.session)))
@@ -180,8 +179,7 @@ object VeloxRuleApi {
     injector.injectPostTransform(_ => CollapseProjectExecTransformer)
     injector.injectPostTransform(c => 
FlushableHashAggregateRule.apply(c.session))
     injector.injectPostTransform(c => 
InsertTransitions.create(c.outputsColumnar, VeloxBatch))
-    injector.injectPostTransform(
-      c => RemoveTopmostColumnarToRow(c.session, c.ac.isAdaptiveContext()))
+    injector.injectPostTransform(c => RemoveTopmostColumnarToRow(c.session, 
c.caller.isAqe()))
     SparkShimLoader.getSparkShims
       .getExtendedColumnarPostRules()
       .foreach(each => injector.injectPostTransform(c => each(c.session)))
diff --git 
a/gluten-core/src/main/scala/org/apache/gluten/extension/caller/CallerInfo.scala
 
b/gluten-core/src/main/scala/org/apache/gluten/extension/caller/CallerInfo.scala
new file mode 100644
index 0000000000..d23d172e46
--- /dev/null
+++ 
b/gluten-core/src/main/scala/org/apache/gluten/extension/caller/CallerInfo.scala
@@ -0,0 +1,71 @@
+/*
+ * 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.gluten.extension.caller
+
+import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanExec
+import org.apache.spark.sql.execution.columnar.InMemoryRelation
+
+/**
+ * Helper API that stores information about the call site of the columnar 
rule. Specific columnar
+ * rules could call the API to check whether this time of rule call was 
initiated for certain
+ * purpose. For example, a rule call could be for AQE optimization, or for 
cached plan optimization,
+ * or for regular executed plan optimization.
+ */
+trait CallerInfo {
+  def isAqe(): Boolean
+  def isCache(): Boolean
+}
+
+object CallerInfo {
+  private val localStorage: ThreadLocal[Option[CallerInfo]] =
+    new ThreadLocal[Option[CallerInfo]]() {
+      override def initialValue(): Option[CallerInfo] = None
+    }
+
+  private class Impl(override val isAqe: Boolean, override val isCache: 
Boolean) extends CallerInfo
+
+  /*
+   * Find the information about the caller that initiated the rule call.
+   */
+  def create(): CallerInfo = {
+    if (localStorage.get().nonEmpty) {
+      return localStorage.get().get
+    }
+    val stack = Thread.currentThread.getStackTrace
+    new Impl(isAqe = inAqeCall(stack), isCache = inCacheCall(stack))
+  }
+
+  private def inAqeCall(stack: Seq[StackTraceElement]): Boolean = {
+    stack.exists(_.getClassName.equals(AdaptiveSparkPlanExec.getClass.getName))
+  }
+
+  private def inCacheCall(stack: Seq[StackTraceElement]): Boolean = {
+    stack.exists(_.getClassName.equals(InMemoryRelation.getClass.getName))
+  }
+
+  /** For testing only. */
+  def withLocalValue[T](isAqe: Boolean, isCache: Boolean)(body: => T): T = {
+    val prevValue = localStorage.get()
+    val newValue = new Impl(isAqe, isCache)
+    localStorage.set(Some(newValue))
+    try {
+      body
+    } finally {
+      localStorage.set(prevValue)
+    }
+  }
+}
diff --git 
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/ColumnarRuleApplier.scala
 
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/ColumnarRuleApplier.scala
index 3e9bd72c45..7257865507 100644
--- 
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/ColumnarRuleApplier.scala
+++ 
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/ColumnarRuleApplier.scala
@@ -17,7 +17,7 @@
 package org.apache.gluten.extension.columnar
 
 import org.apache.gluten.config.GlutenConfig
-import org.apache.gluten.extension.util.AdaptiveContext
+import org.apache.gluten.extension.caller.CallerInfo
 
 import org.apache.spark.sql.SparkSession
 import org.apache.spark.sql.execution.SparkPlan
@@ -29,7 +29,7 @@ trait ColumnarRuleApplier {
 object ColumnarRuleApplier {
   class ColumnarRuleCall(
       val session: SparkSession,
-      val ac: AdaptiveContext,
+      val caller: CallerInfo,
       val outputsColumnar: Boolean) {
     val glutenConf: GlutenConfig = {
       new GlutenConfig(session.sessionState.conf)
diff --git 
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/cost/GlutenCostModel.scala
 
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/cost/GlutenCostModel.scala
index 80edf8919f..bfaf7380c6 100644
--- 
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/cost/GlutenCostModel.scala
+++ 
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/cost/GlutenCostModel.scala
@@ -21,11 +21,15 @@ import org.apache.gluten.component.Component
 import org.apache.spark.internal.Logging
 import org.apache.spark.sql.execution.SparkPlan
 import org.apache.spark.util.SparkReflectionUtil
-
+// format: off
 /**
  * The cost model API of Gluten. Used by:
- *   1. RAS planner for cost-based optimization; 2. Transition graph for 
choosing transition paths.
+ * <p>
+ *   1. RAS planner for cost-based optimization;
+ * <p>
+ *   2. Transition graph for choosing transition paths.
  */
+// format: on
 trait GlutenCostModel {
   def costOf(node: SparkPlan): GlutenCost
   def costComparator(): Ordering[GlutenCost]
@@ -38,14 +42,18 @@ trait GlutenCostModel {
 }
 
 object GlutenCostModel extends Logging {
-  def find(aliasOrClass: String): GlutenCostModel = {
-    val costModelRegistry = LongCostModel.registry()
+  private val costModelRegistry = {
+    val r = LongCostModel.registry()
     // Components should override Backend's costers. Hence, reversed 
registration order is applied.
     Component
       .sorted()
       .reverse
       .flatMap(_.costers())
-      .foreach(coster => costModelRegistry.register(coster))
+      .foreach(coster => r.register(coster))
+    r
+  }
+
+  def find(aliasOrClass: String): GlutenCostModel = {
     val costModel = find(costModelRegistry, aliasOrClass)
     costModel
   }
diff --git 
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedApplier.scala
 
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedApplier.scala
index 04cd70656e..2f5c8c4472 100644
--- 
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedApplier.scala
+++ 
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedApplier.scala
@@ -16,9 +16,9 @@
  */
 package org.apache.gluten.extension.columnar.enumerated
 
+import org.apache.gluten.extension.caller.CallerInfo
 import org.apache.gluten.extension.columnar.{ColumnarRuleApplier, 
ColumnarRuleExecutor}
 import 
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
-import org.apache.gluten.extension.util.AdaptiveContext
 import org.apache.gluten.logging.LogLevelUtil
 
 import org.apache.spark.internal.Logging
@@ -39,27 +39,14 @@ class EnumeratedApplier(
   extends ColumnarRuleApplier
   with Logging
   with LogLevelUtil {
-  private val adaptiveContext = AdaptiveContext(session)
-
   override def apply(plan: SparkPlan, outputsColumnar: Boolean): SparkPlan = {
-    val call = new ColumnarRuleCall(session, adaptiveContext, outputsColumnar)
-    val finalPlan = maybeAqe {
-      apply0(ruleBuilders.map(b => b(call)), plan)
-    }
+    val call = new ColumnarRuleCall(session, CallerInfo.create(), 
outputsColumnar)
+    val finalPlan = apply0(ruleBuilders.map(b => b(call)), plan)
     finalPlan
   }
 
   private def apply0(rules: Seq[Rule[SparkPlan]], plan: SparkPlan): SparkPlan =
     new ColumnarRuleExecutor("ras", rules).execute(plan)
-
-  private def maybeAqe[T](f: => T): T = {
-    adaptiveContext.setAdaptiveContext()
-    try {
-      f
-    } finally {
-      adaptiveContext.resetAdaptiveContext()
-    }
-  }
 }
 
 object EnumeratedApplier {}
diff --git 
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedTransform.scala
 
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedTransform.scala
index 72926407c5..6af9b00134 100644
--- 
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedTransform.scala
+++ 
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/enumerated/EnumeratedTransform.scala
@@ -18,12 +18,12 @@ package org.apache.gluten.extension.columnar.enumerated
 
 import org.apache.gluten.component.Component
 import org.apache.gluten.exception.GlutenException
+import org.apache.gluten.extension.caller.CallerInfo
 import 
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
 import org.apache.gluten.extension.columnar.cost.{GlutenCost, GlutenCostModel}
 import 
org.apache.gluten.extension.columnar.enumerated.planner.GlutenOptimization
 import org.apache.gluten.extension.columnar.enumerated.planner.property.Conv
 import org.apache.gluten.extension.injector.Injector
-import org.apache.gluten.extension.util.AdaptiveContext
 import org.apache.gluten.logging.LogLevelUtil
 import org.apache.gluten.ras.{Cost, CostModel}
 import org.apache.gluten.ras.property.PropertySet
@@ -81,7 +81,7 @@ object EnumeratedTransform {
     val session = SparkSession.getActiveSession.getOrElse(
       throw new GlutenException(
         "HeuristicTransform#static can only be called when an active Spark 
session exists"))
-    val call = new ColumnarRuleCall(session, AdaptiveContext(session), false)
+    val call = new ColumnarRuleCall(session, CallerInfo.create(), false)
     dummyInjector.gluten.ras.createEnumeratedTransform(call)
   }
 
diff --git 
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicApplier.scala
 
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicApplier.scala
index e4825d2eb7..493319d423 100644
--- 
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicApplier.scala
+++ 
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicApplier.scala
@@ -16,9 +16,9 @@
  */
 package org.apache.gluten.extension.columnar.heuristic
 
+import org.apache.gluten.extension.caller.CallerInfo
 import org.apache.gluten.extension.columnar.{ColumnarRuleApplier, 
ColumnarRuleExecutor}
 import 
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
-import org.apache.gluten.extension.util.AdaptiveContext
 import org.apache.gluten.logging.LogLevelUtil
 
 import org.apache.spark.internal.Logging
@@ -33,35 +33,33 @@ import org.apache.spark.sql.execution.SparkPlan
 class HeuristicApplier(
     session: SparkSession,
     transformBuilders: Seq[ColumnarRuleCall => Rule[SparkPlan]],
-    fallbackPolicyBuilders: Seq[ColumnarRuleCall => Rule[SparkPlan]],
+    fallbackPolicyBuilders: Seq[ColumnarRuleCall => SparkPlan => 
Rule[SparkPlan]],
     postBuilders: Seq[ColumnarRuleCall => Rule[SparkPlan]],
     finalBuilders: Seq[ColumnarRuleCall => Rule[SparkPlan]])
   extends ColumnarRuleApplier
   with Logging
   with LogLevelUtil {
-  private val adaptiveContext = AdaptiveContext(session)
-
   override def apply(plan: SparkPlan, outputsColumnar: Boolean): SparkPlan = {
-    val call = new ColumnarRuleCall(session, adaptiveContext, outputsColumnar)
+    val call = new ColumnarRuleCall(session, CallerInfo.create(), 
outputsColumnar)
     makeRule(call).apply(plan)
   }
 
   private def makeRule(call: ColumnarRuleCall): Rule[SparkPlan] = {
-    plan =>
-      prepareFallback(plan) {
-        p =>
-          val suggestedPlan = transformPlan("transform", transformRules(call), 
p)
-          val finalPlan = transformPlan("fallback", fallbackPolicies(call), 
suggestedPlan) match {
-            case FallbackNode(fallbackPlan) =>
-              // we should use vanilla c2r rather than native c2r,
-              // and there should be no `GlutenPlan` any more,
-              // so skip the `postRules()`.
-              fallbackPlan
-            case plan =>
-              transformPlan("post", postRules(call), plan)
-          }
-          transformPlan("final", finalRules(call), finalPlan)
+    originalPlan =>
+      val suggestedPlan = transformPlan("transform", transformRules(call), 
originalPlan)
+      val finalPlan = transformPlan(
+        "fallback",
+        fallbackPolicies(call).map(_(originalPlan)),
+        suggestedPlan) match {
+        case FallbackNode(fallbackPlan) =>
+          // we should use vanilla c2r rather than native c2r,
+          // and there should be no `GlutenPlan` anymore,
+          // so skip the `postRules()`.
+          fallbackPlan
+        case plan =>
+          transformPlan("post", postRules(call), plan)
       }
+      transformPlan("final", finalRules(call), finalPlan)
   }
 
   private def transformPlan(
@@ -70,17 +68,6 @@ class HeuristicApplier(
       plan: SparkPlan): SparkPlan =
     new ColumnarRuleExecutor(phase, rules).execute(plan)
 
-  private def prepareFallback[T](p: SparkPlan)(f: SparkPlan => T): T = {
-    adaptiveContext.setAdaptiveContext()
-    adaptiveContext.setOriginalPlan(p)
-    try {
-      f(p)
-    } finally {
-      adaptiveContext.resetOriginalPlan()
-      adaptiveContext.resetAdaptiveContext()
-    }
-  }
-
   /**
    * Rules to let planner create a suggested Gluten plan being sent to 
`fallbackPolicies` in which
    * the plan will be breakdown and decided to be fallen back or not.
@@ -93,7 +80,7 @@ class HeuristicApplier(
    * Rules to add wrapper `FallbackNode`s on top of the input plan, as hints 
to make planner fall
    * back the whole input plan to the original vanilla Spark plan.
    */
-  private def fallbackPolicies(call: ColumnarRuleCall): Seq[Rule[SparkPlan]] = 
{
+  private def fallbackPolicies(call: ColumnarRuleCall): Seq[SparkPlan => 
Rule[SparkPlan]] = {
     fallbackPolicyBuilders.map(b => b.apply(call))
   }
 
@@ -112,12 +99,6 @@ class HeuristicApplier(
   private def finalRules(call: ColumnarRuleCall): Seq[Rule[SparkPlan]] = {
     finalBuilders.map(b => b.apply(call))
   }
-
-  // Just for test use.
-  def enableAdaptiveContext(): HeuristicApplier = {
-    adaptiveContext.enableAdaptiveContext()
-    this
-  }
 }
 
 object HeuristicApplier {}
diff --git 
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicTransform.scala
 
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicTransform.scala
index 5a0fdfeefe..011453032b 100644
--- 
a/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicTransform.scala
+++ 
b/gluten-core/src/main/scala/org/apache/gluten/extension/columnar/heuristic/HeuristicTransform.scala
@@ -18,12 +18,12 @@ package org.apache.gluten.extension.columnar.heuristic
 
 import org.apache.gluten.component.Component
 import org.apache.gluten.exception.GlutenException
+import org.apache.gluten.extension.caller.CallerInfo
 import 
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
 import org.apache.gluten.extension.columnar.offload.OffloadSingleNode
 import org.apache.gluten.extension.columnar.rewrite.RewriteSingleNode
 import org.apache.gluten.extension.columnar.validator.Validator
 import org.apache.gluten.extension.injector.Injector
-import org.apache.gluten.extension.util.AdaptiveContext
 import org.apache.gluten.logging.LogLevelUtil
 
 import org.apache.spark.internal.Logging
@@ -126,7 +126,7 @@ object HeuristicTransform {
     val session = SparkSession.getActiveSession.getOrElse(
       throw new GlutenException(
         "HeuristicTransform#static can only be called when an active Spark 
session exists"))
-    val call = new ColumnarRuleCall(session, AdaptiveContext(session), false)
+    val call = new ColumnarRuleCall(session, CallerInfo.create(), false)
     dummyInjector.gluten.legacy.createHeuristicTransform(call)
   }
 }
diff --git 
a/gluten-core/src/main/scala/org/apache/gluten/extension/injector/GlutenInjector.scala
 
b/gluten-core/src/main/scala/org/apache/gluten/extension/injector/GlutenInjector.scala
index a208db2c96..fa8704509e 100644
--- 
a/gluten-core/src/main/scala/org/apache/gluten/extension/injector/GlutenInjector.scala
+++ 
b/gluten-core/src/main/scala/org/apache/gluten/extension/injector/GlutenInjector.scala
@@ -57,7 +57,8 @@ object GlutenInjector {
     private val preTransformBuilders = mutable.Buffer.empty[ColumnarRuleCall 
=> Rule[SparkPlan]]
     private val transformBuilders = mutable.Buffer.empty[ColumnarRuleCall => 
Rule[SparkPlan]]
     private val postTransformBuilders = mutable.Buffer.empty[ColumnarRuleCall 
=> Rule[SparkPlan]]
-    private val fallbackPolicyBuilders = mutable.Buffer.empty[ColumnarRuleCall 
=> Rule[SparkPlan]]
+    private val fallbackPolicyBuilders =
+      mutable.Buffer.empty[ColumnarRuleCall => SparkPlan => Rule[SparkPlan]]
     private val postBuilders = mutable.Buffer.empty[ColumnarRuleCall => 
Rule[SparkPlan]]
     private val finalBuilders = mutable.Buffer.empty[ColumnarRuleCall => 
Rule[SparkPlan]]
 
@@ -73,7 +74,7 @@ object GlutenInjector {
       postTransformBuilders += builder
     }
 
-    def injectFallbackPolicy(builder: ColumnarRuleCall => Rule[SparkPlan]): 
Unit = {
+    def injectFallbackPolicy(builder: ColumnarRuleCall => SparkPlan => 
Rule[SparkPlan]): Unit = {
       fallbackPolicyBuilders += builder
     }
 
diff --git 
a/gluten-core/src/main/scala/org/apache/gluten/extension/util/AdaptiveContext.scala
 
b/gluten-core/src/main/scala/org/apache/gluten/extension/util/AdaptiveContext.scala
deleted file mode 100644
index b0f42e7967..0000000000
--- 
a/gluten-core/src/main/scala/org/apache/gluten/extension/util/AdaptiveContext.scala
+++ /dev/null
@@ -1,86 +0,0 @@
-/*
- * 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.gluten.extension.util
-
-import org.apache.spark.sql.SparkSession
-import org.apache.spark.sql.execution.SparkPlan
-import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanExec
-
-import scala.collection.mutable.ListBuffer
-
-// Since: https://github.com/apache/incubator-gluten/pull/3294.
-sealed trait AdaptiveContext {
-  def enableAdaptiveContext(): Unit
-  def isAdaptiveContext(): Boolean
-  def setAdaptiveContext(): Unit
-  def resetAdaptiveContext(): Unit
-  def setOriginalPlan(plan: SparkPlan): Unit
-  def originalPlan(): SparkPlan
-  def resetOriginalPlan(): Unit
-}
-
-object AdaptiveContext {
-  def apply(session: SparkSession): AdaptiveContext =
-    new AdaptiveContextImpl(session)
-
-  private val GLUTEN_IS_ADAPTIVE_CONTEXT = "gluten.isAdaptiveContext"
-
-  // Holds the original plan for possible entire fallback.
-  private val localOriginalPlans: ThreadLocal[ListBuffer[SparkPlan]] =
-    ThreadLocal.withInitial(() => ListBuffer.empty[SparkPlan])
-  private val localIsAdaptiveContextFlags: ThreadLocal[ListBuffer[Boolean]] =
-    ThreadLocal.withInitial(() => ListBuffer.empty[Boolean])
-
-  private class AdaptiveContextImpl(session: SparkSession) extends 
AdaptiveContext {
-    // Just for test use.
-    override def enableAdaptiveContext(): Unit = {
-      session.sparkContext.setLocalProperty(GLUTEN_IS_ADAPTIVE_CONTEXT, "true")
-    }
-
-    override def isAdaptiveContext(): Boolean =
-      Option(session.sparkContext.getLocalProperty(GLUTEN_IS_ADAPTIVE_CONTEXT))
-        .getOrElse("false")
-        .toBoolean ||
-        localIsAdaptiveContextFlags.get().head
-
-    override def setAdaptiveContext(): Unit = {
-      val traceElements = Thread.currentThread.getStackTrace
-      // ApplyColumnarRulesAndInsertTransitions is called by either 
QueryExecution or
-      // AdaptiveSparkPlanExec. So by checking the stack trace, we can know 
whether
-      // columnar rule will be applied in adaptive execution context.
-      localIsAdaptiveContextFlags
-        .get()
-        .prepend(
-          
traceElements.exists(_.getClassName.equals(AdaptiveSparkPlanExec.getClass.getName)))
-    }
-
-    override def resetAdaptiveContext(): Unit =
-      localIsAdaptiveContextFlags.get().remove(0)
-
-    override def setOriginalPlan(plan: SparkPlan): Unit = {
-      localOriginalPlans.get().prepend(plan)
-    }
-
-    override def originalPlan(): SparkPlan = {
-      val plan = localOriginalPlans.get().head
-      assert(plan != null)
-      plan
-    }
-
-    override def resetOriginalPlan(): Unit = localOriginalPlans.get().remove(0)
-  }
-}
diff --git 
a/gluten-ut/spark32/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
 
b/gluten-ut/spark32/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
index 97a779289d..bbbc913bd8 100644
--- 
a/gluten-ut/spark32/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
+++ 
b/gluten-ut/spark32/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
@@ -20,6 +20,7 @@ import org.apache.gluten.backendsapi.BackendsApiManager
 import org.apache.gluten.config.GlutenConfig
 import org.apache.gluten.execution.{BasicScanExecTransformer, GlutenPlan}
 import org.apache.gluten.extension.GlutenSessionExtensions
+import org.apache.gluten.extension.caller.CallerInfo
 import 
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
 import 
org.apache.gluten.extension.columnar.MiscColumnarRules.RemoveTopmostColumnarToRow
 import org.apache.gluten.extension.columnar.RemoveFallbackTagRule
@@ -54,37 +55,39 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
 
   testGluten("Fall back the whole plan if meeting the configured threshold") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"1")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to fall back the entire plan.
-      assert(outputPlan == originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to fall back the entire plan.
+        assert(outputPlan == originalPlan)
+      }
     }
   }
 
   testGluten("Don't fall back the whole plan if NOT meeting the configured 
threshold") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"4")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to get the plan with columnar rule applied.
-      assert(outputPlan != originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to get the plan with columnar rule applied.
+        assert(outputPlan != originalPlan)
+      }
     }
   }
 
@@ -92,19 +95,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
     "Fall back the whole plan if meeting the configured threshold (leaf node 
is" +
       " transformable)") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"2")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to fall back the entire plan.
-      assert(outputPlan == originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to fall back the entire plan.
+        assert(outputPlan == originalPlan)
+      }
     }
   }
 
@@ -112,19 +116,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait 
{
     "Don't Fall back the whole plan if NOT meeting the configured threshold (" 
+
       "leaf node is transformable)") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"3")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to get the plan with columnar rule applied.
-      assert(outputPlan != originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to get the plan with columnar rule applied.
+        assert(outputPlan != originalPlan)
+      }
     }
   }
 
@@ -169,9 +174,9 @@ private object FallbackStrategiesSuite {
     new HeuristicApplier(
       spark,
       transformBuilders,
-      List(c => ExpandFallbackPolicy(c.ac.isAdaptiveContext(), 
c.ac.originalPlan())),
+      List(c => p => ExpandFallbackPolicy(c.caller.isAqe(), p)),
       List(
-        c => RemoveTopmostColumnarToRow(c.session, c.ac.isAdaptiveContext()),
+        c => RemoveTopmostColumnarToRow(c.session, c.caller.isAqe()),
         _ => ColumnarCollapseTransformStages(GlutenConfig.get)
       ),
       List(_ => RemoveFallbackTagRule())
diff --git 
a/gluten-ut/spark33/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
 
b/gluten-ut/spark33/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
index 8b68c36bb3..f51f2721f4 100644
--- 
a/gluten-ut/spark33/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
+++ 
b/gluten-ut/spark33/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
@@ -20,6 +20,7 @@ import org.apache.gluten.backendsapi.BackendsApiManager
 import org.apache.gluten.config.GlutenConfig
 import org.apache.gluten.execution.{BasicScanExecTransformer, GlutenPlan}
 import org.apache.gluten.extension.GlutenSessionExtensions
+import org.apache.gluten.extension.caller.CallerInfo
 import org.apache.gluten.extension.columnar.{FallbackTags, 
RemoveFallbackTagRule}
 import 
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
 import 
org.apache.gluten.extension.columnar.MiscColumnarRules.RemoveTopmostColumnarToRow
@@ -53,37 +54,39 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
 
   testGluten("Fall back the whole plan if meeting the configured threshold") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"1")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to fall back the entire plan.
-      assert(outputPlan == originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to fall back the entire plan.
+        assert(outputPlan == originalPlan)
+      }
     }
   }
 
   testGluten("Don't fall back the whole plan if NOT meeting the configured 
threshold") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"4")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to get the plan with columnar rule applied.
-      assert(outputPlan != originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to get the plan with columnar rule applied.
+        assert(outputPlan != originalPlan)
+      }
     }
   }
 
@@ -91,19 +94,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
     "Fall back the whole plan if meeting the configured threshold (leaf node 
is" +
       " transformable)") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"2")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to fall back the entire plan.
-      assert(outputPlan == originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to fall back the entire plan.
+        assert(outputPlan == originalPlan)
+      }
     }
   }
 
@@ -111,19 +115,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait 
{
     "Don't Fall back the whole plan if NOT meeting the configured threshold (" 
+
       "leaf node is transformable)") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"3")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to get the plan with columnar rule applied.
-      assert(outputPlan != originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to get the plan with columnar rule applied.
+        assert(outputPlan != originalPlan)
+      }
     }
   }
 
@@ -179,9 +184,9 @@ private object FallbackStrategiesSuite {
     new HeuristicApplier(
       spark,
       transformBuilders,
-      List(c => ExpandFallbackPolicy(c.ac.isAdaptiveContext(), 
c.ac.originalPlan())),
+      List(c => p => ExpandFallbackPolicy(c.caller.isAqe(), p)),
       List(
-        c => RemoveTopmostColumnarToRow(c.session, c.ac.isAdaptiveContext()),
+        c => RemoveTopmostColumnarToRow(c.session, c.caller.isAqe()),
         _ => ColumnarCollapseTransformStages(GlutenConfig.get)
       ),
       List(_ => RemoveFallbackTagRule())
diff --git 
a/gluten-ut/spark34/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
 
b/gluten-ut/spark34/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
index 8b68c36bb3..f51f2721f4 100644
--- 
a/gluten-ut/spark34/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
+++ 
b/gluten-ut/spark34/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
@@ -20,6 +20,7 @@ import org.apache.gluten.backendsapi.BackendsApiManager
 import org.apache.gluten.config.GlutenConfig
 import org.apache.gluten.execution.{BasicScanExecTransformer, GlutenPlan}
 import org.apache.gluten.extension.GlutenSessionExtensions
+import org.apache.gluten.extension.caller.CallerInfo
 import org.apache.gluten.extension.columnar.{FallbackTags, 
RemoveFallbackTagRule}
 import 
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
 import 
org.apache.gluten.extension.columnar.MiscColumnarRules.RemoveTopmostColumnarToRow
@@ -53,37 +54,39 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
 
   testGluten("Fall back the whole plan if meeting the configured threshold") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"1")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to fall back the entire plan.
-      assert(outputPlan == originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to fall back the entire plan.
+        assert(outputPlan == originalPlan)
+      }
     }
   }
 
   testGluten("Don't fall back the whole plan if NOT meeting the configured 
threshold") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"4")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to get the plan with columnar rule applied.
-      assert(outputPlan != originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to get the plan with columnar rule applied.
+        assert(outputPlan != originalPlan)
+      }
     }
   }
 
@@ -91,19 +94,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
     "Fall back the whole plan if meeting the configured threshold (leaf node 
is" +
       " transformable)") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"2")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to fall back the entire plan.
-      assert(outputPlan == originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to fall back the entire plan.
+        assert(outputPlan == originalPlan)
+      }
     }
   }
 
@@ -111,19 +115,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait 
{
     "Don't Fall back the whole plan if NOT meeting the configured threshold (" 
+
       "leaf node is transformable)") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"3")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to get the plan with columnar rule applied.
-      assert(outputPlan != originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to get the plan with columnar rule applied.
+        assert(outputPlan != originalPlan)
+      }
     }
   }
 
@@ -179,9 +184,9 @@ private object FallbackStrategiesSuite {
     new HeuristicApplier(
       spark,
       transformBuilders,
-      List(c => ExpandFallbackPolicy(c.ac.isAdaptiveContext(), 
c.ac.originalPlan())),
+      List(c => p => ExpandFallbackPolicy(c.caller.isAqe(), p)),
       List(
-        c => RemoveTopmostColumnarToRow(c.session, c.ac.isAdaptiveContext()),
+        c => RemoveTopmostColumnarToRow(c.session, c.caller.isAqe()),
         _ => ColumnarCollapseTransformStages(GlutenConfig.get)
       ),
       List(_ => RemoveFallbackTagRule())
diff --git 
a/gluten-ut/spark35/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
 
b/gluten-ut/spark35/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
index 898967d4f3..1dd7eccc21 100644
--- 
a/gluten-ut/spark35/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
+++ 
b/gluten-ut/spark35/src/test/scala/org/apache/spark/sql/execution/FallbackStrategiesSuite.scala
@@ -20,6 +20,7 @@ import org.apache.gluten.backendsapi.BackendsApiManager
 import org.apache.gluten.config.GlutenConfig
 import org.apache.gluten.execution.{BasicScanExecTransformer, GlutenPlan}
 import org.apache.gluten.extension.GlutenSessionExtensions
+import org.apache.gluten.extension.caller.CallerInfo
 import org.apache.gluten.extension.columnar.{FallbackTags, 
RemoveFallbackTagRule}
 import 
org.apache.gluten.extension.columnar.ColumnarRuleApplier.ColumnarRuleCall
 import 
org.apache.gluten.extension.columnar.MiscColumnarRules.RemoveTopmostColumnarToRow
@@ -54,37 +55,39 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
 
   testGluten("Fall back the whole plan if meeting the configured threshold") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"1")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to fall back the entire plan.
-      assert(outputPlan == originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to fall back the entire plan.
+        assert(outputPlan == originalPlan)
+      }
     }
   }
 
   testGluten("Don't fall back the whole plan if NOT meeting the configured 
threshold") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"4")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to get the plan with columnar rule applied.
-      assert(outputPlan != originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOp()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to get the plan with columnar rule applied.
+        assert(outputPlan != originalPlan)
+      }
     }
   }
 
@@ -92,19 +95,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait {
     "Fall back the whole plan if meeting the configured threshold (leaf node 
is" +
       " transformable)") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"2")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to fall back the entire plan.
-      assert(outputPlan == originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to fall back the entire plan.
+        assert(outputPlan == originalPlan)
+      }
     }
   }
 
@@ -112,19 +116,20 @@ class FallbackStrategiesSuite extends GlutenSQLTestsTrait 
{
     "Don't Fall back the whole plan if NOT meeting the configured threshold (" 
+
       "leaf node is transformable)") {
     withSQLConf(("spark.gluten.sql.columnar.wholeStage.fallback.threshold", 
"3")) {
-      val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
-      val rule = newRuleApplier(
-        spark,
-        List(
-          _ =>
-            _ => {
-              
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
-            },
-          c => InsertBackendTransitions(c.outputsColumnar)))
-        .enableAdaptiveContext()
-      val outputPlan = rule.apply(originalPlan, false)
-      // Expect to get the plan with columnar rule applied.
-      assert(outputPlan != originalPlan)
+      CallerInfo.withLocalValue(isAqe = true, isCache = false) {
+        val originalPlan = UnaryOp2(UnaryOp1(UnaryOp2(UnaryOp1(LeafOp()))))
+        val rule = newRuleApplier(
+          spark,
+          List(
+            _ =>
+              _ => {
+                
UnaryOp2(UnaryOp1Transformer(UnaryOp2(UnaryOp1Transformer(LeafOpTransformer()))))
+              },
+            c => InsertBackendTransitions(c.outputsColumnar)))
+        val outputPlan = rule.apply(originalPlan, false)
+        // Expect to get the plan with columnar rule applied.
+        assert(outputPlan != originalPlan)
+      }
     }
   }
 
@@ -180,9 +185,9 @@ private object FallbackStrategiesSuite {
     new HeuristicApplier(
       spark,
       transformBuilders,
-      List(c => ExpandFallbackPolicy(c.ac.isAdaptiveContext(), 
c.ac.originalPlan())),
+      List(c => p => ExpandFallbackPolicy(c.caller.isAqe(), p)),
       List(
-        c => RemoveTopmostColumnarToRow(c.session, c.ac.isAdaptiveContext()),
+        c => RemoveTopmostColumnarToRow(c.session, c.caller.isAqe()),
         _ => ColumnarCollapseTransformStages(GlutenConfig.get)
       ),
       List(_ => RemoveFallbackTagRule())


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to