Skip to content
Closed
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 @@ -179,7 +179,8 @@ class SparkConnectGetStatusHandler(responseObserver: StreamObserver[proto.GetSta
SparkConnectPluginRegistry.getStatusRegistry.flatMap { plugin =>
try {
plugin.processRequestExtensions(sessionHolder, requestExtensions).toScala match {
case Some(extensions) => extensions.asScala.toSeq
// Filter nulls inside the isolation boundary so a null element can't NPE addExtensions.
case Some(extensions) => extensions.asScala.iterator.filter(_ != null).toSeq
case None => Seq.empty
}
} catch {
Expand All @@ -202,7 +203,8 @@ class SparkConnectGetStatusHandler(responseObserver: StreamObserver[proto.GetSta
plugin
.processOperationExtensions(operationId, sessionHolder, operationExtensions)
.toScala match {
case Some(extensions) => extensions.asScala.toSeq
// Filter nulls inside the isolation boundary so a null element can't NPE addExtensions.
case Some(extensions) => extensions.asScala.iterator.filter(_ != null).toSeq
case None => Seq.empty
}
} catch {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -88,6 +88,28 @@ class NoOpGetStatusPlugin extends GetStatusPlugin {
Optional.empty()
}

/**
* A plugin that returns a list containing a null element, to test null filtering.
*/
class NullElementGetStatusPlugin extends GetStatusPlugin {
override def processRequestExtensions(
sessionHolder: SessionHolder,
requestExtensions: util.List[protobuf.Any]): Optional[util.List[protobuf.Any]] =
listWithNull()

override def processOperationExtensions(
operationId: String,
sessionHolder: SessionHolder,
operationExtensions: util.List[protobuf.Any]): Optional[util.List[protobuf.Any]] =
listWithNull()

private def listWithNull(): Optional[util.List[protobuf.Any]] = {
val result = new util.ArrayList[protobuf.Any]()
result.add(null)
Optional.of(result)
}
}

/**
* A plugin that always throws a RuntimeException.
*/
Expand Down Expand Up @@ -344,6 +366,32 @@ class GetStatusHandlerSuite extends SharedSparkSession {
assert(opExtValues.contains(s"op-echo:${executeHolder.operationId}:safe"))
assert(opExtValues.contains(s"second-op:${executeHolder.operationId}:safe"))
}

test("GetStatus filters null extension elements returned by a plugin") {
SparkConnectPluginRegistry.setGetStatusPluginsForTesting(
Seq(new NullElementGetStatusPlugin(), new EchoGetStatusPlugin()))
val sessionHolder = SparkConnectTestUtils.createDummySessionHolder(spark)
val command = proto.Command.newBuilder().build()
val executeHolder = SparkConnectTestUtils.createDummyExecuteHolder(sessionHolder, command)

val reqExt = protobuf.Any.pack(StringValue.of("data"))
val opExt = protobuf.Any.pack(StringValue.of("data"))
val response = sendGetOperationStatusRequest(
sessionHolder.sessionId,
operationIds = Seq(executeHolder.operationId),
userId = sessionHolder.userId,
requestExtensions = Seq(reqExt),
operationExtensions = Seq(opExt))

// The null element is dropped; only the healthy plugin's extension survives, at both levels.
val responseExtValues = response.getExtensionsList.asScala
.map(_.unpack(classOf[StringValue]).getValue)
assert(responseExtValues == Seq("request-echo:data"))

val opExtValues = response.getOperationStatusesList.asScala.head.getExtensionsList.asScala
.map(_.unpack(classOf[StringValue]).getValue)
assert(opExtValues == Seq(s"op-echo:${executeHolder.operationId}:data"))
}
}

private class GetStatusResponseObserver extends StreamObserver[proto.GetStatusResponse] {
Expand Down