Copilot commented on code in PR #13109:
URL: https://github.com/apache/gluten/pull/13109#discussion_r4080904928
##########
cpp/velox/substrait/SubstraitToVeloxPlan.cc:
##########
@@ -1566,17 +1567,20 @@ core::PlanNodePtr
SubstraitToVeloxPlanConverter::toVeloxPlan(const ::substrait::
std::vector<std::string> colNameList;
std::vector<TypePtr> veloxTypeList;
std::vector<ColumnType> columnTypes;
- // Convert field names into lower case when not case-sensitive.
+ // Convert field names into lower case when not case-sensitive. Every name
that has to match a
+ // scan column (schema names, required-subfield column and field names) goes
through foldCase.
bool asLowerCase = !veloxCfg_->get<bool>(kCaseSensitive, false);
+ auto foldCase = [asLowerCase](std::string name) {
+ if (asLowerCase) {
+ folly::toLowerAscii(name);
+ }
+ return name;
+ };
Review Comment:
`foldCase` uses `folly::toLowerAscii`, which only lowercases ASCII. On the
JVM side, the rule’s comments/tests indicate case folding should follow
Spark/ConverterUtils normalization (the suite includes a non-ASCII folding
case), so using ASCII-only folding here can cause required-subfield
column/field names (or schema names) to stop matching in case-insensitive mode,
leading to required subfields being ignored or resolving to nulls. Use the same
case-folding strategy as the Substrait schema/type parsing path (and as the
Scala normalization), or remove this extra folding if Substrait names are
already normalized consistently by the producer.
##########
backends-velox/src/main/scala/org/apache/gluten/backendsapi/velox/VeloxTransformerApi.scala:
##########
@@ -152,6 +152,32 @@ class VeloxTransformerApi extends TransformerApi with
Logging {
packPBMessage(extensionBuilder.build())
}
+ override def packRequiredSubfields(subfields: Map[String,
Seq[SubfieldPath]]): Any = {
+ val extension = RequiredSubfieldsExtension.newBuilder()
+ subfields.toSeq.sortBy(_._1).foreach {
+ case (column, paths) =>
+ val columnBuilder =
+
RequiredSubfieldsExtension.ColumnSubfields.newBuilder().setColumn(column)
+ paths.foreach {
+ path =>
+ val subfield = RequiredSubfieldsExtension.Subfield.newBuilder()
+ path.elements.foreach {
+ element =>
+ val pe = RequiredSubfieldsExtension.PathElement.newBuilder()
+ element match {
+ case SubfieldElement.Field(name) => pe.setField(name)
+ case SubfieldElement.StringKey(key) => pe.setStringKey(key)
+ case SubfieldElement.LongKey(key) => pe.setLongKey(key)
+ }
+ subfield.addElements(pe)
+ }
+ columnBuilder.addSubfields(subfield)
+ }
+ extension.addColumns(columnBuilder)
+ }
+ packPBMessage(extension.build())
+ }
Review Comment:
`requiredMapSubfields` redundantly encodes the column name both as the `Map`
key (`column`) and inside each `SubfieldPath` (`path.column`).
`packRequiredSubfields` currently ignores `path.column` entirely, which makes
it easy for callers to accidentally construct inconsistent data that silently
serializes under the `Map` key. Add a defensive check (e.g., require
`path.column == column`) or remove the redundant `column` field from
`SubfieldPath` and treat it purely as “elements under a known column” to
prevent hard-to-debug mismatches.
##########
backends-velox/src/test/scala/org/apache/gluten/execution/ScanMapKeyPruningSuite.scala:
##########
@@ -0,0 +1,630 @@
+/*
+ * 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.execution
+
+import org.apache.gluten.config.{GlutenConfig, VeloxConfig}
+
+import org.apache.spark.SparkConf
+import org.apache.spark.sql.{DataFrame, Row}
+import org.apache.spark.sql.execution.{BaseScriptTransformationExec, SparkPlan}
+import org.apache.spark.sql.execution.exchange.ReusedExchangeExec
+import org.apache.spark.sql.functions.{col, lit, map, struct, sum}
+import org.apache.spark.sql.types.StringType
+import org.apache.spark.unsafe.types.UTF8String
+
+/**
+ * Every case runs with adaptive execution off and on. Most cases are rows of
one table: a fixture,
+ * a query and the exact declaration expected on each native scan, where an
empty declaration means
+ * the scan must be left whole. The class comment of ScanMapKeyPruning lists
which shapes are
+ * supported; the cases here follow that list.
+ */
+class ScanMapKeyPruningSuite extends VeloxWholeStageTransformerSuite {
+ override protected val resourcePath: String = "/tpch-data-parquet"
+ override protected val fileFormat: String = "parquet"
+
+ override protected def sparkConf: SparkConf = super.sparkConf
+ .set("spark.unsafe.exceptionOnMemoryLeak", "true")
+
+ private val pruningFlag = VeloxConfig.SCAN_MAP_KEY_PRUNING_ENABLED.key
+ private val noBroadcast = "spark.sql.autoBroadcastJoinThreshold" -> "-1"
+
+ //
---------------------------------------------------------------------------------------------
+ // Fixtures: each writes a parquet table and hands its path to the body.
+ //
---------------------------------------------------------------------------------------------
+
+ private type Fixture = (String => Unit) => Unit
+
+ private def table(rows: Long, columns: String*): Fixture = f =>
+ withTempPath {
+ path =>
+ spark.range(0, rows).selectExpr(("id" +: columns):
_*).write.parquet(path.getCanonicalPath)
+ f(path.getCanonicalPath)
+ }
+
+ /** m: MAP<STRING, STRUCT<s STRING, t BIGINT>> with keys a, b, c. */
+ private val mapTable = table(
+ 10000,
+ "map('a', named_struct('s', concat('v', cast(id % 7 as string)), 't', id),
" +
+ "'b', named_struct('s', 'x', 't', id * 2), " +
+ "'c', named_struct('s', 'y', 't', id + 1)) as m"
+ )
+
+ /** m: MAP<STRING, BIGINT>, a map whose values are plain numbers. */
+ private val scalarMapTable = table(1000, "map('a', id, 'b', id * 2, 'c', id
+ 1) as m")
+
+ /** m: MAP<BIGINT, STRUCT<s, t>>. */
+ private val longKeyTable = table(
+ 10000,
+ "map(1L, named_struct('s', concat('v', cast(id % 7 as string)), 't', id),
" +
+ "2L, named_struct('s', 'x', 't', id * 2)) as m")
+
+ /** c: STRUCT<m: MAP<STRING, BIGINT>, v: BIGINT>. */
+ private val nestedMapTable =
+ table(1000, "named_struct('m', map('a', id, 'b', id * 2), 'v', id) as c")
+
+ /** c: STRUCT<m: MAP<STRING, BIGINT>, s: STRUCT<n: MAP<STRING, BIGINT>, v:
BIGINT>>. */
+ private val siblingMapTable = table(
+ 1000,
+ "named_struct('m', map('a', id, 'b', id * 2), " +
+ "'s', named_struct('n', map('p', id), 'v', id)) as c")
+
+ /** m: MAP<STRING, MAP<STRING, BIGINT>>. */
+ private val mapOfMapTable =
+ table(1000, "map('a', map('x', id, 'y', id * 2), 'b', map('z', id)) as m")
+
+ /** m: MAP<STRING, ARRAY<BIGINT>>. */
+ private val arrayMapTable = table(500, "map('a', array(id, id + 1), 'b',
array(id * 2)) as m")
+
+ /** m as in mapTable plus n: MAP<STRING, BIGINT>. */
+ private val twoMapTable = table(
+ 10000,
+ "map('a', named_struct('s', concat('v', cast(id % 7 as string)), 't', id),
" +
+ "'b', named_struct('s', 'x', 't', id * 2)) as m",
+ "map('x', id, 'y', id * 2) as n"
+ )
+
+ // Built from code points rather than literals or unicode escapes, which
scalastyle rejects.
+ private val umlautUpper = 0xc4.toChar // A with diaeresis
+ private val umlautLower = 0xe4.toChar // a with diaeresis
+
+ /**
+ * c: STRUCT<[umlautUpper]rger: MAP<STRING, BIGINT>, v: BIGINT>: a field
name that only full
+ * Unicode lowercasing changes, since the ASCII fold leaves the umlaut alone.
+ */
+ private val umlautFieldTable =
+ table(1000, s"named_struct('${umlautUpper}rger', map('a', id, 'b', id *
2), 'v', id) as c")
+
+ /** m: MAP<STRING, BIGINT> whose first key is a single byte that is not
valid UTF-8. */
+ private val invalidUtf8KeyTable =
+ table(100, "map(cast(x'FF' as string), id, 'a', id * 2) as m")
+
+ /** d: MAP<DATE, BIGINT>, e: MAP<DECIMAL(3, 1), BIGINT>: key types without a
Velox subscript. */
+ private val oddKeyTable = table(
+ 1000,
+ "map(DATE '2020-01-01', id, DATE '2020-01-02', id * 2) as d",
+ "map(CAST(1.5 AS DECIMAL(3, 1)), id, CAST(2.5 AS DECIMAL(3, 1)), id * 2)
as e")
+
+ //
---------------------------------------------------------------------------------------------
+ // Plan inspection.
+ //
---------------------------------------------------------------------------------------------
+
+ /**
+ * The declaration of every native scan in the executed plan (query stages
and commands included),
+ * each rendered in Velox Subfield syntax. At least one native scan must
exist: a "nothing
+ * declared" assertion over a plan whose scan fell back to Spark would prove
nothing.
+ */
+ private def declaredPerScan(df: DataFrame): Seq[Map[String, Set[String]]] = {
+ val scans = getExecutedPlan(df).collect { case scan:
BasicScanExecTransformer => scan }
+ assert(scans.nonEmpty, s"expected a native
scan:\n${df.queryExecution.executedPlan}")
+ scans.map(_.requiredMapSubfields.map { case (c, paths) => c ->
paths.map(_.toString).toSet })
+ }
+
+ /** The declaration over all native scans, for single-scan plans. */
+ private def declared(df: DataFrame): Map[String, Set[String]] =
declaredPerScan(df).flatten.toMap
+
+ private def planHas(df: DataFrame, cls: Class[_ <: SparkPlan]): Boolean =
+ getExecutedPlan(df).exists(cls.isInstance)
+
+ /** Per-scan declarations compared regardless of scan order. */
+ private def multiset[T](xs: Seq[T]): Map[T, Int] = xs.groupBy(identity).map {
+ case (k, v) => k -> v.size
+ }
+
+ //
---------------------------------------------------------------------------------------------
+ // Table-driven cases.
+ //
---------------------------------------------------------------------------------------------
+
+ /**
+ * One case: 'sql' over 'fixture', run with the flag on and compared with
vanilla Spark, must
+ * leave exactly 'expected' on the native scans (one entry per scan,
order-free; an empty map is a
+ * scan left whole) and, if given, must contain an operator of type
'requires' in its plan.
+ */
+ private case class Case(
+ name: String,
+ fixture: Fixture,
+ sql: String => String,
+ expected: Seq[Map[String, Set[String]]],
+ conf: Seq[(String, String)] = Nil,
+ requires: Option[Class[_ <: SparkPlan]] = None,
+ noFallBack: Boolean = true,
+ compare: Boolean = true,
+ before: () => Unit = () => ())
+
+ private def m(paths: String*): Map[String, Set[String]] = Map("m" ->
paths.toSet)
+ private def c(paths: String*): Map[String, Set[String]] = Map("c" ->
paths.toSet)
+ private val whole: Map[String, Set[String]] = Map.empty
+
+ private val cases = Seq(
+ // ----- supported shapes: the map is dropped by a Project before anything
else sees it -----
+ Case(
+ "keyed filter with aggregates over keyed values",
+ mapTable,
+ p => s"SELECT count(*), sum(m['a'].t), max(m['a'].s) FROM parquet.`$p`
WHERE m['a'].s = 'v3'",
+ Seq(m("m[\"a\"].s", "m[\"a\"].t"))
+ ),
+ Case(
+ "filter only, map dropped above the filter",
+ mapTable,
+ p => s"SELECT id FROM parquet.`$p` WHERE m['a'].s = 'v3'",
+ Seq(m("m[\"a\"].s"))),
+ Case(
+ "count only",
+ mapTable,
+ p => s"SELECT count(*) FROM parquet.`$p` WHERE m['a'].s = 'v3'",
+ Seq(m("m[\"a\"].s"))),
+ Case(
+ "projected keyed value without aggregation",
+ mapTable,
+ p => s"SELECT id, m['a'].t FROM parquet.`$p` WHERE m['a'].s = 'v3'",
+ Seq(m("m[\"a\"].s", "m[\"a\"].t"))
+ ),
+ Case(
+ "several keys",
+ mapTable,
+ p => s"SELECT sum(m['a'].t + m['b'].t) FROM parquet.`$p` WHERE m['c'].s
= 'y'",
+ Seq(m("m[\"a\"].t", "m[\"b\"].t", "m[\"c\"].s"))),
+ Case(
+ "constant key expressions are evaluated",
+ mapTable,
+ p => s"SELECT sum(m[concat('a', '')].t) FROM parquet.`$p` WHERE
m[upper('a')].s = 'v3'",
+ Seq(m("m[\"a\"].t", "m[\"A\"].s"))
+ ),
+ Case(
+ "element_at with a constant key",
+ mapTable,
+ p =>
+ s"SELECT sum(element_at(m, 'a').t) FROM parquet.`$p` " +
+ s"WHERE element_at(m, 'a').s = 'v5'",
+ Seq(m("m[\"a\"].s", "m[\"a\"].t"))
+ ),
+ Case(
+ "inferred null check next to a keyed filter is dropped",
+ mapTable,
+ p =>
+ s"SELECT count(*), sum(m['a'].t) FROM parquet.`$p` " +
+ s"WHERE m IS NOT NULL AND m['a'].s = 'v3'",
+ Seq(m("m[\"a\"].s", "m[\"a\"].t"))
+ ),
+ Case(
+ "null check on a keyed path alone",
+ mapTable,
+ p => s"SELECT count(*) FROM parquet.`$p` WHERE m['a'] IS NULL",
+ Seq(m("m[\"a\"]"))),
+ Case(
+ "long keys",
+ longKeyTable,
+ p => s"SELECT count(*), sum(m[1].t) FROM parquet.`$p` WHERE m[1].s =
'v3'",
+ Seq(m("m[1].s", "m[1].t"))),
+ Case(
+ "scalar-valued map",
+ scalarMapTable,
+ p => s"SELECT sum(m['a']) FROM parquet.`$p` WHERE m['b'] > 10",
+ Seq(m("m[\"a\"]", "m[\"b\"]"))),
+ Case(
+ "nested map value used whole under a key",
+ mapOfMapTable,
+ p => s"SELECT sum(cardinality(m['a'])), sum(m['a']['y']) FROM
parquet.`$p`",
+ Seq(m("m[\"a\"]", "m[\"a\"][\"y\"]"))
+ ),
+ Case(
+ "map nested in a struct, through Spark's nested-column alias",
+ nestedMapTable,
+ p => s"SELECT sum(c.m['a']), sum(c.m['b']) FROM parquet.`$p`",
+ Seq(c("c.m[\"a\"]", "c.m[\"b\"]"))
+ ),
+ Case(
+ "sibling struct fields are declared alongside the keyed map",
+ siblingMapTable,
+ p => s"SELECT sum(c.m['a']), sum(c.s.v), sum(cardinality(c.s.n)) FROM
parquet.`$p`",
+ Seq(c("c.m[\"a\"]", "c.s.v", "c.s.n"))
+ ),
+ Case(
+ "case-insensitive field access renders the schema's spelling",
+ nestedMapTable,
+ p => s"SELECT sum(c.M['a']) FROM parquet.`$p`",
+ Seq(c("c.m[\"a\"]")),
+ conf = Seq("spark.sql.caseSensitive" -> "false")
+ ),
+ Case(
+ "collect limit at the root",
+ mapTable,
+ p => s"SELECT id FROM parquet.`$p` WHERE m['a'].s = 'v3' LIMIT 5",
+ Seq(m("m[\"a\"].s"))),
+ Case(
+ "take ordered and project",
+ mapTable,
+ p => s"SELECT id FROM parquet.`$p` WHERE m['a'].s = 'v3' ORDER BY id
LIMIT 5",
+ Seq(m("m[\"a\"].s"))),
+ Case(
+ "window partitioned by a keyed value",
+ scalarMapTable,
+ p => s"SELECT id, rank() OVER (PARTITION BY m['a'] ORDER BY id) AS r
FROM parquet.`$p`",
+ Seq(m("m[\"a\"]"))
+ ),
+ Case(
+ "generate over a keyed value, map still read above the generate",
+ arrayMapTable,
+ p => s"SELECT m['b'], explode(m['a']) FROM parquet.`$p`",
+ Seq(m("m[\"a\"]", "m[\"b\"]")),
+ requires = Some(classOf[GenerateExecTransformer])
+ ),
+ Case(
+ "generate over a keyed value ends the chain when nothing above needs the
map",
+ arrayMapTable,
+ p => s"SELECT explode(m['a']) FROM parquet.`$p`",
+ Seq(m("m[\"a\"]")),
+ requires = Some(classOf[GenerateExecTransformer])
+ ),
+ Case(
+ "non-ASCII struct field name folded like the schema, case-insensitive",
+ umlautFieldTable,
+ // Non-ASCII identifiers must be quoted in Spark SQL.
+ p =>
+ s"SELECT sum(c.`${umlautLower}rger`['a']) FROM parquet.`$p` " +
+ s"WHERE c.`${umlautUpper}RGER`['b'] > 10",
+ // Plain concatenation: scalastyle's parser does not accept \" inside
interpolated strings.
+ Seq(c("c." + umlautLower + "rger[\"a\"]", "c." + umlautLower +
"rger[\"b\"]")),
+ conf = Seq("spark.sql.caseSensitive" -> "false")
+ ),
+ Case(
Review Comment:
This case asserts the *declared* required subfields for a non-ASCII field
name under case-insensitive resolution, but it doesn’t validate that the native
reader actually applies the declaration correctly end-to-end (i.e., that the
declaration survives transport and matches Velox/Hive field resolution).
Consider adding an end-to-end assertion similar to the “native reader applies a
declaration” test, specifically for a non-ASCII folded struct field (and/or
nested struct field selection like `c.<field>["a"]`), to catch mismatches
between JVM normalization and native-side case folding.
##########
backends-velox/src/main/scala/org/apache/gluten/extension/ScanMapKeyPruning.scala:
##########
@@ -0,0 +1,420 @@
+/*
+ * 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
+
+import org.apache.gluten.config.VeloxConfig
+import org.apache.gluten.execution.{BasicScanExecTransformer,
FilterExecTransformerBase, GenerateExecTransformer, ProjectExecTransformerBase,
SubfieldElement, SubfieldPath}
+import org.apache.gluten.execution.SubfieldElement.{Field, LongKey, StringKey
=> StrKey}
+import org.apache.gluten.expression.ConverterUtils
+import org.apache.gluten.substrait.rel.LocalFilesNode.ReadFileFormat
+
+import org.apache.spark.sql.catalyst.expressions._
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.execution.{FilterExec, GenerateExec, ProjectExec,
SparkPlan}
+import org.apache.spark.sql.types._
+import org.apache.spark.unsafe.types.UTF8String
+
+import java.nio.charset.StandardCharsets
+
+import scala.collection.mutable
+import scala.util.control.NonFatal
+
+/**
+ * Physical plan rule that derives Velox required subfields for map columns
read by a native Parquet
+ * scan, so that the reader prunes map entries the query never accesses.
+ *
+ * Background. Parquet stores a map as a repeated `key_value` group: one key
column chunk and one
+ * value column chunk per row group, covering every entry of every row.
Evaluating `m['k1']` on the
+ * scan output therefore decodes all entries into a MapVector, although only
one key is needed.
+ * Velox's Hive connector lets a `HiveColumnHandle` declare required
subfields, written in Velox
+ * `Subfield` syntax such as `m["k1"].f1`; `makeScanSpec` turns the map keys
they name into a key
+ * filter on the map's key stream, so the SelectiveMapColumnReader
materializes only matching
+ * entries. Struct fields not covered by any required subfield are filled with
nulls.
+ *
+ * Velox treats the declared set as complete and returns null for anything
outside it, so the
+ * declared paths must cover every read of the column. The rule therefore only
handles one plan
+ * shape, in which that is easy to show:
+ * {{{
+ * Scan [id, m] filters: m["a"].s = 'v3'
+ * -> Filter* more predicates over keyed paths
+ * -> Project* keyed paths under aliases; m
may be forwarded
+ * -> Generate* m forwarded only if still
needed above
+ * -> a Project (or Generate) that no longer outputs m end of the chain
+ * }}}
+ * Once a Project drops the attribute (and every alias that still carries map
data of it), no
+ * operator above can reference it: in a resolved plan an attribute can only
be referenced if some
+ * child outputs it. Whatever sits above the chain is therefore irrelevant and
needs no analysis.
+ * Aliases of keyed paths that leave the chain (`m['a'].t AS x`) are fine as
well: their value
+ * depends only on entries the declared paths keep.
+ *
+ * Supported, that is, pruned:
+ * - filters on constant-key lookups, `m['a'].s = 'v3'`, `m[1].t > 0`,
`element_at(m, 'a')`,
+ * including the `isnotnull(m)` Spark infers next to them;
+ * - keyed values projected or aggregated, `SELECT m['a'].t`,
`sum(m['a'].t)`, also under a
+ * `count(*)`, a `LIMIT`, an `ORDER BY ... LIMIT`, a `UNION ALL` of such
projections, a write of
+ * such values, a UDF or script transformation applied to them, or a join
whose sides only read
+ * keyed values below the shuffle: in all of these Spark or Gluten places
the keyed access in a
+ * Project that drops the map before anything else sees it;
+ * - maps nested in structs, `c.m['a']`, including the `c.m AS _extract_m`
alias Spark's
+ * NestedColumnAliasing inserts; sibling struct fields are declared
alongside;
+ * - generators over a keyed value, `explode(m['a'])`: Gluten's
pre-projection computes the lookup
+ * below the Generate, and the Generate forwards the map only while an
operator above still
+ * needs it, so the chain may end at the Generate itself;
+ * - window partition keys, `OVER (PARTITION BY m['a'])`, which Spark's
analyzer extracts into a
+ * Project below the Window's shuffle;
+ * - null checks on keyed paths, `m['a'] IS NULL`.
+ *
+ * Not supported, that is, the map is read whole:
+ * - the map itself, or an alias still holding it, reaching an exchange, a
join, an aggregate, a
+ * union, a write, a UDTF, the fragment's output or any other operator
that is not a chain node;
+ * this includes `ORDER BY m['a']`, whose pre-projection must forward the
map through the
+ * shuffle because the Sort still outputs it;
+ * - whole-map uses anywhere on the chain, `size(m)`, `map_keys(m)`,
`explode(m)`, `m[key_col]`;
+ * - key types without a Velox subscript form, such as date or decimal, and
string keys whose
+ * bytes are not valid UTF-8.
+ *
+ * A chain ends at the first operator that is not a chain node, an exchange
included, so a scan's
+ * declaration depends only on the operators between it and that point.
Structurally identical
+ * exchange subtrees therefore always receive identical declarations, and
exchange reuse is
+ * unaffected.
+ *
+ * Example. For
+ * {{{
+ * SELECT sum(m['a'].t) FROM t WHERE m['a'].s = 'v3'
+ * }}}
+ * Spark places `m['a'].t AS _extract_t` in a Project below the aggregate, so
the chain is Scan ->
+ * Filter -> Project and the declared paths are `m["a"].s` and `m["a"].t`
(shown in Velox Subfield
+ * syntax, which SubfieldPath.toString renders; the wire format is structural,
one element per step,
+ * see RequiredSubfieldsExtension.proto).
+ */
+object ScanMapKeyPruning extends Rule[SparkPlan] {
+
+ // One element of a subfield path below the attribute, mirroring Velox
Subfield::PathElement:
+ // NestedField, StringSubscript, LongSubscript. `c.m["k1"].f1` is attribute
`c` followed by
+ // Field("m"), StrKey("k1"), Field("f1").
+ private type Elem = SubfieldElement
+
+ /** What the chain collectDeclarations collects for one scan attribute. */
+ private class CollectedPaths {
+ // Paths whose values are read; the map entries they name must be
materialized.
+ val valuePaths: mutable.LinkedHashSet[List[Elem]] =
+ mutable.LinkedHashSet[List[Elem]]()
+ // Paths referenced only by IsNull / IsNotNull; a null check reads the
null bitmap, which the
+ // reader produces for anything it materializes, so it needs no map
entries of its own.
+ val nullCheckPaths: mutable.LinkedHashSet[List[Elem]] =
+ mutable.LinkedHashSet[List[Elem]]()
+ // Set when some reference cannot be declared; the attribute is then left
alone.
+ var abandoned: Boolean = false
+ }
+
+ override def apply(plan: SparkPlan): SparkPlan = {
+ if (!VeloxConfig.get.scanMapKeyPruningEnabled) {
+ return plan
+ }
+ // Walk down from the root, remembering the chain nodes directly above the
current node,
+ // nearest first. Any other operator restarts the chain. Declarations are
kept per scan object,
+ // since two scans of one table are distinct objects.
+ val declarationsByScan =
+ new java.util.IdentityHashMap[SparkPlan, Map[String,
Seq[SubfieldPath]]]()
+ def collectDeclarations(node: SparkPlan, chain: List[SparkPlan]): Unit =
node match {
+ case scan: BasicScanExecTransformer if isPrunableParquetScan(scan) =>
+ val mapAttrs = scan.output.filter(a => containsMapType(a.dataType) &&
a.name.nonEmpty)
+ if (mapAttrs.nonEmpty) {
+ val filters = scan.filterExprs()
+ val subfields = mapAttrs.flatMap {
+ attr =>
+ pathsToDeclare(filters, chain, attr).map(paths =>
renderSubfieldPaths(attr, paths))
+ }.toMap
+ if (subfields.nonEmpty) {
+ declarationsByScan.put(scan, subfields)
+ }
+ }
+ case _ =>
+ val next = if (isChainOperator(node)) node :: chain else Nil
+ node.children.foreach(collectDeclarations(_, next))
+ }
+ collectDeclarations(plan, Nil)
+ if (declarationsByScan.isEmpty) {
+ return plan
+ }
+
+ plan.transformUp {
+ case scan: BasicScanExecTransformer if
declarationsByScan.containsKey(scan) =>
+ val subfields = declarationsByScan.get(scan)
+ logInfo(s"Applying scan map-key pruning: $subfields")
+ val newScan = scan.withRequiredMapSubfields(subfields)
+ newScan.copyTagsFrom(scan)
+ newScan
+ }
+ }
+
+ /** A native Parquet scan that supports the declaration and has none yet. */
+ private def isPrunableParquetScan(scan: BasicScanExecTransformer): Boolean =
+ scan.supportsMapKeyPruning && scan.requiredMapSubfields.isEmpty &&
+ scan.fileFormat == ReadFileFormat.ParquetReadFormat
+
+ /**
+ * The operators a chain may consist of: Filters and Projects, which use the
input rows through
+ * their predicate or project list, and Generate, which forwards its
requiredChildOutput and uses
+ * the input only through the generator. ColumnarPartialProjectExec and
+ * ColumnarPartialGenerateExec are deliberately not among them: part of
their work is in a field
+ * the rule cannot see.
+ */
+ private def isChainOperator(node: SparkPlan): Boolean = node match {
+ case _: FilterExecTransformerBase | _: FilterExec => true
+ case _: ProjectExecTransformerBase | _: ProjectExec => true
+ case _: GenerateExecTransformer | _: GenerateExec => true
+ case _ => false
+ }
+
+ /**
+ * Follows 'attr' from its scan (whose pushed-down 'filters' are visited
first) up the chain and
+ * returns the paths to declare, or None when the attribute must stay whole.
+ *
+ * 'liveIds' holds the ExprIds that still carry map data of the attribute,
with the path from the
+ * attribute to what the id stands for: the attribute itself with an empty
path, a bare alias (`m
+ * AS mm`) with an empty path, or a struct-field alias (`c.m AS _extract_m`,
which Spark's
+ * NestedColumnAliasing inserts) with path [m]. A reference to a live id
resolves by prepending
+ * that path: `_extract_m['a']` is `c.m["a"]`. The chain is complete once no
id is live.
+ */
+ private def pathsToDeclare(
+ filters: Seq[Expression],
+ chain: List[SparkPlan],
+ attr: Attribute): Option[Seq[List[Elem]]] = {
+ val paths = new CollectedPaths
+ var liveIds: Map[ExprId, List[Elem]] = Map(attr.exprId -> Nil)
+
+ def recordUsesIn(e: Expression, nullCheck: Boolean): Unit =
+ recordUses(liveIds, paths, e, nullCheck)
+
+ // A project list decides what stays live. An entry that resolves to a
path under a live id
+ // stays live under its own id when the path is key-free and still holds a
map: a forwarded
+ // `m#3`, a bare alias `m#3 AS mm#12`, or Spark's `c#3.m AS _extract_m#20`
(recording the path
+ // instead would declare all of it and defeat pruning below). A keyed
path, or a path with no
+ // map under it (`c#3.s.v AS _extract_v#8`), is recorded and needs no
further tracking: the
+ // entry then only carries data the declared path keeps. Anything else is
a computation, a use.
+ def applyProjectList(list: Seq[NamedExpression]): Unit = {
+ val next = mutable.LinkedHashMap[ExprId, List[Elem]]()
+ list.foreach {
+ entry =>
+ val expr = entry match {
+ case Alias(child, _) => child
+ case other => other
+ }
+ resolveSubfieldPath(expr) match {
+ case Some((root, path)) if liveIds.contains(root.exprId) =>
+ val full = liveIds(root.exprId) ++ path
+ if (!hasKeyLookup(full) && containsMapType(entry.dataType)) {
+ next(entry.exprId) = full
+ } else {
+ recordPath(paths, full, nullCheck = false)
+ }
+ case _ => recordUsesIn(expr, nullCheck = false)
+ }
+ }
+ liveIds = next.toMap
+ }
+
+ // A Generate uses its input only through the generator
(requiredChildOutput is a forward),
+ // and outputs requiredChildOutput plus the generated columns. Whatever it
does not output
+ // cannot be referenced above it, so it may also end the chain.
+ def applyGenerate(node: SparkPlan, generator: Expression): Unit = {
+ recordUsesIn(generator, nullCheck = false)
+ val forwarded = node.output.map(_.exprId).toSet
+ liveIds = liveIds.filter { case (id, _) => forwarded.contains(id) }
+ }
+
+ // Filters pushed into the scan (PushDownFilterToScan) are evaluated
natively on the scan
+ // output; they read the map like the Filter above them does.
+ filters.foreach(recordUsesIn(_, nullCheck = false))
+ chain.iterator.takeWhile(_ => liveIds.nonEmpty &&
!paths.abandoned).foreach {
+ case f: FilterExecTransformerBase => recordUsesIn(f.cond, nullCheck =
false)
+ case f: FilterExec => recordUsesIn(f.condition, nullCheck = false)
+ case p: ProjectExecTransformerBase => applyProjectList(p.list)
+ case p: ProjectExec => applyProjectList(p.projectList)
+ case g: GenerateExecTransformer => applyGenerate(g, g.generator)
+ case g: GenerateExec => applyGenerate(g, g.generator)
+ }
+ // Still live when the chain ends: the attribute (or an alias holding map
data of it) reaches
+ // an operator the rule does not pathsToDeclare, or the fragment's output.
Stay whole.
Review Comment:
Fix grammar in comment: `does not pathsToDeclare` is ungrammatical and reads
like an accidental symbol reference; rephrase to something like `does not
analyze` / `does not handle`.
--
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]