Skip to content
Open
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 @@ -17,13 +17,14 @@

package org.apache.spark.scheduler.cluster

import java.io.{Externalizable, ObjectInput, ObjectOutput}
import java.nio.ByteBuffer

import org.apache.spark.TaskState.TaskState
import org.apache.spark.resource.{ResourceInformation, ResourceProfile}
import org.apache.spark.resource.{CpuAmount, ResourceInformation, ResourceProfile}
import org.apache.spark.rpc.RpcEndpointRef
import org.apache.spark.scheduler.{ExecutorLossReason, MiscellaneousProcessDetails}
import org.apache.spark.util.SerializableBuffer
import org.apache.spark.util.{SerializableBuffer, Utils}

private[spark] sealed trait CoarseGrainedClusterMessage extends Serializable

Expand Down Expand Up @@ -78,14 +79,54 @@ private[spark] object CoarseGrainedClusterMessages {

case class LaunchedExecutor(executorId: String) extends CoarseGrainedClusterMessage

// Serialized manually (Externalizable): default Java serialization of `state` (a Scala
// Enumeration value) and `taskCpus` (a BigDecimal) each drag in a large class-descriptor graph
// -- together ~1.3KB per message on this hot, per-task-state-change path.
case class StatusUpdate(
executorId: String,
taskId: Long,
state: TaskState,
data: SerializableBuffer,
taskCpus: BigDecimal,
resources: Map[String, Map[String, Long]] = Map.empty)
extends CoarseGrainedClusterMessage
var executorId: String,
var taskId: Long,
var state: TaskState,
var data: SerializableBuffer,
var taskCpus: BigDecimal,
var resources: Map[String, Map[String, Long]] = Map.empty)
extends CoarseGrainedClusterMessage with Externalizable {

def this() = this(null, 0L, null, null, null, Map.empty) // For deserialization only

override def writeExternal(out: ObjectOutput): Unit = Utils.tryOrIOException {
out.writeUTF(executorId)
out.writeLong(taskId)
out.writeByte(state.id)
// taskCpus as a normalized decimal string (like TaskDescription); round-trips exactly.
out.writeUTF(CpuAmount.toDisplayString(taskCpus))
// Reuse SerializableBuffer's channel-based write (no extra copy for large results).
out.writeObject(data)
out.writeInt(resources.size)
resources.foreach { case (rName, addressAmounts) =>
out.writeUTF(rName)
out.writeInt(addressAmounts.size)
addressAmounts.foreach { case (address, amount) =>
out.writeUTF(address)
out.writeLong(amount)
}
}
}

override def readExternal(in: ObjectInput): Unit = Utils.tryOrIOException {
executorId = in.readUTF()
taskId = in.readLong()
state = org.apache.spark.TaskState(in.readByte().toInt)
taskCpus = CpuAmount.normalize(BigDecimal(in.readUTF()))
data = in.readObject().asInstanceOf[SerializableBuffer]
val numResources = in.readInt()
resources = Iterator.fill(numResources) {
val rName = in.readUTF()
val numAddresses = in.readInt()
val addressAmounts = Iterator.fill(numAddresses)(in.readUTF() -> in.readLong()).toMap
rName -> addressAmounts
}.toMap
}
}

object StatusUpdate {
/** Alternate factory method that takes a ByteBuffer directly for the data field */
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
/*
* 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.spark.scheduler.cluster

import java.nio.ByteBuffer

import org.apache.spark.{SparkConf, SparkFunSuite, TaskState}
import org.apache.spark.resource.CpuAmount
import org.apache.spark.scheduler.cluster.CoarseGrainedClusterMessages.StatusUpdate
import org.apache.spark.serializer.JavaSerializer

class CoarseGrainedClusterMessagesSuite extends SparkFunSuite {

private val ser = new JavaSerializer(new SparkConf(false)).newInstance()

private def roundTrip(su: StatusUpdate): StatusUpdate =
ser.deserialize[StatusUpdate](ser.serialize(su))

private def statusUpdate(
state: TaskState.TaskState = TaskState.RUNNING,
data: ByteBuffer = ByteBuffer.allocate(0),
taskCpus: BigDecimal = CpuAmount.normalize(BigDecimal(1)),
resources: Map[String, Map[String, Long]] = Map.empty): StatusUpdate =
StatusUpdate("exec-17", 1234567L, state, data, taskCpus, resources)

test("StatusUpdate round-trips all fields") {
val data = ByteBuffer.wrap(Array[Byte](1, 2, 3, 4, 5))
val resources = Map(
"gpu" -> Map("0" -> 1L, "1" -> 2L),
"fpga" -> Map("addr-a" -> 7L))
val cpus = CpuAmount.normalize(BigDecimal("0.5"))
val su = statusUpdate(TaskState.FINISHED, data, cpus, resources)
val rt = roundTrip(su)

assert(rt.executorId === su.executorId)
assert(rt.taskId === su.taskId)
assert(rt.state === su.state)
assert(rt.taskCpus === cpus)
assert(rt.taskCpus.scale === cpus.scale)
assert(rt.data.value === su.data.value)
assert(rt.resources === resources)
}

test("StatusUpdate round-trips every TaskState") {
TaskState.values.foreach { s =>
assert(roundTrip(statusUpdate(state = s)).state === s)
}
}

test("SPARK-58192: fractional taskCpus round-trips exactly") {
Seq("1", "0.5", "0.333333333", "2.25").foreach { amount =>
val cpus = CpuAmount.normalize(BigDecimal(amount))
val rt = roundTrip(statusUpdate(taskCpus = cpus))
assert(rt.taskCpus === cpus, s"value for $amount")
assert(rt.taskCpus.scale === cpus.scale, s"scale for $amount")
}
}

test("StatusUpdate round-trips with empty data and resources") {
val rt = roundTrip(statusUpdate())
assert(rt.data.value.remaining() === 0)
assert(rt.resources.isEmpty)
}

test("StatusUpdate is compact after manual serialization") {
// With default Java serialization, `state` (a Scala Enumeration value) and `taskCpus` (a
// BigDecimal) alone were ~1.3KB and an empty-payload message was ~1.7KB. Manual
// (Externalizable) encoding brings it to ~191 bytes; assert an ample upper bound.
val size = ser.serialize(statusUpdate()).remaining()
assert(size < 512, s"StatusUpdate serialized to $size bytes")
}
}