From 9a672037d43b4986b16ae4cf75b22708896c92b4 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 09:00:16 +0000 Subject: [PATCH 01/10] [CONNECT] Show ML cache status in Spark Connect UI --- .../apache/spark/sql/connect/ml/MLCache.scala | 19 +++++++++ .../sql/connect/service/SessionHolder.scala | 13 ++++++- .../connect/service/SparkConnectService.scala | 3 +- .../service/SparkConnectSessionManager.scala | 5 +++ .../connect/ui/SparkConnectServerPage.scala | 32 ++++++++++++++- .../connect/ui/SparkConnectServerTab.scala | 11 +++++- .../spark/sql/connect/ui/ToolTips.scala | 3 ++ .../apache/spark/sql/connect/ml/MLSuite.scala | 24 ++++++++++++ .../ui/SparkConnectServerPageSuite.scala | 39 +++++++++++++++++-- 9 files changed, 141 insertions(+), 8 deletions(-) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala index 5deb9ce1c3f19..62b729091ac42 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala @@ -280,4 +280,23 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { } info.result() } + + def getStatus: MLCacheStatus = this.synchronized { + MLCacheStatus( + memoryControlEnabled = getMemoryControlEnabled, + cachedObjectCount = cachedModel.size(), + inMemorySizeBytes = totalMLCacheInMemorySizeBytes.get(), + maxInMemorySizeBytes = sessionHolder.session.conf.get( + Connect.CONNECT_SESSION_CONNECT_ML_CACHE_MEMORY_CONTROL_MAX_IN_MEMORY_SIZE), + totalSizeBytes = totalMLCacheSizeBytes.get(), + maxTotalSizeBytes = getMLCacheMaxSize) + } } + +private[connect] case class MLCacheStatus( + memoryControlEnabled: Boolean, + cachedObjectCount: Int, + inMemorySizeBytes: Long, + maxInMemorySizeBytes: Long, + totalSizeBytes: Long, + maxTotalSizeBytes: Long) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SessionHolder.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SessionHolder.scala index 2276230545e67..e97bcb0c786e3 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SessionHolder.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SessionHolder.scala @@ -38,7 +38,7 @@ import org.apache.spark.sql.classic.SparkSession import org.apache.spark.sql.connect.IllegalStateErrors import org.apache.spark.sql.connect.common.InvalidPlanInput import org.apache.spark.sql.connect.config.Connect -import org.apache.spark.sql.connect.ml.MLCache +import org.apache.spark.sql.connect.ml.{MLCache, MLCacheStatus} import org.apache.spark.sql.connect.pipelines.DataflowGraphRegistry import org.apache.spark.sql.connect.planner.PythonStreamingQueryListener import org.apache.spark.sql.connect.planner.StreamingForeachBatchHelper @@ -132,7 +132,16 @@ case class SessionHolder(userId: String, sessionId: String, session: SparkSessio new ConcurrentHashMap() // ML model cache - private[connect] lazy val mlCache = new MLCache(this) + @volatile private var mlCacheInitialized = false + private[connect] lazy val mlCache = { + val cache = new MLCache(this) + mlCacheInitialized = true + cache + } + + private[connect] def getMLCacheStatus: Option[MLCacheStatus] = { + if (mlCacheInitialized) Some(mlCache.getStatus) else None + } // Mapping from id to StreamingQueryListener. Used for methods like removeListener() in // StreamingQueryManager. diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectService.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectService.scala index c76794e3b6ec1..0c18a7491c1f2 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectService.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectService.scala @@ -386,7 +386,8 @@ object SparkConnectService extends Logging { Some( new SparkConnectServerTab( new SparkConnectServerAppStatusStore(kvStore), - SparkConnectServerTab.getSparkUI(sc))) + SparkConnectServerTab.getSparkUI(sc), + Some(sessionManager))) } else { None } diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectSessionManager.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectSessionManager.scala index d3ddf592e9e7d..615d859839427 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectSessionManager.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectSessionManager.scala @@ -32,6 +32,7 @@ import org.apache.spark.internal.Logging import org.apache.spark.internal.LogKeys.{INTERVAL, SESSION_HOLD_INFO} import org.apache.spark.sql.classic.SparkSession import org.apache.spark.sql.connect.config.Connect.{CONNECT_SESSION_MANAGER_CLOSED_SESSIONS_TOMBSTONES_SIZE, CONNECT_SESSION_MANAGER_DEFAULT_SESSION_TIMEOUT, CONNECT_SESSION_MANAGER_MAINTENANCE_INTERVAL} +import org.apache.spark.sql.connect.ml.MLCacheStatus import org.apache.spark.util.ThreadUtils /** @@ -284,6 +285,10 @@ class SparkConnectSessionManager extends Logging { closedSessionsCache.asMap.asScala.values.toSeq } + private[connect] def getMLCacheStatus(key: SessionKey): Option[MLCacheStatus] = { + Option(sessionStore.get(key)).flatMap(_.getMLCacheStatus) + } + /** * Schedules periodic maintenance checks if it is not already scheduled. * diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala index ea78b2dc59f6a..4f418c09523ee 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala @@ -404,7 +404,8 @@ private[ui] class SessionStatsPagedTable( ("Start Time", true, None), ("Finish Time", true, None), ("Duration", true, Some(SPARK_CONNECT_SESSION_DURATION)), - ("Total Execute", true, Some(SPARK_CONNECT_SESSION_TOTAL_EXECUTE))) + ("Total Execute", true, Some(SPARK_CONNECT_SESSION_TOTAL_EXECUTE)), + ("ML Cache", false, Some(SPARK_CONNECT_SESSION_ML_CACHE))) isSortColumnValid(sessionTableHeadersAndTooltips, sortColumn) @@ -431,8 +432,37 @@ private[ui] class SessionStatsPagedTable( {if (session.finishTimestamp > 0) formatDate(session.finishTimestamp)} {formatDurationVerbose(session.totalTime)} {session.totalExecution.toString} + {renderMLCacheStatus(session)} } + + private def renderMLCacheStatus(session: SessionInfo): Seq[Node] = { + if (!parent.hasLiveMLCacheStatus || session.finishTimestamp > 0) { + N/A + } else { + parent + .getMLCacheStatus(session.userId, session.sessionId) + .map { status => + if (status.memoryControlEnabled) { + val objectLabel = if (status.cachedObjectCount == 1) "object" else "objects" + + {s"${status.cachedObjectCount} $objectLabel in memory"}
+ {s"${Utils.bytesToString(status.inMemorySizeBytes)} / " + + s"${Utils.bytesToString(status.maxInMemorySizeBytes)} memory"}
+ {s"${Utils.bytesToString(status.totalSizeBytes)} / " + + s"${Utils.bytesToString(status.maxTotalSizeBytes)} total"} +
+ } else { + val objectLabel = if (status.cachedObjectCount == 1) "object" else "objects" + + Memory control disabled
+ {s"${status.cachedObjectCount} cached $objectLabel"} +
+ } + } + .getOrElse(Not used) + } + } } private[ui] class SqlStatsTableRow( diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerTab.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerTab.scala index c5ea0bf618b52..181b9fa8e95b6 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerTab.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerTab.scala @@ -21,12 +21,15 @@ import java.util.Date import org.apache.spark.SparkContext import org.apache.spark.internal.Logging +import org.apache.spark.sql.connect.ml.MLCacheStatus +import org.apache.spark.sql.connect.service.{SessionKey, SparkConnectSessionManager} import org.apache.spark.sql.errors.QueryExecutionErrors import org.apache.spark.ui.{SparkUI, SparkUITab} private[connect] class SparkConnectServerTab( val store: SparkConnectServerAppStatusStore, - sparkUI: SparkUI) + sparkUI: SparkUI, + sessionManager: Option[SparkConnectSessionManager] = None) extends SparkUITab(sparkUI, "connect") with Logging { @@ -47,6 +50,12 @@ private[connect] class SparkConnectServerTab( parent.detachTab(this) } + def getMLCacheStatus(userId: String, sessionId: String): Option[MLCacheStatus] = { + sessionManager.flatMap(_.getMLCacheStatus(SessionKey(userId, sessionId))) + } + + def hasLiveMLCacheStatus: Boolean = sessionManager.nonEmpty + override def displayOrder: Int = 3 } diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/ToolTips.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/ToolTips.scala index 9b51ace83c6c1..ff385db60af5b 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/ToolTips.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/ToolTips.scala @@ -36,4 +36,7 @@ private[ui] object ToolTips { val SPARK_CONNECT_SESSION_DURATION = "Elapsed time since session start, or until closed if the session was closed" + val SPARK_CONNECT_SESSION_ML_CACHE = + "Current ML cache usage for an active Spark Connect session" + } diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ml/MLSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ml/MLSuite.scala index 9ba5a499ba8fe..c7bc9c888cb98 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ml/MLSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ml/MLSuite.scala @@ -387,6 +387,30 @@ class MLSuite extends MLHelper { } } + test("MLCache status") { + val sessionHolder = SparkConnectTestUtils.createDummySessionHolder(spark) + sessionHolder.session.conf + .set(Connect.CONNECT_SESSION_CONNECT_ML_CACHE_MEMORY_CONTROL_ENABLED.key, "true") + sessionHolder.session.conf.set( + Connect.CONNECT_SESSION_CONNECT_ML_CACHE_MEMORY_CONTROL_MAX_IN_MEMORY_SIZE.key, + 16384) + sessionHolder.session.conf.set( + Connect.CONNECT_SESSION_CONNECT_ML_CACHE_MEMORY_CONTROL_MAX_STORAGE_SIZE.key, + 65536) + + // Reading UI status should not initialize an unused cache. + assert(sessionHolder.getMLCacheStatus.isEmpty) + + trainLogisticRegressionModel(sessionHolder) + val status = sessionHolder.getMLCacheStatus.get + assert(status.memoryControlEnabled) + assert(status.cachedObjectCount === 1) + assert(status.inMemorySizeBytes > 0) + assert(status.maxInMemorySizeBytes === 16384) + assert(status.totalSizeBytes === status.inMemorySizeBytes) + assert(status.maxTotalSizeBytes === 65536) + } + test("MLCache offloading works") { val sessionHolder = SparkConnectTestUtils.createDummySessionHolder(spark) sessionHolder.session.conf diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala index 7f6af17bc41b3..92693a31fa25d 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala @@ -26,6 +26,7 @@ import org.scalatest.BeforeAndAfter import org.apache.spark.{SharedSparkContext, SparkConf, SparkFunSuite} import org.apache.spark.scheduler.SparkListenerJobStart +import org.apache.spark.sql.connect.ml.MLCacheStatus import org.apache.spark.sql.connect.service._ import org.apache.spark.status.ElementTrackingStore import org.apache.spark.util.kvstore.InMemoryStore @@ -47,7 +48,8 @@ class SparkConnectServerPageSuite /** * Run a dummy session and return the store */ - private def getStatusStore: SparkConnectServerAppStatusStore = { + private def getStatusStore( + closeSession: Boolean = true): SparkConnectServerAppStatusStore = { kvstore = new ElementTrackingStore(new InMemoryStore, new SparkConf()) // val server = mock(classOf[SparkConnectServer], RETURNS_SMART_NULLS) val sparkConf = new SparkConf @@ -74,8 +76,10 @@ class SparkConnectServerPageSuite SparkListenerConnectOperationFinished("jobTag", "operationId", System.currentTimeMillis())) listener.onOtherEvent( SparkListenerConnectOperationClosed("jobTag", "operationId", System.currentTimeMillis())) - listener.onOtherEvent( - SparkListenerConnectSessionClosed("sessionId", "userId", System.currentTimeMillis())) + if (closeSession) { + listener.onOtherEvent( + SparkListenerConnectSessionClosed("sessionId", "userId", System.currentTimeMillis())) + } statusStore } @@ -134,6 +138,35 @@ class SparkConnectServerPageSuite " data-bs-target=\"#aggregated-sqlsessionstat\"")) } + test("Spark Connect Server page should show live ML cache status") { + val store = getStatusStore(closeSession = false) + + val request = mock(classOf[HttpServletRequest]) + val tab = mock(classOf[SparkConnectServerTab], RETURNS_SMART_NULLS) + when(tab.startTime).thenReturn(Calendar.getInstance().getTime) + when(tab.store).thenReturn(store) + when(tab.appName).thenReturn("testing") + when(tab.headerTabs).thenReturn(Seq.empty) + when(tab.hasLiveMLCacheStatus).thenReturn(true) + when(tab.getMLCacheStatus("userId", "sessionId")).thenReturn( + Some( + MLCacheStatus( + memoryControlEnabled = true, + cachedObjectCount = 2, + inMemorySizeBytes = 1024, + maxInMemorySizeBytes = 4096, + totalSizeBytes = 2048, + maxTotalSizeBytes = 8192))) + + val page = new SparkConnectServerPage(tab) + val html = page.render(request).toString().toLowerCase(Locale.ROOT) + + assert(html.contains("ml cache")) + assert(html.contains("2 objects in memory")) + assert(html.contains("1.0 kib / 4.0 kib memory")) + assert(html.contains("2.0 kib / 8.0 kib total")) + } + test("SPARK-58097: session page only shows the requested user's operations") { // Two users share the same session UUID, each running a distinct query. kvstore = new ElementTrackingStore(new InMemoryStore, new SparkConf()) From 4725d7fbc6ae8acf39acbaa09699856c1dcd3863 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 09:34:23 +0000 Subject: [PATCH 02/10] [CONNECT] Add ML cache model statistics to Spark UI --- .../apache/spark/sql/connect/ml/MLCache.scala | 92 +++++-- .../service/SparkConnectSessionManager.scala | 7 +- .../connect/ui/SparkConnectServerPage.scala | 225 +++++++++++++++--- .../connect/ui/SparkConnectServerTab.scala | 6 +- .../spark/sql/connect/ui/ToolTips.scala | 10 +- .../apache/spark/sql/connect/ml/MLSuite.scala | 25 +- .../ui/SparkConnectServerPageSuite.scala | 46 +++- 7 files changed, 342 insertions(+), 69 deletions(-) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala index 62b729091ac42..b8993ed58035f 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala @@ -74,14 +74,29 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { Connect.CONNECT_SESSION_CONNECT_ML_CACHE_MEMORY_CONTROL_OFFLOADING_TIMEOUT) } + private case class ModelMetadata( + uid: String, + className: String, + modelString: String, + estimatedSizeBytes: Option[Long]) + + // Keep lightweight metadata after a model is evicted from memory so the UI can report + // offloaded models without loading them back into memory. + private val cachedModelMetadata = new ConcurrentHashMap[String, ModelMetadata]() + private val inMemoryModelIds = ConcurrentHashMap.newKeySet[String]() + private[ml] case class CacheItem(obj: Object, sizeBytes: Long) private[ml] val cachedModel: ConcurrentMap[String, CacheItem] = { if (getMemoryControlEnabled) { CacheBuilder .newBuilder() .softValues() - .removalListener((removed: RemovalNotification[String, CacheItem]) => - totalMLCacheInMemorySizeBytes.addAndGet(-removed.getValue.sizeBytes)) + .removalListener((removed: RemovalNotification[String, CacheItem]) => { + Option(removed.getValue).foreach { value => + totalMLCacheInMemorySizeBytes.addAndGet(-value.sizeBytes) + } + inMemoryModelIds.remove(removed.getKey) + }) .maximumWeight(getMaxInMemoryCacheSizeKB) .weigher((key: String, value: CacheItem) => { Math.ceil(value.sizeBytes.toDouble / 1024).toInt @@ -142,24 +157,32 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { if (obj.isInstanceOf[Summary]) { cachedModel.put(objectId, CacheItem(obj, 0)) } else if (obj.isInstanceOf[Model[_]]) { - val sizeBytes = if (getMemoryControlEnabled) { - val _sizeBytes = estimateObjectSize(obj) - checkModelSize(_sizeBytes) - _sizeBytes + val model = obj.asInstanceOf[Model[_]] + val estimatedSizeBytes = if (getMemoryControlEnabled) { + val sizeBytes = estimateObjectSize(model) + checkModelSize(sizeBytes) + Some(sizeBytes) } else { - 0L // Don't need to calculate size if disables memory-control. + // Avoid adding model-size estimation overhead when memory control is disabled. + None } - cachedModel.put(objectId, CacheItem(obj, sizeBytes)) if (getMemoryControlEnabled) { val savePath = getModelOffloadingPath(objectId) - obj.asInstanceOf[MLWritable].write.saveToLocal(savePath.toString) - if (obj.isInstanceOf[HasTrainingSummary[_]] - && obj.asInstanceOf[HasTrainingSummary[_]].hasSummary) { - obj + model.asInstanceOf[MLWritable].write.saveToLocal(savePath.toString) + if (model.isInstanceOf[HasTrainingSummary[_]] + && model.asInstanceOf[HasTrainingSummary[_]].hasSummary) { + model .asInstanceOf[HasTrainingSummary[_]] .saveSummary(savePath.resolve("summary").toString) } - Files.writeString(savePath.resolve(modelClassNameFile), obj.getClass.getName) + Files.writeString(savePath.resolve(modelClassNameFile), model.getClass.getName) + } + cachedModelMetadata.put( + objectId, + ModelMetadata(model.uid, model.getClass.getName, model.toString, estimatedSizeBytes)) + inMemoryModelIds.add(objectId) + cachedModel.put(objectId, CacheItem(model, estimatedSizeBytes.getOrElse(0L))) + estimatedSizeBytes.foreach { sizeBytes => totalMLCacheInMemorySizeBytes.addAndGet(sizeBytes) totalMLCacheSizeBytes.addAndGet(sizeBytes) } @@ -219,6 +242,7 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { loadPath.toString, loadFromLocal = true) val sizeBytes = estimateObjectSize(obj) + inMemoryModelIds.add(refId) cachedModel.put(refId, CacheItem(obj, sizeBytes)) totalMLCacheInMemorySizeBytes.addAndGet(sizeBytes) } @@ -231,8 +255,12 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { verifyObjectId(refId) val removedModel = cachedModel.remove(refId) val removedFromMem = removedModel != null - val removedFromDisk = if (!evictOnly && removedModel != null && getMemoryControlEnabled) { - totalMLCacheSizeBytes.addAndGet(-removedModel.sizeBytes) + inMemoryModelIds.remove(refId) + val metadata = Option(cachedModelMetadata.get(refId)) + val removedFromDisk = if (!evictOnly && metadata.nonEmpty && getMemoryControlEnabled) { + metadata.get.estimatedSizeBytes.foreach { sizeBytes => + totalMLCacheSizeBytes.addAndGet(-sizeBytes) + } val removePath = getModelOffloadingPath(refId) val offloadingPath = new File(removePath.toString) if (offloadingPath.exists()) { @@ -244,7 +272,9 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { } else { false } - removedFromMem || removedFromDisk + val removeMetadata = !evictOnly || !getMemoryControlEnabled + val removedMetadata = removeMetadata && cachedModelMetadata.remove(refId) != null + removedFromMem || removedFromDisk || removedMetadata } /** @@ -264,6 +294,9 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { def clear(): Int = this.synchronized { val size = cachedModel.size() cachedModel.clear() + cachedModelMetadata.clear() + inMemoryModelIds.clear() + totalMLCacheInMemorySizeBytes.set(0) totalMLCacheSizeBytes.set(0) if (getMemoryControlEnabled) { SparkFileUtils.cleanDirectory(new File(offloadedModelsDir.toString)) @@ -281,22 +314,41 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { info.result() } + /** Returns a cache snapshot without loading or touching any cached model. */ def getStatus: MLCacheStatus = this.synchronized { + val models = mutable.ArrayBuilder.make[MLCacheModelInfo] + cachedModelMetadata.forEach { case (id, metadata) => + models += MLCacheModelInfo( + id = id, + uid = metadata.uid, + className = metadata.className, + modelString = metadata.modelString, + estimatedSizeBytes = metadata.estimatedSizeBytes, + inMemory = inMemoryModelIds.contains(id)) + } MLCacheStatus( memoryControlEnabled = getMemoryControlEnabled, - cachedObjectCount = cachedModel.size(), inMemorySizeBytes = totalMLCacheInMemorySizeBytes.get(), maxInMemorySizeBytes = sessionHolder.session.conf.get( Connect.CONNECT_SESSION_CONNECT_ML_CACHE_MEMORY_CONTROL_MAX_IN_MEMORY_SIZE), totalSizeBytes = totalMLCacheSizeBytes.get(), - maxTotalSizeBytes = getMLCacheMaxSize) + maxTotalSizeBytes = getMLCacheMaxSize, + models = models.result().sortBy(_.id)) } } +private[connect] case class MLCacheModelInfo( + id: String, + uid: String, + className: String, + modelString: String, + estimatedSizeBytes: Option[Long], + inMemory: Boolean) + private[connect] case class MLCacheStatus( memoryControlEnabled: Boolean, - cachedObjectCount: Int, inMemorySizeBytes: Long, maxInMemorySizeBytes: Long, totalSizeBytes: Long, - maxTotalSizeBytes: Long) + maxTotalSizeBytes: Long, + models: Seq[MLCacheModelInfo]) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectSessionManager.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectSessionManager.scala index 615d859839427..cf1f188342cb3 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectSessionManager.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectSessionManager.scala @@ -285,8 +285,11 @@ class SparkConnectSessionManager extends Logging { closedSessionsCache.asMap.asScala.values.toSeq } - private[connect] def getMLCacheStatus(key: SessionKey): Option[MLCacheStatus] = { - Option(sessionStore.get(key)).flatMap(_.getMLCacheStatus) + // Read live cache state directly without updating the sessions' last-access times. + private[connect] def getMLCacheStatuses: Seq[(SessionKey, MLCacheStatus)] = { + sessionStore.entrySet().asScala.flatMap { entry => + entry.getValue.getMLCacheStatus.map(entry.getKey -> _) + }.toSeq.sortBy { case (key, _) => (key.userId, key.sessionId) } } /** diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala index 4f418c09523ee..46f3744e7fe56 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala @@ -26,6 +26,7 @@ import scala.xml.Node import jakarta.servlet.http.HttpServletRequest import org.apache.spark.internal.Logging +import org.apache.spark.sql.connect.ml.MLCacheModelInfo import org.apache.spark.sql.connect.ui.ToolTips._ import org.apache.spark.ui._ import org.apache.spark.ui.UIUtils._ @@ -63,6 +64,7 @@ private[ui] class SparkConnectServerPage(parent: SparkConnectServerTab) Request(s) ++ generateSessionStatsTable(request) ++ + generateMLCacheStatsTable(request) ++ generateSQLStatsTable(request) } UIUtils.headerSparkPage(request, "Spark Connect", content, parent) @@ -179,6 +181,83 @@ private[ui] class SparkConnectServerPage(parent: SparkConnectServerTab) content } + + /** Generate live ML cache statistics for active Spark Connect sessions. */ + private def generateMLCacheStatsTable(request: HttpServletRequest): Seq[Node] = { + val cacheStatuses = parent.getMLCacheStatuses.filter(_._2.models.nonEmpty) + val models = cacheStatuses.flatMap { case (key, status) => + status.models.map(MLCacheModelTableRow(key.userId, key.sessionId, _)) + } + if (models.isEmpty) { + return Seq.empty + } + + val tableTag = "mlcachemodels" + val tablePage = Option(request.getParameter(s"$tableTag.page")).map(_.toInt).getOrElse(1) + val table = try { + new MLCacheModelStatsPagedTable( + request, + parent, + models, + "connect", + UIUtils.prependBaseUri(request, parent.basePath), + tableTag).table(tablePage) + } catch { + case e @ (_: IllegalArgumentException | _: IndexOutOfBoundsException) => +
+

Error while rendering ML cache table:

+
+            {Utils.exceptionString(e)}
+          
+
+ } + + val inMemoryModels = models.count(_.model.inMemory) + val memoryControlledStatuses = cacheStatuses.map(_._2).filter(_.memoryControlEnabled) + val inMemorySize = memoryControlledStatuses.map(s => BigInt(s.inMemorySizeBytes)).sum + val maxInMemorySize = memoryControlledStatuses.map(s => BigInt(s.maxInMemorySizeBytes)).sum + val totalSize = memoryControlledStatuses.map(s => BigInt(s.totalSizeBytes)).sum + val maxTotalSize = memoryControlledStatuses.map(s => BigInt(s.maxTotalSizeBytes)).sum + val sizeStats = if (memoryControlledStatuses.nonEmpty) { + Seq( +
  • + In-memory size: + {Utils.bytesToString(inMemorySize)} / {Utils.bytesToString(maxInMemorySize)} +
  • , +
  • + Total estimated size: + {Utils.bytesToString(totalSize)} / {Utils.bytesToString(maxTotalSize)} +
  • ) + } else { + Seq.empty + } + + +

    + + ML Cache Statistics ({models.size}) +

    +
    ++ +
    +
      +
    • Sessions with cached models: {cacheStatuses.size}
    • +
    • + Cached models: + {models.size} ({inMemoryModels} in memory, {models.size - inMemoryModels} offloaded) +
    • +
    • + Memory control enabled sessions: + {memoryControlledStatuses.size} / {cacheStatuses.size} +
    • + {sizeStats} +
    +
    Cached Models
    + {table} +
    + } } private[ui] class SqlStatsPagedTable( @@ -404,8 +483,7 @@ private[ui] class SessionStatsPagedTable( ("Start Time", true, None), ("Finish Time", true, None), ("Duration", true, Some(SPARK_CONNECT_SESSION_DURATION)), - ("Total Execute", true, Some(SPARK_CONNECT_SESSION_TOTAL_EXECUTE)), - ("ML Cache", false, Some(SPARK_CONNECT_SESSION_ML_CACHE))) + ("Total Execute", true, Some(SPARK_CONNECT_SESSION_TOTAL_EXECUTE))) isSortColumnValid(sessionTableHeadersAndTooltips, sortColumn) @@ -432,36 +510,127 @@ private[ui] class SessionStatsPagedTable( {if (session.finishTimestamp > 0) formatDate(session.finishTimestamp)} {formatDurationVerbose(session.totalTime)} {session.totalExecution.toString} - {renderMLCacheStatus(session)} } +} - private def renderMLCacheStatus(session: SessionInfo): Seq[Node] = { - if (!parent.hasLiveMLCacheStatus || session.finishTimestamp > 0) { - N/A - } else { - parent - .getMLCacheStatus(session.userId, session.sessionId) - .map { status => - if (status.memoryControlEnabled) { - val objectLabel = if (status.cachedObjectCount == 1) "object" else "objects" - - {s"${status.cachedObjectCount} $objectLabel in memory"}
    - {s"${Utils.bytesToString(status.inMemorySizeBytes)} / " + - s"${Utils.bytesToString(status.maxInMemorySizeBytes)} memory"}
    - {s"${Utils.bytesToString(status.totalSizeBytes)} / " + - s"${Utils.bytesToString(status.maxTotalSizeBytes)} total"} -
    - } else { - val objectLabel = if (status.cachedObjectCount == 1) "object" else "objects" - - Memory control disabled
    - {s"${status.cachedObjectCount} cached $objectLabel"} -
    - } - } - .getOrElse(Not used) +private[ui] case class MLCacheModelTableRow( + userId: String, + sessionId: String, + model: MLCacheModelInfo) + +private[ui] class MLCacheModelStatsPagedTable( + request: HttpServletRequest, + parent: SparkConnectServerTab, + data: Seq[MLCacheModelTableRow], + subPath: String, + basePath: String, + tableTag: String) + extends PagedTable[MLCacheModelTableRow] { + + private val (sortColumn, desc, pageSize) = + getTableParameters(request, tableTag, "Estimated Size") + + private val encodedSortColumn = URLEncoder.encode(sortColumn, UTF_8.name()) + private val parameterPath = s"$basePath/$subPath/?${getParameterOtherTable(request, tableTag)}" + + override val dataSource = + new MLCacheModelTableDataSource(data, pageSize, sortColumn, desc) + + override def tableId: String = tableTag + + override def tableCssClass: String = + "table table-bordered table-sm table-striped table-head-clickable table-cell-width-limited" + + override def pageLink(page: Int): String = { + parameterPath + + s"&$pageNumberFormField=$page" + + s"&$tableTag.sort=$encodedSortColumn" + + s"&$tableTag.desc=$desc" + + s"&$pageSizeFormField=$pageSize" + + s"#$tableTag" + } + + override def pageSizeFormField: String = s"$tableTag.pageSize" + + override def pageNumberFormField: String = s"$tableTag.page" + + override def goButtonFormPath: String = + s"$parameterPath&$tableTag.sort=$encodedSortColumn" + + s"&$tableTag.desc=$desc#$tableTag" + + override def headers: Seq[Node] = { + val headersAndTooltips: Seq[(String, Boolean, Option[String])] = Seq( + ("User", true, None), + ("Session ID", true, None), + ("Model ID", true, None), + ("UID", true, None), + ("Model Class", true, None), + ("Model Details", true, Some(SPARK_CONNECT_ML_CACHE_MODEL_DETAILS)), + ("Estimated Size", true, Some(SPARK_CONNECT_ML_CACHE_ESTIMATED_SIZE)), + ("Storage", true, Some(SPARK_CONNECT_ML_CACHE_STORAGE))) + + isSortColumnValid(headersAndTooltips, sortColumn) + headerRow( + headersAndTooltips, + desc, + pageSize, + sortColumn, + parameterPath, + tableTag, + tableTag) + } + + override def row(row: MLCacheModelTableRow): Seq[Node] = { + val model = row.model + val sessionLink = "%s/%s/session/?id=%s&userId=%s".format( + UIUtils.prependBaseUri(request, parent.basePath), + parent.prefix, + URLEncoder.encode(row.sessionId, UTF_8.name()), + ConnectUiUtils.encodeUserId(row.userId)) + + {row.userId} + {row.sessionId} + {model.id} + {model.uid} + {model.className} + {model.modelString} + {model.estimatedSizeBytes.map(Utils.bytesToString).getOrElse("N/A")} + {if (model.inMemory) "In memory" else "Offloaded"} + + } +} + +private[ui] class MLCacheModelTableDataSource( + info: Seq[MLCacheModelTableRow], + pageSize: Int, + sortColumn: String, + desc: Boolean) + extends PagedDataSource[MLCacheModelTableRow](pageSize) { + + private val data = info.sorted(ordering(sortColumn, desc)) + + override def dataSize: Int = data.size + + override def sliceData(from: Int, to: Int): Seq[MLCacheModelTableRow] = data.slice(from, to) + + private def ordering( + sortColumn: String, + desc: Boolean): Ordering[MLCacheModelTableRow] = { + val ordering: Ordering[MLCacheModelTableRow] = sortColumn match { + case "User" => Ordering.by(_.userId) + case "Session ID" => Ordering.by(_.sessionId) + case "Model ID" => Ordering.by(_.model.id) + case "UID" => Ordering.by(_.model.uid) + case "Model Class" => Ordering.by(_.model.className) + case "Model Details" => Ordering.by(_.model.modelString) + case "Estimated Size" => + Ordering.by((row: MLCacheModelTableRow) => + (row.model.estimatedSizeBytes.isDefined, row.model.estimatedSizeBytes.getOrElse(0L))) + case "Storage" => Ordering.by(_.model.inMemory) + case unknownColumn => throw new IllegalArgumentException(s"Unknown column: $unknownColumn") } + if (desc) ordering.reverse else ordering } } diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerTab.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerTab.scala index 181b9fa8e95b6..d720acb31fee2 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerTab.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerTab.scala @@ -50,12 +50,10 @@ private[connect] class SparkConnectServerTab( parent.detachTab(this) } - def getMLCacheStatus(userId: String, sessionId: String): Option[MLCacheStatus] = { - sessionManager.flatMap(_.getMLCacheStatus(SessionKey(userId, sessionId))) + def getMLCacheStatuses: Seq[(SessionKey, MLCacheStatus)] = { + sessionManager.toSeq.flatMap(_.getMLCacheStatuses) } - def hasLiveMLCacheStatus: Boolean = sessionManager.nonEmpty - override def displayOrder: Int = 3 } diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/ToolTips.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/ToolTips.scala index ff385db60af5b..be2ab4c4903b2 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/ToolTips.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/ToolTips.scala @@ -36,7 +36,13 @@ private[ui] object ToolTips { val SPARK_CONNECT_SESSION_DURATION = "Elapsed time since session start, or until closed if the session was closed" - val SPARK_CONNECT_SESSION_ML_CACHE = - "Current ML cache usage for an active Spark Connect session" + val SPARK_CONNECT_ML_CACHE_ESTIMATED_SIZE = + "Approximate model size recorded when it was added to the Spark Connect ML cache" + + val SPARK_CONNECT_ML_CACHE_MODEL_DETAILS = + "Output of model.toString recorded when the model was added to the cache" + + val SPARK_CONNECT_ML_CACHE_STORAGE = + "Whether the model is currently in driver memory or offloaded to driver-local disk" } diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ml/MLSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ml/MLSuite.scala index c7bc9c888cb98..4dd21bd6069eb 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ml/MLSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ml/MLSuite.scala @@ -401,14 +401,20 @@ class MLSuite extends MLHelper { // Reading UI status should not initialize an unused cache. assert(sessionHolder.getMLCacheStatus.isEmpty) - trainLogisticRegressionModel(sessionHolder) + val modelId = trainLogisticRegressionModel(sessionHolder) val status = sessionHolder.getMLCacheStatus.get assert(status.memoryControlEnabled) - assert(status.cachedObjectCount === 1) assert(status.inMemorySizeBytes > 0) assert(status.maxInMemorySizeBytes === 16384) assert(status.totalSizeBytes === status.inMemorySizeBytes) assert(status.maxTotalSizeBytes === 65536) + assert(status.models.size === 1) + val modelInfo = status.models.head + assert(modelInfo.id === modelId) + assert(modelInfo.className === classOf[LogisticRegressionModel].getName) + assert(modelInfo.modelString.startsWith("LogisticRegressionModel: uid=")) + assert(modelInfo.estimatedSizeBytes.contains(status.totalSizeBytes)) + assert(modelInfo.inMemory) } test("MLCache offloading works") { @@ -444,10 +450,24 @@ class MLSuite extends MLHelper { assert(sessionHolder.mlCache.totalMLCacheInMemorySizeBytes.get() <= memorySizeBytes) } + val status = sessionHolder.getMLCacheStatus.get + assert(status.models.size === modelIdList.size) + assert(status.models.count(_.inMemory) === maxNumModels) + assert(status.models.map(_.id).toSet === modelIdList.toSet) + // Assert all models can be loaded back from disk after they are offloaded. for (modelId <- modelIdList) { assert(sessionHolder.mlCache.get(modelId) != null) } + + val statusBeforeRemove = sessionHolder.getMLCacheStatus.get + val offloadedModel = statusBeforeRemove.models.find(!_.inMemory).get + assert(sessionHolder.mlCache.remove(offloadedModel.id)) + val statusAfterRemove = sessionHolder.getMLCacheStatus.get + assert(!statusAfterRemove.models.exists(_.id == offloadedModel.id)) + assert( + statusAfterRemove.totalSizeBytes === + statusBeforeRemove.totalSizeBytes - offloadedModel.estimatedSizeBytes.get) } test("Model size limit") { @@ -485,6 +505,7 @@ class MLSuite extends MLHelper { assert(mlCache2.get(modelId) != null) mlCache2.close() assert(mlCache2.cachedModel.isEmpty) + assert(mlCache2.getStatus.models.isEmpty) // Test 3: Edge case - register then remove model, close should still run cleanup val edgeCaseSessionHolder = SparkConnectTestUtils.createDummySessionHolder(spark) diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala index 92693a31fa25d..a10a9d3aca047 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala @@ -26,7 +26,7 @@ import org.scalatest.BeforeAndAfter import org.apache.spark.{SharedSparkContext, SparkConf, SparkFunSuite} import org.apache.spark.scheduler.SparkListenerJobStart -import org.apache.spark.sql.connect.ml.MLCacheStatus +import org.apache.spark.sql.connect.ml.{MLCacheModelInfo, MLCacheStatus} import org.apache.spark.sql.connect.service._ import org.apache.spark.status.ElementTrackingStore import org.apache.spark.util.kvstore.InMemoryStore @@ -93,6 +93,7 @@ class SparkConnectServerPageSuite when(tab.store).thenReturn(store) when(tab.appName).thenReturn("testing") when(tab.headerTabs).thenReturn(Seq.empty) + when(tab.getMLCacheStatuses).thenReturn(Seq.empty) val page = new SparkConnectServerPage(tab) val html = page.render(request).toString().toLowerCase(Locale.ROOT) @@ -100,6 +101,7 @@ class SparkConnectServerPageSuite assert(html.contains("session statistics (1)")) assert(html.contains("request statistics (1)")) assert(html.contains("dummy query")) + assert(!html.contains("ml cache statistics")) // Pagination support assert(html.contains("")) @@ -138,7 +140,7 @@ class SparkConnectServerPageSuite " data-bs-target=\"#aggregated-sqlsessionstat\"")) } - test("Spark Connect Server page should show live ML cache status") { + test("Spark Connect Server page should show live ML cache statistics and model details") { val store = getStatusStore(closeSession = false) val request = mock(classOf[HttpServletRequest]) @@ -147,24 +149,46 @@ class SparkConnectServerPageSuite when(tab.store).thenReturn(store) when(tab.appName).thenReturn("testing") when(tab.headerTabs).thenReturn(Seq.empty) - when(tab.hasLiveMLCacheStatus).thenReturn(true) - when(tab.getMLCacheStatus("userId", "sessionId")).thenReturn( - Some( + when(tab.getMLCacheStatuses).thenReturn( + Seq( + SessionKey("userId", "sessionId") -> MLCacheStatus( memoryControlEnabled = true, - cachedObjectCount = 2, inMemorySizeBytes = 1024, maxInMemorySizeBytes = 4096, totalSizeBytes = 2048, - maxTotalSizeBytes = 8192))) + maxTotalSizeBytes = 8192, + models = Seq( + MLCacheModelInfo( + id = "model-id-1", + uid = "logreg-1", + className = "org.apache.spark.ml.classification.LogisticRegressionModel", + modelString = "LogisticRegressionModel: uid=logreg-1", + estimatedSizeBytes = Some(1024), + inMemory = true), + MLCacheModelInfo( + id = "model-id-2", + uid = "logreg-2", + className = "org.apache.spark.ml.classification.LogisticRegressionModel", + modelString = "LogisticRegressionModel: uid=logreg-2", + estimatedSizeBytes = Some(1024), + inMemory = false))))) val page = new SparkConnectServerPage(tab) val html = page.render(request).toString().toLowerCase(Locale.ROOT) - assert(html.contains("ml cache")) - assert(html.contains("2 objects in memory")) - assert(html.contains("1.0 kib / 4.0 kib memory")) - assert(html.contains("2.0 kib / 8.0 kib total")) + val sessionStatsIndex = html.indexOf("session statistics") + val mlCacheStatsIndex = html.indexOf("ml cache statistics (2)") + val requestStatsIndex = html.indexOf("request statistics") + assert(sessionStatsIndex < mlCacheStatsIndex && mlCacheStatsIndex < requestStatsIndex) + assert(html.contains("2 (1 in memory, 1 offloaded)")) + assert(html.contains("1.0 kib / 4.0 kib")) + assert(html.contains("2.0 kib / 8.0 kib")) + assert(html.contains("model-id-1")) + assert(html.contains("logreg-1")) + assert(html.contains("logisticregressionmodel: uid=logreg-1")) + assert(html.contains("in memory")) + assert(html.contains("offloaded")) } test("SPARK-58097: session page only shows the requested user's operations") { From 1ecb1500e07bf4eb013b532ef5805bce6d38d83f Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 09:54:01 +0000 Subject: [PATCH 03/10] [CONNECT] Fix ML cache UI compilation --- .../main/scala/org/apache/spark/sql/connect/ml/MLCache.scala | 5 +++-- .../spark/sql/connect/ui/SparkConnectServerPageSuite.scala | 4 ++-- 2 files changed, 5 insertions(+), 4 deletions(-) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala index b8993ed58035f..29790081aeb96 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala @@ -23,6 +23,7 @@ import java.util.concurrent.{ConcurrentHashMap, ConcurrentMap, TimeUnit} import java.util.concurrent.atomic.{AtomicBoolean, AtomicLong} import scala.collection.mutable +import scala.jdk.CollectionConverters._ import scala.util.control.NonFatal import com.google.common.cache.{CacheBuilder, RemovalNotification} @@ -317,7 +318,7 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { /** Returns a cache snapshot without loading or touching any cached model. */ def getStatus: MLCacheStatus = this.synchronized { val models = mutable.ArrayBuilder.make[MLCacheModelInfo] - cachedModelMetadata.forEach { case (id, metadata) => + cachedModelMetadata.asScala.foreach { case (id, metadata) => models += MLCacheModelInfo( id = id, uid = metadata.uid, @@ -333,7 +334,7 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { Connect.CONNECT_SESSION_CONNECT_ML_CACHE_MEMORY_CONTROL_MAX_IN_MEMORY_SIZE), totalSizeBytes = totalMLCacheSizeBytes.get(), maxTotalSizeBytes = getMLCacheMaxSize, - models = models.result().sortBy(_.id)) + models = models.result().toIndexedSeq.sortBy(_.id)) } } diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala index a10a9d3aca047..6d81e5656c2d7 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala @@ -85,7 +85,7 @@ class SparkConnectServerPageSuite } test("Spark Connect Server page should load successfully") { - val store = getStatusStore + val store = getStatusStore() val request = mock(classOf[HttpServletRequest]) val tab = mock(classOf[SparkConnectServerTab], RETURNS_SMART_NULLS) @@ -113,7 +113,7 @@ class SparkConnectServerPageSuite } test("Spark Connect Server session page should load successfully") { - val store = getStatusStore + val store = getStatusStore() val request = mock(classOf[HttpServletRequest]) when(request.getParameter("id")).thenReturn("sessionId") From afc9a7172590f1f08cb8f209f420e7c204babdce Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 10:08:10 +0000 Subject: [PATCH 04/10] [CONNECT] Remove redundant ML cache model UID --- .../main/scala/org/apache/spark/sql/connect/ml/MLCache.scala | 5 +---- .../apache/spark/sql/connect/ui/SparkConnectServerPage.scala | 3 --- .../spark/sql/connect/ui/SparkConnectServerPageSuite.scala | 3 --- 3 files changed, 1 insertion(+), 10 deletions(-) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala index 29790081aeb96..5c3db4b99e307 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ml/MLCache.scala @@ -76,7 +76,6 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { } private case class ModelMetadata( - uid: String, className: String, modelString: String, estimatedSizeBytes: Option[Long]) @@ -180,7 +179,7 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { } cachedModelMetadata.put( objectId, - ModelMetadata(model.uid, model.getClass.getName, model.toString, estimatedSizeBytes)) + ModelMetadata(model.getClass.getName, model.toString, estimatedSizeBytes)) inMemoryModelIds.add(objectId) cachedModel.put(objectId, CacheItem(model, estimatedSizeBytes.getOrElse(0L))) estimatedSizeBytes.foreach { sizeBytes => @@ -321,7 +320,6 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { cachedModelMetadata.asScala.foreach { case (id, metadata) => models += MLCacheModelInfo( id = id, - uid = metadata.uid, className = metadata.className, modelString = metadata.modelString, estimatedSizeBytes = metadata.estimatedSizeBytes, @@ -340,7 +338,6 @@ private[connect] class MLCache(sessionHolder: SessionHolder) extends Logging { private[connect] case class MLCacheModelInfo( id: String, - uid: String, className: String, modelString: String, estimatedSizeBytes: Option[Long], diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala index 46f3744e7fe56..dcc241f7b93de 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala @@ -564,7 +564,6 @@ private[ui] class MLCacheModelStatsPagedTable( ("User", true, None), ("Session ID", true, None), ("Model ID", true, None), - ("UID", true, None), ("Model Class", true, None), ("Model Details", true, Some(SPARK_CONNECT_ML_CACHE_MODEL_DETAILS)), ("Estimated Size", true, Some(SPARK_CONNECT_ML_CACHE_ESTIMATED_SIZE)), @@ -592,7 +591,6 @@ private[ui] class MLCacheModelStatsPagedTable( {row.userId} {row.sessionId} {model.id} - {model.uid} {model.className} {model.modelString} {model.estimatedSizeBytes.map(Utils.bytesToString).getOrElse("N/A")} @@ -621,7 +619,6 @@ private[ui] class MLCacheModelTableDataSource( case "User" => Ordering.by(_.userId) case "Session ID" => Ordering.by(_.sessionId) case "Model ID" => Ordering.by(_.model.id) - case "UID" => Ordering.by(_.model.uid) case "Model Class" => Ordering.by(_.model.className) case "Model Details" => Ordering.by(_.model.modelString) case "Estimated Size" => diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala index 6d81e5656c2d7..826ce645e7ed2 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala @@ -161,14 +161,12 @@ class SparkConnectServerPageSuite models = Seq( MLCacheModelInfo( id = "model-id-1", - uid = "logreg-1", className = "org.apache.spark.ml.classification.LogisticRegressionModel", modelString = "LogisticRegressionModel: uid=logreg-1", estimatedSizeBytes = Some(1024), inMemory = true), MLCacheModelInfo( id = "model-id-2", - uid = "logreg-2", className = "org.apache.spark.ml.classification.LogisticRegressionModel", modelString = "LogisticRegressionModel: uid=logreg-2", estimatedSizeBytes = Some(1024), @@ -185,7 +183,6 @@ class SparkConnectServerPageSuite assert(html.contains("1.0 kib / 4.0 kib")) assert(html.contains("2.0 kib / 8.0 kib")) assert(html.contains("model-id-1")) - assert(html.contains("logreg-1")) assert(html.contains("logisticregressionmodel: uid=logreg-1")) assert(html.contains("in memory")) assert(html.contains("offloaded")) From ec810fe2836d5376922a50badc05bb063a9ac696 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 10:22:08 +0000 Subject: [PATCH 05/10] [CONNECT] Move ML cache statistics after requests --- .../apache/spark/sql/connect/ui/SparkConnectServerPage.scala | 4 ++-- .../spark/sql/connect/ui/SparkConnectServerPageSuite.scala | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala index dcc241f7b93de..10193b2609343 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala @@ -64,8 +64,8 @@ private[ui] class SparkConnectServerPage(parent: SparkConnectServerTab) Request(s) ++ generateSessionStatsTable(request) ++ - generateMLCacheStatsTable(request) ++ - generateSQLStatsTable(request) + generateSQLStatsTable(request) ++ + generateMLCacheStatsTable(request) } UIUtils.headerSparkPage(request, "Spark Connect", content, parent) } diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala index 826ce645e7ed2..36f7b2da41e02 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala @@ -178,7 +178,7 @@ class SparkConnectServerPageSuite val sessionStatsIndex = html.indexOf("session statistics") val mlCacheStatsIndex = html.indexOf("ml cache statistics (2)") val requestStatsIndex = html.indexOf("request statistics") - assert(sessionStatsIndex < mlCacheStatsIndex && mlCacheStatsIndex < requestStatsIndex) + assert(sessionStatsIndex < requestStatsIndex && requestStatsIndex < mlCacheStatsIndex) assert(html.contains("2 (1 in memory, 1 offloaded)")) assert(html.contains("1.0 kib / 4.0 kib")) assert(html.contains("2.0 kib / 8.0 kib")) From 175b5c3276a74a8c554b7fd45c8909f9b45d5903 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 10:24:46 +0000 Subject: [PATCH 06/10] [CONNECT] Simplify ML cache summary --- .../apache/spark/sql/connect/ui/SparkConnectServerPage.scala | 4 ---- 1 file changed, 4 deletions(-) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala index 10193b2609343..95df3256fd6c8 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala @@ -248,10 +248,6 @@ private[ui] class SparkConnectServerPage(parent: SparkConnectServerTab) Cached models: {models.size} ({inMemoryModels} in memory, {models.size - inMemoryModels} offloaded) -
  • - Memory control enabled sessions: - {memoryControlledStatuses.size} / {cacheStatuses.size} -
  • {sizeStats}
    Cached Models
    From 02057aa7c96005c7d3b7d53a5759ef634ec82a8f Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 10:29:24 +0000 Subject: [PATCH 07/10] [CONNECT] Clarify total ML cache size --- .../apache/spark/sql/connect/ui/SparkConnectServerPage.scala | 2 +- .../spark/sql/connect/ui/SparkConnectServerPageSuite.scala | 1 + 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala index 95df3256fd6c8..0d4c9e48f09a3 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala @@ -225,7 +225,7 @@ private[ui] class SparkConnectServerPage(parent: SparkConnectServerTab) {Utils.bytesToString(inMemorySize)} / {Utils.bytesToString(maxInMemorySize)} ,
  • - Total estimated size: + Total estimated size (in-memory and offloaded data): {Utils.bytesToString(totalSize)} / {Utils.bytesToString(maxTotalSize)}
  • ) } else { diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala index 36f7b2da41e02..f0cae771e6ca0 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala @@ -181,6 +181,7 @@ class SparkConnectServerPageSuite assert(sessionStatsIndex < requestStatsIndex && requestStatsIndex < mlCacheStatsIndex) assert(html.contains("2 (1 in memory, 1 offloaded)")) assert(html.contains("1.0 kib / 4.0 kib")) + assert(html.contains("total estimated size (in-memory and offloaded data)")) assert(html.contains("2.0 kib / 8.0 kib")) assert(html.contains("model-id-1")) assert(html.contains("logisticregressionmodel: uid=logreg-1")) From 371390457d53dfca6d88e6fde2805fb2dd44c5d5 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 10:31:52 +0000 Subject: [PATCH 08/10] [CONNECT] Clarify ML cache size labels --- .../apache/spark/sql/connect/ui/SparkConnectServerPage.scala | 4 ++-- .../spark/sql/connect/ui/SparkConnectServerPageSuite.scala | 3 ++- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala index 0d4c9e48f09a3..2d65b1a72b358 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala @@ -221,11 +221,11 @@ private[ui] class SparkConnectServerPage(parent: SparkConnectServerTab) val sizeStats = if (memoryControlledStatuses.nonEmpty) { Seq(
  • - In-memory size: + Estimated size (In-memory): {Utils.bytesToString(inMemorySize)} / {Utils.bytesToString(maxInMemorySize)}
  • ,
  • - Total estimated size (in-memory and offloaded data): + Estimated size (In-memory + Disk): {Utils.bytesToString(totalSize)} / {Utils.bytesToString(maxTotalSize)}
  • ) } else { diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala index f0cae771e6ca0..ef68ef14b9514 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala @@ -180,8 +180,9 @@ class SparkConnectServerPageSuite val requestStatsIndex = html.indexOf("request statistics") assert(sessionStatsIndex < requestStatsIndex && requestStatsIndex < mlCacheStatsIndex) assert(html.contains("2 (1 in memory, 1 offloaded)")) + assert(html.contains("estimated size (in-memory)")) assert(html.contains("1.0 kib / 4.0 kib")) - assert(html.contains("total estimated size (in-memory and offloaded data)")) + assert(html.contains("estimated size (in-memory + disk)")) assert(html.contains("2.0 kib / 8.0 kib")) assert(html.contains("model-id-1")) assert(html.contains("logisticregressionmodel: uid=logreg-1")) From f44cb985b5d8f7897895520b15f5f8fe6716671c Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 10:47:23 +0000 Subject: [PATCH 09/10] [CONNECT] Clarify offloaded ML cache size label --- .../apache/spark/sql/connect/ui/SparkConnectServerPage.scala | 2 +- .../spark/sql/connect/ui/SparkConnectServerPageSuite.scala | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala index 2d65b1a72b358..34d1d5ae7aece 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala @@ -225,7 +225,7 @@ private[ui] class SparkConnectServerPage(parent: SparkConnectServerTab) {Utils.bytesToString(inMemorySize)} / {Utils.bytesToString(maxInMemorySize)} ,
  • - Estimated size (In-memory + Disk): + Estimated size (In-memory and Offloaded data): {Utils.bytesToString(totalSize)} / {Utils.bytesToString(maxTotalSize)}
  • ) } else { diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala index ef68ef14b9514..060987f8123b6 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala @@ -182,7 +182,7 @@ class SparkConnectServerPageSuite assert(html.contains("2 (1 in memory, 1 offloaded)")) assert(html.contains("estimated size (in-memory)")) assert(html.contains("1.0 kib / 4.0 kib")) - assert(html.contains("estimated size (in-memory + disk)")) + assert(html.contains("estimated size (in-memory and offloaded data)")) assert(html.contains("2.0 kib / 8.0 kib")) assert(html.contains("model-id-1")) assert(html.contains("logisticregressionmodel: uid=logreg-1")) From 878f460536e1bdfe4883d747ca5cab7f36be4a94 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 12:55:28 +0000 Subject: [PATCH 10/10] [SPARK-58556][CONNECT][UI] Fix ML cache UI checks --- .../service/SparkConnectSessionManager.scala | 11 +++-- .../connect/ui/SparkConnectServerPage.scala | 38 +++++++--------- .../apache/spark/sql/connect/ml/MLSuite.scala | 10 ++--- .../ui/SparkConnectServerPageSuite.scala | 43 +++++++++---------- 4 files changed, 48 insertions(+), 54 deletions(-) diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectSessionManager.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectSessionManager.scala index cf1f188342cb3..ca124ba9693fa 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectSessionManager.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/service/SparkConnectSessionManager.scala @@ -287,9 +287,14 @@ class SparkConnectSessionManager extends Logging { // Read live cache state directly without updating the sessions' last-access times. private[connect] def getMLCacheStatuses: Seq[(SessionKey, MLCacheStatus)] = { - sessionStore.entrySet().asScala.flatMap { entry => - entry.getValue.getMLCacheStatus.map(entry.getKey -> _) - }.toSeq.sortBy { case (key, _) => (key.userId, key.sessionId) } + sessionStore + .entrySet() + .asScala + .flatMap { entry => + entry.getValue.getMLCacheStatus.map(entry.getKey -> _) + } + .toSeq + .sortBy { case (key, _) => (key.userId, key.sessionId) } } /** diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala index 34d1d5ae7aece..b3af78f62e361 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPage.scala @@ -194,23 +194,24 @@ private[ui] class SparkConnectServerPage(parent: SparkConnectServerTab) val tableTag = "mlcachemodels" val tablePage = Option(request.getParameter(s"$tableTag.page")).map(_.toInt).getOrElse(1) - val table = try { - new MLCacheModelStatsPagedTable( - request, - parent, - models, - "connect", - UIUtils.prependBaseUri(request, parent.basePath), - tableTag).table(tablePage) - } catch { - case e @ (_: IllegalArgumentException | _: IndexOutOfBoundsException) => -
    + val table = + try { + new MLCacheModelStatsPagedTable( + request, + parent, + models, + "connect", + UIUtils.prependBaseUri(request, parent.basePath), + tableTag).table(tablePage) + } catch { + case e @ (_: IllegalArgumentException | _: IndexOutOfBoundsException) => +

    Error while rendering ML cache table:

                 {Utils.exceptionString(e)}
               
    - } + } val inMemoryModels = models.count(_.model.inMemory) val memoryControlledStatuses = cacheStatuses.map(_._2).filter(_.memoryControlEnabled) @@ -566,14 +567,7 @@ private[ui] class MLCacheModelStatsPagedTable( ("Storage", true, Some(SPARK_CONNECT_ML_CACHE_STORAGE))) isSortColumnValid(headersAndTooltips, sortColumn) - headerRow( - headersAndTooltips, - desc, - pageSize, - sortColumn, - parameterPath, - tableTag, - tableTag) + headerRow(headersAndTooltips, desc, pageSize, sortColumn, parameterPath, tableTag, tableTag) } override def row(row: MLCacheModelTableRow): Seq[Node] = { @@ -608,9 +602,7 @@ private[ui] class MLCacheModelTableDataSource( override def sliceData(from: Int, to: Int): Seq[MLCacheModelTableRow] = data.slice(from, to) - private def ordering( - sortColumn: String, - desc: Boolean): Ordering[MLCacheModelTableRow] = { + private def ordering(sortColumn: String, desc: Boolean): Ordering[MLCacheModelTableRow] = { val ordering: Ordering[MLCacheModelTableRow] = sortColumn match { case "User" => Ordering.by(_.userId) case "Session ID" => Ordering.by(_.sessionId) diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ml/MLSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ml/MLSuite.scala index 4dd21bd6069eb..b78423e8687e9 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ml/MLSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ml/MLSuite.scala @@ -391,12 +391,10 @@ class MLSuite extends MLHelper { val sessionHolder = SparkConnectTestUtils.createDummySessionHolder(spark) sessionHolder.session.conf .set(Connect.CONNECT_SESSION_CONNECT_ML_CACHE_MEMORY_CONTROL_ENABLED.key, "true") - sessionHolder.session.conf.set( - Connect.CONNECT_SESSION_CONNECT_ML_CACHE_MEMORY_CONTROL_MAX_IN_MEMORY_SIZE.key, - 16384) - sessionHolder.session.conf.set( - Connect.CONNECT_SESSION_CONNECT_ML_CACHE_MEMORY_CONTROL_MAX_STORAGE_SIZE.key, - 65536) + sessionHolder.session.conf + .set(Connect.CONNECT_SESSION_CONNECT_ML_CACHE_MEMORY_CONTROL_MAX_IN_MEMORY_SIZE.key, 16384) + sessionHolder.session.conf + .set(Connect.CONNECT_SESSION_CONNECT_ML_CACHE_MEMORY_CONTROL_MAX_STORAGE_SIZE.key, 65536) // Reading UI status should not initialize an unused cache. assert(sessionHolder.getMLCacheStatus.isEmpty) diff --git a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala index 060987f8123b6..8def63046cecf 100644 --- a/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala +++ b/sql/connect/server/src/test/scala/org/apache/spark/sql/connect/ui/SparkConnectServerPageSuite.scala @@ -48,8 +48,7 @@ class SparkConnectServerPageSuite /** * Run a dummy session and return the store */ - private def getStatusStore( - closeSession: Boolean = true): SparkConnectServerAppStatusStore = { + private def getStatusStore(closeSession: Boolean = true): SparkConnectServerAppStatusStore = { kvstore = new ElementTrackingStore(new InMemoryStore, new SparkConf()) // val server = mock(classOf[SparkConnectServer], RETURNS_SMART_NULLS) val sparkConf = new SparkConf @@ -152,25 +151,25 @@ class SparkConnectServerPageSuite when(tab.getMLCacheStatuses).thenReturn( Seq( SessionKey("userId", "sessionId") -> - MLCacheStatus( - memoryControlEnabled = true, - inMemorySizeBytes = 1024, - maxInMemorySizeBytes = 4096, - totalSizeBytes = 2048, - maxTotalSizeBytes = 8192, - models = Seq( - MLCacheModelInfo( - id = "model-id-1", - className = "org.apache.spark.ml.classification.LogisticRegressionModel", - modelString = "LogisticRegressionModel: uid=logreg-1", - estimatedSizeBytes = Some(1024), - inMemory = true), - MLCacheModelInfo( - id = "model-id-2", - className = "org.apache.spark.ml.classification.LogisticRegressionModel", - modelString = "LogisticRegressionModel: uid=logreg-2", - estimatedSizeBytes = Some(1024), - inMemory = false))))) + MLCacheStatus( + memoryControlEnabled = true, + inMemorySizeBytes = 1024, + maxInMemorySizeBytes = 4096, + totalSizeBytes = 2048, + maxTotalSizeBytes = 8192, + models = Seq( + MLCacheModelInfo( + id = "model-id-1", + className = "org.apache.spark.ml.classification.LogisticRegressionModel", + modelString = "LogisticRegressionModel: uid=logreg-1", + estimatedSizeBytes = Some(1024), + inMemory = true), + MLCacheModelInfo( + id = "model-id-2", + className = "org.apache.spark.ml.classification.LogisticRegressionModel", + modelString = "LogisticRegressionModel: uid=logreg-2", + estimatedSizeBytes = Some(1024), + inMemory = false))))) val page = new SparkConnectServerPage(tab) val html = page.render(request).toString().toLowerCase(Locale.ROOT) @@ -181,7 +180,7 @@ class SparkConnectServerPageSuite assert(sessionStatsIndex < requestStatsIndex && requestStatsIndex < mlCacheStatsIndex) assert(html.contains("2 (1 in memory, 1 offloaded)")) assert(html.contains("estimated size (in-memory)")) - assert(html.contains("1.0 kib / 4.0 kib")) + assert(html.contains("1024.0 b / 4.0 kib")) assert(html.contains("estimated size (in-memory and offloaded data)")) assert(html.contains("2.0 kib / 8.0 kib")) assert(html.contains("model-id-1"))