This is an automated email from the ASF dual-hosted git repository.

ethanfeng pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/incubator-celeborn.git


The following commit(s) were added to refs/heads/main by this push:
     new 4c5e6f065 [CELEBORN-1182] Support application dimension 
ActiveConnectionCount metric to record the number of registered connections for 
each application
4c5e6f065 is described below

commit 4c5e6f065c96ffc6366b19becabec0c7b7f6e6db
Author: SteNicholas <[email protected]>
AuthorDate: Fri Feb 2 16:29:10 2024 +0800

    [CELEBORN-1182] Support application dimension ActiveConnectionCount metric 
to record the number of registered connections for each application
    
    ### What changes were proposed in this pull request?
    
    `WorkerSource` supports application dimension `ActiveConnectionCount` 
metric to record the number of registered connections for each application.
    
    ### Why are the changes needed?
    
    `ActiveConnectionCount` metric records the number of registered connections 
at present. It's recommended to support dimension ActiveConnectionCount metric 
to record the number of registered connections for each application in Worker. 
Application dimension `ActiveConnectionCount` metric could provide users with 
the actual number of registered connections for each application.
    
    ### Does this PR introduce _any_ user-facing change?
    
    No.
    
    ### How was this patch tested?
    
    Internal tests.
    
    Closes #2167 from SteNicholas/CELEBORN-1182.
    
    Authored-by: SteNicholas <[email protected]>
    Signed-off-by: mingji <[email protected]>
---
 .../common/metrics/source/AbstractSource.scala     |  2 +
 .../metrics/source/ResourceConsumptionSource.scala |  2 -
 .../celeborn/service/deploy/master/Master.scala    |  4 +-
 .../deploy/worker/storage/CreditStreamManager.java | 24 +++++++++--
 .../service/deploy/worker/FetchHandler.scala       | 39 +++++++++++++-----
 .../service/deploy/worker/PushDataHandler.scala    |  7 +++-
 .../celeborn/service/deploy/worker/Worker.scala    | 10 ++++-
 .../service/deploy/worker/WorkerSource.scala       | 47 ++++++++++++++++++++++
 .../worker/storage/CreditStreamManagerSuiteJ.java  | 13 ++++--
 9 files changed, 123 insertions(+), 25 deletions(-)

diff --git 
a/common/src/main/scala/org/apache/celeborn/common/metrics/source/AbstractSource.scala
 
b/common/src/main/scala/org/apache/celeborn/common/metrics/source/AbstractSource.scala
index 8d76414c1..301d991a2 100644
--- 
a/common/src/main/scala/org/apache/celeborn/common/metrics/source/AbstractSource.scala
+++ 
b/common/src/main/scala/org/apache/celeborn/common/metrics/source/AbstractSource.scala
@@ -67,6 +67,8 @@ abstract class AbstractSource(conf: CelebornConf, role: 
String)
   val staticLabels: Map[String, String] = conf.metricsExtraLabels + roleLabel
   val staticLabelsString: String = MetricLabels.labelString(staticLabels)
 
+  val applicationLabel = "applicationId"
+
   protected val namedGauges: JQueue[NamedGauge[_]] = new 
