Copilot commented on code in PR #162: URL: https://github.com/apache/hbase-connectors/pull/162#discussion_r3966724046
########## spark4/hbase-spark4/src/main/scala/org/apache/hadoop/hbase/spark/datasources/HBasePartitionReader.scala: ########## @@ -0,0 +1,421 @@ +/* + * 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.hadoop.hbase.spark.datasources + +import java.util.ArrayList +import org.apache.hadoop.fs.Path +import org.apache.hadoop.hbase.{CellUtil, HBaseConfiguration, TableName} +import org.apache.hadoop.hbase.client.{Get, Query, Result, ResultScanner, Scan, Table} +import org.apache.hadoop.hbase.spark.{AndLogicExpression, DynamicLogicExpression, + EqualLogicExpression, GreaterThanLogicExpression, GreaterThanOrEqualLogicExpression, + HBaseConnectionCache, IsNullLogicExpression, LessThanLogicExpression, + LessThanOrEqualLogicExpression, Logging, OrLogicExpression, PassThroughLogicExpression, + PushdownMappedField, SmartConnection, SparkSQLPushDownFilter, StartsWithLogicExpression} +import org.apache.hadoop.hbase.util.Bytes +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.GenericInternalRow +import org.apache.spark.sql.catalyst.util.DateTimeUtils +import org.apache.spark.sql.types.Decimal +import org.apache.spark.sql.connector.read.PartitionReader +import org.apache.spark.sql.sources._ +import org.apache.spark.sql.types._ +import org.apache.spark.unsafe.types.UTF8String +import org.apache.yetus.audience.InterfaceAudience +import scala.collection.mutable.ListBuffer +import scala.jdk.CollectionConverters._ + +/** + * This is a new class in the spark4 module. Extends PartitionReader[InternalRow] for reading data from HBase regions. + * The actual execution: opens an HBase scanner on the partition's range, attaches the SparkSQLPushDownFilter, + * reads Result objects, and converts them to InternalRow. Implements next()/get()/close(). + * + * + * In the spark 3 DS V1 model, this logic was inside DefaultSource.buildScan() + * which returned an RDD[Row] with its own compute() method. + * + * Ranges are executed as Scan operations, whilst points are executed as batched Get operations. This mirrors the spark3 + * HBaseTableScanRDD.compute() behavior. + */ [email protected] +class HBasePartitionReader( + partition: HBaseInputPartition, + requiredSchema: StructType, + properties: Map[String, String], + catalog: HBaseTableCatalog, + pushedFilters: Array[Filter], + encoderClsName: String, + usePushDownColumnFilter: Boolean) + extends PartitionReader[InternalRow] + with Logging { + + private val conf = HBaseConfiguration.create() + private val configResources = properties.get(HBaseSparkConf.HBASE_CONFIG_LOCATION) + configResources.foreach(_.split(",").foreach(r => conf.addResource(new Path(r)))) + + private val connection: SmartConnection = HBaseConnectionCache.getConnection(conf) + private val tableName = s"${catalog.namespace}:${catalog.name}" + private val table: Table = connection.getTable(TableName.valueOf(tableName)) + + private val requiredFields = requiredSchema.fieldNames.map(catalog.sMap.getField(_)) + private val filterFields = extractFilterFields(pushedFilters) + private val scanFields = (requiredFields ++ filterFields).distinct.filterNot(_.isRowKey) + private val pushDownFilter: Option[SparkSQLPushDownFilter] = buildPushDownFilter() + + private val bulkGetSize = properties + .get(HBaseSparkConf.BULKGET_SIZE) + .map(_.toInt) + .getOrElse(HBaseSparkConf.DEFAULT_BULKGET_SIZE) + + private val blockCacheEnable = properties + .get(HBaseSparkConf.QUERY_CACHEBLOCKS) + .map(_.toBoolean) + .getOrElse(HBaseSparkConf.DEFAULT_QUERY_CACHEBLOCKS) + + private val scanners = new ListBuffer[ResultScanner]() + + private val resultIterator: Iterator[Result] = { + val scanIterators = partition.scanRanges.map { range => + val scanner = buildScanner(range) + scanners += scanner + scannerToIterator(scanner) + } + val getIterator = if (partition.points.nonEmpty) { + buildGets(partition.points) + } else { + Iterator.empty + } + scanIterators.foldLeft(Iterator.empty: Iterator[Result])(_ ++ _) ++ getIterator + } + + private var currentResult: Result = _ + + override def next(): Boolean = { + if (resultIterator.hasNext) { + currentResult = resultIterator.next() + true + } else { + false + } + } + + override def get(): InternalRow = { + val fields = requiredSchema.fieldNames.map(catalog.sMap.getField(_)) + val rowKey = currentResult.getRow + val keyFields = catalog.getRowKey Review Comment: `get()` recomputes `fields` for every row (`requiredSchema.fieldNames.map(...)`) even though `requiredFields` is already computed once at reader construction. This is in the hot path and can add noticeable overhead. Use the cached `requiredFields` (or precompute the `fields` sequence once) to avoid per-row allocations and catalog lookups. ########## spark4/hbase-spark4/src/main/scala/org/apache/hadoop/hbase/spark/datasources/HBaseBatch.scala: ########## @@ -0,0 +1,120 @@ +/* + * 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.hadoop.hbase.spark.datasources + +import org.apache.hadoop.fs.Path +import org.apache.hadoop.hbase.{HBaseConfiguration, TableName} +import org.apache.hadoop.hbase.spark.{HBaseConnectionCache, Logging} +import org.apache.spark.sql.connector.read.{Batch, InputPartition, PartitionReaderFactory} +import org.apache.spark.sql.sources._ +import org.apache.spark.sql.types.StructType +import org.apache.yetus.audience.InterfaceAudience + +/** + * This is a new class in the spark4 module. Implements Batch. + * Responsible for physical planning: splits the read into partitions. + * Calls RegionLocator.getStartEndKeys() to discover HBase regions, + * intersects them with the row key filter's scan ranges, and produces an array of InputPartition objects. + * + * In the spark 3 DS V1 model, this logic was inside HBaseTableScanRDD.getPartitions(). + * + * @param requiredSchema + * @param properties + * @param catalog + * @param rowKeyFilter + * @param pushedFilters + * @param encoderClsName + */ [email protected] +class HBaseBatch( + requiredSchema: StructType, + properties: Map[String, String], + catalog: HBaseTableCatalog, + rowKeyFilter: RowKeyFilter, + pushedFilters: Array[Filter], + encoderClsName: String) + extends Batch + with Logging { + + override def planInputPartitions(): Array[InputPartition] = { + val conf = HBaseConfiguration.create() + val configResources = properties.get(HBaseSparkConf.HBASE_CONFIG_LOCATION) + configResources.foreach(_.split(",").foreach(r => conf.addResource(new Path(r)))) + + val connection = HBaseConnectionCache.getConnection(conf) + try { + val tableName = s"${catalog.namespace}:${catalog.name}" + val regionLocator = connection.getRegionLocator(TableName.valueOf(tableName)) + try { + val keys = regionLocator.getStartEndKeys + val startKeys = keys.getFirst + val endKeys = keys.getSecond + + val regions = startKeys.zip(endKeys).zipWithIndex.map { case ((start, end), idx) => + HBaseRegion(idx, Some(start), Some(end)) Review Comment: HBase region start/end keys commonly use empty byte arrays to represent unbounded start/end (especially the last region's end key). Wrapping an empty end key in `Some(end)` risks planning scans with `withStopRow(Array.emptyByteArray)`, which can yield empty scans / missing data for the last region. Consider normalizing empty keys to `None` here (and similarly for empty starts if your `Range(region)` expects `None` for unbounded). ########## spark4/hbase-spark4/src/main/scala/org/apache/hadoop/hbase/spark/datasources/ScanRange.scala: ########## @@ -0,0 +1,225 @@ +/* + * 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.hadoop.hbase.spark.datasources + +import org.apache.hadoop.hbase.util.Bytes +import org.apache.yetus.audience.InterfaceAudience +import scala.collection.mutable.ListBuffer + +/** + * This is a new class in the spark4 module. Wraps ScanRange and RowKeyFilter. + * + * Extracted from DefaultSource.scala in spark3 where it was an inner class. Handles merging scan ranges + * from row key predicates (union/intersect). Same logic, just in its own file now for clarity. + * + * @param upperBound + * @param isUpperBoundEqualTo + * @param lowerBound + * @param isLowerBoundEqualTo + */ + [email protected] +class ScanRange( + var upperBound: Array[Byte], + var isUpperBoundEqualTo: Boolean, + var lowerBound: Array[Byte], + var isLowerBoundEqualTo: Boolean) + extends Serializable { + + def mergeIntersect(other: ScanRange): Unit = { + val upperBoundCompare = compareRange(upperBound, other.upperBound) + val lowerBoundCompare = compareRange(lowerBound, other.lowerBound) + + upperBound = if (upperBoundCompare < 0) upperBound else other.upperBound + lowerBound = if (lowerBoundCompare > 0) lowerBound else other.lowerBound + + isLowerBoundEqualTo = + if (lowerBoundCompare == 0) + isLowerBoundEqualTo && other.isLowerBoundEqualTo + else if (lowerBoundCompare < 0) other.isLowerBoundEqualTo + else isLowerBoundEqualTo + + isUpperBoundEqualTo = + if (upperBoundCompare == 0) + isUpperBoundEqualTo && other.isUpperBoundEqualTo + else if (upperBoundCompare < 0) isUpperBoundEqualTo + else other.isUpperBoundEqualTo + } + + def mergeUnion(other: ScanRange): Unit = { + val upperBoundCompare = compareRange(upperBound, other.upperBound) + val lowerBoundCompare = compareRange(lowerBound, other.lowerBound) + + upperBound = if (upperBoundCompare > 0) upperBound else other.upperBound + lowerBound = if (lowerBoundCompare < 0) lowerBound else other.lowerBound + + isLowerBoundEqualTo = + if (lowerBoundCompare == 0) + isLowerBoundEqualTo || other.isLowerBoundEqualTo + else if (lowerBoundCompare < 0) isLowerBoundEqualTo + else other.isLowerBoundEqualTo + + isUpperBoundEqualTo = + if (upperBoundCompare == 0) + isUpperBoundEqualTo || other.isUpperBoundEqualTo + else if (upperBoundCompare < 0) other.isUpperBoundEqualTo + else isUpperBoundEqualTo + } + + def getOverLapScanRange(other: ScanRange): ScanRange = { + var leftRange: ScanRange = null + var rightRange: ScanRange = null + + if (compareRange(lowerBound, other.lowerBound) < 0 || + compareRange(upperBound, other.upperBound) < 0) { + leftRange = this + rightRange = other + } else { + leftRange = other + rightRange = this + } + + if (hasOverlap(leftRange, rightRange)) { + val result = new ScanRange(upperBound, isUpperBoundEqualTo, lowerBound, isLowerBoundEqualTo) + result.mergeIntersect(other) + result + } else { + null + } + } + + def hasOverlap(left: ScanRange, right: ScanRange): Boolean = { + compareRange(left.upperBound, right.lowerBound) >= 0 + } + + def compareRange(left: Array[Byte], right: Array[Byte]): Int = { + if (left == null && right == null) 0 + else if (left == null && right != null) 1 + else if (left != null && right == null) -1 + else Bytes.compareTo(left, right) + } + + def containsPoint(point: Array[Byte]): Boolean = { + val lowerCompare = compareRange(point, lowerBound) + val upperCompare = compareRange(point, upperBound) + + ((isLowerBoundEqualTo && lowerCompare >= 0) || + (!isLowerBoundEqualTo && lowerCompare > 0)) && + ((isUpperBoundEqualTo && upperCompare <= 0) || + (!isUpperBoundEqualTo && upperCompare < 0)) + } + + override def toString: String = { + "ScanRange:(upperBound:" + Bytes.toString(upperBound) + + ",isUpperBoundEqualTo:" + isUpperBoundEqualTo + ",lowerBound:" + + Bytes.toString(lowerBound) + ",isLowerBoundEqualTo:" + isLowerBoundEqualTo + ")" + } +} + [email protected] +class RowKeyFilter( + currentPoint: Array[Byte] = null, + currentRange: ScanRange = new ScanRange(null, true, new Array[Byte](0), true), + var points: ListBuffer[Array[Byte]] = new ListBuffer[Array[Byte]](), + var ranges: ListBuffer[ScanRange] = new ListBuffer[ScanRange]()) + extends Serializable { + + if (currentRange != null) ranges += currentRange + if (currentPoint != null) points += currentPoint + + def mergeUnion(other: RowKeyFilter): RowKeyFilter = { + other.points.foreach(p => points += p) + + other.ranges.foreach { otherR => + var doesOverLap = false + ranges.foreach { r => + if (r.getOverLapScanRange(otherR) != null) { + r.mergeUnion(otherR) + doesOverLap = true + } + } + if (!doesOverLap) ranges += otherR + } + this + } + + def mergeIntersect(other: RowKeyFilter): RowKeyFilter = { + val survivingPoints = new ListBuffer[Array[Byte]]() + val didntSurviveFirstPassPoints = new ListBuffer[Array[Byte]]() + if (points == null || points.isEmpty) { + other.points.foreach(otherP => didntSurviveFirstPassPoints += otherP) + } else { + points.foreach { p => + if (other.points.isEmpty) { + didntSurviveFirstPassPoints += p + } else { + other.points.foreach { otherP => + if (Bytes.equals(p, otherP)) { + survivingPoints += p + } else { + didntSurviveFirstPassPoints += p + } + } Review Comment: RowKeyFilter.mergeIntersect has incorrect point intersection logic: when `other.points` has multiple items, a point `p` that matches one element can still be added to `didntSurviveFirstPassPoints` for the non-matching elements (and may be added multiple times). This can produce duplicate points and incorrect intersection results. Fix by determining survival per `p` (e.g., `if (other.points.exists(Bytes.equals(p, _))) survivingPoints += p else didntSurviveFirstPassPoints += p`) or by using a set-based intersection. ########## spark4/hbase-spark4/src/main/scala/org/apache/hadoop/hbase/spark/datasources/HBasePartitionReader.scala: ########## @@ -0,0 +1,421 @@ +/* + * 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.hadoop.hbase.spark.datasources + +import java.util.ArrayList +import org.apache.hadoop.fs.Path +import org.apache.hadoop.hbase.{CellUtil, HBaseConfiguration, TableName} +import org.apache.hadoop.hbase.client.{Get, Query, Result, ResultScanner, Scan, Table} +import org.apache.hadoop.hbase.spark.{AndLogicExpression, DynamicLogicExpression, + EqualLogicExpression, GreaterThanLogicExpression, GreaterThanOrEqualLogicExpression, + HBaseConnectionCache, IsNullLogicExpression, LessThanLogicExpression, + LessThanOrEqualLogicExpression, Logging, OrLogicExpression, PassThroughLogicExpression, + PushdownMappedField, SmartConnection, SparkSQLPushDownFilter, StartsWithLogicExpression} +import org.apache.hadoop.hbase.util.Bytes +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.GenericInternalRow +import org.apache.spark.sql.catalyst.util.DateTimeUtils +import org.apache.spark.sql.types.Decimal +import org.apache.spark.sql.connector.read.PartitionReader +import org.apache.spark.sql.sources._ +import org.apache.spark.sql.types._ +import org.apache.spark.unsafe.types.UTF8String +import org.apache.yetus.audience.InterfaceAudience +import scala.collection.mutable.ListBuffer +import scala.jdk.CollectionConverters._ + +/** + * This is a new class in the spark4 module. Extends PartitionReader[InternalRow] for reading data from HBase regions. + * The actual execution: opens an HBase scanner on the partition's range, attaches the SparkSQLPushDownFilter, + * reads Result objects, and converts them to InternalRow. Implements next()/get()/close(). + * + * + * In the spark 3 DS V1 model, this logic was inside DefaultSource.buildScan() + * which returned an RDD[Row] with its own compute() method. + * + * Ranges are executed as Scan operations, whilst points are executed as batched Get operations. This mirrors the spark3 + * HBaseTableScanRDD.compute() behavior. + */ [email protected] +class HBasePartitionReader( + partition: HBaseInputPartition, + requiredSchema: StructType, + properties: Map[String, String], + catalog: HBaseTableCatalog, + pushedFilters: Array[Filter], + encoderClsName: String, + usePushDownColumnFilter: Boolean) + extends PartitionReader[InternalRow] + with Logging { + + private val conf = HBaseConfiguration.create() + private val configResources = properties.get(HBaseSparkConf.HBASE_CONFIG_LOCATION) + configResources.foreach(_.split(",").foreach(r => conf.addResource(new Path(r)))) + + private val connection: SmartConnection = HBaseConnectionCache.getConnection(conf) + private val tableName = s"${catalog.namespace}:${catalog.name}" + private val table: Table = connection.getTable(TableName.valueOf(tableName)) + + private val requiredFields = requiredSchema.fieldNames.map(catalog.sMap.getField(_)) + private val filterFields = extractFilterFields(pushedFilters) + private val scanFields = (requiredFields ++ filterFields).distinct.filterNot(_.isRowKey) + private val pushDownFilter: Option[SparkSQLPushDownFilter] = buildPushDownFilter() + + private val bulkGetSize = properties + .get(HBaseSparkConf.BULKGET_SIZE) + .map(_.toInt) + .getOrElse(HBaseSparkConf.DEFAULT_BULKGET_SIZE) + + private val blockCacheEnable = properties + .get(HBaseSparkConf.QUERY_CACHEBLOCKS) + .map(_.toBoolean) + .getOrElse(HBaseSparkConf.DEFAULT_QUERY_CACHEBLOCKS) + + private val scanners = new ListBuffer[ResultScanner]() + + private val resultIterator: Iterator[Result] = { + val scanIterators = partition.scanRanges.map { range => + val scanner = buildScanner(range) + scanners += scanner + scannerToIterator(scanner) + } + val getIterator = if (partition.points.nonEmpty) { + buildGets(partition.points) + } else { + Iterator.empty + } + scanIterators.foldLeft(Iterator.empty: Iterator[Result])(_ ++ _) ++ getIterator + } + + private var currentResult: Result = _ + + override def next(): Boolean = { + if (resultIterator.hasNext) { + currentResult = resultIterator.next() + true + } else { + false + } + } + + override def get(): InternalRow = { + val fields = requiredSchema.fieldNames.map(catalog.sMap.getField(_)) + val rowKey = currentResult.getRow + val keyFields = catalog.getRowKey + + val keyValues = parseRowKey(rowKey, keyFields) + val values = new Array[Any](fields.length) + + fields.zipWithIndex.foreach { case (field, idx) => + if (field.isRowKey) { + values(idx) = convertToInternalRow(keyValues.get(field).orNull, field.dt) + } else { + val cell = currentResult.getColumnLatestCell( + Bytes.toBytes(field.cf), Bytes.toBytes(field.col)) + if (cell == null || cell.getValueLength == 0) { + values(idx) = null + } else { + val v = CellUtil.cloneValue(cell) + val scalaValue = field.dt match { + case BinaryType => v + case _ => Utils.hbaseFieldToScalaType(field, v, 0, v.length) + } + values(idx) = convertToInternalRow(scalaValue, field.dt) + } + } + } + new GenericInternalRow(values) + } + + override def close(): Unit = { + scanners.foreach(s => if (s != null) s.close()) + if (table != null) table.close() + if (connection != null) connection.close() + } + + private def setStopRow(scan: Scan, bound: Bound): Scan = { + if (bound.inc) { + val incremented = Utils.incrementByteArray(bound.b) + if (incremented != null) scan.withStopRow(incremented) + else scan + } else { + scan.withStopRow(bound.b) + } + } + + private def buildScanner(range: Range): ResultScanner = { + val scan = (range.lower, range.upper) match { + case (Some(Bound(a, _)), Some(upper)) => + setStopRow(new Scan().withStartRow(a), upper) + case (None, Some(upper)) => + setStopRow(new Scan(), upper) + case (Some(Bound(a, _)), None) => + new Scan().withStartRow(a) + case (None, None) => + new Scan() + } + + scan.setCacheBlocks(blockCacheEnable) + properties.get(HBaseSparkConf.QUERY_CACHEDROWS).map(_.toInt).foreach { rows => + if (rows > 0) scan.setCaching(rows) + } + properties.get(HBaseSparkConf.QUERY_BATCHSIZE).map(_.toInt).foreach { batch => + if (batch > 0) scan.setBatch(batch) + } + handleTimeSemantics(scan) + + scanFields.foreach { f => + scan.addColumn(f.cfBytes, f.colBytes) + } + pushDownFilter.foreach(scan.setFilter(_)) + + table.getScanner(scan) + } + + private def buildGets(points: Seq[Array[Byte]]): Iterator[Result] = { + points.grouped(bulkGetSize).flatMap { batch => + val gets = new ArrayList[Get](batch.size) + batch.foreach { point => + val g = new Get(point) + handleTimeSemantics(g) + scanFields.foreach { f => + g.addColumn(f.cfBytes, f.colBytes) + } + pushDownFilter.foreach(g.setFilter(_)) + gets.add(g) + } + table.get(gets).toSeq.iterator.filter(r => r != null && !r.isEmpty) + } + } + + private def scannerToIterator(scanner: ResultScanner): Iterator[Result] = { + new Iterator[Result] { + var cur: Option[Result] = None + override def hasNext: Boolean = { + if (cur.isEmpty) { + val r = scanner.next() + if (r != null) cur = Some(r) + } + cur.isDefined + } + override def next(): Result = { + hasNext + val ret = cur.get + cur = None + ret + } + } + } + + private def handleTimeSemantics(query: Query): Unit = { + val timestamp = properties.get(HBaseSparkConf.TIMESTAMP).map(_.toLong) + val minTs = properties.get(HBaseSparkConf.TIMERANGE_START).map(_.toLong) + val maxTs = properties.get(HBaseSparkConf.TIMERANGE_END).map(_.toLong) + (query, timestamp, minTs, maxTs) match { + case (q: Scan, Some(ts), None, None) => q.setTimestamp(ts) + case (q: Get, Some(ts), None, None) => q.setTimestamp(ts) + case (q: Scan, None, Some(min), Some(max)) => q.setTimeRange(min, max) + case (q: Get, None, Some(min), Some(max)) => q.setTimeRange(min, max) + case (_, None, None, None) => + case _ => + throw new IllegalArgumentException( + "Invalid combination of timestamp/time range provided.") + } + val maxVersions = properties.get(HBaseSparkConf.MAX_VERSIONS).map(_.toInt) + maxVersions.foreach { mv => + query match { + case q: Scan => q.readVersions(mv) + case q: Get => q.readVersions(mv) + case _ => + } + } + } + + private def buildPushDownFilter(): Option[SparkSQLPushDownFilter] = { + if (!usePushDownColumnFilter || pushedFilters.isEmpty) return None + val valueArray = buildValueArray() + val dynamicLogicExpression = buildDynamicLogicExpression() + if (dynamicLogicExpression == null) return None + + val allFilterFields = (requiredFields ++ filterFields).distinct + val columnMappings = allFilterFields.map { field => + new PushdownMappedField { + override def colName(): String = field.colName + override def cfBytes(): Array[Byte] = field.cfBytes + override def colBytes(): Array[Byte] = field.colBytes + } + } + Some(new SparkSQLPushDownFilter( + dynamicLogicExpression, + valueArray, + columnMappings.toList.asJava, + encoderClsName)) + } + + private def convertToInternalRow(value: Any, dataType: DataType): Any = { + if (value == null) return null + dataType match { + case StringType => UTF8String.fromString(value.asInstanceOf[String]) + case DateType => + val d = value.asInstanceOf[java.sql.Date] + DateTimeUtils.fromJavaDate(d) + case TimestampType => + val t = value.asInstanceOf[java.sql.Timestamp] + DateTimeUtils.fromJavaTimestamp(t) + case dt: DecimalType => + Decimal(value.asInstanceOf[java.math.BigDecimal], dt.precision, dt.scale) + case _ => value + } + } + + private def parseRowKey(row: Array[Byte], keyFields: Seq[Field]): Map[Field, Any] = { + keyFields + .foldLeft((0, Seq[(Field, Any)]())) { (state, field) => + val idx = state._1 + val parsed = state._2 + if (field.length != -1) { + val value = Utils.hbaseFieldToScalaType(field, row, idx, field.length) + (idx + field.length, parsed :+ (field, value)) + } else { + field.dt match { + case StringType => + val pos = row.indexOf(HBaseTableCatalog.delimiter, idx) + if (pos == -1 || pos > row.length) { + val value = Utils.hbaseFieldToScalaType(field, row, idx, row.length - idx) + (row.length + 1, parsed :+ (field, value)) + } else { + val value = Utils.hbaseFieldToScalaType(field, row, idx, pos - idx) + (pos, parsed :+ (field, value)) + } + case _ => + ( + row.length + 1, + parsed :+ (field, Utils.hbaseFieldToScalaType(field, row, idx, row.length - idx))) + } + } + } + ._2 + .toMap + } + + private def extractFilterFields(filters: Array[Filter]): Array[Field] = { + val fields = new ListBuffer[Field]() + def extract(f: Filter): Unit = f match { + case EqualTo(attr, _) => catalog.sMap.map.get(attr).foreach(fields += _) + case LessThan(attr, _) => catalog.sMap.map.get(attr).foreach(fields += _) + case GreaterThan(attr, _) => catalog.sMap.map.get(attr).foreach(fields += _) + case LessThanOrEqual(attr, _) => catalog.sMap.map.get(attr).foreach(fields += _) + case GreaterThanOrEqual(attr, _) => catalog.sMap.map.get(attr).foreach(fields += _) + case StringStartsWith(attr, _) => catalog.sMap.map.get(attr).foreach(fields += _) + case IsNull(attr) => catalog.sMap.map.get(attr).foreach(fields += _) + case IsNotNull(attr) => catalog.sMap.map.get(attr).foreach(fields += _) + case Or(left, right) => extract(left); extract(right) + case And(left, right) => extract(left); extract(right) + case _ => + } + filters.foreach(extract) + fields.toArray + } + + private def buildValueArray(): Array[Array[Byte]] = { + val values = new ListBuffer[Array[Byte]]() + pushedFilters.foreach(f => collectFilterValues(values, f)) + values.toArray + } + + private def collectFilterValues(values: ListBuffer[Array[Byte]], filter: Filter): Unit = { + val encoder = JavaBytesEncoder.create(encoderClsName) + filter match { Review Comment: collectFilterValues creates a new encoder instance for each recursive call, which is avoidable work when many filters are pushed down. Create the encoder once per reader (or per buildValueArray call) and reuse it throughout the recursion. ########## spark4/hbase-spark4/src/main/scala/org/apache/hadoop/hbase/spark/datasources/HBasePartitionReader.scala: ########## @@ -0,0 +1,421 @@ +/* + * 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.hadoop.hbase.spark.datasources + +import java.util.ArrayList +import org.apache.hadoop.fs.Path +import org.apache.hadoop.hbase.{CellUtil, HBaseConfiguration, TableName} +import org.apache.hadoop.hbase.client.{Get, Query, Result, ResultScanner, Scan, Table} +import org.apache.hadoop.hbase.spark.{AndLogicExpression, DynamicLogicExpression, + EqualLogicExpression, GreaterThanLogicExpression, GreaterThanOrEqualLogicExpression, + HBaseConnectionCache, IsNullLogicExpression, LessThanLogicExpression, + LessThanOrEqualLogicExpression, Logging, OrLogicExpression, PassThroughLogicExpression, + PushdownMappedField, SmartConnection, SparkSQLPushDownFilter, StartsWithLogicExpression} +import org.apache.hadoop.hbase.util.Bytes +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.GenericInternalRow +import org.apache.spark.sql.catalyst.util.DateTimeUtils +import org.apache.spark.sql.types.Decimal +import org.apache.spark.sql.connector.read.PartitionReader +import org.apache.spark.sql.sources._ +import org.apache.spark.sql.types._ +import org.apache.spark.unsafe.types.UTF8String +import org.apache.yetus.audience.InterfaceAudience +import scala.collection.mutable.ListBuffer +import scala.jdk.CollectionConverters._ + +/** + * This is a new class in the spark4 module. Extends PartitionReader[InternalRow] for reading data from HBase regions. + * The actual execution: opens an HBase scanner on the partition's range, attaches the SparkSQLPushDownFilter, + * reads Result objects, and converts them to InternalRow. Implements next()/get()/close(). + * + * + * In the spark 3 DS V1 model, this logic was inside DefaultSource.buildScan() + * which returned an RDD[Row] with its own compute() method. + * + * Ranges are executed as Scan operations, whilst points are executed as batched Get operations. This mirrors the spark3 + * HBaseTableScanRDD.compute() behavior. + */ [email protected] +class HBasePartitionReader( + partition: HBaseInputPartition, + requiredSchema: StructType, + properties: Map[String, String], + catalog: HBaseTableCatalog, + pushedFilters: Array[Filter], + encoderClsName: String, + usePushDownColumnFilter: Boolean) + extends PartitionReader[InternalRow] + with Logging { + + private val conf = HBaseConfiguration.create() + private val configResources = properties.get(HBaseSparkConf.HBASE_CONFIG_LOCATION) + configResources.foreach(_.split(",").foreach(r => conf.addResource(new Path(r)))) + + private val connection: SmartConnection = HBaseConnectionCache.getConnection(conf) + private val tableName = s"${catalog.namespace}:${catalog.name}" + private val table: Table = connection.getTable(TableName.valueOf(tableName)) + + private val requiredFields = requiredSchema.fieldNames.map(catalog.sMap.getField(_)) + private val filterFields = extractFilterFields(pushedFilters) + private val scanFields = (requiredFields ++ filterFields).distinct.filterNot(_.isRowKey) + private val pushDownFilter: Option[SparkSQLPushDownFilter] = buildPushDownFilter() + + private val bulkGetSize = properties + .get(HBaseSparkConf.BULKGET_SIZE) + .map(_.toInt) + .getOrElse(HBaseSparkConf.DEFAULT_BULKGET_SIZE) + + private val blockCacheEnable = properties + .get(HBaseSparkConf.QUERY_CACHEBLOCKS) + .map(_.toBoolean) + .getOrElse(HBaseSparkConf.DEFAULT_QUERY_CACHEBLOCKS) + + private val scanners = new ListBuffer[ResultScanner]() + + private val resultIterator: Iterator[Result] = { + val scanIterators = partition.scanRanges.map { range => + val scanner = buildScanner(range) + scanners += scanner + scannerToIterator(scanner) + } + val getIterator = if (partition.points.nonEmpty) { + buildGets(partition.points) + } else { + Iterator.empty + } + scanIterators.foldLeft(Iterator.empty: Iterator[Result])(_ ++ _) ++ getIterator + } + + private var currentResult: Result = _ + + override def next(): Boolean = { + if (resultIterator.hasNext) { + currentResult = resultIterator.next() + true + } else { + false + } + } + + override def get(): InternalRow = { + val fields = requiredSchema.fieldNames.map(catalog.sMap.getField(_)) + val rowKey = currentResult.getRow + val keyFields = catalog.getRowKey + + val keyValues = parseRowKey(rowKey, keyFields) + val values = new Array[Any](fields.length) + + fields.zipWithIndex.foreach { case (field, idx) => + if (field.isRowKey) { + values(idx) = convertToInternalRow(keyValues.get(field).orNull, field.dt) + } else { + val cell = currentResult.getColumnLatestCell( + Bytes.toBytes(field.cf), Bytes.toBytes(field.col)) + if (cell == null || cell.getValueLength == 0) { + values(idx) = null + } else { + val v = CellUtil.cloneValue(cell) + val scalaValue = field.dt match { + case BinaryType => v + case _ => Utils.hbaseFieldToScalaType(field, v, 0, v.length) + } + values(idx) = convertToInternalRow(scalaValue, field.dt) + } + } + } + new GenericInternalRow(values) + } + + override def close(): Unit = { + scanners.foreach(s => if (s != null) s.close()) + if (table != null) table.close() + if (connection != null) connection.close() + } + + private def setStopRow(scan: Scan, bound: Bound): Scan = { + if (bound.inc) { + val incremented = Utils.incrementByteArray(bound.b) + if (incremented != null) scan.withStopRow(incremented) + else scan + } else { + scan.withStopRow(bound.b) + } + } + + private def buildScanner(range: Range): ResultScanner = { + val scan = (range.lower, range.upper) match { + case (Some(Bound(a, _)), Some(upper)) => + setStopRow(new Scan().withStartRow(a), upper) + case (None, Some(upper)) => + setStopRow(new Scan(), upper) + case (Some(Bound(a, _)), None) => + new Scan().withStartRow(a) + case (None, None) => + new Scan() + } + + scan.setCacheBlocks(blockCacheEnable) + properties.get(HBaseSparkConf.QUERY_CACHEDROWS).map(_.toInt).foreach { rows => + if (rows > 0) scan.setCaching(rows) + } + properties.get(HBaseSparkConf.QUERY_BATCHSIZE).map(_.toInt).foreach { batch => + if (batch > 0) scan.setBatch(batch) + } + handleTimeSemantics(scan) + + scanFields.foreach { f => + scan.addColumn(f.cfBytes, f.colBytes) + } + pushDownFilter.foreach(scan.setFilter(_)) + + table.getScanner(scan) + } + + private def buildGets(points: Seq[Array[Byte]]): Iterator[Result] = { + points.grouped(bulkGetSize).flatMap { batch => + val gets = new ArrayList[Get](batch.size) + batch.foreach { point => + val g = new Get(point) + handleTimeSemantics(g) + scanFields.foreach { f => + g.addColumn(f.cfBytes, f.colBytes) + } + pushDownFilter.foreach(g.setFilter(_)) + gets.add(g) + } + table.get(gets).toSeq.iterator.filter(r => r != null && !r.isEmpty) + } + } + + private def scannerToIterator(scanner: ResultScanner): Iterator[Result] = { + new Iterator[Result] { + var cur: Option[Result] = None + override def hasNext: Boolean = { + if (cur.isEmpty) { + val r = scanner.next() + if (r != null) cur = Some(r) + } + cur.isDefined + } + override def next(): Result = { + hasNext + val ret = cur.get + cur = None + ret + } + } + } + + private def handleTimeSemantics(query: Query): Unit = { + val timestamp = properties.get(HBaseSparkConf.TIMESTAMP).map(_.toLong) + val minTs = properties.get(HBaseSparkConf.TIMERANGE_START).map(_.toLong) + val maxTs = properties.get(HBaseSparkConf.TIMERANGE_END).map(_.toLong) + (query, timestamp, minTs, maxTs) match { + case (q: Scan, Some(ts), None, None) => q.setTimestamp(ts) + case (q: Get, Some(ts), None, None) => q.setTimestamp(ts) + case (q: Scan, None, Some(min), Some(max)) => q.setTimeRange(min, max) + case (q: Get, None, Some(min), Some(max)) => q.setTimeRange(min, max) + case (_, None, None, None) => + case _ => + throw new IllegalArgumentException( + "Invalid combination of timestamp/time range provided.") + } + val maxVersions = properties.get(HBaseSparkConf.MAX_VERSIONS).map(_.toInt) + maxVersions.foreach { mv => + query match { + case q: Scan => q.readVersions(mv) + case q: Get => q.readVersions(mv) + case _ => + } + } + } + + private def buildPushDownFilter(): Option[SparkSQLPushDownFilter] = { + if (!usePushDownColumnFilter || pushedFilters.isEmpty) return None + val valueArray = buildValueArray() + val dynamicLogicExpression = buildDynamicLogicExpression() + if (dynamicLogicExpression == null) return None + + val allFilterFields = (requiredFields ++ filterFields).distinct + val columnMappings = allFilterFields.map { field => + new PushdownMappedField { + override def colName(): String = field.colName + override def cfBytes(): Array[Byte] = field.cfBytes + override def colBytes(): Array[Byte] = field.colBytes + } + } + Some(new SparkSQLPushDownFilter( + dynamicLogicExpression, + valueArray, + columnMappings.toList.asJava, + encoderClsName)) + } + + private def convertToInternalRow(value: Any, dataType: DataType): Any = { + if (value == null) return null + dataType match { + case StringType => UTF8String.fromString(value.asInstanceOf[String]) + case DateType => + val d = value.asInstanceOf[java.sql.Date] + DateTimeUtils.fromJavaDate(d) + case TimestampType => + val t = value.asInstanceOf[java.sql.Timestamp] + DateTimeUtils.fromJavaTimestamp(t) + case dt: DecimalType => + Decimal(value.asInstanceOf[java.math.BigDecimal], dt.precision, dt.scale) + case _ => value + } + } + + private def parseRowKey(row: Array[Byte], keyFields: Seq[Field]): Map[Field, Any] = { + keyFields + .foldLeft((0, Seq[(Field, Any)]())) { (state, field) => + val idx = state._1 + val parsed = state._2 + if (field.length != -1) { + val value = Utils.hbaseFieldToScalaType(field, row, idx, field.length) + (idx + field.length, parsed :+ (field, value)) + } else { + field.dt match { + case StringType => + val pos = row.indexOf(HBaseTableCatalog.delimiter, idx) + if (pos == -1 || pos > row.length) { + val value = Utils.hbaseFieldToScalaType(field, row, idx, row.length - idx) + (row.length + 1, parsed :+ (field, value)) + } else { + val value = Utils.hbaseFieldToScalaType(field, row, idx, pos - idx) + (pos, parsed :+ (field, value)) Review Comment: parseRowKey does not advance past the delimiter after parsing a variable-length `StringType` component: it returns `(pos, ...)` instead of `(pos + 1, ...)`. This will cause the next field to start at the delimiter byte, yielding incorrect decoding (empty strings or corrupted fixed-length fields). Adjust the next index to `pos + 1` when a delimiter is found. ########## spark4/hbase-spark4/src/test/scala/org/apache/hadoop/hbase/spark/datasources/HBaseTableProviderSuite.scala: ########## @@ -0,0 +1,193 @@ +/* + * 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.hadoop.hbase.spark.datasources + +import java.io.{File, FileOutputStream} +import org.apache.hadoop.hbase.{HBaseTestingUtility, TableName} +import org.apache.hadoop.hbase.client.{ConnectionFactory, Put} +import org.apache.hadoop.hbase.spark.Logging +import org.apache.hadoop.hbase.util.Bytes +import org.apache.spark.sql.SparkSession +import org.scalatest.BeforeAndAfterAll +import org.scalatest.funsuite.AnyFunSuite + +class HBaseTableProviderSuite extends AnyFunSuite with BeforeAndAfterAll with Logging { + + val TEST_UTIL = new HBaseTestingUtility + var spark: SparkSession = _ + var configFile: File = _ + + val tableName = "test_provider" + val columnFamily = "cf" + val numRows = 20 + + val catalog: String = s"""{ + |"table":{"namespace":"default", "name":"$tableName"}, + |"rowkey":"key", + |"columns":{ + |"key":{"cf":"rowkey", "col":"key", "type":"string"}, + |"name":{"cf":"$columnFamily", "col":"name", "type":"string"}, + |"age":{"cf":"$columnFamily", "col":"age", "type":"string"}, + |"salary":{"cf":"$columnFamily", "col":"salary", "type":"string"} + |} + |}""".stripMargin + + override def beforeAll(): Unit = { + TEST_UTIL.startMiniCluster() + logInfo(" - minicluster started") + + TEST_UTIL.createTable(TableName.valueOf(tableName), Bytes.toBytes(columnFamily)) + logInfo(s" - created table $tableName") + + populateTestData() + + val tmpDir = new File("target", "test-tmp") + tmpDir.mkdirs() + configFile = File.createTempFile("hbase-site", ".xml", tmpDir) + configFile.deleteOnExit() + val out = new FileOutputStream(configFile) + TEST_UTIL.getConfiguration.writeXml(out) + out.close() Review Comment: The FileOutputStream should be closed in a `finally`/resource-safety construct so it is reliably closed if `writeXml` throws (prevents leaking file handles in CI). Use a try/finally (or Scala `Using.resource`) around the stream. ########## spark4/hbase-spark4/src/main/scala/org/apache/hadoop/hbase/spark/datasources/ScanRange.scala: ########## @@ -0,0 +1,225 @@ +/* + * 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.hadoop.hbase.spark.datasources + +import org.apache.hadoop.hbase.util.Bytes +import org.apache.yetus.audience.InterfaceAudience +import scala.collection.mutable.ListBuffer + +/** + * This is a new class in the spark4 module. Wraps ScanRange and RowKeyFilter. + * + * Extracted from DefaultSource.scala in spark3 where it was an inner class. Handles merging scan ranges + * from row key predicates (union/intersect). Same logic, just in its own file now for clarity. + * + * @param upperBound + * @param isUpperBoundEqualTo + * @param lowerBound + * @param isLowerBoundEqualTo + */ + [email protected] +class ScanRange( + var upperBound: Array[Byte], + var isUpperBoundEqualTo: Boolean, + var lowerBound: Array[Byte], + var isLowerBoundEqualTo: Boolean) + extends Serializable { + + def mergeIntersect(other: ScanRange): Unit = { + val upperBoundCompare = compareRange(upperBound, other.upperBound) + val lowerBoundCompare = compareRange(lowerBound, other.lowerBound) + + upperBound = if (upperBoundCompare < 0) upperBound else other.upperBound + lowerBound = if (lowerBoundCompare > 0) lowerBound else other.lowerBound + + isLowerBoundEqualTo = + if (lowerBoundCompare == 0) + isLowerBoundEqualTo && other.isLowerBoundEqualTo + else if (lowerBoundCompare < 0) other.isLowerBoundEqualTo + else isLowerBoundEqualTo + + isUpperBoundEqualTo = + if (upperBoundCompare == 0) + isUpperBoundEqualTo && other.isUpperBoundEqualTo + else if (upperBoundCompare < 0) isUpperBoundEqualTo + else other.isUpperBoundEqualTo + } + + def mergeUnion(other: ScanRange): Unit = { + val upperBoundCompare = compareRange(upperBound, other.upperBound) + val lowerBoundCompare = compareRange(lowerBound, other.lowerBound) + + upperBound = if (upperBoundCompare > 0) upperBound else other.upperBound + lowerBound = if (lowerBoundCompare < 0) lowerBound else other.lowerBound + + isLowerBoundEqualTo = + if (lowerBoundCompare == 0) + isLowerBoundEqualTo || other.isLowerBoundEqualTo + else if (lowerBoundCompare < 0) isLowerBoundEqualTo + else other.isLowerBoundEqualTo + + isUpperBoundEqualTo = + if (upperBoundCompare == 0) + isUpperBoundEqualTo || other.isUpperBoundEqualTo + else if (upperBoundCompare < 0) other.isUpperBoundEqualTo + else isUpperBoundEqualTo + } + + def getOverLapScanRange(other: ScanRange): ScanRange = { + var leftRange: ScanRange = null + var rightRange: ScanRange = null + + if (compareRange(lowerBound, other.lowerBound) < 0 || + compareRange(upperBound, other.upperBound) < 0) { + leftRange = this + rightRange = other + } else { + leftRange = other + rightRange = this + } + + if (hasOverlap(leftRange, rightRange)) { + val result = new ScanRange(upperBound, isUpperBoundEqualTo, lowerBound, isLowerBoundEqualTo) + result.mergeIntersect(other) + result + } else { + null + } + } + + def hasOverlap(left: ScanRange, right: ScanRange): Boolean = { + compareRange(left.upperBound, right.lowerBound) >= 0 Review Comment: ScanRange.hasOverlap ignores inclusive/exclusive bound semantics. If `left.upperBound == right.lowerBound`, overlap exists only when both `left.isUpperBoundEqualTo` and `right.isLowerBoundEqualTo` are true; otherwise the ranges are disjoint. Update overlap logic to treat the equality case based on the bound inclusivity flags. ########## spark4/hbase-spark4/src/main/scala/org/apache/hadoop/hbase/spark/datasources/HBaseScan.scala: ########## @@ -0,0 +1,153 @@ +/* + * 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.hadoop.hbase.spark.datasources + +import org.apache.hadoop.hbase.spark.Logging +import org.apache.spark.sql.connector.read.{Batch, Scan} +import org.apache.spark.sql.sources._ +import org.apache.spark.sql.types.StructType +import org.apache.yetus.audience.InterfaceAudience +import scala.collection.mutable.ListBuffer + +/** + * This is a new class in the spark4 module. Implements Scan. + * + * An immutable snapshot of the scan plan after negotiation is complete. + * Holds the pushed filters, required schema, and encoder. + * Builds the RowKeyFilter (for scan range narrowing from row key predicates) and produces the Batch. + * + * In the spark 3 V1 model, there was no separation between "plan" and "execution", the buildScan() method did both. + * + * @param requiredSchema + * @param properties + * @param catalog + * @param pushedFilters + * @param encoderClsName + * @param encoder + */ [email protected] +class HBaseScan( + requiredSchema: StructType, + properties: Map[String, String], + catalog: HBaseTableCatalog, + pushedFilters: Array[Filter], + encoderClsName: String, + @transient encoder: BytesEncoder) + extends Scan + with Logging { + + override def readSchema(): StructType = requiredSchema + + override def toBatch: Batch = { + val rowKeyFilter = buildRowKeyFilter() + new HBaseBatch(requiredSchema, properties, catalog, rowKeyFilter, pushedFilters, encoderClsName) + } + + private def buildRowKeyFilter(): RowKeyFilter = { + var superRowKeyFilter: RowKeyFilter = null + val queryValueList = new ListBuffer[Array[Byte]] + + pushedFilters.foreach { f => + val rowKeyFilter = new RowKeyFilter() + traverseFilterTree(rowKeyFilter, queryValueList, f) + if (superRowKeyFilter == null) { + superRowKeyFilter = rowKeyFilter + } else { + superRowKeyFilter.mergeIntersect(rowKeyFilter) + } + } + + if (superRowKeyFilter == null) { + superRowKeyFilter = new RowKeyFilter + } + superRowKeyFilter + } + + private def traverseFilterTree( + parentRowKeyFilter: RowKeyFilter, + valueArray: ListBuffer[Array[Byte]], + filter: Filter): Unit = { Review Comment: `valueArray` is passed through traverseFilterTree but is not used anywhere in the method body. This makes the API confusing and suggests incomplete logic. Either remove the parameter (and call sites) or use it for its intended purpose to keep the method signature minimal and clear. -- 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]
