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

zhouky 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 a86a315bd [CELEBORN-1229] Support for application registration with 
Celeborn Master
a86a315bd is described below

commit a86a315bd56bfa327d80f80ba01b6e73ceca62bd
Author: Chandni Singh <[email protected]>
AuthorDate: Tue Jan 23 09:31:19 2024 +0800

    [CELEBORN-1229] Support for application registration with Celeborn Master
    
    ### What changes were proposed in this pull request?
    This adds support for applications to register with Celeborn Master by 
introducing the `RegistrationClientBootstrap`, `RegistrationServerBootstrap`, 
and `RegistrationRpcHandler` classes, which facilitate the client connection 
setup with the Celeborn Master. The registration protocol details are described 
in the [auth 
proposal](https://docs.google.com/document/d/1D1U2COYhS3ob7l0t2WghRhBk_Fci9RGx-2FBXA3nvXk/edit#heading=h.po9dc3r1kb3k).
    
    ### Why are the changes needed?
    The changes are needed for adding authentication to Celeborn. See 
[CELEBORN-1011](https://issues.apache.org/jira/browse/CELEBORN-1011).
    
    ### Does this PR introduce _any_ user-facing change?
    Add the config `celeborn.auth.enabled`
    
    ### How was this patch tested?
    Added UTs.
    
    Closes #2231 from otterc/CELEBORN-1229.
    
    Authored-by: Chandni Singh <[email protected]>
    Signed-off-by: zky.zhoukeyong <[email protected]>
---
 .../network/client/TransportClientBootstrap.java   |   4 +-
 .../network/client/TransportClientFactory.java     |   3 +-
 .../common/network/protocol/TransportMessage.java  |  12 +
 .../common/network/sasl/CelebornSaslServer.java    |   6 +-
 .../common/network/sasl/SaslClientBootstrap.java   |   6 +-
 .../common/network/sasl/SaslRpcHandler.java        |   2 +-
 .../celeborn/common/network/sasl/SaslUtils.java    |   4 +-
 .../common/network/sasl/SecretRegistry.java        |   4 +
 .../common/network/sasl/SecretRegistryImpl.java    |  12 +-
 .../registration/RegistrationClientBootstrap.java  | 261 +++++++++++++++++++
 .../RegistrationInfo.java}                         |  25 +-
 .../sasl/registration/RegistrationRpcHandler.java  | 281 +++++++++++++++++++++
 .../registration/RegistrationServerBootstrap.java  |  45 ++++
 .../common/network/util/TransportConf.java         |   5 +
 common/src/main/proto/TransportMessages.proto      |  30 ++-
 .../org/apache/celeborn/common/CelebornConf.scala  |  13 +
 .../common/network/sasl/CelebornSaslSuiteJ.java    | 104 +-------
 .../common/network/sasl/RegistrationSuiteJ.java    | 113 +++++++++
 .../celeborn/common/network/sasl/SaslTestBase.java | 140 ++++++++++
 19 files changed, 950 insertions(+), 120 deletions(-)

diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/client/TransportClientBootstrap.java
 
b/common/src/main/java/org/apache/celeborn/common/network/client/TransportClientBootstrap.java
index bdad119fc..a82e204c7 100644
--- 
a/common/src/main/java/org/apache/celeborn/common/network/client/TransportClientBootstrap.java
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/client/TransportClientBootstrap.java
@@ -17,8 +17,6 @@
 
 package org.apache.celeborn.common.network.client;
 
-import io.netty.channel.Channel;
-
 /**
  * A bootstrap which is executed on a TransportClient before it is returned to 
the user. This
  * enables an initial exchange of information (e.g., SASL authentication 
tokens) on a once-per-
@@ -36,5 +34,5 @@ public interface TransportClientBootstrap {
    * @param channel the associated channel with the transport client
    * @throws RuntimeException
    */
-  void doBootstrap(TransportClient client, Channel channel) throws 
RuntimeException;
+  void doBootstrap(TransportClient client) throws RuntimeException;
 }
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/client/TransportClientFactory.java
 
b/common/src/main/java/org/apache/celeborn/common/network/client/TransportClientFactory.java
index a4d4d7515..68bad239b 100644
--- 
a/common/src/main/java/org/apache/celeborn/common/network/client/TransportClientFactory.java
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/client/TransportClientFactory.java
@@ -258,7 +258,6 @@ public class TransportClientFactory implements Closeable {
     }
 
     TransportClient client = clientRef.get();
-    Channel channel = channelRef.get();
     assert client != null : "Channel future completed successfully with null 
client";
 
     // Execute any client bootstraps synchronously before marking the Client 
