Skip to content
Original file line number Diff line number Diff line change
@@ -0,0 +1,44 @@
/*
* Copyright (c) 2026, NVIDIA CORPORATION.
*
* Licensed 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 com.nvidia.spark.rapids

import java.io.FileNotFoundException

/**
* Carries the owning file path across an asynchronous reader boundary.
*
* Spark 4.x needs the path to construct `FAILED_READ_FILE.FILE_NOT_EXIST`, but a
* `Future.get()` otherwise exposes only the reader's `FileNotFoundException`.
*/
object GpuFileNotFoundException {
private final class WithPath(
val filePath: String,
val originalException: FileNotFoundException)
extends FileNotFoundException(originalException.getMessage) {
initCause(originalException)
}

def apply(filePath: String, error: FileNotFoundException): FileNotFoundException = error match {
case pathError: WithPath => pathError
case _ => new WithPath(filePath, error)
}

def unapply(error: FileNotFoundException): Option[(String, FileNotFoundException)] = error match {
case pathError: WithPath => Some((pathError.filePath, pathError.originalException))
case _ => None
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@

package com.nvidia.spark.rapids

import java.io.{File, IOException}
import java.io.{File, FileNotFoundException, IOException}
import java.net.{URI, URISyntaxException}
import java.util.concurrent.{CompletionService, ConcurrentLinkedQueue, ExecutorCompletionService, Future, ThreadPoolExecutor, TimeUnit}
import java.util.concurrent.atomic.{AtomicBoolean, AtomicInteger}
Expand Down Expand Up @@ -282,6 +282,14 @@ object MultiFileReaderUtils {
files: Array[String],
cloudSchemes: Set[String]): Boolean =
!coalescingEnabled || (multiThreadEnabled && hasPathInCloud(files, cloudSchemes))

private[rapids] def attachFilePathToMissingFile[T](runner: AsyncRunner[T], file: Path): Unit = {
runner.addFailureTransformer {
case error: FileNotFoundException =>
GpuFileNotFoundException(file.toUri.toString, error)
case error => error
}
}
}

/**
Expand Down Expand Up @@ -581,6 +589,11 @@ abstract class MultiFileCloudPartitionReaderBase(
// An AsyncRunner wrapper used to update related metrics
val newTaskRunner = (file: PartitionedFile) => {
val runner = getBatchRunner(tc, file, conf, filters)
runner.addFailureTransformer {
case error: FileNotFoundException =>
GpuFileNotFoundException(file.filePath.toString, error)
case error => error
}
val metrics = GpuTaskMetrics.get
val taskId = tc.taskAttemptId()
runner.addPreHook(() => {
Expand Down Expand Up @@ -1472,8 +1485,9 @@ abstract class MultiFileCoalescingPartitionReaderBase(
// use a single buffer and slice it up for different files if we need
val outLocal = hmb.slice(offset, fileBlockSize)
// Third, copy the blocks for each file in parallel using background threads
tasks.add(threadPool.submit(
getBatchRunner(tc, file, outLocal, blocks, offset, batchContext)))
val runner = getBatchRunner(tc, file, outLocal, blocks, offset, batchContext)
MultiFileReaderUtils.attachFilePathToMissingFile(runner, file)
tasks.add(threadPool.submit(runner))
offset += fileBlockSize
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -734,6 +734,8 @@ case class GpuOrcMultiFilePartitionReaderFactory(
} catch {
case e: FileNotFoundException if ignoreMissingFiles =>
logWarning(s"Skipped missing file: ${file.filePath}", e)
case e: FileNotFoundException =>
throw GpuFileNotFoundException(file.filePath.toString, e)
}
}
}
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
/*
* Copyright (c) 2025, NVIDIA CORPORATION.
* Copyright (c) 2025-2026, NVIDIA CORPORATION.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
Expand All @@ -22,6 +22,7 @@ import java.util.concurrent.locks.ReentrantLock
import java.util.function.LongUnaryOperator

import scala.collection.mutable
import scala.util.control.NonFatal

import com.nvidia.spark.rapids.jni.TaskPriority

Expand Down Expand Up @@ -195,6 +196,8 @@ trait AsyncRunner[T] extends Callable[AsyncResult[T]] {
val resultData = try {
beforeExecuteHooks.foreach { hook => hook() }
callImpl()
} catch {
case NonFatal(error) => throw failureTransformer(error)
} finally {
afterExecuteHooks.foreach { hook => hook() }
}
Expand All @@ -206,13 +209,23 @@ trait AsyncRunner[T] extends Callable[AsyncResult[T]] {

private val beforeExecuteHooks = mutable.ArrayBuffer.empty[() => Unit]
private val afterExecuteHooks = mutable.ArrayBuffer.empty[() => Unit]
private var failureTransformer: Throwable => Throwable = identity

// Add hook to be executed right before the task execution.
def addPreHook(hook: () => Unit): Unit = beforeExecuteHooks += hook

// Add hook to be executed right after the task execution.
def addPostHook(hook: () => Unit): Unit = afterExecuteHooks += hook

/**
* Adds a transformer for failures thrown by the runner body. Transformers are applied in the
* order they are added and are not invoked on the successful execution path.
*/
def addFailureTransformer(transformer: Throwable => Throwable): Unit = {
val previousTransformer = failureTransformer
failureTransformer = error => transformer(previousTransformer(error))
}

/**
* This method is called when the required resource has been just acquired from pool.
* It can be overridden by subclasses to perform actions right after the acquisition.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1297,7 +1297,8 @@ abstract class AbstractGpuParquetMultiFilePartitionReaderFactory(
hasInt96Timestamps = false)
BlockMetaWithPartFile(meta, file)
// Throw FileNotFoundException even if `ignoreCorruptFiles` is true
case e: FileNotFoundException if !ignoreMissingFiles => throw e
case e: FileNotFoundException if !ignoreMissingFiles =>
throw GpuFileNotFoundException(file.filePath.toString, e)
// If ignoreMissingFiles=true, this case will never be reached. But it's ok
// to leave this branch here.
case e@(_: RuntimeException | _: IOException) if ignoreCorruptFiles =>
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -45,7 +45,8 @@ abstract class GpuBatchScanExecBase(
// return an empty RDD with 1 partition if dynamic filtering removed the only split
sparkContext.parallelize(Array.empty[InternalRow], 1)
} else {
new GpuDataSourceRDD(sparkContext, filteredPartitions, readerFactory)
new GpuDataSourceRDD(
sparkContext, filteredPartitions, readerFactory, includeRefreshHint = false)
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,16 +16,18 @@

package com.nvidia.spark.rapids.shims

import java.util.concurrent.ConcurrentHashMap
import java.io.FileNotFoundException
import java.util.concurrent.{ConcurrentHashMap, ExecutionException}

import com.nvidia.spark.rapids.{FileSystemBytesReadTracker, GpuFileNotFoundException}
import com.nvidia.spark.rapids.Arm.closeOnExcept
import com.nvidia.spark.rapids.FileSystemBytesReadTracker
import com.nvidia.spark.rapids.ScalableTaskCompletion.onTaskCompletion

import org.apache.spark.{InterruptibleIterator, Partition, SparkContext, SparkException, TaskContext}
import org.apache.spark.rdd.RDD
import org.apache.spark.sql.catalyst.InternalRow
import org.apache.spark.sql.connector.read.{InputPartition, PartitionReader, PartitionReaderFactory}
import org.apache.spark.sql.execution.datasources.FilePartition
import org.apache.spark.sql.rapids.execution.TrampolineUtil
import org.apache.spark.sql.vectorized.ColumnarBatch

Expand Down Expand Up @@ -63,6 +65,7 @@ class GpuDataSourceRDD(
sc: SparkContext,
@transient private val inputPartitions: Seq[Seq[InputPartition]],
partitionReaderFactory: PartitionReaderFactory,
includeRefreshHint: Boolean,
customMetricsFactory: GpuDataSourceCustomMetricsFactory =
NoopGpuDataSourceCustomMetricsFactory
) extends RDD[InternalRow](sc, Nil) {
Expand Down Expand Up @@ -94,27 +97,53 @@ class GpuDataSourceRDD(
private val inputPartitions = castPartition(split).inputPartitions
private var currentIter: Option[Iterator[Object]] = None
private var currentIndex: Int = 0
private var currentInputPartition: InputPartition = _

override def hasNext: Boolean = {
override def hasNext: Boolean = try {
val result = currentIter.exists(_.hasNext) || advanceToNextIter()
if (!result) {
bytesReadTracker.update()
}
result
} catch {
case e: FileNotFoundException =>
throw GpuDataSourceRDD.missingFileError(
e, includeRefreshHint, currentInputPartition)
case e: ExecutionException =>
e.getCause match {
case cause: FileNotFoundException =>
throw GpuDataSourceRDD.missingFileError(
cause, includeRefreshHint, currentInputPartition)
case _ => throw e
}
}

override def next(): Object = {

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Updated.

if (!hasNext) {
throw new NoSuchElementException("No more elements")
try {
if (!hasNext) {
throw new NoSuchElementException("No more elements")
}
currentIter.get.next()
} catch {
case e: FileNotFoundException =>
throw GpuDataSourceRDD.missingFileError(
e, includeRefreshHint, currentInputPartition)
case e: ExecutionException =>
e.getCause match {
case cause: FileNotFoundException =>
throw GpuDataSourceRDD.missingFileError(
cause, includeRefreshHint, currentInputPartition)
case _ => throw e
}
}
currentIter.get.next()
}

private def advanceToNextIter(): Boolean = {
if (currentIndex >= inputPartitions.length) {
false
} else {
val inputPartition = inputPartitions(currentIndex)
currentInputPartition = inputPartition
currentIndex += 1

// TODO: SPARK-25083 remove the type erasure hack in data source scan
Expand Down Expand Up @@ -236,14 +265,35 @@ class GpuDataSourceRDD(
}

object GpuDataSourceRDD {
private def missingFileError(
error: FileNotFoundException,
includeRefreshHint: Boolean,
inputPartition: InputPartition): Throwable = {
val (filePath, originalError) = error match {
case GpuFileNotFoundException(path, originalException) =>
(Some(path), originalException)
case _ =>
(singleFilePath(inputPartition), error)
}
MissingFileErrorShim.convert(filePath, originalError, includeRefreshHint)
}

private def singleFilePath(inputPartition: InputPartition): Option[String] = {
Option(inputPartition).collect { case filePartition: FilePartition =>
SparkShimImpl.getPartitionFiles(filePartition)
}.filter(_.length == 1).map(_.head.filePath.toString)
}

private case class GpuDataSourceRDDPartition(
override val index: Int,
inputPartitions: Seq[InputPartition]) extends Partition

def apply(
sc: SparkContext,
inputPartitions: Seq[InputPartition],
partitionReaderFactory: PartitionReaderFactory): GpuDataSourceRDD = {
new GpuDataSourceRDD(sc, inputPartitions.map(Seq(_)), partitionReaderFactory)
partitionReaderFactory: PartitionReaderFactory,
includeRefreshHint: Boolean): GpuDataSourceRDD = {
new GpuDataSourceRDD(
sc, inputPartitions.map(Seq(_)), partitionReaderFactory, includeRefreshHint)
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -281,7 +281,8 @@ case class GpuAvroMultiFilePartitionReaderFactory(
logWarning(s"Skipped missing file: ${file.filePath}", e)
AvroBlockMeta(null, 0L, Seq.empty)
// Throw FileNotFoundException even if `ignoreCorruptFiles` is true
case e: FileNotFoundException if !ignoreMissingFiles => throw e
case e: FileNotFoundException if !ignoreMissingFiles =>
throw GpuFileNotFoundException(file.filePath.toString, e)
case e@(_: RuntimeException | _: IOException) if ignoreCorruptFiles =>
logWarning(
s"Skipped the rest of the content in the corrupted file: ${file.filePath}", e)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -612,7 +612,11 @@ case class GpuFileSourceScanExec(
logDebug(s"Using Datasource RDD, files are: " +
s"${prunedPartitions.flatMap(FilePartitionShims.getFiles).mkString(",")}")
// note we use the v2 DataSourceRDD instead of FileScanRDD so we don't have to copy more code
GpuDataSourceRDD(relation.sparkSession.sparkContext, locatedPartitions, readerFactory)
GpuDataSourceRDD(
relation.sparkSession.sparkContext,
locatedPartitions,
readerFactory,
includeRefreshHint = true)
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ package org.apache.spark.sql.rapids.shims
import java.io.{FileNotFoundException, IOException}

import com.nvidia.spark.rapids.ScalableTaskCompletion.onTaskCompletion
import com.nvidia.spark.rapids.shims.SparkShimImpl
import com.nvidia.spark.rapids.shims.{MissingFileErrorShim, SparkShimImpl}
import org.apache.parquet.io.ParquetDecodingException

import org.apache.spark.{Partition => RDDPartition, SparkUpgradeException, TaskContext}
Expand Down Expand Up @@ -73,11 +73,19 @@ class GpuFileScanRDD(
// InterruptibleIterator, but we inline it here instead of wrapping the iterator in order
// to avoid performance overhead.
context.killTaskIfInterrupted()
(currentIterator != null && currentIterator.hasNext) || nextIterator()
try {
(currentIterator != null && currentIterator.hasNext) || nextIterator()
} catch {
case e: FileNotFoundException => throw convertMissingFile(e)
}
}

def next(): Object = {
val nextElement = currentIterator.next()
val nextElement = try {
currentIterator.next()
} catch {
case e: FileNotFoundException => throw convertMissingFile(e)
}
// TODO: we should have a better separation of row based and batch based scan, so that we
// don't need to run this `if` for every record.
if (nextElement.isInstanceOf[ColumnarBatch]) {
Expand All @@ -94,19 +102,11 @@ class GpuFileScanRDD(
nextElement
}

private def readCurrentFile(): Iterator[InternalRow] = {
try {
readFunction(currentFile)
} catch {
case e: FileNotFoundException =>
throw new FileNotFoundException(
e.getMessage + "\n" +
"It is possible the underlying files have been updated. " +
"You can explicitly invalidate the cache in Spark by " +
"running 'REFRESH TABLE tableName' command in SQL or " +
"by recreating the Dataset/DataFrame involved.")
}
}
private def readCurrentFile(): Iterator[InternalRow] = readFunction(currentFile)

private def convertMissingFile(error: FileNotFoundException): Throwable =
MissingFileErrorShim.convert(
Some(currentFile.filePath.toString), error, includeRefreshHint = true)

/** Advances to the next file. Returns true if a new non-empty iterator is available. */
private def nextIterator(): Boolean = {
Expand Down Expand Up @@ -201,4 +201,3 @@ class GpuFileScanRDD(
split.asInstanceOf[FilePartition].preferredLocations()
}
}

Loading
Loading