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();
+ }
+ }
+ }
+}