as successful.
@@ -266,7 +265,7 @@ public class TransportClientFactory implements Closeable {
     logger.debug("Running bootstraps for {} ...", address);
     try {
       for (TransportClientBootstrap clientBootstrap : clientBootstraps) {
-        clientBootstrap.doBootstrap(client, channel);
+        clientBootstrap.doBootstrap(client);
       }
     } catch (Exception e) { // catch non-RuntimeExceptions too as bootstrap 
may be written in Scala
       long bootstrapTime = System.nanoTime() - preBootstrap;
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/protocol/TransportMessage.java
 
b/common/src/main/java/org/apache/celeborn/common/network/protocol/TransportMessage.java
index c14f20b5d..4e2978b9d 100644
--- 
a/common/src/main/java/org/apache/celeborn/common/network/protocol/TransportMessage.java
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/protocol/TransportMessage.java
@@ -29,6 +29,8 @@ import org.slf4j.LoggerFactory;
 
 import org.apache.celeborn.common.exception.CelebornIOException;
 import org.apache.celeborn.common.protocol.MessageType;
+import org.apache.celeborn.common.protocol.PbAuthenticationInitiationRequest;
+import org.apache.celeborn.common.protocol.PbAuthenticationInitiationResponse;
 import org.apache.celeborn.common.protocol.PbBacklogAnnouncement;
 import org.apache.celeborn.common.protocol.PbBufferStreamEnd;
 import org.apache.celeborn.common.protocol.PbChunkFetchRequest;
@@ -39,6 +41,8 @@ import 
org.apache.celeborn.common.protocol.PbPushDataHandShake;
 import org.apache.celeborn.common.protocol.PbReadAddCredit;
 import org.apache.celeborn.common.protocol.PbRegionFinish;
 import org.apache.celeborn.common.protocol.PbRegionStart;
+import org.apache.celeborn.common.protocol.PbRegisterApplicationRequest;
+import org.apache.celeborn.common.protocol.PbRegisterApplicationResponse;
 import org.apache.celeborn.common.protocol.PbReportShuffleFetchFailure;
 import org.apache.celeborn.common.protocol.PbReportShuffleFetchFailureResponse;
 import org.apache.celeborn.common.protocol.PbSaslRequest;
@@ -105,6 +109,14 @@ public class TransportMessage implements Serializable {
         return (T) PbReportShuffleFetchFailureResponse.parseFrom(payload);
       case SASL_REQUEST_VALUE:
         return (T) PbSaslRequest.parseFrom(payload);
+      case AUTHENTICATION_INITIATION_REQUEST_VALUE:
+        return (T) PbAuthenticationInitiationRequest.parseFrom(payload);
+      case AUTHENTICATION_INITIATION_RESPONSE_VALUE:
+        return (T) PbAuthenticationInitiationResponse.parseFrom(payload);
+      case REGISTER_APPLICATION_REQUEST_VALUE:
+        return (T) PbRegisterApplicationRequest.parseFrom(payload);
+      case REGISTER_APPLICATION_RESPONSE_VALUE:
+        return (T) PbRegisterApplicationResponse.parseFrom(payload);
       default:
         logger.error("Unexpected type {}", type);
     }
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/CelebornSaslServer.java
 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/CelebornSaslServer.java
index bc2354f05..15153b308 100644
--- 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/CelebornSaslServer.java
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/CelebornSaslServer.java
@@ -105,7 +105,7 @@ public class CelebornSaslServer {
    * Implementation of javax.security.auth.callback.CallbackHandler for SASL 
DIGEST-MD5 mechanism.
    */
   static class DigestCallbackHandler implements CallbackHandler {
-    private final SecretRegistry secretKeyHolder;
+    private final SecretRegistry secretRegistry;
 
     /**
      * The use of 'volatile' is not necessary here because the 'handle' 
invocation includes both the
@@ -115,7 +115,7 @@ public class CelebornSaslServer {
     private String userName = null;
 
     DigestCallbackHandler(SecretRegistry secretRegistry) {
-      this.secretKeyHolder = Preconditions.checkNotNull(secretRegistry);
+      this.secretRegistry = Preconditions.checkNotNull(secretRegistry);
     }
 
     @Override
@@ -132,7 +132,7 @@ public class CelebornSaslServer {
         } else if (callback instanceof PasswordCallback) {
           logger.trace("SASL server callback: setting password");
           PasswordCallback pc = (PasswordCallback) callback;
-          String secret = secretKeyHolder.getSecretKey(userName);
+          String secret = secretRegistry.getSecretKey(userName);
           if (secret == null) {
             // TODO: CELEBORN-1179 Add support for fetching the secret from 
the Celeborn master.
             throw new RuntimeException("Registration information not found for 
" + userName);
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/SaslClientBootstrap.java
 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/SaslClientBootstrap.java
index 99a8054b2..a970e2c7c 100644
--- 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/SaslClientBootstrap.java
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/SaslClientBootstrap.java
@@ -25,7 +25,6 @@ import java.util.concurrent.TimeoutException;
 
 import com.google.common.base.Preconditions;
 import com.google.protobuf.ByteString;
-import io.netty.channel.Channel;
 import org.slf4j.Logger;
 import org.slf4j.LoggerFactory;
 
@@ -69,7 +68,7 @@ public class SaslClientBootstrap implements 
TransportClientBootstrap {
    * to mismatch.
    */
   @Override
-  public void doBootstrap(TransportClient client, Channel channel) {
+  public void doBootstrap(TransportClient client) {
     // TODO: Hardcoding the SASL mechanism to DIGEST-MD5 for Connection 
Authentication. This
     // should be configurable in the future.
     CelebornSaslClient saslClient =
@@ -84,8 +83,9 @@ public class SaslClientBootstrap implements 
TransportClientBootstrap {
       while (!saslClient.isComplete()) {
         PbSaslRequest.Builder builder = PbSaslRequest.newBuilder();
         if (firstToken) {
-          
builder.setMethod(DIGEST_MD5).setAuthType(PbAuthType.CONNECTION_AUTH);
+          builder.setMethod(DIGEST_MD5);
         }
+        builder.setAuthType(PbAuthType.CONNECTION_AUTH);
         TransportMessage msg =
             new TransportMessage(
                 MessageType.SASL_REQUEST,
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/SaslRpcHandler.java
 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/SaslRpcHandler.java
index 5bc72f0f1..127b3e772 100644
--- 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/SaslRpcHandler.java
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/SaslRpcHandler.java
@@ -117,7 +117,7 @@ public class SaslRpcHandler extends AbstractAuthRpcHandler {
     cleanup();
   }
 
-  private void cleanup() {
+  public void cleanup() {
     if (null != saslServer) {
       try {
         saslServer.dispose();
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/SaslUtils.java 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/SaslUtils.java
index 374543499..5203581bd 100644
--- 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/SaslUtils.java
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/SaslUtils.java
@@ -32,7 +32,7 @@ public class SaslUtils {
   static final byte[] EMPTY_BYTE_ARRAY = new byte[0];
 
   /** Sasl Mechanisms */
-  static final String DIGEST_MD5 = "DIGEST-MD5";
+  public static final String DIGEST_MD5 = "DIGEST-MD5";
 
   public static final String ANONYMOUS = "ANONYMOUS";
 
@@ -41,7 +41,7 @@ public class SaslUtils {
 
   static final String DEFAULT_REALM = "default";
 
-  static final Map<String, String> DEFAULT_SASL_CLIENT_PROPS =
+  public static final Map<String, String> DEFAULT_SASL_CLIENT_PROPS =
       ImmutableMap.<String, String>builder().put(Sasl.QOP, QOP_AUTH).build();
 
   static final Map<String, String> DEFAULT_SASL_SERVER_PROPS =
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/SecretRegistry.java
 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/SecretRegistry.java
index 995e7af80..9a933189e 100644
--- 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/SecretRegistry.java
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/SecretRegistry.java
@@ -24,4 +24,8 @@ public interface SecretRegistry {
   String getSecretKey(String appId);
 
   boolean isRegistered(String appId);
+
+  void register(String appId, String secret);
+
+  void unregister(String appId);
 }
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/SecretRegistryImpl.java
 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/SecretRegistryImpl.java
index 92b33bc47..811ccbf98 100644
--- 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/SecretRegistryImpl.java
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/SecretRegistryImpl.java
@@ -33,10 +33,20 @@ public class SecretRegistryImpl implements SecretRegistry {
 
   private final ConcurrentHashMap<String, String> secrets = new 
ConcurrentHashMap<>();
 
+  @Override
   public void register(String appId, String secret) {
-    secrets.put(appId, secret);
+    // TODO: Persist the secret in ratis. See 
https://issues.apache.org/jira/browse/CELEBORN-1234
+    secrets.compute(
+        appId,
+        (id, oldVal) -> {
+          if (oldVal != null) {
+            throw new IllegalArgumentException("AppId " + appId + " is already 
registered.");
+          }
+          return secret;
+        });
   }
 
+  @Override
   public void unregister(String appId) {
     secrets.remove(appId);
   }
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/registration/RegistrationClientBootstrap.java
 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/registration/RegistrationClientBootstrap.java
new file mode 100644
index 000000000..3604e927c
--- /dev/null
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/registration/RegistrationClientBootstrap.java
@@ -0,0 +1,261 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.celeborn.common.network.sasl.registration;
+
+import static org.apache.celeborn.common.network.sasl.SaslUtils.*;
+
+import java.io.IOException;
+import java.nio.ByteBuffer;
+import java.util.HashMap;
+import java.util.List;
+import java.util.Map;
+import java.util.Set;
+import java.util.concurrent.TimeoutException;
+
+import com.google.common.base.Preconditions;
+import com.google.common.collect.Lists;
+import com.google.common.collect.Sets;
+import com.google.protobuf.ByteString;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+
+import org.apache.celeborn.common.exception.CelebornException;
+import org.apache.celeborn.common.network.client.TransportClient;
+import org.apache.celeborn.common.network.client.TransportClientBootstrap;
+import org.apache.celeborn.common.network.protocol.TransportMessage;
+import org.apache.celeborn.common.network.sasl.CelebornSaslClient;
+import org.apache.celeborn.common.network.sasl.SaslClientBootstrap;
+import org.apache.celeborn.common.network.sasl.SaslCredentials;
+import org.apache.celeborn.common.network.sasl.SaslTimeoutException;
+import org.apache.celeborn.common.network.util.TransportConf;
+import org.apache.celeborn.common.protocol.MessageType;
+import org.apache.celeborn.common.protocol.PbAuthType;
+import org.apache.celeborn.common.protocol.PbAuthenticationInitiationRequest;
+import org.apache.celeborn.common.protocol.PbAuthenticationInitiationResponse;
+import org.apache.celeborn.common.protocol.PbRegisterApplicationRequest;
+import org.apache.celeborn.common.protocol.PbRegisterApplicationResponse;
+import org.apache.celeborn.common.protocol.PbSaslMechanism;
+import org.apache.celeborn.common.protocol.PbSaslRequest;
+import org.apache.celeborn.common.util.JavaUtils;
+
+/**
+ * Bootstraps a {@link TransportClient} by registering application (if the 
application is not
+ * registered). If the application is already registered, it will bootstrap 
the client by performing
+ * SASL authentication.
+ */
+public class RegistrationClientBootstrap implements TransportClientBootstrap {
+
+  private static final Logger LOG = 
LoggerFactory.getLogger(RegistrationClientBootstrap.class);
+
+  private static final String VERSION = "1.0";
+
+  /**
+   * TODO: This should be made configurable. For now, we only support 
ANONYMOUS for client-auth and
+   * DIGEST-MD5 for connect-auth.
+   */
+  private static final List<PbSaslMechanism> SASL_MECHANISMS =
+      Lists.newArrayList(
+          PbSaslMechanism.newBuilder()
+              .setMechanism(ANONYMOUS)
+              .addAuthTypes(PbAuthType.CLIENT_AUTH)
+              .build(),
+          PbSaslMechanism.newBuilder()
+              .setMechanism(DIGEST_MD5)
+              .addAuthTypes(PbAuthType.CONNECTION_AUTH)
+              .build());
+
+  private final TransportConf conf;
+  private final String appId;
+  private final SaslCredentials saslCredentials;
+
+  private final RegistrationInfo registrationInfo;
+
+  public RegistrationClientBootstrap(
+      TransportConf conf,
+      String appId,
+      SaslCredentials saslCredentials,
+      RegistrationInfo registrationInfo) {
+    this.conf = Preconditions.checkNotNull(conf, "conf");
+    this.appId = Preconditions.checkNotNull(appId, "appId");
+    this.saslCredentials = Preconditions.checkNotNull(saslCredentials, 
"saslCredentials");
+    this.registrationInfo = Preconditions.checkNotNull(registrationInfo, 
"registrationInfo");
+  }
+
+  @Override
+  public void doBootstrap(TransportClient client) throws RuntimeException {
+    if (registrationInfo.getRegistrationState() == 
RegistrationInfo.RegistrationState.REGISTERED) {
+      LOG.info("client has already registered, skip register.");
+      doSaslBootstrap(client);
+      return;
+    }
+    try {
+      LOG.info("authentication initiation started for {}", appId);
+      doAuthInitiation(client);
+      LOG.info("authentication initiation successful for {}", appId);
+      doClientAuthentication(client);
+      LOG.info("client authenticated for {}", appId);
+      register(client);
+      LOG.info("Registration for {}", appId);
+      
registrationInfo.setRegistrationState(RegistrationInfo.RegistrationState.REGISTERED);
+    } catch (IOException | CelebornException e) {
+      throw new RuntimeException(e);
+    } finally {
+      if (registrationInfo.getRegistrationState()
+          != RegistrationInfo.RegistrationState.REGISTERED) {
+        
registrationInfo.setRegistrationState(RegistrationInfo.RegistrationState.FAILED);
+      }
+    }
+  }
+
+  private void doAuthInitiation(TransportClient client) throws IOException, 
CelebornException {
+    PbAuthenticationInitiationRequest authInitRequest =
+        PbAuthenticationInitiationRequest.newBuilder()
+            .setVersion(VERSION)
+            .setAuthEnabled(true)
+            .addAllSaslMechanisms(SASL_MECHANISMS)
+            .build();
+    TransportMessage msg =
+        new TransportMessage(
+            MessageType.AUTHENTICATION_INITIATION_REQUEST, 
authInitRequest.toByteArray());
+    ByteBuffer authInitResponseBuffer;
+    try {
+      authInitResponseBuffer = client.sendRpcSync(msg.toByteBuffer(), 
conf.saslTimeoutMs());
+    } catch (RuntimeException ex) {
+      if (ex.getCause() instanceof TimeoutException) {
+        throw new SaslTimeoutException(ex.getCause());
+      } else {
+        throw ex;
+      }
+    }
+    PbAuthenticationInitiationResponse authInitResponse =
+        
TransportMessage.fromByteBuffer(authInitResponseBuffer).getParsedPayload();
+    if (!validateServerResponse(authInitResponse)) {
+      String exMsg =
+          "Registration failed due to incompatibility with the server."
+              + " InitRequest: "
+              + authInitRequest
+              + " InitResponse: "
+              + authInitResponse;
+      throw new CelebornException(exMsg);
+    }
+    // TODO: client validates required/supported mechanism is present
+  }
+
+  private void doClientAuthentication(TransportClient client) throws 
IOException {
+    // Client will authenticate itself with the selected SaslMechanism for 
Client Authentication
+    CelebornSaslClient saslClient = new CelebornSaslClient(ANONYMOUS, null, 
null);
+    try {
+      byte[] payload = saslClient.firstToken();
+      while (!saslClient.isComplete()) {
+        TransportMessage msg =
+            new TransportMessage(
+                MessageType.SASL_REQUEST,
+                PbSaslRequest.newBuilder()
+                    .setMethod(ANONYMOUS)
+                    .setAuthType(PbAuthType.CLIENT_AUTH)
+                    .setPayload(ByteString.copyFrom(payload))
+                    .build()
+                    .toByteArray());
+        ByteBuffer response;
+        try {
+          LOG.info("Sending SASL message for client authentication");
+          response = client.sendRpcSync(msg.toByteBuffer(), 
conf.saslTimeoutMs());
+        } catch (RuntimeException ex) {
+          // We know it is a Sasl timeout here if it is a TimeoutException.
+          if (ex.getCause() instanceof TimeoutException) {
+            throw new SaslTimeoutException(ex.getCause());
+          } else {
+            throw ex;
+          }
+        }
+        payload = saslClient.response(JavaUtils.bufferToArray(response));
+      }
+
+    } finally {
+      try { // Once authentication is complete, the server will trust all 
remaining communication.
+        saslClient.dispose();
+      } catch (RuntimeException e) {
+        LOG.warn("Error while disposing SASL client", e);
+      }
+    }
+  }
+
+  private void register(TransportClient client) throws IOException, 
CelebornException {
+    TransportMessage msg =
+        new TransportMessage(
+            MessageType.REGISTER_APPLICATION_REQUEST,
+            PbRegisterApplicationRequest.newBuilder()
+                .setId(appId)
+                .setSecret(saslCredentials.getPassword())
+                .build()
+                .toByteArray());
+    ByteBuffer response;
+    try {
+      response = client.sendRpcSync(msg.toByteBuffer(), conf.saslTimeoutMs());
+    } catch (RuntimeException ex) {
+      // We know it is a Sasl timeout here if it is a TimeoutException.
+      if (ex.getCause() instanceof TimeoutException) {
+        throw new SaslTimeoutException(ex.getCause());
+      } else {
+        throw ex;
+      }
+    }
+    PbRegisterApplicationResponse registerApplicationResponse =
+        TransportMessage.fromByteBuffer(response).getParsedPayload();
+    if (!registerApplicationResponse.getStatus()) {
+      throw new CelebornException("Application registration failed. AppId = " 
+ appId);
+    }
+  }
+
+  private void doSaslBootstrap(TransportClient client) {
+    SaslClientBootstrap bootstrap = new SaslClientBootstrap(conf, appId, 
saslCredentials);
+    bootstrap.doBootstrap(client);
+  }
+
+  private boolean validateServerResponse(PbAuthenticationInitiationResponse 
authInitResponse) {
+    if (!authInitResponse.getVersion().equals(VERSION)) {
+      return false;
+    }
+    Map<PbAuthType, Set<String>> serverSupportedMechs =
+        findSupportedSaslMechs(authInitResponse.getSaslMechanismsList());
+    Set<String> clientAuthMechs = 
serverSupportedMechs.get(PbAuthType.CLIENT_AUTH);
+    if (clientAuthMechs == null) {
+      return false;
+    }
+    if (!clientAuthMechs.contains(ANONYMOUS)) {
+      return false;
+    }
+    Set<String> connectionAuthMechs = 
serverSupportedMechs.get(PbAuthType.CONNECTION_AUTH);
+    if (connectionAuthMechs == null) {
+      return false;
+    }
+    return connectionAuthMechs.contains(DIGEST_MD5);
+  }
+
+  private static Map<PbAuthType, Set<String>> findSupportedSaslMechs(
+      List<PbSaslMechanism> serverSupportedMechs) {
+    Map<PbAuthType, Set<String>> supportedMechs = new HashMap<>();
+    for (PbSaslMechanism mech : serverSupportedMechs) {
+      for (PbAuthType authType : mech.getAuthTypesList()) {
+        Set<String> mechanisms = supportedMechs.computeIfAbsent(authType, k -> 
Sets.newHashSet());
+        mechanisms.add(mech.getMechanism());
+      }
+    }
+    return supportedMechs;
+  }
+}
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/SecretRegistry.java
 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/registration/RegistrationInfo.java
similarity index 60%
copy from 
common/src/main/java/org/apache/celeborn/common/network/sasl/SecretRegistry.java
copy to 
common/src/main/java/org/apache/celeborn/common/network/sasl/registration/RegistrationInfo.java
index 995e7af80..0b9f46df2 100644
--- 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/SecretRegistry.java
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/registration/RegistrationInfo.java
@@ -15,13 +15,26 @@
  * limitations under the License.
  */
 
-package org.apache.celeborn.common.network.sasl;
+package org.apache.celeborn.common.network.sasl.registration;
 
-/** Interface for getting a secret key associated with some application. */
-public interface SecretRegistry {
+import java.util.concurrent.atomic.AtomicReference;
 
-  /** Gets an appropriate SASL secret key for the given appId. */
-  String getSecretKey(String appId);
+public class RegistrationInfo {
 
-  boolean isRegistered(String appId);
+  private final AtomicReference<RegistrationState> state =
+      new AtomicReference<>(RegistrationState.UNREGISTERED);
+
+  public RegistrationState getRegistrationState() {
+    return state.get();
+  }
+
+  public void setRegistrationState(RegistrationState newState) {
+    state.set(newState);
+  }
+
+  public enum RegistrationState {
+    REGISTERED,
+    UNREGISTERED,
+    FAILED
+  }
 }
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/registration/RegistrationRpcHandler.java
 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/registration/RegistrationRpcHandler.java
new file mode 100644
index 000000000..10803f95d
--- /dev/null
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/registration/RegistrationRpcHandler.java
@@ -0,0 +1,281 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.celeborn.common.network.sasl.registration;
+
+import static org.apache.celeborn.common.network.sasl.SaslUtils.*;
+import static org.apache.celeborn.common.protocol.MessageType.*;
+
+import java.io.IOException;
+import java.nio.ByteBuffer;
+import java.util.List;
+
+import com.google.common.base.Throwables;
+import com.google.common.collect.Lists;
+import io.netty.channel.Channel;
+import org.slf4j.Logger;
+import org.slf4j.LoggerFactory;
+
+import org.apache.celeborn.common.network.client.RpcResponseCallback;
+import org.apache.celeborn.common.network.client.TransportClient;
+import org.apache.celeborn.common.network.protocol.RequestMessage;
+import org.apache.celeborn.common.network.protocol.RpcFailure;
+import org.apache.celeborn.common.network.protocol.RpcRequest;
+import org.apache.celeborn.common.network.protocol.TransportMessage;
+import org.apache.celeborn.common.network.sasl.CelebornSaslServer;
+import org.apache.celeborn.common.network.sasl.SaslRpcHandler;
+import org.apache.celeborn.common.network.sasl.SecretRegistry;
+import org.apache.celeborn.common.network.server.BaseMessageHandler;
+import org.apache.celeborn.common.network.util.TransportConf;
+import org.apache.celeborn.common.protocol.PbAuthType;
+import org.apache.celeborn.common.protocol.PbAuthenticationInitiationRequest;
+import org.apache.celeborn.common.protocol.PbAuthenticationInitiationResponse;
+import org.apache.celeborn.common.protocol.PbRegisterApplicationRequest;
+import org.apache.celeborn.common.protocol.PbRegisterApplicationResponse;
+import org.apache.celeborn.common.protocol.PbSaslMechanism;
+import org.apache.celeborn.common.protocol.PbSaslRequest;
+
+/**
+ * RPC Handler which registers an application. If an application is registered 
or the connection is
+ * authenticated, subsequent messages are delegated to a child RPC handler.
+ */
+public class RegistrationRpcHandler extends BaseMessageHandler {
+  private static final Logger LOG = 
LoggerFactory.getLogger(RegistrationRpcHandler.class);
+
+  private static final String VERSION = "1.0";
+
+  /**
+   * TODO: This should be made configurable. For now, we only support 
ANONYMOUS for client-auth and
+   * DIGEST-MD5 for connect-auth.
+   */
+  private static final List<PbSaslMechanism> SASL_MECHANISMS =
+      Lists.newArrayList(
+          PbSaslMechanism.newBuilder()
+              .setMechanism(ANONYMOUS)
+              .addAuthTypes(PbAuthType.CLIENT_AUTH)
+              .build(),
+          PbSaslMechanism.newBuilder()
+              .setMechanism(DIGEST_MD5)
+              .addAuthTypes(PbAuthType.CONNECTION_AUTH)
+              .build());
+
+  /** Transport configuration. */
+  private final TransportConf conf;
+
+  /** The client channel. */
+  private final Channel channel;
+
+  private final BaseMessageHandler delegate;
+
+  private RegistrationState registrationState = RegistrationState.NONE;
+
+  /** Class which provides secret keys which are shared by server and client 
on a per-app basis. */
+  private final SecretRegistry secretRegistry;
+
+  private SaslRpcHandler saslHandler;
+
+  /** Used for client authentication. */
+  private CelebornSaslServer saslServer = null;
+
+  public RegistrationRpcHandler(
+      TransportConf conf,
+      Channel channel,
+      BaseMessageHandler delegate,
+      SecretRegistry secretRegistry) {
+    this.conf = conf;
+    this.channel = channel;
+    this.secretRegistry = secretRegistry;
+    this.delegate = delegate;
+    this.saslHandler = new SaslRpcHandler(conf, channel, delegate, 
secretRegistry);
+  }
+
+  @Override
+  public boolean checkRegistered() {
+    return delegate.checkRegistered();
+  }
+
+  @Override
+  public final void receive(
+      TransportClient client, RequestMessage message, RpcResponseCallback 
callback) {
+    // The message is delegated either if the client is already authenticated 
or if the connection
+    // is authenticated.
+    if (registrationState == RegistrationState.REGISTERED || 
saslHandler.isAuthenticated()) {
+      LOG.trace("Already authenticated. Delegating {}", client.getClientId());
+      delegate.receive(client, message, callback);
+    } else {
+      RpcRequest rpcRequest = (RpcRequest) message;
+      try {
+        processRpcMessage(client, rpcRequest, callback);
+      } catch (Exception e) {
+        LOG.error("Error while invoking RpcHandler#receive() on RPC id " + 
rpcRequest.requestId, e);
+        registrationState = RegistrationState.FAILED;
+        client
+            .getChannel()
+            .writeAndFlush(
+                new RpcFailure(rpcRequest.requestId, 
Throwables.getStackTraceAsString(e)));
+      }
+    }
+  }
+
+  @Override
+  public final void receive(TransportClient client, RequestMessage message) {
+    if (registrationState == RegistrationState.REGISTERED || 
saslHandler.isAuthenticated()) {
+      if (LOG.isTraceEnabled()) {
+        LOG.trace("Already authenticated. Delegating {}", 
client.getClientId());
+      }
+      delegate.receive(client, message);
+    } else {
+      throw new SecurityException("Unauthenticated call to receive().");
+    }
+  }
+
+  private void processRpcMessage(
+      TransportClient client, RpcRequest message, RpcResponseCallback 
callback) throws IOException {
+    TransportMessage pbMsg = 
TransportMessage.fromByteBuffer(message.body().nioByteBuffer());
+    switch (pbMsg.getMessageTypeValue()) {
+      case AUTHENTICATION_INITIATION_REQUEST_VALUE:
+        // TODO: not validating the auth init request. Should we?
+        PbAuthenticationInitiationRequest authInitRequest = 
pbMsg.getParsedPayload();
+        checkRequestAllowed(RegistrationState.NONE);
+        respondToAuthInitialization(callback);
+        registrationState = RegistrationState.INIT;
+        LOG.trace("Authentication initialization completed: rpcId {}", 
message.requestId);
+        break;
+      case SASL_REQUEST_VALUE:
+        PbSaslRequest saslRequest = pbMsg.getParsedPayload();
+        if (saslRequest.getAuthType().equals(PbAuthType.CLIENT_AUTH)) {
+          LOG.trace("Received Sasl Message for client authentication");
+          checkRequestAllowed(RegistrationState.INIT);
+          authenticateClient(saslRequest, callback);
+          if (saslServer.isComplete()) {
+            LOG.debug("SASL authentication successful for channel {}", client);
+            complete();
+            registrationState = RegistrationState.AUTHENTICATED;
+            LOG.trace("Client authenticated: rpcId {}", message.requestId);
+          }
+        } else {
+          // It is a SASL message to authenticate the connection. If the 
application hasn't
+          // registered, then
+          // saslHandler will throw an exception that the app hasn't 
registered.
+          LOG.trace("Delegating to sasl handler: rpcId {}", message.requestId);
+          saslHandler.receive(client, message, callback);
+        }
+        break;
+      case REGISTER_APPLICATION_REQUEST_VALUE:
+        PbRegisterApplicationRequest registerApplicationRequest = 
pbMsg.getParsedPayload();
+        checkRequestAllowed(RegistrationState.AUTHENTICATED);
+        LOG.trace("Application registration started {}", 
registerApplicationRequest.getId());
+        processRegisterApplicationRequest(registerApplicationRequest, 
callback);
+        registrationState = RegistrationState.REGISTERED;
+        LOG.info(
+            "Application registered: appId {} rpcId {}",
+            registerApplicationRequest.getId(),
+            message.requestId);
+        break;
+      default:
+        throw new SecurityException(
+            "The app is not registered and the connection is not authenticated 
"
+                + message.requestId);
+    }
+  }
+
+  private void checkRequestAllowed(RegistrationState expectedState) {
+    if (registrationState != expectedState) {
+      throw new IllegalStateException(
+          "Invalid registration state. Expected: "
+              + expectedState
+              + ", Actual: "
+              + registrationState);
+    }
+  }
+
+  private void respondToAuthInitialization(RpcResponseCallback callback) {
+    PbAuthenticationInitiationResponse response =
+        PbAuthenticationInitiationResponse.newBuilder()
+            .setAuthEnabled(conf.authEnabled())
+            .setVersion(VERSION)
+            .addAllSaslMechanisms(SASL_MECHANISMS)
+            .build();
+    TransportMessage message =
+        new TransportMessage(AUTHENTICATION_INITIATION_RESPONSE, 
response.toByteArray());
+    callback.onSuccess(message.toByteBuffer());
+  }
+
+  private void authenticateClient(PbSaslRequest saslMessage, 
RpcResponseCallback callback) {
+    if (saslServer == null || !saslServer.isComplete()) {
+      if (saslServer == null) {
+        saslServer = new CelebornSaslServer(ANONYMOUS, null, null);
+      }
+      byte[] response = 
saslServer.response(saslMessage.getPayload().toByteArray());
+      callback.onSuccess(ByteBuffer.wrap(response));
+    } else {
+      throw new IllegalArgumentException("Unexpected message type " + 
saslMessage.toString());
+    }
+  }
+
+  private void processRegisterApplicationRequest(
+      PbRegisterApplicationRequest registerApplicationRequest, 
RpcResponseCallback callback) {
+    if (secretRegistry.isRegistered(registerApplicationRequest.getId())) {
+      // Re-registration is not allowed.
+      throw new IllegalStateException(
+          "Application is already registered " + 
registerApplicationRequest.getId());
+    }
+    secretRegistry.register(
+        registerApplicationRequest.getId(), 
registerApplicationRequest.getSecret());
+    PbRegisterApplicationResponse response =
+        PbRegisterApplicationResponse.newBuilder().setStatus(true).build();
+    TransportMessage message =
+        new TransportMessage(REGISTER_APPLICATION_RESPONSE, 
response.toByteArray());
+    callback.onSuccess(message.toByteBuffer());
+  }
+
+  @Override
+  public void channelInactive(TransportClient client) {
+    delegate.channelInactive(client);
+    cleanup();
+  }
+
+  @Override
+  public void exceptionCaught(Throwable cause, TransportClient client) {
+    delegate.exceptionCaught(cause, client);
+  }
+
+  private void complete() {
+    cleanup();
+  }
+
+  private void cleanup() {
+    if (null != saslServer) {
+      try {
+        saslServer.dispose();
+      } catch (RuntimeException e) {
+        LOG.error("Error while disposing SASL server", e);
+      } finally {
+        saslServer = null;
+      }
+    }
+    saslHandler.cleanup();
+  }
+
+  private enum RegistrationState {
+    NONE,
+    INIT,
+    AUTHENTICATED,
+    REGISTERED,
+    FAILED
+  }
+}
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/sasl/registration/RegistrationServerBootstrap.java
 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/registration/RegistrationServerBootstrap.java
new file mode 100644
index 000000000..b4378619c
--- /dev/null
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/sasl/registration/RegistrationServerBootstrap.java
@@ -0,0 +1,45 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.celeborn.common.network.sasl.registration;
+
+import io.netty.channel.Channel;
+
+import org.apache.celeborn.common.network.sasl.SecretRegistry;
+import org.apache.celeborn.common.network.server.BaseMessageHandler;
+import org.apache.celeborn.common.network.server.TransportServerBootstrap;
+import org.apache.celeborn.common.network.util.TransportConf;
+
+/**
+ * A bootstrap which is executed on a TransportServer's (in the Master) client 
channel once a client
+ * connects to the server.
+ */
+public class RegistrationServerBootstrap implements TransportServerBootstrap {
+
+  private final TransportConf conf;
+  private final SecretRegistry secretRegistry;
+
+  public RegistrationServerBootstrap(TransportConf conf, SecretRegistry 
secretRegistry) {
+    this.conf = conf;
+    this.secretRegistry = secretRegistry;
+  }
+
+  @Override
+  public BaseMessageHandler doBootstrap(Channel channel, BaseMessageHandler 
rpcHandler) {
+    return new RegistrationRpcHandler(conf, channel, rpcHandler, 
secretRegistry);
+  }
+}
diff --git 
a/common/src/main/java/org/apache/celeborn/common/network/util/TransportConf.java
 
b/common/src/main/java/org/apache/celeborn/common/network/util/TransportConf.java
index a93b545dc..417e93319 100644
--- 
a/common/src/main/java/org/apache/celeborn/common/network/util/TransportConf.java
+++ 
b/common/src/main/java/org/apache/celeborn/common/network/util/TransportConf.java
@@ -158,4 +158,9 @@ public class TransportConf {
   public int saslTimeoutMs() {
     return celebornConf.networkIoSaslTimoutMs(module);
   }
+
+  /** Whether authentication is enabled or not. */
+  public boolean authEnabled() {
+    return celebornConf.authEnabled();
+  }
 }
diff --git a/common/src/main/proto/TransportMessages.proto 
b/common/src/main/proto/TransportMessages.proto
index b27c06140..963bc57c9 100644
--- a/common/src/main/proto/TransportMessages.proto
+++ b/common/src/main/proto/TransportMessages.proto
@@ -92,6 +92,10 @@ enum MessageType {
   GET_SHUFFLE_ID = 69;
   GET_SHUFFLE_ID_RESPONSE = 70;
   SASL_REQUEST = 71;
+  AUTHENTICATION_INITIATION_REQUEST = 72;
+  AUTHENTICATION_INITIATION_RESPONSE = 73;
+  REGISTER_APPLICATION_REQUEST = 74;
+  REGISTER_APPLICATION_RESPONSE = 75;
 }
 
 enum StreamType {
@@ -636,8 +640,9 @@ message PbTransportableError {
 }
 
 enum PbAuthType {
-  CLIENT_AUTH = 0;
-  CONNECTION_AUTH = 1;
+  UNDEFINED_AUTH = 0;
+  CLIENT_AUTH = 1;
+  CONNECTION_AUTH = 2;
 }
 
 message PbSaslMechanism {
@@ -650,3 +655,24 @@ message PbSaslRequest {
   PbAuthType authType = 2;
   bytes payload = 3;
 }
+
+message PbAuthenticationInitiationRequest {
+  string version = 1;
+  bool authEnabled = 2;
+  repeated PbSaslMechanism saslMechanisms = 3;
+}
+
+message PbAuthenticationInitiationResponse {
+  string version = 1;
+  bool authEnabled = 2;
+  repeated PbSaslMechanism saslMechanisms = 3;
+}
+
+message PbRegisterApplicationRequest {
+  string id = 1;
+  string secret = 2;
+}
+
+message PbRegisterApplicationResponse {
+  bool status = 1;
+}
diff --git 
a/common/src/main/scala/org/apache/celeborn/common/CelebornConf.scala 
b/common/src/main/scala/org/apache/celeborn/common/CelebornConf.scala
index 9c94fa2a0..40e3c743e 100644
--- a/common/src/main/scala/org/apache/celeborn/common/CelebornConf.scala
+++ b/common/src/main/scala/org/apache/celeborn/common/CelebornConf.scala
@@ -1110,6 +1110,11 @@ class CelebornConf(loadDefaults: Boolean) extends 
Cloneable with Logging with Se
   // //////////////////////////////////////////////////////
   def hdfsStorageKerberosPrincipal = get(HDFS_STORAGE_KERBEROS_PRINCIPAL)
   def hdfsStorageKerberosKeytab = get(HDFS_STORAGE_KERBEROS_KEYTAB)
+
+  // //////////////////////////////////////////////////////
+  //               Authentication                        //
+  // //////////////////////////////////////////////////////
+  def authEnabled: Boolean = get(AUTH_ENABLED)
 }
 
 object CelebornConf extends Logging {
@@ -4359,4 +4364,12 @@ object CelebornConf extends Logging {
       .version("0.5.0")
       .timeConf(TimeUnit.MILLISECONDS)
       .createWithDefaultString("30s")
+
+  val AUTH_ENABLED: ConfigEntry[Boolean] =
+    buildConf("celeborn.auth.enabled")
+      .categories("auth")
+      .version("0.5.0")
+      .doc("Whether to enable authentication.")
+      .booleanConf
+      .createWithDefault(false)
 }
diff --git 
a/common/src/test/java/org/apache/celeborn/common/network/sasl/CelebornSaslSuiteJ.java
 
b/common/src/test/java/org/apache/celeborn/common/network/sasl/CelebornSaslSuiteJ.java
index 60ff875ed..b7d328541 100644
--- 
a/common/src/test/java/org/apache/celeborn/common/network/sasl/CelebornSaslSuiteJ.java
+++ 
b/common/src/test/java/org/apache/celeborn/common/network/sasl/CelebornSaslSuiteJ.java
@@ -21,37 +21,16 @@ import static 
org.apache.celeborn.common.network.sasl.SaslUtils.*;
 import static org.junit.Assert.*;
 import static org.mockito.Mockito.*;
 
-import java.nio.ByteBuffer;
-import java.util.ArrayList;
-import java.util.Collections;
-import java.util.List;
-import java.util.concurrent.TimeUnit;
-
-import org.junit.BeforeClass;
 import org.junit.Test;
 
 import org.apache.celeborn.common.CelebornConf;
-import org.apache.celeborn.common.network.TransportContext;
-import org.apache.celeborn.common.network.client.RpcResponseCallback;
-import org.apache.celeborn.common.network.client.TransportClient;
-import org.apache.celeborn.common.network.client.TransportClientBootstrap;
-import org.apache.celeborn.common.network.protocol.RequestMessage;
 import org.apache.celeborn.common.network.server.BaseMessageHandler;
-import org.apache.celeborn.common.network.server.TransportServer;
 import org.apache.celeborn.common.network.util.TransportConf;
-import org.apache.celeborn.common.util.JavaUtils;
 
 /**
  * Jointly tests {@link CelebornSaslClient} and {@link CelebornSaslServer}, as 
both are black boxes.
  */
-public class CelebornSaslSuiteJ {
-  private static final String TEST_USER = "appId";
-  private static final String TEST_SECRET = "secret";
-
-  @BeforeClass
-  public static void setup() {
-    SecretRegistryImpl.getInstance().register(TEST_USER, TEST_SECRET);
-  }
+public class CelebornSaslSuiteJ extends SaslTestBase {
 
   @Test
   public void testDigestMatching() {
@@ -115,43 +94,12 @@ public class CelebornSaslSuiteJ {
 
   @Test
   public void testSaslAuth() throws Throwable {
-    BaseMessageHandler rpcHandler = mock(BaseMessageHandler.class);
-    doAnswer(
-            invocation -> {
-              RequestMessage message = (RequestMessage) 
invocation.getArguments()[1];
-              RpcResponseCallback cb = (RpcResponseCallback) 
invocation.getArguments()[2];
-              assertEquals("Ping", 
JavaUtils.bytesToString(message.body().nioByteBuffer()));
-              cb.onSuccess(JavaUtils.stringToBytes("Pong"));
-              return null;
-            })
-        .when(rpcHandler)
-        .receive(
-            any(TransportClient.class), any(RequestMessage.class), 
any(RpcResponseCallback.class));
-
-    doReturn(true).when(rpcHandler).checkRegistered();
-
-    try (SaslTestCtx ctx = new SaslTestCtx(rpcHandler)) {
-      ByteBuffer response =
-          ctx.client.sendRpcSync(JavaUtils.stringToBytes("Ping"), 
TimeUnit.SECONDS.toMillis(10));
-      assertEquals("Pong", JavaUtils.bytesToString(response));
-    } finally {
-      // There should be 2 terminated events; one for the client, one for the 
server.
-      Throwable error = null;
-      long deadline = System.nanoTime() + TimeUnit.NANOSECONDS.convert(10, 
TimeUnit.SECONDS);
-      while (deadline > System.nanoTime()) {
-        try {
-          verify(rpcHandler, 
times(2)).channelInactive(any(TransportClient.class));
-          error = null;
-          break;
-        } catch (Throwable t) {
-          error = t;
-          TimeUnit.MILLISECONDS.sleep(10);
-        }
-      }
-      if (error != null) {
-        throw error;
-      }
-    }
+    TransportConf conf = new TransportConf("shuffle", new CelebornConf());
+    SaslServerBootstrap serverBootstrap =
+        new SaslServerBootstrap(conf, SecretRegistryImpl.getInstance());
+    SaslClientBootstrap clientBootstrap =
+        new SaslClientBootstrap(conf, TEST_USER, new 
SaslCredentials(TEST_USER, TEST_SECRET));
+    authHelper(conf, serverBootstrap, clientBootstrap);
   }
 
   @Test
@@ -188,42 +136,4 @@ public class CelebornSaslSuiteJ {
     client.dispose();
     assertFalse(client.isComplete());
   }
-
-  private static class SaslTestCtx implements AutoCloseable {
-
-    final TransportClient client;
-    final TransportServer server;
-    final TransportContext ctx;
-
-    SaslTestCtx(BaseMessageHandler rpcHandler) throws Exception {
-      TransportConf conf = new TransportConf("shuffle", new CelebornConf());
-
-      this.ctx = new TransportContext(conf, rpcHandler);
-      this.server =
-          ctx.createServer(
-              Collections.singletonList(
-                  new SaslServerBootstrap(conf, 
SecretRegistryImpl.getInstance())));
-      List<TransportClientBootstrap> clientBootstraps = new ArrayList<>();
-      clientBootstraps.add(
-          new SaslClientBootstrap(conf, "appId", new 
SaslCredentials(TEST_USER, TEST_SECRET)));
-      try {
-        this.client =
-            ctx.createClientFactory(clientBootstraps)
-                .createClient(JavaUtils.getLocalHost(), server.getPort());
-      } catch (Exception e) {
-        close();
-        throw e;
-      }
-    }
-
-    @Override
-    public void close() {
-      if (client != null) {
-        client.close();
-      }
-      if (server != null) {
-        server.close();
-      }
-    }
-  }
 }
diff --git 
a/common/src/test/java/org/apache/celeborn/common/network/sasl/RegistrationSuiteJ.java
 
b/common/src/test/java/org/apache/celeborn/common/network/sasl/RegistrationSuiteJ.java
new file mode 100644
index 000000000..8240324bd
--- /dev/null
+++ 
b/common/src/test/java/org/apache/celeborn/common/network/sasl/RegistrationSuiteJ.java
@@ -0,0 +1,113 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.celeborn.common.network.sasl;
+
+import static org.junit.Assert.*;
+
+import java.io.IOException;
+import java.util.HashMap;
+import java.util.Map;
+
+import com.google.common.base.Throwables;
+import org.junit.Test;
+
+import org.apache.celeborn.common.CelebornConf;
+import 
org.apache.celeborn.common.network.sasl.registration.RegistrationClientBootstrap;
+import org.apache.celeborn.common.network.sasl.registration.RegistrationInfo;
+import 
org.apache.celeborn.common.network.sasl.registration.RegistrationRpcHandler;
+import 
org.apache.celeborn.common.network.sasl.registration.RegistrationServerBootstrap;
+import org.apache.celeborn.common.network.util.TransportConf;
+
+/**
+ * Jointly tests {@link RegistrationClientBootstrap} and {@link 
RegistrationRpcHandler}, as both are
+ * black boxes.
+ */
+public class RegistrationSuiteJ extends SaslTestBase {
+
+  @Test
+  public void testRegistration() throws Throwable {
+    TransportConf conf = new TransportConf("shuffle", new CelebornConf());
+    RegistrationServerBootstrap serverBootstrap =
+        new RegistrationServerBootstrap(conf, new TestSecretRegistry());
+    RegistrationClientBootstrap clientBootstrap =
+        new RegistrationClientBootstrap(
+            conf, TEST_USER, new SaslCredentials(TEST_USER, TEST_SECRET), new 
RegistrationInfo());
+    authHelper(conf, serverBootstrap, clientBootstrap);
+  }
+
+  @Test(expected = IOException.class)
+  public void testReRegisterationFails() throws Throwable {
+    TransportConf conf = new TransportConf("shuffle", new CelebornConf());
+    // The SecretRegistryImpl already has the entry for TEST_USER so 
re-registering the app should
+    // fail.
+    RegistrationServerBootstrap serverBootstrap =
+        new RegistrationServerBootstrap(conf, 
SecretRegistryImpl.getInstance());
+    RegistrationClientBootstrap clientBootstrap =
+        new RegistrationClientBootstrap(
+            conf, TEST_USER, new SaslCredentials(TEST_USER, TEST_SECRET), new 
RegistrationInfo());
+
+    try {
+      authHelper(conf, serverBootstrap, clientBootstrap);
+    } catch (Throwable t) {
+      assertTrue(Throwables.getStackTraceAsString(t).contains("Application is 
already registered"));
+      throw t.getCause();
+    }
+  }
+
+  @Test(expected = IOException.class)
+  public void testConnectionAuthWithoutRegistrationShouldFail() throws 
Throwable {
+    TransportConf conf = new TransportConf("shuffle", new CelebornConf());
+    RegistrationServerBootstrap serverBootstrap =
+        new RegistrationServerBootstrap(conf, new TestSecretRegistry());
+    SaslClientBootstrap clientBootstrap =
+        new SaslClientBootstrap(conf, TEST_USER, new 
SaslCredentials(TEST_USER, TEST_SECRET));
+
+    try {
+      authHelper(conf, serverBootstrap, clientBootstrap);
+    } catch (Throwable t) {
+      assertTrue(
+          Throwables.getStackTraceAsString(t).contains("Registration 
information not found"));
+      throw t.getCause();
+    }
+  }
+
+  static class TestSecretRegistry implements SecretRegistry {
+
+    private final Map<String, String> secrets = new HashMap<>();
+
+    @Override
+    public void register(String appId, String secret) {
+      secrets.put(appId, secret);
+    }
+
+    @Override
+    public void unregister(String appId) {
+      secrets.remove(appId);
+    }
+
+    @Override
+    public boolean isRegistered(String appId) {
+      return secrets.containsKey(appId);
+    }
+
+    @Override
+    public String getSecretKey(String appId) {
+      return secrets.get(appId);
+    }
+  }
+}
diff --git 
a/common/src/test/java/org/apache/celeborn/common/network/sasl/SaslTestBase.java
 
b/common/src/test/java/org/apache/celeborn/common/network/sasl/SaslTestBase.java
new file mode 100644
index 000000000..561af5cd6
--- /dev/null
+++ 
b/common/src/test/java/org/apache/celeborn/common/network/sasl/SaslTestBase.java
@@ -0,0 +1,140 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.celeborn.common.network.sasl;
+
+import static org.junit.Assert.*;
+import static org.mockito.ArgumentMatchers.any;
+import static org.mockito.Mockito.*;
+
+import java.nio.ByteBuffer;
+import java.util.ArrayList;
+import java.util.Collections;
+import java.util.List;
+import java.util.concurrent.TimeUnit;
+
+import org.junit.AfterClass;
+import org.junit.BeforeClass;
+
+import org.apache.celeborn.common.network.TransportContext;
+import org.apache.celeborn.common.network.client.RpcResponseCallback;
+import org.apache.celeborn.common.network.client.TransportClient;
+import org.apache.celeborn.common.network.client.TransportClientBootstrap;
+import org.apache.celeborn.common.network.protocol.RequestMessage;
+import org.apache.celeborn.common.network.server.BaseMessageHandler;
+import org.apache.celeborn.common.network.server.TransportServer;
+import org.apache.celeborn.common.network.server.TransportServerBootstrap;
+import org.apache.celeborn.common.network.util.TransportConf;
+import org.apache.celeborn.common.util.JavaUtils;
+
+public class SaslTestBase {
+
+  @BeforeClass
+  public static void setup() {
+    SecretRegistryImpl.getInstance().register(TEST_USER, TEST_SECRET);
+  }
+
+  @AfterClass
+  public static void teardown() {
+    SecretRegistryImpl.getInstance().unregister(TEST_USER);
+  }
+
+  static final String TEST_USER = "appId";
+  static final String TEST_SECRET = "secret";
+
+  void authHelper(
+      TransportConf conf,
+      TransportServerBootstrap serverBootstrap,
+      TransportClientBootstrap clientBootstrap)
+      throws Throwable {
+    BaseMessageHandler rpcHandler = mock(BaseMessageHandler.class);
+    doAnswer(
+            invocation -> {
+              RequestMessage message = (RequestMessage) 
invocation.getArguments()[1];
+              RpcResponseCallback cb = (RpcResponseCallback) 
invocation.getArguments()[2];
+              assertEquals("Ping", 
JavaUtils.bytesToString(message.body().nioByteBuffer()));
+              cb.onSuccess(JavaUtils.stringToBytes("Pong"));
+              return null;
+            })
+        .when(rpcHandler)
+        .receive(
+            any(TransportClient.class), any(RequestMessage.class), 
any(RpcResponseCallback.class));
+
+    doReturn(true).when(rpcHandler).checkRegistered();
+
+    try (SaslTestCtx ctx = new SaslTestCtx(conf, rpcHandler, serverBootstrap, 
clientBootstrap)) {
+      ByteBuffer response =
+          ctx.client.sendRpcSync(JavaUtils.stringToBytes("Ping"), 
TimeUnit.SECONDS.toMillis(10));
+      assertEquals("Pong", JavaUtils.bytesToString(response));
+    } finally {
+      // There should be 2 terminated events; one for the client, one for the 
server.
+      Throwable error = null;
+      long deadline = System.nanoTime() + TimeUnit.NANOSECONDS.convert(10, 
TimeUnit.SECONDS);
+      while (deadline > System.nanoTime()) {
+        try {
+          verify(rpcHandler, 
times(2)).channelInactive(any(TransportClient.class));
+          error = null;
+          break;
+        } catch (Throwable t) {
+          error = t;
+          TimeUnit.MILLISECONDS.sleep(10);
+        }
+      }
+      if (error != null) {
+        throw error;
+      }
+    }
+  }
+
+  static class SaslTestCtx implements AutoCloseable {
+
+    final TransportClient client;
+    final TransportServer server;
+    final TransportContext ctx;
+
+    SaslTestCtx(
+        TransportConf conf,
+        BaseMessageHandler rpcHandler,
+        TransportServerBootstrap serverBootstrap,
+        TransportClientBootstrap clientBootstrap)
+        throws Exception {
+
+      this.ctx = new TransportContext(conf, rpcHandler);
+      this.server = 
ctx.createServer(Collections.singletonList(serverBootstrap));
+      List<TransportClientBootstrap> clientBootstraps = new ArrayList<>();
+      clientBootstraps.add(clientBootstrap);
+      try {
+        this.client =
+            ctx.createClientFactory(clientBootstraps)
+                .createClient(JavaUtils.getLocalHost(), server.getPort());
+      } catch (Exception e) {
+        close();
+        throw e;
+      }
+    }
+
+    @Override
+    public void close() {
+      if (client != null) {
+        client.close();
+      }
+      if (server != null) {
+        server.close();
+      }
+    }
+  }
+}

Reply via email to