ConcurrentLinkedQueue[NamedGauge[_]]()
 
   def addGauge[T](
diff --git 
a/common/src/main/scala/org/apache/celeborn/common/metrics/source/ResourceConsumptionSource.scala
 
b/common/src/main/scala/org/apache/celeborn/common/metrics/source/ResourceConsumptionSource.scala
index 88a4b9858..df33310bb 100644
--- 
a/common/src/main/scala/org/apache/celeborn/common/metrics/source/ResourceConsumptionSource.scala
+++ 
b/common/src/main/scala/org/apache/celeborn/common/metrics/source/ResourceConsumptionSource.scala
@@ -33,6 +33,4 @@ object ResourceConsumptionSource {
   val HDFS_FILE_COUNT = "hdfsFileCount"
 
   val HDFS_BYTES_WRITTEN = "hdfsBytesWritten"
-
-  val APPLICATION_LABEL = "applicationId"
 }
diff --git 
a/master/src/main/scala/org/apache/celeborn/service/deploy/master/Master.scala 
b/master/src/main/scala/org/apache/celeborn/service/deploy/master/Master.scala
index 79971f22b..b6441d8e0 100644
--- 
a/master/src/main/scala/org/apache/celeborn/service/deploy/master/Master.scala
+++ 
b/master/src/main/scala/org/apache/celeborn/service/deploy/master/Master.scala
@@ -875,7 +875,7 @@ private[celeborn] class Master(
       appId: String): Unit = {
     resourceConsumptionSource.removeGauge(
       resourceConsumptionName,
-      ResourceConsumptionSource.APPLICATION_LABEL,
+      resourceConsumptionSource.applicationLabel,
       appId)
   }
 
@@ -955,7 +955,7 @@ private[celeborn] class Master(
       applicationId: String = null): Unit = {
     val resourceConsumptionLabel =
       if (applicationId == null) userIdentifier.toMap
-      else userIdentifier.toMap + (ResourceConsumptionSource.APPLICATION_LABEL 
-> applicationId)
+      else userIdentifier.toMap + (resourceConsumptionSource.applicationLabel 
-> applicationId)
     resourceConsumptionSource.addGauge(
       ResourceConsumptionSource.DISK_FILE_COUNT,
       resourceConsumptionLabel) { () =>
diff --git 
a/worker/src/main/java/org/apache/celeborn/service/deploy/worker/storage/CreditStreamManager.java
 
b/worker/src/main/java/org/apache/celeborn/service/deploy/worker/storage/CreditStreamManager.java
index 1fb7766a8..69587544f 100644
--- 
a/worker/src/main/java/org/apache/celeborn/service/deploy/worker/storage/CreditStreamManager.java
+++ 
b/worker/src/main/java/org/apache/celeborn/service/deploy/worker/storage/CreditStreamManager.java
@@ -78,6 +78,7 @@ public class CreditStreamManager {
   public long registerStream(
       Consumer<Long> notifyStreamHandlerCallback,
       Channel channel,
+      String shuffleKey,
       int initialCredit,
       int startSubIndex,
       int endSubIndex,
@@ -112,7 +113,7 @@ public class CreditStreamManager {
                 }
               }
               initializeStreamStateAndPartitionReader(
-                  channel, startSubIndex, endSubIndex, fileInfo, streamId, v);
+                  channel, shuffleKey, startSubIndex, endSubIndex, fileInfo, 
streamId, v);
               return v;
             });
     if (exception.get() != null) {
@@ -130,6 +131,7 @@ public class CreditStreamManager {
 
   private void initializeStreamStateAndPartitionReader(
       Channel channel,
+      String shuffleKey,
       int startSubIndex,
       int endSubIndex,
       FileInfo fileInfo,
@@ -137,7 +139,10 @@ public class CreditStreamManager {
       MapPartitionData mapPartitionData) {
     StreamState streamState =
         new StreamState(
-            channel, ((MapFileMeta) fileInfo.getFileMeta()).getBufferSize(), 
mapPartitionData);
+            channel,
+            shuffleKey,
+            ((MapFileMeta) fileInfo.getFileMeta()).getBufferSize(),
+            mapPartitionData);
     streams.put(streamId, streamState);
     mapPartitionData.setupDataPartitionReader(startSubIndex, endSubIndex, 
streamId, channel);
   }
@@ -186,6 +191,10 @@ public class CreditStreamManager {
     return streams;
   }
 
+  public String getStreamShuffleKey(Long streamId) {
+    return streams.get(streamId).getShuffleKey();
+  }
+
   private void startRecycleThread() {
     synchronized (lock) {
       if (recycleThread == null) {
@@ -245,12 +254,17 @@ public class CreditStreamManager {
 
   protected class StreamState {
     private Channel associatedChannel;
+    private String shuffleKey;
     private int bufferSize;
     private MapPartitionData mapPartitionData;
 
     public StreamState(
-        Channel associatedChannel, int bufferSize, MapPartitionData 
mapPartitionData) {
+        Channel associatedChannel,
+        String shuffleKey,
+        int bufferSize,
+        MapPartitionData mapPartitionData) {
       this.associatedChannel = associatedChannel;
+      this.shuffleKey = shuffleKey;
       this.bufferSize = bufferSize;
       this.mapPartitionData = mapPartitionData;
     }
@@ -259,6 +273,10 @@ public class CreditStreamManager {
       return associatedChannel;
     }
 
+    public String getShuffleKey() {
+      return shuffleKey;
+    }
+
     public int getBufferSize() {
       return bufferSize;
     }
diff --git 
a/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/FetchHandler.scala
 
b/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/FetchHandler.scala
index 006a1400d..ceabd9802 100644
--- 
a/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/FetchHandler.scala
+++ 
b/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/FetchHandler.scala
@@ -100,9 +100,9 @@ class FetchHandler(
   override def receive(client: TransportClient, msg: RequestMessage): Unit = {
     msg match {
       case r: BufferStreamEnd =>
-        handleEndStreamFromClient(r.getStreamId)
+        handleEndStreamFromClient(client, r.getStreamId)
       case r: ReadAddCredit =>
-        handleReadAddCredit(r.getCredit, r.getStreamId)
+        handleReadAddCredit(client, r.getCredit, r.getStreamId)
       case r: ChunkFetchRequest =>
         handleChunkFetchRequest(client, r.streamChunkSlice, r)
       case unknown: RequestMessage =>
@@ -137,9 +137,12 @@ class FetchHandler(
           openStream.getReadLocalShuffle,
           callback)
       case bufferStreamEnd: PbBufferStreamEnd =>
-        handleEndStreamFromClient(bufferStreamEnd.getStreamId, 
bufferStreamEnd.getStreamType)
+        handleEndStreamFromClient(
+          client,
+          bufferStreamEnd.getStreamId,
+          bufferStreamEnd.getStreamType)
       case readAddCredit: PbReadAddCredit =>
-        handleReadAddCredit(readAddCredit.getCredit, readAddCredit.getStreamId)
+        handleReadAddCredit(client, readAddCredit.getCredit, 
readAddCredit.getStreamId)
       case chunkFetchRequest: PbChunkFetchRequest =>
         handleChunkFetchRequest(
           client,
@@ -205,6 +208,7 @@ class FetchHandler(
       isLegacy: Boolean,
       readLocalShuffle: Boolean = false,
       callback: RpcResponseCallback): Unit = {
+    workerSource.recordAppActiveConnection(client, shuffleKey)
     workerSource.startTimer(WorkerSource.OPEN_STREAM_TIME, shuffleKey)
     try {
       var fileInfo = getRawDiskFileInfo(shuffleKey, fileName)
@@ -281,6 +285,7 @@ class FetchHandler(
           creditStreamManager.registerStream(
             creditStreamHandler,
             client.getChannel,
+            shuffleKey,
             initialCredit,
             startIndex,
             endIndex,
@@ -350,24 +355,34 @@ class FetchHandler(
     
rpcResponseCallback.onFailure(ExceptionUtils.wrapIOExceptionToUnRetryable(ioe))
   }
 
-  def handleEndStreamFromClient(streamId: Long): Unit = {
-    handleEndStreamFromClient(streamId, StreamType.CreditStream)
+  def handleEndStreamFromClient(client: TransportClient, streamId: Long): Unit 
= {
+    handleEndStreamFromClient(client, streamId, StreamType.CreditStream)
   }
 
-  def handleEndStreamFromClient(streamId: Long, streamType: StreamType): Unit 
= {
+  def handleEndStreamFromClient(
+      client: TransportClient,
+      streamId: Long,
+      streamType: StreamType): Unit = {
     streamType match {
       case StreamType.ChunkStream =>
         val (shuffleKey, fileName) = 
chunkStreamManager.getShuffleKeyAndFileName(streamId)
+        workerSource.recordAppActiveConnection(client, shuffleKey)
         getRawDiskFileInfo(shuffleKey, fileName).closeStream(
           streamId)
       case StreamType.CreditStream =>
+        workerSource.recordAppActiveConnection(
+          client,
+          creditStreamManager.getStreamShuffleKey(streamId))
         creditStreamManager.notifyStreamEndByClient(streamId)
       case _ =>
         logError(s"Received a PbBufferStreamEnd message with unknown type 
$streamType")
     }
   }
 
-  def handleReadAddCredit(credit: Int, streamId: Long): Unit = {
+  def handleReadAddCredit(client: TransportClient, credit: Int, streamId: 
Long): Unit = {
+    workerSource.recordAppActiveConnection(
+      client,
+      creditStreamManager.getStreamShuffleKey(streamId))
     creditStreamManager.addCredit(credit, streamId)
   }
 
@@ -378,6 +393,10 @@ class FetchHandler(
     logDebug(s"Received req from 
${NettyUtils.getRemoteAddress(client.getChannel)}" +
       s" to fetch block $streamChunkSlice")
 
+    workerSource.recordAppActiveConnection(
+      client,
+      
chunkStreamManager.getShuffleKeyAndFileName(streamChunkSlice.streamId)._1)
+
     maxChunkBeingTransferred.foreach { threshold =>
       val chunksBeingTransferred = chunkStreamManager.chunksBeingTransferred 
// take high cpu usage
       if (chunksBeingTransferred > threshold) {
@@ -442,12 +461,12 @@ class FetchHandler(
   /** Invoked when the channel associated with the given client is active. */
   override def channelActive(client: TransportClient): Unit = {
     logDebug(s"channel active ${client.getSocketAddress}")
-    workerSource.incCounter(WorkerSource.ACTIVE_CONNECTION_COUNT)
+    workerSource.connectionActive(client)
     super.channelActive(client)
   }
 
   override def channelInactive(client: TransportClient): Unit = {
-    workerSource.incCounter(WorkerSource.ACTIVE_CONNECTION_COUNT, -1)
+    workerSource.connectionInactive(client)
     creditStreamManager.connectionTerminated(client.getChannel)
     logDebug(s"channel inactive ${client.getSocketAddress}")
   }
diff --git 
a/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/PushDataHandler.scala
 
b/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/PushDataHandler.scala
index a0b0d6b95..37648d627 100644
--- 
a/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/PushDataHandler.scala
+++ 
b/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/PushDataHandler.scala
@@ -110,6 +110,7 @@ class PushDataHandler(val workerSource: WorkerSource) 
extends BaseMessageHandler
   override def receive(client: TransportClient, msg: RequestMessage): Unit =
     msg match {
       case pushData: PushData =>
+        workerSource.recordAppActiveConnection(client, pushData.shuffleKey)
         val callback = new SimpleRpcResponseCallback(
           client,
           pushData.requestId,
@@ -133,6 +134,7 @@ class PushDataHandler(val workerSource: WorkerSource) 
extends BaseMessageHandler
           },
           callback)
       case pushMergedData: PushMergedData =>
+        workerSource.recordAppActiveConnection(client, 
pushMergedData.shuffleKey)
         val callback = new SimpleRpcResponseCallback(
           client,
           pushMergedData.requestId,
@@ -828,6 +830,7 @@ class PushDataHandler(val workerSource: WorkerSource) 
extends BaseMessageHandler
     val requestId = rpcRequest.requestId
     val (pbMsg, msg, isLegacy, messageType, mode, shuffleKey, 
partitionUniqueId, checkSplit) =
       mapPartitionRpcRequest(rpcRequest)
+    workerSource.recordAppActiveConnection(client, shuffleKey)
     handleCore(
       client,
       rpcRequest,
@@ -1293,7 +1296,7 @@ class PushDataHandler(val workerSource: WorkerSource) 
extends BaseMessageHandler
    * Invoked when the channel associated with the given client is active.
    */
   override def channelActive(client: TransportClient): Unit = {
-    workerSource.incCounter(WorkerSource.ACTIVE_CONNECTION_COUNT)
+    workerSource.connectionActive(client)
     super.channelActive(client)
   }
 
@@ -1302,7 +1305,7 @@ class PushDataHandler(val workerSource: WorkerSource) 
extends BaseMessageHandler
    * No further requests will come from this client.
    */
   override def channelInactive(client: TransportClient): Unit = {
-    workerSource.incCounter(WorkerSource.ACTIVE_CONNECTION_COUNT, -1)
+    workerSource.connectionInactive(client)
     super.channelInactive(client)
   }
 }
diff --git 
a/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/Worker.scala 
b/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/Worker.scala
index 383820543..b1b5a72cc 100644
--- 
a/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/Worker.scala
+++ 
b/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/Worker.scala
@@ -470,6 +470,7 @@ private[celeborn] class Worker(
         commitThreadPool.shutdownNow()
         asyncReplyPool.shutdownNow()
       }
+      workerSource.appActiveConnections.clear()
       partitionsSorter.close(exitKind)
       storageManager.close(exitKind)
       memoryManager.close()
@@ -547,7 +548,7 @@ private[celeborn] class Worker(
       applicationId: String = null): Unit = {
     var resourceConsumptionLabel = userIdentifier.toMap
     if (applicationId != null)
-      resourceConsumptionLabel += (ResourceConsumptionSource.APPLICATION_LABEL 
-> applicationId)
+      resourceConsumptionLabel += (resourceConsumptionSource.applicationLabel 
-> applicationId)
     resourceConsumptionSource.addGauge(
       ResourceConsumptionSource.DISK_FILE_COUNT,
       resourceConsumptionLabel) { () =>
@@ -597,6 +598,7 @@ private[celeborn] class Worker(
           // When the running applications does not contain the application 
corresponding to expired shuffle key,
           // resource consumption source should remove lose application gauges.
           removeAppResourceConsumption(applicationId)
+          removeAppActiveConnection(applicationId)
         }
         logInfo(s"Cleaned up expired shuffle $shuffleKey")
       }
@@ -627,10 +629,14 @@ private[celeborn] class Worker(
       applicationId: String): Unit = {
     resourceConsumptionSource.removeGauge(
       resourceConsumptionName,
-      ResourceConsumptionSource.APPLICATION_LABEL,
+      resourceConsumptionSource.applicationLabel,
       applicationId)
   }
 
+  private def removeAppActiveConnection(applicationId: String): Unit = {
+    workerSource.removeAppActiveConnection(applicationId)
+  }
+
   override def getWorkerInfo: String = {
     val sb = new StringBuilder
     sb.append("====================== WorkerInfo of Worker 
===========================\n")
diff --git 
a/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/WorkerSource.scala
 
b/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/WorkerSource.scala
index eb037afc7..0b52e9093 100644
--- 
a/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/WorkerSource.scala
+++ 
b/worker/src/main/scala/org/apache/celeborn/service/deploy/worker/WorkerSource.scala
@@ -17,13 +17,25 @@
 
 package org.apache.celeborn.service.deploy.worker
 
+import java.util
+import java.util.concurrent.ConcurrentHashMap
+
+import scala.collection.JavaConverters._
+
+import com.google.common.collect.Sets
+
 import org.apache.celeborn.common.CelebornConf
 import org.apache.celeborn.common.metrics.MetricsSystem
 import org.apache.celeborn.common.metrics.source.AbstractSource
+import org.apache.celeborn.common.network.client.TransportClient
+import org.apache.celeborn.common.util.{CollectionUtils, JavaUtils, Utils}
 
 class WorkerSource(conf: CelebornConf) extends AbstractSource(conf, 
MetricsSystem.ROLE_WORKER) {
   override val sourceName = "worker"
 
+  val appActiveConnections: ConcurrentHashMap[String, util.Set[String]] =
+    JavaUtils.newConcurrentHashMap[String, util.Set[String]]
+
   import WorkerSource._
   // add counters
   addCounter(OPEN_STREAM_SUCCESS_COUNT)
@@ -69,6 +81,41 @@ class WorkerSource(conf: CelebornConf) extends 
AbstractSource(conf, MetricsSyste
     val metricNameWithLabel = metricNameWithCustomizedLabels(metricsName, 
Map.empty)
     namedCounters.get(metricNameWithLabel).counter.getCount
   }
+
+  def connectionActive(client: TransportClient): Unit = {
+    appActiveConnections.putIfAbsent(
+      client.getChannel.id().asLongText(),
+      Sets.newConcurrentHashSet[String]())
+    incCounter(ACTIVE_CONNECTION_COUNT, 1)
+  }
+
+  def connectionInactive(client: TransportClient): Unit = {
+    appActiveConnections.remove(client.getChannel.id().asLongText())
+    incCounter(ACTIVE_CONNECTION_COUNT, -1)
+  }
+
+  def recordAppActiveConnection(client: TransportClient, shuffleKey: String): 
Unit = {
+    val applicationIds = 
appActiveConnections.get(client.getChannel.id().asLongText())
+    val applicationId = Utils.splitShuffleKey(shuffleKey)._1
+    if (CollectionUtils.isNotEmpty(applicationIds) && 
!applicationIds.contains(applicationId)) {
+      applicationIds.add(applicationId)
+      addGauge(ACTIVE_CONNECTION_COUNT, Map(applicationLabel -> 
applicationId)) { () =>
+        appActiveConnections.asScala.count { case (_, applicationIds) =>
+          applicationIds.contains(applicationId)
+        }
+      }
+    }
+  }
+
+  def removeAppActiveConnection(applicationId: String): Unit = {
+    appActiveConnections.asScala.foreach { case (_, applicationIds) =>
+      if (applicationIds.contains(applicationId)) {
+        applicationIds.remove(applicationId)
+        removeGauge(ACTIVE_CONNECTION_COUNT, Map(applicationLabel -> 
applicationId))
+      }
+    }
+  }
+
   // start cleaner thread
   startCleaner()
 }
diff --git 
a/worker/src/test/java/org/apache/celeborn/service/deploy/worker/storage/CreditStreamManagerSuiteJ.java
 
b/worker/src/test/java/org/apache/celeborn/service/deploy/worker/storage/CreditStreamManagerSuiteJ.java
index 4050691b1..0896f2462 100644
--- 
a/worker/src/test/java/org/apache/celeborn/service/deploy/worker/storage/CreditStreamManagerSuiteJ.java
+++ 
b/worker/src/test/java/org/apache/celeborn/service/deploy/worker/storage/CreditStreamManagerSuiteJ.java
@@ -82,18 +82,23 @@ public class CreditStreamManagerSuiteJ {
     diskFileInfo.replaceFileMeta(mapFileMeta);
     Consumer<Long> streamIdConsumer = streamId -> Assert.assertTrue(streamId > 
0);
 
+    String shuffleKey = "application_1694674023293_0003-0";
     long registerStream1 =
-        creditStreamManager.registerStream(streamIdConsumer, channel, 0, 1, 1, 
diskFileInfo);
+        creditStreamManager.registerStream(
+            streamIdConsumer, channel, shuffleKey, 0, 1, 1, diskFileInfo);
     Assert.assertTrue(registerStream1 > 0);
     Assert.assertEquals(1, creditStreamManager.getStreamsCount());
 
     long registerStream2 =
-        creditStreamManager.registerStream(streamIdConsumer, channel, 0, 1, 1, 
diskFileInfo);
+        creditStreamManager.registerStream(
+            streamIdConsumer, channel, shuffleKey, 0, 1, 1, diskFileInfo);
     Assert.assertNotEquals(registerStream1, registerStream2);
     Assert.assertEquals(2, creditStreamManager.getStreamsCount());
 
-    creditStreamManager.registerStream(streamIdConsumer, channel, 0, 1, 1, 
diskFileInfo);
-    creditStreamManager.registerStream(streamIdConsumer, channel, 0, 1, 1, 
diskFileInfo);
+    creditStreamManager.registerStream(
+        streamIdConsumer, channel, shuffleKey, 0, 1, 1, diskFileInfo);
+    creditStreamManager.registerStream(
+        streamIdConsumer, channel, shuffleKey, 0, 1, 1, diskFileInfo);
 
     MapPartitionData mapPartitionData1 =
         
creditStreamManager.getStreams().get(registerStream1).getMapDataPartition();

Reply via email to