Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@
*/
package org.apache.spark.sql.execution

import org.apache.gluten.sql.shims.SparkShimLoader

import org.apache.spark.TaskContext
import org.apache.spark.internal.Logging
import org.apache.spark.internal.io.{FileCommitProtocol, FileNameSpec, HadoopMapReduceCommitProtocol}
Expand Down Expand Up @@ -44,15 +46,15 @@ class SparkWriteFilesCommitProtocol(
extends Logging {
assert(committer.isInstanceOf[HadoopMapReduceCommitProtocol])

val sparkStageId = TaskContext.get().stageId()
val sparkPartitionId = TaskContext.get().partitionId()
val sparkAttemptNumber = TaskContext.get().taskAttemptId().toInt & Int.MaxValue
val sparkStageId: Int = TaskContext.get().stageId()
val sparkPartitionId: Int = TaskContext.get().partitionId()
val sparkAttemptNumber: Int = TaskContext.get().taskAttemptId().toInt & Int.MaxValue
private val jobId = createJobID(jobTrackerID, sparkStageId)

private val taskId = new TaskID(jobId, TaskType.MAP, sparkPartitionId)
private val taskAttemptId = new TaskAttemptID(taskId, sparkAttemptNumber)

private var fileNames: mutable.Set[String] = null
private var fileNames: mutable.Set[String] = _

// Set up the attempt context required to use in the output committer.
val taskAttemptContext: TaskAttemptContext = {
Expand Down Expand Up @@ -86,7 +88,9 @@ class SparkWriteFilesCommitProtocol(
// Note that %05d does not truncate the split number, so if we have more than 100000 tasks,
// the file name is fine and won't overflow.
val split = taskAttemptContext.getTaskAttemptID.getTaskID.getId
val fileName = f"${spec.prefix}part-$split%05d-${UUID.randomUUID().toString()}${spec.suffix}"
val basename = taskAttemptContext.getConfiguration.get("mapreduce.output.basename", "part")
val fileName = f"${spec.prefix}$basename-$split%05d-${UUID.randomUUID().toString}${spec.suffix}"

fileNames += fileName
fileName
}
Expand All @@ -103,7 +107,13 @@ class SparkWriteFilesCommitProtocol(
stagingDir.toString
}

def commitTask(): Unit = {
private def enrichWriteError[T](path: => String)(f: => T): T = try {
f
} catch {
case t: Throwable => SparkShimLoader.getSparkShims.enrichWriteException(t, description.path)
}

def commitTask(): Unit = enrichWriteError(description.path) {
val (_, taskCommitTime) = Utils.timeTakenMs {
committer.commitTask(taskAttemptContext)
}
Expand All @@ -114,7 +124,7 @@ class SparkWriteFilesCommitProtocol(
}
}

def abortTask(writePath: String): Unit = {
def abortTask(writePath: String): Unit = enrichWriteError(description.path) {
committer.abortTask(taskAttemptContext)

// Deletes the files written by current task.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -475,9 +475,6 @@ class VeloxTestSettings extends BackendTestSettings {
// Velox parquet reader not allow offset zero.
.exclude("SPARK-40128 read DELTA_LENGTH_BYTE_ARRAY encoded strings")
// TODO: fix in Spark-4.0
.exclude("SPARK-49991: Respect 'mapreduce.output.basename' to generate file names")
.exclude("SPARK-6330 regression test")
.exclude("SPARK-7837 Do not close output writer twice when commitTask() fails")
.exclude("explode nested lists crossing a rowgroup boundary")
enableSuite[GlutenParquetV1PartitionDiscoverySuite]
enableSuite[GlutenParquetV2PartitionDiscoverySuite]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,7 @@ trait SparkShims {

def generateFileScanRDD(
sparkSession: SparkSession,
readFunction: (PartitionedFile) => Iterator[InternalRow],
readFunction: PartitionedFile => Iterator[InternalRow],
filePartitions: Seq[FilePartition],
fileSourceScanExec: FileSourceScanExec): FileScanRDD

Expand Down Expand Up @@ -145,7 +145,7 @@ trait SparkShims {
Expression,
Expression,
Int,
Int) => TypedImperativeAggregate[T]): Expression;
Int) => TypedImperativeAggregate[T]): Expression

def replaceMightContain[T](
expr: Expression,
Expand Down Expand Up @@ -343,6 +343,11 @@ trait SparkShims {
t)
}

// Compatibility method for Spark 4.0: rethrows the exception cause to maintain API compatibility
def enrichWriteException(cause: Throwable, path: String): Nothing = {
throw cause
}

def getFileSourceScanStream(scan: FileSourceScanExec): Option[SparkDataStream] = {
None
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -65,8 +65,8 @@ import org.apache.parquet.schema.MessageType
import java.time.ZoneOffset
import java.util.{Map => JMap}

import scala.collection.JavaConverters._
import scala.collection.mutable
import scala.jdk.CollectionConverters._
import scala.reflect.ClassTag

class Spark40Shims extends SparkShims {
Expand Down Expand Up @@ -151,7 +151,7 @@ class Spark40Shims extends SparkShims {
options: CaseInsensitiveStringMap,
partitionFilters: Seq[Expression],
dataFilters: Seq[Expression]): TextScan = {
new TextScan(
TextScan(
sparkSession,
fileIndex,
dataSchema,
Expand Down Expand Up @@ -742,6 +742,9 @@ class Spark40Shims extends SparkShims {
throw t
}

override def enrichWriteException(cause: Throwable, path: String): Nothing = {
GlutenFileFormatWriter.wrapWriteError(cause, path)
}
override def getFileSourceScanStream(scan: FileSourceScanExec): Option[SparkDataStream] = {
scan.stream
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ package org.apache.spark.sql.execution

import org.apache.spark.internal.io.FileCommitProtocol
import org.apache.spark.sql.catalyst.InternalRow
import org.apache.spark.sql.errors.QueryExecutionErrors
import org.apache.spark.sql.execution.datasources.{FileFormatWriter, WriteJobDescription, WriteTaskResult}

object GlutenFileFormatWriter {
Expand All @@ -40,4 +41,9 @@ object GlutenFileFormatWriter {
None
)
}

// Wrapper for throwing standardized write error using QueryExecutionErrors
def wrapWriteError(cause: Throwable, writePath: String): Nothing = {
throw QueryExecutionErrors.taskFailedWhileWritingRowsError(writePath, cause)
}
}